reyden009/speculative-decoding-lab
8
1#!/usr/bin/env python32"""analyze_final.py -- Final analysis for the speculative-decoding empirical study.3 4Consolidates the 67 experimental runs (26 ``final-*``, 20 ``curves-*``, 215``ksweep-*``) into the numbers, tables and figures that the paper cites:6 7 * ``experiments/analysis/summary.json`` consolidated statistics8 * ``experiments/analysis/tables/t2..t6_*.md`` paper-ready markdown tables9 * ``experiments/analysis/curves/acc_by_pos_*.csv`` per-position acceptance10 * ``experiments/analysis/README.md`` methodology notes11 * ``manuscript/figures/F1..F4_*.png`` paper figures (300 dpi)12 13Methodology14-----------15* **Records**: one JSON line in ``results.jsonl`` == one OK completion. Config16 comes from ``config.json`` (model path -> family+quant, ``spec_type`` +17 draft path + ``p_min`` -> drafter id).18* **Sentinel exclusion**: llama-server reports ``predicted_per_second =19 1,000,000`` and ``predicted_ms = 0`` for a handful of Gemma completions20 (timing quirk, not a real speedup). Every record with ``tok_per_s >= 1e5``21 or ``predicted_ms <= 0`` is excluded from ALL statistics (they also carry22 ``alpha/tau/draft_n = None``). ``solo`` configs have ``alpha/tau/draft_n =23 None`` by design (no drafter) and are never treated as an anomaly.24* **Log mapping (curves/ksweep)**: the lines ``draft acceptance = ...`` and25 ``acc per pos = (...)`` in ``server.log`` appear in the SAME order as the26 records with non-None ``alpha`` in ``results.jsonl`` (verified: max27 |log - record| = 0.00005, rounding only). Records without ``alpha``28 (Gemma sentinels) have no log line and are skipped.29* **Speedup vs solo**: per-prompt ratio ``tps_draft / tps_solo`` matched by30 prompt ``id`` against the ``solo`` run of the SAME family+quant; we report31 the mean and median of the per-prompt ratios and the ratio of means.32* **Break-even ``alpha_be``** (paper #32, Bielik et al.): OLS fit33 ``TPS = a + beta*alpha`` over per-prompt observations of a ksweep run34 (per domain and pooled); ``alpha_be = (TPS_base - a) / beta`` is the35 acceptance rate where the regression line crosses a context-compatible36 autoregressive baseline. The baseline must match target, context, prompt37 set, and sampling protocol; its tok/s is averaged over prompt IDs shared38 with the ksweep observations. CI95 via the delta method on the OLS covariance39 of ``a`` and ``beta``. ``beta`` is the paper's ``b`` ("recovery rate"):40 tok/s gained per unit acceptance.41 42Only Python 3.12 stdlib + numpy + matplotlib (repo venv). Idempotent.43"""44 45from __future__ import annotations46 47import argparse48import hashlib49import json50import logging51import os52import re53import sys54from dataclasses import dataclass55from pathlib import Path56from typing import Any, cast57 58import matplotlib59import numpy as np60 61matplotlib.use("Agg") # noqa: E402 (must run before pyplot import)62 63import matplotlib.pyplot as plt # noqa: E40264from matplotlib.lines import Line2D # noqa: E40265 66logger = logging.getLogger("analyze_final")67 68# --------------------------------------------------------------------------- #69# Constants70# --------------------------------------------------------------------------- #71 72SENTINEL_TPS = 1e5 # tok_per_s >= this value => spurious timing record73SENTINEL_MS = 0.0 # predicted_ms <= this value => spurious timing record74BOOTSTRAP_RNG_SEED = 42 # seeded bootstrap for CI95 of per-position alpha75BOOTSTRAP_ITERS = 200076 77# Expected values from session 08 (baseline sanity checks; >5% deviation warns).78# speedup expectations are checked against BOTH per-prompt mean and median.79EXPECTED_CHECKS: list[dict[str, Any]] = [80 {"run": "final-qwen-q4-solo", "metric": "tps_mean", "exp": 53.4, "label": "qwen-q4 solo"},81 {"run": "final-qwen-q4-vanilla17b", "metric": "tps_mean", "exp": 75.5, "label": "vanilla17b"},82 {"run": "final-qwen-q4-vanilla17b", "metric": "speedup", "exp": 1.41, "label": "vanilla17b x"},83 {84 "run": "final-qwen-q4-vanilla17b",85 "metric": "alpha",86 "exp": 0.725,87 "label": "vanilla17b alpha",88 },89 {"run": "final-qwen-q4-eagle3", "metric": "tps_mean", "exp": 74.4, "label": "eagle3"},90 {"run": "final-qwen-q4-eagle3", "metric": "speedup", "exp": 1.39, "label": "eagle3 x"},91 {"run": "final-qwen-q4-eagle3", "metric": "alpha", "exp": 0.440, "label": "eagle3 alpha"},92 {"run": "final-qwen-q4-dspark-p0", "metric": "tps_mean", "exp": 87.7, "label": "dspark-p0"},93 {"run": "final-qwen-q4-dspark-p0", "metric": "speedup", "exp": 1.64, "label": "dspark-p0 x"},94 {"run": "final-qwen-q4-dspark-p0", "metric": "alpha", "exp": 0.616, "label": "dspark-p0 alpha"},95 {"run": "final-qwen-q4-dspark-p6", "metric": "tps_mean", "exp": 80.8, "label": "dspark-p6"},96 {"run": "final-qwen-q4-dspark-p6", "metric": "speedup", "exp": 1.51, "label": "dspark-p6 x"},97 {"run": "final-qwen-q4-dspark-p6", "metric": "alpha", "exp": 0.714, "label": "dspark-p6 alpha"},98 {"run": "final-qwen-q5-solo", "metric": "tps_mean", "exp": 46.8, "label": "qwen-q5 solo"},99 {"run": "final-qwen-q5-eagle3", "metric": "tps_mean", "exp": 68.3, "label": "q5 eagle3"},100 {"run": "final-qwen-q5-eagle3", "metric": "speedup", "exp": 1.46, "label": "q5 eagle3 x"},101 {"run": "final-qwen-q5-dspark-p0", "metric": "tps_mean", "exp": 80.7, "label": "q5 dspark-p0"},102 {"run": "final-qwen-q5-dspark-p0", "metric": "speedup", "exp": 1.72, "label": "q5 dspark-p0 x"},103 {"run": "final-qwen-q8-solo", "metric": "tps_mean", "exp": 32.4, "label": "qwen-q8 solo"},104 {"run": "final-qwen-q8-eagle3", "metric": "tps_mean", "exp": 52.8, "label": "q8 eagle3"},105 {"run": "final-qwen-q8-eagle3", "metric": "speedup", "exp": 1.63, "label": "q8 eagle3 x"},106 {"run": "final-gemma-q4-solo", "metric": "tps_mean", "exp": 34.9, "label": "gemma-q4 solo"},107 {108 "run": "final-gemma-q4-dflash-f16",109 "metric": "speedup",110 "exp": 2.24,111 "label": "gemma-q4 dflash-f16 x",112 },113 {"run": "final-gemma-q4-mtp", "metric": "speedup", "exp": 2.5, "label": "gemma-q4 mtp x"},114]115 116# Short human labels for drafters (figures/tables).117DRAFTER_LABELS: dict[str, str] = {118 "solo": "solo",119 "vanilla17b": "Vanilla-1.7B",120 "eagle3": "EAGLE-3",121 "dflash-f16": "DFlash-F16",122 "dflash-q4": "DFlash-Q4",123 "dflash-q8": "DFlash-Q8",124 "dspark-p0": "DSpark p=0.0",125 "dspark-p2": "DSpark p=0.2",126 "dspark-p4": "DSpark p=0.4",127 "dspark-p6": "DSpark p=0.6",128 "mtp": "MTP",129}130 131DOMAINS = ("math", "code", "chat")132 133# Colorblind-safe categorical palette (Okabe & Ito 2008). The first three134# entries double as the domain colors; figures reuse this palette for the135# drafter families so every figure speaks the same visual language.136OKABE_ITO = (137 "#0072B2", # blue138 "#D55E00", # vermillion139 "#009E73", # bluish green140 "#E69F00", # orange141 "#56B4E9", # sky blue142 "#CC79A7", # reddish purple143 "#000000", # black144 "#F0E442", # yellow (sparingly: low contrast on white)145)146 147# Published M2 Pro break-even acceptance rates (paper #32, Table 5, Bielik et148# al.): ranges across drafter/dataset combinations. The compact 40-77% band149# (k=2..4) is what figures/README use for the comparison.150M2PRO_ABE = {2: (38.0, 52.8), 4: (77.7, 90.1)} # exact Table 5 ranges151M2PRO_BAND = (0.40, 0.77) # compact k=2..4 range used in F3152 153 154# --------------------------------------------------------------------------- #155# Small helpers156# --------------------------------------------------------------------------- #157 158 159def _fmt(x: float | None, nd: int = 2, suffix: str = "") -> str:160 """Format a float for markdown tables, or an em dash when None/NaN."""161 if x is None or (isinstance(x, float) and not np.isfinite(x)):162 return "\u2014"163 return f"{x:.{nd}f}{suffix}"164 165 166def md_table(headers: list[str], rows: list[list[str]]) -> str:167 """Render a GitHub-flavored markdown table."""168 lines = ["| " + " | ".join(headers) + " |"]169 lines.append("|" + "|".join(" " + "-" * (len(h) + 2) + " " for h in headers) + "|")170 for row in rows:171 cells = [str(c) for c in row]172 if len(cells) != len(headers):173 raise ValueError(f"row has {len(cells)} cells, header has {len(headers)}: {cells}")174 lines.append("| " + " | ".join(cells) + " |")175 return "\n".join(lines)176 177 178# --------------------------------------------------------------------------- #179# Config normalization180# --------------------------------------------------------------------------- #181 182 183@dataclass(frozen=True)184class RunInfo:185 run_name: str186 family: str # final | curves | ksweep | baseline187 target: str # qwen-q4 | gemma-q8 ...188 family_name: str # qwen | gemma189 quant: str # q4 | q5 | q8190 drafter: str # normalized drafter id (e.g. "dspark-p0", "dflash-f16", "solo")191 k: int # spec_draft_n_max192 ctx: int193 prompt_path: str194 prompt_set: str195 n_tokens: int196 temperature: float197 top_k: int198 top_p: float199 seed: int200 p_min: float | None201 draft_path: str | None202 203 204_QUANT_RE = re.compile(r"Q([458])[_K]")205 206 207def normalize_target(model_path: str) -> str:208 """Map a model path to '{family}-q{quant}' (e.g. 'qwen-q4')."""209 family = "qwen" if "Qwen" in model_path else "gemma" if "gemma" in model_path else "?"210 m = _QUANT_RE.search(model_path)211 quant = m.group(1) if m else "?"212 if family == "?" or quant == "?":213 raise ValueError(f"cannot normalize model path: {model_path}")214 return f"{family}-q{quant}"215 216 217def normalize_drafter(spec_type: str | None, p_min: float | None, draft_path: str | None) -> str:218 """Map spec_type (+ p_min and draft quant) to the normalized drafter id."""219 st = (spec_type or "none").lower()220 if st in ("none", "", "vanilla"):221 return "solo"222 if st == "draft-simple":223 return "vanilla17b"224 if st == "draft-eagle3":225 return "eagle3"226 if st == "draft-mtp":227 return "mtp"228 if st == "draft-dspark":229 p = 0.0 if p_min is None else p_min230 return f"dspark-p{int(round(p * 10))}"231 if st == "draft-dflash":232 path = draft_path or ""233 if "Q4_K_M" in path or "-Q4" in path:234 return "dflash-q4"235 if "Q8_0" in path or "-Q8" in path:236 return "dflash-q8"237 return "dflash-f16" # F16 (or unnamed block gguf) is the default DFlash238 raise ValueError(f"unknown spec_type: {spec_type}")239 240 241def prompt_set_id(prompt_path: str) -> str:242 """Return a content identity for a prompt file, not just its path spelling."""243 raw = Path(prompt_path)244 candidates = [raw]245 if not raw.is_absolute():246 candidates.append(Path.cwd() / raw)247 for candidate in candidates:248 try:249 if candidate.is_file():250 return "sha256:" + hashlib.sha256(candidate.read_bytes()).hexdigest()251 except OSError:252 continue253 # Missing files remain distinguishable and cannot accidentally match a254 # different file that happens to have the same basename.255 return "path:" + os.path.normpath(prompt_path)256 257 258def parse_run_info(run_name: str, cfg: dict[str, Any]) -> RunInfo:259 """Normalize one run directory into a RunInfo."""260 prefix = run_name.split("-")[0]261 family = "baseline" if prefix in ("baseline", "ctx2048") else prefix262 if family not in ("final", "curves", "ksweep", "baseline"):263 raise ValueError(f"unexpected run prefix: {run_name}")264 target = normalize_target(cfg["model"])265 fam_name, quant = target.split("-")266 p_min = cfg.get("spec_draft_p_min")267 drafter = normalize_drafter(cfg.get("spec_type"), p_min, cfg.get("draft"))268 k = int(cfg.get("spec_draft_n_max") or 0)269 sampling = cfg.get("sampling") or {}270 prompt_path = str(cfg.get("prompts") or "")271 return RunInfo(272 run_name=run_name,273 family=family,274 target=target,275 family_name=fam_name,276 quant=quant,277 drafter=drafter,278 k=k,279 ctx=int(cfg.get("ctx") or 0),280 prompt_path=prompt_path,281 prompt_set=prompt_set_id(prompt_path),282 n_tokens=int(cfg.get("n_tokens") or 0),283 temperature=float(sampling.get("temperature", cfg.get("temperature", 0.0))),284 top_k=int(sampling.get("top_k", cfg.get("top_k", 0))),285 top_p=float(sampling.get("top_p", cfg.get("top_p", 0.0))),286 seed=int(sampling.get("seed", cfg.get("seed", 0))),287 p_min=p_min,288 draft_path=cfg.get("draft"),289 )290 291 292# --------------------------------------------------------------------------- #293# Data loading294# --------------------------------------------------------------------------- #295 296 297def load_records(run_dir: str) -> list[dict[str, Any]]:298 """Load results.jsonl (tolerates stray non-UTF8 bytes / broken lines)."""299 records: list[dict[str, Any]] = []300 corrupt = 0301 with open(os.path.join(run_dir, "results.jsonl"), errors="replace") as fh:302 for line in fh:303 line = line.strip()304 if not line:305 continue306 try:307 records.append(json.loads(line))308 except json.JSONDecodeError:309 corrupt += 1310 if corrupt:311 logger.warning(" %s: skipped %d unparseable lines", os.path.basename(run_dir), corrupt)312 return records313 314 315def count_error_records(run_dir: str) -> int:316 """Count persisted error records, including errors from resumed attempts."""317 path = os.path.join(run_dir, "errors.jsonl")318 if not os.path.isfile(path):319 return 0320 with open(path, errors="replace") as fh:321 return sum(1 for line in fh if line.strip())322 323 324def is_sentinel(rec: dict[str, Any]) -> bool:325 """True when a record carries the spurious 1e6 tok/s timing marker."""326 tps = rec.get("tok_per_s")327 pred = rec.get("predicted_ms")328 return (isinstance(tps, (int, float)) and tps >= SENTINEL_TPS) or (329 isinstance(pred, (int, float)) and pred <= SENTINEL_MS330 )331 332 333@dataclass334class CleanStats:335 kept: list[dict[str, Any]]336 excluded: int337 total: int338 339 340def clean_records(records: list[dict[str, Any]]) -> CleanStats:341 kept = [r for r in records if not is_sentinel(r)]342 return CleanStats(kept=kept, excluded=len(records) - len(kept), total=len(records))343 344 345_ACCEPT_RE = re.compile(r"draft acceptance = ([\d.]+)")346_POS_RE = re.compile(r"acc per pos = \(([\d.,\s]+)\)")347 348 349def parse_server_log(run_dir: str) -> tuple[list[float], list[list[float]]]:350 """Extract per-request acceptance + per-position vectors from server.log."""351 alphas: list[float] = []352 positions: list[list[float]] = []353 with open(os.path.join(run_dir, "server.log"), errors="replace") as fh:354 for line in fh:355 m = _ACCEPT_RE.search(line)356 if m:357 alphas.append(float(m.group(1)))358 m2 = _POS_RE.search(line)359 if m2:360 positions.append([float(x) for x in m2.group(1).split(",")])361 if len(alphas) != len(positions):362 raise RuntimeError(363 f"{os.path.basename(run_dir)}: {len(alphas)} acceptance lines vs "364 f"{len(positions)} per-position lines"365 )366 return alphas, positions367 368 369@dataclass370class LogMatch:371 run_name: str372 k: int373 matched: int374 log_lines: int375 max_alpha_diff: float376 position_by_id: dict[str, list[float]]377 domain_by_id: dict[str, str]378 379 380def match_log_to_records(381 runs_dir_path: str, run_name: str, records: list[dict[str, Any]], k: int382) -> LogMatch:383 """Assign the i-th log line to the i-th record with non-None alpha.384 385 Gemini timing-sentinel records have ``alpha = None`` and no log line, so386 they are skipped while consuming log entries (verified: max |diff| <= 5e-5).387 """388 alphas, positions = parse_server_log(os.path.join(runs_dir_path, run_name))389 rec_with_alpha = [r for r in records if r.get("alpha") is not None]390 if len(alphas) != len(rec_with_alpha):391 raise RuntimeError(392 f"{run_name}: {len(alphas)} log lines vs {len(rec_with_alpha)} records with alpha"393 )394 max_diff = 0.0395 position_by_id: dict[str, list[float]] = {}396 domain_by_id: dict[str, str] = {}397 for rec, log_alpha, pos in zip(rec_with_alpha, alphas, positions, strict=True):398 max_diff = max(max_diff, abs(rec["alpha"] - log_alpha))399 if len(pos) != k:400 logger.warning(" %s: pos line has %d values, expected k=%d", run_name, len(pos), k)401 position_by_id[rec["id"]] = pos402 domain_by_id[rec["id"]] = rec["domain"]403 return LogMatch(404 run_name=run_name,405 k=k,406 matched=len(rec_with_alpha),407 log_lines=len(alphas),408 max_alpha_diff=max_diff,409 position_by_id=position_by_id,410 domain_by_id=domain_by_id,411 )412 413 414# --------------------------------------------------------------------------- #415# Aggregation helpers416# --------------------------------------------------------------------------- #417 418 419def basic_stats(values: list[float]) -> dict[str, float]:420 arr = np.asarray(values, dtype=float)421 return {422 "mean": float(arr.mean()),423 "median": float(np.median(arr)),424 "p95": float(np.percentile(arr, 95)),425 }426 427 428def alpha_stats(records: list[dict[str, Any]]) -> dict[str, float] | None:429 alphas = [r["alpha"] for r in records if r.get("alpha") is not None]430 if not alphas:431 return None432 arr = np.asarray(alphas)433 return {"mean": float(arr.mean()), "median": float(np.median(arr))}434 435 436def tau_mean(records: list[dict[str, Any]]) -> float | None:437 taus = [r["tau"] for r in records if r.get("tau") is not None]438 return float(np.mean(taus)) if taus else None439 440 441def bootstrap_ci(values: list[float], seed: int = BOOTSTRAP_RNG_SEED) -> tuple[float, float]:442 """Percentile-bootstrap 95% CI of the mean."""443 arr = np.asarray(values, dtype=float)444 rng = np.random.default_rng(seed)445 samples = np.empty(BOOTSTRAP_ITERS)446 for i in range(BOOTSTRAP_ITERS):447 samples[i] = rng.choice(arr, size=len(arr), replace=True).mean()448 lo, hi = np.percentile(samples, [2.5, 97.5])449 return float(lo), float(hi)450 451 452def ols_breakeven(alphas: list[float], tps: list[float], baseline: float) -> dict[str, float]:453 """OLS TPS = a + beta*alpha and the break-even acceptance rate.454 455 ``alpha_be = (baseline - a) / beta``; CI95 from the OLS covariance of456 (a, beta) via the delta method (paper #32 uses the same approach).457 Returns raw coefficients even when ``alpha_be`` is outside [0, 1]458 (value > 1 => unreachable, < 0 => always above baseline).459 """460 x = np.asarray(alphas, dtype=float)461 y = np.asarray(tps, dtype=float)462 if len(x) < 3:463 return {464 "n": len(x),465 "beta": np.nan,466 "intercept": np.nan,467 "r2": np.nan,468 "alpha_be": np.nan,469 "ci95": np.nan,470 }471 (beta, a), cov = np.polyfit(x, y, 1, cov=True)472 yhat = a + beta * x473 ss_res = float(np.sum((y - yhat) ** 2))474 ss_tot = float(np.sum((y - y.mean()) ** 2))475 r2 = 1.0 - ss_res / ss_tot if ss_tot > 0 else np.nan476 alpha_be = (baseline - a) / beta477 # delta method: d(alpha_be)/d(a) = -1/beta ; d(alpha_be)/d(beta) = -(base-a)/beta^2478 da, db = -1.0 / beta, -(baseline - a) / (beta * beta)479 var = da**2 * cov[0, 0] + db**2 * cov[1, 1] + 2.0 * da * db * cov[0, 1]480 ci = 1.96 * float(np.sqrt(max(var, 0.0)))481 return {482 "n": len(x),483 "beta": float(beta),484 "intercept": float(a),485 "r2": float(r2),486 "alpha_be": float(alpha_be),487 "ci95": float(ci),488 }489 490 491def speedup_vs(492 recs: list[dict[str, Any]], solo_tps_by_id: dict[str, float]493) -> dict[str, float] | None:494 """Per-prompt matched speedup: mean/median of ratios + ratio of means."""495 ratios = [r["tok_per_s"] / solo_tps_by_id[r["id"]] for r in recs if r["id"] in solo_tps_by_id]496 if not ratios:497 return None498 arr = np.asarray(ratios)499 matched = [r for r in recs if r["id"] in solo_tps_by_id]500 return {501 "mean": float(arr.mean()),502 "median": float(np.median(arr)),503 "agg": float(504 np.mean([r["tok_per_s"] for r in matched]) / np.mean(list(solo_tps_by_id.values()))505 ),506 "n_match": len(ratios),507 }508 509 510# --------------------------------------------------------------------------- #511# final runs512# --------------------------------------------------------------------------- #513 514 515def build_final_stats(runs_dir_path: str) -> dict[str, Any]:516 """Aggregate the 26 final runs: per-config and per-domain statistics."""517 infos: dict[str, RunInfo] = {}518 for name in sorted(os.listdir(runs_dir_path)):519 cfg_path = os.path.join(runs_dir_path, name, "config.json")520 if os.path.isfile(cfg_path):521 with open(cfg_path) as fh:522 infos[name] = parse_run_info(name, json.load(fh))523 524 final_infos = {n: i for n, i in infos.items() if i.family == "final"}525 if len(final_infos) != 26:526 logger.warning("expected 26 final runs, found %d", len(final_infos))527 528 # Load + clean every final run, log exclusions.529 records_by_run: dict[str, list[dict[str, Any]]] = {}530 exclusions: dict[str, dict[str, int]] = {}531 for name in sorted(final_infos):532 cs = clean_records(load_records(os.path.join(runs_dir_path, name)))533 records_by_run[name] = cs.kept534 exclusions[name] = {"total": cs.total, "kept": len(cs.kept), "excluded": cs.excluded}535 if cs.excluded:536 logger.info(537 "excluded %d sentinel records in %s (%d kept)", cs.excluded, name, len(cs.kept)538 )539 540 # Baseline solos: per target -> per domain mean tok/s + per-prompt map.541 solo_by_target: dict[str, dict[str, Any]] = {}542 for name, info in final_infos.items():543 if info.drafter != "solo":544 continue545 recs = records_by_run[name]546 by_domain = {d: [r for r in recs if r["domain"] == d] for d in DOMAINS}547 solo_by_target[info.target] = {548 "by_domain": by_domain,549 "tps_mean": {550 d: float(np.mean([r["tok_per_s"] for r in by_domain[d]])) for d in DOMAINS551 },552 "tps_mean_all": float(np.mean([r["tok_per_s"] for r in recs])),553 "tps_by_id": {r["id"]: r["tok_per_s"] for r in recs},554 }555 if len(solo_by_target) != 6:556 logger.warning(557 "expected 6 solo baselines (qwen/gemma x q4/q5/q8), found %d", len(solo_by_target)558 )559 560 summary: dict[str, Any] = {}561 for name, info in sorted(final_infos.items()):562 recs = records_by_run[name]563 entry: dict[str, Any] = {564 "run": name,565 "target": info.target,566 "family": info.family_name,567 "quant": info.quant,568 "drafter": info.drafter,569 "k": info.k,570 "ctx": info.ctx,571 "n": len(recs),572 "tok_per_s": basic_stats([r["tok_per_s"] for r in recs]),573 "alpha": alpha_stats(recs),574 "tau_mean": tau_mean(recs),575 "ttft_ms": basic_stats([r["prompt_ms"] for r in recs]),576 }577 vram_path = os.path.join(runs_dir_path, name, "vram.json")578 if os.path.isfile(vram_path):579 with open(vram_path) as fh:580 vram = json.load(fh)581 entry["vram_max_mib"] = vram.get("max_gpu_mib")582 entry["power_max_w"] = vram.get("max_power_w")583 metrics_path = os.path.join(runs_dir_path, name, "metrics.json")584 if os.path.isfile(metrics_path):585 with open(metrics_path) as fh:586 metrics = json.load(fh)587 entry["duration_s"] = metrics.get("duration_s")588 entry["errors"] = count_error_records(os.path.join(runs_dir_path, name))589 590 solo = solo_by_target.get(info.target)591 per_domain: dict[str, Any] = {}592 for d in DOMAINS:593 dr = [r for r in recs if r["domain"] == d]594 cell: dict[str, Any] = {595 "n": len(dr),596 "tok_per_s": basic_stats([r["tok_per_s"] for r in dr]),597 "alpha": alpha_stats(dr),598 "tau_mean": tau_mean(dr),599 }600 if solo is not None and info.drafter != "solo":601 cell["speedup"] = speedup_vs(dr, solo["tps_by_id"])602 per_domain[d] = cell603 entry["per_domain"] = per_domain604 if solo is not None and info.drafter != "solo":605 entry["speedup"] = speedup_vs(recs, solo["tps_by_id"])606 summary[name] = entry607 return {"summary": summary, "exclusions": exclusions, "solo_by_target": solo_by_target}608 609 610def baseline_compatibility_key(info: RunInfo) -> tuple[str, int, str, int, float, int, float, int]:611 """Identity used to pair a ksweep run with a target-only baseline."""612 return (613 info.target,614 info.ctx,615 info.prompt_set,616 info.n_tokens,617 info.temperature,618 info.top_k,619 info.top_p,620 info.seed,621 )622 623 624def build_baseline_stats(625 runs_dir_path: str,626) -> tuple[627 dict[tuple[str, int, str, int, float, int, float, int], dict[str, Any]], list[dict[str, Any]]628]:629 """Load explicitly named contextual target-only baseline runs.630 631 Baselines are kept separate from the six final-run solo baselines. A632 baseline can be used for break-even only when its full protocol identity633 matches the ksweep run (target, context, prompt-file content, and sampling634 settings). Duplicate identities are rejected rather than selected635 implicitly.636 """637 by_key: dict[tuple[str, int, str, int, float, int, float, int], dict[str, Any]] = {}638 summaries: list[dict[str, Any]] = []639 for name in sorted(os.listdir(runs_dir_path)):640 # Only the explicit baseline-* namespace is analytical input. A641 # previous interrupted controller left a ctx2048-* scratch run; do642 # not let an accidental rerun compete with the canonical baseline.643 if not name.startswith("baseline-"):644 continue645 cfg_path = os.path.join(runs_dir_path, name, "config.json")646 if not os.path.isfile(cfg_path):647 continue648 with open(cfg_path) as fh:649 info = parse_run_info(name, json.load(fh))650 if info.family != "baseline" or info.drafter != "solo":651 continue652 653 cs = clean_records(load_records(os.path.join(runs_dir_path, name)))654 recs = cs.kept655 by_domain = {d: [r for r in recs if r.get("domain") == d] for d in DOMAINS}656 tps = [r["tok_per_s"] for r in recs if r.get("tok_per_s") is not None]657 tps_by_id = {r["id"]: r["tok_per_s"] for r in recs if r.get("tok_per_s") is not None}658 metrics_path = os.path.join(runs_dir_path, name, "metrics.json")659 errors = count_error_records(os.path.join(runs_dir_path, name))660 if errors == 0 and os.path.isfile(metrics_path):661 with open(metrics_path) as fh:662 errors = json.load(fh).get("errors", 0)663 664 key = baseline_compatibility_key(info)665 if key in by_key:666 previous = by_key[key]["run"]667 raise RuntimeError(668 f"duplicate compatible baselines for {info.target} ctx={info.ctx}: "669 f"{previous} and {name}"670 )671 672 entry: dict[str, Any] = {673 "run": name,674 "target": info.target,675 "ctx": info.ctx,676 "prompt_path": info.prompt_path,677 "prompt_set": info.prompt_set,678 "n_tokens": info.n_tokens,679 "sampling": {680 "temperature": info.temperature,681 "top_k": info.top_k,682 "top_p": info.top_p,683 "seed": info.seed,684 },685 "n": len(recs),686 "excluded": cs.excluded,687 "errors": errors,688 "tok_per_s": basic_stats(tps) if tps else None,689 "tok_per_s_by_id": tps_by_id,690 "tok_per_s_by_domain": {691 d: basic_stats([r["tok_per_s"] for r in by_domain[d]]) if by_domain[d] else None692 for d in DOMAINS693 },694 }695 by_key[key] = entry696 summaries.append({k: v for k, v in entry.items() if k != "tok_per_s_by_id"})697 return by_key, summaries698 699 700# --------------------------------------------------------------------------- #701# curves / ksweep: per-position acceptance + ksweep break-even702# --------------------------------------------------------------------------- #703 704 705def build_position_data(runs_dir_path: str) -> dict[str, Any]:706 """Parse curves/ksweep logs and aggregate per (run, domain, position)."""707 out: dict[str, Any] = {}708 for name in sorted(os.listdir(runs_dir_path)):709 cfg_path = os.path.join(runs_dir_path, name, "config.json")710 if not os.path.isfile(cfg_path):711 continue712 with open(cfg_path) as fh:713 info = parse_run_info(name, json.load(fh))714 if info.family not in ("curves", "ksweep"):715 continue716 cs = clean_records(load_records(os.path.join(runs_dir_path, name)))717 if cs.excluded:718 logger.info("excluded %d sentinel records in %s", cs.excluded, name)719 match = match_log_to_records(runs_dir_path, name, cs.kept, info.k)720 # Group per (domain, position).721 per_domain: dict[str, list[dict[str, Any]]] = {d: [] for d in DOMAINS}722 for rid, pos in match.position_by_id.items():723 dom = match.domain_by_id[rid]724 for p, val in enumerate(pos, start=1):725 per_domain[dom].append({"pos": p, "val": val})726 dom_out: dict[str, Any] = {}727 for d in DOMAINS:728 pos_stats: list[dict[str, Any]] = []729 for p in range(1, info.k + 1):730 vals = [e["val"] for e in per_domain[d] if e["pos"] == p]731 if not vals:732 continue733 lo, hi = bootstrap_ci(vals)734 pos_stats.append(735 {736 "position": p,737 "n": len(vals),738 "alpha_mean": float(np.mean(vals)),739 "alpha_median": float(np.median(vals)),740 "ci95_low": lo,741 "ci95_high": hi,742 }743 )744 dom_out[d] = pos_stats745 out[name] = {746 "family": info.family,747 "target": info.target,748 "drafter": info.drafter,749 "k": info.k,750 "n_excluded": cs.excluded,751 "n_matched": match.matched,752 "log_lines": match.log_lines,753 "max_alpha_diff": match.max_alpha_diff,754 "per_domain": dom_out,755 }756 return out757 758 759def build_breakeven(760 runs_dir_path: str,761 baseline_by_key: dict[tuple[str, int, str, int, float, int, float, int], dict[str, Any]],762) -> dict[str, Any]:763 """OLS alpha_be per ksweep run using a protocol-compatible baseline."""764 out: dict[str, Any] = {}765 for name in sorted(os.listdir(runs_dir_path)):766 cfg_path = os.path.join(runs_dir_path, name, "config.json")767 if not os.path.isfile(cfg_path):768 continue769 with open(cfg_path) as fh:770 info = parse_run_info(name, json.load(fh))771 if info.family != "ksweep":772 continue773 baseline = baseline_by_key.get(baseline_compatibility_key(info))774 if baseline is None:775 logger.warning(776 "skipping %s: no target-only baseline matches target=%s ctx=%d "777 "prompt_set=%s sampling=%s",778 name,779 info.target,780 info.ctx,781 info.prompt_set,782 baseline_compatibility_key(info)[3:],783 )784 continue785 recs = [786 r787 for r in clean_records(load_records(os.path.join(runs_dir_path, name))).kept788 if r.get("alpha") is not None789 ]790 matched = [r for r in recs if r.get("id") in baseline["tok_per_s_by_id"]]791 if len(matched) < 3:792 logger.warning("skipping %s: fewer than 3 baseline-matched observations", name)793 continue794 baseline_by_id = baseline["tok_per_s_by_id"]795 baseline_all = float(np.mean([baseline_by_id[r["id"]] for r in matched]))796 pooled = ols_breakeven(797 [r["alpha"] for r in matched], [r["tok_per_s"] for r in matched], baseline_all798 )799 per_domain: dict[str, Any] = {}800 baseline_by_domain: dict[str, float] = {}801 baseline_n_by_domain: dict[str, int] = {}802 for d in DOMAINS:803 dr = [r for r in matched if r["domain"] == d]804 baseline_values = [baseline_by_id[r["id"]] for r in dr]805 baseline_by_domain[d] = float(np.mean(baseline_values)) if baseline_values else np.nan806 baseline_n_by_domain[d] = len(baseline_values)807 per_domain[d] = ols_breakeven(808 [r["alpha"] for r in dr],809 [r["tok_per_s"] for r in dr],810 baseline_by_domain[d],811 )812 out[name] = {813 "target": info.target,814 "drafter": info.drafter,815 "k": info.k,816 "ctx": info.ctx,817 "prompt_set": info.prompt_set,818 "baseline_run": baseline["run"],819 "baseline_ctx": baseline["ctx"],820 "baseline_prompt_set": baseline["prompt_set"],821 "baseline_n": baseline["n"],822 "baseline_n_matched": len(matched),823 "baseline_n_by_domain": baseline_n_by_domain,824 "baseline_excluded": baseline["excluded"],825 "baseline_all": baseline_all,826 "baseline_by_domain": baseline_by_domain,827 "pooled": pooled,828 "per_domain": per_domain,829 }830 return out831 832 833# --------------------------------------------------------------------------- #834# Markdown tables835# --------------------------------------------------------------------------- #836 837 838def _config_rows(839 final_summary: dict[str, Any], family_name: str840) -> list[tuple[str, dict[str, Any]]]:841 """(run-name, entry) pairs for one family, sorted by quant then drafter."""842 rows = [(n, e) for n, e in final_summary.items() if e["family"] == family_name]843 order = {"q4": 0, "q5": 1, "q8": 2}844 drafter_order = {845 "solo": 0,846 "vanilla17b": 1,847 "eagle3": 2,848 "dflash-f16": 3,849 "dflash-q4": 4,850 "dflash-q8": 5,851 "dspark-p0": 6,852 "dspark-p2": 7,853 "dspark-p4": 8,854 "dspark-p6": 9,855 "mtp": 10,856 }857 rows.sort(key=lambda r: (order.get(r[1]["quant"], 9), drafter_order.get(r[1]["drafter"], 99)))858 return rows859 860 861def _short_name(run_name: str) -> str:862 return run_name.split("-", 2)[-1]863 864 865def write_t2_speedup(final_summary: dict[str, Any], out_dir: Path) -> None:866 """tok/s and speedup vs solo, per config x domain (Qwen and Gemma blocks)."""867 headers = [868 "Config",869 "n",870 "Math tok/s",871 "Math x",872 "Code tok/s",873 "Code x",874 "Chat tok/s",875 "Chat x",876 "All tok/s",877 "All x",878 "alpha all",879 ]880 sections: list[str] = []881 for fam, title in (882 ("qwen", "### Qwen3-8B (baseline: same-quant solo)"),883 ("gemma", "### Gemma 4 12B (baseline: same-quant solo)"),884 ):885 lines = [title, ""]886 rows: list[list[str]] = []887 for name, entry in _config_rows(final_summary, fam):888 cells = [_short_name(name), str(entry["n"])]889 for d in (*DOMAINS, "all"):890 dom = entry if d == "all" else entry["per_domain"][d]891 cells.append(_fmt(dom["tok_per_s"]["mean"], 1))892 sup = entry.get("speedup") if d == "all" else dom.get("speedup")893 sup = sup.get("mean") if isinstance(sup, dict) else None894 cells.append(_fmt(sup, 2, "x") if sup is not None else "\u2014")895 alpha = entry.get("alpha")896 cells.append(_fmt(alpha["mean"], 3) if alpha else "\u2014")897 rows.append(cells)898 lines.append(md_table(headers, rows))899 lines.append("")900 sections.append("\n".join(lines))901 (out_dir / "t2_speedup.md").write_text("\n".join(sections), encoding="utf-8")902 903 904def write_t3_alpha_tau(final_summary: dict[str, Any], out_dir: Path) -> None:905 """alpha and tau per config x domain (drafter configs only)."""906 headers = [907 "Config",908 "Math alpha",909 "Math tau",910 "Code alpha",911 "Code tau",912 "Chat alpha",913 "Chat tau",914 "All alpha",915 "All tau",916 ]917 sections: list[str] = []918 for fam, title in (("qwen", "### Qwen3-8B"), ("gemma", "### Gemma 4 12B")):919 lines = [title, ""]920 rows: list[list[str]] = []921 for name, entry in _config_rows(final_summary, fam):922 if entry["drafter"] == "solo":923 continue924 cells = [_short_name(name)]925 for d in (*DOMAINS, "all"):926 dom = entry if d == "all" else entry["per_domain"][d]927 a = dom.get("alpha")928 cells.append(_fmt(a["mean"], 3) if a else "\u2014")929 cells.append(_fmt(dom.get("tau_mean"), 0))930 rows.append(cells)931 lines.append(md_table(headers, rows))932 lines.append("")933 sections.append("\n".join(lines))934 (out_dir / "t3_alpha_tau.md").write_text("\n".join(sections), encoding="utf-8")935 936 937def _summarize_all(entry: dict[str, Any]) -> tuple[str, str, str, str]:938 tps = _fmt(entry["tok_per_s"]["mean"], 1)939 sup = entry.get("speedup")940 sup_s = _fmt(sup["mean"], 2, "x") if sup else "\u2014"941 a = entry.get("alpha")942 a_s = _fmt(a["mean"], 3) if a else "\u2014"943 return tps, sup_s, a_s, _fmt(entry.get("tau_mean"), 0)944 945 946def write_t4_quantization(final_summary: dict[str, Any], out_dir: Path) -> None:947 """Target-quant x drafter interaction, plus gemma-q4 draft-quant effect."""948 headers = ["Target quant", "Drafter", "n", "tok/s", "x vs solo", "alpha", "tau"]949 lines: list[str] = []950 for fam, title in (951 ("qwen", "### Qwen3-8B โ target quant (q4/q5/q8) x drafter"),952 ("gemma", "### Gemma 4 12B โ target quant x drafter"),953 ):954 lines.append(title)955 lines.append("")956 rows: list[list[str]] = []957 for _, entry in _config_rows(final_summary, fam):958 if entry["drafter"] == "solo":959 continue960 tps, sup, a, tau = _summarize_all(entry)961 rows.append(962 [963 entry["quant"],964 DRAFTER_LABELS[entry["drafter"]],965 str(entry["n"]),966 tps,967 sup,968 a,969 tau,970 ]971 )972 lines.append(md_table(headers, rows))973 lines.append("")974 975 lines.append("### Gemma-4 Q4 โ draft quantization effect (DFlash drafts, final runs)")976 lines.append("")977 rows: list[list[str]] = []978 for _, entry in _config_rows(final_summary, "gemma"):979 if (980 entry["quant"] != "q4"981 or entry["drafter"] == "solo"982 or not entry["drafter"].startswith("dflash")983 ):984 continue985 tps, sup, a, tau = _summarize_all(entry)986 rows.append(987 [988 DRAFTER_LABELS[entry["drafter"]],989 str(entry["n"]),990 tps,991 sup,992 a,993 tau,994 _fmt(entry.get("vram_max_mib"), 0),995 ]996 )997 lines.append(998 md_table(999 ["Drafter (draft quant)", "n", "tok/s", "x vs solo", "alpha", "tau", "max VRAM (MiB)"],1000 rows,1001 )1002 )1003 lines.append("")1004 lines.append("_F16 = 1.47 GB draft, Q4_K_M = 0.44 GB, Q8_0 = 0.79 GB (model-hashes.json)._")1005 lines.append("")1006 (out_dir / "t4_quantization.md").write_text("\n".join(lines), encoding="utf-8")1007 1008 1009def write_t5_hardware(final_summary: dict[str, Any], out_dir: Path) -> None:1010 """TTFT, VRAM, power, duration per config."""1011 headers = [1012 "Config",1013 "TTFT mean (ms)",1014 "TTFT median (ms)",1015 "TTFT p95 (ms)",1016 "max VRAM (MiB)",1017 "max power (W)",1018 "duration (s)",1019 ]1020 rows: list[list[str]] = []1021 for name, entry in sorted(final_summary.items()):1022 tt = entry["ttft_ms"]1023 rows.append(1024 [1025 name,1026 _fmt(tt["mean"], 1),1027 _fmt(tt["median"], 1),1028 _fmt(tt["p95"], 1),1029 _fmt(entry.get("vram_max_mib"), 0),1030 _fmt(entry.get("power_max_w"), 1),1031 _fmt(entry.get("duration_s"), 1),1032 ]1033 )1034 (out_dir / "t5_hardware.md").write_text(1035 md_table(headers, rows)1036 + "\n\n_All timings from clean (non-sentinel) records; VRAM/power from"1037 " vram.json; duration from metrics.json._\n",1038 encoding="utf-8",1039 )1040 1041 1042def write_t6_breakeven(breakeven: dict[str, Any], out_dir: Path) -> None:1043 """ksweep OLS break-even: pooled + per-domain, with baseline provenance."""1044 lines: list[str] = []1045 headers = [1046 "Config",1047 "k",1048 "n",1049 "baseline run",1050 "baseline ctx",1051 "baseline n/match",1052 "baseline (tok/s)",1053 "beta (slope)",1054 "alpha_be",1055 "CI95",1056 "R2",1057 ]1058 rows: list[list[str]] = []1059 for name in sorted(breakeven):1060 be = breakeven[name]1061 pooled = be["pooled"]1062 rows.append(1063 [1064 f"{be['target']}-{be['drafter']}",1065 str(be["k"]),1066 str(pooled["n"]),1067 be["baseline_run"],1068 str(be["baseline_ctx"]),1069 f"{be['baseline_n']}/{be['baseline_n_matched']}",1070 _fmt(be["baseline_all"], 1),1071 _fmt(pooled["beta"], 2),1072 _fmt(pooled["alpha_be"], 3),1073 _fmt(pooled["ci95"], 3),1074 _fmt(pooled["r2"], 3),1075 ]1076 )1077 lines.append("### Pooled (all domains)")1078 lines.append("")1079 lines.append(md_table(headers, rows))1080 lines.append("")1081 1082 lines.append("### Per domain")1083 lines.append("")1084 headers_d = [1085 "Config",1086 "k",1087 "domain",1088 "n",1089 "baseline ctx",1090 "baseline n",1091 "baseline",1092 "beta",1093 "alpha_be",1094 "CI95",1095 "R2",1096 ]1097 rows_d: list[list[str]] = []1098 for name in sorted(breakeven):1099 be = breakeven[name]1100 for d in DOMAINS:1101 b = be["per_domain"][d]1102 rows_d.append(1103 [1104 f"{be['target']}-{be['drafter']}",1105 str(be["k"]),1106 d,1107 str(b["n"]),1108 str(be["baseline_ctx"]),1109 str(be["baseline_n_by_domain"][d]),1110 _fmt(be["baseline_by_domain"][d], 1),1111 _fmt(b["beta"], 2),1112 _fmt(b["alpha_be"], 3),1113 _fmt(b["ci95"], 3),1114 _fmt(b["r2"], 3),1115 ]1116 )1117 lines.append(md_table(headers_d, rows_d))1118 lines.append("")1119 1120 lines.append("### Comparison with M2 Pro (paper #32, Bielik et al., cross-family)")1121 lines.append("")1122 lines.append(1123 "Paper #32 fits `TPS = a + b*alpha` by OLS and defines `alpha_be = (TPS_base - a) / b` "1124 "(its `b` is our `beta`, the 'recovery rate'). Published values (Table 5, ranging "1125 "across drafters and datasets) are **k=2: 38.0-52.8%** and **k=4: 77.7-90.1%**. "1126 "The compact 40-77% range summarizes k=2..4. Our ksweep starts "1127 "at k=5, so the comparison is directional: the k=10 RTX range is below the "1128 "reported k=2 band, while the upper ends at k=5 and k=7 slightly overlap its "1129 "lower edge."1130 )1131 lines.append("")1132 rows_c: list[list[str]] = []1133 for k, (lo, hi) in sorted(M2PRO_ABE.items()):1134 rows_c.append([f"M2 Pro k={k}", f"{lo:.1f}-{hi:.1f}%"])1135 ours: dict[int, list[float]] = {}1136 for be in breakeven.values():1137 ours.setdefault(be["k"], []).append(be["pooled"]["alpha_be"])1138 for k in sorted(ours):1139 vals = [v for v in ours[k] if np.isfinite(v)]1140 if vals:1141 rows_c.append(1142 [1143 f"Ours k={k} (n={len(vals)} configs)",1144 f"{100 * min(vals):.1f}-{100 * max(vals):.1f}%",1145 ]1146 )1147 lines.append(md_table(["Reference", "alpha_be range"], rows_c))1148 lines.append("")1149 lines.append(1150 "_alpha_be > 1.00 = no OLS-reachable break-even; CI95 by delta"1151 " method over the OLS covariance._"1152 )1153 lines.append("")1154 (out_dir / "t6_breakeven.md").write_text("\n".join(lines), encoding="utf-8")1155 1156 1157# --------------------------------------------------------------------------- #1158# Figures1159# --------------------------------------------------------------------------- #1160 1161 1162def _style_figure() -> None:1163 """One visual theme for all figures, authored at final print size.1164 1165 Figures are rendered at ~6.5 in width (the LaTeX textwidth) so pandoc's1166 \\pandocbounded downscale is ~1x and text prints at true 7.5-8 pt instead1167 of the 3-7 pt the old large canvases degraded to.1168 """1169 plt.rcParams.update(1170 {1171 # Serif matches the LaTeX body; STIX covers the math glyphs.1172 "font.family": "serif",1173 "font.serif": ["STIXGeneral", "DejaVu Serif", "Times New Roman"],1174 "mathtext.fontset": "stix",1175 "font.size": 8.0,1176 "axes.titlesize": 8.0,1177 "axes.labelsize": 8.0,1178 "legend.fontsize": 7.5,1179 "xtick.labelsize": 7.5,1180 "ytick.labelsize": 7.5,1181 "figure.dpi": 150,1182 "savefig.dpi": 600, # arXiv wants >= 300; line art is cheap at 6001183 "axes.grid": True,1184 "grid.alpha": 0.25,1185 "grid.linewidth": 0.4,1186 "axes.spines.top": False,1187 "axes.spines.right": False,1188 "axes.linewidth": 0.6,1189 "lines.linewidth": 1.5,1190 "lines.markersize": 4.5,1191 "legend.frameon": False,1192 "savefig.bbox": "tight",1193 "savefig.pad_inches": 0.02,1194 }1195 )1196 1197 1198DOMAIN_COLORS = {"math": OKABE_ITO[0], "code": OKABE_ITO[1], "chat": OKABE_ITO[2]}1199DOMAIN_MARKERS = {"math": "o", "code": "s", "chat": "^"}1200DOMAIN_LINESTYLES = {"math": "-", "code": "--", "chat": ":"}