File size: 2,742 Bytes
664e67e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
75
76
77
78
79
80
81
82
83
84
#!/usr/bin/env python3
"""
Plot the confusion matrix for whispers, voiced, and reverberant
whispers side by side.  Writes to whisper_figures/confusion_3way.png.

Run: python examples/confusion_plot.py
"""
import os
import numpy as np
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
from whisper_decoder import (
    MLP, train_mlp, generate_corpus, FEATURE_DIM, WORD_ORDER,
)

OUT = "whisper_figures"


def confusion(model, corpus, mu, sigma):
    X = (corpus.X_test - mu) / sigma
    pred = model.predict(X)
    cm = np.zeros((10, 10), dtype=int)
    for t, p in zip(corpus.y_test, pred):
        cm[t, p] += 1
    return cm, float((pred == corpus.y_test).mean())


def main():
    os.makedirs(OUT, exist_ok=True)

    print("training...")
    train = generate_corpus('whisper', 150, 30, seed=0,
                            reverb_train_frac=0.25)
    mu = train.X_train.mean(axis=0)
    sigma = train.X_train.std(axis=0) + 1e-9
    X_tr = (train.X_train - mu) / sigma

    model = MLP(FEATURE_DIM, 64, 32, 10, seed=0)
    train_mlp(model, X_tr, train.y_train, epochs=200,
              batch=64, lr=3e-3, seed=0)

    print("evaluating on three conditions...")
    whisper = generate_corpus('whisper', 10, 30, seed=10)
    voiced = generate_corpus('voiced', 10, 30, seed=11)
    reverb = generate_corpus('whisper', 10, 30, seed=12, reverb=True)

    fig, axes = plt.subplots(1, 3, figsize=(18, 5.5))
    conditions = [
        ("whisper  (in-distribution)", whisper),
        ("voiced  (domain shift)", voiced),
        ("reverberant whisper", reverb),
    ]
    for ax, (title, corpus) in zip(axes, conditions):
        cm, acc = confusion(model, corpus, mu, sigma)
        im = ax.imshow(cm, cmap="Blues", aspect="auto",
                       vmin=0, vmax=cm.max())
        ax.set_xticks(np.arange(10))
        ax.set_yticks(np.arange(10))
        ax.set_xticklabels(WORD_ORDER, rotation=45, fontsize=8)
        ax.set_yticklabels(WORD_ORDER, fontsize=8)
        for i in range(10):
            for j in range(10):
                if cm[i, j] > 0:
                    ax.text(j, i, str(cm[i, j]),
                            ha="center", va="center",
                            fontsize=7,
                            color="white" if cm[i, j] > cm.max() / 2
                            else "black")
        ax.set_xlabel("predicted")
        ax.set_ylabel("true")
        ax.set_title(f"{title}\naccuracy {acc*100:.1f}%",
                     fontsize=11)
        plt.colorbar(im, ax=ax, shrink=0.7)

    plt.tight_layout()
    path = os.path.join(OUT, "confusion_3way.png")
    plt.savefig(path, dpi=130, bbox_inches="tight")
    plt.close()
    print(f"saved: {path}")


if __name__ == "__main__":
    main()