Spaces:
Running on Zero
Running on Zero
File size: 5,302 Bytes
f1ef7e2 | 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 132 133 134 135 136 137 138 | """
A/B comparison on the frozen probe set. Compares two result dirs without
re-transcribing. Default: per-channel (A) vs per-channel+preprocess (B).
The 'base' column is the static mono-mix baseline_acc recorded in probe_set.json.
Also reports:
- no_speech_prob diagnostics per dir (mean, #high-ns segments, total words) to
show whether preprocessing reduced Whisper's silence-hallucination signal
- a confidence-filter experiment on B (drop words with segment ns > threshold)
Usage:
python compare_probe.py # chan vs chan+pp
python compare_probe.py --a-dir results --a-label mix --b-dir results_channels --b-label chan
"""
import os
import json
import argparse
import jiwer
from eval_common import normalise
DATA_DIR = r"d:\Desktop\ai-ml-capstone\data\na_testset"
MANIFEST = os.path.join(DATA_DIR, "manifest.json")
PROBE_SET = os.path.join(DATA_DIR, "probe_set.json")
def load_result(root, accent, cid):
p = os.path.join(DATA_DIR, root, accent, cid + ".json")
if not os.path.exists(p):
return None
with open(p, encoding="utf-8") as f:
return json.load(f)
def text_of(words, ns_max=None):
if ns_max is None:
return " ".join(w["word"] for w in words)
return " ".join(w["word"] for w in words if w.get("ns", 0.0) <= ns_max)
def score_call(ref_a, ref_c, hyp_a, hyp_c):
ra, rc = normalise(ref_a), normalise(ref_c)
ha, hc = normalise(hyp_a), normalise(hyp_c)
na, nc = len(ra.split()), len(rc.split())
if na + nc == 0:
return None, 0
wa = jiwer.wer(ra, ha) if na else 0.0
wc = jiwer.wer(rc, hc) if nc else 0.0
return (wa * na + wc * nc) / (na + nc), na + nc
def acc(wer):
return None if wer is None else (1 - wer) * 100
def ns_diag(result):
"""(#high-ns segments, total segments, total hyp words) from captured segments."""
if not result or "segments" not in result:
return None
segs = result["segments"].get("agent", []) + result["segments"].get("customer", [])
hi = sum(1 for s in segs if s.get("no_speech_prob", 0) > 0.6)
words = len(result.get("agent", [])) + len(result.get("customer", []))
return hi, len(segs), words
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--a-dir", default="results_channels"); ap.add_argument("--a-label", default="chan")
ap.add_argument("--b-dir", default="results_channels_pp"); ap.add_argument("--b-label", default="chan+pp")
ap.add_argument("--ns-thresholds", nargs="*", type=float, default=[0.6, 0.8])
args = ap.parse_args()
with open(MANIFEST, encoding="utf-8") as f:
manifest = {m["call_id"]: m for m in json.load(f)}
with open(PROBE_SET, encoding="utf-8") as f:
probe = json.load(f)["calls"]
print(f"{'tier':<13} {'call_id':<32} {'base':>6} {args.a_label:>7} {args.b_label:>8} {'delta':>7}")
print("-" * 80)
agg = {"a": [0.0, 0], "b": [0.0, 0]}
agg_ns = {t: [0.0, 0] for t in args.ns_thresholds}
nsA = [0, 0, 0]; nsB = [0, 0, 0] # hi, segs, words
missing = []
for p in probe:
cid, accent = p["call_id"], p["accent"]
m = manifest[cid]
ra_, rc_ = m["agent_transcript"], m["customer_transcript"]
A = load_result(args.a_dir, accent, cid)
B = load_result(args.b_dir, accent, cid)
if A is None or B is None:
missing.append((cid, args.a_dir if A is None else args.b_dir))
continue
a_wer, n = score_call(ra_, rc_, text_of(A["agent"]), text_of(A["customer"]))
b_wer, _ = score_call(ra_, rc_, text_of(B["agent"]), text_of(B["customer"]))
agg["a"][0] += a_wer * n; agg["a"][1] += n
agg["b"][0] += b_wer * n; agg["b"][1] += n
for t in args.ns_thresholds:
w, _ = score_call(ra_, rc_, text_of(B["agent"], t), text_of(B["customer"], t))
agg_ns[t][0] += w * n; agg_ns[t][1] += n
for store, res in [(nsA, A), (nsB, B)]:
d = ns_diag(res)
if d:
store[0] += d[0]; store[1] += d[1]; store[2] += d[2]
delta = acc(b_wer) - acc(a_wer)
flag = " <<" if delta >= 3 else (" !!" if delta <= -3 else "")
print(f"{p['tier']:<13} {cid:<32} {p['baseline_acc']*100:>5.1f}% "
f"{acc(a_wer):>6.1f}% {acc(b_wer):>7.1f}% {delta:>+6.1f}{flag}")
print("-" * 80)
a_acc = acc(agg["a"][0] / agg["a"][1])
b_acc = acc(agg["b"][0] / agg["b"][1])
print(f"{'OVERALL (word-weighted)':<46} {a_acc:>6.1f}% {b_acc:>7.1f}% {b_acc - a_acc:>+6.1f}")
print(f"\nno_speech_prob diagnostics (hallucination signal):")
print(f" {args.a_label:<10} high-ns segs: {nsA[0]:>4} / {nsA[1]:<5} segments | total hyp words: {nsA[2]}")
print(f" {args.b_label:<10} high-ns segs: {nsB[0]:>4} / {nsB[1]:<5} segments | total hyp words: {nsB[2]}")
print(f"\nConfidence filter on {args.b_label} (drop words where segment no_speech_prob > t):")
print(f" no filter : {b_acc:5.1f}%")
for t in args.ns_thresholds:
a = acc(agg_ns[t][0] / agg_ns[t][1])
print(f" ns <= {t:<4} : {a:5.1f}% ({a - b_acc:+.1f})")
if missing:
print("\nMissing results:")
for cid, which in missing:
print(f" {cid} ({which})")
if __name__ == "__main__":
main()
|