zeechimp commited on
Commit
15b9f68
·
verified ·
1 Parent(s): 5cfcb74

Upload whisper_decoder.py

Browse files
Files changed (1) hide show
  1. whisper_decoder.py +811 -0
whisper_decoder.py ADDED
@@ -0,0 +1,811 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ whisper_decoder.py
4
+ ==================
5
+
6
+ Closed-vocabulary word classification from whispered speech.
7
+
8
+ Fix history
9
+ -----------
10
+ v1 Three bugs.
11
+
12
+ (a) Training data was generated word-by-word in alphabetical
13
+ order. The validation split took the last 20% of samples,
14
+ which was only the last two words ('two' and 'zero'). The
15
+ reported validation accuracy of 0.0% was a split artifact,
16
+ not a training failure.
17
+
18
+ (b) The synthesized words are shorter (0.2-0.5 s) than the
19
+ fixed clip duration (1.2 s). Features were averaged over
20
+ the whole clip, so zero-padding dominated the spectrum.
21
+ Word identity was swamped by silence.
22
+
23
+ (c) The low-band feature (200-1000 Hz) misses F2 for front
24
+ vowels. 'two' and 'three' were nearly identical in this
25
+ feature.
26
+
27
+ v2 Fixes.
28
+
29
+ (a) Shuffle before splitting. Stratified validation.
30
+ (b) Trim to non-silent content. Use trimmed duration as a
31
+ feature. Report the trimmed length.
32
+ (c) Extend the low band to 200-2500 Hz.
33
+ (d) Lower the voiced harmonicity threshold to 0.25.
34
+ (e) Add reverb augmentation during training.
35
+
36
+ Vocabulary: 10 English digits. Code-only, no downloads.
37
+ """
38
+
39
+ from __future__ import annotations
40
+
41
+ import argparse
42
+ import math
43
+ import os
44
+ import time
45
+ import wave
46
+ from dataclasses import dataclass, field
47
+ from pathlib import Path
48
+ from typing import Dict, List, Optional, Tuple
49
+
50
+ import numpy as np
51
+
52
+ try:
53
+ import matplotlib
54
+ matplotlib.use("Agg")
55
+ import matplotlib.pyplot as plt
56
+ HAS_MPL = True
57
+ except ImportError:
58
+ HAS_MPL = False
59
+
60
+
61
+ SR = 16000
62
+ DURATION_S = 1.2
63
+ N_SAMPLES = int(SR * DURATION_S)
64
+
65
+ FRAME_LEN = 512
66
+ HOP = 160
67
+ N_MELS = 24
68
+ N_MFCC = 13
69
+ N_MELS_LOW = 10
70
+ LOW_FMIN = 200.0
71
+ LOW_FMAX = 2500.0 # was 1000 in v1
72
+
73
+ SEED = 0
74
+ N_TRAIN_PER_WORD = 200
75
+ N_TEST_PER_WORD = 40
76
+ REVERB_TRAIN_FRAC = 0.25 # 25% of training samples get reverb
77
+
78
+
79
+ # =====================================================================
80
+ # §1 Phoneme inventory (unchanged)
81
+ # =====================================================================
82
+
83
+ PHONEMES: Dict[str, Dict] = {
84
+ 'i': {'type': 'vowel', 'dur': 0.12,
85
+ 'formants': [(270, 80, 1.0), (2290, 120, 0.8), (3010, 150, 0.5)]},
86
+ 'ɪ': {'type': 'vowel', 'dur': 0.10,
87
+ 'formants': [(390, 80, 1.0), (1990, 120, 0.8), (2550, 150, 0.5)]},
88
+ 'e': {'type': 'vowel', 'dur': 0.12,
89
+ 'formants': [(530, 80, 1.0), (1840, 120, 0.8), (2480, 150, 0.5)]},
90
+ 'ɛ': {'type': 'vowel', 'dur': 0.11,
91
+ 'formants': [(660, 80, 1.0), (1720, 120, 0.8), (2410, 150, 0.5)]},
92
+ 'a': {'type': 'vowel', 'dur': 0.13,
93
+ 'formants': [(730, 80, 1.0), (1090, 120, 0.8), (2440, 150, 0.5)]},
94
+ 'ɑ': {'type': 'vowel', 'dur': 0.13,
95
+ 'formants': [(730, 80, 1.0), (1090, 120, 0.8), (2440, 150, 0.5)]},
96
+ 'ɔ': {'type': 'vowel', 'dur': 0.13,
97
+ 'formants': [(570, 80, 1.0), (840, 120, 0.8), (2410, 150, 0.5)]},
98
+ 'o': {'type': 'vowel', 'dur': 0.13,
99
+ 'formants': [(570, 80, 1.0), (840, 120, 0.8), (2410, 150, 0.5)]},
100
+ 'ʊ': {'type': 'vowel', 'dur': 0.11,
101
+ 'formants': [(440, 80, 1.0), (1020, 120, 0.8), (2240, 150, 0.5)]},
102
+ 'u': {'type': 'vowel', 'dur': 0.12,
103
+ 'formants': [(300, 80, 1.0), (870, 120, 0.8), (2240, 150, 0.5)]},
104
+ 'ʌ': {'type': 'vowel', 'dur': 0.11,
105
+ 'formants': [(640, 80, 1.0), (1190, 120, 0.8), (2390, 150, 0.5)]},
106
+ 's': {'type': 'fricative', 'dur': 0.11,
107
+ 'formants': [(6000, 1000, 1.0)]},
108
+ 'z': {'type': 'fricative', 'dur': 0.10,
109
+ 'formants': [(6000, 1000, 0.7)]},
110
+ 'ʃ': {'type': 'fricative', 'dur': 0.11,
111
+ 'formants': [(3500, 800, 1.0)]},
112
+ 'f': {'type': 'fricative', 'dur': 0.10,
113
+ 'formants': [(5000, 2000, 0.6)]},
114
+ 'v': {'type': 'fricative', 'dur': 0.08,
115
+ 'formants': [(5000, 2000, 0.7)]},
116
+ 'θ': {'type': 'fricative', 'dur': 0.10,
117
+ 'formants': [(5000, 1500, 0.5)]},
118
+ 'p': {'type': 'plosive', 'dur': 0.08, 'silence': 0.06,
119
+ 'formants': [(1000, 400, 1.0)]},
120
+ 't': {'type': 'plosive', 'dur': 0.08, 'silence': 0.06,
121
+ 'formants': [(4000, 800, 1.0)]},
122
+ 'k': {'type': 'plosive', 'dur': 0.08, 'silence': 0.06,
123
+ 'formants': [(2500, 600, 1.0)]},
124
+ 'n': {'type': 'nasal', 'dur': 0.09,
125
+ 'formants': [(300, 100, 1.0), (1500, 200, 0.4), (2500, 250, 0.2)]},
126
+ 'm': {'type': 'nasal', 'dur': 0.09,
127
+ 'formants': [(300, 100, 1.0), (1000, 200, 0.4), (2200, 250, 0.2)]},
128
+ 'r': {'type': 'approx', 'dur': 0.07,
129
+ 'formants': [(300, 100, 1.0), (1100, 200, 0.6), (1600, 200, 0.3)]},
130
+ 'w': {'type': 'approx', 'dur': 0.07,
131
+ 'formants': [(300, 100, 1.0), (600, 200, 0.5), (2200, 250, 0.2)]},
132
+ }
133
+
134
+ WORDS: Dict[str, List[str]] = {
135
+ 'zero': ['z', 'i', 'r', 'o'],
136
+ 'one': ['w', 'ʌ', 'n'],
137
+ 'two': ['t', 'u'],
138
+ 'three': ['θ', 'r', 'i'],
139
+ 'four': ['f', 'o', 'r'],
140
+ 'five': ['f', 'a', 'i', 'v'],
141
+ 'six': ['s', 'ɪ', 'k', 's'],
142
+ 'seven': ['s', 'ɛ', 'v', 'ɛ', 'n'],
143
+ 'eight': ['e', 'i', 't'],
144
+ 'nine': ['n', 'a', 'i', 'n'],
145
+ }
146
+ WORD_ORDER = sorted(WORDS.keys())
147
+
148
+
149
+ # =====================================================================
150
+ # §2 Formant filter
151
+ # =====================================================================
152
+
153
+ _FILTER_CACHE: Dict = {}
154
+
155
+
156
+ def formant_filter(formants, n_fft, sr):
157
+ key = (tuple(formants), n_fft, sr)
158
+ if key in _FILTER_CACHE:
159
+ return _FILTER_CACHE[key]
160
+ freqs = np.fft.rfftfreq(n_fft, d=1.0 / sr)
161
+ resp = np.zeros_like(freqs)
162
+ for f0, bw, amp in formants:
163
+ resp += amp * np.exp(-((freqs - f0) ** 2) / (2.0 * bw * bw))
164
+ _FILTER_CACHE[key] = resp
165
+ return resp
166
+
167
+
168
+ def _next_pow2(n):
169
+ p = 1
170
+ while p < n:
171
+ p *= 2
172
+ return p
173
+
174
+
175
+ # =====================================================================
176
+ # §3 Phoneme and word synthesis
177
+ # =====================================================================
178
+
179
+ def synth_phoneme_whisper(name, sr, seed, jitter=0.03):
180
+ p = PHONEMES[name]
181
+ rng = np.random.default_rng(seed)
182
+ dur = p['dur'] * (1.0 + jitter * rng.uniform(-1, 1))
183
+ n = int(sr * dur)
184
+
185
+ if p['type'] == 'plosive':
186
+ n_sil = int(sr * p.get('silence', 0.06))
187
+ n_burst = max(4, n - n_sil)
188
+ n_fft = _next_pow2(n_burst)
189
+ noise = rng.standard_normal(n_burst)
190
+ H = formant_filter(p['formants'], n_fft, sr)
191
+ X = np.fft.rfft(noise, n=n_fft)
192
+ burst = np.fft.irfft(X * H, n=n_fft)[:n_burst]
193
+ env = np.exp(-np.arange(n_burst) / (0.005 * sr))
194
+ out = np.zeros(n)
195
+ out[n_sil:] = burst * env
196
+ return out.astype(np.float32)
197
+
198
+ n_fft = _next_pow2(n)
199
+ noise = rng.standard_normal(n)
200
+ H = formant_filter(p['formants'], n_fft, sr)
201
+ X = np.fft.rfft(noise, n=n_fft)
202
+ sig = np.fft.irfft(X * H, n=n_fft)[:n]
203
+ ramp = min(int(0.015 * sr), n // 4)
204
+ env = np.ones(n)
205
+ if ramp > 0:
206
+ up = 0.5 - 0.5 * np.cos(np.arange(ramp) * np.pi / ramp)
207
+ env[:ramp] = up
208
+ env[-ramp:] = up[::-1]
209
+ return (sig * env).astype(np.float32)
210
+
211
+
212
+ def synth_phoneme_voiced(name, sr, seed, f0=120.0, jitter=0.03):
213
+ p = PHONEMES[name]
214
+ rng = np.random.default_rng(seed)
215
+ dur = p['dur'] * (1.0 + jitter * rng.uniform(-1, 1))
216
+ n = int(sr * dur)
217
+
218
+ if p['type'] == 'plosive':
219
+ return synth_phoneme_whisper(name, sr, seed, jitter)
220
+
221
+ n_fft = _next_pow2(n)
222
+ period = max(2, int(sr / f0))
223
+ pulse = np.zeros(n)
224
+ for i in range(0, n, period):
225
+ pulse[i] = 1.0
226
+ H = formant_filter(p['formants'], n_fft, sr)
227
+ X = np.fft.rfft(pulse, n=n_fft)
228
+ sig = np.fft.irfft(X * H, n=n_fft)[:n]
229
+ peak = float(np.max(np.abs(sig))) + 1e-12
230
+ sig = sig / peak * 0.5
231
+ ramp = min(int(0.015 * sr), n // 4)
232
+ env = np.ones(n)
233
+ if ramp > 0:
234
+ up = 0.5 - 0.5 * np.cos(np.arange(ramp) * np.pi / ramp)
235
+ env[:ramp] = up
236
+ env[-ramp:] = up[::-1]
237
+ return (sig * env).astype(np.float32)
238
+
239
+
240
+ def synth_word(word, mode, seed, noise_frac=0.002):
241
+ rng = np.random.default_rng(seed)
242
+ pieces = []
243
+ for ph in WORDS[word]:
244
+ s = int(rng.integers(1 << 30))
245
+ if mode == 'whisper':
246
+ pieces.append(synth_phoneme_whisper(ph, SR, s))
247
+ else:
248
+ f0 = 120.0 + 30.0 * rng.uniform(-1, 1)
249
+ pieces.append(synth_phoneme_voiced(ph, SR, s, f0=f0))
250
+ sig = np.concatenate(pieces) if pieces else np.zeros(1)
251
+ if len(sig) >= N_SAMPLES:
252
+ sig = sig[:N_SAMPLES]
253
+ else:
254
+ sig = np.concatenate([sig, np.zeros(N_SAMPLES - len(sig))])
255
+ if noise_frac > 0:
256
+ sig = sig + noise_frac * rng.standard_normal(len(sig))
257
+ return sig.astype(np.float32)
258
+
259
+
260
+ def apply_reverb(sig, rt60_s=0.4, seed=0):
261
+ rng = np.random.default_rng(seed)
262
+ n_ir = int(SR * rt60_s * 1.5)
263
+ decay = np.exp(-3.0 * np.arange(n_ir) / (SR * rt60_s / 6.91))
264
+ ir = rng.standard_normal(n_ir) * decay
265
+ ir[0] = 1.0
266
+ out = np.convolve(sig, ir, mode="same")
267
+ peak = float(np.max(np.abs(out))) + 1e-12
268
+ return (out / peak * 0.9).astype(np.float32)
269
+
270
+
271
+ # =====================================================================
272
+ # §4 Feature extraction -- v2 TRIM + extended low band
273
+ # =====================================================================
274
+
275
+ _MEL_FB_CACHE: Dict = {}
276
+
277
+
278
+ def frame_signal(sig, frame_len, hop):
279
+ if len(sig) < frame_len:
280
+ sig = np.pad(sig, (0, frame_len - len(sig)))
281
+ n = 1 + (len(sig) - frame_len) // hop
282
+ return np.stack([sig[i * hop:i * hop + frame_len]
283
+ for i in range(n)])
284
+
285
+
286
+ def _hz_to_mel(f): return 2595.0 * math.log10(1.0 + f / 700.0)
287
+ def _mel_to_hz(m): return 700.0 * (10.0 ** (m / 2595.0) - 1.0)
288
+
289
+
290
+ def mel_filterbank(n_mels, n_fft, sr, fmin, fmax):
291
+ key = (n_mels, n_fft, sr, fmin, fmax)
292
+ if key in _MEL_FB_CACHE:
293
+ return _MEL_FB_CACHE[key]
294
+ n_bins = n_fft // 2 + 1
295
+ mel_pts = np.linspace(_hz_to_mel(fmin), _hz_to_mel(fmax),
296
+ n_mels + 2)
297
+ hz_pts = np.array([_mel_to_hz(m) for m in mel_pts])
298
+ bin_pts = np.floor((n_fft + 1) * hz_pts / sr).astype(int)
299
+ bin_pts = np.clip(bin_pts, 0, n_bins - 1)
300
+ fb = np.zeros((n_mels, n_bins))
301
+ for k in range(1, n_mels + 1):
302
+ left, centre, right = bin_pts[k - 1], bin_pts[k], bin_pts[k + 1]
303
+ if centre > left:
304
+ fb[k - 1, left:centre] = (
305
+ np.arange(left, centre) - left) / (centre - left)
306
+ if right > centre:
307
+ fb[k - 1, centre:right] = (
308
+ right - np.arange(centre, right)) / (right - centre)
309
+ _MEL_FB_CACHE[key] = fb
310
+ return fb
311
+
312
+
313
+ def dct_matrix(n_mfcc, n_mels):
314
+ k = np.arange(n_mfcc)[:, None]
315
+ n = np.arange(n_mels)[None, :]
316
+ d = np.cos(math.pi * k * (2 * n + 1) / (2 * n_mels))
317
+ d *= math.sqrt(2.0 / n_mels)
318
+ d[0, :] *= math.sqrt(0.5)
319
+ return d
320
+
321
+
322
+ def trim_silence(sig, threshold_frac=0.06):
323
+ """Trim leading and trailing silence. Returns trimmed signal."""
324
+ abs_sig = np.abs(sig)
325
+ peak = float(abs_sig.max())
326
+ if peak < 1e-9:
327
+ return sig, 0.0
328
+ threshold = threshold_frac * peak
329
+ above = abs_sig > threshold
330
+ if not above.any():
331
+ return sig, 0.0
332
+ start = int(np.argmax(above))
333
+ end = int(len(above) - np.argmax(above[::-1]))
334
+ trimmed = sig[start:end]
335
+ return trimmed, len(trimmed) / SR
336
+
337
+
338
+ def mfcc_sequence(sig):
339
+ frames = frame_signal(sig, FRAME_LEN, HOP)
340
+ window = np.hanning(FRAME_LEN).astype(np.float32)
341
+ frames = frames * window[None, :]
342
+ spec = np.abs(np.fft.rfft(frames, n=FRAME_LEN))
343
+ fb = mel_filterbank(N_MELS, FRAME_LEN, SR, 80.0, 7800.0)
344
+ mel = fb @ (spec ** 2).T
345
+ log_mel = np.log(np.maximum(mel, 1e-10))
346
+ dct = dct_matrix(N_MFCC, N_MELS)
347
+ return (dct @ log_mel).astype(np.float32)
348
+
349
+
350
+ def low_band_sequence(sig):
351
+ frames = frame_signal(sig, FRAME_LEN, HOP)
352
+ window = np.hanning(FRAME_LEN).astype(np.float32)
353
+ frames = frames * window[None, :]
354
+ spec = np.abs(np.fft.rfft(frames, n=FRAME_LEN))
355
+ fb = mel_filterbank(N_MELS_LOW, FRAME_LEN, SR, LOW_FMIN, LOW_FMAX)
356
+ mel = fb @ (spec ** 2).T
357
+ return np.log(np.maximum(mel, 1e-10)).astype(np.float32)
358
+
359
+
360
+ def harmonicity_sequence(sig, f0_min=80.0, f0_max=300.0):
361
+ frames = frame_signal(sig, FRAME_LEN, HOP)
362
+ lag_min = max(1, int(SR / f0_max))
363
+ lag_max = min(FRAME_LEN - 1, int(SR / f0_min))
364
+ out = np.zeros(len(frames), dtype=np.float32)
365
+ for i, f in enumerate(frames):
366
+ f = f - f.mean()
367
+ n_fft = _next_pow2(2 * FRAME_LEN)
368
+ X = np.fft.rfft(f, n=n_fft)
369
+ acf = np.fft.irfft(np.abs(X) ** 2, n=n_fft)[:FRAME_LEN]
370
+ if acf[0] < 1e-12:
371
+ continue
372
+ acf = acf / acf[0]
373
+ seg = acf[lag_min:lag_max + 1]
374
+ if len(seg) > 0:
375
+ out[i] = max(0.0, float(seg.max()))
376
+ return out
377
+
378
+
379
+ def extract_features(sig):
380
+ """v2: trim silence, use trimmed signal for MFCC and low band.
381
+ Duration feature is the trimmed duration, not the padded one."""
382
+ trimmed, trimmed_dur = trim_silence(sig)
383
+ mfcc = mfcc_sequence(trimmed)
384
+ low = low_band_sequence(trimmed)
385
+ harm = harmonicity_sequence(trimmed)
386
+ return np.concatenate([
387
+ mfcc.mean(axis=1), mfcc.std(axis=1),
388
+ low.mean(axis=1), low.std(axis=1),
389
+ [harm.mean(), harm.std(), trimmed_dur],
390
+ ]).astype(np.float32)
391
+
392
+
393
+ FEATURE_DIM = N_MFCC * 2 + N_MELS_LOW * 2 + 3
394
+
395
+
396
+ # =====================================================================
397
+ # §5 Data generation
398
+ # =====================================================================
399
+
400
+ @dataclass
401
+ class Corpus:
402
+ X_train: np.ndarray
403
+ y_train: np.ndarray
404
+ X_test: np.ndarray
405
+ y_test: np.ndarray
406
+ word_index: Dict[str, int]
407
+
408
+
409
+ def generate_corpus(mode, n_per_word_train, n_per_word_test,
410
+ seed, reverb=False, reverb_train_frac=0.0,
411
+ verbose=False):
412
+ rng = np.random.default_rng(seed)
413
+ word_index = {w: i for i, w in enumerate(WORD_ORDER)}
414
+ X_tr, y_tr, X_te, y_te = [], [], [], []
415
+
416
+ t0 = time.time()
417
+ for wi, word in enumerate(WORD_ORDER):
418
+ for k in range(n_per_word_train):
419
+ s = int(rng.integers(1 << 30))
420
+ sig = synth_word(word, mode, s)
421
+ # Training reverb augmentation
422
+ if (mode == 'whisper' and reverb_train_frac > 0
423
+ and rng.random() < reverb_train_frac):
424
+ sig = apply_reverb(
425
+ sig, rt60_s=0.35,
426
+ seed=int(rng.integers(1 << 30)))
427
+ X_tr.append(extract_features(sig))
428
+ y_tr.append(wi)
429
+ for k in range(n_per_word_test):
430
+ s = int(rng.integers(1 << 30))
431
+ sig = synth_word(word, mode, s)
432
+ if reverb:
433
+ sig = apply_reverb(
434
+ sig, rt60_s=0.35,
435
+ seed=int(rng.integers(1 << 30)))
436
+ X_te.append(extract_features(sig))
437
+ y_te.append(wi)
438
+ if verbose:
439
+ print(f" {word:<6s} ({time.time() - t0:.1f}s)")
440
+
441
+ return Corpus(
442
+ X_train=np.stack(X_tr), y_train=np.array(y_tr, dtype=np.int64),
443
+ X_test=np.stack(X_te), y_test=np.array(y_te, dtype=np.int64),
444
+ word_index=word_index,
445
+ )
446
+
447
+
448
+ # =====================================================================
449
+ # §6 Classifier
450
+ # =====================================================================
451
+
452
+ class MLP:
453
+ def __init__(self, in_dim, h1=64, h2=32, out_dim=10, seed=0):
454
+ rng = np.random.default_rng(seed)
455
+ def he(shape):
456
+ return rng.standard_normal(shape) * math.sqrt(2.0 / shape[0])
457
+ self.W1 = he((in_dim, h1)); self.b1 = np.zeros(h1)
458
+ self.W2 = he((h1, h2)); self.b2 = np.zeros(h2)
459
+ self.W3 = he((h2, out_dim)); self.b3 = np.zeros(out_dim)
460
+
461
+ def params(self):
462
+ return [self.W1, self.b1, self.W2, self.b2, self.W3, self.b3]
463
+
464
+ def forward(self, X):
465
+ z1 = X @ self.W1 + self.b1
466
+ h1 = np.maximum(z1, 0.0)
467
+ z2 = h1 @ self.W2 + self.b2
468
+ h2 = np.maximum(z2, 0.0)
469
+ logits = h2 @ self.W3 + self.b3
470
+ return z1, h1, z2, h2, logits
471
+
472
+ def predict_proba(self, X):
473
+ _, _, _, _, logits = self.forward(X)
474
+ z = logits - logits.max(axis=1, keepdims=True)
475
+ e = np.exp(z)
476
+ return e / e.sum(axis=1, keepdims=True)
477
+
478
+ def predict(self, X):
479
+ return self.predict_proba(X).argmax(axis=1)
480
+
481
+ def loss_and_grad(self, X, y):
482
+ z1, h1, z2, h2, logits = self.forward(X)
483
+ n = len(y)
484
+ z = logits - logits.max(axis=1, keepdims=True)
485
+ e = np.exp(z)
486
+ p = e / e.sum(axis=1, keepdims=True)
487
+ loss = -np.log(p[np.arange(n), y] + 1e-12).mean()
488
+ dz = p.copy()
489
+ dz[np.arange(n), y] -= 1.0
490
+ dz /= n
491
+ dW3 = h2.T @ dz; db3 = dz.sum(axis=0)
492
+ dh2 = dz @ self.W3.T; dz2 = dh2 * (z2 > 0.0)
493
+ dW2 = h1.T @ dz2; db2 = dz2.sum(axis=0)
494
+ dh1 = dz2 @ self.W2.T; dz1 = dh1 * (z1 > 0.0)
495
+ dW1 = X.T @ dz1; db1 = dz1.sum(axis=0)
496
+ return loss, [dW1, db1, dW2, db2, dW3, db3]
497
+
498
+
499
+ def train_mlp(model, X, y, X_val=None, y_val=None,
500
+ epochs=200, batch=64, lr=3e-3, seed=0,
501
+ verbose=False):
502
+ rng = np.random.default_rng(seed)
503
+ m = [np.zeros_like(p) for p in model.params()]
504
+ v = [np.zeros_like(p) for p in model.params()]
505
+ t = 0
506
+ b1, b2, eps = 0.9, 0.999, 1e-8
507
+ losses = []
508
+ n = len(y)
509
+ for epoch in range(epochs):
510
+ idx = rng.permutation(n)
511
+ ep_loss = 0.0; nb = 0
512
+ for s in range(0, n, batch):
513
+ sel = idx[s:s + batch]
514
+ loss, grads = model.loss_and_grad(X[sel], y[sel])
515
+ t += 1
516
+ for i, (p, g) in enumerate(zip(model.params(), grads)):
517
+ m[i] = b1 * m[i] + (1 - b1) * g
518
+ v[i] = b2 * v[i] + (1 - b2) * g * g
519
+ mh = m[i] / (1 - b1 ** t)
520
+ vh = v[i] / (1 - b2 ** t)
521
+ p -= lr * mh / (np.sqrt(vh) + eps)
522
+ ep_loss += float(loss); nb += 1
523
+ ep_loss /= max(1, nb)
524
+ losses.append(ep_loss)
525
+ if verbose and ((epoch + 1) % 40 == 0 or epoch == 0):
526
+ msg = f" epoch {epoch+1:>4} loss {ep_loss:.4f}"
527
+ if X_val is not None:
528
+ acc = float((model.predict(X_val) == y_val).mean())
529
+ msg += f" val {acc:.3f}"
530
+ print(msg)
531
+ return losses
532
+
533
+
534
+ # =====================================================================
535
+ # §7 Self-test
536
+ # =====================================================================
537
+
538
+ def self_test(verbose=True):
539
+ checks = []
540
+ sig = synth_word('seven', 'whisper', seed=1)
541
+ checks.append(("whisper length", len(sig) == N_SAMPLES))
542
+ checks.append(("whisper finite", bool(np.all(np.isfinite(sig)))))
543
+ checks.append(("whisper has energy",
544
+ float(np.sqrt(np.mean(sig ** 2))) > 1e-4))
545
+ sig_v = synth_word('seven', 'voiced', seed=1)
546
+ checks.append(("voiced length", len(sig_v) == N_SAMPLES))
547
+ checks.append(("voiced finite", bool(np.all(np.isfinite(sig_v)))))
548
+
549
+ harm_w = harmonicity_sequence(sig)
550
+ harm_v = harmonicity_sequence(sig_v)
551
+ checks.append((f"whisper harmonicity < 0.3 "
552
+ f"(got {harm_w.mean():.3f})",
553
+ float(harm_w.mean()) < 0.3))
554
+ checks.append((f"voiced harmonicity > 0.25 "
555
+ f"(got {harm_v.mean():.3f})",
556
+ float(harm_v.mean()) > 0.25))
557
+
558
+ feats = extract_features(sig)
559
+ checks.append((f"feature dim = {FEATURE_DIM}",
560
+ feats.shape == (FEATURE_DIM,)))
561
+ checks.append(("features finite",
562
+ bool(np.all(np.isfinite(feats)))))
563
+
564
+ sig_a = synth_word('zero', 'whisper', seed=11)
565
+ sig_b = synth_word('zero', 'whisper', seed=12)
566
+ m_a = mfcc_sequence(sig_a).mean(axis=1)
567
+ m_b = mfcc_sequence(sig_b).mean(axis=1)
568
+ cos_same = float(np.dot(m_a, m_b)
569
+ / (np.linalg.norm(m_a)
570
+ * np.linalg.norm(m_b) + 1e-12))
571
+ checks.append((f"same word MFCC cos > 0.9 ({cos_same:.3f})",
572
+ cos_same > 0.9))
573
+
574
+ sig_c = synth_word('four', 'whisper', seed=11)
575
+ m_c = mfcc_sequence(sig_c).mean(axis=1)
576
+ cos_diff = float(np.dot(m_a, m_c)
577
+ / (np.linalg.norm(m_a)
578
+ * np.linalg.norm(m_c) + 1e-12))
579
+ checks.append((f"different words cos < same cos "
580
+ f"({cos_diff:.3f} < {cos_same:.3f})",
581
+ cos_diff < cos_same))
582
+
583
+ passed = sum(1 for _, ok in checks if ok)
584
+ if verbose:
585
+ print()
586
+ print("=" * 74)
587
+ print("SELF-TEST")
588
+ print("=" * 74)
589
+ for name, ok in checks:
590
+ mark = "PASS" if ok else "FAIL"
591
+ print(f" [{mark}] {name}")
592
+ print()
593
+ print(f" {passed}/{len(checks)} correct")
594
+ return passed, len(checks)
595
+
596
+
597
+ # =====================================================================
598
+ # §8 Demo
599
+ # =====================================================================
600
+
601
+ def banner(t, w=76):
602
+ print()
603
+ print("=" * w)
604
+ print(t)
605
+ print("=" * w)
606
+
607
+
608
+ def demo():
609
+ print()
610
+ print("=" * 76)
611
+ print("WHISPER DECODER v2")
612
+ print("=" * 76)
613
+ print(f"""
614
+ Vocabulary : {len(WORDS)} words
615
+ Sample rate : {SR} Hz
616
+ Clip duration : {DURATION_S:.1f} s
617
+ Feature dim : {FEATURE_DIM}
618
+ Classifier : MLP 64-32, Adam, 200 epochs
619
+ Low band : {LOW_FMIN:.0f}-{LOW_FMAX:.0f} Hz (was 200-1000)
620
+ Reverb aug. : {REVERB_TRAIN_FRAC*100:.0f}% of training samples
621
+ """)
622
+
623
+ banner("SELF-TEST")
624
+ self_test(verbose=True)
625
+
626
+ banner("GENERATING TRAINING DATA")
627
+ t0 = time.time()
628
+ train = generate_corpus('whisper', N_TRAIN_PER_WORD,
629
+ N_TEST_PER_WORD, seed=SEED,
630
+ reverb_train_frac=REVERB_TRAIN_FRAC,
631
+ verbose=True)
632
+ print(f" train: {train.X_train.shape} "
633
+ f"test: {train.X_test.shape} "
634
+ f"({time.time() - t0:.1f}s)")
635
+
636
+ mu = train.X_train.mean(axis=0)
637
+ sigma = train.X_train.std(axis=0) + 1e-9
638
+ X_tr = (train.X_train - mu) / sigma
639
+ X_te = (train.X_test - mu) / sigma
640
+
641
+ # v2 fix: shuffle before splitting
642
+ rng = np.random.default_rng(SEED + 100)
643
+ perm = rng.permutation(len(X_tr))
644
+ X_tr = X_tr[perm]
645
+ y_tr_all = train.y_train[perm]
646
+ n_val = len(X_tr) // 5
647
+ X_fit, X_val = X_tr[:-n_val], X_tr[-n_val:]
648
+ y_fit, y_val = y_tr_all[:-n_val], y_tr_all[-n_val:]
649
+
650
+ banner("TRAINING")
651
+ model = MLP(FEATURE_DIM, 64, 32, len(WORDS), seed=SEED)
652
+ n_params = sum(p.size for p in model.params())
653
+ print(f" parameters: {n_params}")
654
+ t0 = time.time()
655
+ train_mlp(model, X_fit, y_fit, X_val, y_val,
656
+ epochs=200, batch=64, lr=3e-3, seed=SEED,
657
+ verbose=True)
658
+ print(f" training time: {time.time() - t0:.1f}s")
659
+ acc_val = float((model.predict(X_val) == y_val).mean())
660
+ print(f" held-out val accuracy: {acc_val*100:.1f}%")
661
+
662
+ banner("TEST 1 -- synthetic whispers (in-distribution)")
663
+ pred = model.predict(X_te)
664
+ acc_w = float((pred == train.y_test).mean())
665
+ print(f" accuracy: {acc_w*100:.1f}%")
666
+
667
+ cm = np.zeros((len(WORDS), len(WORDS)), dtype=int)
668
+ for t, p in zip(train.y_test, pred):
669
+ cm[t, p] += 1
670
+ print()
671
+ print(f" {'true \\ pred':<10}" + "".join(
672
+ f"{w[:6]:>7}" for w in WORD_ORDER))
673
+ print(" " + "-" * (10 + 7 * len(WORD_ORDER)))
674
+ for i, w in enumerate(WORD_ORDER):
675
+ row = f" {w:<10}" + "".join(f"{v:>7}" for v in cm[i])
676
+ print(row)
677
+
678
+ banner("TEST 2 -- synthetic voiced speech (domain shift)")
679
+ voiced = generate_corpus('voiced', 20, N_TEST_PER_WORD,
680
+ seed=SEED + 1)
681
+ X_v = (voiced.X_test - mu) / sigma
682
+ pred_v = model.predict(X_v)
683
+ acc_v = float((pred_v == voiced.y_test).mean())
684
+ print(f" accuracy: {acc_v*100:.1f}% "
685
+ f"(chance = {100.0/len(WORDS):.1f}%)")
686
+
687
+ banner("TEST 3 -- reverberant whispers (robustness)")
688
+ reverb = generate_corpus('whisper', 20, N_TEST_PER_WORD,
689
+ seed=SEED + 2, reverb=True)
690
+ X_r = (reverb.X_test - mu) / sigma
691
+ pred_r = model.predict(X_r)
692
+ acc_r = float((pred_r == reverb.y_test).mean())
693
+ print(f" accuracy: {acc_r*100:.1f}%")
694
+
695
+ banner("SUMMARY")
696
+ print(f" {'corpus':<34} {'accuracy':>9}")
697
+ print(" " + "-" * 46)
698
+ print(f" {'whisper (in-distribution)':<34} "
699
+ f"{acc_w*100:>8.1f}%")
700
+ print(f" {'voiced (domain shift)':<34} "
701
+ f"{acc_v*100:>8.1f}%")
702
+ print(f" {'reverberant whisper (robustness)':<34} "
703
+ f"{acc_r*100:>8.1f}%")
704
+ print()
705
+ print(" chance: 10.0%")
706
+
707
+ if HAS_MPL:
708
+ banner("RENDERING")
709
+ out_dir = "whisper_figures"
710
+ os.makedirs(out_dir, exist_ok=True)
711
+ render_demo(train, cm, acc_w, acc_v, acc_r, out_dir)
712
+
713
+
714
+ def render_demo(corpus, cm, acc_w, acc_v, acc_r, out_dir):
715
+ fig = plt.figure(figsize=(15, 9))
716
+
717
+ sig_w = synth_word('seven', 'whisper', seed=42)
718
+ sig_v = synth_word('seven', 'voiced', seed=42)
719
+ ax1 = fig.add_subplot(4, 2, 1)
720
+ ax1.plot(np.arange(len(sig_w)) / SR, sig_w, color="#1f77b4",
721
+ linewidth=0.6)
722
+ ax1.set_title("Whispered 'seven'")
723
+ ax1.set_xlabel("time (s)")
724
+ ax1.grid(alpha=0.3)
725
+
726
+ ax2 = fig.add_subplot(4, 2, 2)
727
+ ax2.plot(np.arange(len(sig_v)) / SR, sig_v, color="#d62728",
728
+ linewidth=0.6)
729
+ ax2.set_title("Voiced 'seven'")
730
+ ax2.set_xlabel("time (s)")
731
+ ax2.grid(alpha=0.3)
732
+
733
+ harm_w = harmonicity_sequence(sig_w)
734
+ harm_v = harmonicity_sequence(sig_v)
735
+ t_h = np.arange(len(harm_w)) * HOP / SR
736
+ ax3 = fig.add_subplot(4, 2, 3)
737
+ ax3.plot(t_h, harm_w, color="#1f77b4", label="whisper")
738
+ ax3.plot(t_h, harm_v, color="#d62728", label="voiced")
739
+ ax3.axhline(0.25, color="#888", linestyle=":",
740
+ label="voiced threshold")
741
+ ax3.set_title("Harmonicity over time")
742
+ ax3.set_ylim(-0.05, 1.05)
743
+ ax3.legend(fontsize=8)
744
+ ax3.grid(alpha=0.3)
745
+
746
+ ax4 = fig.add_subplot(4, 2, 4)
747
+ mf = mfcc_sequence(sig_w)
748
+ ax4.imshow(mf, aspect="auto", origin="lower", cmap="magma",
749
+ extent=[0, len(sig_w) / SR, 0, N_MFCC])
750
+ ax4.set_title("MFCC of whispered 'seven'")
751
+ ax4.set_xlabel("time (s)")
752
+ plt.colorbar(ax4.images[0], ax=ax4, shrink=0.7)
753
+
754
+ ax5 = fig.add_subplot(4, 2, 5)
755
+ im = ax5.imshow(cm, cmap="Blues", aspect="auto")
756
+ ax5.set_xticks(np.arange(len(WORD_ORDER)))
757
+ ax5.set_yticks(np.arange(len(WORD_ORDER)))
758
+ ax5.set_xticklabels(WORD_ORDER, rotation=45, fontsize=8)
759
+ ax5.set_yticklabels(WORD_ORDER, fontsize=8)
760
+ for i in range(len(WORD_ORDER)):
761
+ for j in range(len(WORD_ORDER)):
762
+ if cm[i, j] > 0:
763
+ ax5.text(j, i, str(cm[i, j]),
764
+ ha="center", va="center",
765
+ color="white" if cm[i, j] > cm.max() / 2
766
+ else "black", fontsize=8)
767
+ ax5.set_title(f"Confusion matrix "
768
+ f"(accuracy {acc_w*100:.1f}%)")
769
+ plt.colorbar(im, ax=ax5, shrink=0.7)
770
+
771
+ ax6 = fig.add_subplot(4, 2, 6)
772
+ names = ["whisper", "voiced", "reverb"]
773
+ vals = [acc_w, acc_v, acc_r]
774
+ colors = ["#1f77b4", "#d62728", "#2ca02c"]
775
+ bars = ax6.bar(names, vals, color=colors, edgecolor="#222")
776
+ for bar, v in zip(bars, vals):
777
+ ax6.text(bar.get_x() + bar.get_width() / 2,
778
+ v + 0.02, f"{v*100:.1f}%",
779
+ ha="center", fontsize=10)
780
+ ax6.axhline(0.1, color="#888", linestyle=":",
781
+ label="chance (10%)")
782
+ ax6.set_ylim(0, 1.1)
783
+ ax6.set_title("Accuracy by corpus condition")
784
+ ax6.legend(fontsize=9)
785
+ ax6.grid(alpha=0.3, axis="y")
786
+
787
+ fig.suptitle("Whisper Decoder v2", fontsize=13)
788
+ plt.tight_layout()
789
+ path = os.path.join(out_dir, "whisper_demo.png")
790
+ plt.savefig(path, dpi=130, bbox_inches="tight")
791
+ plt.close()
792
+ print(f" saved: {path}")
793
+
794
+
795
+ # =====================================================================
796
+ # §9 Entry point
797
+ # =====================================================================
798
+
799
+ def main():
800
+ p = argparse.ArgumentParser()
801
+ p.add_argument("--self-test", action="store_true")
802
+ p.add_argument("--epochs", type=int, default=200)
803
+ args = p.parse_args()
804
+ if args.self_test:
805
+ self_test(verbose=True)
806
+ else:
807
+ demo()
808
+
809
+
810
+ if __name__ == "__main__":
811
+ main()