Download whisper_decoder.py from zeechimp/whisper-decoder: direct link, hf CLI and curl.
- Browser
- Download file 29.2 kB
-
https://huggingface.co/zeechimp/whisper-decoder/resolve/main/whisper_decoder.py
- Command line
-
hf download hf://zeechimp/whisper-decoder/whisper_decoder.py
-
curl -L -o whisper_decoder.py https://huggingface.co/zeechimp/whisper-decoder/resolve/main/whisper_decoder.py
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 | |
| # ===================================================================== | |
| 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() |