{
 "cells": [
  {
   "cell_type": "markdown",
   "id": "34ec0f8d",
   "metadata": {},
   "source": [
    "<link rel=\"stylesheet\" href=\"berkeley.css\">\n",
    "\n",
    "<h1 class=\"cal cal-h1\">Lecture 05: Density Estimation and Gaussian Mixture Models – CS 189, Fall 2026</h1>\n",
    "\n",
    "**Accompanying demonstration: fitting continuous data, one Gaussian and then several.**\n",
    "\n",
    "Lecture 4 estimated the parameters of a Bernoulli, a distribution over a single binary variable.\n",
    "This notebook applies the same machinery to continuous data. It fits a Gaussian by maximum\n",
    "likelihood, shows that the resulting variance is biased, then meets data that no single Gaussian\n",
    "describes and fits a mixture with the EM algorithm.\n",
    "\n",
    "Everything is one-dimensional, matching the slides. The EM implementation in section 6 is written\n",
    "for general $D$, so it also works on multivariate data, but every example here uses $D = 1$."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "1700c334",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-09T06:49:41.938212Z",
     "iopub.status.busy": "2026-09-09T06:49:41.937945Z",
     "iopub.status.idle": "2026-09-09T06:49:42.850490Z",
     "shell.execute_reply": "2026-09-09T06:49:42.850248Z"
    }
   },
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "import pandas as pd\n",
    "import plotly.express as px\n",
    "import plotly.graph_objects as go\n",
    "from scipy.stats import multivariate_normal\n",
    "\n",
    "rng = np.random.default_rng(0)\n",
    "\n",
    "BLUE, GOLD, GREEN, ORANGE = '#002675', '#FDB515', '#028842', '#DF5327'\n",
    "\n",
    "# A separate stream for building the dataset, so the figures on the slides and the\n",
    "# numbers here stay identical no matter what else is drawn along the way.\n",
    "data_rng = np.random.default_rng(189)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "95760990",
   "metadata": {},
   "source": [
    "## 1. Fitting a Gaussian by maximum likelihood\n",
    "\n",
    "The feature is the loudness of a one-second audio segment in decibels. The data is simulated,\n",
    "since a recording of the lecture hall cannot be distributed with the notebook, but the shape is\n",
    "representative of a quiet room."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "7fa2d8cc",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-09T06:49:42.852032Z",
     "iopub.status.busy": "2026-09-09T06:49:42.851906Z",
     "iopub.status.idle": "2026-09-09T06:49:43.076299Z",
     "shell.execute_reply": "2026-09-09T06:49:43.076003Z"
    }
   },
   "outputs": [],
   "source": [
    "loudness = data_rng.normal(-52, 4.0, size=4000)     # background segments, in dB\n",
    "\n",
    "px.histogram(loudness, nbins=60, histnorm='probability density',\n",
    "             color_discrete_sequence=[BLUE], labels={'value': 'loudness (dB)'},\n",
    "             title='Loudness of 4,000 background segments',\n",
    "             width=850, height=430, template='plotly_white')"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "5f6d2d7a",
   "metadata": {},
   "source": [
    "The distribution is unimodal and roughly symmetric, which suggests a Gaussian. The maximum\n",
    "likelihood estimates derived in lecture are the sample mean and the sample variance."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "18c8bf5a",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-09T06:49:43.077731Z",
     "iopub.status.busy": "2026-09-09T06:49:43.077634Z",
     "iopub.status.idle": "2026-09-09T06:49:43.079755Z",
     "shell.execute_reply": "2026-09-09T06:49:43.079517Z"
    }
   },
   "outputs": [],
   "source": [
    "mu_ml = loudness.mean()\n",
    "var_ml = ((loudness - mu_ml) ** 2).mean()        # note: divide by N, not N - 1\n",
    "\n",
    "print(f'mu_ML     = {mu_ml:.3f}')\n",
    "print(f'sigma2_ML = {var_ml:.3f}   (sigma = {np.sqrt(var_ml):.3f})')"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "d5a82d48",
   "metadata": {},
   "source": [
    "Two numbers, computed in a single pass. They are the sufficient statistics: given $\\sum_n x_n$ and\n",
    "$\\sum_n x_n^2$, the 4,000 observations are not needed to write down the fitted density."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "e059e689",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-09T06:49:43.081069Z",
     "iopub.status.busy": "2026-09-09T06:49:43.080982Z",
     "iopub.status.idle": "2026-09-09T06:49:43.107907Z",
     "shell.execute_reply": "2026-09-09T06:49:43.107620Z"
    }
   },
   "outputs": [],
   "source": [
    "def gaussian(x, mu, var):\n",
    "    return np.exp(-0.5 * (x - mu) ** 2 / var) / np.sqrt(2 * np.pi * var)\n",
    "\n",
    "grid = np.linspace(loudness.min() - 3, loudness.max() + 3, 400)\n",
    "\n",
    "fig = px.histogram(loudness, nbins=60, histnorm='probability density',\n",
    "                   color_discrete_sequence=['#c9d3e6'], labels={'value': 'loudness (dB)'})\n",
    "fig.add_trace(go.Scatter(x=grid, y=gaussian(grid, mu_ml, var_ml),\n",
    "                         line=dict(color=BLUE, width=3), name='fitted Gaussian'))\n",
    "fig.update_layout(title='Maximum likelihood fit', width=850, height=430,\n",
    "                  template='plotly_white', showlegend=False)\n",
    "fig"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "16ffce2a",
   "metadata": {},
   "source": [
    "### The optimization surface\n",
    "\n",
    "The closed form conceals a two-dimensional optimization over $\\mu$ and $\\sigma$. Evaluating the\n",
    "log likelihood on a grid shows the surface whose maximum it locates."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "32f72841",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-09T06:49:43.109216Z",
     "iopub.status.busy": "2026-09-09T06:49:43.109108Z",
     "iopub.status.idle": "2026-09-09T06:49:43.124777Z",
     "shell.execute_reply": "2026-09-09T06:49:43.124546Z"
    }
   },
   "outputs": [],
   "source": [
    "mus = np.linspace(mu_ml - 1.0, mu_ml + 1.0, 120)\n",
    "sigmas = np.linspace(np.sqrt(var_ml) - 0.8, np.sqrt(var_ml) + 0.8, 120)\n",
    "\n",
    "MU, SIG = np.meshgrid(mus, sigmas)\n",
    "n = len(loudness)\n",
    "\n",
    "# Sum_n (x_n - mu)^2 = Sum x_n^2 - 2 mu Sum x_n + n mu^2, so the two sufficient\n",
    "# statistics are all that is required; the observations never enter the grid.\n",
    "s1, s2 = loudness.sum(), (loudness ** 2).sum()\n",
    "sq = s2 - 2 * MU * s1 + n * MU ** 2\n",
    "LL = -0.5 * n * np.log(2 * np.pi * SIG ** 2) - sq / (2 * SIG ** 2)\n",
    "\n",
    "fig = go.Figure(go.Contour(x=mus, y=sigmas, z=LL, ncontours=40,\n",
    "                           colorscale='Blues', showscale=False))\n",
    "fig.add_trace(go.Scatter(x=[mu_ml], y=[np.sqrt(var_ml)], mode='markers',\n",
    "                         marker=dict(color=GOLD, size=14, line=dict(color='black', width=1))))\n",
    "fig.update_xaxes(title='mu'); fig.update_yaxes(title='sigma')\n",
    "fig.update_layout(title='Log likelihood surface, with the closed-form solution marked',\n",
    "                  width=760, height=520, template='plotly_white', showlegend=False)\n",
    "fig"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "fb4d7626",
   "metadata": {},
   "source": [
    "## 2. The maximum likelihood variance is biased\n",
    "\n",
    "$\\sigma^2_{ML}$ divides by $N$. Lecture established that\n",
    "$\\mathbb{E}[\\sigma^2_{ML}] = \\frac{N-1}{N}\\sigma^2$, so on average it underestimates. The\n",
    "deviations are measured from $\\mu_{ML}$, which was fitted to the same data and has moved toward\n",
    "those observations.\n",
    "\n",
    "The claim concerns an expectation over datasets, so it is checked by drawing many datasets."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "ac1df222",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-09T06:49:43.126004Z",
     "iopub.status.busy": "2026-09-09T06:49:43.125914Z",
     "iopub.status.idle": "2026-09-09T06:49:43.149154Z",
     "shell.execute_reply": "2026-09-09T06:49:43.148946Z"
    }
   },
   "outputs": [],
   "source": [
    "TRUE_MU, TRUE_VAR = -52.0, 16.0\n",
    "\n",
    "def average_variance_estimates(N, reps=4000):\n",
    "    s = rng.normal(TRUE_MU, np.sqrt(TRUE_VAR), size=(reps, N))\n",
    "    return s.var(axis=1).mean(), s.var(axis=1, ddof=1).mean()\n",
    "\n",
    "Ns = np.array([2, 3, 5, 10, 20, 50, 100])\n",
    "results = np.array([average_variance_estimates(N) for N in Ns])\n",
    "\n",
    "fig = go.Figure()\n",
    "fig.add_trace(go.Scatter(x=Ns, y=results[:, 0], mode='lines+markers',\n",
    "                         name='sigma2_ML  (divide by N)', line=dict(color=ORANGE)))\n",
    "fig.add_trace(go.Scatter(x=Ns, y=results[:, 1], mode='lines+markers',\n",
    "                         name='unbiased  (divide by N-1)', line=dict(color=GREEN)))\n",
    "fig.add_trace(go.Scatter(x=Ns, y=TRUE_VAR * (Ns - 1) / Ns, mode='lines',\n",
    "                         name='theory: sigma2 (N-1)/N', line=dict(color=BLUE, dash='dot')))\n",
    "fig.add_hline(y=TRUE_VAR, line_dash='dash', annotation_text='true sigma2 = 16')\n",
    "fig.update_xaxes(type='log', title='N  (samples per dataset)', tickmode='array',\n",
    "                 tickvals=Ns, ticktext=[str(x) for x in Ns], minor_showgrid=False)\n",
    "fig.update_yaxes(title='average estimate over 4,000 datasets')\n",
    "fig.update_layout(title='The MLE variance is low by a factor of (N-1)/N',\n",
    "                  width=850, height=480, template='plotly_white')\n",
    "fig"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "849bc7a6",
   "metadata": {},
   "source": [
    "The simulated averages follow the theoretical curve. The bias vanishes as $N$ grows and matters\n",
    "only for small datasets. It is also the origin of the `ddof` argument: `np.var` with the default\n",
    "`ddof=0` is the maximum likelihood estimate, `ddof=1` is the unbiased one, and Pandas `.var()`\n",
    "defaults to `ddof=1`."
   ]
  },
  {
   "cell_type": "markdown",
   "id": "86e4db77",
   "metadata": {},
   "source": [
    "---\n",
    "\n",
    "## 3. When one Gaussian is the wrong model\n",
    "\n",
    "Everything above assumed the correct family. Here is a full day of loudness measurements,\n",
    "including the segments where somebody was speaking. This is the dataset on the transition slide."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "a9c2c6d3",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-09T06:49:43.150519Z",
     "iopub.status.busy": "2026-09-09T06:49:43.150438Z",
     "iopub.status.idle": "2026-09-09T06:49:43.152736Z",
     "shell.execute_reply": "2026-09-09T06:49:43.152488Z"
    }
   },
   "outputs": [],
   "source": [
    "speech = data_rng.normal(-30, 4.5, size=1800)      # someone talking\n",
    "day = np.concatenate([loudness, speech])\n",
    "data_rng.shuffle(day)\n",
    "\n",
    "mu_day, var_day = day.mean(), day.var()\n",
    "print(f'N = {len(day)}   mu_ML = {mu_day:.1f}   sigma_ML = {np.sqrt(var_day):.1f}')"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "5649ca0f",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-09T06:49:43.153865Z",
     "iopub.status.busy": "2026-09-09T06:49:43.153772Z",
     "iopub.status.idle": "2026-09-09T06:49:43.180945Z",
     "shell.execute_reply": "2026-09-09T06:49:43.180734Z"
    }
   },
   "outputs": [],
   "source": [
    "grid2 = np.linspace(day.min() - 3, day.max() + 3, 600)\n",
    "\n",
    "fig = px.histogram(day, nbins=90, histnorm='probability density',\n",
    "                   color_discrete_sequence=['#c3ccdd'], labels={'value': 'loudness (dB)'})\n",
    "fig.add_trace(go.Scatter(x=grid2, y=gaussian(grid2, mu_day, var_day),\n",
    "                         line=dict(color=ORANGE, width=4), name='best single Gaussian'))\n",
    "fig.add_vline(x=mu_day, line_dash='dot', line_color=ORANGE)\n",
    "fig.update_layout(title='The best single Gaussian puts its peak where the data is sparsest',\n",
    "                  width=880, height=470, template='plotly_white', showlegend=False)\n",
    "fig"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "6c9f2fef",
   "metadata": {},
   "source": [
    "The fitted density places its mode between the two modes of the data, at a loudness that is rarely\n",
    "observed. This is not a failure of the optimization: it is the best available single Gaussian, and\n",
    "the family was the wrong choice.\n",
    "\n",
    "One remedy abandons the parametric family and lets the data determine the shape."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "d0dc805a",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-09T06:49:43.182253Z",
     "iopub.status.busy": "2026-09-09T06:49:43.182164Z",
     "iopub.status.idle": "2026-09-09T06:49:43.244953Z",
     "shell.execute_reply": "2026-09-09T06:49:43.244681Z"
    }
   },
   "outputs": [],
   "source": [
    "from scipy.stats import gaussian_kde\n",
    "\n",
    "fig = px.histogram(day, nbins=90, histnorm='probability density',\n",
    "                   color_discrete_sequence=['#c3ccdd'], labels={'value': 'loudness (dB)'})\n",
    "for bw, color, dash in [(0.05, GOLD, 'dot'), (0.25, GREEN, 'solid'), (1.0, ORANGE, 'dash')]:\n",
    "    fig.add_trace(go.Scatter(x=grid2, y=gaussian_kde(day, bw_method=bw)(grid2),\n",
    "                             line=dict(color=color, dash=dash), name=f'bandwidth {bw}'))\n",
    "fig.update_layout(title='Kernel density estimation at three bandwidths',\n",
    "                  width=880, height=470, template='plotly_white')\n",
    "fig"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "5ae7f56d",
   "metadata": {},
   "source": [
    "A small bandwidth tracks individual observations; a large one merges the two modes. The bandwidth\n",
    "plays the role that bin width plays in a histogram. The cost is that the estimate retains all\n",
    "5,800 observations, where the Gaussian needed two numbers.\n",
    "\n",
    "The other remedy keeps a parametric model but enlarges the family, which is the Gaussian mixture.\n",
    "\n",
    "---\n",
    "\n",
    "## 4. A mixture of two Gaussians\n",
    "\n",
    "A mixture assigns each component a weight $\\pi_k$, a mean $\\mu_k$, and a variance $\\sigma_k^2$:\n",
    "\n",
    "$$p(x) = \\sum_{k=1}^{K} \\pi_k\\, \\mathcal{N}(x \\mid \\mu_k, \\sigma_k^2)$$\n",
    "\n",
    "Fitting one with scikit-learn takes a single call. The data must be shaped $(N, D)$ even when\n",
    "$D = 1$."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "9724dbc4",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-09T06:49:43.246291Z",
     "iopub.status.busy": "2026-09-09T06:49:43.246194Z",
     "iopub.status.idle": "2026-09-09T06:49:43.705581Z",
     "shell.execute_reply": "2026-09-09T06:49:43.705224Z"
    }
   },
   "outputs": [],
   "source": [
    "from sklearn.mixture import GaussianMixture\n",
    "\n",
    "X = day.reshape(-1, 1)                       # (N, 1): one feature\n",
    "gmm = GaussianMixture(n_components=2, random_state=0).fit(X)\n",
    "\n",
    "for k in np.argsort(gmm.means_.ravel()):\n",
    "    print(f'component {k}:  pi = {gmm.weights_[k]:.3f}   '\n",
    "          f'mu = {gmm.means_[k, 0]:7.2f}   sigma = {np.sqrt(gmm.covariances_[k, 0, 0]):.2f}')"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "a917c3b2",
   "metadata": {},
   "source": [
    "Compare those against the values the data was generated from: weights 4000/5800 and 1800/5800,\n",
    "means -52 and -30, standard deviations 4.0 and 4.5."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "5467ce60",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-09T06:49:43.707346Z",
     "iopub.status.busy": "2026-09-09T06:49:43.707124Z",
     "iopub.status.idle": "2026-09-09T06:49:43.739637Z",
     "shell.execute_reply": "2026-09-09T06:49:43.739321Z"
    }
   },
   "outputs": [],
   "source": [
    "mix = np.exp(gmm.score_samples(grid2.reshape(-1, 1)))\n",
    "\n",
    "fig = px.histogram(day, nbins=90, histnorm='probability density',\n",
    "                   color_discrete_sequence=['#c3ccdd'], labels={'value': 'loudness (dB)'})\n",
    "fig.add_trace(go.Scatter(x=grid2, y=gaussian(grid2, mu_day, var_day),\n",
    "                         line=dict(color=ORANGE, width=3, dash='dash'), name='single Gaussian'))\n",
    "fig.add_trace(go.Scatter(x=grid2, y=mix, line=dict(color=BLUE, width=4), name='mixture of 2'))\n",
    "for k, c in zip(range(2), [GREEN, GOLD]):\n",
    "    comp = gmm.weights_[k] * gaussian(grid2, gmm.means_[k, 0], gmm.covariances_[k, 0, 0])\n",
    "    fig.add_trace(go.Scatter(x=grid2, y=comp, line=dict(color=c, dash='dot'),\n",
    "                             name=f'component {k}'))\n",
    "fig.update_layout(title='One Gaussian against a mixture of two',\n",
    "                  width=880, height=470, template='plotly_white')\n",
    "fig"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "3a6f69c3",
   "metadata": {},
   "source": [
    "### Soft assignments\n",
    "\n",
    "`predict_proba` gives the responsibility of each component for each point. In one dimension we can\n",
    "plot it against $x$ directly, which shows exactly where the model is uncertain."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "f231a3cd",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-09T06:49:43.741284Z",
     "iopub.status.busy": "2026-09-09T06:49:43.741078Z",
     "iopub.status.idle": "2026-09-09T06:49:43.756557Z",
     "shell.execute_reply": "2026-09-09T06:49:43.756293Z"
    }
   },
   "outputs": [],
   "source": [
    "resp = gmm.predict_proba(grid2.reshape(-1, 1))\n",
    "order = np.argsort(gmm.means_.ravel())\n",
    "\n",
    "fig = go.Figure()\n",
    "for j, k in enumerate(order):\n",
    "    fig.add_trace(go.Scatter(x=grid2, y=resp[:, k], line=dict(color=[GREEN, GOLD][j], width=3),\n",
    "                             name=f'p(z = {j} | x)'))\n",
    "fig.add_hline(y=0.5, line_dash='dot', line_color='#888')\n",
    "fig.update_xaxes(title='loudness (dB)')\n",
    "fig.update_yaxes(title='responsibility', range=[-0.03, 1.03])\n",
    "fig.update_layout(title='Responsibilities: the crossover is where a hard assignment would be a guess',\n",
    "                  width=880, height=440, template='plotly_white')\n",
    "fig"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "dc2f7992",
   "metadata": {},
   "source": [
    "The curves are near 0 or 1 almost everywhere and swap over a narrow band around -40 dB. Points in\n",
    "that band are genuinely ambiguous, and a hard assignment would discard exactly that information.\n",
    "This is the difference between k-means and a mixture model.\n",
    "\n",
    "---\n",
    "\n",
    "## 5. A GMM is a generative model\n",
    "\n",
    "The mixture factorizes as $p(x, z) = p(z)\\,p(x \\mid z)$, so it can be sampled by **ancestor\n",
    "sampling**: draw the component first, then draw the point from that component."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "fce8784a",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-09T06:49:43.757966Z",
     "iopub.status.busy": "2026-09-09T06:49:43.757875Z",
     "iopub.status.idle": "2026-09-09T06:49:43.762749Z",
     "shell.execute_reply": "2026-09-09T06:49:43.762428Z"
    }
   },
   "outputs": [],
   "source": [
    "pi_true = np.array([0.2, 0.5, 0.3])\n",
    "mu_true = np.array([-1.0, 2.0, 5.0])\n",
    "var_true = np.array([0.2, 0.5, 0.1])\n",
    "\n",
    "N = 3000\n",
    "z = rng.choice(3, size=N, p=pi_true)                  # first the latent component\n",
    "x = rng.normal(mu_true[z], np.sqrt(var_true[z]))      # then the observation\n",
    "\n",
    "pd.Series(z).value_counts(normalize=True).sort_index().round(3)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "697012d3",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-09T06:49:43.764224Z",
     "iopub.status.busy": "2026-09-09T06:49:43.764058Z",
     "iopub.status.idle": "2026-09-09T06:49:43.792055Z",
     "shell.execute_reply": "2026-09-09T06:49:43.791778Z"
    }
   },
   "outputs": [],
   "source": [
    "grid3 = np.linspace(x.min() - 1, x.max() + 1, 600)\n",
    "density = sum(pi_true[k] * gaussian(grid3, mu_true[k], var_true[k]) for k in range(3))\n",
    "\n",
    "fig = px.histogram(x, nbins=90, histnorm='probability density',\n",
    "                   color_discrete_sequence=['#c9d3e6'], labels={'value': 'x'})\n",
    "fig.add_trace(go.Scatter(x=grid3, y=density, line=dict(color=BLUE, width=3), name='mixture'))\n",
    "for k, c in enumerate([GOLD, GREEN, ORANGE]):\n",
    "    fig.add_trace(go.Scatter(x=grid3, y=pi_true[k] * gaussian(grid3, mu_true[k], var_true[k]),\n",
    "                             line=dict(color=c, dash='dash'), name=f'component {k}'))\n",
    "fig.update_layout(title='Samples drawn by ancestor sampling, with the density that generated them',\n",
    "                  width=880, height=470, template='plotly_white')\n",
    "fig"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "257edfda",
   "metadata": {},
   "source": [
    "The dashed curves are the weighted components and the solid curve is their sum.\n",
    "\n",
    "Note what we needed in order to sample: the component $z$ of every point. When fitting, that is\n",
    "precisely what we do not have."
   ]
  },
  {
   "cell_type": "markdown",
   "id": "8e8308d3",
   "metadata": {},
   "source": [
    "---\n",
    "\n",
    "## 6. Implementing EM\n",
    "\n",
    "We now fit a mixture without knowing $z$. The two steps from lecture translate directly into code.\n",
    "\n",
    "The functions below are written for general $D$, taking data of shape $(N, D)$ and covariance\n",
    "matrices of shape $(K, D, D)$. Every example in this notebook uses $D = 1$, which is what the\n",
    "slides derive, but the same code fits multivariate mixtures unchanged.\n",
    "\n",
    "**E-step.** With the parameters fixed, compute the responsibility of each component for each\n",
    "point, which is the posterior over $z$ obtained by Bayes' theorem:\n",
    "\n",
    "$$\\gamma_{nk} = \\frac{\\pi_k\\,\\mathcal{N}(x_n \\mid \\mu_k, \\sigma_k^2)}\n",
    "{\\sum_{k'} \\pi_{k'}\\,\\mathcal{N}(x_n \\mid \\mu_{k'}, \\sigma_{k'}^2)}$$"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "4feaa38b",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-09T06:49:43.793568Z",
     "iopub.status.busy": "2026-09-09T06:49:43.793469Z",
     "iopub.status.idle": "2026-09-09T06:49:43.795632Z",
     "shell.execute_reply": "2026-09-09T06:49:43.795379Z"
    }
   },
   "outputs": [],
   "source": [
    "# Posterior over the latent component for every point: returns an (N, K) array.\n",
    "# Written for general D; with D = 1 each Sigma[k] is a 1x1 matrix holding sigma_k^2.\n",
    "def E_step(x, mu, Sigma, pi):\n",
    "    N, D = x.shape\n",
    "    K = len(pi)\n",
    "    gamma = np.zeros((N, K))\n",
    "    for k in range(K):\n",
    "        gamma[:, k] = pi[k] * multivariate_normal(mu[k], Sigma[k]).pdf(x)\n",
    "    return gamma / gamma.sum(axis=1, keepdims=True)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "30d54e19",
   "metadata": {},
   "source": [
    "**M-step.** With the responsibilities fixed, maximize the expected complete-data log likelihood.\n",
    "Each parameter has a closed form, derived in lecture. In one dimension the last line reads\n",
    "$\\sigma_k^2 = \\frac{1}{N_k}\\sum_n \\gamma_{nk}(x_n - \\mu_k)^2$, the same weighted average of\n",
    "squared deviations we derived for a single Gaussian.\n",
    "\n",
    "$$N_k = \\sum_n \\gamma_{nk} \\qquad \\pi_k = \\frac{N_k}{N} \\qquad\n",
    "\\mu_k = \\frac{1}{N_k}\\sum_n \\gamma_{nk} x_n$$"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "af80dfea",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-09T06:49:43.796966Z",
     "iopub.status.busy": "2026-09-09T06:49:43.796874Z",
     "iopub.status.idle": "2026-09-09T06:49:43.799238Z",
     "shell.execute_reply": "2026-09-09T06:49:43.798990Z"
    }
   },
   "outputs": [],
   "source": [
    "# Closed-form parameter updates. `reg` keeps a component from collapsing onto one point.\n",
    "def M_step(x, gamma, reg=1e-6):\n",
    "    N, D = x.shape\n",
    "    K = gamma.shape[1]\n",
    "    mu = np.zeros((K, D))\n",
    "    Sigma = np.zeros((K, D, D))\n",
    "    pi = np.zeros(K)\n",
    "    for k in range(K):\n",
    "        N_k = gamma[:, k].sum()\n",
    "        mu[k] = gamma[:, k] @ x / N_k\n",
    "        centred = x - mu[k]\n",
    "        Sigma[k] = (gamma[:, k] * centred.T) @ centred / N_k + reg * np.eye(D)\n",
    "        pi[k] = N_k / N\n",
    "    return mu, Sigma, pi"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "11298efd",
   "metadata": {},
   "source": [
    "The `reg` term addresses the singularity from lecture. Without it a component can shrink onto a\n",
    "single point, its variance collapsing toward zero and the likelihood diverging.\n",
    "\n",
    "Initialization uses k-means, as `sklearn` does."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "3a60584d",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-09T06:49:43.800514Z",
     "iopub.status.busy": "2026-09-09T06:49:43.800421Z",
     "iopub.status.idle": "2026-09-09T06:49:43.803643Z",
     "shell.execute_reply": "2026-09-09T06:49:43.803373Z"
    }
   },
   "outputs": [],
   "source": [
    "from sklearn.cluster import KMeans\n",
    "\n",
    "def initialize(x, K, seed=0):\n",
    "    D = x.shape[1]\n",
    "    centers = KMeans(n_clusters=K, n_init=10, random_state=seed).fit(x).cluster_centers_\n",
    "    Sigma = np.array([np.atleast_2d(np.cov(x.T)) for _ in range(K)])\n",
    "    return centers, Sigma, np.ones(K) / K\n",
    "\n",
    "\n",
    "def log_likelihood(x, mu, Sigma, pi):\n",
    "    per_component = np.column_stack(\n",
    "        [pi[k] * multivariate_normal(mu[k], Sigma[k]).pdf(x) for k in range(len(pi))])\n",
    "    return np.log(per_component.sum(axis=1)).sum()\n",
    "\n",
    "\n",
    "def em(x, K, iters=50, seed=0):\n",
    "    mu, Sigma, pi = initialize(x, K, seed)\n",
    "    history = [log_likelihood(x, mu, Sigma, pi)]\n",
    "    for _ in range(iters):\n",
    "        gamma = E_step(x, mu, Sigma, pi)\n",
    "        mu, Sigma, pi = M_step(x, gamma)\n",
    "        history.append(log_likelihood(x, mu, Sigma, pi))\n",
    "    return mu, Sigma, pi, history"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "fcff3c2a",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-09T06:49:43.804934Z",
     "iopub.status.busy": "2026-09-09T06:49:43.804861Z",
     "iopub.status.idle": "2026-09-09T06:49:43.849038Z",
     "shell.execute_reply": "2026-09-09T06:49:43.848759Z"
    }
   },
   "outputs": [],
   "source": [
    "mu_em, Sigma_em, pi_em, history = em(X, K=2)      # X is day.reshape(-1, 1)\n",
    "\n",
    "for k in np.argsort(mu_em.ravel()):\n",
    "    print(f'component {k}:  pi = {pi_em[k]:.3f}   mu = {mu_em[k,0]:7.2f}   '\n",
    "          f'sigma = {np.sqrt(Sigma_em[k,0,0]):.2f}')\n",
    "print(f'\\nlog likelihood: {history[0]:.1f} -> {history[-1]:.1f}')"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "f2499b5b",
   "metadata": {},
   "source": [
    "Those are the same parameters scikit-learn found, recovered from an implementation that is about\n",
    "twenty lines long.\n",
    "\n",
    "### Each iteration increases the log likelihood\n",
    "\n",
    "This is the property that guarantees convergence, and it is worth seeing rather than taking on\n",
    "faith."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "ca4e07af",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-09T06:49:43.850484Z",
     "iopub.status.busy": "2026-09-09T06:49:43.850407Z",
     "iopub.status.idle": "2026-09-09T06:49:43.867956Z",
     "shell.execute_reply": "2026-09-09T06:49:43.867698Z"
    }
   },
   "outputs": [],
   "source": [
    "fig = px.line(y=history, markers=True, labels={'x': 'iteration', 'y': 'log likelihood'},\n",
    "              title='EM increases the log likelihood at every step',\n",
    "              width=820, height=430, template='plotly_white')\n",
    "fig.update_traces(line_color=BLUE)\n",
    "fig"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "78f826b4",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-09T06:49:43.869277Z",
     "iopub.status.busy": "2026-09-09T06:49:43.869184Z",
     "iopub.status.idle": "2026-09-09T06:49:43.871775Z",
     "shell.execute_reply": "2026-09-09T06:49:43.871433Z"
    }
   },
   "outputs": [],
   "source": [
    "print('non-decreasing? ', bool(np.all(np.diff(history) >= -1e-6)))\n",
    "print('largest single-step decrease:', np.diff(history).min())\n",
    "print('\\nours vs sklearn (means, sorted):',\n",
    "      np.round(np.sort(mu_em.ravel()), 3), np.round(np.sort(gmm.means_.ravel()), 3))"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "48303a5a",
   "metadata": {},
   "source": [
    "### Local optima\n",
    "\n",
    "EM converges to a local maximum, so the starting point matters. Initializing at random rather than\n",
    "from k-means makes that visible."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "d050f4a1",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-09T06:49:43.873142Z",
     "iopub.status.busy": "2026-09-09T06:49:43.873056Z",
     "iopub.status.idle": "2026-09-09T06:49:44.538727Z",
     "shell.execute_reply": "2026-09-09T06:49:44.538445Z"
    }
   },
   "outputs": [],
   "source": [
    "def em_random_init(x, K, iters=50, seed=0):\n",
    "    g = np.random.default_rng(seed)\n",
    "    mu = x[g.choice(len(x), size=K, replace=False)] + g.normal(0, 8, size=(K, x.shape[1]))\n",
    "    Sigma = np.array([np.atleast_2d(np.cov(x.T)) for _ in range(K)])\n",
    "    pi = np.ones(K) / K\n",
    "    for _ in range(iters):\n",
    "        mu, Sigma, pi = M_step(x, E_step(x, mu, Sigma, pi))\n",
    "    return log_likelihood(x, mu, Sigma, pi)\n",
    "\n",
    "random_finals = [em_random_init(X, 3, seed=s) for s in range(8)]\n",
    "kmeans_finals = [em(X, K=3, iters=50, seed=s)[3][-1] for s in range(8)]\n",
    "\n",
    "pd.DataFrame({'random init': np.round(random_finals, 1),\n",
    "              'k-means init': np.round(kmeans_finals, 1)})"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "f74eeb6f",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-09T06:49:44.540136Z",
     "iopub.status.busy": "2026-09-09T06:49:44.540038Z",
     "iopub.status.idle": "2026-09-09T06:49:44.542142Z",
     "shell.execute_reply": "2026-09-09T06:49:44.541766Z"
    }
   },
   "outputs": [],
   "source": [
    "for name, vals in [('random init', random_finals), ('k-means init', kmeans_finals)]:\n",
    "    print(f'{name:13s} best {max(vals):10.1f}   worst {min(vals):10.1f}   '\n",
    "          f'spread {max(vals) - min(vals):8.1f}')"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "e2fdc2a6",
   "metadata": {},
   "source": [
    "Fitting three components to data that really has two leaves genuine ambiguity about how to split\n",
    "it, and the random starts settle into two distinct optima a couple of nats apart. The k-means\n",
    "starts reach the same solution every time.\n",
    "\n",
    "Two honest caveats. The gap here is small, because well-separated components in one dimension\n",
    "make EM fairly robust; local optima bite much harder with more components and in higher\n",
    "dimensions. And k-means initialization is not automatically better, only consistent: one random\n",
    "start here found a marginally higher likelihood. What it buys is reproducibility, which is why it\n",
    "is the default in `sklearn.mixture.GaussianMixture`, with `n_init` available for the cases where\n",
    "consistency is not enough.\n",
    "\n",
    "## 7. k-means is EM in a limit\n",
    "\n",
    "Lecture claimed that constraining every component to the same fixed variance\n",
    "$\\sigma_k^2 = \\varepsilon$ and letting $\\varepsilon \\to 0$ turns the E-step into nearest-center\n",
    "assignment. We can watch the responsibilities harden as $\\varepsilon$ shrinks."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "83161819",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-09T06:49:44.543784Z",
     "iopub.status.busy": "2026-09-09T06:49:44.543688Z",
     "iopub.status.idle": "2026-09-09T06:49:44.566552Z",
     "shell.execute_reply": "2026-09-09T06:49:44.566147Z"
    }
   },
   "outputs": [],
   "source": [
    "# At small epsilon every density underflows to zero, so the responsibilities have to be\n",
    "# computed in log space: subtract the row maximum before exponentiating.\n",
    "def responsibilities_logspace(x, mu, Sigma, pi):\n",
    "    logs = np.column_stack([np.log(pi[k]) + multivariate_normal(mu[k], Sigma[k]).logpdf(x)\n",
    "                            for k in range(len(pi))])\n",
    "    logs -= logs.max(axis=1, keepdims=True)\n",
    "    r = np.exp(logs)\n",
    "    return r / r.sum(axis=1, keepdims=True)\n",
    "\n",
    "centers = KMeans(n_clusters=2, n_init=10, random_state=0).fit(X).cluster_centers_\n",
    "pi_eq = np.ones(2) / 2\n",
    "nearest = np.argmin(((X[:, None, :] - centers[None]) ** 2).sum(-1), axis=1)\n",
    "\n",
    "rows = []\n",
    "for eps in (100.0, 25.0, 5.0, 1.0, 0.1, 0.01):\n",
    "    Sig = np.array([eps * np.eye(1)] * 2)\n",
    "    g = responsibilities_logspace(X, centers, Sig, pi_eq)\n",
    "    rows.append({'epsilon': eps,\n",
    "                 'mean largest responsibility': g.max(axis=1).mean(),\n",
    "                 'fraction above 0.99': (g.max(axis=1) > 0.99).mean(),\n",
    "                 'matches nearest center': (g.argmax(axis=1) == nearest).mean()})\n",
    "pd.DataFrame(rows)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "c88e4e21",
   "metadata": {},
   "source": [
    "As $\\varepsilon$ falls, the largest responsibility approaches 1 for every point: the soft\n",
    "assignment becomes a hard one.\n",
    "\n",
    "The last column is 1.0 throughout, which is worth noticing. With equal weights and equal variances\n",
    "the most probable component is always the nearest center, whatever $\\varepsilon$ is. What\n",
    "$\\varepsilon$ controls is not *which* component wins but how confidently it wins.\n",
    "\n",
    "k-means is therefore not a separate algorithm that resembles EM. It is EM for a mixture of equal,\n",
    "infinitely narrow Gaussians, which is also why it cannot represent components of different widths\n",
    "and a full mixture can.\n",
    "\n",
    "---\n",
    "\n",
    "## Summary\n",
    "\n",
    "| | |\n",
    "|---|---|\n",
    "| Gaussian MLE | the sample mean and sample variance, depending on the data only through $\\sum_n x_n$ and $\\sum_n x_n^2$ |\n",
    "| Bias | $\\sigma^2_{ML}$ is low by $(N-1)/N$, which is where `ddof` comes from |\n",
    "| Model choice | maximum likelihood optimizes within a family and cannot rescue the wrong family |\n",
    "| Mixtures | a weighted sum of Gaussians, with a latent $z$ naming the component |\n",
    "| EM | alternate the posterior over $z$ with closed-form parameter updates; the log likelihood increases every step |\n",
    "| k-means | EM for equal-variance components in the zero-variance limit |\n",
    "\n",
    "Lecture 6 applies the Gaussian again, this time as a noise model, and shows that least squares is\n",
    "maximum likelihood."
   ]
  }
 ],
 "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
}
