File size: 2,286 Bytes
6ca0d2e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
#!/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()