"""The reference engine: task-agnostic document prefix, float32 arithmetic, chunked prefill. Contract v2 ----------- prefix = chat template(system=GENERIC, user="Document:\n" + document ...) <- cached once branch = "\n\n" + task block (instructions, question, options, answer cue) + generation prompt The document state therefore carries no task instructions: one state serves Boolean and choice questions. Historical sources (qwen38/, decision_service/, backbone_comparison/) are not modified. """ import copy import hashlib import json from pathlib import Path import numpy as np ROOT = Path(__file__).resolve().parents[1] MODEL = ROOT / 'runtime/qwen38-8bit' ART = ROOT / 'artifacts/reference' CONTRACT = 'solomon-document-prefix-v2' CLASSES = ['yes_only', 'no_only', 'neither', 'both'] LETTERS = 'ABCDEFGH' CAPTURE_LAYERS = (31, 47, 55) # zero-based decoder layers; final normed state is always captured CHUNK = 2048 TOKEN_CAP = 40960 SYSTEM = ('You answer questions about the supplied document. Use only the document. ' 'Task instructions follow the document; follow them exactly.') BOOLEAN_TASK = ('Task: classify the evidence for the question using only the document and its explicit rules. ' 'A = Yes only. B = No only. C = neither Yes nor No is established. D = both Yes and No are established. ' 'A missing fact is not a negative fact. Evidence about another person or subject does not contradict ' 'the queried one. Apply explicit time and replacement rules before deciding. ' 'Respond with exactly one letter: A, B, C, or D. Do not explain.') CHOICE_TASK = ('Task: choose the single option that the document best supports. ' 'Respond with exactly one letter. Do not explain.') def sha(path): return hashlib.sha256(Path(path).read_bytes()).hexdigest() def dump(path, value): path = Path(path) path.parent.mkdir(parents=True, exist_ok=True) tmp = path.with_suffix(path.suffix + '.tmp') tmp.write_text(json.dumps(value, indent=2) + '\n') tmp.replace(path) def read_rows(path): return [json.loads(s) for s in Path(path).read_text().splitlines() if s.strip()] def softmax(a): a = np.asarray(a, dtype=np.float64) p = np.exp(a - a.max()) return p / p.sum() def boolean_block(question): return BOOLEAN_TASK + '\nQuestion: ' + question + '\nAnswer (one letter):' def choice_block(question, options): if not 2 <= len(options) <= len(LETTERS): raise ValueError('Choice needs 2-8 options') lines = [f'{LETTERS[i]}. {text}' for i, text in enumerate(options)] return CHOICE_TASK + '\nQuestion: ' + question + '\nOptions:\n' + '\n'.join(lines) + '\nAnswer (one letter):' def block_for(row): """Task block and number of answer letters for a benchmark row.""" inp = row['input'] if row.get('task', 'boolean') == 'choice': return choice_block(inp['question'], inp['options']), len(inp['options']) return boolean_block(inp['question']), 4 def messages(document, block): return [{'role': 'system', 'content': SYSTEM}, {'role': 'user', 'content': 'Document:\n' + document + '\n\n' + block}] class Ledger: def __init__(self, path, caps=None): self.path = Path(path) self.caps = caps or {} self.counts = json.loads(self.path.read_text()) if self.path.exists() else { 'forwards': 0, 'input_tokens': 0, 'generated_tokens': 0} def add(self, forwards=0, input_tokens=0, generated_tokens=0): c = self.counts c['forwards'] += forwards c['input_tokens'] += input_tokens c['generated_tokens'] += generated_tokens for key, cap in self.caps.items(): if c[key] > cap: raise RuntimeError(f'Budget cap exceeded: {key} {c[key]} > {cap}') dump(self.path, c) class Engine: def __init__(self, arithmetic='float32', ledger=None, adapter=None): import mlx.core as mx from mlx_vlm import load if arithmetic not in ('float32', 'bfloat16'): raise ValueError('arithmetic must be float32 or bfloat16') self.mx = mx self.arithmetic = arithmetic self.model, self.processor = load(str(MODEL), trust_remote_code=False, local_files_only=True) self.model.eval() self.t = self.processor.tokenizer if hasattr(self.processor, 'tokenizer') else self.processor self.lm = self.model.language_model self.adapter = None if adapter is not None: from solomon.lora import apply_lora, load_adapter apply_lora(self.lm) load_adapter(self.lm, adapter) self.adapter = str(adapter) if arithmetic == 'float32': # Same packed 8-bit weights; floating parameters and activations promoted. self.model.apply(lambda x: x.astype(mx.float32) if mx.issubdtype(x.dtype, mx.floating) else x) mx.eval(self.model.parameters()) mx.clear_cache() self.ledger = ledger self._letters = {} # ---- rendering ------------------------------------------------------------------- def render(self, document, block): text = self.t.apply_chat_template(messages(document, block), tokenize=False, add_generation_prompt=True, enable_thinking=False) return text, self.t.encode(text, add_special_tokens=False) def prefix_ids(self, document): """Token ids of the task-agnostic prefix: everything up to the end of the document.""" text, _ = self.render(document, 'X') marker = '\n\nX' end = text.rfind(marker) if end < 0: raise ValueError('Document boundary missing') # Leave the boundary token uncached: its tokenisation can depend on what follows. ids = self.t.encode(text[:end], add_special_tokens=False)[:-1] if not ids: raise ValueError('Empty prefix') return ids def letter_ids(self, text, ids, n): out = [] for letter in LETTERS[:n]: key = letter if key not in self._letters: ext = self.t.encode(text + letter, add_special_tokens=False) if len(ext) != len(ids) + 1 or ext[:-1] != ids: raise ValueError('Unstable answer-letter continuation') self._letters[key] = ext[-1] out.append(self._letters[key]) return out # ---- execution ------------------------------------------------------------------- def _run(self, ids, cache, offset, capture, chunk=CHUNK): """Process ids through the language model. Returns (last logits, {layer: vector}).""" mx = self.mx if offset + len(ids) > TOKEN_CAP: raise ValueError(f'{offset + len(ids)} tokens exceeds the {TOKEN_CAP}-token scope cap') if self.ledger is not None: self.ledger.add(forwards=1, input_tokens=len(ids)) n = len(ids) step = n if chunk is None else chunk if cache is None and step < n: raise ValueError('Chunked execution needs a cache') start = 0 logits = feats = None while start < n: stop = min(start + step, n) last = stop == n piece = mx.array([ids[start:stop]]) pos = mx.broadcast_to(mx.arange(offset + start, offset + stop)[None, None, :], (3, 1, stop - start)) kwargs = dict(cache=cache, position_ids=pos, skip_logits=True, return_hidden=True) if last and capture: kwargs['capture_layer_ids'] = list(CAPTURE_LAYERS) out = self.lm(piece, **kwargs) hidden = out.hidden_states if last: final = hidden[-1][:, -1:, :] lg = self.lm.lm_head(final)[0, -1].astype(mx.float32) vecs = {} if capture: for layer, h in zip(CAPTURE_LAYERS, hidden[:-1]): vecs[layer] = h[0, -1].astype(mx.float32) vecs['final'] = final[0, -1].astype(mx.float32) mx.eval(lg, *vecs.values()) logits = np.array(lg) feats = {k: np.array(v) for k, v in vecs.items()} elif cache is not None: mx.eval([c.state for c in cache]) del out, hidden start = stop if not np.isfinite(logits).all() or any(not np.isfinite(v).all() for v in feats.values()): raise ValueError('Non-finite model output') if mx.get_peak_memory() > 100 * 2**30: raise MemoryError('100 GiB MLX allocation cap exceeded') return logits, feats def prefill(self, document, chunk=CHUNK): """Encode the task-agnostic document prefix once.""" ids = self.prefix_ids(document) cache = self.lm.make_cache() mx = self.mx if self.ledger is not None: self.ledger.add(forwards=1, input_tokens=len(ids)) if len(ids) > TOKEN_CAP: raise ValueError('Prefix exceeds scope cap') start = 0 while start < len(ids): stop = min(start + chunk, len(ids)) pos = mx.broadcast_to(mx.arange(start, stop)[None, None, :], (3, 1, stop - start)) self.lm(mx.array([ids[start:stop]]), cache=cache, position_ids=pos, skip_logits=True) mx.eval([c.state for c in cache]) start = stop mx.clear_cache() return {'prefix_ids': ids, 'cache': cache} def score(self, document, block, n_letters, state=None, capture=True, execution='cached', full_chunk=None): """Answer-letter distribution (and features) for one task block.""" text, ids = self.render(document, block) letters = self.letter_ids(text, ids, n_letters) fallback = None if execution == 'cached': if state is None: raise ValueError('cached execution needs a prefilled state') prefix = state['prefix_ids'] if ids[:len(prefix)] != prefix: execution, fallback = 'full', 'prefix token mismatch' if execution == 'cached': branch = copy.deepcopy(state['cache']) logits, feats = self._run(ids[len(prefix):], branch, len(prefix), capture, chunk=None) del branch reused = len(prefix) elif full_chunk: # Long prompts: unchunked float32 attention scores would not fit in memory. The whole # prompt is still processed from scratch, in chunks at different split points. logits, feats = self._run(ids, self.lm.make_cache(), 0, capture, chunk=full_chunk) reused = 0 else: logits, feats = self._run(ids, None, 0, capture, chunk=None) reused = 0 full = softmax(logits) self.mx.clear_cache() return {'letter_logits': logits[letters].astype(np.float64), 'probabilities': softmax(logits[letters]), 'mass': float(full[letters].sum()), 'top_token': int(logits.argmax()), 'top_is_letter': int(logits.argmax()) in letters, 'features': feats, 'execution': execution, 'fallback': fallback, 'prompt_tokens': len(ids), 'reused_prefix_tokens': reused, 'branch_tokens': len(ids) - reused} def generate(self, document, block, max_tokens=2048, thinking=True): """Reasoning-enabled reference. Uses the library generator on the same contract.""" from mlx_vlm import stream_generate text = self.t.apply_chat_template(messages(document, block), tokenize=False, add_generation_prompt=True, enable_thinking=thinking) pieces, last = [], None for x in stream_generate(self.model, self.processor, prompt=text, max_tokens=max_tokens, temperature=0., top_p=1., top_k=0, repetition_penalty=1.): pieces.append(x.text) last = x n = getattr(last, 'generation_tokens', 0) if self.ledger is not None: self.ledger.add(generated_tokens=n) return {'text': ''.join(pieces), 'generation_tokens': n, 'finish_reason': getattr(last, 'finish_reason', None)} def identity(engine): """Runtime fingerprint for state handles: no task prompt inside it any more.""" from importlib.metadata import version payload = {'contract': CONTRACT, 'system_sha256': hashlib.sha256(SYSTEM.encode()).hexdigest(), 'arithmetic': engine.arithmetic, 'adapter': engine.adapter and sha(engine.adapter), 'engine_sha256': sha(Path(__file__)), 'capture_layers': list(CAPTURE_LAYERS), 'model_config_sha256': sha(MODEL / 'config.json'), 'versions': {p: version(p) for p in ('mlx', 'mlx-vlm', 'transformers')}} payload['fingerprint'] = hashlib.sha256(json.dumps(payload, sort_keys=True).encode()).hexdigest() return payload