Code / experiments/candidate_01/implementation/runtime.py
experiments/candidate_01/implementation/runtime.py
167 lines
# =============================================================================
# Project : localvm-research
# File : experiments/candidate_01/implementation/runtime.py
# Purpose : Margin-gated deferred-refinement generation loop (q4 resident
# base + windowed q8 verification sweeps with rollback)
# Author : Simon-Pierre Boucher
# Contact : contact@spboucher.ai
# Created : 2026-08-12
# Modified : 2026-08-12
# Platform : macOS / Apple Silicon (arm64) — MLX / Metal
# License : All rights reserved (research code)
# =============================================================================
"""Candidate-01 runtime.
Two modes:
* ``verify-all`` — every window position is checked against the q8 sweep
(exact q8-greedy output; precision-tiered speculative decoding).
* ``margin`` — only positions whose base margin < tau are checked
(approximate; fewer rollbacks — the novel relaxation under test).
"""
from __future__ import annotations
import time
from dataclasses import dataclass, field
import mlx.core as mx
from mlx_lm.models.cache import make_prompt_cache, trim_prompt_cache
@dataclass
class RunStats:
tokens_out: int = 0
deferred: int = 0
sweeps: int = 0
rollbacks: int = 0
sweep_time_s: float = 0.0
gen_time_s: float = 0.0
sweep_logical_bytes: int = 0
events: list = field(default_factory=list)
def _greedy_step(model, inp, cache):
logits = model(inp, cache=cache)[0, -1].astype(mx.float32)
top2 = mx.topk(logits, 2)
nxt = mx.argmax(logits)
mx.eval(top2, nxt)
v = top2.tolist()
return int(nxt.item()), abs(v[1] - v[0])
def generate_deferred(base_model, verify_model, tokenizer, prompt_ids: list[int],
max_tokens: int, tau: float, window: int,
mode: str, q8_bytes: int) -> tuple[list[int], RunStats]:
"""Generate with the resident base model, verifying windows with the q8
model. Returns (tokens, stats). `verify_model` weights are assumed
mmap-backed; each sweep accounts `q8_bytes` logical bytes."""
assert mode in ("verify-all", "margin")
stats = RunStats()
tokens: list[int] = [] # accepted generated tokens
flags: list[bool] = [] # deferred flag per generated token
verified_upto = 0 # tokens before this index are settled
cache = make_prompt_cache(base_model)
# Incremental verify cache: sweeps feed only the unseen suffix, cutting
# sweep compute from O(context) to O(window). v_pos = tokens of the
# current sequence already ingested; invariant at sweep: v_pos <= base-1.
# verify_model may be a StreamingVerifier (layer-streamed weights for
# models that do not fit beside the resident base).
is_streaming = hasattr(verify_model, "forward_chunk")
vcache = make_prompt_cache(verify_model.model if is_streaming else verify_model)
v_pos = 0
t0 = time.perf_counter()
inp = mx.array(prompt_ids)[None]
def sweep() -> bool:
"""Verify tokens[verified_upto:] with q8; rollback on first flip.
Returns True if a rollback happened."""
nonlocal verified_upto, tokens, flags, inp, cache, v_pos
pending = len(tokens) - verified_upto
if pending <= 0:
return False
s0 = time.perf_counter()
stats.sweeps += 1
stats.sweep_logical_bytes += q8_bytes
full = prompt_ids + tokens
base_of_window = len(prompt_ids) + verified_upto
# feed full[v_pos : end-1]; logits rows are absolute v_pos..end-2,
# predictions for tokens v_pos+1..end-1 ⊇ the pending window
chunk = full[v_pos : len(full) - 1]
if is_streaming:
logits = verify_model.forward_chunk(chunk, vcache)
else:
logits = verify_model(mx.array(chunk)[None], cache=vcache)[0]
sel = logits[base_of_window - 1 - v_pos :].astype(mx.float32)
v_pos = len(full) - 1
q8_argmax = mx.argmax(sel, axis=-1)
mx.eval(q8_argmax)
q8_next = q8_argmax.tolist()
flip_at = None
for j in range(pending):
i = verified_upto + j
check = mode == "verify-all" or flags[i]
if check and q8_next[j] != tokens[i]:
flip_at = (i, q8_next[j])
break
if flip_at is None:
verified_upto = len(tokens)
stats.sweep_time_s += time.perf_counter() - s0
return False
i, corrected = flip_at
stats.rollbacks += 1
# Cache invariant at sweep time: cache = prompt + tokens[:-1] (the
# newest token was appended but not yet fed). To leave the cache at
# prompt + tokens[:i], trim (len(tokens)-1 - i) entries.
trim_prompt_cache(cache, len(tokens) - 1 - i)
# verify cache must hold a prefix of the corrected sequence:
# keep exactly prompt + tokens[:i] (= base_next - 1, invariant holds)
keep = len(prompt_ids) + i
if v_pos > keep:
trim_prompt_cache(vcache, v_pos - keep)
v_pos = keep
tokens = tokens[:i] + [corrected]
flags = flags[:i] + [False]
verified_upto = len(tokens)
# do NOT feed `corrected` here — the main loop's next _greedy_step
# feeds tokens[-1], preserving the invariant
stats.sweep_time_s += time.perf_counter() - s0
stats.events.append({"rollback_at": i})
return True
# Loop until the full budget is generated AND the tail is verified —
# a rollback in the final window must resume generation, not truncate.
max_steps = 6 * max_tokens # safety bound against rollback ping-pong
steps = 0
while steps < max_steps:
steps += 1
hit_eos = False
if len(tokens) < max_tokens:
nxt, margin = _greedy_step(base_model, inp, cache)
hit_eos = nxt == tokenizer.eos_token_id
# EOS is a decision too: append it (always checked) so the sweep
# can veto a premature stop; stripped before returning.
tokens.append(nxt)
flags.append(hit_eos or margin < tau)
if not hit_eos and margin < tau:
stats.deferred += 1
inp = mx.array([[nxt]])
if not hit_eos and len(tokens) - verified_upto < window and len(tokens) < max_tokens:
continue
rolled = sweep()
if rolled:
if tokens and tokens[-1] == tokenizer.eos_token_id:
tokens.pop() # q8 corrected the decision to "stop here"
break
inp = mx.array([[tokens[-1]]]) # corrected token, not yet fed
continue
if len(tokens) and tokens[-1] == tokenizer.eos_token_id:
tokens.pop()
break
if len(tokens) >= max_tokens:
break
stats.tokens_out = len(tokens)
stats.gen_time_s = time.perf_counter() - t0 - stats.sweep_time_s
return tokens, stats