zeechimp commited on
Commit
6ca0d2e
·
verified ·
1 Parent(s): 5691579

Create examples/feature_importance.py

Browse files
Files changed (1) hide show
  1. examples/feature_importance.py +74 -0
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()