Spaces:
Running on Zero
Running on Zero
File size: 5,955 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 | """
VAD-recall experiment: does a more sensitive VAD (or no VAD) recover the dropped
speech segments we saw in error analysis, WITHOUT adding hallucination?
Runs a named decode config on a small mini-set with small.en (fast), writes to
results_vad_<config>/, and prints per-call accuracy vs the frozen small.en baseline
(results_channels). Includes a good-banking call as a REGRESSION GUARD -- a config
that helps hard calls but hurts easy ones is no good.
Usage:
python run_vad_exp.py --config sensitive
python run_vad_exp.py --config novad
python run_vad_exp.py --config midvad --full # all 12 probe calls
"""
import os, json, time, argparse, jiwer
from faster_whisper import WhisperModel
import transcribe_channels
from eval_common import normalise
DATA = r"d:\Desktop\ai-ml-capstone\data\na_testset"
MANIFEST = os.path.join(DATA, "manifest.json")
PROBE_SET = os.path.join(DATA, "probe_set.json")
BASELINE = "results_channels" # small.en Phase-1 baseline
# decode configs (merged onto the Phase-1 defaults inside transcribe_channels)
CONFIGS = {
"baseline": {}, # sanity: should reproduce results_channels
"sensitive": {"vad_parameters": {"threshold": 0.3, "min_silence_duration_ms": 200,
"speech_pad_ms": 600}},
"midvad": {"vad_parameters": {"threshold": 0.4, "min_silence_duration_ms": 300,
"speech_pad_ms": 500}},
"novad": {"vad_filter": False},
# no-VAD but lean on Whisper's own silence suppression to limit hallucination
"novad_strict": {"vad_filter": False, "no_speech_threshold": 0.5,
"condition_on_previous_text": False},
}
# mini-set: 2 filler-heavy catastrophic + 2 bad-banking + 1 good-banking (guard)
MINISET = [
"en_US_General_Health_1587175", # catastrophic, lots of dropped filler/segments
"en_CA_Aviation_1588678", # catastrophic, API regressed here
"en_CA_Banking_1588683", # bad_banking
"en_US_General_Banking_1584540", # bad_banking, big API gain
"en_CA_Banking_1592237", # good_banking -- REGRESSION GUARD (94.8%)
]
def acc(ref, hyp):
r, h = normalise(ref), normalise(hyp)
n = len(r.split())
return (1 - jiwer.wer(r, h)) * 100 if n else 100.0, n
def overall(m, ag, cu):
aa, na = acc(m["agent_transcript"], ag)
ac, nc = acc(m["customer_transcript"], cu)
return (aa*na + ac*nc) / (na+nc), na+nc
def baseline_acc(m, accent, cid):
with open(os.path.join(DATA, BASELINE, accent, cid + ".json"), encoding="utf-8") as f:
b = json.load(f)
return overall(m, " ".join(w["word"] for w in b["agent"]),
" ".join(w["word"] for w in b["customer"]))[0]
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--config", required=True, choices=list(CONFIGS))
ap.add_argument("--model", default="small.en")
ap.add_argument("--full", action="store_true", help="all 12 probe calls, not mini-set")
args = ap.parse_args()
with open(MANIFEST, encoding="utf-8") as f:
manifest = {m["call_id"]: m for m in json.load(f)}
if args.full:
with open(PROBE_SET, encoding="utf-8") as f:
ids = [c["call_id"] for c in json.load(f)["calls"]]
else:
ids = MINISET
decode = CONFIGS[args.config]
model_tag = args.model.replace(".", "")
# keep small.en config-only dirs (backward compat with the first mini-set runs);
# tag non-small models so medium/large results never collide with small ones.
out_dir = (f"results_vad_{args.config}" if args.model == "small.en"
else f"results_vad_{args.config}_{model_tag}")
print(f"VAD experiment | config={args.config} | model={args.model} | {len(ids)} calls")
print(f" decode overrides: {decode}\n")
model = WhisperModel(args.model, device="cpu", compute_type="int8")
print(f" {'call_id':<32} {'base':>6} {'new':>6} {'delta':>6} {'A_wd':>5} {'C_wd':>5}")
print(" " + "-" * 68)
deltas, ws, news = [], [], []
t_start = time.time()
for cid in ids:
m = manifest[cid]
accent = m["accent"]
a_wav = os.path.join(DATA, m["agent_wav"])
c_wav = os.path.join(DATA, m["customer_wav"])
out = os.path.join(DATA, out_dir, accent, cid + ".json")
os.makedirs(os.path.dirname(out), exist_ok=True)
if os.path.exists(out) and os.path.getsize(out) > 0:
with open(out, encoding="utf-8") as f:
r = json.load(f)
aw, cw = r["agent"], r["customer"]
else:
aw, cw, segs, dur = transcribe_channels.transcribe_call(
model, a_wav, c_wav, decode=decode)
with open(out, "w", encoding="utf-8") as f:
json.dump({"call_id": cid, "accent": accent, "domain": m["domain"],
"config": args.config, "agent": aw, "customer": cw}, f, indent=2)
new_o, w = overall(m, " ".join(x["word"] for x in aw),
" ".join(x["word"] for x in cw))
base_o = baseline_acc(m, accent, cid)
d = new_o - base_o
deltas.append(d * w); ws.append(w); news.append(new_o * w)
print(f" {cid:<32} {base_o:>5.1f}% {new_o:>5.1f}% {d:>+5.1f} {len(aw):>5} {len(cw):>5}")
print(" " + "-" * 68)
wavg = sum(deltas) / sum(ws)
overall_abs = sum(news) / sum(ws)
print(f" OVERALL absolute accuracy: {overall_abs:.2f}% "
f"(word-weighted mean delta vs small.en baseline: {wavg:+.2f})")
print(f" {'>> HELPS' if wavg > 0.1 else '>> NEUTRAL' if abs(wavg) <= 0.1 else '>> HURTS'}")
if args.full:
tag = ("CLEARS +cushion" if overall_abs >= 93.5 else
"clears (thin)" if overall_abs >= 93.0 else "short of 93%")
print(f" GOAL (>=93% over 12 calls): {overall_abs:.2f}% -> {tag}")
print(f" ({(time.time()-t_start)/60:.1f} min)")
if __name__ == "__main__":
main()
|