Files
RL-Study/Notebooks/C6.ipynb
T

258 lines
96 KiB
Plaintext
Raw Normal View History

2026-02-28 00:18:43 +08:00
{
"cells": [
{
"cell_type": "markdown",
"id": "f2aa1842",
"metadata": {},
"source": [
"# 第 6 章:随机近似 (Stochastic Approximation)\n",
"\n",
"在学习 TD 算法之前,我们必须掌握如何“增量式”地更新数据。\n",
"\n",
"## 1. 增量式均值估计 (Incremental Mean Estimation)\n",
"求均值不必等所有数据到齐。如果我们依次收到数据 $x_k$,我们可以通过以下公式不断更新当前的均值估计 $w_{k+1}$\n",
"$$w_{k+1} = w_k - \\alpha_k (w_k - x_k)$$\n",
"\n",
"## 2. Robbins-Monro (RM) 算法\n",
"当我们想解方程 $g(w) = 0$,但我们不知道 $g(w)$ 的具体表达式,且每次观测到的结果都有随机噪声时(即观测值为 $\\tilde{g}(w_k, \\eta_k) = g(w_k) + \\eta_k$),我们可以使用 RM 算法:\n",
"$$w_{k+1} = w_k - a_k \\tilde{g}(w_k, \\eta_k)$$\n",
"\n",
"**收敛的魔法条件**\n",
"步长 $a_k$ 必须满足 $\\sum a_k = \\infty$(不能衰减太快,保证能走到终点)且 $\\sum a_k^2 < \\infty$(最终要衰减到 0,保证能收敛),最经典的步长就是 $a_k = 1/k$。\n",
"\n",
"## 3. 随机梯度下降 (SGD)\n",
"大名鼎鼎的 SGD 其实就是 RM 算法的特例\n",
"SGD 用于最小化目标函数 $J(w) = \\mathbb{E}[f(w, X)]$。它无需计算真实的期望梯度,而是直接使用单个样本的随机梯度来更新:\n",
"$$w_{k+1} = w_k - \\alpha_k \\nabla_w f(w_k, x_k)$$"
]
},
{
"cell_type": "markdown",
"id": "34f1c8c7",
"metadata": {},
"source": [
"增量式均值估计 vs 传统均值估计"
]
},
{
"cell_type": "code",
"execution_count": 1,
"id": "50d3faaa",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"--- 传统方法 (非增量式) ---\n",
"一次性求平均结果: 10.0387\n",
"\n",
"--- 增量式均值估计 (Incremental) ---\n",
"收到 10 个样本后的均值估计: 10.8961\n",
"收到 100 个样本后的均值估计: 9.7923\n",
"收到 1000 个样本后的均值估计: 10.0387\n",
"\n",
"结论:增量式更新完美逼近了真实均值,且不需要保存历史样本!\n"
]
}
],
"source": [
"import numpy as np\n",
"\n",
"# 生成 1000 个随机样本,假设它们是某个状态的奖励\n",
"np.random.seed(42)\n",
"samples = np.random.normal(loc=10.0, scale=2.0, size=1000)\n",
"\n",
"print(\"--- 传统方法 (非增量式) ---\")\n",
"true_mean = np.mean(samples)\n",
"print(f\"一次性求平均结果: {true_mean:.4f}\")\n",
"\n",
"print(\"\\n--- 增量式均值估计 (Incremental) ---\")\n",
"w_k = 0.0 # 初始猜测\n",
"for k, x_k in enumerate(samples, start=1):\n",
" alpha_k = 1 / k # 步长 (对应公式 1/k)\n",
" # 核心公式: 新估计 = 老估计 - 步长 * (老估计 - 新样本)\n",
" w_k = w_k - alpha_k * (w_k - x_k)\n",
" \n",
" if k in [10, 100, 1000]:\n",
" print(f\"收到 {k} 个样本后的均值估计: {w_k:.4f}\")\n",
"\n",
"print(\"\\n结论:增量式更新完美逼近了真实均值,且不需要保存历史样本!\")"
]
},
{
"cell_type": "markdown",
"id": "b0ac2d46",
"metadata": {},
"source": [
"运行一个 Robbins-Monro (RM) 算法"
]
},
{
"cell_type": "code",
"execution_count": 11,
"id": "c2e13c9c",
"metadata": {},
"outputs": [
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAr8AAAGMCAYAAADTH8pwAAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjEwLjcsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvTLEjVAAAAAlwSFlzAAAPYQAAD2EBqD+naQAAaxRJREFUeJzt3XlcVFX/B/DPsC8K4oYoqLiviYm7iKSgaGaauVTu+jxmi0pq+itzq0xLBXNrUdE0pRS1xQUq3ErNBU3TzJJEEUJcAEFgGM7vj3lmZJiFmWFW+Lxfr3nBnPnOuefOuTN8OXPuuRIhhAARERERURXgYO0GEBERERFZCpNfIiIiIqoymPwSERERUZXB5JeIiIiIqgwmv0RERERUZTD5JSIiIqIqg8kvEREREVUZTH6JiIiIqMpg8ktEREREVQaTXyIiIiKqMpj8EhEREVGVweSXiIjIhowaNQq+vr7w8vLCE088ge+++87aTSKqVCRCCGHtRhAREZHc77//jubNm8PFxQW//vorwsPDcf36ddSqVcvaTSOqFDjyS0REZEPatm0LFxcXAICTkxOKioqQlpZm5VYRVR5MfquQ2NhYSCQS5c3JyQl+fn4YNWoUrl27ZlSdCxcuhEQiQVZWlsliFe38559/jGqTsUq/PocPH1Z7XAiBZs2aQSKRoE+fPhZtm7msXr0aEokE7dq10/i4tfpC1/Z/+eUXLFy4EA8ePFCLN+R4tJSy77vSt1mzZqnFmfK1Vrwe5tyGrdJ1nFRUXFwc2rZtC3d3d0gkEpw/f97k23jxxRfh5uaGTp064amnnkL79u1Nvo3SDh8+rPU4PXnypFm3DQAPHz7EjBkzUL9+fbi5uSEoKAg7d+402/YM+Wyx5c+VqvBeNgcnazeALG/z5s1o1aoVCgoK8PPPP+O9995DUlIS/vjjD/j4+Fi7eRg0aBBOnDgBPz8/q2y/evXq2Lhxo1qCe+TIEfz999+oXr26VdplDps2bQIg/5r11KlT6Nq1q5VbpErTsfDLL79g0aJFGD9+PGrUqGG9xhlI8b4rrX79+srfLXHcW/u9ZUnmOk7u3LmDMWPGYMCAAVi3bh1cXV3RokULk9WvsH37dmzZsgU//fQTrly5ovJPjDm9//77CAsLUynT9s+xKQ0bNgynT5/GBx98gBYtWuDLL7/E6NGjUVJSghdeeMHk27P3z5aq9F42Bya/VVC7du0QHBwMAOjTpw9kMhkWLFiAvXv3YsKECVZuHVCnTh3UqVPHatsfOXIktm/fjrVr18LLy0tZvnHjRnTv3h05OTlWaVd+fj48PDxMVt+ZM2dw4cIFDBo0CN9//z02btxoU8lvfn6+1Y8FUyr9vtPEEvtqi6+nqY9rc/vzzz8hlUrx0ksvITQ01KzbcnJyQkREBFavXo3mzZtj4MCBZt0eADRv3hzdunUz+3ZK279/PxITE5UJLwCEhYXhxo0bmD17NkaOHAlHR0eTbtMW3wuGsPf2WxunPZDyD/K///6rUn78+HH07dsX1atXh4eHB3r06IHvv/9eYx03b97EsGHD4OXlBW9vb7z00ku4c+eOUbGavs5RfO30+++/Y/To0fD29oavry8mTpyI7Oxslfrv3LmD//znPwgICICrqyvq1KmDnj174ocfftDr9VB8+O7YsUNZlp2djd27d2PixIlan6fP66Xvfijizp07h+HDh8PHxwdNmzY1aFvl2bhxIwDggw8+QI8ePbBz507k5+fr9dx9+/bhiSeegKurK5o0aYKYmBi1r9gNaae2/S17LCxcuBCzZ88GAAQGBmqdpvLvv//q9fr+9ttveP755+Ht7Y2aNWsiKioKxcXFuHr1KgYMGIDq1aujcePGWL58ub4vq9E07au+xzwAfP/99wgKCoKrqysCAwPx0UcflbsNY7ajb99rouu4NuSYLi9W3+PE0HrHjx+PXr16AZD/k1zeFKhZs2ahbt26KmVvvPEGJBKJSv9kZGTA1dUVGzZs0FiPTCbDX3/9pbPt9mzPnj2oVq0ann/+eZXyCRMm4Pbt2zh16pTW5/7++++QSCT4+uuvlWVnz56FRCJB27ZtVWKfeeYZdOrUCYDm91t5x0x5nyvaGPoe0+e9oOm9rM/fvmvXruGFF15A3bp14erqitatW2Pt2rXl7kNlw+SXkJKSAgAqX90dOXIETz31FLKzs7Fx40bs2LED1atXx+DBgxEXF6dWx9ChQ9GsWTPs2rULCxcuxN69e9G/f39IpdIKxZb13HPPoUWLFti9ezfmzp2LL7/8EjNnzlSJGTNmDPbu3Yt33nkHCQkJ+Pzzz9GvXz/cvXtXr9fDy8sLw4cPV04JAOSJsIODA0aOHKnxOYa+XvrsByD/KrBZs2b4+uuvlX8YDd2WJo8ePcKOHTvQuXNntGvXDhMnTkRubq7KHxBtDh48iGHDhqFWrVqIi4vD8uXLsWPHDmzZsqVCr4m2/S1t8uTJeO211wAA8fHxOHHiBE6cOIEnn3xSJU7f13fEiBHo0KEDdu/ejSlTpmDVqlWYOXMmnn32WQwaNAh79uzBU089hTfffBPx8fEqzzV07rdMJkNxcbHKTR/67MuPP/6IIUOGoHr16ti5cyc+/PBDfPXVV9i8ebPe7dNnO/r2fXnK9rMhx4o+sfoeJ4bWO3/+fGWi8P777+PEiRNYt26d1jpr1qyp8k3R/fv38emnn8LLywv37t1Tlq9ZswY1atTA+PHjkZGRgd27dyMvLw/FxcX46quvkJSUZPZRZoVXXnkFTk5O8PLyQv/+/XH8+HGd8UIIteNa202bS5cuoXXr1nByUv0y+oknnlA+rk3btm3h5+enkuD98MMPcHd3x+XLl3H79m0AQHFxMY4cOYJ+/fpprEefY0bfzxVt9Hl+RT7fy/vbd/nyZXTu3BmXLl3CihUr8N1332HQoEF4/fXXsWjRIr33o1IQVGVs3rxZABAnT54UUqlU5ObmioMHD4p69eqJ3r17C6lUqozt1q2bqFu3rsjNzVWWFRcXi3bt2gl/f39RUlIihBBiwYIFAoCYOXOmyra2b98uAIht27Ypy/SNVbQzJSVF7bnLly9Xee60adOEm5ubsj1CCFGtWjUxY8YMo1+f06dPi6SkJAFAXLp0SQghROfOncX48eOFEEK0bdtWhIaGqjzX0NervP1QxL3zzjtq7dR3W7ps3bpVABAbNmwQQgiRm5srqlWrJkJCQjS+JqX7onPnziIgIEAUFhYqy3Jzc0WtWrVE6Y8UQ9qpbX81bf/DDz9UKytbj76v74oVK1TigoKCBAARHx+vLJNKpaJOnTpi2LBhKrGOjo7iqaeeUmtDWYp90HQr/Z4ru6+GHPNdu3YV9evXF48ePVKW5eTkiJo1a6r0SUXfW/r2vTba+tmQY0XfWF3HiSb61qv4bPj666/LrXPdunUCgPL1WrRokWjbtq14/vnnxX//+18hhBD5+fmiVq1aYsmSJUIIIdLT00WvXr2El5eX8Pb2FsHBwWLfvn167UNFnDt3TkyfPl3s2bNHHD16VGzatEm0bt1aODo6ioMHD2p9nuL10OemrS+aN28u+vfvr1Z++/ZtAUC8//77Otv+0ksviSZNmijv9+vXT0yZMkX4+PiILVu2CCGE+PnnnwUAkZCQIIQw7LPFkPeIJoY8X9/jUFP7y/vb179/f+Hv7y+ys7NVyl999VXh5uYm7t27p3M/KhOO/FZB3bp1g7OzM6pXr44BAwbAx8cH+/btU/7XnZeXh1OnTmH48OGoVq2a8nmOjo4YM2YMbt26hatXr6rU+eKLL6rcHzFiBJycnJCUlKS2fUNiy3rmmWdU7j/xxBMoKChAZmamsqxLly6IjY3Fu+++i5MnT6qNKJcdjRAalroODQ1F06ZNsWnTJly8eBGnT5/WOuXBmNdLn/0A5CMFFd2WJhs3boS7uztGjRoFAMqvHI8dO6Zz5Y+8vDycOXMGzz77rHIpJsXzBw8eXOF2lt1fY+n7+j799NMq91u3bg2JRILIyEhlmZOTE5o1a4YbN26oxBYXF+PHH3/Uu01bt27F6dOnVW5lR7qM2Ze
"text/plain": [
"<Figure size 800x400 with 1 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"import matplotlib.pyplot as plt\n",
"import numpy as np\n",
"\n",
"# 真实函数 (我们假装不知道它的表达式)\n",
"def g(w):\n",
" return w**3 - 5\n",
"\n",
"# 带噪声的黑盒观测器\n",
"def noisy_observation(w):\n",
" noise = np.random.normal(0, 1) # 标准正态分布噪声\n",
" return g(w) + noise\n",
"\n",
"w_k = 0.0 # 初始猜测 w1 = 0\n",
"w_history = [w_k]\n",
"\n",
"# 迭代 100 次\n",
"for k in range(1, 101):\n",
" a_k = 0.1 / k # 满足 RM 定理的收敛步长,减小初始步长防止发散\n",
" \n",
" # 观测带噪声的输出\n",
" g_tilde = noisy_observation(w_k)\n",
" \n",
" # RM 核心更新公式\n",
" w_k = w_k - a_k * g_tilde\n",
" w_history.append(w_k)\n",
"\n",
"# 绘图展示收敛过程 (类似图 6.3)\n",
"plt.figure(figsize=(8, 4))\n",
"plt.plot(range(101), w_history, marker='o', linestyle='-', color='b')\n",
"plt.axhline(y=5**(1/3), color='r', linestyle='--', label='True Root (~1.71)')\n",
"plt.title('Robbins-Monro Algorithm: Finding root of $w^3 - 5 = 0$ with noise')\n",
"plt.xlabel('Iteration index k')\n",
"plt.ylabel('Estimated root $w_k$')\n",
"plt.legend()\n",
"plt.grid(True)\n",
"plt.show()"
]
},
{
"cell_type": "markdown",
"id": "3cc50d1f",
"metadata": {},
"source": [
"BGD, MBGD, 与 SGD 收敛性对比"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "c91caa6f",
"metadata": {},
"outputs": [
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAq8AAAGHCAYAAACedrtbAAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjEwLjcsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvTLEjVAAAAAlwSFlzAAAPYQAAD2EBqD+naQAAmWtJREFUeJzs3Xl8VNX9//HXnX0m+54QIGGTHcQNQWWpBQXFfcWquNVWW2utXdSfFluL1bbWWqtWq6hfBXHBXVGqgiBocUEBEdnCEhIC2ZPZZ87vjzszZCeTTJghfJ6PxzyS3Llz77lnbpL3nHvuOZpSSiGEEEIIIcRhwBDvAgghhBBCCNFZEl6FEEIIIcRhQ8KrEEIIIYQ4bEh4FUIIIYQQhw0Jr0IIIYQQ4rAh4VUIIYQQQhw2JLwKIYQQQojDhoRXIYQQQghx2JDwKoQQQgghDhsSXoUQQgghxGFDwqs44jz99NNomhZ52Gw28vPzmTp1Kvfeey8VFRWtXjN37lw0TYtqP06nk7lz57Js2bIYlTy+euJ45syZ0+y9aO8xZ86cmO2zK5qeM20dv1KKwYMHo2kaU6ZMOeTli5bP52PYsGH8+c9/bra8oaGBm2++mT59+mCz2Tj66KN54YUXOr3dzr5+0qRJ3Hzzzd09jIQzZ84ciouLmy2bN28er732WlzK05lyLFu2rN3zWoiEpYQ4wsyfP18Bav78+Wr16tXq448/Vi+//LK6+eabVVpamsrMzFRLly5t9ppdu3ap1atXR7Wfffv2KUD9/ve/j2Hp46cnjmfLli1q9erVkce//vUvBah58+Y1W75ly5aY7bMrwudMSkqK+tGPftTq+Y8++ijy/OTJkw99AaP04IMPqtzcXNXQ0NBs+bRp01R6erp67LHH1IcffqiuvfZaBajnn3++U9vt7OuXLVumzGaz+u6772J2TIlgy5Yt6ssvv2y2LCkpSV155ZXxKVAnylFbW6tWr16tamtrD32hhOgiCa/iiBMOImvWrGn13I4dO1S/fv1USkqKKi8v79Z+JLxGLxwCX3rppQ7XczqdKhgM9lg5WgqfM9dee62y2+2t/tH/6Ec/UhMmTFAjR45M+PDq8/lUYWGh+t3vftds+dtvv60AtWDBgmbLp02bpvr06aP8fn+H24329aNGjVLXXXddN44kPpxOZ1Tr90R49fv9yu12x70cQsSLdBsQoon+/fvzt7/9jfr6ev79739HlrfVbeDDDz9kypQpZGVlYbfb6d+/P+effz5Op5OSkhJycnIAuPvuu1td/t6yZQtXXXUVQ4YMweFwUFhYyKxZs1i3bl2zfYQv6S1cuJA77riDPn36kJqayg9/+EM2bdrUqvxLlizh1FNPJS0tDYfDwfDhw7n33nubrfP5559z1llnkZmZic1mY9y4cbz44osd1svBjgdg5cqVnHrqqaSkpOBwOJg4cSJvv/12xxXeCeFL9u+//z5XX301OTk5OBwOPB5Pm5dpoe33SynFI488wtFHH43dbicjI4MLLriAbdu2dbosl156KQALFy6MLKutreWVV17h6quvbvM1Xq+Xe+65h2HDhmG1WsnJyeGqq65i3759zdZbtGgR06dPp6CgALvdzvDhw/nd735HY2Njs/XmzJlDcnIyW7ZsYebMmSQnJ9OvXz9+9atf4fF4DnoMb7zxBqWlpVx++eXNlr/66qskJydz4YUXNlt+1VVXsWfPHj777LMOtxvt6y+//HIWLFhAfX39QctcVVXFDTfcQGFhIRaLhYEDB3LHHXc0O95x48ZxyimntHptIBCgsLCQ8847L7Kss+9JcXExZ555JosXL2bcuHHYbDbuvvvudsvZ8nzUNI3GxkaeeeaZyO9M024l5eXlXH/99fTt2xeLxcKAAQO4++678fv9kXVKSkrQNI3777+fe+65hwEDBmC1Wvnoo49wu9386le/4uijjyYtLY3MzEwmTJjA66+/3qxcHZWjvW4Db7zxBhMmTMDhcJCSksK0adNYvXp1s3XCv2cbNmzg0ksvJS0tjby8PK6++mpqa2ubrfvSSy8xfvz4yN+mgQMHtvs7I8TBSHgVooWZM2diNBr5+OOP212npKSEM844A4vFwlNPPcWSJUv485//TFJSEl6vl4KCApYsWQLANddcw+rVq1m9ejV33nknAHv27CErK4s///nPLFmyhH/961+YTCbGjx/fZii9/fbb2bFjB//5z394/PHH2bx5M7NmzSIQCETWefLJJ5k5cybBYJDHHnuMN998k5tuuondu3dH1vnoo4846aSTqKmp4bHHHuP111/n6KOP5uKLL+bpp59u93gPdjzLly/nBz/4AbW1tTz55JMsXLiQlJQUZs2axaJFizpf+R24+uqrMZvN/N///R8vv/wyZrM5qtdff/313Hzzzfzwhz/ktdde45FHHmHDhg1MnDiRvXv3dmobqampXHDBBTz11FORZQsXLsRgMHDxxRe3Wj8YDHL22Wfz5z//mdmzZ/P222/z5z//maVLlzJlyhRcLldk3c2bNzNz5kyefPJJlixZws0338yLL77IrFmzWm3X5/Nx1llnceqpp/L6669z9dVX8/e//5377rvvoMfw9ttvk5uby4gRI5otX79+PcOHD8dkMjVbPmbMmMjzHYn29VOmTKGxsfGgfS3dbjdTp07l2Wef5ZZbbuHtt9/mRz/6Effff3+zQHrVVVexcuVKNm/e3Oz177//Pnv27OGqq64ContPAL788kt+/etfc9NNN7FkyRLOP//8Dsvb1OrVq7Hb7cycOTPyO/PII48AenA94YQTeO+997jrrrt49913ueaaa7j33nu57rrrWm3roYce4sMPP+Svf/0r7777LsOGDcPj8VBVVcWtt97Ka6+9xsKFCzn55JM577zzePbZZztVjrYsWLCAs88+m9TUVBYuXMiTTz5JdXU1U6ZMYeXKla3WP//88znqqKN45ZVX+N3vfseCBQv45S9/2Wz/F198MQMHDuSFF17g7bff5q677moW0oWISrybfoU41DrqNhCWl5enhg8fHvn597//vWr66/Lyyy8rQK1du7bdbURzmd3v9yuv16uGDBmifvnLX0aWhy+jz5w5s9n6L774ogIi/XDr6+tVamqqOvnkkzu8nD5s2DA1btw45fP5mi0/88wzVUFBgQoEAl06nhNPPFHl5uaq+vr6Zsc0atQo1bdv305f4m+r20D4/briiitarX/llVeqoqKiVstbvl+rV69WgPrb3/7WbL1du3Ypu92ufvOb33RYrqbnTLiM69evV0opdfzxx6s5c+YopVSrbgMLFy5UgHrllVeabW/NmjUKUI888kib+wsGg8rn86nly5crQH399dfNjhlQL774YrPXzJw5Uw0dOrTD41BKqeHDh6vTTz+91fIhQ4ao0047rdXyPXv2RPohdyTa13u9XqVpmvrtb3/b4XYfe+yxNo/3vvvuU4B6//33lVJK7d+/X1ksFnX77bc3W++iiy5SeXl5kXM+mvekqKhIGY1GtWnTpg7LGNbW+dje5frrr79eJScnqx07djRb/te//lUBasOGDUoppbZv364ANWjQIOX1ejvcv9/vVz6fT11zzTVq3LhxnSpH+Hz+6KOPlFJKBQIB1adPHzV69Ohmfw/q6+tVbm6umjhxYmRZ+Pfs/vvvb7bNG264QdlstsjvffiYampqOiy/EJ0lLa9CtEEp1eHzRx99NBaLhR//+Mc888wzUV16BvD7/cybN48RI0ZgsVgwmUxYLBY2b97Mxo0bW61/1llnNfs53Jq1Y8cOAFatWkVdXR033HBDu6MibNmyhe+++47LLrssUobwY+bMmZSVlbXZ6nswjY2NfPbZZ1xwwQUkJydHlhuNRi6//HJ2797dpe22FE2LV0tvvfUWmqbxox/9qNlx5+fnM3bs2KjutJ48eTKDBg3iqaeeYt26daxZs6bdy59vvfUW6enpzJo1q9l+jz76aPLz85vtd9u2bcyePZv8/HyMRiNms5nJkycDtDonNE1r1SI7ZsyYyPnQkT179pCbm9vmcx2NqNGZ0Taieb3ZbCY
"text/plain": [
"<Figure size 800x400 with 1 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"结论:离真值越远,SGD下降得越快;靠近真值时,SGD 会有一定的波动,而 MBGD (使用更多样本) 波动更小 [cite: 7129-7134]。\n"
]
2026-02-28 15:27:59 +08:00
},
{
"ename": "",
"evalue": "",
"output_type": "error",
"traceback": [
"\u001b[1;31m在当前单元格或上一个单元格中执行代码时 Kernel 崩溃。\n",
"\u001b[1;31m请查看单元格中的代码,以确定故障的可能原因。\n",
"\u001b[1;31m单击<a href='https://aka.ms/vscodeJupyterKernelCrash'>此处</a>了解详细信息。\n",
"\u001b[1;31m有关更多详细信息,请查看 Jupyter <a href='command:jupyter.viewOutput'>log</a>。"
]
2026-02-28 00:18:43 +08:00
}
],
"source": [
"# 目标: 寻找二维平面上一堆散点的中心 (均值)\n",
"# 对应优化问题: 最小化均方误差 J(w)\n",
"np.random.seed(0)\n",
"n_samples = 100\n",
"# 在边长 20 的正方形内均匀分布采样,中心(期望)为 [0,0]\n",
"X = np.random.uniform(-10, 10, size=(n_samples, 2)) \n",
"\n",
"# 初始化\n",
"w_init = np.array([20.0, 20.0]) # 故意从很远的地方开始\n",
"w_sgd = w_init.copy()\n",
"w_mbgd_5 = w_init.copy()\n",
"\n",
"dist_sgd = [np.linalg.norm(w_sgd)]\n",
"dist_mbgd_5 = [np.linalg.norm(w_mbgd_5)]\n",
"\n",
"# 模拟前 30 步迭代\n",
"for k in range(1, 31):\n",
" alpha_k = 1 / k\n",
" \n",
" # SGD: 每次随机抽 1 个样本 [cite: 7173]\n",
" idx_sgd = np.random.choice(n_samples, 1)\n",
" grad_sgd = w_sgd - X[idx_sgd[0]] # 梯度: w - x\n",
" w_sgd = w_sgd - alpha_k * grad_sgd\n",
" dist_sgd.append(np.linalg.norm(w_sgd))\n",
" \n",
" # MBGD (m=5): 每次随机抽 5 个样本\n",
" idx_mbgd = np.random.choice(n_samples, 5)\n",
" grad_mbgd = w_mbgd_5 - np.mean(X[idx_mbgd], axis=0) \n",
" w_mbgd_5 = w_mbgd_5 - alpha_k * grad_mbgd\n",
" dist_mbgd_5.append(np.linalg.norm(w_mbgd_5))\n",
"\n",
"plt.figure(figsize=(8, 4))\n",
"plt.plot(dist_sgd, marker='d', label='SGD (m=1)', alpha=0.7)\n",
"plt.plot(dist_mbgd_5, marker='>', label='MBGD (m=5)', alpha=0.7)\n",
"plt.title('Distance to True Mean (0,0) over iterations')\n",
"plt.xlabel('Iteration step')\n",
"plt.ylabel('Distance to mean')\n",
"plt.legend()\n",
"plt.grid(True)\n",
"plt.show()\n",
"\n",
"print(\"结论:离真值越远,SGD下降得越快;靠近真值时,SGD 会有一定的波动,而 MBGD (使用更多样本) 波动更小。\")"
]
}
],
"metadata": {
"kernelspec": {
"display_name": "GymRL",
"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.13.9"
}
},
"nbformat": 4,
"nbformat_minor": 5
}