erik-at-boson commited on
Commit
db49668
·
verified ·
1 Parent(s): 95e6211

Add leaderboard reproduction helpers

Browse files
Files changed (1) hide show
  1. open_asr_leaderboard/aggregate.py +125 -0
open_asr_leaderboard/aggregate.py ADDED
@@ -0,0 +1,125 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Aggregate per-shard checkpoints into a leaderboard-style summary.
2
+
3
+ Reads every ``CHK_*.json`` snapshot written by run_eval.py under a given
4
+ results directory, concatenates per-dataset predictions/references, computes
5
+ WER per dataset with evaluate.load('wer'), and prints the macro-average WER
6
+ across all 8 datasets so we can gauge progress against the 5.30% target
7
+ without waiting for the run to finish.
8
+ """
9
+
10
+ import argparse
11
+ import glob
12
+ import json
13
+ import os
14
+ import sys
15
+ from collections import defaultdict
16
+
17
+ import evaluate
18
+
19
+
20
+ wer_metric = evaluate.load("wer")
21
+
22
+
23
+ DATASETS = {
24
+ "ami": ["test"],
25
+ "earnings22": ["test"],
26
+ "gigaspeech": ["test"],
27
+ "librispeech": ["test.clean", "test.other"],
28
+ "spgispeech": ["test"],
29
+ "tedlium": ["test"],
30
+ "voxpopuli": ["test"],
31
+ }
32
+
33
+
34
+ def parse_chk_name(fn):
35
+ # CHK_<model>_<dataset>_<split>_shard<i>-<n>.json
36
+ name = os.path.basename(fn)[:-5]
37
+ if not name.startswith("CHK_"):
38
+ return None
39
+ parts = name[4:].rsplit("_shard", 1)
40
+ head = parts[0]
41
+ shard_part = parts[1]
42
+ # head = <model>_<dataset>_<split>
43
+ # pop the split off the right; split may have a dot (test.clean)
44
+ toks = head.split("_")
45
+ # model contains dashes, dataset is single token (ami, earnings22, ...)
46
+ # try matching known datasets from right
47
+ for ds in DATASETS:
48
+ for split in DATASETS[ds]:
49
+ suffix = f"{ds}_{split}"
50
+ if head.endswith(suffix):
51
+ model = head[: -len(suffix) - 1]
52
+ return model, ds, split, shard_part
53
+ return None
54
+
55
+
56
+ def main():
57
+ ap = argparse.ArgumentParser()
58
+ ap.add_argument("--dir", required=True)
59
+ args = ap.parse_args()
60
+
61
+ chks = sorted(glob.glob(os.path.join(args.dir, "CHK_*.json")))
62
+ if not chks:
63
+ print(f"No CHK_*.json under {args.dir}", file=sys.stderr)
64
+ sys.exit(1)
65
+
66
+ # (model, dataset, split) -> merged {predictions, references, audio_len, time}
67
+ buckets = defaultdict(lambda: {"predictions": [], "references": [],
68
+ "audio_length_s": [], "transcription_time_s": []})
69
+ for chk in chks:
70
+ meta = parse_chk_name(chk)
71
+ if not meta:
72
+ print("skip", chk); continue
73
+ model, ds, split, _ = meta
74
+ with open(chk) as f:
75
+ d = json.load(f)
76
+ key = (model, ds, split)
77
+ for k in buckets[key]:
78
+ buckets[key][k].extend(d.get(k, []))
79
+
80
+ # Per-dataset WER (collapse librispeech clean/other into two entries, avg'd).
81
+ print(f"{'dataset':<24}{'n':>8} {'WER%':>7} {'audio_h':>8} {'RTFx':>7}")
82
+ print("-" * 60)
83
+ per_ds = {}
84
+ for (model, ds, split), v in sorted(buckets.items()):
85
+ n = len(v["predictions"])
86
+ if n == 0:
87
+ continue
88
+ try:
89
+ wer = wer_metric.compute(
90
+ references=v["references"], predictions=v["predictions"]
91
+ )
92
+ wer_pct = round(100 * wer, 2)
93
+ except Exception as e:
94
+ wer_pct = float("nan")
95
+ aud_h = sum(v["audio_length_s"]) / 3600.0
96
+ tot_t = sum(v["transcription_time_s"]) or 1.0
97
+ rtfx = sum(v["audio_length_s"]) / tot_t
98
+ label = f"{ds}/{split}"
99
+ print(f"{label:<24}{n:>8} {wer_pct:>7.2f} {aud_h:>8.2f} {rtfx:>7.2f}")
100
+ per_ds[(ds, split)] = wer_pct
101
+
102
+ # Macro-average across 8 leaderboard datasets (ami, e22, gs, ls.c, ls.o, spg, ted, vp)
103
+ lb8 = [
104
+ ("ami", "test"),
105
+ ("earnings22", "test"),
106
+ ("gigaspeech", "test"),
107
+ ("librispeech", "test.clean"),
108
+ ("librispeech", "test.other"),
109
+ ("spgispeech", "test"),
110
+ ("tedlium", "test"),
111
+ ("voxpopuli", "test"),
112
+ ]
113
+ vals = [per_ds.get(k) for k in lb8]
114
+ have = [v for v in vals if v is not None]
115
+ if len(have) == 8:
116
+ avg = round(sum(have) / 8, 2)
117
+ print(f"\nMacro avg WER over 8 datasets: {avg}%")
118
+ else:
119
+ present = [f"{k[0]}/{k[1]}" for k, v in zip(lb8, vals) if v is not None]
120
+ missing = [f"{k[0]}/{k[1]}" for k, v in zip(lb8, vals) if v is None]
121
+ print(f"\nPartial: {len(have)}/8 datasets scored. present={present} missing={missing}")
122
+
123
+
124
+ if __name__ == "__main__":
125
+ main()