{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-09-11T15:37:23.176669Z",
     "iopub.status.busy": "2025-09-11T15:37:23.176540Z",
     "iopub.status.idle": "2025-09-11T15:37:23.586259Z",
     "shell.execute_reply": "2025-09-11T15:37:23.586020Z"
    }
   },
   "outputs": [],
   "source": [
    "# Import necessary libraries\n",
    "%matplotlib inline\n",
    "import numpy as np\n",
    "import matplotlib.pyplot as plt\n",
    "\n",
    "%load_ext autoreload\n",
    "%autoreload 2\n",
    "\n",
    "# Load test module for sanity check\n",
    "from test_utils import test"
   ]
  },
  {
   "cell_type": "markdown",
   "execution_count": null,
   "metadata": {
    "id": "TYyZPqnPmhYC"
   },
   "source": [
    "Data Generation\n",
    "==="
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-09-11T15:37:23.587686Z",
     "iopub.status.busy": "2025-09-11T15:37:23.587576Z",
     "iopub.status.idle": "2025-09-11T15:37:23.601092Z",
     "shell.execute_reply": "2025-09-11T15:37:23.600852Z"
    }
   },
   "outputs": [],
   "source": [
    "from numpy.random import rand, randn"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-09-11T15:37:23.602156Z",
     "iopub.status.busy": "2025-09-11T15:37:23.602083Z",
     "iopub.status.idle": "2025-09-11T15:37:23.612574Z",
     "shell.execute_reply": "2025-09-11T15:37:23.612368Z"
    }
   },
   "outputs": [],
   "source": [
    "n, d, k = 100, 2, 2"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-09-11T15:37:23.613672Z",
     "iopub.status.busy": "2025-09-11T15:37:23.613606Z",
     "iopub.status.idle": "2025-09-11T15:37:23.625922Z",
     "shell.execute_reply": "2025-09-11T15:37:23.625727Z"
    }
   },
   "outputs": [],
   "source": [
    "np.random.seed(20)\n",
    "X = rand(n, d)\n",
    "\n",
    "# means = [rand(d)  for _ in range(k)]  # works for any k\n",
    "means = [rand(d) * 0.5 + 0.5, -rand(d) * 0.5 + 0.5]  # for better plotting when k = 2\n",
    "\n",
    "S = np.diag(rand(d))\n",
    "\n",
    "sigmas = [S] * k  # we'll use the same Sigma for all clusters for better visual results\n",
    "\n",
    "print(means)\n",
    "print(sigmas)"
   ]
  },
  {
   "cell_type": "markdown",
   "execution_count": null,
   "metadata": {},
   "source": [
    "## Computing the probability density"
   ]
  },
  {
   "cell_type": "markdown",
   "execution_count": null,
   "metadata": {},
   "source": [
    "### Recall the math\n",
    "\n",
    "For each data point $x_i$, the Gaussian exponent needs\n",
    "\n",
    "$$\n",
    "Q_i \\;=\\; (x_i - \\mu)^\\top \\Sigma^{-1} (x_i - \\mu).\n",
    "$$\n",
    "\n",
    "Let\n",
    "\n",
    "- $Y = X - \\mu$  (shape $n \\times d$)  \n",
    "- $A = \\Sigma^{-1}$  (shape $d \\times d$)  \n",
    "\n",
    "so that\n",
    "\n",
    "$$\n",
    "Q_i = Y_i^\\top A Y_i.\n",
    "$$\n",
    "\n",
    "where $Y_i$ is the feature vector of data point $i$.\n",
    "\n",
    "How can we get $Q_i$?\n",
    "\n",
    "Option 1: $YA$ is of shape $(n,d)$ and $Y$ is of shape $(n,d)$. Elementwise multiplication gives $(n,d)$, and row sum is the scalar quadratic form.\n",
    "\n",
    "$$\n",
    "\\sum_{k=1}^d Y_{ik}(Y_iA)_k = Y_i^\\top A Y_i\n",
    "$$\n",
    "\n",
    "Option 2: Take the diagonal entries of the gram matrix:\n",
    "$$\n",
    "[((YA)Y^\\top)]_{ii} = (Y_iA)\\cdot Y = Y_i^\\top AY_i\n",
    "$$\n",
    "\n",
    "Option 3:\n",
    "Iterate each data point and calculate point-wise $Q_i$"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-09-11T15:37:23.641792Z",
     "iopub.status.busy": "2025-09-11T15:37:23.641687Z",
     "iopub.status.idle": "2025-09-11T15:37:23.659293Z",
     "shell.execute_reply": "2025-09-11T15:37:23.659066Z"
    }
   },
   "outputs": [],
   "source": [
    "def compute_p(X, mean, sigma):\n",
    "    \"\"\"\n",
    "    Compute the probability of each data point in X under a Gaussian distribution\n",
    "\n",
    "    Args:\n",
    "        X: (n, d) numpy array, where each row corresponds to a data point\n",
    "        mean: (d, ) numpy array, the mean of the Gaussian distribution\n",
    "        sigma: (d, d) numpy array, the covariance matrix of the Gaussian distribution\n",
    "\n",
    "    Returns:\n",
    "        p: (n, ) numpy array, the probability of each data point\n",
    "\n",
    "    >>> compute_p(np.array([[0, 0], [1, 1]]), np.array([0, 0]), np.eye(2))\n",
    "    array([0.15915494, 0.05854983])\n",
    "    \"\"\"\n",
    "\n",
    "    d = X.shape[1]\n",
    "    dxm = X - mean\n",
    "    const = 1 / np.sqrt((2 * np.pi) ** d * np.linalg.det(sigma))\n",
    "\n",
    "    ###############################\n",
    "    # Option 1: elementwise (Hadamard) multiplication via * operation\n",
    "    ###############################\n",
    "    \"\"\"\n",
    "    np.dot(dxm, np.linalg.inv(sigma) gives a matrix of shape (n,d)\n",
    "    dxm is of shape (n,d)\n",
    "    * is elementwise (Hadamard) multiplication, which gives elementwise multiplication\n",
    "    so dxm * np.dot(dxm, np.linalg.inv(sigma)) is a (n,d) matrix,\n",
    "    with each \n",
    "    \"\"\"\n",
    "    # exponent = -0.5 * np.sum(dxm * np.dot(dxm, np.linalg.inv(sigma)), axis=1)\n",
    "\n",
    "    ###############################\n",
    "    # Option 2: matrix multiplication\n",
    "    ###############################\n",
    "    \"\"\"\n",
    "    Note after matrix multiplication, we have a gram matrix of shape (n,n)\n",
    "    we only need the diagonal entries \n",
    "    \"\"\"\n",
    "    # 1) using np.dot or dot()\n",
    "    # exponent = -0.5 * (np.dot(np.dot(dxm, np.linalg.inv(sigma)), dxm.transpose())).diagonal()\n",
    "    # equivalently,\n",
    "    # exponent = -0.5 * (dxm.dot(np.linalg.inv(sigma)).dot(dxm.transpose())).diagonal()\n",
    "    # 2) using @ operator\n",
    "    exponent = -0.5 * ((dxm @ np.linalg.inv(sigma)) @ dxm.transpose()).diagonal()\n",
    "    return const * np.exp(exponent)\n",
    "\n",
    "    ###############################\n",
    "    # Option 3: iteration through all data points\n",
    "    ###############################\n",
    "    # [n, d] = np.shape(X)\n",
    "    # invSigma = np.linalg.inv(sigma)\n",
    "\n",
    "    # result = np.zeros((n,))\n",
    "    # for i in range(n):\n",
    "    #     xmu = X[i] - mean # shape (d,)\n",
    "    #     result[i] = const * np.exp(-0.5 * (xmu).T.dot(invSigma).dot(xmu))\n",
    "    # return result\n",
    "\n",
    "    ### TEMPLATE\n",
    "    # # ***************************************************\n",
    "    # # INSERT YOUR CODE HERE\n",
    "    # # ***************************************************\n",
    "    # raise NotImplementedError\n",
    "    ### END SOLUTION\n",
    "\n",
    "\n",
    "test(compute_p)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-09-11T15:37:23.660379Z",
     "iopub.status.busy": "2025-09-11T15:37:23.660306Z",
     "iopub.status.idle": "2025-09-11T15:37:23.672094Z",
     "shell.execute_reply": "2025-09-11T15:37:23.671873Z"
    }
   },
   "outputs": [],
   "source": [
    "ps = [\n",
    "    compute_p(X, m, s) for m, s in zip(means, sigmas)\n",
    "]  # exercise: try to do this without looping"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-09-11T15:37:23.673204Z",
     "iopub.status.busy": "2025-09-11T15:37:23.673135Z",
     "iopub.status.idle": "2025-09-11T15:37:23.684741Z",
     "shell.execute_reply": "2025-09-11T15:37:23.684551Z"
    }
   },
   "outputs": [],
   "source": [
    "assignments = np.argmax(ps, axis=0)\n",
    "print(assignments)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-09-11T15:37:23.685741Z",
     "iopub.status.busy": "2025-09-11T15:37:23.685662Z",
     "iopub.status.idle": "2025-09-11T15:37:23.752693Z",
     "shell.execute_reply": "2025-09-11T15:37:23.752398Z"
    }
   },
   "outputs": [],
   "source": [
    "colors = np.array([\"red\", \"green\"])[assignments]\n",
    "plt.scatter(X[:, 0], X[:, 1], c=colors, s=100)\n",
    "plt.scatter(np.array(means)[:, 0], np.array(means)[:, 1], marker=\"*\", s=200)\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "markdown",
   "execution_count": null,
   "metadata": {
    "id": "VsIOpA8QmhYI"
   },
   "source": [
    "Solution\n",
    "==="
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-09-11T15:37:23.754022Z",
     "iopub.status.busy": "2025-09-11T15:37:23.753905Z",
     "iopub.status.idle": "2025-09-11T15:37:23.767073Z",
     "shell.execute_reply": "2025-09-11T15:37:23.766797Z"
    }
   },
   "outputs": [],
   "source": [
    "def compute_log_p(X, mean, sigma):\n",
    "    \"\"\"\n",
    "    Compute the log probability of each data point in X under a Gaussian distribution\n",
    "\n",
    "    Args:\n",
    "        X: (n, d) numpy array, where each row corresponds to a data point\n",
    "        mean: (d, ) numpy array, the mean of the Gaussian distribution\n",
    "        sigma: (d, d) numpy array, the covariance matrix of the Gaussian distribution\n",
    "\n",
    "    Returns:\n",
    "        log_p: (n, ) numpy array, the log probability of each data point\n",
    "\n",
    "    >>> compute_log_p(np.array([[0, 0], [1, 1]]), np.array([0, 0]), np.eye(2))\n",
    "    array([-1.83787707, -2.83787707])\n",
    "    \"\"\"\n",
    "    # ***************************************************\n",
    "    # INSERT YOUR CODE HERE\n",
    "    # ***************************************************\n",
    "    raise NotImplementedError\n",
    "\n",
    "\n",
    "test(compute_log_p)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-09-11T15:37:23.768101Z",
     "iopub.status.busy": "2025-09-11T15:37:23.768029Z",
     "iopub.status.idle": "2025-09-11T15:37:23.780535Z",
     "shell.execute_reply": "2025-09-11T15:37:23.780302Z"
    }
   },
   "outputs": [],
   "source": [
    "log_ps = [\n",
    "    compute_log_p(X, m, s) for m, s in zip(means, sigmas)\n",
    "]  # exercise: try to do this without looping"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-09-11T15:37:23.781702Z",
     "iopub.status.busy": "2025-09-11T15:37:23.781593Z",
     "iopub.status.idle": "2025-09-11T15:37:23.793495Z",
     "shell.execute_reply": "2025-09-11T15:37:23.793278Z"
    }
   },
   "outputs": [],
   "source": [
    "assignments = np.argmax(log_ps, axis=0)\n",
    "print(assignments)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {
    "execution": {
     "iopub.execute_input": "2025-09-11T15:37:23.794521Z",
     "iopub.status.busy": "2025-09-11T15:37:23.794455Z",
     "iopub.status.idle": "2025-09-11T15:37:23.850300Z",
     "shell.execute_reply": "2025-09-11T15:37:23.850018Z"
    }
   },
   "outputs": [],
   "source": [
    "colors = np.array([\"red\", \"green\"])[assignments]\n",
    "plt.scatter(X[:, 0], X[:, 1], c=colors, s=100)\n",
    "plt.scatter(np.array(means)[:, 0], np.array(means)[:, 1], marker=\"*\", s=200)\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "base",
   "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.12.2"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 1
}
