File size: 5,039 Bytes
29f25be
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
"""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()