Download examples/feature_importance.py from zeechimp/whisper-decoder: direct link, hf CLI and curl.
- Browser
- Download file 2.29 kB
-
https://huggingface.co/zeechimp/whisper-decoder/resolve/main/examples/feature_importance.py
- Command line
-
hf download hf://zeechimp/whisper-decoder/examples/feature_importance.py
-
curl -L -o feature_importance.py https://huggingface.co/zeechimp/whisper-decoder/resolve/main/examples/feature_importance.py
2.29 kB
| #!/usr/bin/env python3 | |
| """ | |
| Which features carry the word identity? | |
| Train the classifier, then ablate individual feature families | |
| and report the drop in test accuracy. | |
| Run: python examples/feature_importance.py | |
| """ | |
| import numpy as np | |
| from whisper_decoder import ( | |
| MLP, train_mlp, generate_corpus, FEATURE_DIM, | |
| ) | |
| # Feature layout: | |
| # [0:13] MFCC mean | |
| # [13:26] MFCC std | |
| # [26:36] low-band mel mean | |
| # [36:46] low-band mel std | |
| # [46] harmonicity mean | |
| # [47] harmonicity std | |
| # [48] trimmed duration | |
| GROUPS = { | |
| 'MFCC mean': list(range(0, 13)), | |
| 'MFCC std': list(range(13, 26)), | |
| 'low-band mean': list(range(26, 36)), | |
| 'low-band std': list(range(36, 46)), | |
| 'harmonicity': [46, 47], | |
| 'duration': [48], | |
| } | |
| def train_and_eval(X_tr, y_tr, X_te, y_te, cols): | |
| mu = X_tr[:, cols].mean(axis=0) | |
| sigma = X_tr[:, cols].std(axis=0) + 1e-9 | |
| Xtr = (X_tr[:, cols] - mu) / sigma | |
| Xte = (X_te[:, cols] - mu) / sigma | |
| model = MLP(len(cols), 64, 32, 10, seed=0) | |
| train_mlp(model, Xtr, y_tr, epochs=150, batch=64, | |
| lr=3e-3, seed=0) | |
| return float((model.predict(Xte) == y_te).mean()) | |
| def main(): | |
| print("generating corpus...") | |
| train = generate_corpus('whisper', 150, 40, seed=0, | |
| reverb_train_frac=0.25) | |
| test = generate_corpus('whisper', 10, 40, seed=99) | |
| all_cols = list(range(FEATURE_DIM)) | |
| base = train_and_eval(train.X_train, train.y_train, | |
| test.X_test, test.y_test, all_cols) | |
| print(f"\nall {FEATURE_DIM} features: accuracy " | |
| f"{base*100:.1f}%\n") | |
| print(f" {'ablation':<22} {'accuracy':>10} {'drop':>8}") | |
| print(" " + "-" * 44) | |
| for name, cols in GROUPS.items(): | |
| remaining = [c for c in all_cols if c not in cols] | |
| acc = train_and_eval(train.X_train, train.y_train, | |
| test.X_test, test.y_test, remaining) | |
| drop = base - acc | |
| print(f" {'without ' + name:<22} " | |
| f"{acc*100:>9.1f}% {drop*100:>+7.1f}%") | |
| print() | |
| print(" A large drop means the removed family was " | |
| "load-bearing.") | |
| print(" A near-zero drop means the classifier did not use it.") | |
| if __name__ == "__main__": | |
| main() |