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()