{
 "cells": [
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": "# Letter confusion matrix from ISOLET\n\nThis notebook reproduces the measured confusion matrix shown on the speller page.\n\n**Recognizer:** a support vector machine (RBF kernel, C=10, gamma='scale') on the 617 acoustic features that ship with ISOLET, standardised first.\n\n**Evaluation:** 5-fold cross-validation where each fold is one ISOLET speaker group (30 speakers), so the recognizer is always tested on speakers it never trained on.\n\nRequirements: `pip install numpy scikit-learn matplotlib`"
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": "## 1. Download ISOLET (UCI Machine Learning Repository)"
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": "import urllib.request, zipfile, os, subprocess\nURL = 'https://archive.ics.uci.edu/static/public/54/isolet.zip'\nif not os.path.exists('isolet5.data'):\n    urllib.request.urlretrieve(URL, 'isolet.zip')\n    zipfile.ZipFile('isolet.zip').extractall('.')\n    # the files are Unix-compressed (.Z); gzip can unpack them\n    for f in ['isolet1+2+3+4.data.Z', 'isolet5.data.Z']:\n        subprocess.run(['gzip', '-d', '-f', f], check=True)\nprint(sorted(f for f in os.listdir('.') if f.endswith('.data')))"
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": "## 2. Load features and labels\n\nEach row is one recording: 617 features (spectral, contour, sonorant and post-sonorant features computed by the ISOLET authors) plus the letter label 1\u201326. `isolet1+2+3+4.data` holds speaker groups 1\u20134, `isolet5.data` holds group 5."
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": "import numpy as np\na = np.loadtxt('isolet1+2+3+4.data', delimiter=',')\nb = np.loadtxt('isolet5.data', delimiter=',')\nX = np.vstack([a[:, :-1], b[:, :-1]])\ny = np.concatenate([a[:, -1], b[:, -1]]).astype(int) - 1\nprint(X.shape, y.shape)"
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": "## 3. Speaker groups\n\nThe first file is the four groups concatenated in order (about 1,560 rows each). A few recordings are missing in the original data, so splitting it into four equal blocks can put a handful of rows from one speaker in the neighbouring fold. Group 5 is exact."
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": "g = np.concatenate([np.repeat([0, 1, 2, 3], len(a) // 4 + 1)[:len(a)], np.full(len(b), 4)])\nprint(np.bincount(g))"
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": "## 4. Train and predict with speaker-independent cross-validation"
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": "from sklearn.svm import SVC\nfrom sklearn.preprocessing import StandardScaler\nfrom sklearn.pipeline import make_pipeline\nfrom sklearn.model_selection import cross_val_predict, GroupKFold\n\nclf = make_pipeline(StandardScaler(), SVC(C=10, gamma='scale'))\npred = cross_val_predict(clf, X, y, groups=g, cv=GroupKFold(5), n_jobs=-1)\nprint('accuracy', (pred == y).mean())"
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": "## 5. Confusion matrix (row = letter spoken, column = letter recognised)"
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": "L = 'ABCDEFGHIJKLMNOPQRSTUVWXYZ'\nM = np.zeros((26, 26), int)\nfor t, p in zip(y, pred):\n    M[t, p] += 1\nR = M / M.sum(1, keepdims=True)\n\nimport matplotlib.pyplot as plt\nfig, ax = plt.subplots(figsize=(8, 8))\noff = R.copy(); np.fill_diagonal(off, 0)\nax.imshow(off, cmap='Reds', vmax=0.08)\nax.set_xticks(range(26), L); ax.set_yticks(range(26), L)\nax.set_xlabel('heard as'); ax.set_ylabel('spoken')\nplt.show()"
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": "## 6. Most confused pairs"
  },
  {
   "cell_type": "code",
   "metadata": {},
   "execution_count": null,
   "outputs": [],
   "source": "pairs = sorted(((R[i, j], L[i], L[j]) for i in range(26) for j in range(26) if i != j and M[i, j]), reverse=True)\nfor r, s, h in pairs[:20]:\n    print(f'{s} heard as {h}: {r:.1%}')"
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "Python 3",
   "language": "python",
   "name": "python3"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}