vimeml-tiny-ja-v2.1 / source /scripts /benchmarks /compare_expanded_ime.py
Voltline's picture
Release VimeML V2.1 step40000 FP32 and Core ML INT8 (GPL-2.0)
29f25be verified
Raw History Blame Contribute Delete
5.04 kB
"""Compare frozen development orders without opening the blind LM split."""
import argparse
import json
import sys
from pathlib import Path
ROOT = Path(__file__).resolve().parents[2]
sys.path.insert(0, str(ROOT / "src"))
from vimeml.benchmarks.comparison import paired, rows
from vimeml.benchmarks.evaluate_ajimee import metrics
from vimeml.training.data import write_json
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--v1", type=Path, default=ROOT / "outputs/ime-eval/expanded-v21-dev-draft-v1"
)
parser.add_argument(
"--v2", type=Path, default=ROOT / "outputs/ime-eval/expanded-v21-dev-draft-v2"
)
parser.add_argument(
"--output", type=Path, default=ROOT / "outputs/ime-eval/expanded-v21-dev-draft-comparison"
)
args = parser.parse_args()
if args.output.exists():
parser.error("Preserve earlier comparisons; choose a fresh output.")
reports = {
name: json.loads((path / "metrics.json").read_text(encoding="utf-8"))
for name, path in [("v1", args.v1), ("v2", args.v2)]
}
for report in reports.values():
if (
report.get("result_role") != "provisional_expanded_development"
or report["benchmark_manifest"]["split"] != "development"
):
raise ValueError("Only expanded development diagnostic scores allowed here.")
before, after = rows(args.v1 / "scores.jsonl"), rows(args.v2 / "scores.jsonl")
if len(before) != 2000 or len(after) != 2000:
raise ValueError("Expected all 2000 development cases.")
result = paired(before, after)
result.update(
{
"format": "expanded_ime_development_comparison_v1",
"labels_formal_gold": False,
"blind_lm_scored": False,
"result_role": "provisional_label_diagnostic",
"candidate_label_version": "ime-expanded-v21-candidates-v1",
"models": {name: report["model"] for name, report in reports.items()},
"methods": {name: report["metrics"] for name, report in reports.items()},
"elapsed_seconds": {
name: report["elapsed_seconds"] for name, report in reports.items()
},
"strata": {},
"policy": "Same actual candidate pool, full denominator, fixed suffix logP sum, FP32 CUDA, no EOS/truncation.",
"limitations": [
"Provisional references; spelling gaps and corpus noise still need audit",
"Development diagnostics, not final blind inference or production selection",
"Near-duplicate overlap is not fully audited",
],
}
)
dimensions = {
"reading_length": lambda r: "<=16"
if len(r["query"]) <= 16
else "17-32"
if len(r["query"]) <= 32
else ">32",
"candidate_count": lambda r: "1-5"
if len(r["candidates"]) <= 5
else "6-10"
if len(r["candidates"]) <= 10
else "11-20",
}
for dim, group in dimensions.items():
result["strata"][dim] = {}
for key in sorted({group(r) for r in before}):
a = [r for r in before if group(r) == key]
b = [r for r in after if group(r) == key]
result["strata"][dim][key] = {
"azookey": metrics(a, "azookey"),
"v1": metrics(a, "lm_context_sum"),
"v2": metrics(b, "lm_context_sum"),
}
records = []
for a, b in zip(before, after):
assert a["id"] == b["id"]
records.append(
{
"id": a["id"],
"query": a["query"],
"context": a["left_context"],
"answers": a["answers"],
"covered_exact": bool(set(a["answers"]).intersection(a["orders"]["azookey"])),
"azookey_top1": a["orders"]["azookey"][0] if a["orders"]["azookey"] else "",
"v1_top1": a["orders"]["lm_context_sum"][0]
if a["orders"]["lm_context_sum"]
else "",
"v2_top1": b["orders"]["lm_context_sum"][0]
if b["orders"]["lm_context_sum"]
else "",
}
)
args.output.mkdir(parents=True)
write_json(args.output / "comparison.json", result)
write_json(args.output / "case-outcomes.json", records)
print(
json.dumps(
{
"v1": result["v1"],
"v2": result["v2"],
"paired": result["paired"],
"context": {
n: {
g: r["metrics"][g]["lm_context_sum"]["top1_correct"]
for g in ("with_context", "without_context")
}
for n, r in reports.items()
},
"elapsed_seconds": result["elapsed_seconds"],
},
ensure_ascii=False,
indent=2,
)
)
if __name__ == "__main__":
main()