Spaces:
Running on Zero
Running on Zero
File size: 9,093 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 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 | """
Error analysis over the Phase 1 per-channel probe results. No new transcription.
Answers the two questions that decide the next phase:
Q1 (Phase 5 / LLM): WHAT are the remaining errors? Categorize normalized
substitutions/deletions/insertions (numbers / morphology / function words /
lexical-content) and surface the most frequent lexical ref->hyp pairs --
i.e. exactly what an LLM cleanup pass would need to fix.
Q2 (Phase 3 / routing): does word CONFIDENCE predict errors? Bucket hypothesis
words by their decoding probability (and segment no_speech_prob) and show the
error rate per bucket. If errors concentrate in low-confidence words, we can
gate any correction pass to those spans (cheaper, less over-correction).
Pass 1 (categorization) runs at the normalized level so number/contraction
formatting is collapsed -> only genuine errors are counted.
Pass 2 (confidence) runs at the raw level so hypothesis tokens map 1:1 onto the
original word objects that carry prob/ns.
"""
import os
import json
import re
import jiwer
from collections import Counter, defaultdict
from eval_common import normalise, normalise_raw
DATA = r"d:\Desktop\ai-ml-capstone\data\na_testset"
RES = "results_channels" # Phase 1
FUNCTION_WORDS = set("a an the is are was were be been being am to of in on at for "
"and or but so it its i you he she we they me him her us them my "
"your his our their this that these those as if then than do does "
"did has have had will would can could may might shall should must "
"not no yes ok okay oh uh um mm yeah".split())
NUMBER_WORDS = set("zero one two three four five six seven eight nine ten eleven twelve "
"thirteen fourteen fifteen sixteen seventeen eighteen nineteen twenty "
"thirty forty fifty sixty seventy eighty ninety hundred thousand million "
"first second third fourth fifth".split())
def is_number(tok):
return bool(re.search(r"\d", tok)) or tok in NUMBER_WORDS
def same_stem(a, b):
"""Crude morphology check: same word, different inflection (plural/tense)."""
if a == b:
return False
for suf in ("s", "es", "d", "ed", "ing", "'s"):
if a == b + suf or b == a + suf:
return True
# shared stem of length >=4
i = 0
while i < min(len(a), len(b)) and a[i] == b[i]:
i += 1
return i >= 4 and abs(len(a) - len(b)) <= 3
def categorize(ref_w, hyp_w):
if is_number(ref_w) or is_number(hyp_w):
return "number"
if same_stem(ref_w, hyp_w):
return "morphology"
if ref_w in FUNCTION_WORDS or hyp_w in FUNCTION_WORDS:
return "function"
return "lexical"
def align_chunks(ref_str, hyp_str):
out = jiwer.process_words(ref_str, hyp_str)
refs, hyps = out.references[0], out.hypotheses[0]
return refs, hyps, out.alignments[0]
def main():
# ββ Pass 1: normalized categorization βββββββββββββββββββββββββββββββββββββ
with open(os.path.join(DATA, "manifest.json"), encoding="utf-8") as f:
manifest = {m["call_id"]: m for m in json.load(f)}
with open(os.path.join(DATA, "probe_set.json"), encoding="utf-8") as f:
probe = json.load(f)["calls"]
cat_counts = Counter()
lexical_pairs = Counter()
del_words = Counter()
ins_words = Counter()
per_call_err = {}
tot_ref = tot_S = tot_D = tot_I = 0
# ββ Pass 2 accumulators (raw level) βββββββββββββββββββββββββββββββββββββββ
prob_buckets = [(0.0, 0.3), (0.3, 0.5), (0.5, 0.7), (0.7, 0.85), (0.85, 1.01)]
ns_buckets = [(0.0, 0.1), (0.1, 0.3), (0.3, 0.6), (0.6, 1.01)]
prob_stat = {b: [0, 0] for b in prob_buckets} # [n_words, n_errors]
ns_stat = {b: [0, 0] for b in ns_buckets}
err_lowprob = err_total = words_lowprob = words_total = 0
for p in probe:
cid, accent = p["call_id"], p["accent"]
m = manifest[cid]
with open(os.path.join(DATA, RES, accent, cid + ".json"), encoding="utf-8") as f:
hyp = json.load(f)
call_err = 0
for ref_text, words in [(m["agent_transcript"], hyp["agent"]),
(m["customer_transcript"], hyp["customer"])]:
# Pass 1: normalized
refs, hyps, chunks = align_chunks(normalise(ref_text),
normalise(" ".join(w["word"] for w in words)))
for ch in chunks:
if ch.type == "equal":
continue
rspan = refs[ch.ref_start_idx:ch.ref_end_idx]
hspan = hyps[ch.hyp_start_idx:ch.hyp_end_idx]
if ch.type == "substitute":
tot_S += len(rspan); call_err += len(rspan)
for rw, hw in zip(rspan, hspan):
c = categorize(rw, hw)
cat_counts[c] += 1
if c == "lexical":
lexical_pairs[f"{rw} -> {hw}"] += 1
elif ch.type == "delete":
tot_D += len(rspan); call_err += len(rspan)
for rw in rspan:
del_words[rw] += 1
cat_counts["number" if is_number(rw) else
"function" if rw in FUNCTION_WORDS else "del_other"] += 1
elif ch.type == "insert":
tot_I += len(hspan); call_err += len(hspan)
for hw in hspan:
ins_words[hw] += 1
tot_ref += len(refs)
# Pass 2: raw-level confidence mapping (1:1 token<->word meta)
toks, meta = [], []
for w in words:
for t in normalise_raw(w["word"]).split():
toks.append(t); meta.append((w.get("prob", 1.0), w.get("ns", 0.0)))
_, _, rchunks = align_chunks(normalise_raw(ref_text), " ".join(toks))
err_flag = [False] * len(toks)
for ch in rchunks:
if ch.type in ("substitute", "insert"):
for hi in range(ch.hyp_start_idx, ch.hyp_end_idx):
if hi < len(err_flag):
err_flag[hi] = True
for (prob, ns), is_err in zip(meta, err_flag):
words_total += 1; err_total += is_err
if prob < 0.5:
words_lowprob += 1; err_lowprob += is_err
for b in prob_buckets:
if b[0] <= prob < b[1]:
prob_stat[b][0] += 1; prob_stat[b][1] += is_err; break
for b in ns_buckets:
if b[0] <= ns < b[1]:
ns_stat[b][0] += 1; ns_stat[b][1] += is_err; break
per_call_err[cid] = call_err
# ββ Report ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
print("=== ERROR ANALYSIS (Phase 1 per-channel, 12 probe calls) ===\n")
acc = (1 - (tot_S + tot_D + tot_I) / tot_ref) * 100
print(f"Normalized: {acc:.1f}% accuracy | {tot_ref} ref words | "
f"S={tot_S} D={tot_D} I={tot_I}\n")
print("Q1 -- Error type breakdown (normalized S+D+I):")
total_cat = sum(cat_counts.values())
for c, n in cat_counts.most_common():
print(f" {c:<12} {n:>4} ({n/total_cat*100:4.1f}%)")
print("\n Top lexical substitutions (ref -> hyp) -- the LLM-fixable content errors:")
for pair, n in lexical_pairs.most_common(20):
print(f" {pair:<32} x{n}")
print("\n Top deletions (ref words missing):", ", ".join(f"{w}({n})" for w, n in del_words.most_common(10)))
print(" Top insertions (extra hyp words): ", ", ".join(f"{w}({n})" for w, n in ins_words.most_common(10)))
print("\nQ2 -- Does confidence predict errors? (raw-level hyp words)")
print(f" {'prob bucket':<14} {'#words':>7} {'#err':>6} {'err-rate':>9}")
for b in prob_buckets:
n, e = prob_stat[b]
print(f" [{b[0]:.2f},{b[1]-0.01:.2f}] {n:>7} {e:>6} {(e/n*100 if n else 0):>8.1f}%")
print(f" {'no_speech':<14} {'#words':>7} {'#err':>6} {'err-rate':>9}")
for b in ns_buckets:
n, e = ns_stat[b]
print(f" ns[{b[0]:.1f},{b[1]-0.01:.2f}] {n:>7} {e:>6} {(e/n*100 if n else 0):>8.1f}%")
if err_total:
print(f"\n Words with prob<0.5 hold {err_lowprob/err_total*100:.0f}% of all errors "
f"while being only {words_lowprob/words_total*100:.0f}% of all words.")
print("\nPer-call error count (normalized S+D+I), worst first:")
for cid, n in sorted(per_call_err.items(), key=lambda x: -x[1]):
tier = next(p["tier"] for p in probe if p["call_id"] == cid)
print(f" {tier:<13} {cid:<32} {n}")
if __name__ == "__main__":
main()
|