0b37d286065ecbf5f6dfcfafa2d9274d1d21dba3 / bench/tasks/needle.py · 11702 bytes · raw
#!/usr/bin/env python3
"""
needle.py — repeatable needle-in-a-haystack battery for quant/lever A/Bs.
Replaces the ad-hoc one-off probes (EMERALD-7742 @66.5k etc.) with a committed
instrument. Measures retrieval quality at depth AND speed telemetry per depth,
so a candidate quant can be compared against the incumbent on one run.
Method (mirrors halogen-flash's published battery + our 2026-09-03 probes):
- Filler text is neutral prose paragraphs (seeded, deterministic per depth).
- A needle is a single sentence with a random code token: "The maintenance
code for the {place} is {CODE}-{NNNN}." Spliced into filler at a target
depth, the document continues, then asks for the code. Exact string match
on the retrieved code. Temp 0 (greedy) — quant differences only, no
sampling noise.
- Speed: prefill tok/s (prompt_eval from usage), decode tok/s (eval), wall.
- Multiple needles per depth, different insertion positions -> a distribution,
not a single anecdote. Also: one NEAR-MISS needle per depth (a second code
exists elsewhere in the doc; the model must return the RIGHT one) to catch
confabulation.
Usage:
python bench.py needle run # full battery
python bench.py needle run --depths 1000,16000 # quick subset
python bench.py needle run --needles 3 --repeats 2 # smoke config
Env (via .env / bench.core.client): ENDPOINT, MODEL, API_KEY, MAX_TOKENS,
TEMPERATURE, HTTP_TIMEOUT. Plus:
LABEL tag for the JSON/report (default MODEL)
N_PREDICT max answer tokens (default 512; reasoning models think first)
Output: results/needle-<label>.json + printed table. Exit 0 if every needle
retrieved exact (any miss -> exit 1, so CI-style gating can catch it).
"""
import json
import os
import random
import sys
import time
from bench.core import client
from bench.core.reporting import write_json, RESULTS_DIR
LABEL = os.environ.get("LABEL", client.MODEL)
N_PREDICT = int(os.environ.get("N_PREDICT", "512"))
DEPTH_GRID = [1000, 4000, 16000, 32000, 64000]
# Deterministic per-depth filler. Neutral office-log prose, no code tokens.
FILLER_SENTENCES = [
"The quarterly inventory review proceeded without any material discrepancies noted by the floor staff.",
"Facilities confirmed the heating schedule for the west corridor will revert to the winter profile next week.",
"The procurement team circulated updated vendor terms for the spring contract cycle.",
"Attendance at the all-hands was recorded at eighty-four percent, which is typical for this month.",
"The backup generator passed its monthly load test and the transfer switch responded within specification.",
"Reception logged three courier deliveries before noon and routed each to the appropriate department.",
"The finance office reminded everyone that expense submissions close on the fifteenth of the month.",
"A scheduled patch window for the document management system was announced for Thursday evening.",
"The landscaping crew completed the seasonal pruning along the south entrance without incident.",
"Warehouse slotting for the new product line was finalized and communicated to the picking team.",
"The safety committee published minutes from its monthly walkthrough of the loading dock area.",
"IT reported that the printer fleet firmware update completed on all but two devices.",
]
PLACES = ["north stairwell", "server room", "loading dock", "rooftop hatch", "maintenance office",
"freight elevator", "parking garage level two", "electrical closet", "boiler room", "west fire exit"]
CODE_WORDS = ["EMERALD", "COPPER", "GRANITE", "LANTERN", "HARBOR", "MERIDIAN", "SAPPHIRE", "THISTLE"]
def _para(rng):
return " ".join(rng.choice(FILLER_SENTENCES) for _ in range(4)) + "\n\n"
def build_doc(target_tokens, needle, needle_at_frac, rng):
"""Filler document with `needle` spliced at needle_at_frac of target size."""
parts, approx = [], 0
while approx < target_tokens:
parts.append(_para(rng))
approx += 64 # ~64 tok/paragraph (3 sentences avg)
doc_parts_before = []
tok_before = int(target_tokens * needle_at_frac)
idx = max(1, tok_before // 64)
head = "".join(parts[:idx])
tail = "".join(parts[idx:])
return head + needle + "\n\n" + tail
def make_needle(rng):
place = rng.choice(PLACES)
word = rng.choice(CODE_WORDS)
num = rng.randint(1000, 9999)
code = f"{word}-{num}"
sentence = f"Maintenance note: the access code for the {place} is {code} until the end of the quarter."
return sentence, code, place
def make_decoy(rng):
place = rng.choice(PLACES)
word = rng.choice(CODE_WORDS)
num = rng.randint(1000, 9999)
sentence = (f"Maintenance note: the access code for the {place} was {word}-{num} "
f"until it was rotated last month.")
return sentence
def call_model(prompt, n_predict=N_PREDICT, max_retries=1):
"""Delegate to bench.core.client.call_model and extract needle-relevant telemetry.
needle_bench needs the `timings` field from llama-server responses for
prefill/decode tok/s, so we keep the local parsing of resp.get("timings")
and resp.get("usage") on top of the shared client.
"""
last_err = None
for attempt in range(max_retries + 1):
try:
resp, wall = client.call_model(
[{"role": "user", "content": prompt}],
max_tokens=n_predict,
temperature=0.0,
http_timeout=1800,
)
usage = resp.get("usage", {})
timings = resp.get("timings", {})
ctok = usage.get("completion_tokens", 0)
ptok = usage.get("prompt_tokens", 0)
# llama-server reports rates top-level in `timings` (prompt_per_second /
# predicted_per_second); OpenAI-style usage has no rate fields.
pd = timings.get("prompt_per_second")
ev = timings.get("predicted_per_second")
msg = resp.get("choices", [{}])[0].get("message", {})
content = msg.get("content") or ""
reasoning = msg.get("reasoning_content") or ""
return {"content": content, "reasoning": reasoning, "ptok": ptok, "ctok": ctok,
"prefill_ps": pd, "decode_ps": ev, "wall": wall,
"finish": resp.get("choices", [{}])[0].get("finish_reason", "")}
except Exception as e:
last_err = str(e)[:300]
time.sleep(2)
return {"error": last_err}
def run_battery(needles, depths, repeats, decoy_probe=True):
results = []
for depth in depths:
for rep in range(repeats):
# ---- normal needles at varying positions
for n in range(needles):
rng = random.Random(f"{depth}-{rep}-{n}")
sentence, code, place = make_needle(rng)
pos_frac = (n + 0.5) / needles # spread across the doc
doc = build_doc(depth, sentence, pos_frac, rng)
prompt = (doc + "\n\nQUESTION: According to the maintenance notes above, "
"what is the access code for the " + place +
"? Answer with the code exactly as written, nothing else.")
r = call_model(prompt)
ok = code in r.get("content", "")
results.append({"depth": depth, "rep": rep, "needle_idx": n,
"code": code, "exact": ok, **{k: r.get(k) for k in
("ptok", "ctok", "prefill_ps", "decode_ps", "wall", "error")}})
_print_row(results[-1])
# ---- decoy probe: right needle + rotated-out old code in the same doc
if decoy_probe:
rng = random.Random(f"{depth}-{rep}-decoy")
sentence, code, place = make_needle(rng)
decoy_sent = make_decoy(rng)
rng2 = random.Random(f"{depth}-{rep}-decoy-tail")
head = build_doc(int(depth * 0.4), sentence, 0.5, rng)
tail = build_doc(int(depth * 0.6), decoy_sent, 0.5, rng2)
prompt = (head + tail + "\n\nQUESTION: According to the maintenance notes above, "
"what is the CURRENT access code for the " + place +
"? Answer with the code exactly as written, nothing else.")
r = call_model(prompt)
ok = code in r.get("content", "")
results.append({"depth": depth, "rep": rep, "needle_idx": "decoy",
"code": code, "exact": ok, **{k: r.get(k) for k in
("ptok", "ctok", "prefill_ps", "decode_ps", "wall", "error")}})
_print_row(results[-1])
return results
def _print_row(r):
if r.get("error"):
print(f" depth={r['depth']:>6} rep={r['rep']} needle={r['needle_idx']!s:>5} ERROR: {r['error']}")
return
print(f" depth={r['depth']:>6} rep={r['rep']} needle={r['needle_idx']!s:>5} "
f"{'PASS' if r['exact'] else 'MISS'} ptok={r.get('ptok', 0):>6} prefill={r.get('prefill_ps') or 0:>7.1f} "
f"decode={r.get('decode_ps') or 0:>5.1f} wall={r.get('wall', 0):>6.1f}s")
def summarize(results):
by_depth = {}
for r in results:
d = by_depth.setdefault(r["depth"], {"pass": 0, "total": 0, "walls": [], "prefills": [], "decodes": []})
d["total"] += 1
if r.get("exact"):
d["pass"] += 1
if not r.get("error"):
if r.get("wall") is not None:
d["walls"].append(r["wall"])
if r.get("prefill_ps"):
d["prefills"].append(r["prefill_ps"])
if r.get("decode_ps"):
d["decodes"].append(r["decode_ps"])
print("\n== Summary ==")
print(f"{'depth':>8} {'retrieved':>10} {'wall mean':>10} {'prefill t/s':>12} {'decode t/s':>11}")
for depth in sorted(by_depth):
d = by_depth[depth]
walls = d["walls"]
prefills = [p for p in d["prefills"] if p]
decodes = d["decodes"]
def mean(xs):
return sum(xs) / len(xs) if xs else 0.0
print(f"{depth:>8} {d['pass']:>4}/{d['total']:<4} {mean(walls):>9.1f}s {mean(prefills):>11.1f} {mean(decodes):>10.1f}")
total_pass = sum(1 for r in results if r.get("exact"))
print(f"\nTOTAL: {total_pass}/{len(results)} exact")
return total_pass == len(results)
def run():
args = sys.argv[1:]
if not args or args[0] not in ("run",):
print(__doc__)
sys.exit(2)
needles = 3
repeats = 1
depths = list(DEPTH_GRID)
rest = args[1:]
if "--needles" in rest:
needles = int(rest[rest.index("--needles") + 1])
if "--repeats" in rest:
repeats = int(rest[rest.index("--repeats") + 1])
if "--depths" in rest:
depths = [int(x) for x in rest[rest.index("--depths") + 1].split(",")]
print(f"needle_bench: model={client.MODEL} endpoint={client.ENDPOINT} label={LABEL}")
print(f"depths={depths} needles/depth={needles} repeats={repeats} decoy=on")
t0 = time.time()
results = run_battery(needles, depths, repeats)
all_ok = summarize(results)
payload = {"label": LABEL, "model": client.MODEL, "endpoint": client.ENDPOINT,
"ts": time.strftime("%Y-%m-%d %H:%M:%S"),
"needles_per_depth": needles, "repeats": repeats, "depths": depths, "results": results}
out_path = write_json(payload, LABEL, prefix="needle")
print(f"wrote {out_path} ({time.time() - t0:.0f}s total)")
sys.exit(0 if all_ok else 1)
def main():
run()
if __name__ == "__main__":
main()