reyden009/speculative-decoding-lab
8
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 