"""Evaluate by execution and produce the accuracy-vs-latency curve (EVALUATION.md). For each eval record the model predicts the hole; we reconstruct prefix + hole + suffix and verify by execution against the held-out tests (pass@1). Latency is wall-clock at batch 1. The diffusion model is swept over denoising steps N (the test-time-compute knob); the AR baseline is a single sequential decode. Results are grouped by difficulty so we can see where, at iso-latency, the diffusion substrate wins. Usage: python -m ml.evaluate --diff runs/diff --ar runs/ar --eval data/eval.jsonl \ --steps 1,2,4,8,16,32 --out runs/curve.csv """ from __future__ import annotations import argparse import csv import time import torch from . import ar, block_diffusion from .ar import build_prompt import numpy as np from .config import ModelConfig, TaskConfig from .data import (_ids_canvas, block_for_eval, encode_ar, encode_diffusion, load_records, make_block, make_block_lua) from .model import Transformer, amp_ctx from .tokenizer import Tokenizer from .verify import verify_batch def load_model(path: str, device: str): ckpt = torch.load(f"{path}/model.pt", map_location=device, weights_only=False) mcfg = ModelConfig(**ckpt["model_cfg"]) model = Transformer(mcfg, causal=(ckpt["mode"] == "ar")).to(device) model.load_state_dict(ckpt["model"]) model.eval() tok = Tokenizer.load(f"{path}/tokenizer.json") task = TaskConfig(**ckpt["task_cfg"]) return model, tok, task, ckpt["mode"] def pick_device() -> str: if torch.backends.mps.is_available(): return "mps" if torch.cuda.is_available(): return "cuda" return "cpu" def rec_seed(rec) -> int: """Stable per-record seed for the eval block (independent of list position).""" return rec.get("seed", 0) % 2147483647 def get_block(tok, source, frac, seed): """(pre, blk, suf) for the record — id lists in lua mode, strings in char.""" rng = np.random.RandomState(seed + 1) return make_block_lua(source, frac, rng, tok) if tok.mode == "lua" else make_block(source, frac, rng) def reconstruct(tok, blk, pred_ids): """Full program from prefix + predicted block + suffix.""" if tok.mode == "lua": return tok.decode(list(blk[0]) + list(pred_ids) + list(blk[2])) return blk[0] + tok.decode(pred_ids) + blk[2] def eval_diffusion(model, tok, task, records, n_inner, device, frac): """Block-diffusion eval at one n_inner setting (latency knob).""" cands, lats, kept = [], [], [] for i, rec in enumerate(records): blk = get_block(tok, rec["source"], frac, rec_seed(rec)) if blk is None: continue enc = _ids_canvas(tok, blk[0], blk[1], blk[2], task, ar=False) if tok.mode == "lua" \ else encode_diffusion(tok, blk[0], blk[1], blk[2], task) if enc is None: continue ids, region, _block_id, attn = enc ids_row = torch.from_numpy(ids).to(device) region_row = torch.from_numpy(region).to(device) attn_row = torch.from_numpy(attn).to(device) t0 = time.perf_counter() with amp_ctx(device): toks = block_diffusion.sample(model, ids_row, region_row, attn_row, tok, task, n_inner) if device == "mps": torch.mps.synchronize() lat = (time.perf_counter() - t0) * 1000.0 cands.append({"source": reconstruct(tok, blk, toks), "tests": rec["tests"]}) lats.append(lat) kept.append(i) return cands, lats, kept def eval_ar(model, tok, task, records, device, frac): cands, lats, kept = [], [], [] for i, rec in enumerate(records): blk = get_block(tok, rec["source"], frac, rec_seed(rec)) if blk is None: continue head = build_prompt(tok, blk[0], blk[2], task) if head is None: continue head_ids = torch.tensor(head, device=device) t0 = time.perf_counter() with amp_ctx(device): ids_out = model.generate(head_ids, max_new=task.max_decode, eos_id=tok.eos_id) if device == "mps": torch.mps.synchronize() lat = (time.perf_counter() - t0) * 1000.0 cands.append({"source": reconstruct(tok, blk, ids_out), "tests": rec["tests"]}) lats.append(lat) kept.append(i) return cands, lats, kept def summarize(rows, records, kept, passes, lats, model_name, steps, frac): """Group pass@1 and latency by difficulty, for one masking fraction.""" by_diff = {} for k, p, lat in zip(kept, passes, lats): d = records[k]["difficulty"] by_diff.setdefault(d, []).append((p, lat)) for d in sorted(by_diff): ps = [p for p, _ in by_diff[d]] ls = [lat for _, lat in by_diff[d]] rows.append({ "model": model_name, "frac": frac, "steps": steps, "difficulty": d, "n": len(ps), "pass@1": round(sum(ps) / len(ps), 4), "mean_latency_ms": round(sum(ls) / len(ls), 2), }) rows.append({ "model": model_name, "frac": frac, "steps": steps, "difficulty": "all", "n": len(passes), "pass@1": round(sum(passes) / max(1, len(passes)), 4), "mean_latency_ms": round(sum(lats) / max(1, len(lats)), 2), }) def main(): ap = argparse.ArgumentParser() ap.add_argument("--diff", help="diffusion run dir") ap.add_argument("--ar", help="AR run dir") ap.add_argument("--eval", default="data/eval.jsonl") ap.add_argument("--steps", default="1,2,4,8", help="block-diffusion n_inner values to sweep") ap.add_argument("--fracs", default="0.15,0.3,0.45", help="block fractions to sweep") ap.add_argument("--tile_size", type=int, default=0, help="override remask tile size (0=keep checkpoint)") ap.add_argument("--limit", type=int, default=0, help="cap eval records (0=all)") ap.add_argument("--out", default="runs/curve.csv") args = ap.parse_args() device = pick_device() all_records = load_records(args.eval) if args.limit: all_records = all_records[: args.limit] print(f"eval records={len(all_records)} device={device}") def warmup(model, task): ids = torch.zeros(1, task.seq_len, dtype=torch.long, device=device) keep = torch.ones(1, task.seq_len, dtype=torch.bool, device=device) with amp_ctx(device): model(ids, keep) if device == "mps": torch.mps.synchronize() rows = [] step_list = [int(s) for s in args.steps.split(",") if s] frac_list = [float(x) for x in args.fracs.split(",") if x] # Load both models once; reuse across fractions. def warmup_block(model, tok, task): # Warm the cached block-decode kernels so the first timed record is clean. ids = torch.zeros(task.seq_len, dtype=torch.long, device=device) region = torch.zeros(task.seq_len, dtype=torch.bool, device=device) region[4 : 4 + 2 * task.block_len] = True attn = torch.ones(task.seq_len, dtype=torch.bool, device=device) with amp_ctx(device): block_diffusion.sample(model, ids, region, attn, tok, task, 2) if device == "mps": torch.mps.synchronize() dmodel = dtok = dtask = None if args.diff: dmodel, dtok, dtask, m = load_model(args.diff, device) assert m == "diffusion" if args.tile_size: dtask.tile_size = args.tile_size warmup(dmodel, dtask) warmup_block(dmodel, dtok, dtask) amodel = atok = atask = None if args.ar: amodel, atok, atask, m = load_model(args.ar, device) assert m == "ar" warmup(amodel, atask) tok0 = dtok if dtok is not None else atok task0 = dtask if dtask is not None else atask for frac in frac_list: # Common-fit at THIS fraction: both encodings must fit the same records. def fits_both(rec): blk = get_block(tok0, rec["source"], frac, rec_seed(rec)) if blk is None: return False if tok0.mode == "lua": d = _ids_canvas(tok0, blk[0], blk[1], blk[2], task0, ar=False) a = _ids_canvas(tok0, blk[0], blk[1], blk[2], task0, ar=True) else: d = encode_diffusion(tok0, blk[0], blk[1], blk[2], task0) a = encode_ar(tok0, blk[0], blk[1], blk[2], task0) return d is not None and a is not None records = [r for r in all_records if fits_both(r)] if (args.diff and args.ar) else all_records print(f"\n--- frac={frac} common-fit={len(records)} ---") if args.diff: for N in step_list: cands, lats, kept = eval_diffusion(dmodel, dtok, dtask, records, N, device, frac) passes = verify_batch(cands) summarize(rows, records, kept, passes, lats, "diffusion", N, frac) o = rows[-1] print(f" diffusion n_inner={N:>2} pass@1={o['pass@1']:.3f} lat={o['mean_latency_ms']:.1f}ms") if args.ar: cands, lats, kept = eval_ar(amodel, atok, atask, records, device, frac) passes = verify_batch(cands) summarize(rows, records, kept, passes, lats, "ar", "NA", frac) o = rows[-1] print(f" ar pass@1={o['pass@1']:.3f} lat={o['mean_latency_ms']:.1f}ms") with open(args.out, "w", newline="") as f: w = csv.DictWriter(f, fieldnames=["model", "frac", "steps", "difficulty", "n", "pass@1", "mean_latency_ms"]) w.writeheader() w.writerows(rows) print(f"\nwrote {args.out}") if __name__ == "__main__": main()