#!/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()