"""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("", 1)[1] if "" 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, }