Spaces:
Running on Zero
Running on Zero
File size: 5,616 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 | """
Find the optimal loudness-gate threshold WITHOUT any new transcription run.
Key idea: the gated pipeline's output for each channel is already known:
- channel active-RMS >= gate -> identical to Phase 1 (results_channels/, no PP)
- channel active-RMS < gate -> identical to Phase 2 (results_channels_pp/, full PP)
(high-pass is near-neutral on this data, so "above gate" ~ Phase 1 is exact enough.)
So we assemble each call per-channel from the two existing result sets according
to a candidate gate, score it, and sweep the gate to find the threshold that
keeps the quiet-channel gains while removing the over-amplification regressions.
This is pure analysis over cached results -> runs in seconds, no model, no CPU run.
Usage:
python simulate_gate.py
python simulate_gate.py --gate -26 # per-call detail at one gate
"""
import os
import json
import argparse
import numpy as np
import soundfile as sf
import jiwer
from eval_common import normalise
from audio_preprocess import highpass, active_rms
DATA = r"d:\Desktop\ai-ml-capstone\data\na_testset"
P1 = "results_channels" # no preprocessing (Phase 1)
P2 = "results_channels_pp" # full preprocessing (Phase 2, target -20 dBFS)
SR = 16000
def load_result(root, accent, cid):
with open(os.path.join(DATA, root, accent, cid + ".json"), encoding="utf-8") as f:
return json.load(f)
def channel_rms_db(wav_rel):
a, sr = sf.read(os.path.join(DATA, wav_rel), dtype="float32")
if a.ndim > 1:
a = a.mean(axis=1)
a = highpass(a, sr) # match preprocess(): measure post-highpass
return 20 * np.log10(active_rms(a, sr) + 1e-12)
def wer(ref, hyp):
r, h = normalise(ref), normalise(hyp)
n = len(r.split())
if n == 0:
return 0.0, 0
return jiwer.wer(r, h), n
def text(words):
return " ".join(w["word"] for w in words)
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--gate", type=float, default=None, help="show per-call detail at this gate")
ap.add_argument("--sweep", nargs="*", type=float,
default=[-99, -34, -32, -30, -28, -26, -24, -22, -20, 0])
args = ap.parse_args()
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"]
# Pre-load everything once: refs, both result sets, per-channel RMS.
calls = []
for p in probe:
cid, accent = p["call_id"], p["accent"]
m = manifest[cid]
calls.append({
"cid": cid, "tier": p["tier"],
"ref_a": m["agent_transcript"], "ref_c": m["customer_transcript"],
"rms_a": channel_rms_db(m["agent_wav"]),
"rms_c": channel_rms_db(m["customer_wav"]),
"p1_a": text(load_result(P1, accent, cid)["agent"]),
"p1_c": text(load_result(P1, accent, cid)["customer"]),
"p2_a": text(load_result(P2, accent, cid)["agent"]),
"p2_c": text(load_result(P2, accent, cid)["customer"]),
})
def score_at(gate):
"""Word-weighted accuracy across probe at this gate, + per-call accs + #boosted."""
num = den = 0
per_call = {}
boosted = 0
for c in calls:
ha = c["p2_a"] if c["rms_a"] < gate else c["p1_a"]
hc = c["p2_c"] if c["rms_c"] < gate else c["p1_c"]
boosted += (c["rms_a"] < gate) + (c["rms_c"] < gate)
wa, na = wer(c["ref_a"], ha)
wc, nc = wer(c["ref_c"], hc)
cw = (wa * na + wc * nc) / (na + nc)
per_call[c["cid"]] = (1 - cw) * 100
num += wa * na + wc * nc
den += na + nc
return (1 - num / den) * 100, per_call, boosted
# ββ Gate sweep ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
print("Gate sweep (accuracy assembled from existing P1/P2 results, no re-run):\n")
print(f" {'gate':>6} {'boosted':>7} {'overall':>8} {'Bank_1586157':>12} {'Health_1587175':>14}")
print(" " + "-" * 56)
reg_cid, quiet_cid = "en_US_General_Banking_1586157", "en_US_General_Health_1587175"
for g in args.sweep:
acc, per_call, boosted = score_at(g)
label = {-99: "none", 0: "all"}.get(g, f"{g:.0f}")
print(f" {label:>6} {boosted:>7} {acc:>7.1f}% {per_call[reg_cid]:>11.1f}% {per_call[quiet_cid]:>13.1f}%")
print("\n ('none' = pure Phase 1 / no PP; 'all' = pure Phase 2 / full PP)")
print(" Bank_1586157 = the regression call; Health_1587175 = a quiet-channel call")
# ββ Per-call detail at one gate βββββββββββββββββββββββββββββββββββββββββββ
if args.gate is not None:
acc, per_call, boosted = score_at(args.gate)
print(f"\nPer-call at gate = {args.gate} dBFS (overall {acc:.1f}%, {boosted} channels boosted):\n")
print(f" {'tier':<13} {'call_id':<32} {'rmsA':>6} {'rmsC':>6} {'boost':>9} {'acc':>6}")
print(" " + "-" * 74)
for c in calls:
ba = "A" if c["rms_a"] < args.gate else "-"
bc = "C" if c["rms_c"] < args.gate else "-"
print(f" {c['tier']:<13} {c['cid']:<32} {c['rms_a']:>6.1f} {c['rms_c']:>6.1f} "
f"{ba+bc:>9} {per_call[c['cid']]:>5.1f}%")
if __name__ == "__main__":
main()
|