Download scripts/analyze_final.py from reyden009/speculative-decoding-lab: direct link, hf CLI and curl.
- Browser
- Download file 74.3 kB
-
https://huggingface.co/reyden009/speculative-decoding-lab/resolve/main/scripts/analyze_final.py
- Command line
-
hf download hf://reyden009/speculative-decoding-lab/scripts/analyze_final.py
-
curl -L -o analyze_final.py https://huggingface.co/reyden009/speculative-decoding-lab/resolve/main/scripts/analyze_final.py
74.3 kB
| #!/usr/bin/env python3 | |
| """analyze_final.py -- Final analysis for the speculative-decoding empirical study. | |
| Consolidates the 67 experimental runs (26 ``final-*``, 20 ``curves-*``, 21 | |
| ``ksweep-*``) into the numbers, tables and figures that the paper cites: | |
| * ``experiments/analysis/summary.json`` consolidated statistics | |
| * ``experiments/analysis/tables/t2..t6_*.md`` paper-ready markdown tables | |
| * ``experiments/analysis/curves/acc_by_pos_*.csv`` per-position acceptance | |
| * ``experiments/analysis/README.md`` methodology notes | |
| * ``manuscript/figures/F1..F4_*.png`` paper figures (300 dpi) | |
| Methodology | |
| ----------- | |
| * **Records**: one JSON line in ``results.jsonl`` == one OK completion. Config | |
| comes from ``config.json`` (model path -> family+quant, ``spec_type`` + | |
| draft path + ``p_min`` -> drafter id). | |
| * **Sentinel exclusion**: llama-server reports ``predicted_per_second = | |
| 1,000,000`` and ``predicted_ms = 0`` for a handful of Gemma completions | |
| (timing quirk, not a real speedup). Every record with ``tok_per_s >= 1e5`` | |
| or ``predicted_ms <= 0`` is excluded from ALL statistics (they also carry | |
| ``alpha/tau/draft_n = None``). ``solo`` configs have ``alpha/tau/draft_n = | |
| None`` by design (no drafter) and are never treated as an anomaly. | |
| * **Log mapping (curves/ksweep)**: the lines ``draft acceptance = ...`` and | |
| ``acc per pos = (...)`` in ``server.log`` appear in the SAME order as the | |
| records with non-None ``alpha`` in ``results.jsonl`` (verified: max | |
| |log - record| = 0.00005, rounding only). Records without ``alpha`` | |
| (Gemma sentinels) have no log line and are skipped. | |
| * **Speedup vs solo**: per-prompt ratio ``tps_draft / tps_solo`` matched by | |
| prompt ``id`` against the ``solo`` run of the SAME family+quant; we report | |
| the mean and median of the per-prompt ratios and the ratio of means. | |
| * **Break-even ``alpha_be``** (paper #32, Bielik et al.): OLS fit | |
| ``TPS = a + beta*alpha`` over per-prompt observations of a ksweep run | |
| (per domain and pooled); ``alpha_be = (TPS_base - a) / beta`` is the | |
| acceptance rate where the regression line crosses a context-compatible | |
| autoregressive baseline. The baseline must match target, context, prompt | |
| set, and sampling protocol; its tok/s is averaged over prompt IDs shared | |
| with the ksweep observations. CI95 via the delta method on the OLS covariance | |
| of ``a`` and ``beta``. ``beta`` is the paper's ``b`` ("recovery rate"): | |
| tok/s gained per unit acceptance. | |
| Only Python 3.12 stdlib + numpy + matplotlib (repo venv). Idempotent. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import hashlib | |
| import json | |
| import logging | |
| import os | |
| import re | |
| import sys | |
| from dataclasses import dataclass | |
| from pathlib import Path | |
| from typing import Any, cast | |
| import matplotlib | |
| import numpy as np | |
| matplotlib.use("Agg") # noqa: E402 (must run before pyplot import) | |
| import matplotlib.pyplot as plt # noqa: E402 | |
| from matplotlib.lines import Line2D # noqa: E402 | |
| logger = logging.getLogger("analyze_final") | |
| # --------------------------------------------------------------------------- # | |
| # Constants | |
| # --------------------------------------------------------------------------- # | |
| SENTINEL_TPS = 1e5 # tok_per_s >= this value => spurious timing record | |
| SENTINEL_MS = 0.0 # predicted_ms <= this value => spurious timing record | |
| BOOTSTRAP_RNG_SEED = 42 # seeded bootstrap for CI95 of per-position alpha | |
| BOOTSTRAP_ITERS = 2000 | |
| # Expected values from session 08 (baseline sanity checks; >5% deviation warns). | |
| # speedup expectations are checked against BOTH per-prompt mean and median. | |
| EXPECTED_CHECKS: list[dict[str, Any]] = [ | |
| {"run": "final-qwen-q4-solo", "metric": "tps_mean", "exp": 53.4, "label": "qwen-q4 solo"}, | |
| {"run": "final-qwen-q4-vanilla17b", "metric": "tps_mean", "exp": 75.5, "label": "vanilla17b"}, | |
| {"run": "final-qwen-q4-vanilla17b", "metric": "speedup", "exp": 1.41, "label": "vanilla17b x"}, | |
| { | |
| "run": "final-qwen-q4-vanilla17b", | |
| "metric": "alpha", | |
| "exp": 0.725, | |
| "label": "vanilla17b alpha", | |
| }, | |
| {"run": "final-qwen-q4-eagle3", "metric": "tps_mean", "exp": 74.4, "label": "eagle3"}, | |
| {"run": "final-qwen-q4-eagle3", "metric": "speedup", "exp": 1.39, "label": "eagle3 x"}, | |
| {"run": "final-qwen-q4-eagle3", "metric": "alpha", "exp": 0.440, "label": "eagle3 alpha"}, | |
| {"run": "final-qwen-q4-dspark-p0", "metric": "tps_mean", "exp": 87.7, "label": "dspark-p0"}, | |
| {"run": "final-qwen-q4-dspark-p0", "metric": "speedup", "exp": 1.64, "label": "dspark-p0 x"}, | |
| {"run": "final-qwen-q4-dspark-p0", "metric": "alpha", "exp": 0.616, "label": "dspark-p0 alpha"}, | |
| {"run": "final-qwen-q4-dspark-p6", "metric": "tps_mean", "exp": 80.8, "label": "dspark-p6"}, | |
| {"run": "final-qwen-q4-dspark-p6", "metric": "speedup", "exp": 1.51, "label": "dspark-p6 x"}, | |
| {"run": "final-qwen-q4-dspark-p6", "metric": "alpha", "exp": 0.714, "label": "dspark-p6 alpha"}, | |
| {"run": "final-qwen-q5-solo", "metric": "tps_mean", "exp": 46.8, "label": "qwen-q5 solo"}, | |
| {"run": "final-qwen-q5-eagle3", "metric": "tps_mean", "exp": 68.3, "label": "q5 eagle3"}, | |
| {"run": "final-qwen-q5-eagle3", "metric": "speedup", "exp": 1.46, "label": "q5 eagle3 x"}, | |
| {"run": "final-qwen-q5-dspark-p0", "metric": "tps_mean", "exp": 80.7, "label": "q5 dspark-p0"}, | |
| {"run": "final-qwen-q5-dspark-p0", "metric": "speedup", "exp": 1.72, "label": "q5 dspark-p0 x"}, | |
| {"run": "final-qwen-q8-solo", "metric": "tps_mean", "exp": 32.4, "label": "qwen-q8 solo"}, | |
| {"run": "final-qwen-q8-eagle3", "metric": "tps_mean", "exp": 52.8, "label": "q8 eagle3"}, | |
| {"run": "final-qwen-q8-eagle3", "metric": "speedup", "exp": 1.63, "label": "q8 eagle3 x"}, | |
| {"run": "final-gemma-q4-solo", "metric": "tps_mean", "exp": 34.9, "label": "gemma-q4 solo"}, | |
| { | |
| "run": "final-gemma-q4-dflash-f16", | |
| "metric": "speedup", | |
| "exp": 2.24, | |
| "label": "gemma-q4 dflash-f16 x", | |
| }, | |
| {"run": "final-gemma-q4-mtp", "metric": "speedup", "exp": 2.5, "label": "gemma-q4 mtp x"}, | |
| ] | |
| # Short human labels for drafters (figures/tables). | |
| DRAFTER_LABELS: dict[str, str] = { | |
| "solo": "solo", | |
| "vanilla17b": "Vanilla-1.7B", | |
| "eagle3": "EAGLE-3", | |
| "dflash-f16": "DFlash-F16", | |
| "dflash-q4": "DFlash-Q4", | |
| "dflash-q8": "DFlash-Q8", | |
| "dspark-p0": "DSpark p=0.0", | |
| "dspark-p2": "DSpark p=0.2", | |
| "dspark-p4": "DSpark p=0.4", | |
| "dspark-p6": "DSpark p=0.6", | |
| "mtp": "MTP", | |
| } | |
| DOMAINS = ("math", "code", "chat") | |
| # Colorblind-safe categorical palette (Okabe & Ito 2008). The first three | |
| # entries double as the domain colors; figures reuse this palette for the | |
| # drafter families so every figure speaks the same visual language. | |
| OKABE_ITO = ( | |
| "#0072B2", # blue | |
| "#D55E00", # vermillion | |
| "#009E73", # bluish green | |
| "#E69F00", # orange | |
| "#56B4E9", # sky blue | |
| "#CC79A7", # reddish purple | |
| "#000000", # black | |
| "#F0E442", # yellow (sparingly: low contrast on white) | |
| ) | |
| # Published M2 Pro break-even acceptance rates (paper #32, Table 5, Bielik et | |
| # al.): ranges across drafter/dataset combinations. The compact 40-77% band | |
| # (k=2..4) is what figures/README use for the comparison. | |
| M2PRO_ABE = {2: (38.0, 52.8), 4: (77.7, 90.1)} # exact Table 5 ranges | |
| M2PRO_BAND = (0.40, 0.77) # compact k=2..4 range used in F3 | |
| # --------------------------------------------------------------------------- # | |
| # Small helpers | |
| # --------------------------------------------------------------------------- # | |
| def _fmt(x: float | None, nd: int = 2, suffix: str = "") -> str: | |
| """Format a float for markdown tables, or an em dash when None/NaN.""" | |
| if x is None or (isinstance(x, float) and not np.isfinite(x)): | |
| return "\u2014" | |
| return f"{x:.{nd}f}{suffix}" | |
| def md_table(headers: list[str], rows: list[list[str]]) -> str: | |
| """Render a GitHub-flavored markdown table.""" | |
| lines = ["| " + " | ".join(headers) + " |"] | |
| lines.append("|" + "|".join(" " + "-" * (len(h) + 2) + " " for h in headers) + "|") | |
| for row in rows: | |
| cells = [str(c) for c in row] | |
| if len(cells) != len(headers): | |
| raise ValueError(f"row has {len(cells)} cells, header has {len(headers)}: {cells}") | |
| lines.append("| " + " | ".join(cells) + " |") | |
| return "\n".join(lines) | |
| # --------------------------------------------------------------------------- # | |
| # Config normalization | |
| # --------------------------------------------------------------------------- # | |
| class RunInfo: | |
| run_name: str | |
| family: str # final | curves | ksweep | baseline | |
| target: str # qwen-q4 | gemma-q8 ... | |
| family_name: str # qwen | gemma | |
| quant: str # q4 | q5 | q8 | |
| drafter: str # normalized drafter id (e.g. "dspark-p0", "dflash-f16", "solo") | |
| k: int # spec_draft_n_max | |
| ctx: int | |
| prompt_path: str | |
| prompt_set: str | |
| n_tokens: int | |
| temperature: float | |
| top_k: int | |
| top_p: float | |
| seed: int | |
| p_min: float | None | |
| draft_path: str | None | |
| _QUANT_RE = re.compile(r"Q([458])[_K]") | |
| def normalize_target(model_path: str) -> str: | |
| """Map a model path to '{family}-q{quant}' (e.g. 'qwen-q4').""" | |
| family = "qwen" if "Qwen" in model_path else "gemma" if "gemma" in model_path else "?" | |
| m = _QUANT_RE.search(model_path) | |
| quant = m.group(1) if m else "?" | |
| if family == "?" or quant == "?": | |
| raise ValueError(f"cannot normalize model path: {model_path}") | |
| return f"{family}-q{quant}" | |
| def normalize_drafter(spec_type: str | None, p_min: float | None, draft_path: str | None) -> str: | |
| """Map spec_type (+ p_min and draft quant) to the normalized drafter id.""" | |
| st = (spec_type or "none").lower() | |
| if st in ("none", "", "vanilla"): | |
| return "solo" | |
| if st == "draft-simple": | |
| return "vanilla17b" | |
| if st == "draft-eagle3": | |
| return "eagle3" | |
| if st == "draft-mtp": | |
| return "mtp" | |
| if st == "draft-dspark": | |
| p = 0.0 if p_min is None else p_min | |
| return f"dspark-p{int(round(p * 10))}" | |
| if st == "draft-dflash": | |
| path = draft_path or "" | |
| if "Q4_K_M" in path or "-Q4" in path: | |
| return "dflash-q4" | |
| if "Q8_0" in path or "-Q8" in path: | |
| return "dflash-q8" | |
| return "dflash-f16" # F16 (or unnamed block gguf) is the default DFlash | |
| raise ValueError(f"unknown spec_type: {spec_type}") | |
| def prompt_set_id(prompt_path: str) -> str: | |
| """Return a content identity for a prompt file, not just its path spelling.""" | |
| raw = Path(prompt_path) | |
| candidates = [raw] | |
| if not raw.is_absolute(): | |
| candidates.append(Path.cwd() / raw) | |
| for candidate in candidates: | |
| try: | |
| if candidate.is_file(): | |
| return "sha256:" + hashlib.sha256(candidate.read_bytes()).hexdigest() | |
| except OSError: | |
| continue | |
| # Missing files remain distinguishable and cannot accidentally match a | |
| # different file that happens to have the same basename. | |
| return "path:" + os.path.normpath(prompt_path) | |
| def parse_run_info(run_name: str, cfg: dict[str, Any]) -> RunInfo: | |
| """Normalize one run directory into a RunInfo.""" | |
| prefix = run_name.split("-")[0] | |
| family = "baseline" if prefix in ("baseline", "ctx2048") else prefix | |
| if family not in ("final", "curves", "ksweep", "baseline"): | |
| raise ValueError(f"unexpected run prefix: {run_name}") | |
| target = normalize_target(cfg["model"]) | |
| fam_name, quant = target.split("-") | |
| p_min = cfg.get("spec_draft_p_min") | |
| drafter = normalize_drafter(cfg.get("spec_type"), p_min, cfg.get("draft")) | |
| k = int(cfg.get("spec_draft_n_max") or 0) | |
| sampling = cfg.get("sampling") or {} | |
| prompt_path = str(cfg.get("prompts") or "") | |
| return RunInfo( | |
| run_name=run_name, | |
| family=family, | |
| target=target, | |
| family_name=fam_name, | |
| quant=quant, | |
| drafter=drafter, | |
| k=k, | |
| ctx=int(cfg.get("ctx") or 0), | |
| prompt_path=prompt_path, | |
| prompt_set=prompt_set_id(prompt_path), | |
| n_tokens=int(cfg.get("n_tokens") or 0), | |
| temperature=float(sampling.get("temperature", cfg.get("temperature", 0.0))), | |
| top_k=int(sampling.get("top_k", cfg.get("top_k", 0))), | |
| top_p=float(sampling.get("top_p", cfg.get("top_p", 0.0))), | |
| seed=int(sampling.get("seed", cfg.get("seed", 0))), | |
| p_min=p_min, | |
| draft_path=cfg.get("draft"), | |
| ) | |
| # --------------------------------------------------------------------------- # | |
| # Data loading | |
| # --------------------------------------------------------------------------- # | |
| def load_records(run_dir: str) -> list[dict[str, Any]]: | |
| """Load results.jsonl (tolerates stray non-UTF8 bytes / broken lines).""" | |
| records: list[dict[str, Any]] = [] | |
| corrupt = 0 | |
| with open(os.path.join(run_dir, "results.jsonl"), errors="replace") as fh: | |
| for line in fh: | |
| line = line.strip() | |
| if not line: | |
| continue | |
| try: | |
| records.append(json.loads(line)) | |
| except json.JSONDecodeError: | |
| corrupt += 1 | |
| if corrupt: | |
| logger.warning(" %s: skipped %d unparseable lines", os.path.basename(run_dir), corrupt) | |
| return records | |
| def count_error_records(run_dir: str) -> int: | |
| """Count persisted error records, including errors from resumed attempts.""" | |
| path = os.path.join(run_dir, "errors.jsonl") | |
| if not os.path.isfile(path): | |
| return 0 | |
| with open(path, errors="replace") as fh: | |
| return sum(1 for line in fh if line.strip()) | |
| def is_sentinel(rec: dict[str, Any]) -> bool: | |
| """True when a record carries the spurious 1e6 tok/s timing marker.""" | |
| tps = rec.get("tok_per_s") | |
| pred = rec.get("predicted_ms") | |
| return (isinstance(tps, (int, float)) and tps >= SENTINEL_TPS) or ( | |
| isinstance(pred, (int, float)) and pred <= SENTINEL_MS | |
| ) | |
| class CleanStats: | |
| kept: list[dict[str, Any]] | |
| excluded: int | |
| total: int | |
| def clean_records(records: list[dict[str, Any]]) -> CleanStats: | |
| kept = [r for r in records if not is_sentinel(r)] | |
| return CleanStats(kept=kept, excluded=len(records) - len(kept), total=len(records)) | |
| _ACCEPT_RE = re.compile(r"draft acceptance = ([\d.]+)") | |
| _POS_RE = re.compile(r"acc per pos = \(([\d.,\s]+)\)") | |
| def parse_server_log(run_dir: str) -> tuple[list[float], list[list[float]]]: | |
| """Extract per-request acceptance + per-position vectors from server.log.""" | |
| alphas: list[float] = [] | |
| positions: list[list[float]] = [] | |
| with open(os.path.join(run_dir, "server.log"), errors="replace") as fh: | |
| for line in fh: | |
| m = _ACCEPT_RE.search(line) | |
| if m: | |
| alphas.append(float(m.group(1))) | |
| m2 = _POS_RE.search(line) | |
| if m2: | |
| positions.append([float(x) for x in m2.group(1).split(",")]) | |
| if len(alphas) != len(positions): | |
| raise RuntimeError( | |
| f"{os.path.basename(run_dir)}: {len(alphas)} acceptance lines vs " | |
| f"{len(positions)} per-position lines" | |
| ) | |
| return alphas, positions | |
| class LogMatch: | |
| run_name: str | |
| k: int | |
| matched: int | |
| log_lines: int | |
| max_alpha_diff: float | |
| position_by_id: dict[str, list[float]] | |
| domain_by_id: dict[str, str] | |
| def match_log_to_records( | |
| runs_dir_path: str, run_name: str, records: list[dict[str, Any]], k: int | |
| ) -> LogMatch: | |
| """Assign the i-th log line to the i-th record with non-None alpha. | |
| Gemini timing-sentinel records have ``alpha = None`` and no log line, so | |
| they are skipped while consuming log entries (verified: max |diff| <= 5e-5). | |
| """ | |
| alphas, positions = parse_server_log(os.path.join(runs_dir_path, run_name)) | |
| rec_with_alpha = [r for r in records if r.get("alpha") is not None] | |
| if len(alphas) != len(rec_with_alpha): | |
| raise RuntimeError( | |
| f"{run_name}: {len(alphas)} log lines vs {len(rec_with_alpha)} records with alpha" | |
| ) | |
| max_diff = 0.0 | |
| position_by_id: dict[str, list[float]] = {} | |
| domain_by_id: dict[str, str] = {} | |
| for rec, log_alpha, pos in zip(rec_with_alpha, alphas, positions, strict=True): | |
| max_diff = max(max_diff, abs(rec["alpha"] - log_alpha)) | |
| if len(pos) != k: | |
| logger.warning(" %s: pos line has %d values, expected k=%d", run_name, len(pos), k) | |
| position_by_id[rec["id"]] = pos | |
| domain_by_id[rec["id"]] = rec["domain"] | |
| return LogMatch( | |
| run_name=run_name, | |
| k=k, | |
| matched=len(rec_with_alpha), | |
| log_lines=len(alphas), | |
| max_alpha_diff=max_diff, | |
| position_by_id=position_by_id, | |
| domain_by_id=domain_by_id, | |
| ) | |
| # --------------------------------------------------------------------------- # | |
| # Aggregation helpers | |
| # --------------------------------------------------------------------------- # | |
| def basic_stats(values: list[float]) -> dict[str, float]: | |
| arr = np.asarray(values, dtype=float) | |
| return { | |
| "mean": float(arr.mean()), | |
| "median": float(np.median(arr)), | |
| "p95": float(np.percentile(arr, 95)), | |
| } | |
| def alpha_stats(records: list[dict[str, Any]]) -> dict[str, float] | None: | |
| alphas = [r["alpha"] for r in records if r.get("alpha") is not None] | |
| if not alphas: | |
| return None | |
| arr = np.asarray(alphas) | |
| return {"mean": float(arr.mean()), "median": float(np.median(arr))} | |
| def tau_mean(records: list[dict[str, Any]]) -> float | None: | |
| taus = [r["tau"] for r in records if r.get("tau") is not None] | |
| return float(np.mean(taus)) if taus else None | |
| def bootstrap_ci(values: list[float], seed: int = BOOTSTRAP_RNG_SEED) -> tuple[float, float]: | |
| """Percentile-bootstrap 95% CI of the mean.""" | |
| arr = np.asarray(values, dtype=float) | |
| rng = np.random.default_rng(seed) | |
| samples = np.empty(BOOTSTRAP_ITERS) | |
| for i in range(BOOTSTRAP_ITERS): | |
| samples[i] = rng.choice(arr, size=len(arr), replace=True).mean() | |
| lo, hi = np.percentile(samples, [2.5, 97.5]) | |
| return float(lo), float(hi) | |
| def ols_breakeven(alphas: list[float], tps: list[float], baseline: float) -> dict[str, float]: | |
| """OLS TPS = a + beta*alpha and the break-even acceptance rate. | |
| ``alpha_be = (baseline - a) / beta``; CI95 from the OLS covariance of | |
| (a, beta) via the delta method (paper #32 uses the same approach). | |
| Returns raw coefficients even when ``alpha_be`` is outside [0, 1] | |
| (value > 1 => unreachable, < 0 => always above baseline). | |
| """ | |
| x = np.asarray(alphas, dtype=float) | |
| y = np.asarray(tps, dtype=float) | |
| if len(x) < 3: | |
| return { | |
| "n": len(x), | |
| "beta": np.nan, | |
| "intercept": np.nan, | |
| "r2": np.nan, | |
| "alpha_be": np.nan, | |
| "ci95": np.nan, | |
| } | |
| (beta, a), cov = np.polyfit(x, y, 1, cov=True) | |
| yhat = a + beta * x | |
| ss_res = float(np.sum((y - yhat) ** 2)) | |
| ss_tot = float(np.sum((y - y.mean()) ** 2)) | |
| r2 = 1.0 - ss_res / ss_tot if ss_tot > 0 else np.nan | |
| alpha_be = (baseline - a) / beta | |
| # delta method: d(alpha_be)/d(a) = -1/beta ; d(alpha_be)/d(beta) = -(base-a)/beta^2 | |
| da, db = -1.0 / beta, -(baseline - a) / (beta * beta) | |
| var = da**2 * cov[0, 0] + db**2 * cov[1, 1] + 2.0 * da * db * cov[0, 1] | |
| ci = 1.96 * float(np.sqrt(max(var, 0.0))) | |
| return { | |
| "n": len(x), | |
| "beta": float(beta), | |
| "intercept": float(a), | |
| "r2": float(r2), | |
| "alpha_be": float(alpha_be), | |
| "ci95": float(ci), | |
| } | |
| def speedup_vs( | |
| recs: list[dict[str, Any]], solo_tps_by_id: dict[str, float] | |
| ) -> dict[str, float] | None: | |
| """Per-prompt matched speedup: mean/median of ratios + ratio of means.""" | |
| ratios = [r["tok_per_s"] / solo_tps_by_id[r["id"]] for r in recs if r["id"] in solo_tps_by_id] | |
| if not ratios: | |
| return None | |
| arr = np.asarray(ratios) | |
| matched = [r for r in recs if r["id"] in solo_tps_by_id] | |
| return { | |
| "mean": float(arr.mean()), | |
| "median": float(np.median(arr)), | |
| "agg": float( | |
| np.mean([r["tok_per_s"] for r in matched]) / np.mean(list(solo_tps_by_id.values())) | |
| ), | |
| "n_match": len(ratios), | |
| } | |
| # --------------------------------------------------------------------------- # | |
| # final runs | |
| # --------------------------------------------------------------------------- # | |
| def build_final_stats(runs_dir_path: str) -> dict[str, Any]: | |
| """Aggregate the 26 final runs: per-config and per-domain statistics.""" | |
| infos: dict[str, RunInfo] = {} | |
| for name in sorted(os.listdir(runs_dir_path)): | |
| cfg_path = os.path.join(runs_dir_path, name, "config.json") | |
| if os.path.isfile(cfg_path): | |
| with open(cfg_path) as fh: | |
| infos[name] = parse_run_info(name, json.load(fh)) | |
| final_infos = {n: i for n, i in infos.items() if i.family == "final"} | |
| if len(final_infos) != 26: | |
| logger.warning("expected 26 final runs, found %d", len(final_infos)) | |
| # Load + clean every final run, log exclusions. | |
| records_by_run: dict[str, list[dict[str, Any]]] = {} | |
| exclusions: dict[str, dict[str, int]] = {} | |
| for name in sorted(final_infos): | |
| cs = clean_records(load_records(os.path.join(runs_dir_path, name))) | |
| records_by_run[name] = cs.kept | |
| exclusions[name] = {"total": cs.total, "kept": len(cs.kept), "excluded": cs.excluded} | |
| if cs.excluded: | |
| logger.info( | |
| "excluded %d sentinel records in %s (%d kept)", cs.excluded, name, len(cs.kept) | |
| ) | |
| # Baseline solos: per target -> per domain mean tok/s + per-prompt map. | |
| solo_by_target: dict[str, dict[str, Any]] = {} | |
| for name, info in final_infos.items(): | |
| if info.drafter != "solo": | |
| continue | |
| recs = records_by_run[name] | |
| by_domain = {d: [r for r in recs if r["domain"] == d] for d in DOMAINS} | |
| solo_by_target[info.target] = { | |
| "by_domain": by_domain, | |
| "tps_mean": { | |
| d: float(np.mean([r["tok_per_s"] for r in by_domain[d]])) for d in DOMAINS | |
| }, | |
| "tps_mean_all": float(np.mean([r["tok_per_s"] for r in recs])), | |
| "tps_by_id": {r["id"]: r["tok_per_s"] for r in recs}, | |
| } | |
| if len(solo_by_target) != 6: | |
| logger.warning( | |
| "expected 6 solo baselines (qwen/gemma x q4/q5/q8), found %d", len(solo_by_target) | |
| ) | |
| summary: dict[str, Any] = {} | |
| for name, info in sorted(final_infos.items()): | |
| recs = records_by_run[name] | |
| entry: dict[str, Any] = { | |
| "run": name, | |
| "target": info.target, | |
| "family": info.family_name, | |
| "quant": info.quant, | |
| "drafter": info.drafter, | |
| "k": info.k, | |
| "ctx": info.ctx, | |
| "n": len(recs), | |
| "tok_per_s": basic_stats([r["tok_per_s"] for r in recs]), | |
| "alpha": alpha_stats(recs), | |
| "tau_mean": tau_mean(recs), | |
| "ttft_ms": basic_stats([r["prompt_ms"] for r in recs]), | |
| } | |
| vram_path = os.path.join(runs_dir_path, name, "vram.json") | |
| if os.path.isfile(vram_path): | |
| with open(vram_path) as fh: | |
| vram = json.load(fh) | |
| entry["vram_max_mib"] = vram.get("max_gpu_mib") | |
| entry["power_max_w"] = vram.get("max_power_w") | |
| metrics_path = os.path.join(runs_dir_path, name, "metrics.json") | |
| if os.path.isfile(metrics_path): | |
| with open(metrics_path) as fh: | |
| metrics = json.load(fh) | |
| entry["duration_s"] = metrics.get("duration_s") | |
| entry["errors"] = count_error_records(os.path.join(runs_dir_path, name)) | |
| solo = solo_by_target.get(info.target) | |
| per_domain: dict[str, Any] = {} | |
| for d in DOMAINS: | |
| dr = [r for r in recs if r["domain"] == d] | |
| cell: dict[str, Any] = { | |
| "n": len(dr), | |
| "tok_per_s": basic_stats([r["tok_per_s"] for r in dr]), | |
| "alpha": alpha_stats(dr), | |
| "tau_mean": tau_mean(dr), | |
| } | |
| if solo is not None and info.drafter != "solo": | |
| cell["speedup"] = speedup_vs(dr, solo["tps_by_id"]) | |
| per_domain[d] = cell | |
| entry["per_domain"] = per_domain | |
| if solo is not None and info.drafter != "solo": | |
| entry["speedup"] = speedup_vs(recs, solo["tps_by_id"]) | |
| summary[name] = entry | |
| return {"summary": summary, "exclusions": exclusions, "solo_by_target": solo_by_target} | |
| def baseline_compatibility_key(info: RunInfo) -> tuple[str, int, str, int, float, int, float, int]: | |
| """Identity used to pair a ksweep run with a target-only baseline.""" | |
| return ( | |
| info.target, | |
| info.ctx, | |
| info.prompt_set, | |
| info.n_tokens, | |
| info.temperature, | |
| info.top_k, | |
| info.top_p, | |
| info.seed, | |
| ) | |
| def build_baseline_stats( | |
| runs_dir_path: str, | |
| ) -> tuple[ | |
| dict[tuple[str, int, str, int, float, int, float, int], dict[str, Any]], list[dict[str, Any]] | |
| ]: | |
| """Load explicitly named contextual target-only baseline runs. | |
| Baselines are kept separate from the six final-run solo baselines. A | |
| baseline can be used for break-even only when its full protocol identity | |
| matches the ksweep run (target, context, prompt-file content, and sampling | |
| settings). Duplicate identities are rejected rather than selected | |
| implicitly. | |
| """ | |
| by_key: dict[tuple[str, int, str, int, float, int, float, int], dict[str, Any]] = {} | |
| summaries: list[dict[str, Any]] = [] | |
| for name in sorted(os.listdir(runs_dir_path)): | |
| # Only the explicit baseline-* namespace is analytical input. A | |
| # previous interrupted controller left a ctx2048-* scratch run; do | |
| # not let an accidental rerun compete with the canonical baseline. | |
| if not name.startswith("baseline-"): | |
| continue | |
| cfg_path = os.path.join(runs_dir_path, name, "config.json") | |
| if not os.path.isfile(cfg_path): | |
| continue | |
| with open(cfg_path) as fh: | |
| info = parse_run_info(name, json.load(fh)) | |
| if info.family != "baseline" or info.drafter != "solo": | |
| continue | |
| cs = clean_records(load_records(os.path.join(runs_dir_path, name))) | |
| recs = cs.kept | |
| by_domain = {d: [r for r in recs if r.get("domain") == d] for d in DOMAINS} | |
| tps = [r["tok_per_s"] for r in recs if r.get("tok_per_s") is not None] | |
| tps_by_id = {r["id"]: r["tok_per_s"] for r in recs if r.get("tok_per_s") is not None} | |
| metrics_path = os.path.join(runs_dir_path, name, "metrics.json") | |
| errors = count_error_records(os.path.join(runs_dir_path, name)) | |
| if errors == 0 and os.path.isfile(metrics_path): | |
| with open(metrics_path) as fh: | |
| errors = json.load(fh).get("errors", 0) | |
| key = baseline_compatibility_key(info) | |
| if key in by_key: | |
| previous = by_key[key]["run"] | |
| raise RuntimeError( | |
| f"duplicate compatible baselines for {info.target} ctx={info.ctx}: " | |
| f"{previous} and {name}" | |
| ) | |
| entry: dict[str, Any] = { | |
| "run": name, | |
| "target": info.target, | |
| "ctx": info.ctx, | |
| "prompt_path": info.prompt_path, | |
| "prompt_set": info.prompt_set, | |
| "n_tokens": info.n_tokens, | |
| "sampling": { | |
| "temperature": info.temperature, | |
| "top_k": info.top_k, | |
| "top_p": info.top_p, | |
| "seed": info.seed, | |
| }, | |
| "n": len(recs), | |
| "excluded": cs.excluded, | |
| "errors": errors, | |
| "tok_per_s": basic_stats(tps) if tps else None, | |
| "tok_per_s_by_id": tps_by_id, | |
| "tok_per_s_by_domain": { | |
| d: basic_stats([r["tok_per_s"] for r in by_domain[d]]) if by_domain[d] else None | |
| for d in DOMAINS | |
| }, | |
| } | |
| by_key[key] = entry | |
| summaries.append({k: v for k, v in entry.items() if k != "tok_per_s_by_id"}) | |
| return by_key, summaries | |
| # --------------------------------------------------------------------------- # | |
| # curves / ksweep: per-position acceptance + ksweep break-even | |
| # --------------------------------------------------------------------------- # | |
| def build_position_data(runs_dir_path: str) -> dict[str, Any]: | |
| """Parse curves/ksweep logs and aggregate per (run, domain, position).""" | |
| out: dict[str, Any] = {} | |
| for name in sorted(os.listdir(runs_dir_path)): | |
| cfg_path = os.path.join(runs_dir_path, name, "config.json") | |
| if not os.path.isfile(cfg_path): | |
| continue | |
| with open(cfg_path) as fh: | |
| info = parse_run_info(name, json.load(fh)) | |
| if info.family not in ("curves", "ksweep"): | |
| continue | |
| cs = clean_records(load_records(os.path.join(runs_dir_path, name))) | |
| if cs.excluded: | |
| logger.info("excluded %d sentinel records in %s", cs.excluded, name) | |
| match = match_log_to_records(runs_dir_path, name, cs.kept, info.k) | |
| # Group per (domain, position). | |
| per_domain: dict[str, list[dict[str, Any]]] = {d: [] for d in DOMAINS} | |
| for rid, pos in match.position_by_id.items(): | |
| dom = match.domain_by_id[rid] | |
| for p, val in enumerate(pos, start=1): | |
| per_domain[dom].append({"pos": p, "val": val}) | |
| dom_out: dict[str, Any] = {} | |
| for d in DOMAINS: | |
| pos_stats: list[dict[str, Any]] = [] | |
| for p in range(1, info.k + 1): | |
| vals = [e["val"] for e in per_domain[d] if e["pos"] == p] | |
| if not vals: | |
| continue | |
| lo, hi = bootstrap_ci(vals) | |
| pos_stats.append( | |
| { | |
| "position": p, | |
| "n": len(vals), | |
| "alpha_mean": float(np.mean(vals)), | |
| "alpha_median": float(np.median(vals)), | |
| "ci95_low": lo, | |
| "ci95_high": hi, | |
| } | |
| ) | |
| dom_out[d] = pos_stats | |
| out[name] = { | |
| "family": info.family, | |
| "target": info.target, | |
| "drafter": info.drafter, | |
| "k": info.k, | |
| "n_excluded": cs.excluded, | |
| "n_matched": match.matched, | |
| "log_lines": match.log_lines, | |
| "max_alpha_diff": match.max_alpha_diff, | |
| "per_domain": dom_out, | |
| } | |
| return out | |
| def build_breakeven( | |
| runs_dir_path: str, | |
| baseline_by_key: dict[tuple[str, int, str, int, float, int, float, int], dict[str, Any]], | |
| ) -> dict[str, Any]: | |
| """OLS alpha_be per ksweep run using a protocol-compatible baseline.""" | |
| out: dict[str, Any] = {} | |
| for name in sorted(os.listdir(runs_dir_path)): | |
| cfg_path = os.path.join(runs_dir_path, name, "config.json") | |
| if not os.path.isfile(cfg_path): | |
| continue | |
| with open(cfg_path) as fh: | |
| info = parse_run_info(name, json.load(fh)) | |
| if info.family != "ksweep": | |
| continue | |
| baseline = baseline_by_key.get(baseline_compatibility_key(info)) | |
| if baseline is None: | |
| logger.warning( | |
| "skipping %s: no target-only baseline matches target=%s ctx=%d " | |
| "prompt_set=%s sampling=%s", | |
| name, | |
| info.target, | |
| info.ctx, | |
| info.prompt_set, | |
| baseline_compatibility_key(info)[3:], | |
| ) | |
| continue | |
| recs = [ | |
| r | |
| for r in clean_records(load_records(os.path.join(runs_dir_path, name))).kept | |
| if r.get("alpha") is not None | |
| ] | |
| matched = [r for r in recs if r.get("id") in baseline["tok_per_s_by_id"]] | |
| if len(matched) < 3: | |
| logger.warning("skipping %s: fewer than 3 baseline-matched observations", name) | |
| continue | |
| baseline_by_id = baseline["tok_per_s_by_id"] | |
| baseline_all = float(np.mean([baseline_by_id[r["id"]] for r in matched])) | |
| pooled = ols_breakeven( | |
| [r["alpha"] for r in matched], [r["tok_per_s"] for r in matched], baseline_all | |
| ) | |
| per_domain: dict[str, Any] = {} | |
| baseline_by_domain: dict[str, float] = {} | |
| baseline_n_by_domain: dict[str, int] = {} | |
| for d in DOMAINS: | |
| dr = [r for r in matched if r["domain"] == d] | |
| baseline_values = [baseline_by_id[r["id"]] for r in dr] | |
| baseline_by_domain[d] = float(np.mean(baseline_values)) if baseline_values else np.nan | |
| baseline_n_by_domain[d] = len(baseline_values) | |
| per_domain[d] = ols_breakeven( | |
| [r["alpha"] for r in dr], | |
| [r["tok_per_s"] for r in dr], | |
| baseline_by_domain[d], | |
| ) | |
| out[name] = { | |
| "target": info.target, | |
| "drafter": info.drafter, | |
| "k": info.k, | |
| "ctx": info.ctx, | |
| "prompt_set": info.prompt_set, | |
| "baseline_run": baseline["run"], | |
| "baseline_ctx": baseline["ctx"], | |
| "baseline_prompt_set": baseline["prompt_set"], | |
| "baseline_n": baseline["n"], | |
| "baseline_n_matched": len(matched), | |
| "baseline_n_by_domain": baseline_n_by_domain, | |
| "baseline_excluded": baseline["excluded"], | |
| "baseline_all": baseline_all, | |
| "baseline_by_domain": baseline_by_domain, | |
| "pooled": pooled, | |
| "per_domain": per_domain, | |
| } | |
| return out | |
| # --------------------------------------------------------------------------- # | |
| # Markdown tables | |
| # --------------------------------------------------------------------------- # | |
| def _config_rows( | |
| final_summary: dict[str, Any], family_name: str | |
| ) -> list[tuple[str, dict[str, Any]]]: | |
| """(run-name, entry) pairs for one family, sorted by quant then drafter.""" | |
| rows = [(n, e) for n, e in final_summary.items() if e["family"] == family_name] | |
| order = {"q4": 0, "q5": 1, "q8": 2} | |
| drafter_order = { | |
| "solo": 0, | |
| "vanilla17b": 1, | |
| "eagle3": 2, | |
| "dflash-f16": 3, | |
| "dflash-q4": 4, | |
| "dflash-q8": 5, | |
| "dspark-p0": 6, | |
| "dspark-p2": 7, | |
| "dspark-p4": 8, | |
| "dspark-p6": 9, | |
| "mtp": 10, | |
| } | |
| rows.sort(key=lambda r: (order.get(r[1]["quant"], 9), drafter_order.get(r[1]["drafter"], 99))) | |
| return rows | |
| def _short_name(run_name: str) -> str: | |
| return run_name.split("-", 2)[-1] | |
| def write_t2_speedup(final_summary: dict[str, Any], out_dir: Path) -> None: | |
| """tok/s and speedup vs solo, per config x domain (Qwen and Gemma blocks).""" | |
| headers = [ | |
| "Config", | |
| "n", | |
| "Math tok/s", | |
| "Math x", | |
| "Code tok/s", | |
| "Code x", | |
| "Chat tok/s", | |
| "Chat x", | |
| "All tok/s", | |
| "All x", | |
| "alpha all", | |
| ] | |
| sections: list[str] = [] | |
| for fam, title in ( | |
| ("qwen", "### Qwen3-8B (baseline: same-quant solo)"), | |
| ("gemma", "### Gemma 4 12B (baseline: same-quant solo)"), | |
| ): | |
| lines = [title, ""] | |
| rows: list[list[str]] = [] | |
| for name, entry in _config_rows(final_summary, fam): | |
| cells = [_short_name(name), str(entry["n"])] | |
| for d in (*DOMAINS, "all"): | |
| dom = entry if d == "all" else entry["per_domain"][d] | |
| cells.append(_fmt(dom["tok_per_s"]["mean"], 1)) | |
| sup = entry.get("speedup") if d == "all" else dom.get("speedup") | |
| sup = sup.get("mean") if isinstance(sup, dict) else None | |
| cells.append(_fmt(sup, 2, "x") if sup is not None else "\u2014") | |
| alpha = entry.get("alpha") | |
| cells.append(_fmt(alpha["mean"], 3) if alpha else "\u2014") | |
| rows.append(cells) | |
| lines.append(md_table(headers, rows)) | |
| lines.append("") | |
| sections.append("\n".join(lines)) | |
| (out_dir / "t2_speedup.md").write_text("\n".join(sections), encoding="utf-8") | |
| def write_t3_alpha_tau(final_summary: dict[str, Any], out_dir: Path) -> None: | |
| """alpha and tau per config x domain (drafter configs only).""" | |
| headers = [ | |
| "Config", | |
| "Math alpha", | |
| "Math tau", | |
| "Code alpha", | |
| "Code tau", | |
| "Chat alpha", | |
| "Chat tau", | |
| "All alpha", | |
| "All tau", | |
| ] | |
| sections: list[str] = [] | |
| for fam, title in (("qwen", "### Qwen3-8B"), ("gemma", "### Gemma 4 12B")): | |
| lines = [title, ""] | |
| rows: list[list[str]] = [] | |
| for name, entry in _config_rows(final_summary, fam): | |
| if entry["drafter"] == "solo": | |
| continue | |
| cells = [_short_name(name)] | |
| for d in (*DOMAINS, "all"): | |
| dom = entry if d == "all" else entry["per_domain"][d] | |
| a = dom.get("alpha") | |
| cells.append(_fmt(a["mean"], 3) if a else "\u2014") | |
| cells.append(_fmt(dom.get("tau_mean"), 0)) | |
| rows.append(cells) | |
| lines.append(md_table(headers, rows)) | |
| lines.append("") | |
| sections.append("\n".join(lines)) | |
| (out_dir / "t3_alpha_tau.md").write_text("\n".join(sections), encoding="utf-8") | |
| def _summarize_all(entry: dict[str, Any]) -> tuple[str, str, str, str]: | |
| tps = _fmt(entry["tok_per_s"]["mean"], 1) | |
| sup = entry.get("speedup") | |
| sup_s = _fmt(sup["mean"], 2, "x") if sup else "\u2014" | |
| a = entry.get("alpha") | |
| a_s = _fmt(a["mean"], 3) if a else "\u2014" | |
| return tps, sup_s, a_s, _fmt(entry.get("tau_mean"), 0) | |
| def write_t4_quantization(final_summary: dict[str, Any], out_dir: Path) -> None: | |
| """Target-quant x drafter interaction, plus gemma-q4 draft-quant effect.""" | |
| headers = ["Target quant", "Drafter", "n", "tok/s", "x vs solo", "alpha", "tau"] | |
| lines: list[str] = [] | |
| for fam, title in ( | |
| ("qwen", "### Qwen3-8B — target quant (q4/q5/q8) x drafter"), | |
| ("gemma", "### Gemma 4 12B — target quant x drafter"), | |
| ): | |
| lines.append(title) | |
| lines.append("") | |
| rows: list[list[str]] = [] | |
| for _, entry in _config_rows(final_summary, fam): | |
| if entry["drafter"] == "solo": | |
| continue | |
| tps, sup, a, tau = _summarize_all(entry) | |
| rows.append( | |
| [ | |
| entry["quant"], | |
| DRAFTER_LABELS[entry["drafter"]], | |
| str(entry["n"]), | |
| tps, | |
| sup, | |
| a, | |
| tau, | |
| ] | |
| ) | |
| lines.append(md_table(headers, rows)) | |
| lines.append("") | |
| lines.append("### Gemma-4 Q4 — draft quantization effect (DFlash drafts, final runs)") | |
| lines.append("") | |
| rows: list[list[str]] = [] | |
| for _, entry in _config_rows(final_summary, "gemma"): | |
| if ( | |
| entry["quant"] != "q4" | |
| or entry["drafter"] == "solo" | |
| or not entry["drafter"].startswith("dflash") | |
| ): | |
| continue | |
| tps, sup, a, tau = _summarize_all(entry) | |
| rows.append( | |
| [ | |
| DRAFTER_LABELS[entry["drafter"]], | |
| str(entry["n"]), | |
| tps, | |
| sup, | |
| a, | |
| tau, | |
| _fmt(entry.get("vram_max_mib"), 0), | |
| ] | |
| ) | |
| lines.append( | |
| md_table( | |
| ["Drafter (draft quant)", "n", "tok/s", "x vs solo", "alpha", "tau", "max VRAM (MiB)"], | |
| rows, | |
| ) | |
| ) | |
| lines.append("") | |
| lines.append("_F16 = 1.47 GB draft, Q4_K_M = 0.44 GB, Q8_0 = 0.79 GB (model-hashes.json)._") | |
| lines.append("") | |
| (out_dir / "t4_quantization.md").write_text("\n".join(lines), encoding="utf-8") | |
| def write_t5_hardware(final_summary: dict[str, Any], out_dir: Path) -> None: | |
| """TTFT, VRAM, power, duration per config.""" | |
| headers = [ | |
| "Config", | |
| "TTFT mean (ms)", | |
| "TTFT median (ms)", | |
| "TTFT p95 (ms)", | |
| "max VRAM (MiB)", | |
| "max power (W)", | |
| "duration (s)", | |
| ] | |
| rows: list[list[str]] = [] | |
| for name, entry in sorted(final_summary.items()): | |
| tt = entry["ttft_ms"] | |
| rows.append( | |
| [ | |
| name, | |
| _fmt(tt["mean"], 1), | |
| _fmt(tt["median"], 1), | |
| _fmt(tt["p95"], 1), | |
| _fmt(entry.get("vram_max_mib"), 0), | |
| _fmt(entry.get("power_max_w"), 1), | |
| _fmt(entry.get("duration_s"), 1), | |
| ] | |
| ) | |
| (out_dir / "t5_hardware.md").write_text( | |
| md_table(headers, rows) | |
| + "\n\n_All timings from clean (non-sentinel) records; VRAM/power from" | |
| " vram.json; duration from metrics.json._\n", | |
| encoding="utf-8", | |
| ) | |
| def write_t6_breakeven(breakeven: dict[str, Any], out_dir: Path) -> None: | |
| """ksweep OLS break-even: pooled + per-domain, with baseline provenance.""" | |
| lines: list[str] = [] | |
| headers = [ | |
| "Config", | |
| "k", | |
| "n", | |
| "baseline run", | |
| "baseline ctx", | |
| "baseline n/match", | |
| "baseline (tok/s)", | |
| "beta (slope)", | |
| "alpha_be", | |
| "CI95", | |
| "R2", | |
| ] | |
| rows: list[list[str]] = [] | |
| for name in sorted(breakeven): | |
| be = breakeven[name] | |
| pooled = be["pooled"] | |
| rows.append( | |
| [ | |
| f"{be['target']}-{be['drafter']}", | |
| str(be["k"]), | |
| str(pooled["n"]), | |
| be["baseline_run"], | |
| str(be["baseline_ctx"]), | |
| f"{be['baseline_n']}/{be['baseline_n_matched']}", | |
| _fmt(be["baseline_all"], 1), | |
| _fmt(pooled["beta"], 2), | |
| _fmt(pooled["alpha_be"], 3), | |
| _fmt(pooled["ci95"], 3), | |
| _fmt(pooled["r2"], 3), | |
| ] | |
| ) | |
| lines.append("### Pooled (all domains)") | |
| lines.append("") | |
| lines.append(md_table(headers, rows)) | |
| lines.append("") | |
| lines.append("### Per domain") | |
| lines.append("") | |
| headers_d = [ | |
| "Config", | |
| "k", | |
| "domain", | |
| "n", | |
| "baseline ctx", | |
| "baseline n", | |
| "baseline", | |
| "beta", | |
| "alpha_be", | |
| "CI95", | |
| "R2", | |
| ] | |
| rows_d: list[list[str]] = [] | |
| for name in sorted(breakeven): | |
| be = breakeven[name] | |
| for d in DOMAINS: | |
| b = be["per_domain"][d] | |
| rows_d.append( | |
| [ | |
| f"{be['target']}-{be['drafter']}", | |
| str(be["k"]), | |
| d, | |
| str(b["n"]), | |
| str(be["baseline_ctx"]), | |
| str(be["baseline_n_by_domain"][d]), | |
| _fmt(be["baseline_by_domain"][d], 1), | |
| _fmt(b["beta"], 2), | |
| _fmt(b["alpha_be"], 3), | |
| _fmt(b["ci95"], 3), | |
| _fmt(b["r2"], 3), | |
| ] | |
| ) | |
| lines.append(md_table(headers_d, rows_d)) | |
| lines.append("") | |
| lines.append("### Comparison with M2 Pro (paper #32, Bielik et al., cross-family)") | |
| lines.append("") | |
| lines.append( | |
| "Paper #32 fits `TPS = a + b*alpha` by OLS and defines `alpha_be = (TPS_base - a) / b` " | |
| "(its `b` is our `beta`, the 'recovery rate'). Published values (Table 5, ranging " | |
| "across drafters and datasets) are **k=2: 38.0-52.8%** and **k=4: 77.7-90.1%**. " | |
| "The compact 40-77% range summarizes k=2..4. Our ksweep starts " | |
| "at k=5, so the comparison is directional: the k=10 RTX range is below the " | |
| "reported k=2 band, while the upper ends at k=5 and k=7 slightly overlap its " | |
| "lower edge." | |
| ) | |
| lines.append("") | |
| rows_c: list[list[str]] = [] | |
| for k, (lo, hi) in sorted(M2PRO_ABE.items()): | |
| rows_c.append([f"M2 Pro k={k}", f"{lo:.1f}-{hi:.1f}%"]) | |
| ours: dict[int, list[float]] = {} | |
| for be in breakeven.values(): | |
| ours.setdefault(be["k"], []).append(be["pooled"]["alpha_be"]) | |
| for k in sorted(ours): | |
| vals = [v for v in ours[k] if np.isfinite(v)] | |
| if vals: | |
| rows_c.append( | |
| [ | |
| f"Ours k={k} (n={len(vals)} configs)", | |
| f"{100 * min(vals):.1f}-{100 * max(vals):.1f}%", | |
| ] | |
| ) | |
| lines.append(md_table(["Reference", "alpha_be range"], rows_c)) | |
| lines.append("") | |
| lines.append( | |
| "_alpha_be > 1.00 = no OLS-reachable break-even; CI95 by delta" | |
| " method over the OLS covariance._" | |
| ) | |
| lines.append("") | |
| (out_dir / "t6_breakeven.md").write_text("\n".join(lines), encoding="utf-8") | |
| # --------------------------------------------------------------------------- # | |
| # Figures | |
| # --------------------------------------------------------------------------- # | |
| def _style_figure() -> None: | |
| """One visual theme for all figures, authored at final print size. | |
| Figures are rendered at ~6.5 in width (the LaTeX textwidth) so pandoc's | |
| \\pandocbounded downscale is ~1x and text prints at true 7.5-8 pt instead | |
| of the 3-7 pt the old large canvases degraded to. | |
| """ | |
| plt.rcParams.update( | |
| { | |
| # Serif matches the LaTeX body; STIX covers the math glyphs. | |
| "font.family": "serif", | |
| "font.serif": ["STIXGeneral", "DejaVu Serif", "Times New Roman"], | |
| "mathtext.fontset": "stix", | |
| "font.size": 8.0, | |
| "axes.titlesize": 8.0, | |
| "axes.labelsize": 8.0, | |
| "legend.fontsize": 7.5, | |
| "xtick.labelsize": 7.5, | |
| "ytick.labelsize": 7.5, | |
| "figure.dpi": 150, | |
| "savefig.dpi": 600, # arXiv wants >= 300; line art is cheap at 600 | |
| "axes.grid": True, | |
| "grid.alpha": 0.25, | |
| "grid.linewidth": 0.4, | |
| "axes.spines.top": False, | |
| "axes.spines.right": False, | |
| "axes.linewidth": 0.6, | |
| "lines.linewidth": 1.5, | |
| "lines.markersize": 4.5, | |
| "legend.frameon": False, | |
| "savefig.bbox": "tight", | |
| "savefig.pad_inches": 0.02, | |
| } | |
| ) | |
| DOMAIN_COLORS = {"math": OKABE_ITO[0], "code": OKABE_ITO[1], "chat": OKABE_ITO[2]} | |
| DOMAIN_MARKERS = {"math": "o", "code": "s", "chat": "^"} | |
| DOMAIN_LINESTYLES = {"math": "-", "code": "--", "chat": ":"} | |
| # Drafter families in F2: DSpark p-min variants share one color + legend entry | |
| # (they are the same drafter at different confidence thresholds). | |
| F2_FAMILY_COLORS = { | |
| "vanilla17b": OKABE_ITO[6], # black | |
| "eagle3": OKABE_ITO[0], # blue | |
| "dflash-f16": OKABE_ITO[1], # vermillion | |
| "dflash-q4": OKABE_ITO[3], # orange | |
| "dflash-q8": OKABE_ITO[5], # reddish purple | |
| "dspark": OKABE_ITO[4], # sky blue | |
| "mtp": OKABE_ITO[2], # bluish green | |
| } | |
| F2_FAMILY_LABELS = { | |
| "vanilla17b": "Vanilla-1.7B", | |
| "eagle3": "EAGLE-3", | |
| "dflash-f16": "DFlash-F16", | |
| "dflash-q4": "DFlash-Q4", | |
| "dflash-q8": "DFlash-Q8", | |
| "dspark": "DSpark (p=0.0\u20130.6)", | |
| "mtp": "MTP", | |
| } | |
| def _drafter_family(drafter: str) -> str: | |
| """Group the DSpark p-min sweep under one family key.""" | |
| return "dspark" if drafter.startswith("dspark") else drafter | |
| def _panel_letter(idx: int) -> str: | |
| """Lowercase subfigure letter for panel index (a, b, c, ...).""" | |
| return chr(ord("a") + idx) | |
| def make_f1(position_data: dict[str, Any], figs_dir: str) -> None: | |
| """Per-position acceptance curves (curves-*, k=3, ctx 2048). | |
| Split per target at final print size: F1a = Qwen3-8B-Q4 (7 drafter | |
| panels), F1b = Gemma 4 12B-Q4 (4 drafter panels). Both share the domain | |
| legend and the Okabe-Ito theme from ``_style_figure``. | |
| """ | |
| _style_figure() | |
| # Resolve (target, drafter) -> actual run key (qwen-dflash normalizes to | |
| # "dflash-f16" but its run folder is named "curves-qwen-q4-dflash"). | |
| run_by_key: dict[tuple[str, str], str] = {} | |
| for run, data in position_data.items(): | |
| if data["family"] == "curves": | |
| run_by_key[(data["target"], data["drafter"])] = run | |
| groups = [ | |
| ( | |
| "F1a_acceptance_qwen.png", | |
| "Qwen3-8B Q4", | |
| "qwen-q4", | |
| [ | |
| "vanilla17b", | |
| "eagle3", | |
| "dflash-f16", | |
| "dspark-p0", | |
| "dspark-p2", | |
| "dspark-p4", | |
| "dspark-p6", | |
| ], | |
| (4, 2), | |
| (6.5, 3.0), | |
| ), | |
| ( | |
| "F1b_acceptance_gemma.png", | |
| "Gemma 4 12B Q4", | |
| "gemma-q4", | |
| ["mtp", "dflash-f16", "dflash-q4", "dflash-q8"], | |
| (4, 1), | |
| (6.5, 2.1), | |
| ), | |
| ] | |
| domain_handles = [ | |
| Line2D( | |
| [], | |
| [], | |
| color=DOMAIN_COLORS[d], | |
| lw=1.5, | |
| linestyle=cast(Any, DOMAIN_LINESTYLES[d]), | |
| marker=DOMAIN_MARKERS[d], | |
| label=d, | |
| ) | |
| for d in DOMAINS | |
| ] | |
| for fname, target_title, target, drafter_ids, (ncols, nrows), figsize in groups: | |
| panels: list[tuple[str, str, str]] = [] | |
| for dr in drafter_ids: | |
| key = run_by_key.get((target, dr)) | |
| if key is None: | |
| logger.warning("F1: missing run for %s-%s", target, dr) | |
| continue | |
| panels.append((target, dr, key)) | |
| if not panels: | |
| continue | |
| fig, axes = plt.subplots(nrows, ncols, figsize=figsize, sharey=True) | |
| axes = np.atleast_2d(axes) # single-row grids come back 1-D | |
| for idx, (_, dr, run) in enumerate(panels): | |
| row, col = divmod(idx, ncols) | |
| _f1_panel(axes[row, col], position_data[run], DRAFTER_LABELS[dr], _panel_letter(idx)) | |
| for ax in axes.flat[len(panels) :]: | |
| ax.set_axis_off() | |
| for ax in axes[:, 0]: | |
| ax.set_ylabel(r"token acceptance $\alpha$") | |
| fig.supxlabel("draft position (k = 3, ctx 2,048)", y=0.02) | |
| title_y = 0.995 | |
| fig.suptitle(target_title, y=title_y, fontsize=9, weight="bold") | |
| # Domain legend in the reserved bottom margin (NOT over the panels). | |
| fig.legend( | |
| handles=domain_handles, | |
| loc="lower center", | |
| ncol=3, | |
| frameon=False, | |
| bbox_to_anchor=(0.5, 0.10), | |
| ) | |
| # The one-row Gemma figure needs extra headroom for its suptitle. | |
| bottom = 0.28 if nrows > 1 else 0.36 | |
| top = 0.88 if nrows > 1 else 0.80 | |
| fig.subplots_adjust(left=0.06, right=0.99, top=top, bottom=bottom, hspace=0.55) | |
| path = os.path.join(figs_dir, fname) | |
| fig.savefig(path, bbox_inches="tight", pad_inches=0.05) | |
| plt.close(fig) | |
| logger.info("wrote %s", path) | |
| def _f1_panel(ax: Any, data: dict[str, Any], label: str, letter: str) -> None: | |
| """One per-position acceptance panel: domain lines + 95% CI bands.""" | |
| xs = np.arange(1, data["k"] + 1) | |
| for d in DOMAINS: | |
| stats = {s["position"]: s for s in data["per_domain"][d]} | |
| ys = np.array([stats[p]["alpha_mean"] for p in range(1, data["k"] + 1)]) | |
| lo = np.array([stats[p]["ci95_low"] for p in range(1, data["k"] + 1)]) | |
| hi = np.array([stats[p]["ci95_high"] for p in range(1, data["k"] + 1)]) | |
| ax.plot( | |
| xs, | |
| ys, | |
| color=DOMAIN_COLORS[d], | |
| marker=DOMAIN_MARKERS[d], | |
| linestyle=DOMAIN_LINESTYLES[d], | |
| ms=4.5, | |
| lw=1.5, | |
| label=d, | |
| ) | |
| ax.fill_between(xs, lo, hi, color=DOMAIN_COLORS[d], alpha=0.18) | |
| ax.set_title(f"({letter}) {label}") | |
| ax.set_xticks(xs) | |
| ax.set_ylim(0, 1.0) | |
| def make_f2(final_summary: dict[str, Any], figs_dir: str) -> None: | |
| """Speedup vs mean acceptance, per config x domain (26 final runs).""" | |
| _style_figure() | |
| fig, axes = plt.subplots(1, 2, figsize=(6.5, 3.4), sharex=True, sharey=True) | |
| # DSpark p-min variants share one color + one legend entry (family). | |
| families_present = { | |
| fam | |
| for e in final_summary.values() | |
| if e["drafter"] != "solo" | |
| for fam in (_drafter_family(e["drafter"]),) | |
| } | |
| drafter_color = { | |
| fam: F2_FAMILY_COLORS[fam] for fam in F2_FAMILY_COLORS if fam in families_present | |
| } | |
| for entry in final_summary.values(): | |
| if entry["drafter"] == "solo": | |
| continue | |
| ax = axes[0 if entry["family"] == "qwen" else 1] | |
| for d in DOMAINS: | |
| dom = entry["per_domain"][d] | |
| alpha = dom.get("alpha") | |
| sup = dom.get("speedup") | |
| if alpha is not None and sup is not None: | |
| ax.scatter( | |
| alpha["mean"], | |
| sup["mean"], | |
| color=drafter_color[_drafter_family(entry["drafter"])], | |
| marker=DOMAIN_MARKERS[d], | |
| s=46, | |
| alpha=0.88, | |
| edgecolors="white", | |
| linewidths=0.5, | |
| ) | |
| for ax, (letter, title) in zip(axes, (("a", "Qwen3-8B"), ("b", "Gemma 4 12B")), strict=True): | |
| ax.axhspan(0.4, 1.0, color="#b00020", alpha=0.045, zorder=0) | |
| ax.axhline(1.0, color="black", lw=1.1, ls="--", zorder=1) | |
| ax.text( | |
| 0.98, | |
| 0.04, | |
| r"loss ($<1\times$)", | |
| transform=ax.transAxes, | |
| color="#8f001b", | |
| ha="right", | |
| ) | |
| ax.set_title(f"({letter}) {title}") | |
| ax.set_xlim(-0.02, 1.02) | |
| ax.set_ylim(0.4, 2.7) | |
| key_labels = { | |
| "final-qwen-q4-dflash": ("DFlash Q4", (10, 12)), | |
| "final-qwen-q4-dspark-p0": ("DSpark p=0", (7, 8)), | |
| "final-gemma-q4-dflash-f16": ("DFlash F16", (8, -14)), | |
| } | |
| for name, (label, offset) in key_labels.items(): | |
| entry = final_summary.get(name) | |
| if entry is None: | |
| continue | |
| a = entry.get("alpha") | |
| s = entry.get("speedup") | |
| if a is not None and s is not None: | |
| ax = axes[0 if entry["family"] == "qwen" else 1] | |
| ax.annotate( | |
| label, | |
| (a["mean"], s["mean"]), | |
| textcoords="offset points", | |
| xytext=offset, | |
| color="#333333", | |
| arrowprops={"arrowstyle": "-", "color": "#666666", "lw": 0.7}, | |
| ) | |
| handles_d = [ | |
| Line2D([], [], marker=DOMAIN_MARKERS[d], color="#444444", ls="", label=d) for d in DOMAINS | |
| ] | |
| handles_c = [ | |
| Line2D( | |
| [], | |
| [], | |
| color=drafter_color[fam], | |
| ls="", | |
| marker="o", | |
| label=F2_FAMILY_LABELS[fam], | |
| ) | |
| for fam in drafter_color | |
| ] | |
| fig.legend( | |
| handles=handles_c, | |
| loc="center left", | |
| bbox_to_anchor=(0.72, 0.58), | |
| title="drafter (colour)", | |
| frameon=False, | |
| ) | |
| # Keep both legends in the right-side margin so neither covers the x-label. | |
| fig.legend( | |
| handles=handles_d, | |
| loc="upper left", | |
| bbox_to_anchor=(0.72, 0.30), | |
| title="domain (marker)", | |
| ncol=3, | |
| frameon=False, | |
| ) | |
| fig.supxlabel(r"mean acceptance $\alpha$ (per config $\times$ domain, final runs)", y=0.05) | |
| fig.supylabel("mean speedup vs. same-quant solo (per-prompt ratio)", x=0.03) | |
| fig.subplots_adjust(left=0.10, right=0.70, bottom=0.14, top=0.90, wspace=0.22) | |
| path = os.path.join(figs_dir, "F2_speedup.png") | |
| fig.savefig(path, bbox_inches="tight", pad_inches=0.05) | |
| plt.close(fig) | |
| logger.info("wrote %s", path) | |
| def make_f3(breakeven: dict[str, Any], figs_dir: str) -> None: | |
| """alpha_be vs k from ksweep, with M2 Pro reference band.""" | |
| _style_figure() | |
| fig, ax = plt.subplots(figsize=(6.5, 3.5)) | |
| configs: dict[str, list[tuple[int, float, float]]] = {} | |
| for be in breakeven.values(): | |
| key = f"{be['target']}-{be['drafter']}" | |
| configs.setdefault(key, []).append( | |
| (be["k"], be["pooled"]["alpha_be"], be["pooled"]["ci95"]) | |
| ) | |
| ymax = 1.0 | |
| for pts in configs.values(): | |
| for _, abe, _ in pts: | |
| if np.isfinite(abe): | |
| ymax = max(ymax, abe * 1.15) | |
| ymax = min(ymax, 2.0) | |
| # Unique per-config labels for the legend (avoid collapsing drafters | |
| # that appear with different targets/quants, e.g. qwen-q4 vs qwen-q8 EAGLE-3). | |
| f3_labels: dict[str, str] = { | |
| "qwen-q4-dflash-f16": "Qwen3-8B Q4 DFlash", | |
| "qwen-q4-dspark-p0": "Qwen3-8B Q4 DSpark p0", | |
| "qwen-q4-eagle3": "Qwen3-8B Q4 EAGLE-3", | |
| "qwen-q4-vanilla17b": "Qwen3-8B Q4 Vanilla-1.7B", | |
| "qwen-q8-eagle3": "Qwen3-8B Q8 EAGLE-3", | |
| "gemma-q4-dflash-f16": "Gemma-4-12B Q4 DFlash-F16", | |
| "gemma-q4-mtp": "Gemma-4-12B Q4 MTP", | |
| } | |
| for i, (key, pts) in enumerate(sorted(configs.items())): | |
| pts.sort() | |
| ks = [p[0] for p in pts] | |
| abe = np.array([p[1] for p in pts]) | |
| ci = np.array([p[2] for p in pts]) | |
| color = OKABE_ITO[i % len(OKABE_ITO)] | |
| # Clamp for display; annotate values above the visible range. | |
| show = np.minimum(abe, ymax * 0.98) | |
| ax.plot(ks, show, marker="o", color=color, lw=1.5, label=f3_labels.get(key, key)) | |
| ax.errorbar( | |
| ks, | |
| show, | |
| yerr=np.minimum(ci, ymax * 0.95), | |
| fmt="none", | |
| color=color, | |
| alpha=0.5, | |
| capsize=2, | |
| ) | |
| for k, a in zip(ks, abe, strict=True): | |
| if np.isfinite(a) and a > ymax * 0.98: | |
| ax.annotate(f"{a:.2f}", (k, ymax * 0.96), fontsize=7, color=color, ha="center") | |
| ax.axhspan( | |
| M2PRO_BAND[0], | |
| M2PRO_BAND[1], | |
| facecolor="gold", | |
| edgecolor="#B8860B", | |
| alpha=0.16, | |
| hatch="//", | |
| label="M2 Pro reference band (40\u201377%)", | |
| ) | |
| ax.axhline(1.0, color="black", lw=0.9, ls="--", label=r"$\alpha_{be} = 1.0$ (no break-even)") | |
| ax.set_xticks([5, 7, 10]) | |
| ax.set_xlim(4.7, 10.3) | |
| ax.set_xlabel(r"draft length $k$") | |
| ax.set_ylabel(r"break-even acceptance $\alpha_{be}$ (OLS)") | |
| ax.set_ylim(0.0, ymax) | |
| legend_handles, legend_labels = ax.get_legend_handles_labels() | |
| fig.legend( | |
| legend_handles, | |
| legend_labels, | |
| loc="lower center", | |
| bbox_to_anchor=(0.5, 0.025), | |
| frameon=False, | |
| ncol=3, | |
| fontsize=7.5, | |
| ) | |
| fig.subplots_adjust(left=0.10, right=0.97, bottom=0.37, top=0.92) | |
| path = os.path.join(figs_dir, "F3_breakeven.png") | |
| fig.savefig(path, bbox_inches="tight", pad_inches=0.05) | |
| plt.close(fig) | |
| logger.info("wrote %s", path) | |
| def make_f4(position_data: dict[str, Any], figs_dir: str) -> None: | |
| """Per-position alpha across k=5/7/10 for two representative qwen-q4 configs.""" | |
| _style_figure() | |
| pairs = [("ksweep-qwen-q4-eagle3", "EAGLE-3"), ("ksweep-qwen-q4-dspark-p0", "DSpark p=0.0")] | |
| fig, axes = plt.subplots(1, 2, figsize=(6.5, 2.8), sharey=True) | |
| k_styles = (("-", "o"), ("--", "s"), (":", "^")) | |
| for ax, (letter, (base, title)) in zip(axes, (("a", pairs[0]), ("b", pairs[1])), strict=True): | |
| for j, k in enumerate((5, 7, 10)): | |
| run = f"{base}-k{k}" | |
| if run not in position_data: | |
| continue | |
| data = position_data[run] | |
| xs: list[int] = [] | |
| ys: list[float] = [] | |
| for p in range(1, data["k"] + 1): | |
| pooled = [ | |
| s["alpha_mean"] | |
| for d in DOMAINS | |
| for s in data["per_domain"][d] | |
| if s["position"] == p | |
| ] | |
| if pooled: | |
| xs.append(p) | |
| ys.append(np.mean(pooled)) | |
| linestyle, marker = k_styles[j] | |
| ax.plot( | |
| xs, | |
| ys, | |
| color=OKABE_ITO[j], | |
| linestyle=cast(Any, linestyle), | |
| marker=marker, | |
| ms=4.5, | |
| lw=1.5, | |
| label=f"k={k}", | |
| ) | |
| ax.set_title(f"({letter}) {title}") | |
| ax.set_xlabel("draft position") | |
| ax.set_ylim(0, 0.9) | |
| axes[0].set_ylabel(r"$\alpha$ by position (pooled domains)") | |
| k_handles = [ | |
| Line2D( | |
| [], | |
| [], | |
| color=OKABE_ITO[j], | |
| linestyle=cast(Any, k_styles[j][0]), | |
| marker=k_styles[j][1], | |
| ms=4.5, | |
| lw=1.5, | |
| label=f"k={k}", | |
| ) | |
| for j, k in enumerate((5, 7, 10)) | |
| ] | |
| fig.suptitle("Acceptance curve vs. $k$ \u2014 Qwen3-8B Q4 (ctx 2,048)", y=0.99, fontsize=9) | |
| fig.legend( | |
| handles=k_handles, | |
| loc="lower center", | |
| ncol=3, | |
| frameon=False, | |
| bbox_to_anchor=(0.5, 0.05), | |
| ) | |
| fig.subplots_adjust(left=0.09, right=0.97, top=0.82, bottom=0.30, wspace=0.18) | |
| path = os.path.join(figs_dir, "F4_alpha_position_ksweep.png") | |
| fig.savefig(path, bbox_inches="tight", pad_inches=0.05) | |
| plt.close(fig) | |
| logger.info("wrote %s", path) | |
| # --------------------------------------------------------------------------- # | |
| # Checks + README | |
| # --------------------------------------------------------------------------- # | |
| def run_checks(final_summary: dict[str, Any]) -> list[dict[str, Any]]: | |
| """Compare computed values against session-08 expectations (>5% => warn).""" | |
| items: list[dict[str, Any]] = [] | |
| for chk in EXPECTED_CHECKS: | |
| entry = final_summary.get(chk["run"]) | |
| if entry is None: | |
| items.append({**chk, "calc": None, "pass": False, "note": "run missing"}) | |
| continue | |
| metric = chk["metric"] | |
| if metric == "tps_mean": | |
| calc = entry["tok_per_s"]["mean"] | |
| ok = calc is not None and abs(calc - chk["exp"]) / chk["exp"] <= 0.05 | |
| items.append({**chk, "calc": calc, "pass": ok, "note": ""}) | |
| elif metric == "alpha": | |
| a = entry.get("alpha") | |
| calc = a["mean"] if a else None | |
| ok = calc is not None and abs(calc - chk["exp"]) / chk["exp"] <= 0.05 | |
| items.append({**chk, "calc": calc, "pass": ok, "note": ""}) | |
| elif metric == "speedup": | |
| s = entry.get("speedup") | |
| if not s: | |
| items.append({**chk, "calc": None, "pass": False, "note": "no speedup"}) | |
| continue | |
| rel_mean = abs(s["mean"] - chk["exp"]) / chk["exp"] | |
| rel_med = abs(s["median"] - chk["exp"]) / chk["exp"] | |
| ok = min(rel_mean, rel_med) <= 0.05 | |
| items.append( | |
| { | |
| **chk, | |
| "calc": s["mean"], | |
| "calc_median": s["median"], | |
| "pass": ok, | |
| "note": f"mean {rel_mean:.1%} off, median {rel_med:.1%} off", | |
| } | |
| ) | |
| else: | |
| items.append({**chk, "calc": None, "pass": False, "note": f"unknown metric {metric}"}) | |
| return items | |
| def write_checks_md(items: list[dict[str, Any]], out_dir: Path) -> None: | |
| lines = [ | |
| "# CHECKS - expected values (session 08) vs computed", | |
| "", | |
| "Tolerance: >5% relative deviation -> warning. Speedups are compared against", | |
| "the mean and median of the per-prompt ratio; either one passing is enough.", | |
| "", | |
| ] | |
| headers = ["Check", "Expected", "Computed", "Pass?", "Note"] | |
| rows: list[list[str]] = [] | |
| for it in items: | |
| if it["metric"] == "speedup": | |
| calc = f"{_fmt(it['calc'], 2)} (med {_fmt(it.get('calc_median'), 2)})" | |
| else: | |
| calc = _fmt(it["calc"], 3) if it["calc"] is not None else "\u2014" | |
| rows.append( | |
| [ | |
| it["label"], | |
| _fmt(it["exp"], 3), | |
| calc, | |
| "OK" if it["pass"] else "WARN", | |
| it.get("note", ""), | |
| ] | |
| ) | |
| lines.append(md_table(headers, rows)) | |
| lines.append("") | |
| n_fail = sum(1 for it in items if not it["pass"]) | |
| lines.append(f"**{len(items) - n_fail}/{len(items)} checks pass.**") | |
| lines.append("") | |
| (out_dir / "checks.md").write_text("\n".join(lines), encoding="utf-8") | |
| def write_readme(out_dir: Path) -> None: | |
| text = """# Final analysis - `experiments/analysis/` | |
| Generated by `scripts/analyze_final.py` (Python 3.12, stdlib only + numpy + matplotlib). | |
| ## Methodology | |
| ### Sources | |
| - `experiments/runs/<run>/results.jsonl` - one OK record per line (schema in the | |
| inventory). `config.json` normalizes the run: family (`final/curves/ksweep`), | |
| target (`{qwen|gemma}-{q4|q5|q8}` from the model path), drafter from | |
| `spec_type` + `p_min` (dspark) + draft quantization (dflash-f16/q4/q8). | |
| - `server.log` of `curves-*`/`ksweep-*` runs - `draft acceptance` and | |
| `acc per pos` lines (only place with per-position curves; the `final-*` runs | |
| rotated logs and their original segment has no acceptance lines). | |
| ### Exclusions (timing sentinels) | |
| llama-server reports `tok_per_s = 1,000,000` / `predicted_ms = 0` in 271 records | |
| of Gemma runs (timing quirk, not real speedup). **Rule applied to all | |
| computations**: every record with `tok_per_s >= 1e5` or `predicted_ms <= 0` is | |
| excluded (also carrying `alpha/tau/draft_n = None`). Qwen has no sentinels. The | |
| Gemma `metrics.json` files are contaminated by this (mean 12,291 tok/s) and are | |
| **not** used for tok/s; only `duration_s` and `errors` are used. Counts per run | |
| in `summary.json` -> `meta.exclusions_per_run` (q4: 18/run, q5: 25/run, q8: 20/run | |
| in final; 3-4/run in curves/ksweep gemma). | |
| ### Log <-> record mapping (curves/ksweep) | |
| The order of the `draft acceptance` lines in `server.log` == order of records | |
| with non-None `alpha` in `results.jsonl` (verified in 41/41 runs, max |diff| = | |
| 5e-5, rounding only). Sentinel Gemma records (alpha=None) have no line and are | |
| skipped. This assigns a domain to each line. No mismatches in any run. | |
| ### Speedup vs solo | |
| Per-prompt ratio `tps_draft / tps_solo` matched by `id` against the `solo` run | |
| of the same family+quant; the mean and median of the ratio are reported, plus | |
| the ratio of means (`agg`). The 6 solos (qwen q4/q5/q8, gemma q4/q5/q8) exist | |
| and are the baseline. | |
| ### Break-even alpha_be (paper #32, Bielik et al.) | |
| Per `ksweep-*` run and domain (and pooled): OLS of `TPS = a + beta*alpha` over | |
| the run records (per-prompt alpha and tok/s, clean records with alpha). | |
| `alpha_be = (TPS_base - a) / beta` where `TPS_base` = mean tok/s of the | |
| corresponding `solo` (same family+quant+domain, final run). CI95 by delta method | |
| over the OLS covariance of (a, beta). alpha_be > 1 = no reachable break-even; | |
| < 0 = always above the baseline. Paper #32 publishes k=2: 38.0-52.8% and | |
| k=4: 77.7-90.1% (Table 5); the compact 40-77% range covers | |
| k=2..4. | |
| ## Files | |
| - `summary.json` - consolidated numbers (per config x domain and breakeven block). | |
| - `tables/t2_speedup.md` ... `tables/t6_breakeven.md` - markdown tables for the paper. | |
| - `curves/acc_by_pos_<run>.csv` - per-position acceptance by domain (41 runs). | |
| - `checks.md` - result of the checks vs session 08. | |
| - `manuscript/figures/F1..F4_*.png` - figures (300 dpi). | |
| ## Execution | |
| ``` | |
| .venv/bin/python scripts/analyze_final.py \\ | |
| --runs-dir experiments/runs --out-dir experiments/analysis \\ | |
| --figs-dir manuscript/figures | |
| ``` | |
| """ | |
| (out_dir / "README.md").write_text(text, encoding="utf-8") | |
| # --------------------------------------------------------------------------- # | |
| # main | |
| # --------------------------------------------------------------------------- # | |
| def main() -> int: | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument( | |
| "--runs-dir", default="experiments/runs", help="directory with the run folders" | |
| ) | |
| parser.add_argument( | |
| "--out-dir", default="experiments/analysis", help="output directory (summary/tables/curves)" | |
| ) | |
| parser.add_argument( | |
| "--figs-dir", default="manuscript/figures", help="output directory for figures" | |
| ) | |
| args = parser.parse_args() | |
| logging.basicConfig(level=logging.INFO, format="%(levelname)s: %(message)s") | |
| runs_dir_path = args.runs_dir | |
| out_dir = Path(args.out_dir) | |
| figs_dir = Path(args.figs_dir) | |
| curves_dir = out_dir / "curves" | |
| tables_dir = out_dir / "tables" | |
| for d in (out_dir, tables_dir, curves_dir, figs_dir): | |
| d.mkdir(parents=True, exist_ok=True) | |
| logger.info("loading final runs ...") | |
| final = build_final_stats(runs_dir_path) | |
| final_summary, exclusions = ( | |
| final["summary"], | |
| final["exclusions"], | |
| ) | |
| logger.info("loading contextual target-only baselines ...") | |
| baseline_by_key, baseline_summary = build_baseline_stats(runs_dir_path) | |
| logger.info("parsing curves/ksweep logs ...") | |
| position_data = build_position_data(runs_dir_path) | |
| breakeven = build_breakeven(runs_dir_path, baseline_by_key) | |
| # Per-position CSVs. | |
| logger.info("writing per-position CSVs ...") | |
| csv_meta: dict[str, Any] = {} | |
| for run, data in sorted(position_data.items()): | |
| path = curves_dir / f"acc_by_pos_{run}.csv" | |
| with path.open("w", encoding="utf-8") as fh: | |
| fh.write("config,domain,position,n,alpha_mean,alpha_median,ci95_low,ci95_high\n") | |
| for d in DOMAINS: | |
| for s in data["per_domain"][d]: | |
| fh.write( | |
| f"{data['target']}-{data['drafter']},{d},{s['position']},{s['n']}," | |
| f"{s['alpha_mean']:.4f},{s['alpha_median']:.4f},{s['ci95_low']:.4f},{s['ci95_high']:.4f}\n" | |
| ) | |
| csv_meta[run] = { | |
| "csv": str(path.relative_to(out_dir)), | |
| "k": data["k"], | |
| "n_matched": data["n_matched"], | |
| } | |
| # summary.json | |
| logger.info("writing summary.json ...") | |
| total_excluded = sum(e["excluded"] for e in exclusions.values()) + sum( | |
| d["n_excluded"] for d in position_data.values() | |
| ) | |
| ksweep_payload: dict[str, Any] = {} | |
| for run, be in breakeven.items(): | |
| pos = position_data.get(run, {}) | |
| ksweep_payload[run] = { | |
| "family": pos.get("family", "ksweep"), | |
| "target": be["target"], | |
| "drafter": be["drafter"], | |
| "k": be["k"], | |
| "n_excluded": pos.get("n_excluded", 0), | |
| "n_matched": pos.get("n_matched"), | |
| "log_lines": pos.get("log_lines"), | |
| "max_alpha_diff": pos.get("max_alpha_diff"), | |
| "ctx": be["ctx"], | |
| "prompt_set": be["prompt_set"], | |
| "baseline_run": be["baseline_run"], | |
| "baseline_ctx": be["baseline_ctx"], | |
| "baseline_prompt_set": be["baseline_prompt_set"], | |
| "baseline_n": be["baseline_n"], | |
| "baseline_n_matched": be["baseline_n_matched"], | |
| "baseline_n_by_domain": be["baseline_n_by_domain"], | |
| "baseline_excluded": be["baseline_excluded"], | |
| "baseline_all": be["baseline_all"], | |
| "baseline_by_domain": be["baseline_by_domain"], | |
| "per_domain_positions": pos.get("per_domain"), | |
| "breakeven_pooled": be["pooled"], | |
| "breakeven_per_domain": be["per_domain"], | |
| } | |
| meta: dict[str, Any] = { | |
| "runs_total": len(exclusions) + len(position_data), | |
| "final_runs": len(exclusions), | |
| "curves_runs": sum(1 for r in position_data if r.startswith("curves-")), | |
| "ksweep_runs": sum(1 for r in position_data if r.startswith("ksweep-")), | |
| "baseline_runs": len(baseline_summary), | |
| "baseline_records_excluded_total": sum(b["excluded"] for b in baseline_summary), | |
| "records_excluded_total": total_excluded, | |
| "exclusions_per_run": exclusions, | |
| "log_mapping": { | |
| run: { | |
| "log_lines": d["log_lines"], | |
| "n_matched": d["n_matched"], | |
| "max_alpha_diff": d["max_alpha_diff"], | |
| } | |
| for run, d in position_data.items() | |
| }, | |
| "csv_files": csv_meta, | |
| "method": ( | |
| "sentinel exclusion: tok_per_s >= 1e5 or predicted_ms <= 0;" | |
| " log<->record by order of appearance; OLS alpha_be per paper #32" | |
| ), | |
| } | |
| summary_payload = { | |
| "meta": meta, | |
| "final": final_summary, | |
| "baselines": baseline_summary, | |
| "curves": position_data, | |
| "ksweep": ksweep_payload, | |
| } | |
| with (out_dir / "summary.json").open("w", encoding="utf-8") as fh: | |
| json.dump(summary_payload, fh, indent=2, default=str) | |
| # Tables. | |
| logger.info("writing markdown tables ...") | |
| write_t2_speedup(final_summary, tables_dir) | |
| write_t3_alpha_tau(final_summary, tables_dir) | |
| write_t4_quantization(final_summary, tables_dir) | |
| write_t5_hardware(final_summary, tables_dir) | |
| write_t6_breakeven(breakeven, tables_dir) | |
| # Figures. | |
| logger.info("building figures ...") | |
| make_f1(position_data, str(figs_dir)) | |
| make_f2(final_summary, str(figs_dir)) | |
| make_f3(breakeven, str(figs_dir)) | |
| make_f4(position_data, str(figs_dir)) | |
| # Checks. | |
| logger.info("running CHECKS ...") | |
| items = run_checks(final_summary) | |
| write_checks_md(items, out_dir) | |
| print("\n==================== CHECKS ====================") | |
| n_fail = 0 | |
| for it in items: | |
| ok = it["pass"] | |
| n_fail += 0 if ok else 1 | |
| mark = "PASS" if ok else "WARN" | |
| calc = it.get("calc") | |
| extra = it.get("note", "") | |
| if it["metric"] == "speedup": | |
| med = it.get("calc_median") | |
| print( | |
| f" [{mark}] {it['label']:<30} exp={it['exp']:<5} calc={calc:.3f}" | |
| f" (med={med:.3f}) {extra}" | |
| ) | |
| elif calc is None: | |
| print(f" [{mark}] {it['label']:<30} exp={it['exp']:<5} calc=--- {extra}") | |
| else: | |
| print(f" [{mark}] {it['label']:<30} exp={it['exp']:<5} calc={calc:.3f} {extra}") | |
| print("=================================================") | |
| print(f"{len(items) - n_fail}/{len(items)} checks pass; {n_fail} deviate >5%.") | |
| write_readme(out_dir) | |
| logger.info("done. outputs in %s and %s", out_dir, figs_dir) | |
| return 0 | |
| if __name__ == "__main__": | |
| sys.exit(main()) | |