Team Ai
Modelpublic

reyden009/speculative-decoding-lab

sourceHugging Facemitupdated 2mo agoView on Hugging Face
8likes
bench_accept.py926 linesDownload Raw Back to scripts
1#!/usr/bin/env python32"""Acceptance (alpha/tau) benchmark via llama-server (requires gaps G1/G3/G4).3 4Why a server and not llama-cli:5- llama-cli does NOT expose speculative-decoding statistics; llama-server DOES:6  every /completion response returns timings.draft_n and timings.draft_n_accepted7  (server-task.cpp result_timings::to_json), plus predicted_per_second.8- Additionally, with --verbose it writes per-request log lines:9    "draft acceptance = 0.xxxxx (N accepted / M generated), mean len = X.XX"10  and at TRC level "acc per pos = (r1, r2, ...)" (per-position curve, optional:11  parse <out>/server.log after the run).12 13Protocol parity with F1/F2/F3 (llama-cli, default n_max = 3 in common.h):14- --spec-draft-n-max 3, -c 2048, -t 8, seed 42, n_predict 256.15- Explicit sampling in the body: temperature 0.0 (pure greedy argmax; ignores16  top_k/top_p), top_k 40, top_p 0.95 — the final run corrects F1/F2/F3 which17  ran at T=0.8 (llama-cli defaults). Non-thinking: --reasoning off in the18  server command (if raw mode rejects it, the documented escape is in --help).19- The numbers from this script are ONLY comparable among themselves20  (server-consistent methodology: each config starts its own server). DO NOT21  mix with the F1/F2/F3 tok/s (fresh process per prompt vs persistent server).22 23Anti-rerun design (identical to bench_spec.py):24- <out>/results.jsonl (one line per OK prompt), <out>/errors.jsonl (failures,25  retried with --resume), exclusive lock (exit 3 if another runner writes the26  same --out), vram.json (max_gpu_mib + max_power_w sampled during the run),27  config.json (args + sampling + reproducibility: llama.cpp commit, version,28  sha256/size of the GGUFs), metrics.json (total and per-domain aggregates)29  and results.csv (without text column) at the end.30- exit 0 = clean; 2 = failures; 3 = lock busy; 1 = server did not start or31  died mid-run (abort after 2 consecutive connection failures).32 33Usage (example):34    # 1) create a per-domain stratified subset (seed 42):35    python scripts/bench_accept.py --make-subset-from experiments/prompts/f1-sample.jsonl \36        --prompts experiments/prompts/acc-sample.jsonl --subset-per-domain 6037 38    # 2) run one config (target + drafter):39    python scripts/bench_accept.py --model models/Qwen3-8B-Q4_K_M.gguf \40        --draft models/drafts/dflash_qwen3_8b_block7.gguf --spec-type draft-dflash \41        --config-name q4-dflash-q4 --prompts experiments/prompts/acc-sample.jsonl \42        --out experiments/runs/acc-q4-dflash-q4 --resume43 44    # 3) target-solo (no drafter): omit --draft.45    # 4) smoke/mini-runs: --max-prompts N (first N pending prompts).46    RUN ONLY WITH THE GPU FREE (never while the F3 chain is measuring).47 48Per-record output: {id, domain, text, config, tok_per_s, alpha, tau, draft_n,49prompt_ms, predicted_ms, elapsed_s, attempts, ts} with alpha = draft_n_accepted /50draft_n (None if no draft), tau = draft_n_accepted (accepted tokens),51prompt_ms = TTFT (prefill), predicted_ms = total generation time.52"""53 54from __future__ import annotations55 56import argparse57import csv58import fcntl59import hashlib60import json61import math62import os63import random64import re65import shlex66import socket67import subprocess68import sys69import threading70import time71import urllib.request72from pathlib import Path73 74LLAMA_BIN = Path(os.environ.get("LLAMA_CPP_BIN", os.path.expanduser("~/llama.cpp/build/bin")))75SERVER_BIN = LLAMA_BIN / "llama-server"76 77# llama-cli default n_max in this build (common.h) — parity with F1/F2/F3.78DEFAULT_N_MAX = 379 80# Sampling of the final run: true greedy (T=0 = pure argmax, ignores top_k/top_p),81# fixed seed 42. F1/F2/F3 ran at T=0.8 (llama-cli defaults) — not greedy.82DEFAULT_SAMPLING = {"temperature": 0.0, "top_k": 40, "top_p": 0.95}83 84# Fixed results.csv header (no text column; None → empty cell).85CSV_HEADER = [86    "id",87    "domain",88    "config",89    "tok_per_s",90    "alpha",91    "tau",92    "draft_n",93    "prompt_ms",94    "predicted_ms",95    "elapsed_s",96    "attempts",97    "ts",98]99 100# Shared sha256 cache of GGUFs (keyed path|size_bytes|mtime_ns) — avoids101# re-hashing ~12 GB per config (30-60 s) across the 25-config chain.102SHARED_HASH_CACHE = Path("experiments/runs/model-hashes.json")103 104 105class VramSampler:106    """Samples VRAM (memory.used) and power draw (power.draw) every 3 s in a daemon thread."""107 108    def __init__(self, out_dir: Path) -> None:109        self.out_dir = out_dir110        self._max_mib = 0111        self._max_power = 0.0112        self._lock = threading.Lock()113        self._stop = threading.Event()114        self._thread = threading.Thread(target=self._run, daemon=True)115 116    def start(self) -> None:117        self._thread.start()118 119    def stop(self) -> None:120        self._stop.set()121        self._thread.join(timeout=5)122        with self._lock:123            max_mib = self._max_mib124            max_power = self._max_power125        (self.out_dir / "vram.json").write_text(126            json.dumps(127                {128                    "max_gpu_mib": max_mib,129                    "max_power_w": max_power,130                    "sample_interval_s": 3,131                    "timestamp": time.strftime("%Y-%m-%dT%H:%M:%S"),132                }133            )134        )135 136    def _run(self) -> None:137        while not self._stop.is_set():138            try:139                out = subprocess.run(  # noqa: S603140                    [141                        "/usr/bin/nvidia-smi",142                        "--query-gpu=memory.used,power.draw",143                        "--format=csv,noheader,nounits",144                    ],145                    capture_output=True,146                    text=True,147                    timeout=10,148                    check=True,149                )150                parts = [p.strip() for p in out.stdout.split(",")]151                used = int(parts[0])152                pwr_raw = parts[1] if len(parts) > 1 else ""153                pwr = float(pwr_raw) if pwr_raw not in ("", "[N/A]") else 0.0154                with self._lock:155                    self._max_mib = max(self._max_mib, used)156                    self._max_power = max(self._max_power, pwr)157            except (subprocess.SubprocessError, ValueError, IndexError):158                pass159            self._stop.wait(3)160 161 162def _free_port() -> int:163    with socket.socket() as s:164        s.bind(("127.0.0.1", 0))165        return s.getsockname()[1]166 167 168def _http_json(url: str, payload: dict | None = None, timeout: float = 30.0) -> tuple[int, dict]:169    """GET (payload=None) or POST JSON; returns (status, json)."""170    if payload is None:171        req = urllib.request.Request(url)  # noqa: S310 — localhost172    else:173        req = urllib.request.Request(  # noqa: S310 — localhost174            url,175            data=json.dumps(payload).encode("utf-8"),176            headers={"Content-Type": "application/json"},177        )178    with urllib.request.urlopen(req, timeout=timeout) as resp:  # noqa: S310 — localhost179        return resp.status, json.loads(resp.read().decode("utf-8"))180 181 182class Server:183    """llama-server lifecycle: start → wait_health → stop (always)."""184 185    def __init__(self, cmd: list[str]) -> None:186        self.cmd = cmd187        self.proc: subprocess.Popen | None = None188        self.url = ""189 190    def start(self, port: int) -> None:191        self.url = f"http://127.0.0.1:{port}"192        self.proc = subprocess.Popen(  # noqa: S603 — internally controlled command193            self.cmd,194            stdout=subprocess.DEVNULL,195            stderr=subprocess.DEVNULL,196        )197 198    def wait_health(self, timeout: float = 300.0) -> str:199        """Wait for /health == 200. Returns "" if OK, otherwise the reason."""200        deadline = time.time() + timeout201        while time.time() < deadline:202            if self.proc is not None and self.proc.poll() is not None:203                return f"server died on startup (rc={self.proc.returncode})"204            try:205                status, _ = _http_json(self.url + "/health", timeout=5)206                if status == 200:207                    return ""208            except Exception:  # noqa: BLE001, S110 — still loading209                pass210            time.sleep(2)211        return f"timeout waiting for /health ({int(timeout)}s)"212 213    def stop(self) -> None:214        if self.proc is not None and self.proc.poll() is None:215            self.proc.terminate()216            try:217                self.proc.wait(timeout=10)218            except subprocess.TimeoutExpired:219                self.proc.kill()220        self.proc = None221 222 223def _tail(text: str, n: int = 300) -> str:224    return text[-n:] if text else ""225 226 227def _parse_completion(228    data: dict,229) -> tuple[float | None, float | None, int | None, int | None, float | None, float | None]:230    """(tok_per_s, alpha, tau, draft_n, prompt_ms, predicted_ms) from /completion.231 232    alpha = draft_n_accepted / draft_n; tau = draft_n_accepted. With draft_n 0233    or missing (target-solo configs) alpha/tau/draft_n are None (spec). prompt_ms234    (TTFT) and predicted_ms are draft-independent and always present.235    """236    tim = data.get("timings", {})237    tok_s = tim.get("predicted_per_second")238    draft_n = tim.get("draft_n")239    draft_acc = tim.get("draft_n_accepted")240    alpha: float | None = None241    tau: int | None = None242    if draft_n:243        if draft_acc is not None:244            alpha = round(draft_acc / draft_n, 4)245            tau = draft_acc246    else:247        draft_n = None248    prompt_ms = tim.get("prompt_ms")249    predicted_ms = tim.get("predicted_ms")250    return tok_s, alpha, tau, draft_n, prompt_ms, predicted_ms251 252 253def read_git_commit(repo: str) -> str | None:254    """HEAD commit of the pinned repo (fixed -C); None non-fatal if it fails.255 256    The run proceeds without commit reproducibility (threat matrix: git).257    """258    try:259        out = subprocess.run(  # noqa: S603 — fixed internal repo260            ["git", "-C", str(repo), "rev-parse", "HEAD"],  # noqa: S607 — git from the env PATH261            capture_output=True,262            text=True,263            timeout=10,264            check=True,265        )266    except (subprocess.SubprocessError, FileNotFoundError):267        return None268    return out.stdout.strip() or None269 270 271def parse_llama_version(bin_path: Path) -> str | None:272    """Binary version: 'version: 22 (0713275)' → '22'; 'build: 10249' → '10249'.273 274    None if the binary does not answer or the format is not recognized.275    """276    try:277        out = subprocess.run(  # noqa: S603, S607 — internally controlled binary278            [str(bin_path), "--version"],279            capture_output=True,280            text=True,281            timeout=15,282            check=True,283        )284    except (subprocess.SubprocessError, FileNotFoundError):285        return None286    combined = (out.stdout or "") + "\n" + (out.stderr or "")287    m = re.search(r"version:\s*(\S+)", combined)288    if m:289        return m.group(1)290    m = re.search(r"build:\s*(\d+)", combined)291    return m.group(1) if m else None292 293 294def file_sha256(path: Path, cache: dict[tuple[str, int, int], str] | None = None) -> str | None:295    """SHA-256 of a file; cache keyed (path, size_bytes, mtime_ns) (hit → no re-hash).296 297    None if the file does not exist or is not readable.298    """299    try:300        st = path.stat()301    except OSError:302        return None303    key = (str(path), st.st_size, st.st_mtime_ns)304    if cache is not None and key in cache:305        return cache[key]306    h = hashlib.sha256()307    try:308        with path.open("rb") as f:309            for chunk in iter(lambda: f.read(1 << 20), b""):310                h.update(chunk)311    except OSError:312        return None313    digest = h.hexdigest()314    if cache is not None:315        cache[key] = digest316    return digest317 318 319def _cache_key_str(key: tuple[str, int, int]) -> str:320    return f"{key[0]}|{key[1]}|{key[2]}"321 322 323def load_hash_cache(path: Path) -> dict[tuple[str, int, int], str]:324    """Load the shared cache (JSON with 'path|size|mtime_ns' keys) → dict of tuples."""325    cache: dict[tuple[str, int, int], str] = {}326    try:327        if path.exists():328            raw = json.loads(path.read_text())329            for k, v in raw.items():330                p, s, m = k.rsplit("|", 2)331                cache[(p, int(s), int(m))] = v332    except (OSError, ValueError):333        pass334    return cache335 336 337def save_hash_cache(path: Path, cache: dict[tuple[str, int, int], str]) -> bool:338    """Save the shared cache; True if OK, False if it failed (per-out fallback)."""339    try:340        path.parent.mkdir(parents=True, exist_ok=True)341        path.write_text(json.dumps({_cache_key_str(k): v for k, v in cache.items()}, indent=2))342        return True343    except OSError:344        return False345 346 347def _gguf_meta(path: str, cache: dict[tuple[str, int, int], str]) -> dict:348    p = Path(path)349    st = p.stat() if p.exists() else None350    return {"path": path, "size_bytes": st.st_size if st else None, "sha256": file_sha256(p, cache)}351 352 353def read_reproducibility(354    model: str, draft: str | None, cache: dict[tuple[str, int, int], str]355) -> dict:356    """Reproducibility metadata: llama.cpp commit + version + sha256/size of GGUFs.357 358    Failed git → commit None + warn (the run continues; null reproducibility).359    """360    repo = os.path.expanduser("~/llama.cpp")361    commit = read_git_commit(repo)362    if commit is None:363        print(364            f"[bench] WARN: could not read the commit of {repo} → incomplete reproducibility",365            file=sys.stderr,366        )367    return {368        "llama_cpp_commit": commit,369        "llama_cpp_version": parse_llama_version(SERVER_BIN),370        "host": socket.gethostname(),371        "ts": time.strftime("%Y-%m-%dT%H:%M:%S"),372        "model": _gguf_meta(model, cache),373        "draft": _gguf_meta(draft, cache) if draft else None,374    }375 376 377def build_server_cmd(378    args: argparse.Namespace, port: int, out_dir: Path, log_name: str = "server.log"379) -> list[str]:380    """llama-server command: target (± drafter), non-thinking (--reasoning off).381 382    --reasoning off goes after --verbose and before the extras; no auto-retry383    (if raw mode rejects it, the escape is --extra reasoning_effort:"none",384    see --help of --extra).385    """386    cmd = [387        str(SERVER_BIN),388        "-m",389        args.model,390        "-ngl",391        str(args.n_gpu_layers),392        "-c",393        str(args.ctx),394        "-t",395        str(args.threads),396        "-np",397        "1",398        "--host",399        "127.0.0.1",400        "--port",401        str(port),402        "--log-file",403        str(out_dir / log_name),404        "--verbose",405    ]406    if args.draft:407        cmd += [408            "-md",409            args.draft,410            "--spec-type",411            args.spec_type,412            "--spec-draft-n-max",413            str(args.spec_draft_n_max),414            "-ngld",415            str(args.draft_ngl),416        ]417        if args.spec_draft_p_min is not None:418            cmd += ["--spec-draft-p-min", str(args.spec_draft_p_min)]419    cmd += ["--reasoning", "off"]420    cmd += [a for pair in args.extra for a in shlex.split(pair)]421    return cmd422 423 424def resume_done(results_path: Path, errors_path: Path) -> set[str]:425    """Ids already processed (results.jsonl) or known failures (errors.jsonl) for resume.426 427    Includes failures so prompts that always fail are not retried (e.g. the 8428    arena-hard-v2 ones > ctx 2048) and errors.jsonl entries are not duplicated.429    Tolerates corrupt lines (e.g. power loss) and missing files.430    """431    done: set[str] = set()432    for path in (results_path, errors_path):433        if not path.exists():434            continue435        for line in path.read_text(encoding="utf-8", errors="replace").splitlines():436            try:437                done.add(json.loads(line)["id"])438            except (json.JSONDecodeError, KeyError):439                continue440    return done441 442 443def pending_prompts(records: list[dict], done: set[str], max_prompts: int | None) -> list[dict]:444    """Pending = records not completed (resume), truncated to max_prompts (N>0)."""445    pending = [p for p in records if p["id"] not in done]446    if max_prompts is not None and max_prompts > 0:447        pending = pending[:max_prompts]448    return pending449 450 451def read_results(results_path: Path) -> list[dict]:452    """Read results.jsonl tolerating corrupt lines (power loss) and a missing file."""453    records: list[dict] = []454    if results_path.exists():455        for line in results_path.read_text(encoding="utf-8", errors="replace").splitlines():456            try:457                records.append(json.loads(line))458            except json.JSONDecodeError:459                continue460    return records461 462 463def server_log_path(out_dir: Path) -> Path:464    """llama-server session log: server.log the 1st time, server-2.log/3.log on465    resumes. llama-server TRUNCATES --log-file (fopen "w", log.cpp:322) → a resume466    of a partial config must not erase the "acc per pos" curves of its previous session."""467    first = out_dir / "server.log"468    if not first.exists():469        return first470    n = 2471    while (out_dir / f"server-{n}.log").exists():472        n += 1473    return out_dir / f"server-{n}.log"474 475 476def _server_alive(url: str) -> bool:477    """Probe /health: True if the server answers 200 (not down)."""478    try:479        status, _ = _http_json(url + "/health", timeout=5)480        return status == 200481    except Exception:  # noqa: BLE001, S110 — server down482        return False483 484 485def _mean(xs: list[float]) -> float | None:486    return round(sum(xs) / len(xs), 4) if xs else None487 488 489def _median(xs: list[float]) -> float | None:490    if not xs:491        return None492    s = sorted(xs)493    n = len(s)494    mid = n // 2495    if n % 2 == 1:496        return round(s[mid], 4)497    return round((s[mid - 1] + s[mid]) / 2, 4)498 499 500def _pct(xs: list[float], p: float) -> float | None:501    if not xs:502        return None503    s = sorted(xs)504    idx = min(len(s) - 1, max(0, math.ceil(p / 100 * len(s)) - 1))505    return round(s[idx], 4)506 507 508def write_metrics(out_dir: Path, records: list[dict], failed: int, duration_s: float) -> Path:509    """Write metrics.json with total and per-domain aggregates (math/code/chat).510 511    Resume-safe: the caller passes ALL records (this run + those re-read from512    results.jsonl, which include the previous ones). Aggregates ignore None513    (alpha/tau/draft_n of target-solo configs). sampling/spec/reproducibility514    are copied from config.json; vram/power from vram.json.515    """516    cfg: dict = {}517    vram: dict = {}518    try:519        cfg = json.loads((out_dir / "config.json").read_text())520    except (OSError, ValueError):521        pass522    try:523        vram = json.loads((out_dir / "vram.json").read_text())524    except (OSError, ValueError):525        pass526 527    toks = [r["tok_per_s"] for r in records if r.get("tok_per_s") is not None]528    alphas = [r["alpha"] for r in records if r.get("alpha") is not None]529    taus = [r["tau"] for r in records if r.get("tau") is not None]530    draft_ns = [r["draft_n"] for r in records if r.get("draft_n") is not None]531    ttfts = [r["prompt_ms"] for r in records if r.get("prompt_ms") is not None]532 533    per_domain: dict[str, dict] = {}534    for dom in ("math", "code", "chat"):535        dr = [r for r in records if r.get("domain") == dom]536        dtoks = [r["tok_per_s"] for r in dr if r.get("tok_per_s") is not None]537        dalphas = [r["alpha"] for r in dr if r.get("alpha") is not None]538        dtaus = [r["tau"] for r in dr if r.get("tau") is not None]539        dttfts = [r["prompt_ms"] for r in dr if r.get("prompt_ms") is not None]540        per_domain[dom] = {541            "n": len(dr),542            "tok_per_s.mean": _mean(dtoks),543            "alpha.mean": _mean(dalphas),544            "tau.mean": _mean(dtaus),545            "ttft.mean": _mean(dttfts),546        }547 548    metrics = {549        "prompts": {"total": len(records) + failed, "ok": len(records), "failed": failed},550        "tok_per_s": {551            "mean": _mean(toks),552            "median": _median(toks),553            "p50": _median(toks),554            "p95": _pct(toks, 95),555            "min": round(min(toks), 4) if toks else None,556            "max": round(max(toks), 4) if toks else None,557        },558        "alpha": {"mean": _mean(alphas), "median": _median(alphas)},559        "tau": {"mean": _mean(taus), "median": _median(taus)},560        "draft_n": {"total": sum(draft_ns), "mean": _mean(draft_ns), "median": _median(draft_ns)},561        "ttft": {562            "prompt_ms.mean": _mean(ttfts),563            "prompt_ms.median": _median(ttfts),564            "prompt_ms.p95": _pct(ttfts, 95),565        },566        "vram": {"max_gpu_mib": vram.get("max_gpu_mib")},567        "power": {"max_power_w": vram.get("max_power_w")},568        "duration_s": round(duration_s, 2),569        "errors": failed,570        "per_domain": per_domain,571        "sampling": cfg.get("sampling", {}),572        "spec": cfg.get("spec", {}),573        "reproducibility": cfg.get("reproducibility", {}),574    }575    path = out_dir / "metrics.json"576    path.write_text(json.dumps(metrics, indent=2, ensure_ascii=False) + "\n")577    return path578 579 580def export_csv(results_path: Path) -> Path:581    """Export results.jsonl → results.csv (fixed header, no text column)."""582    csv_path = results_path.with_suffix(".csv")583    with csv_path.open("w", newline="") as f:584        w = csv.writer(f)585        w.writerow(CSV_HEADER)586        for line in results_path.read_text().splitlines():587            if not line.strip():588                continue589            try:590                r = json.loads(line)591            except ValueError:592                continue593            w.writerow([r.get(h) for h in CSV_HEADER])594    return csv_path595 596 597def make_subset(source: Path, dest: Path, per_domain: int, seed: int) -> int:598    """Stratified subset: up to per_domain prompts per domain (shuffle seed)."""599    by_domain: dict[str, list[dict]] = {}600    for line in source.read_text().splitlines():601        if not line.strip():602            continue603        p = json.loads(line)604        by_domain.setdefault(p.get("domain", "unknown"), []).append(p)605    rng = random.Random(seed)  # noqa: S311 — deterministic shuffle with fixed seed606    total = 0607    with dest.open("w") as f:608        for dom in sorted(by_domain):609            chosen = by_domain[dom][:]610            rng.shuffle(chosen)611            chosen = chosen[:per_domain]612            for p in chosen:613                f.write(json.dumps(p, ensure_ascii=False) + "\n")614            total += len(chosen)615            print(f"[subset] {dom}: {len(chosen)}/{len(by_domain[dom])}", file=sys.stderr)616    print(f"[subset] {total} prompts → {dest}", file=sys.stderr)617    return 0618 619 620def main() -> int:621    ap = argparse.ArgumentParser(622        description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter623    )624    ap.add_argument("--model", default=None, help="Target GGUF")625    ap.add_argument("--draft", default=None, help="Drafter GGUF (omit = target-solo)")626    ap.add_argument(627        "--spec-type",628        default="none",629        help="draft-simple|draft-eagle3|draft-mtp|draft-dflash|draft-dspark",630    )631    ap.add_argument(632        "--spec-draft-n-max", type=int, default=DEFAULT_N_MAX, help="llama-cli parity n_max"633    )634    ap.add_argument("--spec-draft-p-min", type=float, default=None, help="DSpark confidence cutoff")635    ap.add_argument("--config-name", default=None, help="Label in every record (default: stems)")636    ap.add_argument("--prompts", required=True, type=Path)637    ap.add_argument("--out", default=None, type=Path)638    ap.add_argument("--n-tokens", type=int, default=256)639    ap.add_argument("--seed", type=int, default=42)640    ap.add_argument("--temperature", type=float, default=DEFAULT_SAMPLING["temperature"])641    ap.add_argument("--top-k", type=int, default=DEFAULT_SAMPLING["top_k"])642    ap.add_argument("--top-p", type=float, default=DEFAULT_SAMPLING["top_p"])643    ap.add_argument("--n-gpu-layers", type=int, default=99)644    ap.add_argument("--draft-ngl", type=int, default=99, help="GPU layers of the drafter (-ngld)")645    ap.add_argument("--ctx", type=int, default=2048)646    ap.add_argument("--threads", type=int, default=8)647    ap.add_argument("--retries", type=int, default=2)648    ap.add_argument("--server-timeout", type=float, default=300.0, help="Wait for /health (s)")649    ap.add_argument(650        "--max-prompts",651        type=int,652        default=None,653        help="Only the first N pending prompts (smoke/mini-runs/OOM-CHECK)",654    )655    ap.add_argument("--resume", action="store_true")656    ap.add_argument(657        "--extra",658        action="append",659        default=[],660        help=(661            "Extra flags for llama-server (shlex). If raw mode rejects --reasoning off, "662            'documented escape: --extra reasoning_effort:"none"'663        ),664    )665    ap.add_argument(666        "--make-subset-from", type=Path, default=None, help="Subset mode: source f1-sample.jsonl"667    )668    ap.add_argument("--subset-per-domain", type=int, default=60)669    args = ap.parse_args()670 671    if args.make_subset_from is not None:672        return make_subset(args.make_subset_from, args.prompts, args.subset_per_domain, args.seed)673 674    if args.model is None or args.out is None:675        ap.error("--model and --out are required outside --make-subset-from mode")676 677    if not SERVER_BIN.exists():678        print(f"[bench] ERROR: {SERVER_BIN} does not exist (did llama.cpp build?)", file=sys.stderr)679        return 1680 681    out_dir = args.out682    out_dir.mkdir(parents=True, exist_ok=True)683    results_path = out_dir / "results.jsonl"684    errors_path = out_dir / "errors.jsonl"685 686    # Reproducibility: commit + version + sha256/size of GGUFs with the shared687    # cache (experiments/runs/model-hashes.json); per-out fallback.688    hash_cache = load_hash_cache(SHARED_HASH_CACHE)689    reproducibility = read_reproducibility(args.model, args.draft, hash_cache)690    if not save_hash_cache(SHARED_HASH_CACHE, hash_cache):691        save_hash_cache(out_dir / "model-hashes.json", hash_cache)692 693    cfg = {k: (str(v) if isinstance(v, Path) else v) for k, v in vars(args).items()}694    cfg["sampling"] = {695        "temperature": args.temperature,696        "top_k": args.top_k,697        "top_p": args.top_p,698        "seed": args.seed,699    }700    cfg["spec"] = {701        "type": args.spec_type,702        "draft_n_max": args.spec_draft_n_max,703        "p_min": args.spec_draft_p_min,704    }705    cfg["reproducibility"] = reproducibility706    (out_dir / "config.json").write_text(json.dumps(cfg, indent=2, ensure_ascii=False))707 708    done: set[str] = set()709    if args.resume:710        done = resume_done(results_path, errors_path)711        print(712            f"[bench] resume: {len(done)} prompts already processed (results + known errors)",713            file=sys.stderr,714        )715 716    records_all: list[dict] = []717    for line in args.prompts.read_text().splitlines():718        if not line.strip():719            continue720        try:721            records_all.append(json.loads(line))722        except json.JSONDecodeError:723            continue724    pending = pending_prompts(records_all, done, args.max_prompts)725    if args.max_prompts is not None:726        print(727            f"[bench] max-prompts={args.max_prompts}: {len(pending)} prompts in this run",728            file=sys.stderr,729        )730 731    if not pending:732        # R3-003: with no pending prompts we do not start a server — protects733        # server.log from truncation ("acc per pos" curves of already-completed734        # configs) and avoids loading model/VRAM during resume walks.735        print(736            "[bench] resume: no pending prompts — not starting the server "737            "(protects server.log, avoids model/VRAM load)",738            file=sys.stderr,739        )740        # Power loss: if the config ended up without metrics/csv (crash before741        # the final aggregation), they are regenerated from results.jsonl — no server/GPU.742        if not (out_dir / "metrics.json").exists() or not (out_dir / "results.csv").exists():743            recs = read_results(results_path)744            export_csv(results_path)745            write_metrics(out_dir, recs, 0, 0.0)746            print(747                "[bench] early-exit: metrics.json/results.csv regenerated (were missing)",748                file=sys.stderr,749            )750        return 0751 752    port = _free_port()753    cmd = build_server_cmd(args, port, out_dir, log_name=server_log_path(out_dir).name)754 755    config_name = args.config_name or (756        Path(args.model).stem757        if not args.draft758        else f"{Path(args.model).stem}+{Path(args.draft).stem}"759    )760 761    server = Server(cmd)762    # Anti-rerun lock BEFORE starting the server (R4-001): if another runner763    # writes this --out, we abort without spawning processes or using the GPU.764    guard = results_path.open("a")765    try:766        fcntl.flock(guard.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB)767    except OSError:768        guard.close()769        print(770            "[bench] ERROR: another runner is writing this --out (lock busy). "771            "Wait for it to finish or check it with ps.",772            file=sys.stderr,773        )774        return 3775 776    t_start = time.time()777    print(f"[bench] starting llama-server at {server.url} (pid launched)", file=sys.stderr)778    server.start(port)779    try:780        reason = server.wait_health(args.server_timeout)781        if reason:782            server.stop()783            guard.close()784            print(f"[bench] ERROR: {reason} → see {out_dir / 'server.log'}", file=sys.stderr)785            return 1786        print("[bench] server OK", file=sys.stderr)787    except Exception:  # noqa: BLE001788        server.stop()789        guard.close()790        raise791 792    n_new = 0793    n_failed = 0794    aborted = False795    vram = VramSampler(out_dir)796    vram.start()797    try:798        with results_path.open("a") as fout, errors_path.open("a") as ferr:799            body = {800                "prompt": None,  # per prompt801                "n_predict": args.n_tokens,802                "seed": args.seed,803                "temperature": args.temperature,804                "top_k": args.top_k,805                "top_p": args.top_p,806                "cache_prompt": False,807                "stream": False,808            }809            consec_conn = 0810            for p in pending:811                if p["id"] in done:812                    continue813                n_new += 1814                record: dict | None = None815                last_err = ""816                detail_text = ""817                for attempt in range(1, args.retries + 1):818                    t0 = time.time()819                    try:820                        body["prompt"] = p["text"]821                        status, data = _http_json(server.url + "/completion", body, timeout=900)822                        tok_s, alpha, tau, draft_n, prompt_ms, predicted_ms = _parse_completion(823                            data824                        )825                        if status != 200:826                            detail_text = _tail(str(data), 300)827                    except Exception as e:  # noqa: BLE001 — overnight hardening828                        status = 0829                        tok_s, alpha, tau, draft_n, prompt_ms, predicted_ms = (830                            None,831                            None,832                            None,833                            None,834                            None,835                            None,836                        )837                        last_err = repr(e)838                    elapsed = time.time() - t0839                    if status == 200 and tok_s is not None:840                        consec_conn = 0841                        record = {842                            "id": p["id"],843                            "domain": p.get("domain"),844                            "text": p["text"],845                            "config": config_name,846                            "tok_per_s": round(tok_s, 3),847                            "alpha": alpha,848                            "tau": tau,849                            "draft_n": draft_n,850                            "prompt_ms": round(prompt_ms, 1) if prompt_ms is not None else None,851                            "predicted_ms": round(predicted_ms, 1)852                            if predicted_ms is not None853                            else None,854                            "elapsed_s": round(elapsed, 2),855                            "attempts": attempt,856                            "ts": time.strftime("%Y-%m-%dT%H:%M:%S"),857                        }858                        break859                    if status == 0:860                        consec_conn += 1861                    else:862                        consec_conn = 0863                    if consec_conn >= 2 and not _server_alive(server.url):864                        print(865                            f"[bench] ERROR: server not responding (2 consecutive "866                            f"connection failures: {last_err}) → abort",867                            file=sys.stderr,868                        )869                        aborted = True870                        break871                    last_err = f"status={status} {last_err}"872                    print(873                        f"[bench] {p['id']} attempt {attempt} failed ({last_err})", file=sys.stderr874                    )875                if record is not None:876                    fout.write(json.dumps(record, ensure_ascii=False) + "\n")877                    fout.flush()878                    print(879                        f"[bench] {p['id']}: {record['tok_per_s']} tok/s "880                        f"α={record['alpha']} τ={record['tau']} ({record['elapsed_s']}s)",881                        file=sys.stderr,882                    )883                else:884                    n_failed += 1885                    ferr.write(886                        json.dumps(887                            {888                                "id": p["id"],889                                "domain": p.get("domain"),890                                "error": last_err,891                                "detail": detail_text or last_err,892                                "attempts": args.retries,893                            },894                            ensure_ascii=False,895                        )896                        + "\n"897                    )898                    ferr.flush()899                    print(900                        f"[bench] {p['id']} FAILED after {args.retries} attempts → errors.jsonl",901                        file=sys.stderr,902                    )903                if aborted:904                    break905    finally:906        server.stop()907        vram.stop()908        guard.close()909 910    # Final metrics: re-read the full results.jsonl (resume-safe) → CSV + JSON.911    records = read_results(results_path)912    export_csv(results_path)913    write_metrics(out_dir, records, n_failed, time.time() - t_start)914 915    print(916        f"[bench] done: {n_new - n_failed} new valid, {n_failed} failures, "917        f"{len(done)} previous → {results_path}"918    )919    if aborted:920        return 1921    return 0 if n_failed == 0 else 2922 923 924if __name__ == "__main__":925    sys.exit(main())926