{
 "cells": [
  {
   "cell_type": "markdown",
   "id": "a4799e95",
   "metadata": {},
   "source": [
    "\n",
    "<link rel=\"stylesheet\" href=\"berkeley.css\">\n",
    "\n",
    "<h1 class=\"cal cal-h1\">Lecture 03: PCA and Linear Algebra Review – CS 189, Fall 2026</h1>\n",
    "\n",
    "**Demonstration: making sense of congressional votes.**\n",
    "\n",
    "We use PCA today before we understand it. By the end of this notebook we will have taken a\n",
    "441 x 41 table of votes, reduced it to 441 x 2, and recovered a property of Congress that was\n",
    "never supplied to the algorithm. The remainder of the lecture explains why this works."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "download-lecture-data",
   "metadata": {},
   "outputs": [],
   "source": [
    "# Download the lecture data when it is not already available (for example, in Colab).\n",
    "from pathlib import Path\n",
    "from urllib.request import urlretrieve\n",
    "\n",
    "DATA_BASE_URL = (\n",
    "    \"https://raw.githubusercontent.com/BerkeleyML/fa26-student/\"\n",
    "    \"main/lecture/lec03/data\"\n",
    ")\n",
    "for filename in [\"votes.csv\", \"legislators-2019.yaml\"]:\n",
    "    path = Path(\"data\") / filename\n",
    "    if path.exists():\n",
    "        print(f\"Found {path}\")\n",
    "    else:\n",
    "        path.parent.mkdir(parents=True, exist_ok=True)\n",
    "        urlretrieve(f\"{DATA_BASE_URL}/{filename}\", path)\n",
    "        print(f\"Downloaded {path}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "fb698065",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-02T19:12:34.064561Z",
     "iopub.status.busy": "2026-09-02T19:12:34.064375Z",
     "iopub.status.idle": "2026-09-02T19:12:34.407241Z",
     "shell.execute_reply": "2026-09-02T19:12:34.406060Z"
    }
   },
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "import pandas as pd\n",
    "import yaml\n",
    "from datetime import datetime\n",
    "import plotly.express as px\n",
    "# # Uncomment for HTML Export\n",
    "# import plotly.io as pio\n",
    "# pio.renderers.default = \"notebook_connected\""
   ]
  },
  {
   "cell_type": "markdown",
   "id": "0e14a03c",
   "metadata": {},
   "source": [
    "## Congressional Vote Records\n",
    "\n",
    "Let's examine how the House of Representatives (of the 116th Congress, 1st session) voted in the month of **September 2019**.\n",
    "\n",
    "From the [U.S. Senate website](https://www.senate.gov/reference/Index/Votes.htm):\n",
    "\n",
    "> Roll call votes occur when a representative or senator votes \"yea\" or \"nay,\" so that the names of members voting on each side are recorded. A voice vote is a vote in which those in favor or against a measure say \"yea\" or \"nay,\" respectively, without the names or tallies of members voting on each side being recorded.\n",
    "\n",
    "The data, compiled from ProPublica [source](https://github.com/eyeseast/propublica-congress), is a \"skinny\" table of data where each record is a single vote by a member across any roll call in the 116th Congress, 1st session, as downloaded in February 2020. The member of the House, whom we'll call **legislator**, is denoted by their bioguide alphanumeric ID in http://bioguide.congress.gov/."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "37988ffb",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-02T19:12:34.409975Z",
     "iopub.status.busy": "2026-09-02T19:12:34.409585Z",
     "iopub.status.idle": "2026-09-02T19:12:34.435965Z",
     "shell.execute_reply": "2026-09-02T19:12:34.434820Z"
    }
   },
   "outputs": [],
   "source": [
    "votes = pd.read_csv('data/votes.csv')\n",
    "votes = votes.astype({\"roll call\": str})\n",
    "votes"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "fa52f65f",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-02T19:12:34.438055Z",
     "iopub.status.busy": "2026-09-02T19:12:34.437860Z",
     "iopub.status.idle": "2026-09-02T19:12:34.446714Z",
     "shell.execute_reply": "2026-09-02T19:12:34.445685Z"
    }
   },
   "outputs": [],
   "source": [
    "votes['vote'].value_counts()"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "4bd55dad",
   "metadata": {},
   "source": [
    "This is a \"skinny\" table, with one row per (member, roll call). To treat each legislator as a\n",
    "**datapoint**, we pivot so that each row is a legislator and each column is a roll call.\n",
    "We record a `1` for a Yes vote and a `0` otherwise."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "08b108e5",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-02T19:12:34.448834Z",
     "iopub.status.busy": "2026-09-02T19:12:34.448641Z",
     "iopub.status.idle": "2026-09-02T19:12:34.788568Z",
     "shell.execute_reply": "2026-09-02T19:12:34.787604Z"
    }
   },
   "outputs": [],
   "source": [
    "def was_yes(s):\n",
    "    return 1 if s.iloc[0] == \"Yes\" else 0\n",
    "\n",
    "vote_pivot = votes.pivot_table(index='member',\n",
    "                               columns='roll call',\n",
    "                               values='vote',\n",
    "                               aggfunc=was_yes,\n",
    "                               fill_value=0)\n",
    "print(vote_pivot.shape)\n",
    "vote_pivot.head()"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "dd85e178",
   "metadata": {},
   "source": [
    "So our data matrix $X$ has **441 rows** (legislators) and **41 columns** (roll calls).\n",
    "\n",
    "Each legislator is a point in 41-dimensional space."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "b07aa45d",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-02T19:12:34.791081Z",
     "iopub.status.busy": "2026-09-02T19:12:34.790872Z",
     "iopub.status.idle": "2026-09-02T19:12:34.796765Z",
     "shell.execute_reply": "2026-09-02T19:12:34.795299Z"
    }
   },
   "outputs": [],
   "source": [
    "X = vote_pivot.to_numpy()\n",
    "X.shape"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "11cbcd41",
   "metadata": {},
   "source": [
    "## A first attempt: plotting the raw columns\n",
    "\n",
    "We can only look at two dimensions at a time on a screen. The obvious approach is to select two\n",
    "columns and produce a scatter plot."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "bb814259",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-02T19:12:34.798749Z",
     "iopub.status.busy": "2026-09-02T19:12:34.798545Z",
     "iopub.status.idle": "2026-09-02T19:12:35.027913Z",
     "shell.execute_reply": "2026-09-02T19:12:35.026979Z"
    }
   },
   "outputs": [],
   "source": [
    "px.scatter(vote_pivot, x='555', y='553',\n",
    "           title='Two roll calls at a time', width=700, height=500)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "d09cf311",
   "metadata": {},
   "source": [
    "That is four dots.\n",
    "\n",
    "Every legislator sits at one of four corners, so 441 points collapse onto 4 visible marks. Even\n",
    "if we jittered them apart, there would be $\\binom{41}{2} = 820$ such plots to inspect.\n",
    "\n",
    "**Selecting two of the original columns is not adequate.** We require two *new* columns,\n",
    "constructed from all 41."
   ]
  },
  {
   "cell_type": "markdown",
   "id": "3309bbc8",
   "metadata": {},
   "source": [
    "## PCA in three lines"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a882eec0",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-02T19:12:35.030325Z",
     "iopub.status.busy": "2026-09-02T19:12:35.030066Z",
     "iopub.status.idle": "2026-09-02T19:12:35.982834Z",
     "shell.execute_reply": "2026-09-02T19:12:35.980618Z"
    }
   },
   "outputs": [],
   "source": [
    "from sklearn.decomposition import PCA\n",
    "\n",
    "model = PCA(n_components=2)\n",
    "Z = model.fit_transform(vote_pivot)\n",
    "Z.shape"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "5b8fe044",
   "metadata": {},
   "source": [
    "Each of the 441 legislators is now described by **2 numbers rather than 41**. We plot them\n",
    "below."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "19f2ff7a",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-02T19:12:35.987558Z",
     "iopub.status.busy": "2026-09-02T19:12:35.987154Z",
     "iopub.status.idle": "2026-09-02T19:12:36.030976Z",
     "shell.execute_reply": "2026-09-02T19:12:36.029704Z"
    }
   },
   "outputs": [],
   "source": [
    "px.scatter(x=Z[:, 0], y=Z[:, 1],\n",
    "           labels={'x': 'z1', 'y': 'z2'},\n",
    "           title='Vote data projected onto 2 dimensions',\n",
    "           width=800, height=600, opacity=0.7)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "7dae9972",
   "metadata": {},
   "source": [
    "Two clusters appear, although nothing in the input indicated that clusters should exist. The\n",
    "algorithm never saw a party label, only zeros and ones.\n",
    "\n",
    "We now bring in the identity of each legislator, from\n",
    "[unitedstates/congress-legislators](https://github.com/unitedstates/congress-legislators)."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "10862aa0",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-02T19:12:36.032858Z",
     "iopub.status.busy": "2026-09-02T19:12:36.032668Z",
     "iopub.status.idle": "2026-09-02T19:12:39.375659Z",
     "shell.execute_reply": "2026-09-02T19:12:39.374617Z"
    }
   },
   "outputs": [],
   "source": [
    "# Static copy of the 2019 roster, so it matches our voting data.\n",
    "legislators_data = yaml.safe_load(open('data/legislators-2019.yaml'))\n",
    "\n",
    "def to_date(s):\n",
    "    return datetime.strptime(s, '%Y-%m-%d')\n",
    "\n",
    "legs = pd.DataFrame(\n",
    "    columns=['leg_id', 'first', 'last', 'state', 'chamber', 'party', 'birthday'],\n",
    "    data=[[x['id']['bioguide'],\n",
    "           x['name']['first'],\n",
    "           x['name']['last'],\n",
    "           x['terms'][-1]['state'],\n",
    "           x['terms'][-1]['type'],\n",
    "           x['terms'][-1]['party'],\n",
    "           to_date(x['bio']['birthday'])] for x in legislators_data])\n",
    "legs['age'] = 2024 - legs['birthday'].dt.year\n",
    "legs.head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "959a63e7",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-02T19:12:39.377848Z",
     "iopub.status.busy": "2026-09-02T19:12:39.377591Z",
     "iopub.status.idle": "2026-09-02T19:12:39.392046Z",
     "shell.execute_reply": "2026-09-02T19:12:39.390840Z"
    }
   },
   "outputs": [],
   "source": [
    "vote_2d = pd.DataFrame(Z, index=vote_pivot.index, columns=['z1', 'z2'])\n",
    "vote_2d = vote_2d.join(legs.set_index('leg_id'))\n",
    "vote_2d.head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "15b85e70",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-02T19:12:39.394890Z",
     "iopub.status.busy": "2026-09-02T19:12:39.394612Z",
     "iopub.status.idle": "2026-09-02T19:12:39.451997Z",
     "shell.execute_reply": "2026-09-02T19:12:39.451131Z"
    }
   },
   "outputs": [],
   "source": [
    "px.scatter(vote_2d, x='z1', y='z2', color='party',\n",
    "           title='Vote data, colored by party (PCA never saw this column)',\n",
    "           width=800, height=600, opacity=0.7,\n",
    "           color_discrete_map={'Democrat': 'blue', 'Republican': 'red', 'Independent': 'green'},\n",
    "           hover_data=['first', 'last', 'state'],\n",
    "           render_mode='svg')"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "0af87923",
   "metadata": {},
   "source": [
    "The structure that PCA found corresponds to party affiliation.\n",
    "\n",
    "There is substantial overplotting, since many legislators vote identically and therefore land on\n",
    "exactly the same point. We add jitter to reveal the density."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "37c9b0ec",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-02T19:12:39.454491Z",
     "iopub.status.busy": "2026-09-02T19:12:39.454163Z",
     "iopub.status.idle": "2026-09-02T19:12:39.512206Z",
     "shell.execute_reply": "2026-09-02T19:12:39.510838Z"
    }
   },
   "outputs": [],
   "source": [
    "rng = np.random.default_rng(42)\n",
    "vote_2d['z1_jittered'] = vote_2d['z1'] + rng.normal(0, 0.1, len(vote_2d))\n",
    "vote_2d['z2_jittered'] = vote_2d['z2'] + rng.normal(0, 0.1, len(vote_2d))\n",
    "\n",
    "px.scatter(vote_2d, x='z1_jittered', y='z2_jittered', color='party', size='age',\n",
    "           title='Vote data (jittered)',\n",
    "           width=800, height=600, opacity=0.7, size_max=10,\n",
    "           color_discrete_map={'Democrat': 'blue', 'Republican': 'red', 'Independent': 'green'},\n",
    "           hover_data=['first', 'last', 'state', 'party'])"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "b688fefc",
   "metadata": {},
   "source": [
    "How far does this go? If we use only the **sign of the first coordinate**, how often does it\n",
    "agree with party affiliation?"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "5f232a75",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-02T19:12:39.514603Z",
     "iopub.status.busy": "2026-09-02T19:12:39.514377Z",
     "iopub.status.idle": "2026-09-02T19:12:39.523595Z",
     "shell.execute_reply": "2026-09-02T19:12:39.522535Z"
    }
   },
   "outputs": [],
   "source": [
    "labeled = vote_2d.dropna(subset=['party'])\n",
    "guess = np.where(labeled['z1'] > 0, 'Democrat', 'Republican')\n",
    "(guess == labeled['party']).mean()"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "9ae439ff",
   "metadata": {},
   "source": [
    "A single number per legislator, derived from all 41, recovers party affiliation for\n",
    "approximately 98% of the House."
   ]
  },
  {
   "cell_type": "markdown",
   "id": "438e2369",
   "metadata": {},
   "source": [
    "## How much information was discarded?\n",
    "\n",
    "We replaced 41 columns with 2. The attribute `explained_variance_ratio_` reports the fraction of\n",
    "the spread in the data accounted for by each new coordinate."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "21bc8d95",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-02T19:12:39.526249Z",
     "iopub.status.busy": "2026-09-02T19:12:39.526043Z",
     "iopub.status.idle": "2026-09-02T19:12:39.531742Z",
     "shell.execute_reply": "2026-09-02T19:12:39.530840Z"
    }
   },
   "outputs": [],
   "source": [
    "model.explained_variance_ratio_"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2e36d24f",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-02T19:12:39.533890Z",
     "iopub.status.busy": "2026-09-02T19:12:39.533662Z",
     "iopub.status.idle": "2026-09-02T19:12:39.538938Z",
     "shell.execute_reply": "2026-09-02T19:12:39.537995Z"
    }
   },
   "outputs": [],
   "source": [
    "model.explained_variance_ratio_.sum()"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "03626fdc",
   "metadata": {},
   "source": [
    "The first coordinate alone accounts for roughly 80%. The full profile is shown below."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "99c8b5ea",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-02T19:12:39.540823Z",
     "iopub.status.busy": "2026-09-02T19:12:39.540643Z",
     "iopub.status.idle": "2026-09-02T19:12:39.588226Z",
     "shell.execute_reply": "2026-09-02T19:12:39.587070Z"
    }
   },
   "outputs": [],
   "source": [
    "model10 = PCA(n_components=10).fit(vote_pivot)\n",
    "\n",
    "px.line(y=model10.explained_variance_ratio_, markers=True,\n",
    "        labels={'x': 'component', 'y': 'fraction of total spread'},\n",
    "        title='Spread accounted for by each component',\n",
    "        width=700, height=450)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "f9052139",
   "metadata": {},
   "source": [
    "There is a sharp decrease after the first component, followed by a long flat tail. This shape is\n",
    "what makes the two-dimensional plot trustworthy: no third direction accounts for an appreciable\n",
    "share of the spread.\n",
    "\n",
    "Note that the data matrix is *not* actually low rank:"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "f0753937",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-02T19:12:39.590134Z",
     "iopub.status.busy": "2026-09-02T19:12:39.589955Z",
     "iopub.status.idle": "2026-09-02T19:12:39.594543Z",
     "shell.execute_reply": "2026-09-02T19:12:39.593898Z"
    }
   },
   "outputs": [],
   "source": [
    "np.linalg.matrix_rank(X - X.mean(axis=0))"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "c006089b",
   "metadata": {},
   "source": [
    "The rank is 41, which is full. This is therefore not a case of exactly redundant columns that\n",
    "may be deleted. The data is only **approximately** low dimensional, and that distinction is\n",
    "central to what follows."
   ]
  },
  {
   "cell_type": "markdown",
   "id": "8bff6cfe",
   "metadata": {},
   "source": [
    "## What the model consists of\n",
    "\n",
    "The call to `fit` estimated something. We inspect it below."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "60d47975",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-02T19:12:39.596076Z",
     "iopub.status.busy": "2026-09-02T19:12:39.595790Z",
     "iopub.status.idle": "2026-09-02T19:12:39.599253Z",
     "shell.execute_reply": "2026-09-02T19:12:39.598625Z"
    }
   },
   "outputs": [],
   "source": [
    "model.components_.shape"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "de9800b4",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-02T19:12:39.600732Z",
     "iopub.status.busy": "2026-09-02T19:12:39.600431Z",
     "iopub.status.idle": "2026-09-02T19:12:39.604525Z",
     "shell.execute_reply": "2026-09-02T19:12:39.603715Z"
    }
   },
   "outputs": [],
   "source": [
    "model.components_"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "04fa69b4",
   "metadata": {},
   "source": [
    "**Two rows of 41 numbers.** This is the entire model.\n",
    "\n",
    "Each row is a set of weights over the 41 roll calls, and a legislator's new coordinates are the\n",
    "dot products of their voting record with these two rows:"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2c3563a0",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-02T19:12:39.605850Z",
     "iopub.status.busy": "2026-09-02T19:12:39.605635Z",
     "iopub.status.idle": "2026-09-02T19:12:39.609105Z",
     "shell.execute_reply": "2026-09-02T19:12:39.608527Z"
    }
   },
   "outputs": [],
   "source": [
    "w1 = model.components_[0]\n",
    "Xc = X - X.mean(axis=0)          # PCA centers the data internally\n",
    "\n",
    "# sklearn's answer for the first legislator, vs. a dot product we compute ourselves\n",
    "print(Z[0, 0], Xc[0] @ w1)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "6259eb25",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-02T19:12:39.611420Z",
     "iopub.status.busy": "2026-09-02T19:12:39.611221Z",
     "iopub.status.idle": "2026-09-02T19:12:39.650880Z",
     "shell.execute_reply": "2026-09-02T19:12:39.650217Z"
    }
   },
   "outputs": [],
   "source": [
    "px.bar(x=vote_pivot.columns, y=w1,\n",
    "       labels={'x': 'roll call', 'y': 'weight in the first component'},\n",
    "       title='The first row of components_',\n",
    "       width=900, height=400)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "0fe64144",
   "metadata": {},
   "outputs": [],
   "source": [
    "# Center each column: subtract the House-wide Yes rate for that roll call.\n",
    "vote_pivot_centered = vote_pivot - vote_pivot.mean()\n",
    "\n",
    "# Per party, the average deviation from the House-wide Yes rate on each roll call.\n",
    "party_yes_deviation = vote_pivot_centered.join(labeled['party']).groupby('party').mean()\n",
    "\n",
    "party_yes_deviation_long = (party_yes_deviation\n",
    "                            .reset_index()\n",
    "                            .melt(id_vars='party',\n",
    "                                  var_name='roll call',\n",
    "                                  value_name='Yes Rate Centered'))\n",
    "\n",
    "fig = px.bar(party_yes_deviation_long,\n",
    "             x='roll call', y='Yes Rate Centered',\n",
    "             facet_row='party', color='party',\n",
    "             color_discrete_map={'Democrat': 'blue', 'Republican': 'red', 'Independent': 'green'},\n",
    "             title='Party Yes rate relative to the House average, by roll call',\n",
    "             width=900, height=800)\n",
    "fig.for_each_annotation(lambda a: a.update(text=a.text.split('=')[-1]))  # 'party=Democrat' -> 'Democrat'\n",
    "fig.update_layout(showlegend=False)  # the facet titles already name the party\n",
    "fig"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "a145522e",
   "metadata": {},
   "source": [
    "The method therefore reduces to a single question: **how do we find the right $k$ rows of\n",
    "length $d$?**\n",
    "\n",
    "- How are those rows determined?\n",
    "- In what sense are they the *optimal* choice?\n",
    "- Why does projecting onto them preserve the structure of interest?\n",
    "\n",
    "These are the subject of the remainder of the lecture."
   ]
  },
  {
   "cell_type": "markdown",
   "id": "f35e1afc",
   "metadata": {},
   "source": [
    "## Compressing images: PCA on Fashion-MNIST\n",
    "\n",
    "The congressional votes were 41-dimensional. We now apply the same three lines to a\n",
    "dataset where each datapoint has 784 dimensions and, unlike a voting record, can be looked\n",
    "at directly. [Fashion-MNIST](https://github.com/zalandoresearch/fashion-mnist) is 60,000\n",
    "grayscale 28 x 28 photographs of clothing in 10 categories."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "b1c446c3",
   "metadata": {},
   "outputs": [],
   "source": [
    "# Fetch the Data\n",
    "import torchvision\n",
    "data = torchvision.datasets.FashionMNIST(root='data', train=True, download=True)\n",
    "\n",
    "# Preprocess the data into numpy arrays\n",
    "images = data.data.numpy().astype(float)\n",
    "targets = data.targets.numpy() # integer encoding of class labels\n",
    "class_dict = {i:class_name for i,class_name in enumerate(data.classes)}\n",
    "labels = np.array([class_dict[t] for t in targets]) # raw class labels\n",
    "n = len(images)\n",
    "\n",
    "print(\"Loaded FashionMNIST dataset with {} samples.\".format(n))\n",
    "print(\"Classes: {}\".format(class_dict))\n",
    "print(\"Image shape: {}\".format(images[0].shape))\n",
    "print(\"Image dtype: {}\".format(images[0].dtype))\n",
    "print(\"Image 0:\\n\", images[0])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "e156f4d3",
   "metadata": {},
   "outputs": [],
   "source": [
    "px.imshow(images[0], color_continuous_scale='gray_r') "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "0f3ba1be",
   "metadata": {},
   "outputs": [],
   "source": [
    "def show_images(images, max_images=40, ncols=5, labels = None):\n",
    "    \"\"\"Visualize a subset of images from the dataset.\n",
    "    Args:\n",
    "        images (np.ndarray): Array of images to visualize [img,row,col].\n",
    "        max_images (int): Maximum number of images to display.\n",
    "        ncols (int): Number of columns in the grid.\n",
    "        labels (np.ndarray, optional): Labels for the images, used for facet titles.\n",
    "    Returns:\n",
    "        plotly.graph_objects.Figure: A Plotly figure object containing the images.\n",
    "    \"\"\"\n",
    "    n = min(images.shape[0], max_images) # number of images to show\n",
    "    px_height = 220 # height of each image in pixels\n",
    "    fig = px.imshow(images[:n, :, :], color_continuous_scale='gray_r', \n",
    "                    facet_col = 0, facet_col_wrap=ncols,\n",
    "                    height = px_height * int(np.ceil(n/ncols)))\n",
    "    fig.update_layout(coloraxis_showscale=False)\n",
    "    if labels is not None:\n",
    "        # Extract the facet number and replace with the label.\n",
    "        fig.for_each_annotation(lambda a: a.update(text=labels[int(a.text.split(\"=\")[-1])]))\n",
    "    return fig"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "c12e0cca",
   "metadata": {},
   "outputs": [],
   "source": [
    "show_images(images, 20, labels=labels)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "1439da45",
   "metadata": {},
   "source": [
    "### Each image is a datapoint in 784-dimensional space\n",
    "\n",
    "The voting data gave us one row per legislator and one column per roll call. We do the same\n",
    "thing here: one row per image, one column per **pixel**. A 28 x 28 image becomes a single row of\n",
    "$28 \\times 28 = 784$ numbers."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "0d35dc23",
   "metadata": {},
   "outputs": [],
   "source": [
    "X_img = images.reshape(n, -1)   # (60000, 28, 28) -> (60000, 784)\n",
    "X_img.shape"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "e4bcf2c0",
   "metadata": {},
   "source": [
    "PCA centers the data before it does anything else, so the first thing it computes is the mean of\n",
    "those 60,000 rows. Reshaped back to 28 x 28, the mean is itself an image."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "f1373ec2",
   "metadata": {},
   "outputs": [],
   "source": [
    "mean_image = X_img.mean(axis=0)\n",
    "\n",
    "px.imshow(mean_image.reshape(28, 28), color_continuous_scale='gray_r',\n",
    "          title='The average of all 60,000 images', width=400, height=400)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "bf821648",
   "metadata": {},
   "source": [
    "### The principal components are also images\n",
    "\n",
    "For the votes, `components_` was a $k \\times 41$ matrix, and each row was a set of weights over\n",
    "the 41 roll calls. Here it is a $k \\times 784$ matrix, and each row is a set of weights over the\n",
    "784 pixels. **A row of 784 numbers can be reshaped into a 28 x 28 picture**, so we can look\n",
    "directly at the model.\n",
    "\n",
    "We fit 200 components once, and use the leading $k$ of them below. Because the components are\n",
    "nested, `components_[:k]` is exactly what `PCA(n_components=k)` would have produced."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "3ba07ee3",
   "metadata": {},
   "outputs": [],
   "source": [
    "pca_img = PCA(n_components=200).fit(X_img)\n",
    "pca_img.components_.shape"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "ab790f19",
   "metadata": {},
   "outputs": [],
   "source": [
    "n_show = 10\n",
    "comps = pca_img.components_[:n_show].reshape(n_show, 28, 28)\n",
    "\n",
    "fig = px.imshow(comps, facet_col=0, facet_col_wrap=5,\n",
    "                color_continuous_scale='RdBu_r', color_continuous_midpoint=0,\n",
    "                height=440, title='The first 10 principal components, viewed as images')\n",
    "fig.for_each_annotation(lambda a: a.update(text=f\"PC {int(a.text.split('=')[-1]) + 1}\"))\n",
    "fig.update_layout(coloraxis_showscale=False)\n",
    "fig"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "754a1a4c",
   "metadata": {},
   "source": [
    "Red is a positive weight and blue is a negative one. These are not garments; they are\n",
    "**contrasts**. PC 1 separates wide dark regions from narrow ones, which is roughly the\n",
    "distinction between a shirt and a shoe. Later components encode sleeves, straps, and the gap\n",
    "between trouser legs. The overall sign of each component is arbitrary.\n",
    "\n",
    "### How many components do we need?"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "4d77c38c",
   "metadata": {},
   "outputs": [],
   "source": [
    "cumvar = np.cumsum(pca_img.explained_variance_ratio_)\n",
    "\n",
    "px.line(x=np.arange(1, 201), y=cumvar,\n",
    "        labels={'x': 'number of components k', 'y': 'cumulative fraction of spread'},\n",
    "        title='Spread captured by the first k components', width=750, height=450)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "bd8a9a3c",
   "metadata": {},
   "source": [
    "The curve rises steeply and then flattens: 50 of the 784 directions account for 86% of the\n",
    "spread, and 200 account for 95%. Compare this to the votes, where a *single* component captured\n",
    "80%. Images are low dimensional, but not nearly as aggressively so.\n",
    "\n",
    "### Reconstruction\n",
    "\n",
    "Compression is only useful if we can get the image back. Keeping $k$ scores and then\n",
    "undoing the projection gives\n",
    "\n",
    "$$\\hat{x} = \\bar{x} + \\sum_{j=1}^{k} z_j w_j$$\n",
    "\n",
    "a picture rebuilt as the mean image plus a weighted sum of $k$ component images."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "36a39591",
   "metadata": {},
   "outputs": [],
   "source": [
    "Z_img = pca_img.transform(X_img)   # (60000, 200) scores\n",
    "\n",
    "def reconstruct(k, rows):\n",
    "    \"\"\"Rebuild images from only their first k scores.\"\"\"\n",
    "    return Z_img[rows, :k] @ pca_img.components_[:k] + mean_image"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "ca5dcf64",
   "metadata": {},
   "outputs": [],
   "source": [
    "ks = [1, 2, 5, 10, 25, 50, 100, 200]\n",
    "i = 0   # the ankle boot from the top of this section\n",
    "\n",
    "ladder = np.vstack([X_img[i]] + [reconstruct(k, [i]) for k in ks])\n",
    "\n",
    "# Reconstructions can fall slightly outside [0, 255], so we clip them for display.\n",
    "show_images(np.clip(ladder, 0, 255).reshape(-1, 28, 28), ncols=3,\n",
    "            labels=['original (784)'] + [f'k = {k}' for k in ks])"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "f52efb12",
   "metadata": {},
   "source": [
    "One number produces a dark blob. Ten produce something identifiable as a boot. By 50 the\n",
    "silhouette and the shading are right, and the remaining 734 dimensions mostly carry\n",
    "texture and noise.\n",
    "\n",
    "Below we do the same at $k = 50$ for eight random images. The top row is the original,\n",
    "the bottom row is 50 numbers."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "1059e1e0",
   "metadata": {},
   "outputs": [],
   "source": [
    "rng_img = np.random.default_rng(189)\n",
    "rows = rng_img.choice(n, 8, replace=False)\n",
    "k = 50\n",
    "\n",
    "side_by_side = np.vstack([X_img[rows], reconstruct(k, rows)])\n",
    "show_images(np.clip(side_by_side, 0, 255).reshape(-1, 28, 28), ncols=8,\n",
    "            labels=[labels[r] for r in rows] + [f'k = {k}' for _ in rows])"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "2e3d3cb8",
   "metadata": {},
   "source": [
    "### What did that actually save?\n",
    "\n",
    "To store the whole dataset we need the $n \\times k$ table of scores, plus the basis we need in\n",
    "order to decode it: the $k \\times 784$ components and the 784-pixel mean. The basis is paid for\n",
    "**once**, no matter how many images we compress."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "058be6df",
   "metadata": {},
   "outputs": [],
   "source": [
    "d = X_img.shape[1]\n",
    "\n",
    "for k in [10, 50, 100]:\n",
    "    stored = n * k + k * d + d\n",
    "    rmse = np.sqrt(((reconstruct(k, slice(None)) - X_img) ** 2).mean())\n",
    "    print(f\"k = {k:3d} | scores {n*k:>9,} + basis {k*d + d:>7,} = {stored:>9,} numbers \"\n",
    "          f\"| {n*d/stored:5.1f}x smaller | RMSE {rmse:5.1f}\")"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "324aa31f",
   "metadata": {},
   "source": [
    "At $k = 50$ the dataset is **15x smaller** and the images survive. The basis is only 1.3% of the\n",
    "stored bytes, so essentially all of the cost is the 50 numbers per image.\n",
    "\n",
    "This is lossy compression of the same general kind as JPEG. The difference is that JPEG uses a\n",
    "fixed, universal basis (cosines), while PCA **learns** a basis from this particular collection of\n",
    "images. That is why the components above look like clothing contrasts rather than generic\n",
    "ripples, and it is also why the basis only compresses images that resemble the training set."
   ]
  },
  {
   "cell_type": "markdown",
   "id": "32307715",
   "metadata": {},
   "source": [
    "## Can we run the decoder backwards to invent new clothes?\n",
    "\n",
    "The reconstruction step $\\hat{x} = \\bar{x} + \\sum_j z_j w_j$ turns 50 numbers into an image, and\n",
    "it does not care where those numbers came from. So here is a tempting idea: instead of taking\n",
    "$z$ from a real image, **make $z$ up**, and see what comes out.\n",
    "\n",
    "For this to work, the made-up $z$ has to look like the $z$ of a real image. So first we\n",
    "look at how the real scores are distributed."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "46a0f2dd",
   "metadata": {},
   "outputs": [],
   "source": [
    "sub = np.random.default_rng(189).choice(n, 4000, replace=False)\n",
    "scores_2d = pd.DataFrame({'z1': Z_img[sub, 0], 'z2': Z_img[sub, 1], 'class': labels[sub]})\n",
    "\n",
    "px.scatter(scores_2d, x='z1', y='z2', color='class', opacity=0.6,\n",
    "           title='The first two scores of 4,000 images', width=850, height=600)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "a81e8c87",
   "metadata": {},
   "source": [
    "This is not one cloud. Footwear sits on the left, tops sit on the upper right, trousers hang\n",
    "below, bags sit on top. The score distribution is **multimodal**, and there are wide empty\n",
    "regions between the groups.\n",
    "\n",
    "Let us ignore that for a moment and do the simplest thing: fit a single Gaussian to the 50\n",
    "scores, draw from it, and decode."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "56765448",
   "metadata": {},
   "outputs": [],
   "source": [
    "k = 50\n",
    "Zk = Z_img[:, :k]\n",
    "\n",
    "def sample_images(Z_ref, n_samples, rng):\n",
    "    \"\"\"Fit one Gaussian to the scores in Z_ref, draw from it, and decode into images.\"\"\"\n",
    "    draws = rng.multivariate_normal(Z_ref.mean(axis=0), np.cov(Z_ref.T), n_samples)\n",
    "    return draws @ pca_img.components_[:Z_ref.shape[1]] + mean_image\n",
    "\n",
    "fake = sample_images(Zk, 16, np.random.default_rng(0))\n",
    "show_images(np.clip(fake, 0, 255).reshape(-1, 28, 28), ncols=8, max_images=16)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "20ddf8ca",
   "metadata": {},
   "source": [
    "These are garment-shaped smudges. Several are two items at once: a sleeve fading into a trouser\n",
    "leg, a shoe ghosted over a shirt. They have the *statistics* of the dataset without being\n",
    "plausible members of it.\n",
    "\n",
    "The scatter plot above explains why. A single Gaussian is one blob, so most of its mass lands in\n",
    "the empty space *between* the clusters, and a point halfway between \"sneaker\" and \"pullover\"\n",
    "decodes to a superposition of the two. We can check that the samples really are landing in\n",
    "unoccupied territory by measuring how far each one is from the nearest real image."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "58271150",
   "metadata": {},
   "outputs": [],
   "source": [
    "rng_nn = np.random.default_rng(1)\n",
    "reference = Zk[rng_nn.choice(n, 8000, replace=False)]\n",
    "real_pts  = Zk[rng_nn.choice(n, 200, replace=False)]\n",
    "fake_pts  = rng_nn.multivariate_normal(Zk.mean(axis=0), np.cov(Zk.T), 200)\n",
    "\n",
    "def nn_distance(query, reference):\n",
    "    \"\"\"Distance from each query point to its closest neighbour in reference.\"\"\"\n",
    "    sq = ((query[:, None, :] - reference[None, :, :]) ** 2).sum(axis=2)\n",
    "    return np.sqrt(sq.min(axis=1))\n",
    "\n",
    "print(f\"real image  -> nearest real image: {np.median(nn_distance(real_pts, reference)):.0f}\")\n",
    "print(f\"fake sample -> nearest real image: {np.median(nn_distance(fake_pts, reference)):.0f}\")"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "49f731b3",
   "metadata": {},
   "source": [
    "A real image has a real neighbour roughly twice as close. The samples are not near the data;\n",
    "they are in the gaps.\n",
    "\n",
    "The fix follows directly from the diagnosis. The problem was fitting **one** blob to **ten**\n",
    "clusters, so we fit one Gaussian per class instead, still in the same 50-dimensional score\n",
    "space, and still decoding with the same components."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cd40fd52",
   "metadata": {},
   "outputs": [],
   "source": [
    "rng_cls = np.random.default_rng(0)\n",
    "chosen = [0, 7, 8, 1]   # T-shirt/top, Sneaker, Bag, Trouser\n",
    "\n",
    "per_class = np.vstack([sample_images(Zk[targets == c], 4, rng_cls) for c in chosen])\n",
    "\n",
    "show_images(np.clip(per_class, 0, 255).reshape(-1, 28, 28), ncols=4, max_images=16,\n",
    "            labels=[class_dict[c] for c in chosen for _ in range(4)])"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "e3876337",
   "metadata": {},
   "source": [
    "**These work.** Each row is four garments that do not exist in the dataset, and they are\n",
    "recognizably sneakers, t-shirts, bags, and trousers, with varied heights, widths, and shading.\n",
    "They are blurry, because 50 components cannot represent a sharp edge and because a Gaussian is\n",
    "still only an approximation of a class, but they are plausible items rather than superpositions.\n",
    "\n",
    "So the honest answer is: **the decoder is fine, and the hard part is knowing which $z$ to\n",
    "feed it.** Sampling works exactly as well as our model of the score distribution does.\n",
    "\n",
    "### One caution about what the subspace can do\n",
    "\n",
    "It is tempting to read the score space as a space of *concepts*, where moving from one image to\n",
    "another should morph a shirt into a shoe. It cannot, and the reason is that the map from $z$ to\n",
    "pixels is **linear**. Interpolating between two images in score space is algebraically identical\n",
    "to cross-fading the two images in pixel space."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "0f1aea8b",
   "metadata": {},
   "outputs": [],
   "source": [
    "a = np.where(targets == 0)[0][0]   # a t-shirt\n",
    "b = np.where(targets == 7)[0][0]   # a sneaker\n",
    "\n",
    "t = np.linspace(0, 1, 8)[:, None]\n",
    "path = (1 - t) * Zk[a] + t * Zk[b]\n",
    "blend = path @ pca_img.components_[:k] + mean_image\n",
    "\n",
    "show_images(np.clip(blend, 0, 255).reshape(-1, 28, 28), ncols=8, max_images=8)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "b197cb11",
   "metadata": {},
   "outputs": [],
   "source": [
    "# The same path, computed instead by blending pixels and then projecting. Identical.\n",
    "pixel_blend = (1 - t) * X_img[a] + t * X_img[b]\n",
    "projected = (pixel_blend - mean_image) @ pca_img.components_[:k].T @ pca_img.components_[:k] + mean_image\n",
    "\n",
    "np.abs(blend - projected).max()"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "f669b050",
   "metadata": {},
   "source": [
    "The shirt does not become a shoe; it dissolves while a shoe appears underneath. Linearity is\n",
    "what makes PCA cheap to fit, easy to interpret, and provably optimal in the sense we are about to\n",
    "define, and it is also precisely what stops it from being a generative model of images. Getting\n",
    "a genuine morph requires a decoder that is *not* restricted to a linear subspace.\n",
    "\n",
    "### Summary of this section\n",
    "\n",
    "- Each image is a point in $\\mathbb{R}^{784}$; PCA finds a $k$-dimensional subspace it nearly lies in.\n",
    "- The components are pictures, and reconstruction is the mean image plus a weighted sum of them.\n",
    "- $k = 50$ compresses the dataset 15x with the content of the images intact.\n",
    "- Decoding invented scores does generate new clothing, but only once the score distribution is\n",
    "  modeled per class. A single Gaussian samples the empty space between clusters.\n",
    "- The subspace is linear, so interpolation is a cross-fade, not a morph."
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "py_3_11",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "codemirror_mode": {
    "name": "ipython",
    "version": 3
   },
   "file_extension": ".py",
   "mimetype": "text/x-python",
   "name": "python",
   "nbconvert_exporter": "python",
   "pygments_lexer": "ipython3",
   "version": "3.11.13"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
