File size: 2,821 Bytes
33bdf0e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
from __future__ import annotations

import argparse
import json
from pathlib import Path
from typing import Dict


def _norm(text: str) -> str:
    return " ".join(text.strip().lower().split())


def _load_gold(path: Path) -> Dict[str, str]:
    gold: Dict[str, str] = {}
    with path.open("r", encoding="utf-8", errors="ignore") as f:
        for i, line in enumerate(f, 1):
            s = line.strip()
            if not s:
                continue
            try:
                obj = json.loads(s)
            except Exception as exc:  # noqa: BLE001
                raise SystemExit(f"Gold line {i} is not valid JSON: {exc}") from exc
            for key in ("qhash", "gold_answer"):
                if key not in obj:
                    raise SystemExit(f"Gold line {i} missing field: {key}")
            qhash = str(obj["qhash"])
            if qhash in gold:
                raise SystemExit(f"Duplicate qhash in gold at line {i}: {qhash}")
            gold[qhash] = str(obj["gold_answer"])
    return gold


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--gold", required=True, help="Path to gold JSONL")
    ap.add_argument(
        "--predictions",
        required=True,
        help="Path to prediction JSONL (fields: qhash, predicted)",
    )
    args = ap.parse_args()

    gold = _load_gold(Path(args.gold))
    if not gold:
        raise SystemExit("Gold file is empty after filtering blank lines.")

    total = 0
    exact = 0
    missing = 0

    with Path(args.predictions).open("r", encoding="utf-8", errors="ignore") as f:
        for i, line in enumerate(f, 1):
            s = line.strip()
            if not s:
                continue
            try:
                obj = json.loads(s)
            except Exception as exc:  # noqa: BLE001
                raise SystemExit(f"Prediction line {i} is not valid JSON: {exc}") from exc

            if "qhash" not in obj:
                raise SystemExit(f"Prediction line {i} missing field: qhash")
            qhash = str(obj["qhash"])
            if qhash not in gold:
                # Ignore out-of-set predictions so users can score subset files.
                continue

            total += 1
            pred = _norm(str(obj.get("predicted", "")))
            target = _norm(gold[qhash])
            if not pred:
                missing += 1
            if pred == target:
                exact += 1

    if total == 0:
        raise SystemExit("No predictions matched qhash entries in the gold file.")

    print(json.dumps(
        {
            "n_scored": total,
            "exact_match": round(exact / total, 4),
            "n_exact": exact,
            "n_missing_prediction": missing,
        },
        ensure_ascii=False,
    ))


if __name__ == "__main__":
    main()