Files

418 lines
128 KiB
Plaintext
Raw Permalink Normal View History

2026-02-28 15:27:59 +08:00
{
"cells": [
{
"cell_type": "markdown",
"id": "255e09c4",
"metadata": {},
"source": [
"# 第 8 章:价值函数近似 (Value Function Methods)\n",
"\n",
"## 1. 为什么抛弃表格,拥抱函数?\n",
"* **存储高效**:不需要存成千上万个状态的值,只需要存函数的几个参数 $w$。\n",
"* **泛化能力 (Generalization)**:在表格法中,没见过的状态价值永远是初始值;而在函数法中,更新一个状态的参数 $w$,也会同时改善对周围相似状态的预测。\n",
"\n",
"## 2. 线性函数近似 (Linear Function Approximation)\n",
"最简单的函数就是线性函数: $\\hat{v}(s, w) = \\phi^T(s) w$。\n",
"* **$\\phi(s)$**:状态的**特征向量 (Feature Vector)**。比如在网格中,如果状态坐标是 $(x, y)$,特征向量可以是简单的 $[1, x, y]^T$(代表一个平面),也可以是多项式如 $[1, x, y, x^2, y^2, xy]^T$(代表一个曲面)。\n",
"* **TD-Linear 算法**:我们不再直接更新 $V(s)$,而是通过梯度下降更新参数 $w$\n",
" $$w_{t+1} = w_t + \\alpha [r_{t+1} + \\gamma \\phi^T(s_{t+1})w_t - \\phi^T(s_t)w_t] \\phi(s_t)$$"
]
},
{
"cell_type": "code",
"execution_count": 1,
"id": "a3119ab2",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"--- 体验 TD-Linear 的一次更新过程 ---\n",
"提取的当前特征 phi_t: [ 1. -0.5 -0.5]\n",
"TD 误差: -1.0\n",
"更新后的权重 w: [-0.01 0.005 0.005]\n",
"结论:我们没有单独记录状态 (1,1) 的价值,而是更新了统管全局的参数 w!这就是函数近似的精髓。\n"
]
}
],
"source": [
"import numpy as np\n",
"\n",
"# --- 1. 定义特征提取器 (Feature Extractor) ---\n",
"def get_features(state_x, state_y):\n",
" # 按照书中 8.2.4 的建议,将 x 和 y 归一化到 [-1, 1] 区间\n",
" # 假设网格是 5x5 (x, y 在 0 到 4 之间)\n",
" norm_x = (state_x - 2.0) / 2.0\n",
" norm_y = (state_y - 2.0) / 2.0\n",
" \n",
" # 提取简单的 3D 特征向量: [1, x, y]^T\n",
" return np.array([1.0, norm_x, norm_y])\n",
"\n",
"# --- 2. TD-Linear 参数初始化 ---\n",
"# 权重向量 w 初始设为 0 (对应特征维度 3)\n",
"w = np.zeros(3)\n",
"alpha = 0.01\n",
"gamma = 0.9\n",
"\n",
"print(\"--- 体验 TD-Linear 的一次更新过程 ---\")\n",
"# 假设智能体走了一步:从 (1, 1) 走到 (2, 1),获得奖励 -1\n",
"s_t_x, s_t_y = 1, 1\n",
"s_next_x, s_next_y = 2, 1\n",
"reward = -1.0\n",
"\n",
"# 1. 计算当前状态和下一个状态的特征\n",
"phi_t = get_features(s_t_x, s_t_y)\n",
"phi_next = get_features(s_next_x, s_next_y)\n",
"\n",
"# 2. 用当前权重 w 计算预测价值 v = w^T * phi\n",
"v_t = np.dot(w, phi_t)\n",
"v_next = np.dot(w, phi_next)\n",
"\n",
"# 3. 计算 TD 误差 (TD Error)\n",
"td_target = reward + gamma * v_next\n",
"td_error = td_target - v_t\n",
"\n",
"# 4. 更新权重 w\n",
"w = w + alpha * td_error * phi_t\n",
"\n",
"print(f\"提取的当前特征 phi_t: {phi_t}\")\n",
"print(f\"TD 误差: {td_error}\")\n",
"print(f\"更新后的权重 w: {w}\")\n",
"print(\"结论:我们没有单独记录状态 (1,1) 的价值,而是更新了统管全局的参数 w!这就是函数近似的精髓。\")"
]
},
{
"cell_type": "markdown",
"id": "5bf7d1dd",
"metadata": {},
"source": [
"## 3. 深度 Q 网络 (Deep Q-Learning, DQN)\n",
"当特征极度复杂(比如游戏画面的像素)时,人工设计 $\\phi(s)$ 根本行不通。于是我们让**神经网络**来代替线性特征,直接输入状态,输出 Q 值。\n",
"\n",
"但把神经网络和 TD 算法结合极容易崩溃。DQN 提出了两个伟大的技巧来稳定训练:\n",
"\n",
"1. **经验回放 (Experience Replay)** \n",
" * 智能体边走边把经历 $(s, a, r, s')$ 存入一个“回放缓冲区 (Replay Buffer)”。\n",
" * 每次训练时,从缓冲区**随机打乱抽取**一个小批量 (Mini-batch) 数据。\n",
" * **为什么?** 因为神经网络最怕输入“高度相关”的连续数据。随机抽样打破了数据的时间相关性,满足了数据独立同分布的假设,同时也让数据得到了高效复用。\n",
"\n",
"2. **目标网络 (Target Network)**\n",
" * 建立两个一模一样的网络:**主网络 (Main Network)** 和 **目标网络 (Target Network)**。\n",
" * 计算 TD 目标 $y_T = R + \\gamma \\max_a \\hat{q}(S', a, w_T)$ 时,用的是被冻结的**目标网络**参数 $w_T$。\n",
" * **为什么?** 如果用同一个网络,你在追逐目标的同时,目标也在狂奔(因为参数更新了)。固定目标网络一段时间,能让算法有稳定的方向,不易崩溃。"
]
},
{
"cell_type": "code",
"execution_count": 4,
"id": "628f60f5",
"metadata": {},
"outputs": [
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAAA04AAAIhCAYAAAB5deq6AAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjEwLjcsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvTLEjVAAAAAlwSFlzAAAPYQAAD2EBqD+naQAA+DRJREFUeJzs3Xdc1PUfwPHXsTeKKIiA4N57m6lpSm7NUWku3FqucpWCO7UMG5aDYeasHGmK4s7ULMuVhokoDlBxgMo6uO/vj/txeixBOQ/w/Xw87lH3+a73fb53eO/7fr6ft0pRFAUhhBBCCCGEENkyMXYAQgghhBBCCFHQSeIkhBBCCCGEEE8hiZMQQgghhBBCPIUkTkIIIYQQQgjxFJI4CSGEEEIIIcRTSOIkhBBCCCGEEE8hiZMQQgghhBBCPIUkTkIIIYQQQgjxFJI4CSGEEEIIIcRTSOIkRAGiUqly9Thw4ACXL1/WazM3N6dEiRI0bNiQ8ePH888//+T6uOn7+vTTT3Ncz8vLi4EDBz7nqyxY/P39UalUxMbGGjuUPDPm+Xj06BELFiygdu3aODg4YG9vT/ny5enduzcHDx7UrXfu3Dn8/f25fPnyMx/ryJEj+Pv7c//+/ecP/P/Gjx+PSqXi33//zXadjz76CJVKxV9//ZXr/RbFz0irVq1o1aqV7nlCQgL+/v4cOHAg07rG/jxdvXqVUaNGUalSJaytrXFycqJmzZoMHTqUq1ev6tbbsWMH/v7+BosjY5/lxcCBA7Gzs3vqejmdh6y8zH/nhcgvZsYOQAjx2NGjR/Wez549m/3797Nv3z699mrVqnH37l0A3nvvPd555x00Gg3379/n77//JigoiC+//JL58+fz4Ycf5lt8mzdvxsHBId/2J56Psc5HWloa7dq148yZM3z44Yc0atQIgP/++49t27bx66+/0rJlS0CbOM2cOZNWrVrh5eX1TMc7cuQIM2fOZODAgRQrVixfXoOvry8BAQEEBQWxcOHCTMs1Gg3fffcdderUoV69evlyzMJq6dKles8TEhKYOXMmwDMnB4Zw7do16tWrR7FixZg4cSKVK1cmLi6Oc+fOsXHjRi5duoSHhwegTZy+/vprgyVPGfvMEAx1HuTvvBDZk8RJiAKkSZMmes9LliyJiYlJpnZAlzh5enrqLe/QoQMTJkygR48eTJo0iRo1avDGG2/kS3x169bNl/0YilqtRqVSYWZW+P60paWlkZqaiqWlZa63Mdb5OHToEEeOHCEoKIhBgwbp2tu3b8+YMWPQaDRGiSsvatSoQaNGjVi9ejXz5s3L9J7ZvXs3165dY/LkyUaKsOCoVq2asUPIlRUrVhAbG8vx48fx9vbWtXfr1o1p06Y98/tSURSSkpKwtrbO9TaFpc+yUtD/zgthTDJUT4giyNramsDAQMzNzVm0aFG+7TfjEI4DBw6gUqlYt24dH330EW5ubjg4ONC2bVvCw8Mzbb9nzx7atGmDg4MDNjY2NG/enL179+qtc/HiRQYNGkTFihWxsbGhTJkydO7cmTNnzuitl37s1atXM3HiRMqUKYOlpSUXL17Mt9f7pD///JMuXbrg5OSElZUVdevWZePGjXrr3L59m1GjRlGtWjXs7OwoVaoUr732Gr/++qveeulDZhYuXMicOXPw9vbG0tKS/fv364Y6/fPPP7z99ts4Ojri4uLC4MGDiYuL09vP85wPRVGYN28eZcuWxcrKigYNGhAWFparIUZ37twBoHTp0lkuNzHR/tMSEhJCr169AGjdurVuWGlISAgAYWFhdO3aFXd3d6ysrKhQoQLDhw/XG+bl7++vu2rq7e2tN1w13YYNG2jatCm2trbY2dnRvn17/v777xxfA2ivOsXExLBz585My4KDg7G0tKRv374kJSUxceJE6tSpg6OjI05OTjRt2pStW7c+9RghISGoVKpMQxXTz1XGYVa5+Yzcvn2bYcOG4eHhgaWlJSVLlqR58+bs2bMn2zj++ecfVCoVP/zwg67txIkTqFQqqlevrrduly5dqF+/vu75k++Jy5cvU7JkSQBmzpypOx8Zh3bdvHnzqe/frOTmPZGdO3fuYGJiQqlSpbJcnv6+HDhwIF9//TWgPzw6/RypVCrGjBnDt99+S9WqVbG0tGTVqlW619y4cWOcnJxwcHCgXr16BAYGoiiK3rGy+hxdu3aNnj17Ym9vT7Fixejbty9//PGH3mfiSRcvXqRDhw7Y2dnh4eHBxIkTSU5OBnJ/Hp7Fi/g7L0RhJYmTEEWUm5sb9evX58iRI6Smphr0WNOmTePKlSusXLmS5cuX899//9G5c2fS0tJ063z//fe0a9cOBwcHVq1axcaNG3FycqJ9+/Z6/6jeuHGDEiVK8MknnxAaGsrXX3+NmZkZjRs3zvIf6alTpxIVFcW3337Ltm3bKFWqlC4xya9x+vv376d58+bcv3+fb7/9lq1bt1KnTh369Omj94Un/Sqgn58fv/zyC8HBwZQrV45WrVpleR/CF198wb59+/j000/ZuXMnVapU0S178803qVSpEj/99BNTpkxh7dq1jB8/Plfx5uZ8fPTRR3z00Uf4+PiwdetWRowYwZAhQ7hw4cJT99+gQQPMzc0ZO3Ysa9asITo6Osv1OnbsyLx58wD4+uuvOXr0KEePHqVjx44ARERE0LRpU7755ht2797NjBkz+P3333nllVdQq9UADBkyhPfeew+ATZs26faRPnxu3rx5vP3221SrVo2NGzeyevVqHjx4QIsWLTh37lyOr+Ptt9/GxsaGoKAgvfZ79+6xdetWunfvTvHixUlOTubu3bt88MEHbNmyhXXr1vHKK6/Qo0cPvvvuu6f2V27l9jPy7rvvsmXLFmbMmMHu3btZuXIlbdu21SW0WalevTqlS5fWS6727NmDtbU1586d48aNGwCkpqZy8OBB2rZtm+V+SpcuTWhoKKBNPNPPx/Tp0/XWe9b3b27eE9lp2rQpGo2GHj16sGvXLuLj47Ncb/r06fTs2RNAF//Ro0f1fgjYsmUL33zzDTNmzGDXrl20aNEC0CYsw4cPZ+PGjWzatIkePXrw3nvvMXv27Bxje/ToEa1bt2b//v0sWLCAjRs34uLiQp8+fbJcX61W06VLF9q0acPWrVsZPHgwn3/+OQsWLAByfx7yU37+nRei0FKEEAXWgAEDFFtb2yyXRUZGKoCyaNGibLfv06ePAig3b97M8Ti52ZeiKErZsmWVAQMG6J7v379fAZQOHTrorbdx40YFUI4ePaooiqI8evRIcXJyUjp37qy3XlpamlK7dm2lUaNG2R4zNTVVSUlJUSpWrKiMHz8+07FfffXVTNtcvnxZMTU1VQYPHpzj61EURfHz81MA5fbt29muU6VKFaVu3bqKWq3Wa+/UqZNSunRpJS0tLdvY1Wq10qZNG6V79+669vT+Ll++vJKSkpJlPAsXLtRrHzVqlGJlZaVoNBpd27Oej7t37yqWlpZKnz599NY7evSoAigtW7bMti/SBQYGKnZ2dgqgAErp0qWV/v37K4cOHdJb74cfflAAZf/+/TnuT6PRKGq1Wrly5YoCKFu3btUtW7RokQIokZGRettERUUpZmZmynvvvafX/uDBA8XV1VXp3bv3U1/HgAEDFHNzc73PyJdffqkASlhYWJbbpJ9XX19fpW7dunrLMp6T4ODgLGNPP1fp/ZKXz4idnZ0ybty4p762jPr166eUK1dO97xt27bK0KFDleLFiyurVq1SFEVRfvvtNwVQdu/erVuvZcuWeu+J27dvK4Di5+eX6Rh5ef8+TU7viezWHz58uGJiYqIAikqlUqpWraqMHz8+U/+PHj1aye4rEKA4Ojoqd+/ezfF4aWlpilqtVmbNmqWUKFFC77Vl7LOvv/5aAZSdO3fq7WP48OEKoAQHB+vaBgwYoADKxo0b9dbt0KGDUrlyZd3znM5DVgry33khCgu54iREEaZkGD6Smpqq98i4/Fl16dJF73mtWrUAuHLlCqC9uf/u3bsMGDBA7/gajQYfHx/++OMPHj16pItx3rx5VKtWDQsLC8zMzLCwsOC
"text/plain": [
"<Figure size 1000x600 with 1 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"最终学到的参数 w: [0.177 0.888]\n"
]
}
],
"source": [
"import numpy as np\n",
"import matplotlib.pyplot as plt\n",
"\n",
"# --- 1. 定义环境 ---\n",
"def step(state):\n",
" # 随机向左或向右走\n",
" action = np.random.choice([-1, 1])\n",
" next_state = state + action\n",
" # 到达最右侧(6)奖励为1,其余为0\n",
" reward = 1.0 if next_state == 6 else 0.0\n",
" done = next_state == 0 or next_state == 6\n",
" return next_state, reward, done\n",
"\n",
"# --- 2. 定义特征提取器 (Feature Extractor) ---\n",
"# 对应书中公式:phi(s) = [1, x]^T\n",
"def get_features(state):\n",
" # 将状态归一化到 [0, 1] 之间,这是机器学习的好习惯\n",
" norm_s = state / 6.0\n",
" return np.array([1.0, norm_s])\n",
"\n",
"# --- 3. 运行 TD-Linear 算法 ---\n",
"# 初始化参数 w 为 0 (只有两个参数!不管走廊有多长,都只需要2个参数)\n",
"w = np.zeros(2) \n",
"alpha = 0.1 # 学习率\n",
"gamma = 1.0 # 无折扣\n",
"\n",
"# 记录每个 episode 后的状态价值,用于画图\n",
"history_v = []\n",
"\n",
"num_episodes = 200\n",
"for ep in range(num_episodes):\n",
" state = 3 # 从中间开始\n",
" while True:\n",
" next_state, reward, done = step(state)\n",
" \n",
" # 提取当前状态和下一个状态的特征 phi\n",
" phi_t = get_features(state)\n",
" phi_next = get_features(next_state)\n",
" \n",
" # 计算当前的预测价值 v = w^T * phi \n",
" v_t = np.dot(w, phi_t)\n",
" # 如果游戏结束,下一个状态的价值必定是 0\n",
" v_next = 0.0 if done else np.dot(w, phi_next)\n",
" \n",
" # 计算 TD 误差\n",
" td_target = reward + gamma * v_next\n",
" td_error = td_target - v_t\n",
" \n",
" # 核心:使用 TD-Linear 公式更新参数 w\n",
" # w_{t+1} = w_t + alpha * (TD_target - v_t) * phi(s_t)\n",
" w = w + alpha * td_error * phi_t\n",
" \n",
" state = next_state\n",
" if done:\n",
" break\n",
" \n",
" # 每局结束后,把当前 w 眼中的“所有状态价值”记录下来\n",
" current_estimated_v = [np.dot(w, get_features(s)) for s in range(1, 6)]\n",
" history_v.append(current_estimated_v)\n",
"\n",
"# --- 4. 可视化学习过程 ---\n",
"true_values = [1/6, 2/6, 3/6, 4/6, 5/6]\n",
"\n",
"plt.figure(figsize=(10, 6))\n",
"plt.plot(range(1, 6), true_values, 'r-o', linewidth=2, label='True Values')\n",
"\n",
"# 画出第 10, 50, 199 局的拟合结果\n",
"for ep in [10, 50, 199]:\n",
" plt.plot(range(1, 6), history_v[ep], '--', label=f'Estimated Values @ Ep {ep}')\n",
"\n",
"plt.title('TD-Linear: Learning State Values with a Straight Line')\n",
"plt.xlabel('State')\n",
"plt.ylabel('Value Estimate')\n",
"plt.legend()\n",
"plt.grid(True)\n",
"plt.show()\n",
"\n",
"print(f\"最终学到的参数 w: {np.round(w, 3)}\")"
]
},
{
"cell_type": "code",
"execution_count": 5,
"id": "0da2810e",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"开始训练 DQN,让神经网络去感受陷阱和宝箱...\n",
"训练完成!\n",
"\n",
"--- 揭晓神经网络眼中的 Q 值 ---\n",
"状态 0: Q(向左)= 0.34, Q(向右)= 0.13 => 最优决策: 向左\n",
"状态 1: Q(向左)= 0.32, Q(向右)= -3.04 => 最优决策: 向左\n",
"状态 2: Q(向左)= -2.48, Q(向右)= 8.08 => 最优决策: 向右\n",
"状态 3: Q(向左)= -0.92, Q(向右)= 8.95 => 最优决策: 向右\n",
"状态 4: Q(向左)= 4.09, Q(向右)= 9.81 => 最优决策: 向右\n",
"状态 5: Q(向左)= 9.18, Q(向右)= 10.66 => 最优决策: 向右\n"
]
},
{
"data": {
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAjMAAAHFCAYAAAAHcXhbAAAAOnRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjEwLjcsIGh0dHBzOi8vbWF0cGxvdGxpYi5vcmcvTLEjVAAAAAlwSFlzAAAPYQAAD2EBqD+naQAAVT5JREFUeJzt3XlYVNUfBvB3RBkBYcyFTRH3FXdTUXPNlUxTy7RMs7JSK39mrplkKqZpZou2KG7lUi4tmoIiuIALKG4oaiKigAiyL8N2fn8QN0aGZYaBmQvv53nmibnrd67kvJ5z7rkKIYQAERERkUxVM3YBRERERGXBMENERESyxjBDREREssYwQ0RERLLGMENERESyxjBDREREssYwQ0RERLLGMENERESyxjBDREREssYwQ2SCtmzZAoVCIb1q1qwJe3t7DBgwAB4eHoiJiSly38OHD8PNzQ3169eHUqlEo0aN8PrrryM0NLTQtu7u7lAoFLC1tUVycnKh9Y0bN8Zzzz1X6jqLejVu3Fiv6/BkLVOmTNFr3ylTphikBn3PXatWLaOcm6iqqG7sAoioaJ6enmjdujWysrIQExODU6dO4fPPP8cXX3yB3bt349lnn9XYfu7cuVi9ejWGDRuG7777DnZ2drh58ybWrl2Lzp07Y8+ePVrDyaNHj7Bq1Sp89tlnOtXn5uaGgIAAjWWurq4YN24cPvzwQ2mZUqnU6bja7N+/HzY2Nnrtu3jxYnzwwQdlroGITBPDDJEJc3FxQbdu3aT3Y8eOxf/+9z/06dMHY8aMwa1bt2BnZwcA2LlzJ1avXo13330X3333nbRP3759MWHCBPTr1w8TJ07E1atX0ahRI43zDBs2DF9++SVmzJgBe3v7UtdXv3591K9fv9ByOzs79OzZs8j9cnJykJ2drVPI6dy5c6m3fVKzZs303peITB+7mYhkplGjRlizZg2Sk5Px/fffS8uXL1+Op556Cl988UWhfaysrPD1118jOTkZ69atK7R+2bJlyM7Ohru7u8HrvXv3LhQKBVatWoVly5ahSZMmUCqVOH78ODIyMvDhhx+iU6dOUKlUqFOnDlxdXfH7778XOs6T3Uy+vr5QKBTYuXMnFi1aBEdHR9jY2ODZZ58t1KWmrZtJoVBg5syZ2L59O9q0aQNLS0t07NgRf/31V6Fz//777+jQoQOUSiWaNm2Kr776SuqiM5TNmzejY8eOqFmzJurUqYMXXngB169f19jmzp07ePnll+Ho6AilUgk7OzsMGjQIwcHB0jY+Pj7o378/6tatCwsLCzRq1Ahjx45FWlqawWolMjVsmSGSoREjRsDMzAwnTpwAAERFReHatWsYP348LC0tte7j6uoKW1tbHDlypNA6Z2dnTJ8+HV9//TVmz56Nli1bGrzm9evXo2XLlvjiiy9gY2ODFi1aQK1W4/Hjx5gzZw4aNGiAzMxMHD16FGPGjIGnpydee+21Eo+7cOFC9O7dGz/99BOSkpIwb948jBw5EtevX4eZmVmx+x48eBDnz5/H0qVLUatWLaxatQovvPACQkND0bRpUwB5Y5DGjBmDvn37Yvfu3cjOzsYXX3yBhw8fGuS6AICHhwcWLlyICRMmwMPDA3FxcXB3d4erqyvOnz+PFi1aAMj7c8/JycGqVavQqFEjxMbGwt/fHwkJCQDygqObmxueeeYZbN68GbVr18aDBw9w+PBhZGZmFvm7QSR7gohMjqenpwAgzp8/X+Q2dnZ2ok2bNkIIIc6cOSMAiPnz5xd73B49eggrKyvp/ZIlSwQA8ejRIxEbGytUKpUYO3astN7Z2Vm4ubnpVDsAMWPGDOl9WFiYACCaNWsmMjMzi903OztbZGVliTfeeEN07txZY52zs7OYPHmy9P748eMCgBgxYoTGdnv27BEAREBAgLRs8uTJwtnZuVCddnZ2IikpSVoWHR0tqlWrJjw8PKRlTz/9tHBychJqtVpalpycLOrWrStK81fo5MmTNa75k+Lj44WFhUWhz3Hv3j2hVCrFxIkThRBCxMbGCgBi3bp1RR7rt99+EwBEcHBwiXURVSbsZiKSKSGEXvsU1TVSt25dzJs3D3v37sXZs2fLWl4hzz//PGrUqFFo+a+//orevXujVq1aqF69OmrUqIFNmzYV6mIp7rgFdejQAQAQHh5e4r4DBgyAtbW19N7Ozg62trbSvqmpqQgMDMTo0aNhbm4ubVerVi2MHDmyVPWVJCAgAOnp6YXu1HJycsLAgQNx7NgxAECdOnXQrFkzrF69GmvXrsXFixeRm5ursU+nTp1gbm6OadOmYevWrbhz545BaiQydQwzRDKUmpqKuLg4ODo6AoA0oDcsLKzY/cLDw+Hk5FTk+lmzZsHR0RFz5841XLH/cnBwKLRs3759eOmll9CgQQPs2LEDAQEBOH/+PKZOnYqMjIxSHbdu3boa7/MHFaenp+u8b/7++fvGx8dDCCENsi5I2zJ9xMXFAdB+fRwdHaX1CoUCx44dw9ChQ7Fq1Sp06dIF9evXx/vvvy/dVt+sWTMcPXoUtra2mDFjBpo1a4ZmzZrhq6++MkitRKaKY2aIZOjgwYPIyclB//79AeR9Ebq4uMDLywtpaWlax0YEBATg4cOHGDduXJHHtbCwgLu7O6ZNm4aDBw8atGZtLUI7duxAkyZNsHv3bo31arXaoOfW11NPPQWFQqF1fEx0dLRBzpEfqKKiogqti4yMRL169aT3zs7O2LRpEwDg5s2b2LNnD9zd3ZGZmYmNGzcCAJ555hk888wzyMnJQWBgIL7++mvMmjULdnZ2ePnllw1SM5GpYcsMkczcu3cPc+bMgUqlwttvvy0tX7RoEeLj4zFnzpxC+6SmpuL999+Hubk5pk+fXuzxp06dijZt2mD+/PmFujEMTaFQwNzcXCPIREdHa72byRisrKzQrVs3HDhwAJmZmdLylJQUrXc96cPV1RUWFhbYsWOHxvL79+/Dx8cHgwYN0rpfy5Yt8fHHH6N9+/a4cOFCofVmZmbo0aMHvv32WwDQug1RZcGWGSITdvXqVWRnZyM7OxsxMTE4efIkPD09YWZmhv3792vM8fLyyy8jKCgIX3zxBe7evYupU6fCzs4OoaGh+PLLL3Hjxg1s2rQJbdu2LfacZmZmWLFiBV544QUA/41BKQ/PPfcc9u3bh+nTp2PcuHGIiIjAZ599BgcHB9y6davczquLpUuXws3NDUOHDsUHH3yAnJwcrF69GrVq1cLjx49LdYycnBz89ttvhZZbWVlh+PDhWLx4MRYuXIjXXnsNEyZMQFxcHD799FPUrFkTS5YsAQBcvnwZM2fOxIsvvogWLVrA3NwcPj4+uHz5MubPnw8A2LhxI3x8fODm5oZGjRohIyMDmzdvBoBCEywSVSYMM0Qm7PXXXwcAmJubo3bt2mjTpg3mzZuHN998U+tkdatXr8aAAQPwzTff4O2335bGfNja2sLf3x89evQo1XlHjx6NXr16wd/f36Cf50mvv/46YmJisHHjRmzevBlNmzbF/Pnzcf/+fXz66afleu7SGjZsGPbu3YtPPvkE48ePh729PaZPn47IyEhs3769VMfIyMjAiy++WGi5s7Mz7t69iwULFsDW1hbr16/H7t27YWFhgf79+2PFihXSbdn29vZo1qwZvvvuO0REREChUKBp06ZYs2YN3nvvPQB5A4C9vLywZMkSREdHo1atWnBxccEff/yBIUOGGO6iEJkYhdDnlggiko2lS5diyZIl+Pbbb0vsYqLSycrKQqdOndCgQQN4eXkZuxyiKo8tM0SV3CeffIKoqCjMnDkTVlZWmDx5srFLkp033ngDgwcPhoODA6Kjo7Fx40Zcv36ddwkRmQi2zBARleCll16Cv78/Hj16hBo1aqBLly5YuHAhhg0bZuzSiAgMM0RERCRzvDWbiIiIZI1hhoiIiGSNYYaIiIhkrdLfzZSbm4vIyEhYW1sX+YA9IiIiMi1CCCQnJ8PR0RHVqhXf9lLpw0xkZGSxD9YjIiIi0xUREYGGDRsWu02lDzPW1tYA8i6GjY2NkashIiKi0khKSoKTk5P0PV6
"text/plain": [
"<Figure size 640x480 with 1 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"import torch\n",
"import torch.nn as nn\n",
"import torch.optim as optim\n",
"import numpy as np\n",
"import random\n",
"from collections import deque\n",
"import matplotlib.pyplot as plt\n",
"\n",
"# --- 1. 定义非线性环境 ---\n",
"class NonLinearCorridor:\n",
" def __init__(self):\n",
" self.state = 0\n",
" def step(self, action):\n",
" # action: 0向左, 1向右\n",
" move = -1 if action == 0 else 1\n",
" self.state = max(0, min(5, self.state + move))\n",
" \n",
" # 奖励机制:状态 5 是宝箱,状态 2 是陷阱\n",
" if self.state == 5:\n",
" return self.state, 10.0, True\n",
" elif self.state == 2:\n",
" return self.state, -10.0, False\n",
" else:\n",
" return self.state, 0.0, False\n",
" def reset(self):\n",
" self.state = 0\n",
" return self.state\n",
"\n",
"# --- 2. 定义神经网络 Q-Network ---\n",
"# 输入是状态(1维),隐藏层提取非线性特征,输出是2个动作的Q值\n",
"class QNetwork(nn.Module):\n",
" def __init__(self):\n",
" super(QNetwork, self).__init__()\n",
" # 浅层网络足以解决简单的网格世界问题\n",
" self.fc1 = nn.Linear(1, 16)\n",
" self.relu = nn.ReLU()\n",
" self.fc2 = nn.Linear(16, 2)\n",
"\n",
" def forward(self, x):\n",
" x = self.relu(self.fc1(x))\n",
" return self.fc2(x)\n",
"\n",
"# --- 3. 经验回放缓冲区 (Experience Replay)---\n",
"class ReplayBuffer:\n",
" def __init__(self, capacity=1000):\n",
" self.buffer = deque(maxlen=capacity)\n",
" def add(self, state, action, reward, next_state, done):\n",
" self.buffer.append((state, action, reward, next_state, done))\n",
" def sample(self, batch_size):\n",
" # 均匀随机抽样,打破数据的时间相关性\n",
" return random.sample(self.buffer, batch_size)\n",
" def __len__(self):\n",
" return len(self.buffer)\n",
"\n",
"# --- 4. DQN 核心训练逻辑 ---\n",
"env = NonLinearCorridor()\n",
"\n",
"# 初始化主网络和目标网络,并让它们参数一致\n",
"main_net = QNetwork()\n",
"target_net = QNetwork()\n",
"target_net.load_state_dict(main_net.state_dict())\n",
"\n",
"optimizer = optim.Adam(main_net.parameters(), lr=0.01)\n",
"loss_fn = nn.MSELoss()\n",
"buffer = ReplayBuffer(capacity=2000)\n",
"\n",
"batch_size = 32\n",
"gamma = 0.9\n",
"epsilon = 0.3 # 探索率\n",
"update_target_every = 20 # 每隔 C 步更新一次目标网络\n",
"\n",
"episodes = 300\n",
"loss_history = []\n",
"step_count = 0\n",
"\n",
"print(\"开始训练 DQN,让神经网络去感受陷阱和宝箱...\")\n",
"\n",
"for ep in range(episodes):\n",
" state = env.reset()\n",
" done = False\n",
" \n",
" while not done:\n",
" step_count += 1\n",
" \n",
" # --- 策略:Epsilon-Greedy ---\n",
" if random.random() < epsilon:\n",
" action = random.choice([0, 1])\n",
" else:\n",
" state_tensor = torch.FloatTensor([[state]])\n",
" with torch.no_grad():\n",
" q_values = main_net(state_tensor)\n",
" action = torch.argmax(q_values).item()\n",
" \n",
" next_state, reward, done = env.step(action)\n",
" \n",
" # 存入经验回放池\n",
" buffer.add(state, action, reward, next_state, done)\n",
" state = next_state\n",
" \n",
" # --- 训练阶段 ---\n",
" if len(buffer) >= batch_size:\n",
" # 1. 抽取 Mini-batch\n",
" batch = buffer.sample(batch_size)\n",
" b_s = torch.FloatTensor([[x[0]] for x in batch])\n",
" b_a = torch.LongTensor([[x[1]] for x in batch])\n",
" b_r = torch.FloatTensor([[x[2]] for x in batch])\n",
" b_ns = torch.FloatTensor([[x[3]] for x in batch])\n",
" b_d = torch.FloatTensor([[x[4]] for x in batch])\n",
" \n",
" # 2. 计算当前 Q 值预测: main_net(S)[A]\n",
" q_pred = main_net(b_s).gather(1, b_a)\n",
" \n",
" # 3. 计算目标 Q 值 (使用 Target Network)\n",
" # y_T = R + gamma * max_a Q_target(S', a)\n",
" with torch.no_grad():\n",
" max_q_next = target_net(b_ns).max(1, keepdim=True)[0]\n",
" q_target = b_r + gamma * max_q_next * (1 - b_d)\n",
" \n",
" # 4. 反向传播更新网络\n",
" loss = loss_fn(q_pred, q_target)\n",
" optimizer.zero_grad()\n",
" loss.backward()\n",
" optimizer.step()\n",
" loss_history.append(loss.item())\n",
" \n",
" # 5. 定期同步目标网络\n",
" if step_count % update_target_every == 0:\n",
" target_net.load_state_dict(main_net.state_dict())\n",
"\n",
"print(\"训练完成!\")\n",
"\n",
"# --- 5. 验证神经网络学到了什么 ---\n",
"print(\"\\n--- 揭晓神经网络眼中的 Q 值 ---\")\n",
"states_to_test = torch.FloatTensor([[s] for s in range(6)])\n",
"with torch.no_grad():\n",
" learned_q = main_net(states_to_test).numpy()\n",
"\n",
"for s in range(6):\n",
" q_left, q_right = learned_q[s]\n",
" best_act = \"向左\" if q_left > q_right else \"向右\"\n",
" print(f\"状态 {s}: Q(向左)={q_left:6.2f}, Q(向右)={q_right:6.2f} => 最优决策: {best_act}\")\n",
"\n",
"# 画出 Loss 曲线 (证明它在收敛)\n",
"plt.plot(loss_history)\n",
"plt.title(\"DQN Training Loss\")\n",
"plt.xlabel(\"Training Steps\")\n",
"plt.ylabel(\"Loss (MSE)\")\n",
"plt.show()"
]
}
],
"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
}