#!/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()