whisper-decoder / whisper_decoder.py
zeechimp's picture
Upload whisper_decoder.py
15b9f68 verified
Raw History Blame Contribute Delete
29.2 kB
#!/usr/bin/env python3
"""
whisper_decoder.py
==================
Closed-vocabulary word classification from whispered speech.
Fix history
-----------
v1 Three bugs.
(a) Training data was generated word-by-word in alphabetical
order. The validation split took the last 20% of samples,
which was only the last two words ('two' and 'zero'). The
reported validation accuracy of 0.0% was a split artifact,
not a training failure.
(b) The synthesized words are shorter (0.2-0.5 s) than the
fixed clip duration (1.2 s). Features were averaged over
the whole clip, so zero-padding dominated the spectrum.
Word identity was swamped by silence.
(c) The low-band feature (200-1000 Hz) misses F2 for front
vowels. 'two' and 'three' were nearly identical in this
feature.
v2 Fixes.
(a) Shuffle before splitting. Stratified validation.
(b) Trim to non-silent content. Use trimmed duration as a
feature. Report the trimmed length.
(c) Extend the low band to 200-2500 Hz.
(d) Lower the voiced harmonicity threshold to 0.25.
(e) Add reverb augmentation during training.
Vocabulary: 10 English digits. Code-only, no downloads.
"""
from __future__ import annotations
import argparse
import math
import os
import time
import wave
from dataclasses import dataclass, field
from pathlib import Path
from typing import Dict, List, Optional, Tuple
import numpy as np
try:
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
HAS_MPL = True
except ImportError:
HAS_MPL = False
SR = 16000
DURATION_S = 1.2
N_SAMPLES = int(SR * DURATION_S)
FRAME_LEN = 512
HOP = 160
N_MELS = 24
N_MFCC = 13
N_MELS_LOW = 10
LOW_FMIN = 200.0
LOW_FMAX = 2500.0 # was 1000 in v1
SEED = 0
N_TRAIN_PER_WORD = 200
N_TEST_PER_WORD = 40
REVERB_TRAIN_FRAC = 0.25 # 25% of training samples get reverb
# =====================================================================
# §1 Phoneme inventory (unchanged)
# =====================================================================
PHONEMES: Dict[str, Dict] = {
'i': {'type': 'vowel', 'dur': 0.12,
'formants': [(270, 80, 1.0), (2290, 120, 0.8), (3010, 150, 0.5)]},
'ɪ': {'type': 'vowel', 'dur': 0.10,
'formants': [(390, 80, 1.0), (1990, 120, 0.8), (2550, 150, 0.5)]},
'e': {'type': 'vowel', 'dur': 0.12,
'formants': [(530, 80, 1.0), (1840, 120, 0.8), (2480, 150, 0.5)]},
'ɛ': {'type': 'vowel', 'dur': 0.11,
'formants': [(660, 80, 1.0), (1720, 120, 0.8), (2410, 150, 0.5)]},
'a': {'type': 'vowel', 'dur': 0.13,
'formants': [(730, 80, 1.0), (1090, 120, 0.8), (2440, 150, 0.5)]},
'ɑ': {'type': 'vowel', 'dur': 0.13,
'formants': [(730, 80, 1.0), (1090, 120, 0.8), (2440, 150, 0.5)]},
'ɔ': {'type': 'vowel', 'dur': 0.13,
'formants': [(570, 80, 1.0), (840, 120, 0.8), (2410, 150, 0.5)]},
'o': {'type': 'vowel', 'dur': 0.13,
'formants': [(570, 80, 1.0), (840, 120, 0.8), (2410, 150, 0.5)]},
'ʊ': {'type': 'vowel', 'dur': 0.11,
'formants': [(440, 80, 1.0), (1020, 120, 0.8), (2240, 150, 0.5)]},
'u': {'type': 'vowel', 'dur': 0.12,
'formants': [(300, 80, 1.0), (870, 120, 0.8), (2240, 150, 0.5)]},
'ʌ': {'type': 'vowel', 'dur': 0.11,
'formants': [(640, 80, 1.0), (1190, 120, 0.8), (2390, 150, 0.5)]},
's': {'type': 'fricative', 'dur': 0.11,
'formants': [(6000, 1000, 1.0)]},
'z': {'type': 'fricative', 'dur': 0.10,
'formants': [(6000, 1000, 0.7)]},
'ʃ': {'type': 'fricative', 'dur': 0.11,
'formants': [(3500, 800, 1.0)]},
'f': {'type': 'fricative', 'dur': 0.10,
'formants': [(5000, 2000, 0.6)]},
'v': {'type': 'fricative', 'dur': 0.08,
'formants': [(5000, 2000, 0.7)]},
'θ': {'type': 'fricative', 'dur': 0.10,
'formants': [(5000, 1500, 0.5)]},
'p': {'type': 'plosive', 'dur': 0.08, 'silence': 0.06,
'formants': [(1000, 400, 1.0)]},
't': {'type': 'plosive', 'dur': 0.08, 'silence': 0.06,
'formants': [(4000, 800, 1.0)]},
'k': {'type': 'plosive', 'dur': 0.08, 'silence': 0.06,
'formants': [(2500, 600, 1.0)]},
'n': {'type': 'nasal', 'dur': 0.09,
'formants': [(300, 100, 1.0), (1500, 200, 0.4), (2500, 250, 0.2)]},
'm': {'type': 'nasal', 'dur': 0.09,
'formants': [(300, 100, 1.0), (1000, 200, 0.4), (2200, 250, 0.2)]},
'r': {'type': 'approx', 'dur': 0.07,
'formants': [(300, 100, 1.0), (1100, 200, 0.6), (1600, 200, 0.3)]},
'w': {'type': 'approx', 'dur': 0.07,
'formants': [(300, 100, 1.0), (600, 200, 0.5), (2200, 250, 0.2)]},
}
WORDS: Dict[str, List[str]] = {
'zero': ['z', 'i', 'r', 'o'],
'one': ['w', 'ʌ', 'n'],
'two': ['t', 'u'],
'three': ['θ', 'r', 'i'],
'four': ['f', 'o', 'r'],
'five': ['f', 'a', 'i', 'v'],
'six': ['s', 'ɪ', 'k', 's'],
'seven': ['s', 'ɛ', 'v', 'ɛ', 'n'],
'eight': ['e', 'i', 't'],
'nine': ['n', 'a', 'i', 'n'],
}
WORD_ORDER = sorted(WORDS.keys())
# =====================================================================
# §2 Formant filter
# =====================================================================
_FILTER_CACHE: Dict = {}
def formant_filter(formants, n_fft, sr):
key = (tuple(formants), n_fft, sr)
if key in _FILTER_CACHE:
return _FILTER_CACHE[key]
freqs = np.fft.rfftfreq(n_fft, d=1.0 / sr)
resp = np.zeros_like(freqs)
for f0, bw, amp in formants:
resp += amp * np.exp(-((freqs - f0) ** 2) / (2.0 * bw * bw))
_FILTER_CACHE[key] = resp
return resp
def _next_pow2(n):
p = 1
while p < n:
p *= 2
return p
# =====================================================================
# §3 Phoneme and word synthesis
# =====================================================================
def synth_phoneme_whisper(name, sr, seed, jitter=0.03):
p = PHONEMES[name]
rng = np.random.default_rng(seed)
dur = p['dur'] * (1.0 + jitter * rng.uniform(-1, 1))
n = int(sr * dur)
if p['type'] == 'plosive':
n_sil = int(sr * p.get('silence', 0.06))
n_burst = max(4, n - n_sil)
n_fft = _next_pow2(n_burst)
noise = rng.standard_normal(n_burst)
H = formant_filter(p['formants'], n_fft, sr)
X = np.fft.rfft(noise, n=n_fft)
burst = np.fft.irfft(X * H, n=n_fft)[:n_burst]
env = np.exp(-np.arange(n_burst) / (0.005 * sr))
out = np.zeros(n)
out[n_sil:] = burst * env
return out.astype(np.float32)
n_fft = _next_pow2(n)
noise = rng.standard_normal(n)
H = formant_filter(p['formants'], n_fft, sr)
X = np.fft.rfft(noise, n=n_fft)
sig = np.fft.irfft(X * H, n=n_fft)[:n]
ramp = min(int(0.015 * sr), n // 4)
env = np.ones(n)
if ramp > 0:
up = 0.5 - 0.5 * np.cos(np.arange(ramp) * np.pi / ramp)
env[:ramp] = up
env[-ramp:] = up[::-1]
return (sig * env).astype(np.float32)
def synth_phoneme_voiced(name, sr, seed, f0=120.0, jitter=0.03):
p = PHONEMES[name]
rng = np.random.default_rng(seed)
dur = p['dur'] * (1.0 + jitter * rng.uniform(-1, 1))
n = int(sr * dur)
if p['type'] == 'plosive':
return synth_phoneme_whisper(name, sr, seed, jitter)
n_fft = _next_pow2(n)
period = max(2, int(sr / f0))
pulse = np.zeros(n)
for i in range(0, n, period):
pulse[i] = 1.0
H = formant_filter(p['formants'], n_fft, sr)
X = np.fft.rfft(pulse, n=n_fft)
sig = np.fft.irfft(X * H, n=n_fft)[:n]
peak = float(np.max(np.abs(sig))) + 1e-12
sig = sig / peak * 0.5
ramp = min(int(0.015 * sr), n // 4)
env = np.ones(n)
if ramp > 0:
up = 0.5 - 0.5 * np.cos(np.arange(ramp) * np.pi / ramp)
env[:ramp] = up
env[-ramp:] = up[::-1]
return (sig * env).astype(np.float32)
def synth_word(word, mode, seed, noise_frac=0.002):
rng = np.random.default_rng(seed)
pieces = []
for ph in WORDS[word]:
s = int(rng.integers(1 << 30))
if mode == 'whisper':
pieces.append(synth_phoneme_whisper(ph, SR, s))
else:
f0 = 120.0 + 30.0 * rng.uniform(-1, 1)
pieces.append(synth_phoneme_voiced(ph, SR, s, f0=f0))
sig = np.concatenate(pieces) if pieces else np.zeros(1)
if len(sig) >= N_SAMPLES:
sig = sig[:N_SAMPLES]
else:
sig = np.concatenate([sig, np.zeros(N_SAMPLES - len(sig))])
if noise_frac > 0:
sig = sig + noise_frac * rng.standard_normal(len(sig))
return sig.astype(np.float32)
def apply_reverb(sig, rt60_s=0.4, seed=0):
rng = np.random.default_rng(seed)
n_ir = int(SR * rt60_s * 1.5)
decay = np.exp(-3.0 * np.arange(n_ir) / (SR * rt60_s / 6.91))
ir = rng.standard_normal(n_ir) * decay
ir[0] = 1.0
out = np.convolve(sig, ir, mode="same")
peak = float(np.max(np.abs(out))) + 1e-12
return (out / peak * 0.9).astype(np.float32)
# =====================================================================
# §4 Feature extraction -- v2 TRIM + extended low band
# =====================================================================
_MEL_FB_CACHE: Dict = {}
def frame_signal(sig, frame_len, hop):
if len(sig) < frame_len:
sig = np.pad(sig, (0, frame_len - len(sig)))
n = 1 + (len(sig) - frame_len) // hop
return np.stack([sig[i * hop:i * hop + frame_len]
for i in range(n)])
def _hz_to_mel(f): return 2595.0 * math.log10(1.0 + f / 700.0)
def _mel_to_hz(m): return 700.0 * (10.0 ** (m / 2595.0) - 1.0)
def mel_filterbank(n_mels, n_fft, sr, fmin, fmax):
key = (n_mels, n_fft, sr, fmin, fmax)
if key in _MEL_FB_CACHE:
return _MEL_FB_CACHE[key]
n_bins = n_fft // 2 + 1
mel_pts = np.linspace(_hz_to_mel(fmin), _hz_to_mel(fmax),
n_mels + 2)
hz_pts = np.array([_mel_to_hz(m) for m in mel_pts])
bin_pts = np.floor((n_fft + 1) * hz_pts / sr).astype(int)
bin_pts = np.clip(bin_pts, 0, n_bins - 1)
fb = np.zeros((n_mels, n_bins))
for k in range(1, n_mels + 1):
left, centre, right = bin_pts[k - 1], bin_pts[k], bin_pts[k + 1]
if centre > left:
fb[k - 1, left:centre] = (
np.arange(left, centre) - left) / (centre - left)
if right > centre:
fb[k - 1, centre:right] = (
right - np.arange(centre, right)) / (right - centre)
_MEL_FB_CACHE[key] = fb
return fb
def dct_matrix(n_mfcc, n_mels):
k = np.arange(n_mfcc)[:, None]
n = np.arange(n_mels)[None, :]
d = np.cos(math.pi * k * (2 * n + 1) / (2 * n_mels))
d *= math.sqrt(2.0 / n_mels)
d[0, :] *= math.sqrt(0.5)
return d
def trim_silence(sig, threshold_frac=0.06):
"""Trim leading and trailing silence. Returns trimmed signal."""
abs_sig = np.abs(sig)
peak = float(abs_sig.max())
if peak < 1e-9:
return sig, 0.0
threshold = threshold_frac * peak
above = abs_sig > threshold
if not above.any():
return sig, 0.0
start = int(np.argmax(above))
end = int(len(above) - np.argmax(above[::-1]))
trimmed = sig[start:end]
return trimmed, len(trimmed) / SR
def mfcc_sequence(sig):
frames = frame_signal(sig, FRAME_LEN, HOP)
window = np.hanning(FRAME_LEN).astype(np.float32)
frames = frames * window[None, :]
spec = np.abs(np.fft.rfft(frames, n=FRAME_LEN))
fb = mel_filterbank(N_MELS, FRAME_LEN, SR, 80.0, 7800.0)
mel = fb @ (spec ** 2).T
log_mel = np.log(np.maximum(mel, 1e-10))
dct = dct_matrix(N_MFCC, N_MELS)
return (dct @ log_mel).astype(np.float32)
def low_band_sequence(sig):
frames = frame_signal(sig, FRAME_LEN, HOP)
window = np.hanning(FRAME_LEN).astype(np.float32)
frames = frames * window[None, :]
spec = np.abs(np.fft.rfft(frames, n=FRAME_LEN))
fb = mel_filterbank(N_MELS_LOW, FRAME_LEN, SR, LOW_FMIN, LOW_FMAX)
mel = fb @ (spec ** 2).T
return np.log(np.maximum(mel, 1e-10)).astype(np.float32)
def harmonicity_sequence(sig, f0_min=80.0, f0_max=300.0):
frames = frame_signal(sig, FRAME_LEN, HOP)
lag_min = max(1, int(SR / f0_max))
lag_max = min(FRAME_LEN - 1, int(SR / f0_min))
out = np.zeros(len(frames), dtype=np.float32)
for i, f in enumerate(frames):
f = f - f.mean()
n_fft = _next_pow2(2 * FRAME_LEN)
X = np.fft.rfft(f, n=n_fft)
acf = np.fft.irfft(np.abs(X) ** 2, n=n_fft)[:FRAME_LEN]
if acf[0] < 1e-12:
continue
acf = acf / acf[0]
seg = acf[lag_min:lag_max + 1]
if len(seg) > 0:
out[i] = max(0.0, float(seg.max()))
return out
def extract_features(sig):
"""v2: trim silence, use trimmed signal for MFCC and low band.
Duration feature is the trimmed duration, not the padded one."""
trimmed, trimmed_dur = trim_silence(sig)
mfcc = mfcc_sequence(trimmed)
low = low_band_sequence(trimmed)
harm = harmonicity_sequence(trimmed)
return np.concatenate([
mfcc.mean(axis=1), mfcc.std(axis=1),
low.mean(axis=1), low.std(axis=1),
[harm.mean(), harm.std(), trimmed_dur],
]).astype(np.float32)
FEATURE_DIM = N_MFCC * 2 + N_MELS_LOW * 2 + 3
# =====================================================================
# §5 Data generation
# =====================================================================
@dataclass
class Corpus:
X_train: np.ndarray
y_train: np.ndarray
X_test: np.ndarray
y_test: np.ndarray
word_index: Dict[str, int]
def generate_corpus(mode, n_per_word_train, n_per_word_test,
seed, reverb=False, reverb_train_frac=0.0,
verbose=False):
rng = np.random.default_rng(seed)
word_index = {w: i for i, w in enumerate(WORD_ORDER)}
X_tr, y_tr, X_te, y_te = [], [], [], []
t0 = time.time()
for wi, word in enumerate(WORD_ORDER):
for k in range(n_per_word_train):
s = int(rng.integers(1 << 30))
sig = synth_word(word, mode, s)
# Training reverb augmentation
if (mode == 'whisper' and reverb_train_frac > 0
and rng.random() < reverb_train_frac):
sig = apply_reverb(
sig, rt60_s=0.35,
seed=int(rng.integers(1 << 30)))
X_tr.append(extract_features(sig))
y_tr.append(wi)
for k in range(n_per_word_test):
s = int(rng.integers(1 << 30))
sig = synth_word(word, mode, s)
if reverb:
sig = apply_reverb(
sig, rt60_s=0.35,
seed=int(rng.integers(1 << 30)))
X_te.append(extract_features(sig))
y_te.append(wi)
if verbose:
print(f" {word:<6s} ({time.time() - t0:.1f}s)")
return Corpus(
X_train=np.stack(X_tr), y_train=np.array(y_tr, dtype=np.int64),
X_test=np.stack(X_te), y_test=np.array(y_te, dtype=np.int64),
word_index=word_index,
)
# =====================================================================
# §6 Classifier
# =====================================================================
class MLP:
def __init__(self, in_dim, h1=64, h2=32, out_dim=10, seed=0):
rng = np.random.default_rng(seed)
def he(shape):
return rng.standard_normal(shape) * math.sqrt(2.0 / shape[0])
self.W1 = he((in_dim, h1)); self.b1 = np.zeros(h1)
self.W2 = he((h1, h2)); self.b2 = np.zeros(h2)
self.W3 = he((h2, out_dim)); self.b3 = np.zeros(out_dim)
def params(self):
return [self.W1, self.b1, self.W2, self.b2, self.W3, self.b3]
def forward(self, X):
z1 = X @ self.W1 + self.b1
h1 = np.maximum(z1, 0.0)
z2 = h1 @ self.W2 + self.b2
h2 = np.maximum(z2, 0.0)
logits = h2 @ self.W3 + self.b3
return z1, h1, z2, h2, logits
def predict_proba(self, X):
_, _, _, _, logits = self.forward(X)
z = logits - logits.max(axis=1, keepdims=True)
e = np.exp(z)
return e / e.sum(axis=1, keepdims=True)
def predict(self, X):
return self.predict_proba(X).argmax(axis=1)
def loss_and_grad(self, X, y):
z1, h1, z2, h2, logits = self.forward(X)
n = len(y)
z = logits - logits.max(axis=1, keepdims=True)
e = np.exp(z)
p = e / e.sum(axis=1, keepdims=True)
loss = -np.log(p[np.arange(n), y] + 1e-12).mean()
dz = p.copy()
dz[np.arange(n), y] -= 1.0
dz /= n
dW3 = h2.T @ dz; db3 = dz.sum(axis=0)
dh2 = dz @ self.W3.T; dz2 = dh2 * (z2 > 0.0)
dW2 = h1.T @ dz2; db2 = dz2.sum(axis=0)
dh1 = dz2 @ self.W2.T; dz1 = dh1 * (z1 > 0.0)
dW1 = X.T @ dz1; db1 = dz1.sum(axis=0)
return loss, [dW1, db1, dW2, db2, dW3, db3]
def train_mlp(model, X, y, X_val=None, y_val=None,
epochs=200, batch=64, lr=3e-3, seed=0,
verbose=False):
rng = np.random.default_rng(seed)
m = [np.zeros_like(p) for p in model.params()]
v = [np.zeros_like(p) for p in model.params()]
t = 0
b1, b2, eps = 0.9, 0.999, 1e-8
losses = []
n = len(y)
for epoch in range(epochs):
idx = rng.permutation(n)
ep_loss = 0.0; nb = 0
for s in range(0, n, batch):
sel = idx[s:s + batch]
loss, grads = model.loss_and_grad(X[sel], y[sel])
t += 1
for i, (p, g) in enumerate(zip(model.params(), grads)):
m[i] = b1 * m[i] + (1 - b1) * g
v[i] = b2 * v[i] + (1 - b2) * g * g
mh = m[i] / (1 - b1 ** t)
vh = v[i] / (1 - b2 ** t)
p -= lr * mh / (np.sqrt(vh) + eps)
ep_loss += float(loss); nb += 1
ep_loss /= max(1, nb)
losses.append(ep_loss)
if verbose and ((epoch + 1) % 40 == 0 or epoch == 0):
msg = f" epoch {epoch+1:>4} loss {ep_loss:.4f}"
if X_val is not None:
acc = float((model.predict(X_val) == y_val).mean())
msg += f" val {acc:.3f}"
print(msg)
return losses
# =====================================================================
# §7 Self-test
# =====================================================================
def self_test(verbose=True):
checks = []
sig = synth_word('seven', 'whisper', seed=1)
checks.append(("whisper length", len(sig) == N_SAMPLES))
checks.append(("whisper finite", bool(np.all(np.isfinite(sig)))))
checks.append(("whisper has energy",
float(np.sqrt(np.mean(sig ** 2))) > 1e-4))
sig_v = synth_word('seven', 'voiced', seed=1)
checks.append(("voiced length", len(sig_v) == N_SAMPLES))
checks.append(("voiced finite", bool(np.all(np.isfinite(sig_v)))))
harm_w = harmonicity_sequence(sig)
harm_v = harmonicity_sequence(sig_v)
checks.append((f"whisper harmonicity < 0.3 "
f"(got {harm_w.mean():.3f})",
float(harm_w.mean()) < 0.3))
checks.append((f"voiced harmonicity > 0.25 "
f"(got {harm_v.mean():.3f})",
float(harm_v.mean()) > 0.25))
feats = extract_features(sig)
checks.append((f"feature dim = {FEATURE_DIM}",
feats.shape == (FEATURE_DIM,)))
checks.append(("features finite",
bool(np.all(np.isfinite(feats)))))
sig_a = synth_word('zero', 'whisper', seed=11)
sig_b = synth_word('zero', 'whisper', seed=12)
m_a = mfcc_sequence(sig_a).mean(axis=1)
m_b = mfcc_sequence(sig_b).mean(axis=1)
cos_same = float(np.dot(m_a, m_b)
/ (np.linalg.norm(m_a)
* np.linalg.norm(m_b) + 1e-12))
checks.append((f"same word MFCC cos > 0.9 ({cos_same:.3f})",
cos_same > 0.9))
sig_c = synth_word('four', 'whisper', seed=11)
m_c = mfcc_sequence(sig_c).mean(axis=1)
cos_diff = float(np.dot(m_a, m_c)
/ (np.linalg.norm(m_a)
* np.linalg.norm(m_c) + 1e-12))
checks.append((f"different words cos < same cos "
f"({cos_diff:.3f} < {cos_same:.3f})",
cos_diff < cos_same))
passed = sum(1 for _, ok in checks if ok)
if verbose:
print()
print("=" * 74)
print("SELF-TEST")
print("=" * 74)
for name, ok in checks:
mark = "PASS" if ok else "FAIL"
print(f" [{mark}] {name}")
print()
print(f" {passed}/{len(checks)} correct")
return passed, len(checks)
# =====================================================================
# §8 Demo
# =====================================================================
def banner(t, w=76):
print()
print("=" * w)
print(t)
print("=" * w)
def demo():
print()
print("=" * 76)
print("WHISPER DECODER v2")
print("=" * 76)
print(f"""
Vocabulary : {len(WORDS)} words
Sample rate : {SR} Hz
Clip duration : {DURATION_S:.1f} s
Feature dim : {FEATURE_DIM}
Classifier : MLP 64-32, Adam, 200 epochs
Low band : {LOW_FMIN:.0f}-{LOW_FMAX:.0f} Hz (was 200-1000)
Reverb aug. : {REVERB_TRAIN_FRAC*100:.0f}% of training samples
""")
banner("SELF-TEST")
self_test(verbose=True)
banner("GENERATING TRAINING DATA")
t0 = time.time()
train = generate_corpus('whisper', N_TRAIN_PER_WORD,
N_TEST_PER_WORD, seed=SEED,
reverb_train_frac=REVERB_TRAIN_FRAC,
verbose=True)
print(f" train: {train.X_train.shape} "
f"test: {train.X_test.shape} "
f"({time.time() - t0:.1f}s)")
mu = train.X_train.mean(axis=0)
sigma = train.X_train.std(axis=0) + 1e-9
X_tr = (train.X_train - mu) / sigma
X_te = (train.X_test - mu) / sigma
# v2 fix: shuffle before splitting
rng = np.random.default_rng(SEED + 100)
perm = rng.permutation(len(X_tr))
X_tr = X_tr[perm]
y_tr_all = train.y_train[perm]
n_val = len(X_tr) // 5
X_fit, X_val = X_tr[:-n_val], X_tr[-n_val:]
y_fit, y_val = y_tr_all[:-n_val], y_tr_all[-n_val:]
banner("TRAINING")
model = MLP(FEATURE_DIM, 64, 32, len(WORDS), seed=SEED)
n_params = sum(p.size for p in model.params())
print(f" parameters: {n_params}")
t0 = time.time()
train_mlp(model, X_fit, y_fit, X_val, y_val,
epochs=200, batch=64, lr=3e-3, seed=SEED,
verbose=True)
print(f" training time: {time.time() - t0:.1f}s")
acc_val = float((model.predict(X_val) == y_val).mean())
print(f" held-out val accuracy: {acc_val*100:.1f}%")
banner("TEST 1 -- synthetic whispers (in-distribution)")
pred = model.predict(X_te)
acc_w = float((pred == train.y_test).mean())
print(f" accuracy: {acc_w*100:.1f}%")
cm = np.zeros((len(WORDS), len(WORDS)), dtype=int)
for t, p in zip(train.y_test, pred):
cm[t, p] += 1
print()
print(f" {'true \\ pred':<10}" + "".join(
f"{w[:6]:>7}" for w in WORD_ORDER))
print(" " + "-" * (10 + 7 * len(WORD_ORDER)))
for i, w in enumerate(WORD_ORDER):
row = f" {w:<10}" + "".join(f"{v:>7}" for v in cm[i])
print(row)
banner("TEST 2 -- synthetic voiced speech (domain shift)")
voiced = generate_corpus('voiced', 20, N_TEST_PER_WORD,
seed=SEED + 1)
X_v = (voiced.X_test - mu) / sigma
pred_v = model.predict(X_v)
acc_v = float((pred_v == voiced.y_test).mean())
print(f" accuracy: {acc_v*100:.1f}% "
f"(chance = {100.0/len(WORDS):.1f}%)")
banner("TEST 3 -- reverberant whispers (robustness)")
reverb = generate_corpus('whisper', 20, N_TEST_PER_WORD,
seed=SEED + 2, reverb=True)
X_r = (reverb.X_test - mu) / sigma
pred_r = model.predict(X_r)
acc_r = float((pred_r == reverb.y_test).mean())
print(f" accuracy: {acc_r*100:.1f}%")
banner("SUMMARY")
print(f" {'corpus':<34} {'accuracy':>9}")
print(" " + "-" * 46)
print(f" {'whisper (in-distribution)':<34} "
f"{acc_w*100:>8.1f}%")
print(f" {'voiced (domain shift)':<34} "
f"{acc_v*100:>8.1f}%")
print(f" {'reverberant whisper (robustness)':<34} "
f"{acc_r*100:>8.1f}%")
print()
print(" chance: 10.0%")
if HAS_MPL:
banner("RENDERING")
out_dir = "whisper_figures"
os.makedirs(out_dir, exist_ok=True)
render_demo(train, cm, acc_w, acc_v, acc_r, out_dir)
def render_demo(corpus, cm, acc_w, acc_v, acc_r, out_dir):
fig = plt.figure(figsize=(15, 9))
sig_w = synth_word('seven', 'whisper', seed=42)
sig_v = synth_word('seven', 'voiced', seed=42)
ax1 = fig.add_subplot(4, 2, 1)
ax1.plot(np.arange(len(sig_w)) / SR, sig_w, color="#1f77b4",
linewidth=0.6)
ax1.set_title("Whispered 'seven'")
ax1.set_xlabel("time (s)")
ax1.grid(alpha=0.3)
ax2 = fig.add_subplot(4, 2, 2)
ax2.plot(np.arange(len(sig_v)) / SR, sig_v, color="#d62728",
linewidth=0.6)
ax2.set_title("Voiced 'seven'")
ax2.set_xlabel("time (s)")
ax2.grid(alpha=0.3)
harm_w = harmonicity_sequence(sig_w)
harm_v = harmonicity_sequence(sig_v)
t_h = np.arange(len(harm_w)) * HOP / SR
ax3 = fig.add_subplot(4, 2, 3)
ax3.plot(t_h, harm_w, color="#1f77b4", label="whisper")
ax3.plot(t_h, harm_v, color="#d62728", label="voiced")
ax3.axhline(0.25, color="#888", linestyle=":",
label="voiced threshold")
ax3.set_title("Harmonicity over time")
ax3.set_ylim(-0.05, 1.05)
ax3.legend(fontsize=8)
ax3.grid(alpha=0.3)
ax4 = fig.add_subplot(4, 2, 4)
mf = mfcc_sequence(sig_w)
ax4.imshow(mf, aspect="auto", origin="lower", cmap="magma",
extent=[0, len(sig_w) / SR, 0, N_MFCC])
ax4.set_title("MFCC of whispered 'seven'")
ax4.set_xlabel("time (s)")
plt.colorbar(ax4.images[0], ax=ax4, shrink=0.7)
ax5 = fig.add_subplot(4, 2, 5)
im = ax5.imshow(cm, cmap="Blues", aspect="auto")
ax5.set_xticks(np.arange(len(WORD_ORDER)))
ax5.set_yticks(np.arange(len(WORD_ORDER)))
ax5.set_xticklabels(WORD_ORDER, rotation=45, fontsize=8)
ax5.set_yticklabels(WORD_ORDER, fontsize=8)
for i in range(len(WORD_ORDER)):
for j in range(len(WORD_ORDER)):
if cm[i, j] > 0:
ax5.text(j, i, str(cm[i, j]),
ha="center", va="center",
color="white" if cm[i, j] > cm.max() / 2
else "black", fontsize=8)
ax5.set_title(f"Confusion matrix "
f"(accuracy {acc_w*100:.1f}%)")
plt.colorbar(im, ax=ax5, shrink=0.7)
ax6 = fig.add_subplot(4, 2, 6)
names = ["whisper", "voiced", "reverb"]
vals = [acc_w, acc_v, acc_r]
colors = ["#1f77b4", "#d62728", "#2ca02c"]
bars = ax6.bar(names, vals, color=colors, edgecolor="#222")
for bar, v in zip(bars, vals):
ax6.text(bar.get_x() + bar.get_width() / 2,
v + 0.02, f"{v*100:.1f}%",
ha="center", fontsize=10)
ax6.axhline(0.1, color="#888", linestyle=":",
label="chance (10%)")
ax6.set_ylim(0, 1.1)
ax6.set_title("Accuracy by corpus condition")
ax6.legend(fontsize=9)
ax6.grid(alpha=0.3, axis="y")
fig.suptitle("Whisper Decoder v2", fontsize=13)
plt.tight_layout()
path = os.path.join(out_dir, "whisper_demo.png")
plt.savefig(path, dpi=130, bbox_inches="tight")
plt.close()
print(f" saved: {path}")
# =====================================================================
# §9 Entry point
# =====================================================================
def main():
p = argparse.ArgumentParser()
p.add_argument("--self-test", action="store_true")
p.add_argument("--epochs", type=int, default=200)
args = p.parse_args()
if args.self_test:
self_test(verbose=True)
else:
demo()
if __name__ == "__main__":
main()