Download examples/confusion_plot.py from zeechimp/whisper-decoder: direct link, hf CLI and curl.
- Browser
- Download file 2.74 kB
-
https://huggingface.co/zeechimp/whisper-decoder/resolve/main/examples/confusion_plot.py
- Command line
-
hf download hf://zeechimp/whisper-decoder/examples/confusion_plot.py
-
curl -L -o confusion_plot.py https://huggingface.co/zeechimp/whisper-decoder/resolve/main/examples/confusion_plot.py
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() |