lfm2-5-350m-rlcd-coreml / preprocessing.py
alexwengg's picture
Publish validated LFM2.5-350M-RLCD FP16 Core ML conversion
616b1e2 verified
Raw
History Blame Contribute Delete
5.13 kB
"""Pinned RLCD schema rendering and fixed-shape candidate preprocessing."""
from __future__ import annotations
import json
from dataclasses import dataclass
import jsonschema
import numpy as np
@dataclass(frozen=True)
class Shape:
length: int = 256
candidates: int = 8
max_value_tokens: int = 16
def validate_schema(schema: dict) -> None:
"""Enforce the pinned RLCD engine's flat, closed schema contract."""
jsonschema.Draft202012Validator.check_schema(schema)
if schema.get("type") != "object" or schema.get("additionalProperties") is not False:
raise ValueError("Only closed, flat object schemas are supported")
fields = schema.get("properties", {})
if not fields or set(schema.get("required", [])) != set(fields):
raise ValueError("All fields must be required")
if set(schema) - {"type", "properties", "required", "additionalProperties"}:
raise ValueError("Unsupported object constraints")
for spec in fields.values():
if set(spec) - {"type", "enum", "description"}:
raise ValueError("Unsupported field constraints")
if spec.get("type") == "boolean" and "enum" not in spec:
continue
values = spec.get("enum", [])
if spec.get("type") != "string" or not values or any(type(value) is not str for value in values):
raise ValueError("Fields must be booleans or nonempty string enums")
if len(set(values)) != len(values):
raise ValueError("Duplicate candidates")
def prompt(tokenizer, context: str, schema: dict) -> str:
"""Build the pinned RLCD prompt exactly."""
validate_schema(schema)
messages = [
{
"role": "system",
"content": "Extract the attributes from the text. Return only a JSON object matching this schema. "
"Use the exact allowed values. No explanation or markdown.\n" + json.dumps(schema, ensure_ascii=False),
},
{"role": "user", "content": context},
]
return tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) + "{\n"
@dataclass(frozen=True)
class Candidate:
field: str
value: str | bool
ids: list[int]
positions: list[int]
targets: list[int]
def prepare_candidates(tokenizer, context: str, schema: dict, shape: Shape) -> list[Candidate]:
def encode(text: str) -> list[int]:
return tokenizer.encode(text, add_special_tokens=False)
prefix = encode(prompt(tokenizer, context, schema))
candidates = []
for name, spec in schema["properties"].items():
values = [True, False] if spec["type"] == "boolean" else spec["enum"]
suffix = encode(" " + json.dumps(name, ensure_ascii=False) + ": ")
for value in values:
value_ids = encode(json.dumps(value, ensure_ascii=False) + "\n")
ids = prefix + suffix + value_ids
if len(ids) > shape.length:
raise ValueError(f"candidate exceeds L{shape.length}: {name}={value!r}, {len(ids)} tokens")
if len(value_ids) > shape.max_value_tokens:
raise ValueError(f"value exceeds {shape.max_value_tokens} tokens: {name}={value!r}")
first = len(prefix) + len(suffix) - 1
candidates.append(Candidate(name, value, ids, list(range(first, first + len(value_ids))), value_ids))
return candidates
def batch_arrays(tokenizer, candidates: list[Candidate], shape: Shape) -> dict[str, np.ndarray]:
if len(candidates) > shape.candidates:
raise ValueError(f"batch has {len(candidates)} candidates, max is {shape.candidates}")
if not candidates:
raise ValueError("empty candidate batch")
pad = tokenizer.pad_token_id
if pad is None:
raise ValueError("tokenizer has no pad token")
ids = np.full((shape.candidates, shape.length), pad, dtype=np.int32)
attention = np.zeros_like(ids)
positions = np.zeros((shape.candidates, shape.max_value_tokens), dtype=np.int32)
targets = np.zeros_like(positions)
value_mask = np.zeros((shape.candidates, shape.max_value_tokens), dtype=np.float32)
for row, candidate in enumerate(candidates):
length = len(candidate.ids)
count = len(candidate.targets)
ids[row, :length] = candidate.ids
attention[row, :length] = 1
positions[row, :count] = candidate.positions
targets[row, :count] = candidate.targets
value_mask[row, :count] = 1.0
return {
"input_ids": ids,
"attention_mask": attention,
"value_positions": positions,
"value_targets": targets,
"value_mask": value_mask,
}
def select_values(candidates: list[Candidate], scores: list[float]) -> dict:
if len(candidates) != len(scores):
raise ValueError("one score required per candidate")
selected = {}
for candidate, score in zip(candidates, scores):
previous = selected.get(candidate.field)
if previous is None or score > previous[1]:
selected[candidate.field] = (candidate.value, score)
return {field: value for field, (value, _) in selected.items()}