Add leaderboard reproduction helpers
Browse files
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()
|