Solomon / src /solomon /engine_reference.py
Archer Hume
Solomon v1.1
105f9ef
Raw
History Blame Contribute Delete
13 kB
"""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