Glazkov commited on
Commit
f371181
·
verified ·
1 Parent(s): 0d148f0

Add score_lenient.py

Browse files
Files changed (1) hide show
  1. score_lenient.py +117 -0
score_lenient.py ADDED
@@ -0,0 +1,117 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Lenient tuple-F1 scoring with unit aliases and date-year normalization.
2
+
3
+ Replicates the offline `+0.020 t_f1` lift reported on the eval set. Use to
4
+ re-score a predictions JSONL against a reference annotations JSONL.
5
+
6
+ Usage::
7
+
8
+ python score_lenient.py preds.jsonl annotations_test.jsonl
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import json
14
+ import re
15
+ import sys
16
+ from pathlib import Path
17
+
18
+
19
+ UNIT_ALIASES = {
20
+ # English plural ↔ singular
21
+ "millions": "million",
22
+ "millions dollars": "million dollars",
23
+ "millions of dollars": "million dollars",
24
+ "thousands": "thousand",
25
+ "thousands dollars": "thousand dollars",
26
+ "billions": "billion",
27
+ "billions dollars": "billion dollars",
28
+ "percent": "%",
29
+ # Russian abbreviations
30
+ "млн. руб.": "млн руб",
31
+ "млн руб.": "млн руб",
32
+ "млн.руб.": "млн руб",
33
+ "млн руб ": "млн руб",
34
+ "тыс. руб.": "тыс руб",
35
+ "тыс руб.": "тыс руб",
36
+ "тыс.руб.": "тыс руб",
37
+ "млрд. руб.": "млрд руб",
38
+ "млрд руб.": "млрд руб",
39
+ }
40
+
41
+ _YEAR_RE = re.compile(r"\b(19|20)\d{2}\b")
42
+
43
+
44
+ def normalize_unit(u: str) -> str:
45
+ s = (u or "").strip().lower().rstrip(".")
46
+ return UNIT_ALIASES.get(s, s)
47
+
48
+
49
+ def normalize_date(d: str) -> str:
50
+ """Return the first 4-digit year in the string, else lowercased stripped."""
51
+ m = _YEAR_RE.search(d or "")
52
+ if m:
53
+ return m.group(0)
54
+ return (d or "").strip().lower()
55
+
56
+
57
+ def lenient_key(p: dict) -> tuple[str, str, str, str]:
58
+ return (
59
+ (p.get("parameter_name") or "").strip().lower(),
60
+ (p.get("parameter_value") or "").strip(),
61
+ normalize_date(p.get("parameter_date", "")),
62
+ normalize_unit(p.get("parameter_unit", "")),
63
+ )
64
+
65
+
66
+ def strict_key(p: dict) -> tuple[str, str, str, str]:
67
+ return (
68
+ (p.get("parameter_name") or ""),
69
+ (p.get("parameter_value") or ""),
70
+ (p.get("parameter_date") or ""),
71
+ (p.get("parameter_unit") or ""),
72
+ )
73
+
74
+
75
+ def f1(pred_set: set, ref_set: set) -> tuple[float, float, float]:
76
+ if not pred_set or not ref_set:
77
+ return 0.0, 0.0, 0.0
78
+ overlap = pred_set & ref_set
79
+ prec = len(overlap) / max(1, len(pred_set))
80
+ rec = len(overlap) / max(1, len(ref_set))
81
+ return (
82
+ (2 * prec * rec / (prec + rec)) if (prec + rec) > 0 else 0.0,
83
+ prec,
84
+ rec,
85
+ )
86
+
87
+
88
+ def load_jsonl(path: Path) -> list[dict]:
89
+ return [json.loads(line) for line in path.open("r", encoding="utf-8")]
90
+
91
+
92
+ def main() -> None:
93
+ if len(sys.argv) != 3:
94
+ sys.exit(f"usage: {sys.argv[0]} preds.jsonl annotations_test.jsonl")
95
+
96
+ preds = {r["image"]: r.get("parameters", []) for r in load_jsonl(Path(sys.argv[1]))}
97
+ refs = {r["image"]: r.get("parameters", []) for r in load_jsonl(Path(sys.argv[2]))}
98
+
99
+ strict_f1s, strict_ps, strict_rs = [], [], []
100
+ lenient_f1s, lenient_ps, lenient_rs = [], [], []
101
+ matched_keys = sorted(set(preds) & set(refs))
102
+ for k in matched_keys:
103
+ ref = refs[k]
104
+ pred = preds[k]
105
+ s_f1, s_p, s_r = f1({strict_key(p) for p in pred}, {strict_key(p) for p in ref})
106
+ l_f1, l_p, l_r = f1({lenient_key(p) for p in pred}, {lenient_key(p) for p in ref})
107
+ strict_f1s.append(s_f1); strict_ps.append(s_p); strict_rs.append(s_r)
108
+ lenient_f1s.append(l_f1); lenient_ps.append(l_p); lenient_rs.append(l_r)
109
+
110
+ n = max(1, len(matched_keys))
111
+ print(f"Scored {len(matched_keys)} matched samples.\n")
112
+ print(f"STRICT tuple_f1 = {sum(strict_f1s)/n:.4f} precision = {sum(strict_ps)/n:.4f} recall = {sum(strict_rs)/n:.4f}")
113
+ print(f"LENIENT tuple_f1 = {sum(lenient_f1s)/n:.4f} precision = {sum(lenient_ps)/n:.4f} recall = {sum(lenient_rs)/n:.4f}")
114
+
115
+
116
+ if __name__ == "__main__":
117
+ main()