{
 "cells": [
  {
   "cell_type": "markdown",
   "id": "37ff2b7b",
   "metadata": {},
   "source": "# Линейная регрессия"
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-001",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:10.870623Z",
     "iopub.status.busy": "2026-09-25T10:17:10.870210Z",
     "iopub.status.idle": "2026-09-25T10:17:12.453280Z",
     "shell.execute_reply": "2026-09-25T10:17:12.452907Z"
    }
   },
   "outputs": [],
   "source": [
    "import warnings\n",
    "\n",
    "import matplotlib as mpl\n",
    "import matplotlib.pyplot as plt\n",
    "import seaborn as sns\n",
    "import numpy as np\n",
    "import scipy as sp\n",
    "import scipy.stats as st\n",
    "import scipy.integrate as integrate\n",
    "\n",
    "from sklearn import linear_model\n",
    "from sklearn.base import BaseEstimator, TransformerMixin\n",
    "from sklearn.pipeline import make_pipeline\n",
    "from scipy.stats import multivariate_normal\n",
    "\n",
    "## Ложные предупреждения matmul от Apple Accelerate (NumPy >= 1.25)\n",
    "warnings.filterwarnings(\"ignore\", message=r\".*encountered in matmul\",\n",
    "                        category=RuntimeWarning)\n",
    "\n",
    "sns.set_style(\"whitegrid\")\n",
    "sns.set_palette(\"colorblind\")\n",
    "palette = sns.color_palette()\n",
    "figsize = (11, 6)\n",
    "legend_fontsize = 16\n",
    "\n",
    "from matplotlib import rc\n",
    "rc('font', **{'family': 'sans-serif'})\n",
    "rc('text', usetex=True)\n",
    "rc('text.latex', preamble=r'\\usepackage[utf8]{inputenc}\\usepackage[russian]{babel}')\n",
    "rc('figure', **{'dpi': 200})\n",
    "\n",
    "SEED = 2026\n",
    "np.random.seed(SEED)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "cell-018",
   "metadata": {},
   "source": "# 1. Линейная и полиномиальная регрессия"
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-019",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:12.457017Z",
     "iopub.status.busy": "2026-09-25T10:17:12.456726Z",
     "iopub.status.idle": "2026-09-25T10:17:12.458505Z",
     "shell.execute_reply": "2026-09-25T10:17:12.458277Z"
    }
   },
   "outputs": [],
   "source": [
    "np.random.seed(6754)\n",
    "\n",
    "## Исходная функция\n",
    "orig = lambda x: np.sin(2 * x)\n",
    "\n",
    "SIGMA_NOISE = .25\n",
    "\n",
    "## Небольшая выборка\n",
    "xd = np.array([-3, -2, -1, -0.5, 0, 0.5, 1, 1.5, 2.5, 3, 4]) / 2\n",
    "num_points = len(xd)\n",
    "data = orig(xd) + np.random.normal(0, SIGMA_NOISE, num_points)\n",
    "\n",
    "## Большая выборка\n",
    "xd_large = np.arange(-1.5, 2, 0.05)\n",
    "num_points_l = len(xd_large)\n",
    "data_large = orig(xd_large) + np.random.normal(0, SIGMA_NOISE, num_points_l)\n",
    "\n",
    "## Для рисования\n",
    "xs = np.arange(xd[0] - 1.5, xd[-1] + 1.5, 0.01)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "cell-020",
   "metadata": {},
   "source": [
    "## Оверфиттинг"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-021",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:12.462541Z",
     "iopub.status.busy": "2026-09-25T10:17:12.462279Z",
     "iopub.status.idle": "2026-09-25T10:17:17.515549Z",
     "shell.execute_reply": "2026-09-25T10:17:17.515737Z"
    }
   },
   "outputs": [],
   "source": [
    "## Выделение полиномиальных признаков\n",
    "xs_d = np.vstack([xs ** i for i in range(1, num_points + 1)]).transpose()\n",
    "xd_d = np.vstack([xd ** i for i in range(1, num_points + 1)]).transpose()\n",
    "\n",
    "## Какие степени многочлена будем обучать и рисовать\n",
    "set_of_powers = [1, 3, 5, 7, 10]\n",
    "\n",
    "fig = plt.figure(figsize=figsize)\n",
    "ax = fig.add_subplot(111)\n",
    "ax.set_xlim((xs[0], xs[-1]))\n",
    "ax.set_ylim((-2, 2))\n",
    "ax.scatter(xd, data, marker='*', s=120)\n",
    "ax.plot(xs, orig(xs), linewidth=1, label=\"Исходная функция\", color=\"black\")\n",
    "\n",
    "for d in set_of_powers:\n",
    "    if d == 0:\n",
    "        print(np.mean(data))\n",
    "        ax.hlines(np.mean(data), xmin=xs[0], xmax=xs[-1], label=\"$d=0$\", linestyle=\"dashed\")\n",
    "    else:\n",
    "        cur_model = linear_model.LinearRegression(fit_intercept=True).fit(xd_d[:, :d], data)\n",
    "        print(\"d = %2d, коэффициенты: %s\" % (d, np.array2string(cur_model.coef_, precision=2)))\n",
    "        ax.plot(xs, cur_model.predict(xs_d[:, :d]), linewidth=2, label=\"$d=%d$\" % d)\n",
    "\n",
    "ax.legend(loc=\"upper right\", fontsize=legend_fontsize)\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-022",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:17.519835Z",
     "iopub.status.busy": "2026-09-25T10:17:17.519598Z",
     "iopub.status.idle": "2026-09-25T10:17:17.521750Z",
     "shell.execute_reply": "2026-09-25T10:17:17.521551Z"
    }
   },
   "outputs": [],
   "source": [
    "## Обусловленность матрицы плана\n",
    "for d in [1, 3, 5, 10]:\n",
    "    print(\"d = %2d:  cond(X) = %.3e\" % (d, np.linalg.cond(xd_d[:, :d])))"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "cell-023",
   "metadata": {},
   "source": [
    "## Локальные признаки в линейной регрессии"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-024",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:17.525205Z",
     "iopub.status.busy": "2026-09-25T10:17:17.524977Z",
     "iopub.status.idle": "2026-09-25T10:17:17.526703Z",
     "shell.execute_reply": "2026-09-25T10:17:17.526476Z"
    }
   },
   "outputs": [],
   "source": [
    "class GaussianFeatures(BaseEstimator, TransformerMixin):\n",
    "    def __init__(self, N, width_factor=.5):\n",
    "        self.N = N\n",
    "        self.width_factor = width_factor\n",
    "\n",
    "    @staticmethod\n",
    "    def _gauss_basis(x, y, width, axis=None):\n",
    "        arg = (x - y) / width\n",
    "        return np.exp(-0.5 * np.sum(arg ** 2, axis))\n",
    "\n",
    "    def fit(self, X, y=None):\n",
    "        self.centers_ = np.linspace(X.min(), X.max(), self.N)\n",
    "        self.width_ = self.width_factor * (self.centers_[1] - self.centers_[0])\n",
    "        return self\n",
    "\n",
    "    def transform(self, X):\n",
    "        return self._gauss_basis(X[:, :, np.newaxis], self.centers_, self.width_, axis=1)\n",
    "\n",
    "\n",
    "def plural_features(n):\n",
    "    if 11 <= n % 100 <= 14:\n",
    "        return \"признаков\"\n",
    "    return {1: \"признак\", 2: \"признака\", 3: \"признака\", 4: \"признака\"}.get(n % 10, \"признаков\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-025",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:17.537134Z",
     "iopub.status.busy": "2026-09-25T10:17:17.536035Z",
     "iopub.status.idle": "2026-09-25T10:17:17.942261Z",
     "shell.execute_reply": "2026-09-25T10:17:17.942049Z"
    }
   },
   "outputs": [],
   "source": [
    "nums_gauss = [2 ]\n",
    "gauss_xd, gauss_yd = xd, data\n",
    "\n",
    "fig = plt.figure(figsize=figsize)\n",
    "ax = fig.add_subplot(111)\n",
    "ax.set_xlim((xs[0], xs[-1]))\n",
    "ax.set_ylim((-1.5, 1.5))\n",
    "ax.scatter(gauss_xd, gauss_yd, marker='*', s=120)\n",
    "ax.plot(xs, orig(xs), linewidth=1, label=\"Исходная функция\", color=\"black\")\n",
    "\n",
    "for num_gauss in nums_gauss:\n",
    "    gauss_model = make_pipeline(GaussianFeatures(num_gauss), linear_model.LinearRegression())\n",
    "    gauss_model.fit(gauss_xd[:, np.newaxis], gauss_yd)\n",
    "    yfit = gauss_model.predict(xs[:, np.newaxis])\n",
    "    coefs = gauss_model.get_params()['linearregression'].coef_\n",
    "    print(\"%d гауссовских %s, коэффициенты: %s\"\n",
    "          % (num_gauss, plural_features(num_gauss), \" \".join(\"%.4f\" % x for x in coefs)))\n",
    "    ax.plot(xs, yfit, linewidth=2,\n",
    "            label=\"%d гауссовских %s\" % (num_gauss, plural_features(num_gauss)))\n",
    "\n",
    "ax.legend(loc=\"upper left\", fontsize=legend_fontsize)\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "cell-026",
   "metadata": {},
   "source": [
    "### Из чего складывается предсказание"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-027",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:17.946381Z",
     "iopub.status.busy": "2026-09-25T10:17:17.946095Z",
     "iopub.status.idle": "2026-09-25T10:17:17.948682Z",
     "shell.execute_reply": "2026-09-25T10:17:17.948937Z"
    }
   },
   "outputs": [],
   "source": [
    "num_gauss = 10\n",
    "gauss_model = make_pipeline(GaussianFeatures(num_gauss), linear_model.LinearRegression())\n",
    "gauss_model.fit(gauss_xd[:, np.newaxis], gauss_yd)\n",
    "yfit = gauss_model.predict(xs[:, np.newaxis])\n",
    "mfeat = gauss_model.get_params()['gaussianfeatures']\n",
    "mregr = gauss_model.get_params()['linearregression']\n",
    "\n",
    "print(\"%d гауссовских %s, коэффициенты: %s\"\n",
    "      % (num_gauss, plural_features(num_gauss), \" \".join(\"%.4f\" % x for x in mregr.coef_)))\n",
    "print(\"Свободный член: %.4f\" % mregr.intercept_)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-028",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:17.970983Z",
     "iopub.status.busy": "2026-09-25T10:17:17.962585Z",
     "iopub.status.idle": "2026-09-25T10:17:20.296261Z",
     "shell.execute_reply": "2026-09-25T10:17:20.296004Z"
    }
   },
   "outputs": [],
   "source": [
    "fig = plt.figure(figsize=figsize)\n",
    "ax = fig.add_subplot(111)\n",
    "ax.set_xlim((xs[0], xs[-1]))\n",
    "ax.scatter(gauss_xd, gauss_yd, marker='*', s=120)\n",
    "\n",
    "for i in range(mfeat.N):\n",
    "    cur_yfit = mregr.coef_[i] * np.array(\n",
    "        [mfeat._gauss_basis(x, mfeat.centers_[i], mfeat.width_) for x in xs])\n",
    "    ax.plot(xs, cur_yfit, color=\"0.4\", linewidth=1,\n",
    "            label=\"Один взвешенный признак\" if i == 0 else None)\n",
    "ax.axhline(mregr.intercept_, color=\"0.6\", linewidth=1, label=\"Свободный член\")\n",
    "ax.plot(xs, yfit, linewidth=2, label=\"Регрессия с гауссовскими признаками\")\n",
    "ax.legend(loc=\"upper center\", fontsize=legend_fontsize - 2)\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-029",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:20.331191Z",
     "iopub.status.busy": "2026-09-25T10:17:20.311358Z",
     "iopub.status.idle": "2026-09-25T10:17:21.240678Z",
     "shell.execute_reply": "2026-09-25T10:17:21.240426Z"
    }
   },
   "outputs": [],
   "source": [
    "fig = plt.figure(figsize=figsize)\n",
    "ax = fig.add_subplot(111)\n",
    "ax.set_xlim((xs[0], xs[-1]))\n",
    "\n",
    "for i in range(mfeat.N):\n",
    "    cur_yfit = [mfeat._gauss_basis(x, mfeat.centers_[i], mfeat.width_) for x in xs]\n",
    "    ax.plot(xs, cur_yfit, color=\"0.6\", linewidth=1)\n",
    "\n",
    "ax.plot(xs, yfit, linewidth=2, label=\"Регрессия\")\n",
    "ax.legend(loc=\"upper center\", fontsize=legend_fontsize)\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "cell-030",
   "metadata": {},
   "source": [
    "## Добавим ещё данных"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-031",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:21.244788Z",
     "iopub.status.busy": "2026-09-25T10:17:21.244533Z",
     "iopub.status.idle": "2026-09-25T10:17:21.445607Z",
     "shell.execute_reply": "2026-09-25T10:17:21.445403Z"
    }
   },
   "outputs": [],
   "source": [
    "xs_big = np.arange(xd_large[0] - .5, xd_large[-1] + .5, 0.01)\n",
    "xs_d_big = np.vstack([xs_big ** i for i in range(1, num_points + 1)]).transpose()\n",
    "xd_d_large = np.vstack([xd_large ** i for i in range(1, num_points + 1)]).transpose()\n",
    "\n",
    "fig = plt.figure(figsize=figsize)\n",
    "ax = fig.add_subplot(111)\n",
    "ax.set_xlim((xs_big[0], xs_big[-1]))\n",
    "ax.set_ylim((-2, 2))\n",
    "ax.scatter(xd_large, data_large, marker='*', s=120)\n",
    "ax.plot(xs_big, orig(xs_big), linewidth=2, label=\"Исходная функция\", color=\"black\")\n",
    "\n",
    "for d in [1, 3, 10]:\n",
    "    cur_model = linear_model.LinearRegression(fit_intercept=True).fit(xd_d_large[:, :d], data_large)\n",
    "    ax.plot(xs_big, cur_model.predict(xs_d_big[:, :d]), linewidth=2, label=\"$d=%d$\" % d)\n",
    "\n",
    "ax.legend(loc=\"upper left\", fontsize=legend_fontsize)\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-032",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:21.456625Z",
     "iopub.status.busy": "2026-09-25T10:17:21.456369Z",
     "iopub.status.idle": "2026-09-25T10:17:21.690695Z",
     "shell.execute_reply": "2026-09-25T10:17:21.690431Z"
    }
   },
   "outputs": [],
   "source": [
    "num_gauss = 20\n",
    "\n",
    "fig = plt.figure(figsize=figsize)\n",
    "ax = fig.add_subplot(111)\n",
    "ax.set_xlim((xs_big[0], xs_big[-1]))\n",
    "ax.scatter(xd_large, data_large, marker='*', s=120)\n",
    "\n",
    "gauss_model = make_pipeline(GaussianFeatures(num_gauss), linear_model.LinearRegression())\n",
    "gauss_model.fit(xd_large[:, np.newaxis], data_large)\n",
    "yfit_big = gauss_model.predict(xs_big[:, np.newaxis])\n",
    "mfeat_b = gauss_model.get_params()['gaussianfeatures']\n",
    "mregr_b = gauss_model.get_params()['linearregression']\n",
    "\n",
    "for i in range(mfeat_b.N):\n",
    "    \n",
    "    cur_yfit = mregr_b.coef_[i] * np.array(\n",
    "        [mfeat_b._gauss_basis(x, mfeat_b.centers_[i], mfeat_b.width_) for x in xs_big])\n",
    "    ax.plot(xs_big, cur_yfit, color=\"0.4\", linewidth=1)\n",
    "ax.axhline(mregr_b.intercept_, color=\"0.6\", linewidth=1)\n",
    "\n",
    "ax.plot(xs_big, yfit_big, color=\"C1\", linewidth=2, label=\"Регрессия\")\n",
    "ax.plot(xs_big, orig(xs_big), linewidth=1, color=\"black\", label=\"Исходная функция\")\n",
    "ax.legend(loc=\"upper left\", fontsize=legend_fontsize)\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "cell-033",
   "metadata": {},
   "source": [
    "## Регуляризация"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-034",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:21.695022Z",
     "iopub.status.busy": "2026-09-25T10:17:21.694730Z",
     "iopub.status.idle": "2026-09-25T10:17:21.696388Z",
     "shell.execute_reply": "2026-09-25T10:17:21.696187Z"
    }
   },
   "outputs": [],
   "source": [
    "def train_model(features, ys, alpha, use_lasso):\n",
    "    if alpha == 0:\n",
    "        return linear_model.LinearRegression(fit_intercept=True).fit(features, ys)\n",
    "    if use_lasso:\n",
    "        return linear_model.Lasso(alpha=alpha, fit_intercept=True, max_iter=100000).fit(features, ys)\n",
    "    return linear_model.Ridge(alpha=alpha, fit_intercept=True).fit(features, ys)\n",
    "\n",
    "\n",
    "def plot_regularization(alpha_values, use_lasso, d=10):\n",
    "    fig = plt.figure(figsize=figsize)\n",
    "    ax = fig.add_subplot(111)\n",
    "    # ax.set_xlim((xs[0], xs[-1]))\n",
    "    ax.set_xlim((-2, 2.5))\n",
    "    ax.set_ylim((-3, 3))\n",
    "    ax.scatter(xd, data, marker='*', s=120)\n",
    "    ax.plot(xs, orig(xs), linewidth=2, label=\"Исходная функция\", color=\"black\")\n",
    "\n",
    "    for alpha in alpha_values:\n",
    "        m = train_model(xd_d[:, :d], data, alpha, use_lasso)\n",
    "        print(\"alpha = %-9g -> %s\" % (alpha, np.array2string(m.coef_, precision=3)))\n",
    "        ax.plot(xs, m.predict(xs_d[:, :d]), linewidth=2, label=r\"$\\alpha=%g$\" % alpha)\n",
    "\n",
    "    ax.set_title(\"Lasso ($L_1$)\" if use_lasso else \"Ridge ($L_2$)\", fontsize=legend_fontsize)\n",
    "    ax.legend(loc=\"upper center\", fontsize=legend_fontsize)\n",
    "    plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-035",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:21.705437Z",
     "iopub.status.busy": "2026-09-25T10:17:21.705140Z",
     "iopub.status.idle": "2026-09-25T10:17:23.019274Z",
     "shell.execute_reply": "2026-09-25T10:17:23.019495Z"
    }
   },
   "outputs": [],
   "source": [
    "plot_regularization([0, 1e-6, 1e-3, 1e0], use_lasso=True)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-036",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:23.028816Z",
     "iopub.status.busy": "2026-09-25T10:17:23.021921Z",
     "iopub.status.idle": "2026-09-25T10:17:23.470853Z",
     "shell.execute_reply": "2026-09-25T10:17:23.470598Z"
    }
   },
   "outputs": [],
   "source": [
    "plot_regularization([0, .01, 1.], use_lasso=True)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "cell-037",
   "metadata": {},
   "source": [
    "## Усреднение предсказаний"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-038",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:23.481548Z",
     "iopub.status.busy": "2026-09-25T10:17:23.480365Z",
     "iopub.status.idle": "2026-09-25T10:17:24.105771Z",
     "shell.execute_reply": "2026-09-25T10:17:24.106040Z"
    }
   },
   "outputs": [],
   "source": [
    "np.random.seed(SEED)\n",
    "\n",
    "N = 100\n",
    "alpha = 1e-5\n",
    "use_lasso = False\n",
    "\n",
    "fig = plt.figure(figsize=figsize)\n",
    "ax = fig.add_subplot(111)\n",
    "ax.set_xlim((xs[0], xs[-1]))\n",
    "ax.set_ylim((-5, 5))\n",
    "\n",
    "res = []\n",
    "for _ in range(N):\n",
    "    cur_data = orig(xd) + np.random.normal(0, SIGMA_NOISE, num_points)\n",
    "    cur_model = train_model(xd_d, cur_data, alpha, use_lasso)\n",
    "    res.append(cur_model.predict(xs_d))\n",
    "    ax.plot(xs, res[-1], linewidth=.1, color=\"0.3\")\n",
    "\n",
    "ax.plot(xs, orig(xs), linewidth=2, label=\"Исходная функция\", color=palette[0])\n",
    "ax.scatter(xd, orig(xd), marker='*', s=150, color=palette[0])\n",
    "ax.plot(xs, np.mean(res, axis=0), linewidth=2, label=\"Усреднённые предсказания\", color=\"red\")\n",
    "ax.legend(loc=\"upper center\", fontsize=legend_fontsize)\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "cell-041",
   "metadata": {},
   "source": "# 2. Байесовский вывод в линейной регрессии"
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-042",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:24.117877Z",
     "iopub.status.busy": "2026-09-25T10:17:24.115629Z",
     "iopub.status.idle": "2026-09-25T10:17:24.691283Z",
     "shell.execute_reply": "2026-09-25T10:17:24.691508Z"
    }
   },
   "outputs": [],
   "source": [
    "np.random.seed(SEED)\n",
    "\n",
    "## Теперь восстанавливаем прямую\n",
    "TRUE_W = np.array([-0.5, 0.5])\n",
    "orig = lambda x: TRUE_W[0] + TRUE_W[1] * x\n",
    "\n",
    "xd = np.array([-3, -2, -1, -0.5, 0, 0.5, 1, 1.5, 2.5, 3, 4]) / 2\n",
    "num_points = len(xd)\n",
    "data = orig(xd) + np.random.normal(0, SIGMA_NOISE, num_points)\n",
    "\n",
    "xs = np.linspace(-3, 3, 250)\n",
    "\n",
    "fig = plt.figure(figsize=figsize)\n",
    "ax = fig.add_subplot(111)\n",
    "ax.set_xlim((xs[0], xs[-1]))\n",
    "ax.set_ylim((-2, 2))\n",
    "ax.plot(xs, orig(xs), linewidth=2, label=\"Правильный ответ\")\n",
    "ax.scatter(xd, data, marker='*', s=120, label=\"Данные\")\n",
    "ax.legend(fontsize=legend_fontsize)\n",
    "plt.show()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-043",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:24.697126Z",
     "iopub.status.busy": "2026-09-25T10:17:24.696838Z",
     "iopub.status.idle": "2026-09-25T10:17:24.699254Z",
     "shell.execute_reply": "2026-09-25T10:17:24.699014Z"
    }
   },
   "outputs": [],
   "source": [
    "## Сетка в пространстве весов (w0, w1)\n",
    "N_GRID = 250\n",
    "W0, W1 = np.meshgrid(np.linspace(-1, 1, N_GRID), np.linspace(-1, 1, N_GRID))\n",
    "pos = np.dstack((W0, W1))\n",
    "\n",
    "\n",
    "def myplot_heatmap(Z, title=None, points=None):\n",
    "    fig = plt.figure(figsize=(6, 6))\n",
    "    ax = fig.add_subplot(111)\n",
    "    ax.pcolormesh(W0, W1, Z, cmap=plt.cm.jet, shading=\"auto\")\n",
    "    ax.scatter([TRUE_W[0]], [TRUE_W[1]], marker='*', s=200, color=\"white\",\n",
    "               edgecolors=\"black\", zorder=5)\n",
    "    ax.set_xlim((-1, 1))\n",
    "    ax.set_ylim((-1, 1))\n",
    "    ax.set_aspect('equal', adjustable='box')\n",
    "    ax.set_xlabel(r\"$w_0$\", fontsize=legend_fontsize)\n",
    "    ax.set_ylabel(r\"$w_1$\", fontsize=legend_fontsize)\n",
    "    if title is not None:\n",
    "        ax.set_title(title, fontsize=legend_fontsize)\n",
    "    ax.grid(False)\n",
    "    plt.show()\n",
    "\n",
    "\n",
    "def myplot_sample_lines(mu, sigma, n=20, points=None, title=None):\n",
    "    my_w = np.random.multivariate_normal(mu, sigma, n)\n",
    "    fig = plt.figure(figsize=figsize)\n",
    "    ax = fig.add_subplot(111)\n",
    "    for w in my_w:\n",
    "        ax.plot(xs, w[0] + w[1] * xs, 'k-', lw=.4)\n",
    "    ax.plot(xs, orig(xs), linewidth=2, color=palette[0], label=\"Правильный ответ\")\n",
    "    ax.set_ylim((-3, 3))\n",
    "    ax.set_xlim((-3, 3))\n",
    "    if points is not None:\n",
    "        ax.scatter(points[0], points[1], marker='*', s=200, zorder=5)\n",
    "    if title is not None:\n",
    "        ax.set_title(title, fontsize=legend_fontsize)\n",
    "    ax.legend(loc=\"upper left\", fontsize=legend_fontsize)\n",
    "    plt.show()"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "cell-044",
   "metadata": {},
   "source": [
    "## Априорное распределение"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-045",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:24.701961Z",
     "iopub.status.busy": "2026-09-25T10:17:24.701687Z",
     "iopub.status.idle": "2026-09-25T10:17:27.316631Z",
     "shell.execute_reply": "2026-09-25T10:17:27.316390Z"
    }
   },
   "outputs": [],
   "source": [
    "np.random.seed(SEED)\n",
    "\n",
    "PRIOR_MU = np.array([0., 0.])\n",
    "PRIOR_SIGMA = 2 * np.eye(2)\n",
    "\n",
    "cur_mu, cur_sigma = PRIOR_MU.copy(), PRIOR_SIGMA.copy()\n",
    "\n",
    "Z = multivariate_normal.pdf(pos, mean=cur_mu, cov=cur_sigma)\n",
    "myplot_heatmap(Z, title=r\"Априорное распределение $p(w)$\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-046",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:27.326588Z",
     "iopub.status.busy": "2026-09-25T10:17:27.326303Z",
     "iopub.status.idle": "2026-09-25T10:17:27.821205Z",
     "shell.execute_reply": "2026-09-25T10:17:27.821480Z"
    }
   },
   "outputs": [],
   "source": [
    "myplot_sample_lines(cur_mu, cur_sigma, 200, title=r\"Прямые из априорного распределения\")"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "cell-047",
   "metadata": {},
   "source": [
    "## Правдоподобие одной точки"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-048",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:27.824929Z",
     "iopub.status.busy": "2026-09-25T10:17:27.824646Z",
     "iopub.status.idle": "2026-09-25T10:17:28.237676Z",
     "shell.execute_reply": "2026-09-25T10:17:28.237436Z"
    }
   },
   "outputs": [],
   "source": [
    "def get_likelihood(px, py, sigma=SIGMA_NOISE):\n",
    "    return lambda w: (np.exp(-((w[..., 0] + w[..., 1] * px - py) ** 2) / (2 * sigma ** 2))\n",
    "                      / (sigma * np.sqrt(2. * np.pi)))\n",
    "\n",
    "\n",
    "def likelihood_grid(px, py):\n",
    "    return get_likelihood(px, py)(pos)\n",
    "\n",
    "\n",
    "px, py = xd[5], data[5]\n",
    "print(\"Первое наблюдение: x = %.2f, y = %.4f\" % (px, py))\n",
    "myplot_heatmap(likelihood_grid(px, py), title=r\"Правдоподобие одной точки\")"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "cell-049",
   "metadata": {},
   "source": [
    "## Байесовское обновление"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-050",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:28.241545Z",
     "iopub.status.busy": "2026-09-25T10:17:28.241214Z",
     "iopub.status.idle": "2026-09-25T10:17:28.663465Z",
     "shell.execute_reply": "2026-09-25T10:17:28.663205Z"
    }
   },
   "outputs": [],
   "source": [
    "def bayesian_update(mu, sigma, x, y, sigma_noise=SIGMA_NOISE):\n",
    "    x_matrix = np.array([[1., x]])\n",
    "    precision = np.linalg.inv(sigma) + (1 / sigma_noise ** 2) * (x_matrix.T @ x_matrix)\n",
    "    sigma_n = np.linalg.inv(precision)\n",
    "    mu_n = sigma_n @ (np.linalg.inv(sigma) @ mu + (1 / sigma_noise ** 2) * x_matrix.T @ np.array([y]))\n",
    "    return mu_n, sigma_n\n",
    "\n",
    "\n",
    "cur_mu, cur_sigma = bayesian_update(cur_mu, cur_sigma, px, py)\n",
    "print(\"mu    =\", np.array2string(cur_mu, precision=4))\n",
    "print(\"sigma =\", np.array2string(cur_sigma, precision=4).replace(\"\\n\", \"\\n        \"))\n",
    "\n",
    "Z = multivariate_normal.pdf(pos, mean=cur_mu, cov=cur_sigma)\n",
    "myplot_heatmap(Z, title=r\"Апостериорное распределение после 1 точки\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-051",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:28.666263Z",
     "iopub.status.busy": "2026-09-25T10:17:28.665919Z",
     "iopub.status.idle": "2026-09-25T10:17:29.107519Z",
     "shell.execute_reply": "2026-09-25T10:17:29.107716Z"
    }
   },
   "outputs": [],
   "source": [
    "np.random.seed(SEED)\n",
    "myplot_sample_lines(cur_mu, cur_sigma, 240, points=[[px], [py]],\n",
    "                    title=r\"Прямые из апостериорного распределения, 1 точка\")"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "cell-052",
   "metadata": {},
   "source": [
    "## Предсказательное распределение"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-053",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:29.112665Z",
     "iopub.status.busy": "2026-09-25T10:17:29.112419Z",
     "iopub.status.idle": "2026-09-25T10:17:30.398922Z",
     "shell.execute_reply": "2026-09-25T10:17:30.398648Z"
    }
   },
   "outputs": [],
   "source": [
    "def sample_statistics(mu, sigma, xs, n=2000):\n",
    "    my_w = np.random.multivariate_normal(mu, sigma, n)\n",
    "    return my_w[:, 0][:, None] + my_w[:, 1][:, None] * xs[None, :]\n",
    "\n",
    "\n",
    "def plot_predictions(xs, mu, preds, points, title=None):\n",
    "    mean_pred = mu[0] + mu[1] * xs\n",
    "    std_pred = np.std(preds, axis=0)\n",
    "\n",
    "    fig = plt.figure(figsize=figsize)\n",
    "    ax = fig.add_subplot(111)\n",
    "    ax.set_xlim((xs[0], xs[-1]))\n",
    "    ax.set_ylim((-2, 2))\n",
    "    ax.plot(xs, orig(xs), label=\"Правильный ответ\")\n",
    "    ax.plot(xs, mean_pred, color=\"red\", label=\"MAP гипотеза\")\n",
    "    ax.fill_between(xs, mean_pred - SIGMA_NOISE, mean_pred + SIGMA_NOISE,\n",
    "                    color=palette[1], alpha=.3, label=r\"$\\pm$ дисперсия шума\")\n",
    "    ax.fill_between(xs, mean_pred - std_pred - SIGMA_NOISE, mean_pred + std_pred + SIGMA_NOISE,\n",
    "                    color=palette[5], alpha=.2, label=r\"$\\pm$ дисперсия предсказаний\")\n",
    "    ax.scatter(points[0], points[1], marker='*', s=200, zorder=5)\n",
    "    if title is not None:\n",
    "        ax.set_title(title, fontsize=legend_fontsize)\n",
    "    ax.legend(fontsize=legend_fontsize - 2)\n",
    "    plt.show()\n",
    "\n",
    "\n",
    "np.random.seed(SEED)\n",
    "preds = sample_statistics(cur_mu, cur_sigma, xs, n=2000)\n",
    "plot_predictions(xs, cur_mu, preds, [[px], [py]], title=r\"После 1 точки\")"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "cell-054",
   "metadata": {},
   "source": [
    "## Вторая точка"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-055",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:30.401618Z",
     "iopub.status.busy": "2026-09-25T10:17:30.401260Z",
     "iopub.status.idle": "2026-09-25T10:17:30.827402Z",
     "shell.execute_reply": "2026-09-25T10:17:30.827624Z"
    }
   },
   "outputs": [],
   "source": [
    "px2, py2 = xd[7], data[7]\n",
    "print(\"Второе наблюдение: x = %.2f, y = %.4f\" % (px2, py2))\n",
    "myplot_heatmap(likelihood_grid(px2, py2), title=r\"Правдоподобие второй точки\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-056",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:30.830206Z",
     "iopub.status.busy": "2026-09-25T10:17:30.829977Z",
     "iopub.status.idle": "2026-09-25T10:17:31.235790Z",
     "shell.execute_reply": "2026-09-25T10:17:31.236134Z"
    }
   },
   "outputs": [],
   "source": [
    "cur_mu, cur_sigma = bayesian_update(cur_mu, cur_sigma, px2, py2)\n",
    "Z = multivariate_normal.pdf(pos, mean=cur_mu, cov=cur_sigma)\n",
    "myplot_heatmap(Z, title=r\"Апостериорное распределение после 2 точек\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-057",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:31.247281Z",
     "iopub.status.busy": "2026-09-25T10:17:31.241865Z",
     "iopub.status.idle": "2026-09-25T10:17:32.198105Z",
     "shell.execute_reply": "2026-09-25T10:17:32.197827Z"
    }
   },
   "outputs": [],
   "source": [
    "np.random.seed(SEED)\n",
    "myplot_sample_lines(cur_mu, cur_sigma, n=240, points=[[px, px2], [py, py2]],\n",
    "                    title=r\"Прямые из апостериорного распределения, 2 точки\")\n",
    "\n",
    "preds = sample_statistics(cur_mu, cur_sigma, xs, n=2000)\n",
    "plot_predictions(xs, cur_mu, preds, [[px, px2], [py, py2]], title=r\"После 2 точек\")"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "cell-058",
   "metadata": {},
   "source": [
    "## Третья точка"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-059",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:32.201006Z",
     "iopub.status.busy": "2026-09-25T10:17:32.200772Z",
     "iopub.status.idle": "2026-09-25T10:17:33.028387Z",
     "shell.execute_reply": "2026-09-25T10:17:33.028577Z"
    }
   },
   "outputs": [],
   "source": [
    "px3, py3 = xd[1], data[1]\n",
    "print(\"Третье наблюдение: x = %.2f, y = %.4f\" % (px3, py3))\n",
    "myplot_heatmap(likelihood_grid(px3, py3), title=r\"Правдоподобие третьей точки\")\n",
    "\n",
    "cur_mu, cur_sigma = bayesian_update(cur_mu, cur_sigma, px3, py3)\n",
    "Z = multivariate_normal.pdf(pos, mean=cur_mu, cov=cur_sigma)\n",
    "myplot_heatmap(Z, title=r\"Апостериорное распределение после 3 точек\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-060",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:33.039340Z",
     "iopub.status.busy": "2026-09-25T10:17:33.037111Z",
     "iopub.status.idle": "2026-09-25T10:17:33.869480Z",
     "shell.execute_reply": "2026-09-25T10:17:33.869217Z"
    }
   },
   "outputs": [],
   "source": [
    "np.random.seed(SEED)\n",
    "myplot_sample_lines(cur_mu, cur_sigma, n=200, points=[[px, px2, px3], [py, py2, py3]],\n",
    "                    title=r\"Прямые из апостериорного распределения, 3 точки\")\n",
    "\n",
    "preds = sample_statistics(cur_mu, cur_sigma, xs, n=2000)\n",
    "plot_predictions(xs, cur_mu, preds, [[px, px2, px3], [py, py2, py3]],\n",
    "                 title=r\"После 3 точек\")"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "cell-061",
   "metadata": {},
   "source": [
    "## Все точки сразу"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-062",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:33.874543Z",
     "iopub.status.busy": "2026-09-25T10:17:33.874305Z",
     "iopub.status.idle": "2026-09-25T10:17:34.291958Z",
     "shell.execute_reply": "2026-09-25T10:17:34.291752Z"
    }
   },
   "outputs": [],
   "source": [
    "cur_mu, cur_sigma = PRIOR_MU.copy(), PRIOR_SIGMA.copy()\n",
    "for x_i, y_i in zip(xd, data):\n",
    "    cur_mu, cur_sigma = bayesian_update(cur_mu, cur_sigma, x_i, y_i)\n",
    "\n",
    "print(\"истинные веса         :\", np.array2string(TRUE_W, precision=4))\n",
    "print(\"апостериорное среднее :\", np.array2string(cur_mu, precision=4))\n",
    "print(\"апостериорные СКО     :\", np.array2string(np.sqrt(np.diag(cur_sigma)), precision=4))\n",
    "\n",
    "ols = linear_model.LinearRegression().fit(xd[:, None], data)\n",
    "print(\"\\nМНК (для сравнения)   : [%.4f %.4f]\" % (ols.intercept_, ols.coef_[0]))\n",
    "\n",
    "Z = multivariate_normal.pdf(pos, mean=cur_mu, cov=cur_sigma)\n",
    "myplot_heatmap(Z, title=r\"Апостериорное распределение по всем %d точкам\" % num_points)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "cell-063",
   "metadata": {
    "execution": {
     "iopub.execute_input": "2026-09-25T10:17:34.301149Z",
     "iopub.status.busy": "2026-09-25T10:17:34.294428Z",
     "iopub.status.idle": "2026-09-25T10:17:35.122217Z",
     "shell.execute_reply": "2026-09-25T10:17:35.121907Z"
    }
   },
   "outputs": [],
   "source": [
    "np.random.seed(SEED)\n",
    "myplot_sample_lines(cur_mu, cur_sigma, n=200, points=[xd, data],\n",
    "                    title=r\"Прямые из апостериорного распределения, все точки\")\n",
    "\n",
    "preds = sample_statistics(cur_mu, cur_sigma, xs, n=2000)\n",
    "plot_predictions(xs, cur_mu, preds, [xd, data], title=r\"После всех %d точек\" % num_points)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "782b5938-316a-46c9-8a3a-5dd16a51b021",
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "Python 3 (ipykernel)",
   "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.10.12"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}