Team Ai
Modelpublic

reyden009/speculative-decoding-lab

sourceHugging Facemitupdated 2mo agoView on Hugging Face
8likes
analyze_final.py1920 linesDownload Raw Back to scripts
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": ":"}

Showing the first 1,200 of 1920 lines. Download the file for the rest.