{
 "cells": [
  {
   "cell_type": "markdown",
   "id": "f46b21b1",
   "metadata": {},
   "source": [
    "# MCB128 — Can a Multilayer Perceptron Recognize Splice Sites?\n",
    "\n",
    "**60-minute hands-on lesson**\n",
    "\n",
    "### Biological question\n",
    "\n",
    "> **Given a short DNA sequence, can a neural network determine whether it contains a splice donor, a splice acceptor, or neither?**\n",
    "\n",
    "We will train a small **multilayer perceptron (MLP)** on real primate DNA sequences from the UCI *Molecular Biology (Splice-junction Gene Sequences)* dataset.\n",
    "\n",
    "### Learning goals\n",
    "\n",
    "By the end of this notebook, you should be able to:\n",
    "\n",
    "1. explain what a splice donor and splice acceptor are;\n",
    "2. represent a DNA sequence using **one-hot encoding**;\n",
    "3. construct a fully connected MLP in PyTorch;\n",
    "4. explain the roles of a **minibatch, epoch, forward pass, loss, backpropagation, and optimizer step**;\n",
    "5. use training and validation curves to diagnose learning;\n",
    "6. evaluate a classifier on unseen data;\n",
    "7. use the trained model to ask **which sequence positions matter most for splice-site recognition**.\n",
    "\n",
    "This notebook deliberately uses the same ingredients as the MLP lecture:\n",
    "**one-hot inputs → fully connected layer → nonlinear activation → output logits → cross-entropy → backpropagation → optimizer update**."
   ]
  },
  {
   "cell_type": "markdown",
   "id": "9a970108",
   "metadata": {},
   "source": [
    "## Lesson roadmap\n",
    "\n",
    "| Approx. time | Activity |\n",
    "|---:|---|\n",
    "| 0–10 min | Biology of pre-mRNA splicing |\n",
    "| 10–18 min | Inspect and one-hot encode DNA sequences |\n",
    "| 18–28 min | Build the MLP |\n",
    "| 28–43 min | Train it with minibatches and backpropagation |\n",
    "| 43–50 min | Evaluate generalization |\n",
    "| 50–58 min | Ask what sequence positions the network uses |\n",
    "| 58–60 min | Wrap-up |"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "0403a031",
   "metadata": {},
   "source": [
    "# 1. Biology: how does pre-mRNA splicing work?\n",
    "\n",
    "Eukaryotic genes are usually interrupted by **introns**.\n",
    "\n",
    "After a gene is transcribed, the initial RNA molecule — the **pre-mRNA** — contains both:\n",
    "\n",
    "- **exons**, which are retained in the mature RNA;\n",
    "- **introns**, which are removed.\n",
    "\n",
    "A large RNA–protein complex called the **spliceosome** recognizes the appropriate boundaries, removes the intron, and joins the neighboring exons.\n",
    "\n",
    "A simplified gene looks like this:\n",
    "\n",
    "```text\n",
    "DNA / pre-mRNA\n",
    "\n",
    "        5' splice site                         3' splice site\n",
    "             donor                                acceptor\n",
    "               |                                     |\n",
    " exon 1        |----------- intron -----------------|       exon 2\n",
    "==============>|=====================================|>==============\n",
    "               GU                                   AG\n",
    "               ^                                     ^\n",
    "          start of intron                       end of intron\n",
    "```\n",
    "\n",
    "In DNA sequence, the RNA `GU` at the major 5′ splice site appears as **GT**.\n",
    "\n",
    "## Splice-site sequence signals\n",
    "\n",
    "Most introns processed by the major spliceosome follow the **GU–AG rule**:\n",
    "\n",
    "- the intron usually begins with **GU** in RNA (`GT` in DNA);\n",
    "- the intron usually ends with **AG**.\n",
    "\n",
    "But the spliceosome does **not** recognize splice sites using only two nucleotides.\n",
    "\n",
    "Recognition also depends on surrounding sequence:\n",
    "\n",
    "- bases surrounding the **5′ donor site** contribute to U1 snRNP recognition;\n",
    "- the **3′ acceptor region** includes a branch-point signal and a nearby **polypyrimidine tract**, which is enriched for C/U;\n",
    "- many GT and AG dinucleotides occur in genomes without being used as splice sites.\n",
    "\n",
    "So splice-site recognition is naturally a **sequence classification problem**:\n",
    "\n",
    "> Given the local nucleotide context, is this sequence centered on a real donor, a real acceptor, or neither?\n",
    "\n",
    "That is the problem we will give our MLP."
   ]
  },
  {
   "cell_type": "markdown",
   "id": "357fcb44",
   "metadata": {},
   "source": [
    "### Think before coding\n",
    "\n",
    "Suppose you encounter `GT` somewhere in a gene.\n",
    "\n",
    "**Why isn't “contains GT” sufficient to conclude that it is a splice donor?**\n",
    "\n",
    "Keep your answer in mind — by the end of the notebook we will ask whether the neural network uses nucleotides outside the central splice-site motif."
   ]
  },
  {
   "cell_type": "markdown",
   "id": "da023dfa",
   "metadata": {},
   "source": [
    "# 2. The splice-junction dataset\n",
    "\n",
    "The UCI dataset contains **3,190 primate DNA sequences**, each **60 nucleotides long**.\n",
    "\n",
    "Each sequence belongs to one of three classes:\n",
    "\n",
    "| Dataset label | Meaning | Biological name |\n",
    "|---|---|---|\n",
    "| `EI` | exon → intron boundary | **splice donor / 5′ splice site** |\n",
    "| `IE` | intron → exon boundary | **splice acceptor / 3′ splice site** |\n",
    "| `N` | neither boundary | **negative example** |\n",
    "\n",
    "The splice junction is near the center of the 60-nt window.\n",
    "\n",
    "We will treat the 60 nucleotide positions as 60 categorical input variables."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "c4317313",
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "import pandas as pd\n",
    "import matplotlib.pyplot as plt\n",
    "\n",
    "import torch\n",
    "from torch import nn\n",
    "from torch.utils.data import TensorDataset, DataLoader\n",
    "\n",
    "from sklearn.model_selection import train_test_split\n",
    "from sklearn.metrics import (\n",
    "    accuracy_score,\n",
    "    balanced_accuracy_score,\n",
    "    confusion_matrix,\n",
    "    ConfusionMatrixDisplay,\n",
    ")\n",
    "\n",
    "SEED = 128\n",
    "np.random.seed(SEED)\n",
    "torch.manual_seed(SEED)\n",
    "\n",
    "print(\"PyTorch version:\", torch.__version__)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "27185d9b",
   "metadata": {},
   "outputs": [],
   "source": [
    "DATA_URLS = [\n",
    "    \"https://archive.ics.uci.edu/ml/machine-learning-databases/\"\n",
    "    \"molecular-biology/splice-junction-gene-sequences/splice.data\",\n",
    "    \"https://www.inf.u-szeged.hu/~tothl/gepitan/uci%20repository/\"\n",
    "    \"molecular-biology/splice-junction-gene-sequences/splice.data\",\n",
    "    \"https://huggingface.co/datasets/mstz/splice/resolve/main/splice.data\",\n",
    "]\n",
    "\n",
    "df = None\n",
    "last_error = None\n",
    "\n",
    "for url in DATA_URLS:\n",
    "    try:\n",
    "        df = pd.read_csv(\n",
    "            url,\n",
    "            header=None,\n",
    "            names=[\"class\", \"name\", \"sequence\"],\n",
    "            skipinitialspace=True,\n",
    "        )\n",
    "        print(\"Loaded data from:\", url)\n",
    "        break\n",
    "    except Exception as e:\n",
    "        last_error = e\n",
    "\n",
    "if df is None:\n",
    "    raise RuntimeError(\n",
    "        \"Could not download splice.data. Download 'splice.data' from the UCI \"\n",
    "        \"dataset page and place it next to this notebook, then replace \"\n",
    "        \"DATA_URLS with ['splice.data'].\"\n",
    "    ) from last_error\n",
    "\n",
    "for col in [\"class\", \"name\", \"sequence\"]:\n",
    "    df[col] = df[col].astype(str).str.strip()\n",
    "\n",
    "df[\"sequence\"] = df[\"sequence\"].str.upper()\n",
    "\n",
    "print(\"Dataset shape:\", df.shape)\n",
    "display(df.head())"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "8da3d3b3",
   "metadata": {},
   "outputs": [],
   "source": [
    "print(\"Sequence lengths:\")\n",
    "display(df[\"sequence\"].str.len().value_counts().sort_index())\n",
    "\n",
    "print(\"\\nClass counts:\")\n",
    "display(df[\"class\"].value_counts())\n",
    "\n",
    "df[\"class\"].value_counts().plot(kind=\"bar\")\n",
    "plt.ylabel(\"Number of sequences\")\n",
    "plt.xlabel(\"Class\")\n",
    "plt.title(\"Class distribution\")\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "58112df1",
   "metadata": {},
   "outputs": [],
   "source": [
    "examples = (\n",
    "    df.groupby(\"class\", group_keys=False)\n",
    "      .sample(3, random_state=SEED)\n",
    "      .copy()\n",
    ")\n",
    "\n",
    "examples[\"central 16 nt\"] = examples[\"sequence\"].str.slice(22, 38)\n",
    "display(examples[[\"class\", \"central 16 nt\"]])"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "7cbe587d",
   "metadata": {},
   "source": [
    "# 3. Turning DNA into numbers\n",
    "\n",
    "Neural networks operate on numbers, not characters.\n",
    "\n",
    "For a canonical DNA base we use **one-hot encoding**:\n",
    "\n",
    "\\[\n",
    "A=[1,0,0,0],\\quad\n",
    "C=[0,1,0,0],\\quad\n",
    "G=[0,0,1,0],\\quad\n",
    "T=[0,0,0,1]\n",
    "\\]\n",
    "\n",
    "A 60-nt sequence therefore becomes a:\n",
    "\n",
    "\\[\n",
    "60 \\times 4\n",
    "\\]\n",
    "\n",
    "matrix, which we flatten into a vector of:\n",
    "\n",
    "\\[\n",
    "60 \\times 4 = 240\n",
    "\\]\n",
    "\n",
    "input features.\n",
    "\n",
    "The historical dataset contains a small number of ambiguous IUPAC nucleotide codes. We encode those as probability-like mixtures. For example,\n",
    "\n",
    "\\[\n",
    "R=A/G=[0.5,0,0.5,0].\n",
    "\\]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "8751ba2a",
   "metadata": {},
   "outputs": [],
   "source": [
    "BASE_VECTOR = {\n",
    "    \"A\": [1.0, 0.0, 0.0, 0.0],\n",
    "    \"C\": [0.0, 1.0, 0.0, 0.0],\n",
    "    \"G\": [0.0, 0.0, 1.0, 0.0],\n",
    "    \"T\": [0.0, 0.0, 0.0, 1.0],\n",
    "    \"R\": [0.5, 0.0, 0.5, 0.0],\n",
    "    \"Y\": [0.0, 0.5, 0.0, 0.5],\n",
    "    \"S\": [0.0, 0.5, 0.5, 0.0],\n",
    "    \"W\": [0.5, 0.0, 0.0, 0.5],\n",
    "    \"K\": [0.0, 0.0, 0.5, 0.5],\n",
    "    \"M\": [0.5, 0.5, 0.0, 0.0],\n",
    "    \"B\": [0.0, 1/3, 1/3, 1/3],\n",
    "    \"D\": [1/3, 0.0, 1/3, 1/3],\n",
    "    \"H\": [1/3, 1/3, 0.0, 1/3],\n",
    "    \"V\": [1/3, 1/3, 1/3, 0.0],\n",
    "    \"N\": [0.25, 0.25, 0.25, 0.25],\n",
    "}\n",
    "\n",
    "def one_hot_encode(sequence):\n",
    "    return np.array([BASE_VECTOR[base] for base in sequence], dtype=np.float32)\n",
    "\n",
    "example_matrix = one_hot_encode(df.loc[0, \"sequence\"])\n",
    "\n",
    "print(\"One sequence:\", df.loc[0, \"sequence\"][:20] + \"...\")\n",
    "print(\"Encoded shape:\", example_matrix.shape)\n",
    "print(\"\\nFirst five encoded bases:\")\n",
    "print(example_matrix[:5])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "5dbb5f79",
   "metadata": {},
   "outputs": [],
   "source": [
    "X = np.stack([one_hot_encode(seq).reshape(-1) for seq in df[\"sequence\"]])\n",
    "\n",
    "label_to_int = {\"N\": 0, \"EI\": 1, \"IE\": 2}\n",
    "int_to_label = {0: \"Neither\", 1: \"Donor (EI)\", 2: \"Acceptor (IE)\"}\n",
    "y = df[\"class\"].map(label_to_int).to_numpy(dtype=np.int64)\n",
    "\n",
    "print(\"X shape:\", X.shape)\n",
    "print(\"y shape:\", y.shape)\n",
    "print(\"Number of input features per sequence:\", X.shape[1])"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "3a8c2161",
   "metadata": {},
   "source": [
    "# 4. Training, validation, and test sets\n",
    "\n",
    "We split the data into:\n",
    "\n",
    "- **70% training** — used to update the parameters;\n",
    "- **15% validation** — used to monitor generalization during training;\n",
    "- **15% test** — untouched until final evaluation.\n",
    "\n",
    "`stratify=y` keeps approximately the same class proportions in all three sets."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "6ecb5e24",
   "metadata": {},
   "outputs": [],
   "source": [
    "indices = np.arange(len(y))\n",
    "\n",
    "X_train, X_temp, y_train, y_temp, idx_train, idx_temp = train_test_split(\n",
    "    X, y, indices,\n",
    "    test_size=0.30,\n",
    "    random_state=SEED,\n",
    "    stratify=y,\n",
    ")\n",
    "\n",
    "X_val, X_test, y_val, y_test, idx_val, idx_test = train_test_split(\n",
    "    X_temp, y_temp, idx_temp,\n",
    "    test_size=0.50,\n",
    "    random_state=SEED,\n",
    "    stratify=y_temp,\n",
    ")\n",
    "\n",
    "print(\"Train:\", X_train.shape)\n",
    "print(\"Validation:\", X_val.shape)\n",
    "print(\"Test:\", X_test.shape)\n",
    "\n",
    "majority_class = np.bincount(y_train).argmax()\n",
    "majority_validation_accuracy = np.mean(y_val == majority_class)\n",
    "\n",
    "print(\n",
    "    f\"\\nMajority-class baseline validation accuracy: \"\n",
    "    f\"{majority_validation_accuracy:.3f}\"\n",
    ")"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "46064387",
   "metadata": {},
   "source": [
    "### Minibatches\n",
    "\n",
    "We convert the arrays to PyTorch tensors and create a `DataLoader`.\n",
    "\n",
    "With a batch size of 64, the network sees 64 sequences, computes one gradient from them, and then updates its parameters."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "d5ce8709",
   "metadata": {},
   "outputs": [],
   "source": [
    "X_train_t = torch.tensor(X_train, dtype=torch.float32)\n",
    "X_val_t   = torch.tensor(X_val, dtype=torch.float32)\n",
    "X_test_t  = torch.tensor(X_test, dtype=torch.float32)\n",
    "\n",
    "y_train_t = torch.tensor(y_train, dtype=torch.long)\n",
    "y_val_t   = torch.tensor(y_val, dtype=torch.long)\n",
    "y_test_t  = torch.tensor(y_test, dtype=torch.long)\n",
    "\n",
    "BATCH_SIZE = 64\n",
    "\n",
    "train_loader = DataLoader(\n",
    "    TensorDataset(X_train_t, y_train_t),\n",
    "    batch_size=BATCH_SIZE,\n",
    "    shuffle=True,\n",
    ")\n",
    "\n",
    "print(\"Training examples:\", len(X_train_t))\n",
    "print(\"Minibatches per epoch:\", len(train_loader))"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "7e0a0172",
   "metadata": {},
   "source": [
    "# 5. Build the MLP\n",
    "\n",
    "Our model has one hidden layer:\n",
    "\n",
    "\\[\n",
    "240\n",
    "\\;\\longrightarrow\\;\n",
    "64\n",
    "\\;\\xrightarrow{\\mathrm{ReLU}}\\;\n",
    "3\n",
    "\\]\n",
    "\n",
    "For input \\(x\\),\n",
    "\n",
    "\\[\n",
    "h = \\mathrm{ReLU}(W_1 x + b_1)\n",
    "\\]\n",
    "\n",
    "and\n",
    "\n",
    "\\[\n",
    "z = W_2 h + b_2.\n",
    "\\]\n",
    "\n",
    "The three output values \\(z\\) are **logits** for:\n",
    "\n",
    "1. neither;\n",
    "2. donor;\n",
    "3. acceptor.\n",
    "\n",
    "`CrossEntropyLoss` will compare these logits with the correct class labels."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "019b46cc",
   "metadata": {},
   "outputs": [],
   "source": [
    "INPUT_DIM = X_train.shape[1]\n",
    "HIDDEN_DIM = 64\n",
    "OUTPUT_DIM = 3\n",
    "\n",
    "model = nn.Sequential(\n",
    "    nn.Linear(INPUT_DIM, HIDDEN_DIM),\n",
    "    nn.ReLU(),\n",
    "    nn.Linear(HIDDEN_DIM, OUTPUT_DIM),\n",
    ")\n",
    "\n",
    "print(model)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "625192ea",
   "metadata": {},
   "outputs": [],
   "source": [
    "criterion = nn.CrossEntropyLoss()\n",
    "optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)\n",
    "\n",
    "n_parameters = sum(\n",
    "    p.numel() for p in model.parameters() if p.requires_grad\n",
    ")\n",
    "print(\"Trainable parameters:\", n_parameters)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "9b9055ee",
   "metadata": {},
   "outputs": [],
   "source": [
    "def evaluate(model, X, y, criterion):\n",
    "    model.eval()\n",
    "\n",
    "    with torch.no_grad():\n",
    "        logits = model(X)\n",
    "        loss = criterion(logits, y).item()\n",
    "        pred = logits.argmax(dim=1).cpu().numpy()\n",
    "\n",
    "    y_np = y.cpu().numpy()\n",
    "\n",
    "    accuracy = accuracy_score(y_np, pred)\n",
    "    balanced_accuracy = balanced_accuracy_score(y_np, pred)\n",
    "\n",
    "    return loss, accuracy, balanced_accuracy"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "b8c70c2e",
   "metadata": {},
   "outputs": [],
   "source": [
    "initial_loss, initial_acc, initial_bal_acc = evaluate(\n",
    "    model, X_val_t, y_val_t, criterion\n",
    ")\n",
    "\n",
    "print(\"Before training:\")\n",
    "print(f\"  validation loss:              {initial_loss:.3f}\")\n",
    "print(f\"  validation accuracy:          {initial_acc:.3f}\")\n",
    "print(f\"  validation balanced accuracy: {initial_bal_acc:.3f}\")"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "38d02a68",
   "metadata": {},
   "source": [
    "# 6. Train the network\n",
    "\n",
    "For **every minibatch**, training follows the same sequence:\n",
    "\n",
    "1. clear the old gradients;\n",
    "2. perform a **forward pass**;\n",
    "3. calculate the loss;\n",
    "4. run **backpropagation**;\n",
    "5. update the weights and biases.\n",
    "\n",
    "One pass through all training examples is an **epoch**.\n",
    "\n",
    "### Before running the next cell\n",
    "\n",
    "Which line will:\n",
    "\n",
    "- calculate the prediction?\n",
    "- calculate the gradient?\n",
    "- actually change the weights?"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "fe5a045b",
   "metadata": {},
   "outputs": [],
   "source": [
    "EPOCHS = 30\n",
    "\n",
    "train_losses = []\n",
    "val_losses = []\n",
    "train_accs = []\n",
    "val_accs = []\n",
    "train_bal_accs = []\n",
    "val_bal_accs = []\n",
    "\n",
    "for epoch in range(EPOCHS):\n",
    "    model.train()\n",
    "\n",
    "    for X_batch, y_batch in train_loader:\n",
    "        optimizer.zero_grad()\n",
    "        logits = model(X_batch)\n",
    "        loss = criterion(logits, y_batch)\n",
    "        loss.backward()\n",
    "        optimizer.step()\n",
    "\n",
    "    train_loss, train_acc, train_bal_acc = evaluate(\n",
    "        model, X_train_t, y_train_t, criterion\n",
    "    )\n",
    "    val_loss, val_acc, val_bal_acc = evaluate(\n",
    "        model, X_val_t, y_val_t, criterion\n",
    "    )\n",
    "\n",
    "    train_losses.append(train_loss)\n",
    "    val_losses.append(val_loss)\n",
    "    train_accs.append(train_acc)\n",
    "    val_accs.append(val_acc)\n",
    "    train_bal_accs.append(train_bal_acc)\n",
    "    val_bal_accs.append(val_bal_acc)\n",
    "\n",
    "    if epoch % 5 == 0 or epoch == EPOCHS - 1:\n",
    "        print(\n",
    "            f\"Epoch {epoch:02d} | \"\n",
    "            f\"train loss {train_loss:.3f} | \"\n",
    "            f\"val loss {val_loss:.3f} | \"\n",
    "            f\"train bal.acc {train_bal_acc:.3f} | \"\n",
    "            f\"val bal.acc {val_bal_acc:.3f}\"\n",
    "        )"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "2b5b34b4",
   "metadata": {},
   "source": [
    "### Learning curves\n",
    "\n",
    "Compare the training and validation curves.\n",
    "\n",
    "Questions:\n",
    "\n",
    "1. Does the network beat the majority-class baseline?\n",
    "2. At what epoch does validation performance begin to plateau?\n",
    "3. Is there evidence of overfitting?"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a4c78f2c",
   "metadata": {},
   "outputs": [],
   "source": [
    "fig, axes = plt.subplots(1, 2, figsize=(11, 4))\n",
    "\n",
    "axes[0].plot(train_losses, label=\"training\")\n",
    "axes[0].plot(val_losses, label=\"validation\")\n",
    "axes[0].set_xlabel(\"Epoch\")\n",
    "axes[0].set_ylabel(\"Cross-entropy loss\")\n",
    "axes[0].set_title(\"Loss\")\n",
    "axes[0].legend()\n",
    "\n",
    "axes[1].plot(train_bal_accs, label=\"training\")\n",
    "axes[1].plot(val_bal_accs, label=\"validation\")\n",
    "axes[1].axhline(1/3, linestyle=\"--\", label=\"chance balanced accuracy\")\n",
    "axes[1].set_xlabel(\"Epoch\")\n",
    "axes[1].set_ylabel(\"Balanced accuracy\")\n",
    "axes[1].set_ylim(0, 1.02)\n",
    "axes[1].set_title(\"Classification performance\")\n",
    "axes[1].legend()\n",
    "\n",
    "plt.tight_layout()\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "9b98c0a3",
   "metadata": {},
   "source": [
    "# 7. Final evaluation on unseen sequences\n",
    "\n",
    "Now we use the held-out **test set**.\n",
    "\n",
    "Because the three classes are not equally frequent, we report both:\n",
    "\n",
    "- ordinary **accuracy**;\n",
    "- **balanced accuracy**, which gives each class equal weight.\n",
    "\n",
    "We also inspect the confusion matrix to see whether donors and acceptors behave differently."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "92b66bc7",
   "metadata": {},
   "outputs": [],
   "source": [
    "test_loss, test_acc, test_bal_acc = evaluate(\n",
    "    model, X_test_t, y_test_t, criterion\n",
    ")\n",
    "\n",
    "print(f\"Test loss:              {test_loss:.3f}\")\n",
    "print(f\"Test accuracy:          {test_acc:.3f}\")\n",
    "print(f\"Test balanced accuracy: {test_bal_acc:.3f}\")\n",
    "\n",
    "model.eval()\n",
    "with torch.no_grad():\n",
    "    test_pred = model(X_test_t).argmax(dim=1).numpy()\n",
    "\n",
    "cm = confusion_matrix(y_test, test_pred)\n",
    "\n",
    "ConfusionMatrixDisplay(\n",
    "    cm,\n",
    "    display_labels=[\"Neither\", \"Donor\", \"Acceptor\"],\n",
    ").plot()\n",
    "\n",
    "plt.title(\"Splice-junction classification\")\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "2a3a2946",
   "metadata": {},
   "source": [
    "# 8. What sequence positions does the MLP use?\n",
    "\n",
    "A neural network can classify well without immediately telling us **why**.\n",
    "\n",
    "We will use a simple perturbation experiment.\n",
    "\n",
    "For each nucleotide position:\n",
    "\n",
    "1. replace that position with an “unknown” base:\n",
    "   \\([0.25,0.25,0.25,0.25]\\);\n",
    "2. recompute test performance;\n",
    "3. ask how much balanced accuracy decreases.\n",
    "\n",
    "If masking a position strongly hurts performance, information at that position was useful to the classifier.\n",
    "\n",
    "This is a simple form of **in-silico perturbation**."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "11274eec",
   "metadata": {},
   "outputs": [],
   "source": [
    "baseline_bal_acc = test_bal_acc\n",
    "\n",
    "X_test_3d = X_test.reshape(-1, 60, 4)\n",
    "importance = []\n",
    "\n",
    "for pos in range(60):\n",
    "    X_masked = X_test_3d.copy()\n",
    "    X_masked[:, pos, :] = 0.25\n",
    "\n",
    "    X_masked_t = torch.tensor(\n",
    "        X_masked.reshape(-1, 60 * 4),\n",
    "        dtype=torch.float32,\n",
    "    )\n",
    "\n",
    "    _, _, masked_bal_acc = evaluate(\n",
    "        model, X_masked_t, y_test_t, criterion\n",
    "    )\n",
    "\n",
    "    importance.append(baseline_bal_acc - masked_bal_acc)\n",
    "\n",
    "relative_position = np.arange(-30, 30)\n",
    "\n",
    "plt.figure(figsize=(12, 4))\n",
    "plt.plot(relative_position, importance)\n",
    "plt.axhline(0, linewidth=1)\n",
    "plt.axvline(0, linestyle=\"--\", linewidth=1)\n",
    "plt.xlabel(\"Position relative to center of sequence window\")\n",
    "plt.ylabel(\"Drop in balanced accuracy when masked\")\n",
    "plt.title(\"Which nucleotide positions matter to the MLP?\")\n",
    "plt.tight_layout()\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "b7af68da",
   "metadata": {},
   "source": [
    "## Compare the network's answer with the biology\n",
    "\n",
    "Finally, let's visualize nucleotide frequencies near the center of donor and acceptor sequences.\n",
    "\n",
    "If the MLP importance plot makes biological sense, its important positions should overlap sequence patterns characteristic of genuine splice junctions."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a4d4392b",
   "metadata": {},
   "outputs": [],
   "source": [
    "def nucleotide_frequencies(sequences):\n",
    "    encoded = np.stack([one_hot_encode(seq) for seq in sequences])\n",
    "    return encoded.mean(axis=0).T\n",
    "\n",
    "fig, axes = plt.subplots(1, 2, figsize=(13, 4), sharey=True)\n",
    "\n",
    "for ax, cls, title in [\n",
    "    (axes[0], \"EI\", \"Donor (exon → intron)\"),\n",
    "    (axes[1], \"IE\", \"Acceptor (intron → exon)\"),\n",
    "]:\n",
    "    freq = nucleotide_frequencies(df.loc[df[\"class\"] == cls, \"sequence\"])\n",
    "    center = freq[:, 20:40]\n",
    "\n",
    "    im = ax.imshow(\n",
    "        center,\n",
    "        aspect=\"auto\",\n",
    "        vmin=0,\n",
    "        vmax=1,\n",
    "        interpolation=\"nearest\",\n",
    "    )\n",
    "    ax.set_yticks(range(4))\n",
    "    ax.set_yticklabels([\"A\", \"C\", \"G\", \"T\"])\n",
    "    ax.set_xticks(range(0, 20, 2))\n",
    "    ax.set_xticklabels(range(-10, 10, 2))\n",
    "    ax.set_xlabel(\"Position relative to center\")\n",
    "    ax.set_title(title)\n",
    "\n",
    "axes[0].set_ylabel(\"Nucleotide\")\n",
    "fig.colorbar(im, ax=axes, label=\"Frequency\")\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "123fcb77",
   "metadata": {},
   "source": [
    "# 9. Discussion\n",
    "\n",
    "You have now trained an MLP to solve a real molecular-biology sequence-recognition problem.\n",
    "\n",
    "### Core questions\n",
    "\n",
    "1. **Did the MLP learn?**  \n",
    "   Compare test performance with the baseline.\n",
    "\n",
    "2. **Did it generalize?**  \n",
    "   Compare training and validation curves.\n",
    "\n",
    "3. **What did it learn biologically?**  \n",
    "   Which positions caused the largest performance drop when masked?\n",
    "\n",
    "4. **Donor vs acceptor:**  \n",
    "   Do their sequence-frequency patterns look identical? Why might recognition of the two boundaries require different sequence information?\n",
    "\n",
    "5. **Why use a hidden layer?**  \n",
    "   A single linear classifier can assign weights to individual nucleotide-position combinations. A hidden layer can combine those inputs into intermediate features before classification.\n",
    "\n",
    "6. **What is missing from this model?**  \n",
    "   Real splicing depends on more than a 60-nt local sequence: distant regulatory elements, RNA-binding proteins, transcriptional context, tissue state, RNA structure, and other features can matter.\n",
    "\n",
    "### Main ML takeaway\n",
    "\n",
    "The training loop is always the same:\n",
    "\n",
    "\\[\n",
    "\\boxed{\n",
    "\\text{minibatch}\n",
    "\\rightarrow\n",
    "\\text{forward pass}\n",
    "\\rightarrow\n",
    "\\text{loss}\n",
    "\\rightarrow\n",
    "\\text{backpropagation}\n",
    "\\rightarrow\n",
    "\\text{parameter update}\n",
    "}\n",
    "\\]\n",
    "\n",
    "The biological representation and prediction target change from problem to problem, but the mechanics of training the MLP do not."
   ]
  },
  {
   "cell_type": "markdown",
   "id": "db18e556",
   "metadata": {},
   "source": [
    "# Optional challenge: does the hidden layer help?\n",
    "\n",
    "If you have extra time, replace the MLP with a model containing **only one linear layer**:\n",
    "\n",
    "```python\n",
    "linear_model = nn.Linear(240, 3)\n",
    "```\n",
    "\n",
    "Train it with the same loss and optimizer.\n",
    "\n",
    "Questions:\n",
    "\n",
    "- How does its validation/test performance compare with the MLP?\n",
    "- Are donor and acceptor sites equally easy for the linear model?\n",
    "- What kinds of sequence rules can a hidden layer represent that a single linear layer cannot?"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "2be5f4e9",
   "metadata": {},
   "source": [
    "# References\n",
    "\n",
    "- **UCI Machine Learning Repository:** *Molecular Biology (Splice-junction Gene Sequences)*, dataset ID 69, DOI: `10.24432/C5M888`.\n",
    "- Alberts *et al.*, *Molecular Biology of the Cell*, section on pre-mRNA splicing.\n",
    "- The MCB128 MLP lecture introduces one-hot representations, fully connected layers, cross-entropy loss, minibatches/epochs, automatic backpropagation, SGD/Adam, and generalization."
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "Python 3",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "name": "python",
   "version": "3"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
