whisper-decoder / examples /confusion_plot.py
zeechimp's picture
Create examples/confusion_plot.py
664e67e verified
Raw History Blame Contribute Delete
2.74 kB
#!/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()