LFM2.5-2.6B-RLCD / pcd /engine.py
monotykamary's picture
feat: add verified inference-only parallel constrained decoding
3145545 verified
Raw History Blame Contribute Delete
19.9 kB
"""Fast finite-choice inference on unchanged LFM2 weights.
Token mode: one branch per field, atomic label codes, restricted LM-head projection.
Sequence mode: one branch per candidate, full-vocabulary normalized likelihood.
Neither mode is trained or calibrated. See README-PCD.md for limits and evidence.
"""
import json
import threading
import time
from collections import OrderedDict
from dataclasses import asdict
import jsonschema
import torch
import torch.nn.functional as F
from transformers import AutoModelForCausalLM, AutoTokenizer
from .cache import fork_cache, state_tensors
from .config import MODEL_ID, MODEL_REVISION, Limits
from .prompting import compile_schema, prompt_tokens, question_suffix
class Engine:
def __init__(
self,
model_id=MODEL_ID,
revision=MODEL_REVISION,
device="cuda",
dtype="bfloat16",
attention="sdpa",
limits=None,
local_files_only=False,
model=None,
tokenizer=None,
):
self.limits = limits or Limits()
self.device = torch.device(device)
self.model_id, self.revision = model_id, revision
self.tokenizer = tokenizer or AutoTokenizer.from_pretrained(
model_id,
revision=revision,
local_files_only=local_files_only,
trust_remote_code=False,
)
self.model = (
model
if model is not None
else AutoModelForCausalLM.from_pretrained(
model_id,
revision=revision,
dtype=getattr(torch, dtype),
attn_implementation=attention,
local_files_only=local_files_only,
trust_remote_code=False,
)
)
if self.model.config.model_type != "lfm2":
raise ValueError("this engine supports the LFM2 hybrid backbone only")
self.model.to(self.device).eval().requires_grad_(False)
if self.tokenizer.pad_token_id is None:
raise ValueError("a tokenizer with an explicit pad token is required")
self._schemas = OrderedDict()
self._lock = threading.RLock()
def synchronize(self):
if self.device.type == "cuda":
torch.cuda.synchronize(self.device)
elif self.device.type == "mps":
torch.mps.synchronize()
def tensor(self, value):
return torch.tensor(value, dtype=torch.long, device=self.device)
def compile(self, schema, mode):
key = (mode, json.dumps(schema, ensure_ascii=False, sort_keys=False))
if len(key[1]) > self.limits.max_schema_chars:
raise ValueError("schema exceeds character budget")
if key not in self._schemas:
compiled = compile_schema(self.tokenizer, schema, self.limits, mode)
self._schemas[key] = compiled
if len(self._schemas) > 32:
self._schemas.popitem(last=False)
self._schemas.move_to_end(key)
return self._schemas[key]
def _prefill(self, prefix):
return self.model.model(self.tensor([prefix]), use_cache=True).past_key_values
def _branches(self, cache, prefix_length, sequences):
width = max(map(len, sequences))
ids = self.tensor(
[list(s) + [self.tokenizer.pad_token_id] * (width - len(s)) for s in sequences]
)
mask = self.tensor(
[[1] * (prefix_length + len(s)) + [0] * (width - len(s)) for s in sequences]
)
output = self.model.model(
ids,
attention_mask=mask,
past_key_values=fork_cache(cache, len(sequences)),
use_cache=True,
)
return output.last_hidden_state
def _admit(self, prefix, batch_size, width, score_positions):
if len(prefix) + width > self.model.config.max_position_embeddings:
raise ValueError("prefix and suffix exceed model context")
if self.device.type != "cuda":
return
config = self.model.config
head_dim = config.hidden_size // config.num_attention_heads
attn_layers = config.layer_types.count("full_attention")
size = next(self.model.parameters()).element_size()
kv = 2 * attn_layers * config.num_key_value_heads * head_dim * size * (len(prefix) + width)
logits = min(score_positions, self.limits.projection_batch_size) * config.vocab_size * 6
attention_workspace = (
config.num_attention_heads
* max(len(prefix) ** 2, batch_size * width * (len(prefix) + width))
* 8
)
# Deliberately conservative, including eager-attention/FP32 temporaries.
estimate = kv * (batch_size + 3) + logits + attention_workspace + 1024**3
free, _ = torch.cuda.mem_get_info(self.device)
if estimate > free * 0.8:
raise ValueError(
"request exceeds conservative GPU memory budget; reduce prompt or branch batch size"
)
def _token_scores(self, cache, prefix, compiled):
fields = compiled.fields
unique = sorted({tid for f in fields for tid in f.token_ids})
index = {tid: i for i, tid in enumerate(unique)}
head = self.model.get_output_embeddings()
weights = head.weight.index_select(0, self.tensor(unique))
bias = head.bias.index_select(0, self.tensor(unique)) if head.bias is not None else None
scores, calls = [], 1
for start in range(0, len(fields), self.limits.branch_batch_size):
batch = fields[start : start + self.limits.branch_batch_size]
hidden = self._branches(cache, len(prefix), [f.suffix for f in batch])
decisions = hidden[
torch.arange(len(batch), device=self.device),
self.tensor([len(f.suffix) - 1 for f in batch]),
]
logits = F.linear(decisions, weights, bias).float()
# A small result transfer per microbatch, never a per-token GPU sync.
rows = logits.cpu().tolist()
scores.extend(
[[row[index[t]] for t in field.token_ids] for row, field in zip(rows, batch)]
)
calls += 1
return scores, calls
def _sequence_scores(self, cache, prefix, compiled):
branches = [
(fi, ci, field.suffix, value)
for fi, field in enumerate(compiled.fields)
for ci, value in enumerate(field.value_tokens)
]
scores = [[0.0] * len(field.values) for field in compiled.fields]
calls = 1
for start in range(0, len(branches), self.limits.branch_batch_size):
batch = branches[start : start + self.limits.branch_batch_size]
hidden = self._branches(
cache, len(prefix), [suffix + value for _, _, suffix, value in batch]
)
rows, positions, targets, owners = [], [], [], []
for row, (_, _, suffix, value) in enumerate(batch):
rows.extend([row] * len(value))
positions.extend(range(len(suffix) - 1, len(suffix) + len(value) - 1))
targets.extend(value)
owners.extend([row] * len(value))
selected = hidden[self.tensor(rows), self.tensor(positions)]
target_ids, owner_ids = self.tensor(targets), self.tensor(owners)
sums = torch.zeros(len(batch), device=self.device, dtype=torch.float32)
for pos in range(0, len(targets), self.limits.projection_batch_size):
end = pos + self.limits.projection_batch_size
logits = self.model.get_output_embeddings()(selected[pos:end]).float()
logp = -F.cross_entropy(logits, target_ids[pos:end], reduction="none")
sums.scatter_add_(0, owner_ids[pos:end], logp)
for (fi, ci, _, _), score in zip(batch, sums.cpu().tolist()):
scores[fi][ci] = score
calls += 1
return scores, calls
def _result(self, compiled, scores, prefix, calls, elapsed_ms, timings):
selected, telemetry = {}, {}
for field, row in zip(compiled.fields, scores):
if not all(torch.isfinite(torch.tensor(row))):
raise RuntimeError("non-finite candidate scores")
winner = max(range(len(row)), key=row.__getitem__)
probs = torch.tensor(row, dtype=torch.float64).softmax(-1).tolist()
selected[field.name] = field.values[winner]
ordered = sorted(row, reverse=True)
telemetry[field.name] = {
"value": selected[field.name],
"selected_probability": probs[winner],
"score_margin": ordered[0] - ordered[1] if len(row) > 1 else None,
"candidates": [
{"value": v, "score": s, "probability": p}
for v, s, p in zip(field.values, row, probs)
],
}
jsonschema.Draft202012Validator(compiled.schema).validate(selected)
return {
"object": selected,
"text": json.dumps(selected, ensure_ascii=False, allow_nan=False),
"fields": telemetry,
"mode": compiled.mode,
"calibrated": False,
"probability_kind": "restricted_label_softmax"
if compiled.mode == "token"
else "normalized_sequence_likelihood",
"prompt_mode": "explicit_empty_thinking",
"prompt_tokens": len(prefix),
"branches": len(compiled.fields)
if compiled.mode == "token"
else sum(len(f.values) for f in compiled.fields),
"forward_calls": calls,
"elapsed_ms": elapsed_ms,
"timings_ms": timings,
"model_id": self.model_id,
"model_revision": self.revision,
}
@torch.inference_mode()
def constrained(self, context, schema, mode="token"):
with self._lock:
self.synchronize()
begin = time.perf_counter()
compiled = self.compile(schema, mode)
prefix = prompt_tokens(self.tokenizer, compiled, context, self.limits)
branches = (
len(compiled.fields)
if mode == "token"
else sum(len(f.values) for f in compiled.fields)
)
width = max(
len(f.suffix) + (max(map(len, f.value_tokens)) if mode == "sequence" else 0)
for f in compiled.fields
)
self._admit(
prefix,
min(branches, self.limits.branch_batch_size),
width,
width * branches,
)
compiled_at = time.perf_counter()
cache = self._prefill(prefix)
self.synchronize()
prefilled_at = time.perf_counter()
scores, calls = (
self._token_scores(cache, prefix, compiled)
if mode == "token"
else self._sequence_scores(cache, prefix, compiled)
)
self.synchronize()
scored_at = time.perf_counter()
result = self._result(
compiled,
scores,
prefix,
calls,
0,
{
"prepare": (compiled_at - begin) * 1000,
"prefill": (prefilled_at - compiled_at) * 1000,
"branch_and_score": (scored_at - prefilled_at) * 1000,
},
)
result["timings_ms"]["assembly"] = (time.perf_counter() - scored_at) * 1000
result["elapsed_ms"] = (time.perf_counter() - begin) * 1000
return result
@torch.inference_mode()
def reference(self, context, schema, mode="sequence"):
"""Uncached full-LM-head oracle; deliberately slow, never the serving path."""
with self._lock:
compiled = self.compile(schema, mode)
prefix = prompt_tokens(self.tokenizer, compiled, context, self.limits)
scores = []
for field in compiled.fields:
if mode == "token":
tokens = prefix + list(field.suffix)
logits = (
self.model(self.tensor([tokens]), use_cache=False, logits_to_keep=1)
.logits[0, -1]
.float()
)
scores.append(logits[self.tensor(field.token_ids)].cpu().tolist())
else:
row = []
for value in field.value_tokens:
tokens = prefix + list(field.suffix + value)
start = len(prefix) + len(field.suffix) - 1
positions = self.tensor(list(range(start, start + len(value))))
logits = (
self.model(
self.tensor([tokens]),
use_cache=False,
logits_to_keep=positions,
)
.logits[0]
.float()
)
row.append(
float(-F.cross_entropy(logits, self.tensor(value), reduction="sum"))
)
scores.append(row)
return {field.name: row for field, row in zip(compiled.fields, scores)}
@torch.inference_mode()
def autoregressive(self, context, schema, mode="sequence", native=False, max_new_tokens=384):
with self._lock:
if type(max_new_tokens) is not int or not 1 <= max_new_tokens <= 1024:
raise ValueError("max_new_tokens must be in [1, 1024]")
if native and mode != "sequence":
raise ValueError("native baseline uses original JSON values")
self.synchronize()
begin = time.perf_counter()
compiled = self.compile(schema, mode)
prefix = prompt_tokens(self.tokenizer, compiled, context, self.limits, native=native)
if mode == "token":
suffix = (
question_suffix(
"Return a JSON object mapping every field name to its chosen option code STRING. No markdown."
)
+ "{\n"
)
prefix += self.tokenizer.encode(suffix, add_special_tokens=False)
if len(prefix) > self.limits.max_prompt_tokens:
raise ValueError("AR prompt exceeds token budget")
self._admit(prefix, 1, max_new_tokens, 1)
ids = self.tensor([prefix])
output = self.model.generate(
ids,
attention_mask=torch.ones_like(ids),
do_sample=False,
max_new_tokens=max_new_tokens,
pad_token_id=self.tokenizer.pad_token_id,
eos_token_id=self.tokenizer.eos_token_id,
)
continuation = output[0, len(prefix) :].cpu().tolist()
raw = self.tokenizer.decode(continuation, skip_special_tokens=False)
text = self.tokenizer.decode(continuation, skip_special_tokens=True)
if native:
# Explicit reasoning/final protocol separation, never brace extraction or JSON repair.
text = text.split("</think>", 1)[1] if "</think>" in text else text
else:
text = "{\n" + text
if mode == "token":
try:
coded = json.loads(text)
expected_keys = {f.name for f in compiled.fields}
if not isinstance(coded, dict) or set(coded) != expected_keys:
raise ValueError("invalid coded JSON fields")
decoded = {
f.name: f.values[f.labels.index(coded[f.name])] for f in compiled.fields
}
text = json.dumps(decoded, ensure_ascii=False)
except (ValueError, KeyError, TypeError):
# Keep invalid raw output for strict benchmark evaluation.
pass
self.synchronize()
return {
"text": text,
"raw_generation": raw,
"mode": "native_ar" if native else "ar_" + mode,
"elapsed_ms": (time.perf_counter() - begin) * 1000,
"generated_tokens": len(continuation),
"prompt_tokens": len(prefix),
"hit_token_limit": len(continuation) == max_new_tokens
and (not continuation or continuation[-1] != self.tokenizer.eos_token_id),
}
def metadata(self):
import importlib.metadata
return {
"model_id": self.model_id,
"revision": self.revision,
"dtype": str(next(self.model.parameters()).dtype),
"device": str(self.device),
"gpu": torch.cuda.get_device_name(self.device) if self.device.type == "cuda" else None,
"cuda": torch.version.cuda,
"attention": self.model.config._attn_implementation,
"convolution": "transformers reference (no optional causal-conv1d installed)",
"limits": asdict(self.limits),
"parameter_count": sum(p.numel() for p in self.model.parameters()),
"versions": {
name: importlib.metadata.version(name)
for name in [
"torch",
"transformers",
"huggingface-hub",
"jsonschema",
"jinja2",
]
},
}
@torch.inference_mode()
def validate_cache(self):
"""Cheap runtime guard: both cached convolution paths and state immutability."""
compiled = self.compile(
{
"type": "object",
"properties": {"yes": {"type": "boolean"}},
"required": ["yes"],
"additionalProperties": False,
},
"sequence",
)
prefix = prompt_tokens(self.tokenizer, compiled, "The answer is yes.", self.limits)
cache = self._prefill(prefix)
original = [t.clone() for t in state_tensors(cache)]
suffix = list(compiled.fields[0].suffix + compiled.fields[0].value_tokens[0])
errors = []
for parts in [[suffix], [[suffix[0]], suffix[1:]]]:
fork = fork_cache(cache, 2)
pieces = []
for part in parts:
pieces.append(
self.model.model(
self.tensor([part, part]), past_key_values=fork, use_cache=True
).last_hidden_state[0]
)
cached = torch.cat(pieces)
full = self.model.model(
self.tensor([prefix + suffix]), use_cache=False
).last_hidden_state[0, len(prefix) :]
errors.append(float((cached - full).abs().max()))
torch.testing.assert_close(
cached.float(),
full.float(),
atol={torch.bfloat16: 0.15, torch.float16: 0.05}.get(cached.dtype, 1e-4),
rtol={torch.bfloat16: 0.08, torch.float16: 0.02}.get(cached.dtype, 1e-4),
)
for before, after in zip(original, state_tensors(cache)):
assert torch.equal(before, after), "prefix state mutated"
layers = self.model.config.layer_types
return {
"max_hidden_errors": errors,
"attention_layers": layers.count("full_attention"),
"convolution_layers": layers.count("conv"),
"prefix_unchanged": True,
}