LFM2.5-2.6B-RLCD / pcd /benchmark.py
monotykamary's picture
feat: add verified inference-only parallel constrained decoding
3145545 verified
Raw History Blame Contribute Delete
8.95 kB
"""Bounded GPU validation and raw benchmark evidence. Never trains the model."""
import hashlib
import json
import platform
import random
import time
import uuid
from pathlib import Path
import numpy as np
import torch
from .schema import evaluate
from .tasks import audit, diagnostic, stress
def validate(engine):
cache = engine.validate_cache()
comparisons = []
case = diagnostic()[4]
for mode in ("token", "sequence"):
actual = engine.constrained(case["context"], case["schema"], mode=mode)
reference = engine.reference(case["context"], case["schema"], mode=mode)
max_error = 0.0
for field, expected in reference.items():
scores = [v["score"] for v in actual["fields"][field]["candidates"]]
error = max(abs(a - b) for a, b in zip(scores, expected))
max_error = max(max_error, error)
# BF16 checks also retain raw errors/margins rather than claiming bitwise parity.
tolerance = {torch.bfloat16: 0.35, torch.float16: 0.10}.get(
next(engine.model.parameters()).dtype, 0.002
)
if error > tolerance:
raise AssertionError(f"{mode}/{field}: cached score error {error} > {tolerance}")
if (
sorted(expected, reverse=True)[0] - sorted(expected, reverse=True)[1]
> 2 * tolerance
):
assert np.argmax(scores) == np.argmax(expected), "non-ambiguous winner changed"
comparisons.append({"mode": mode, "max_score_error": max_error, "result": actual})
return {"passed": True, "cache": cache, "score_comparisons": comparisons}
def summarize(rows):
result = {}
for mode in sorted({row["method"] for row in rows}):
values = [r for r in rows if r["method"] == mode]
times = [r["result"]["elapsed_ms"] for r in values]
unique = {r["case_id"]: r for r in values}
evaluations = [r["evaluation"] for r in unique.values()]
result[mode] = {
"requests": len(values),
"unique_cases": len(unique),
"mean_ms": float(np.mean(times)),
"median_ms": float(np.median(times)),
"p95_ms": float(np.percentile(times, 95)),
"schema_compliance": sum(v["schema_compliant"] for v in evaluations) / len(evaluations),
"field_accuracy": sum(v["field_correct"] for v in evaluations)
/ sum(v["field_total"] for v in evaluations),
"exact_accuracy": sum(v["exact_match"] for v in evaluations) / len(evaluations),
"peak_allocated_gib": max(r["peak_allocated_bytes"] for r in values) / 2**30,
"truncations": sum(r["result"].get("hit_token_limit", False) for r in values),
}
return result
def execute(engine, task="validate", suite="diagnostic", repeats=3):
if (
task not in {"validate", "benchmark", "probe"}
or suite not in {"diagnostic", "stress"}
or not 1 <= repeats <= 10
):
raise ValueError("invalid task, suite, or repeats")
begin = time.perf_counter()
run_id = (
time.strftime("%Y%m%d-%H%M%S", time.gmtime())
+ "-"
+ task
+ "-"
+ suite
+ "-"
+ uuid.uuid4().hex[:6]
)
report = {
"run_id": run_id,
"task": task,
"suite": suite,
"metadata": engine.metadata(),
"platform": platform.platform(),
"source_sha256": {
p.name: hashlib.sha256(p.read_bytes()).hexdigest()
for p in sorted(Path(__file__).parent.glob("*.py"))
},
"validation": validate(engine),
"rows": [],
"limitations": [
"small hand-authored diagnostic set, not held-out production data",
"SDPA and reference convolution, not optimized vLLM/SGLang",
"no statistical calibration or Jev parity evaluation",
"GPU warm request time excludes model load, network and container startup",
],
}
if task == "probe":
from .prompting import prompt_tokens
case = diagnostic()[0]
compiled = engine.compile(case["schema"], "token")
prefix = prompt_tokens(engine.tokenizer, compiled, case["context"], engine.limits)
report["label_probe"] = []
with torch.inference_mode():
for field in compiled.fields:
ids = engine.tensor([prefix + list(field.suffix)])
logits = engine.model(ids, use_cache=False, logits_to_keep=1).logits[0, -1].float()
top = logits.topk(5)
output = engine.model.generate(
ids,
attention_mask=torch.ones_like(ids),
do_sample=False,
max_new_tokens=8,
pad_token_id=engine.tokenizer.pad_token_id,
eos_token_id=engine.tokenizer.eos_token_id,
)
report["label_probe"].append(
{
"field": field.name,
"labels": field.labels,
"top_tokens": [
{"text": engine.tokenizer.decode([int(t)]), "logit": float(s)}
for t, s in zip(top.indices, top.values)
],
"greedy": engine.tokenizer.decode(output[0, ids.shape[1] :].tolist()),
}
)
if task in {"benchmark", "probe"}:
cases = (
(
[{**c, "split": "development"} for c in diagnostic()]
+ [{**c, "split": "audit"} for c in audit()]
)
if suite == "diagnostic"
else [{**c, "split": "synthetic"} for c in stress()]
)
random.Random(17).shuffle(cases)
methods = (
["token", "sequence"]
if task == "probe"
else ["token", "sequence", "ar_token", "ar_sequence"]
)
for case_index, case in enumerate(cases):
# Native reasoning was measured separately in the exploratory run.
# Do not pay for repeating it in every final benchmark.
case_methods = list(methods)
if case_index % 2:
case_methods = list(reversed(case_methods))
for method in case_methods:
for repeat in range(-1, repeats):
if time.perf_counter() - begin > 650:
report["budget_exhausted"] = True
report["summary"] = summarize(report["rows"])
report["wall_seconds"] = time.perf_counter() - begin
return report
if engine.device.type == "cuda":
torch.cuda.reset_peak_memory_stats(engine.device)
if method in {"token", "sequence"}:
output = engine.constrained(case["context"], case["schema"], mode=method)
else:
output = engine.autoregressive(
case["context"],
case["schema"],
mode="token" if method == "ar_token" else "sequence",
native=method == "native_ar",
max_new_tokens=512 if suite == "stress" else 256,
)
if repeat >= 0:
row = {
"case_id": case["id"],
"split": case["split"],
"method": method,
"repeat": repeat,
"result": output,
"evaluation": evaluate(
output["text"], case["schema"], case["expected"]
),
"peak_allocated_bytes": torch.cuda.max_memory_allocated(engine.device)
if engine.device.type == "cuda"
else 0,
}
report["rows"].append(row)
# Persist partial evidence as we go; abrupt failures do not erase earlier cases.
if Path("/results").is_dir():
Path(f"/results/{run_id}.partial.json").write_text(
json.dumps(report, ensure_ascii=False)
)
print(f"completed {case['id']}", flush=True)
report["cases"] = cases
report["summary"] = summarize(report["rows"]) if report["rows"] else {"validation_passed": True}
report["split_summary"] = {
split: summarize([row for row in report["rows"] if row["split"] == split])
for split in sorted({row["split"] for row in report["rows"]})
}
report["wall_seconds"] = time.perf_counter() - begin
return report