Create examples/feature_importance.py
Browse files
examples/feature_importance.py
ADDED
|
@@ -0,0 +1,74 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python3
|
| 2 |
+
"""
|
| 3 |
+
Which features carry the word identity?
|
| 4 |
+
|
| 5 |
+
Train the classifier, then ablate individual feature families
|
| 6 |
+
and report the drop in test accuracy.
|
| 7 |
+
|
| 8 |
+
Run: python examples/feature_importance.py
|
| 9 |
+
"""
|
| 10 |
+
import numpy as np
|
| 11 |
+
from whisper_decoder import (
|
| 12 |
+
MLP, train_mlp, generate_corpus, FEATURE_DIM,
|
| 13 |
+
)
|
| 14 |
+
|
| 15 |
+
# Feature layout:
|
| 16 |
+
# [0:13] MFCC mean
|
| 17 |
+
# [13:26] MFCC std
|
| 18 |
+
# [26:36] low-band mel mean
|
| 19 |
+
# [36:46] low-band mel std
|
| 20 |
+
# [46] harmonicity mean
|
| 21 |
+
# [47] harmonicity std
|
| 22 |
+
# [48] trimmed duration
|
| 23 |
+
|
| 24 |
+
GROUPS = {
|
| 25 |
+
'MFCC mean': list(range(0, 13)),
|
| 26 |
+
'MFCC std': list(range(13, 26)),
|
| 27 |
+
'low-band mean': list(range(26, 36)),
|
| 28 |
+
'low-band std': list(range(36, 46)),
|
| 29 |
+
'harmonicity': [46, 47],
|
| 30 |
+
'duration': [48],
|
| 31 |
+
}
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
def train_and_eval(X_tr, y_tr, X_te, y_te, cols):
|
| 35 |
+
mu = X_tr[:, cols].mean(axis=0)
|
| 36 |
+
sigma = X_tr[:, cols].std(axis=0) + 1e-9
|
| 37 |
+
Xtr = (X_tr[:, cols] - mu) / sigma
|
| 38 |
+
Xte = (X_te[:, cols] - mu) / sigma
|
| 39 |
+
model = MLP(len(cols), 64, 32, 10, seed=0)
|
| 40 |
+
train_mlp(model, Xtr, y_tr, epochs=150, batch=64,
|
| 41 |
+
lr=3e-3, seed=0)
|
| 42 |
+
return float((model.predict(Xte) == y_te).mean())
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
def main():
|
| 46 |
+
print("generating corpus...")
|
| 47 |
+
train = generate_corpus('whisper', 150, 40, seed=0,
|
| 48 |
+
reverb_train_frac=0.25)
|
| 49 |
+
test = generate_corpus('whisper', 10, 40, seed=99)
|
| 50 |
+
|
| 51 |
+
all_cols = list(range(FEATURE_DIM))
|
| 52 |
+
base = train_and_eval(train.X_train, train.y_train,
|
| 53 |
+
test.X_test, test.y_test, all_cols)
|
| 54 |
+
print(f"\nall {FEATURE_DIM} features: accuracy "
|
| 55 |
+
f"{base*100:.1f}%\n")
|
| 56 |
+
|
| 57 |
+
print(f" {'ablation':<22} {'accuracy':>10} {'drop':>8}")
|
| 58 |
+
print(" " + "-" * 44)
|
| 59 |
+
for name, cols in GROUPS.items():
|
| 60 |
+
remaining = [c for c in all_cols if c not in cols]
|
| 61 |
+
acc = train_and_eval(train.X_train, train.y_train,
|
| 62 |
+
test.X_test, test.y_test, remaining)
|
| 63 |
+
drop = base - acc
|
| 64 |
+
print(f" {'without ' + name:<22} "
|
| 65 |
+
f"{acc*100:>9.1f}% {drop*100:>+7.1f}%")
|
| 66 |
+
|
| 67 |
+
print()
|
| 68 |
+
print(" A large drop means the removed family was "
|
| 69 |
+
"load-bearing.")
|
| 70 |
+
print(" A near-zero drop means the classifier did not use it.")
|
| 71 |
+
|
| 72 |
+
|
| 73 |
+
if __name__ == "__main__":
|
| 74 |
+
main()
|