Solomon / mlx /src /solomon_mlx /evaluation.py
kelseyway's picture
Make MLX adapter-only and include the reproducible BF16 converter
5c0a4a8
Raw
History Blame Contribute Delete
13.9 kB
"""Resumable panel scoring, fit-only calibration and one-shot held-out reports."""
import gzip
import hashlib
import json
from collections import defaultdict
from pathlib import Path
import numpy as np
from scipy.optimize import minimize_scalar
from ._vendor.semantics import listed_probs, p_yes
from .api import TASKS
from .artifacts import ADAPTER_SHA, HEADS_SHA, digest, sha256
def load_panel(directory):
directory = Path(directory)
manifest = json.loads((directory / "manifest.json").read_text())
if (
manifest["adapter_sha256"] != ADAPTER_SHA
or manifest["heads_sha256"] != HEADS_SHA
or manifest["readout_mode"] != "four_collapsed"
):
raise ValueError("Panel belongs to a different checkpoint or answer semantics")
raw = gzip.decompress((directory / "jobs.json.gz").read_bytes())
if hashlib.sha256(raw).hexdigest() != manifest["jobs_sha256"]:
raise ValueError("Panel jobs checksum mismatch")
jobs = json.loads(raw)
if len({r["id"] for r in jobs}) != len(jobs):
raise ValueError("Duplicate panel branch IDs")
return jobs, manifest
def score_panel(model, panel, output):
"""Atomically persist each document so interruption never requires rescoring it."""
jobs, manifest = load_panel(panel)
output = Path(output)
output.mkdir(parents=True, exist_ok=True)
identity = {
"runtime": model.identity,
"panel_sha256": manifest["jobs_sha256"],
"panel_role": Path(panel).name.split("-")[0],
}
meta = output / "identity.json"
if meta.exists() and json.loads(meta.read_text()) != identity:
raise ValueError("Cannot resume with different model code, weights or panel")
meta.write_text(json.dumps(identity, indent=2))
documents = defaultdict(list)
for row in jobs:
documents[row["document_key"]].append(row)
for key, group in documents.items():
path = output / (key + ".json")
if path.exists():
record = json.loads(path.read_text())
body = {k: v for k, v in record.items() if k != "sha256"}
if (
record["sha256"] != digest(body)
or record["identity"] != digest(identity)
or [r["id"] for r in record["rows"]] != [r["id"] for r in group]
):
raise ValueError("Corrupt or mismatched resumed document")
continue
parts = group[0].get("parts") or [{"text": group[0]["doc"]}]
if any((r.get("parts") or [{"text": r["doc"]}]) != parts for r in group):
raise ValueError("Document key aliases different sources")
with model.prefill(parts) as state:
rows = []
for job in group:
result = model.engine.ask(state._data, job["block"], job["n"], job["head_key"])
rows.append({**result, **{k: job[k] for k in ("id", "task", "gold", "n", "question_id")}})
body = {
"identity": digest(identity),
"rows": rows,
"prefix_tokens": state.prefix_tokens,
"prefill_seconds": state._data["prefill_seconds"],
}
temp = path.with_suffix(".tmp")
temp.write_text(json.dumps({**body, "sha256": digest(body)}))
temp.replace(path)
print("Scored " + key + " " + str(len(rows)) + " branches", flush=True)
completed = {
"identity": digest(identity),
"documents": len(documents),
"branches": len(jobs),
"files": {key + ".json": sha256(output / (key + ".json")) for key in documents},
}
(output / "complete.json").write_text(json.dumps(completed, indent=2))
def read_scores(directory):
directory = Path(directory)
identity = json.loads((directory / "identity.json").read_text())
completed = json.loads((directory / "complete.json").read_text())
if completed["identity"] != digest(identity):
raise ValueError("Score identity mismatch")
rows = []
for name, checksum in completed["files"].items():
p = directory / name
if not p.resolve().is_relative_to(directory.resolve()) or sha256(p) != checksum:
raise ValueError("Score checksum mismatch")
record = json.loads(p.read_text())
rows.extend(record["rows"])
if len(rows) != completed["branches"]:
raise ValueError("Incomplete score set")
return rows, identity
def unit(row, temperature=1.0):
logits = row["letter_logits"]
task = row["task"]
gold = row["gold"]
if task in ("boolean", "entity", "multilabel"):
p = p_yes(logits, temperature)
return [1 - p, p], int(gold == 0)
width = row["n"] - 2 if row["head_key"].endswith("choiceR") else row["n"]
if not isinstance(gold, int) or not 0 <= gold < width:
return None, None
return listed_probs(logits, width, temperature).tolist(), gold
def fit_calibration(scores, output, *, panel_role):
if panel_role != "fit":
raise ValueError("Temperature fitting accepts fit panels only")
rows, identity = read_scores(scores)
if identity["panel_role"] != "fit":
raise ValueError("Scores were not generated from a fit panel")
output = Path(output)
if output.exists():
raise FileExistsError("Calibration artifacts are immutable")
temperatures, losses = {}, {}
for task in TASKS:
selected = [r for r in rows if r["task"] == task and unit(r)[0] is not None]
if not selected:
raise ValueError("No fit examples for " + task)
def loss(log_t, selected=selected):
t = float(np.exp(log_t))
return float(np.mean([-np.log(max(unit(r, t)[0][unit(r, t)[1]], 1e-300)) for r in selected]))
fit = minimize_scalar(loss, bounds=(np.log(0.05), np.log(20)), method="bounded")
temperatures[task] = float(np.exp(fit.x))
losses[task] = {"before": loss(0.0), "after": float(fit.fun), "units": len(selected)}
payload = {
"schema": "solomon-mlx-temperature-v1",
"runtime": identity["runtime"]["fingerprint"],
"temperatures": temperatures,
"fit_panel_sha256": identity["panel_sha256"],
"losses": losses,
"selection_role": "fit",
"heldout_used": False,
}
output.write_text(json.dumps({**payload, "sha256": digest(payload)}, indent=2))
return payload
def compare_rows(mlx_rows, cuda_rows, *, temperatures=None, reference_temperatures=None):
temperatures = temperatures or dict.fromkeys(TASKS, 1.0)
reference_temperatures = reference_temperatures or dict.fromkeys(TASKS, 1.0)
reference = {r["id"]: r for r in cuda_rows}
if len(reference) != len(cuda_rows) or set(reference) != {r["id"] for r in mlx_rows}:
raise ValueError("Comparison panels have different or duplicate branch IDs")
units = []
questions = defaultdict(list)
for row in mlx_rows:
other = {**row, "letter_logits": reference[row["id"]]["letter_logits"]}
p, gold = unit(row, temperatures[row["task"]])
q, _ = unit(other, reference_temperatures[row["task"]])
if p is None:
continue
left, right = int(np.argmax(p)), int(np.argmax(q))
item = {
"agreement": left == right,
"mlx_correct": left == gold,
"cuda_correct": right == gold,
"probability_drift": float(np.max(np.abs(np.asarray(p) - q))),
}
units.append(item)
questions[row["question_id"]].append(item)
if not units:
raise ValueError("No defined comparison targets")
agreement = float(np.mean([r["agreement"] for r in units]))
question_agreement = float(np.mean([all(x["agreement"] for x in r) for r in questions.values()]))
mlx_accuracy = float(np.mean([all(x["mlx_correct"] for x in r) for r in questions.values()]))
cuda_accuracy = float(np.mean([all(x["cuda_correct"] for x in r) for r in questions.values()]))
return {
"units": len(units),
"questions": len(questions),
"unit_decision_agreement": agreement,
"question_decision_agreement": float(
np.mean([all(x["agreement"] for x in r) for r in questions.values()])
),
"mlx_whole_question_accuracy": mlx_accuracy,
"cuda_whole_question_accuracy": cuda_accuracy,
"accuracy_degradation_percentage_points": 100 * (cuda_accuracy - mlx_accuracy),
"max_probability_drift": max(r["probability_drift"] for r in units),
"probability_comparison": {
"mlx_temperatures": temperatures,
"cuda_temperatures": reference_temperatures,
},
"mean_probability_drift": float(np.mean([r["probability_drift"] for r in units])),
"quality_gate_passed": agreement >= 0.999
and question_agreement >= 0.999
and cuda_accuracy - mlx_accuracy <= 0.0025,
}
def read_cuda_scores(directory, panel, reference_identity):
"""Reuse only scores bound to the exact pinned CUDA runtime and panel."""
jobs, manifest = load_panel(panel)
directory = Path(directory)
result = {}
for file in sorted(directory.glob("scores*.json.gz")):
payload = json.loads(gzip.decompress(file.read_bytes()))
identity = payload["identity"]
if not payload["complete"] or identity["runtime"] != reference_identity:
raise ValueError("Existing CUDA scores do not match the fresh reference runtime")
if identity["manifest"]["jobs_sha256"] != manifest["jobs_sha256"]:
raise ValueError("CUDA scores use another panel")
for key, row in payload["scores"].items():
if key in result:
raise ValueError("Duplicate CUDA score ID")
result[key] = row
if set(result) != {r["id"] for r in jobs}:
raise ValueError("CUDA score set is incomplete")
return [{**r, **result[r["id"]]} for r in jobs]
def select_calibration(fitted, dev_scores, output):
fitted, output = Path(fitted), Path(output)
if output.exists():
raise FileExistsError("Selected calibration is immutable")
fit = json.loads(fitted.read_text())
fit_payload = {k: v for k, v in fit.items() if k != "sha256"}
rows, identity = read_scores(dev_scores)
if (
fit["sha256"] != digest(fit_payload)
or identity["runtime"]["fingerprint"] != fit["runtime"]
or identity["panel_role"] != "dev"
):
raise ValueError("Calibration or development identity mismatch")
temperatures, selection = {}, {}
for task in TASKS:
selected = [r for r in rows if r["task"] == task and unit(r)[0] is not None]
if not selected:
raise ValueError("Missing development task " + task)
def loss(t, selected=selected):
values = [unit(r, t) for r in selected]
return float(np.mean([-np.log(max(p[g], 1e-300)) for p, g in values]))
original, candidate = loss(1.0), loss(fit["temperatures"][task])
temperatures[task] = fit["temperatures"][task] if candidate < original else 1.0
selection[task] = {"untempered_nll": original, "fit_temperature_nll": candidate}
payload = {
**fit_payload,
"temperatures": temperatures,
"selection_role": "dev_selected",
"fit_artifact_sha256": sha256(fitted),
"dev_panel_sha256": identity["panel_sha256"],
"development_selection": selection,
}
output.write_text(json.dumps({**payload, "sha256": digest(payload)}, indent=2))
return payload
def heldout_report(
scores,
cuda_directory,
panel,
calibration,
reference,
output,
*,
reference_binding="evaluations/cuda-acceptance/input/serving-binding.json",
):
"""Evaluate a frozen configuration once; an existing output cannot be replaced."""
output = Path(output)
if output.exists():
raise FileExistsError("Held-out report already exists; do not reuse it for selection")
rows, identity = read_scores(scores)
cal = json.loads(Path(calibration).read_text())
payload = {k: v for k, v in cal.items() if k != "sha256"}
if (
cal["sha256"] != digest(payload)
or cal["runtime"] != identity["runtime"]["fingerprint"]
or cal["selection_role"] != "dev_selected"
or identity["panel_role"] != "cert"
):
raise ValueError(
"Held-out evaluation requires frozen development-selected calibration and cert scores"
)
ref = json.loads(Path(reference).read_text())
cuda = read_cuda_scores(cuda_directory, panel, ref["identity"])
binding_path = Path(reference_binding)
source_manifest = json.loads((binding_path.parent / "manifest.json").read_text())
if sha256(binding_path) != source_manifest["files"][binding_path.name]:
raise ValueError("CUDA acceptance binding checksum mismatch")
binding = json.loads(binding_path.read_text())
for key in (
"adapter_sha256",
"trained_heads_sha256",
"model_sha256",
"numerics",
"placement",
"arithmetic",
):
if binding["runtime"][key] != ref["identity"][key]:
raise ValueError("CUDA calibration belongs to another reference runtime")
reference_temperatures = {task: binding["temperatures"]["models"][task]["temperature"] for task in TASKS}
report = {
**compare_rows(
rows, cuda, temperatures=cal["temperatures"], reference_temperatures=reference_temperatures
),
"cuda_calibration_binding_sha256": sha256(binding_path),
"runtime": identity["runtime"],
"panel_sha256": identity["panel_sha256"],
"calibration_sha256": sha256(calibration),
"reference_sha256": sha256(reference),
"scope": "text-only held-out panel",
"image_qualification": False,
}
output.write_text(json.dumps(report, indent=2))
return report