提交代码

This commit is contained in:
2026-02-28 00:18:43 +08:00
parent 1f0096be07
commit 2537436fb7
3 changed files with 829 additions and 0 deletions
+302
View File
File diff suppressed because one or more lines are too long
+246
View File
File diff suppressed because one or more lines are too long
+281
View File
@@ -0,0 +1,281 @@
{
"cells": [
{
"cell_type": "markdown",
"id": "b83792ff",
"metadata": {},
"source": [
"# 第 7 章:时间差分方法 (Temporal-Difference Methods) - 上篇\n",
"\n",
"## 1. TD(0):走一步,学一步\n",
"在不知道环境模型(无模型)的情况下,我们要评估一个策略的好坏 $V(s)$。\n",
"相比于蒙特卡洛(MC)必须等一个回合结束(比如游戏 Game Over)才能复盘,**TD 算法允许智能体每走一步就更新一次自己的认知**。\n",
"\n",
"## 2. 核心更新公式\n",
"$$V(s_t) \\leftarrow V(s_t) + \\alpha [r_{t+1} + \\gamma V(s_{t+1}) - V(s_t)]$$\n",
"\n",
"我们来拆解这个美妙的公式:\n",
"* **$r_{t+1} + \\gamma V(s_{t+1})$**:这被称为 **TD 目标 (TD Target)**。它用实际走这一步拿到的“真实奖励”,加上对下一个状态的“现有估值”,来作为当前状态的新目标。\n",
"* **$r_{t+1} + \\gamma V(s_{t+1}) - V(s_t)$**:这被称为 **TD 误差 (TD Error)**。它衡量了“我刚刚经历的真实情况与我对未来的预测”加上“我对下一步的预测”,和“我原来对当前步的预测”之间的偏差。\n",
"* **$\\alpha$**:学习率(相当于第六章的步长)。\n",
"\n",
"**核心思想(自举 Bootstrapping)**:用自己对下一步的估计,来更新对当前步的估计。就像是“左脚踩右脚上天”。"
]
},
{
"cell_type": "code",
"execution_count": 1,
"id": "e46747a5",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"=== TD(0) 算法:边走边学演示 ===\n",
"初始状态 V 表: [0. 0. 0. 0. 0.]\n",
"\n",
"[第 1 步] 在状态 3 决定 向右,进入了状态 4。拿到真实奖励 0.0。\n",
" -> 我猜状态 4 的价值是 0.000。\n",
" -> 所以我的 TD 目标是 0.0 + 1.0 * 0.000 = 0.000。\n",
" -> 我把状态 3 的价值从 0.000 更新为了 0.000。\n",
"\n",
"[第 2 步] 在状态 4 决定 向右,进入了状态 5。拿到真实奖励 0.0。\n",
" -> 我猜状态 5 的价值是 0.000。\n",
" -> 所以我的 TD 目标是 0.0 + 1.0 * 0.000 = 0.000。\n",
" -> 我把状态 4 的价值从 0.000 更新为了 0.000。\n",
"\n",
"[第 3 步] 在状态 5 决定 向右,进入了状态 6。拿到真实奖励 1.0。\n",
" -> 我猜状态 6 的价值是 0.000。\n",
" -> 所以我的 TD 目标是 1.0 + 1.0 * 0.000 = 1.000。\n",
" -> 我把状态 5 的价值从 0.000 更新为了 0.100。\n",
"\n",
"游戏结束!最终到达终点 6。\n",
"\n",
"跑完这 1 局后的最新状态 V 表 (状态1到5): [0. 0. 0. 0. 0.1]\n",
"仔细看:相比于 MC 要等游戏结束,TD 在游戏过程中就已经把前面的状态价值更新了!\n"
]
}
],
"source": [
"import numpy as np\n",
"\n",
"# --- 1. 定义简单的 1D 随机游走环境 ---\n",
"# 状态: 0, 1, 2, 3, 4, 5, 6 (0 和 6 是终点)\n",
"# 开始状态是 3 (正中间)\n",
"# 规则: 每次有 50% 概率向左,50% 概率向右\n",
"# 奖励: 只有走到 6 (右边终点) 时奖励为 1,走到 0 (左边终点) 奖励为 0,其他全是 0。\n",
"def step(state):\n",
" action = np.random.choice([-1, 1]) # -1向左, 1向右\n",
" next_state = state + action\n",
" \n",
" # 判断奖励和是否结束\n",
" reward = 1.0 if next_state == 6 else 0.0\n",
" done = (next_state == 0 or next_state == 6)\n",
" \n",
" return next_state, reward, done\n",
"\n",
"# --- 2. 运行 TD(0) 算法 ---\n",
"print(\"=== TD(0) 算法:边走边学演示 ===\")\n",
"\n",
"# 初始化状态价值 V(s) 全为 0 (终点 0 和 6 的价值始终为 0)\n",
"V_td = np.zeros(7)\n",
"alpha = 0.1 # 学习率\n",
"gamma = 1.0 # 假设无折扣因子,简化理解\n",
"\n",
"# 我们只跑 1 局游戏,仔细看看里面发生了什么!\n",
"state = 3 # 从中间开始\n",
"step_count = 0\n",
"\n",
"print(f\"初始状态 V 表: {np.round(V_td[1:6], 3)}\")\n",
"\n",
"while True:\n",
" step_count += 1\n",
" next_state, reward, done = step(state)\n",
" \n",
" # 【核心!】TD 走完这一步立刻开始算账\n",
" td_target = reward + gamma * V_td[next_state]\n",
" td_error = td_target - V_td[state]\n",
" \n",
" # 记录下更新前的 V 值,方便打印对比\n",
" old_v = V_td[state]\n",
" \n",
" # 更新 V(s)\n",
" V_td[state] = V_td[state] + alpha * td_error\n",
" \n",
" # 打印超级详细的“内心独白”\n",
" action_str = \"向右\" if next_state > state else \"向左\"\n",
" print(f\"\\n[第 {step_count} 步] 在状态 {state} 决定 {action_str},进入了状态 {next_state}。拿到真实奖励 {reward}。\")\n",
" print(f\" -> 我猜状态 {next_state} 的价值是 {V_td[next_state]:.3f}。\")\n",
" print(f\" -> 所以我的 TD 目标是 {reward} + {gamma} * {V_td[next_state]:.3f} = {td_target:.3f}。\")\n",
" print(f\" -> 我把状态 {state} 的价值从 {old_v:.3f} 更新为了 {V_td[state]:.3f}。\")\n",
" \n",
" state = next_state\n",
" if done:\n",
" print(f\"\\n游戏结束!最终到达终点 {state}。\")\n",
" break\n",
"\n",
"print(f\"\\n跑完这 1 局后的最新状态 V 表 (状态1到5): {np.round(V_td[1:6], 3)}\")\n",
"print(\"仔细看:相比于 MC 要等游戏结束,TD 在游戏过程中就已经把前面的状态价值更新了!\")"
]
},
{
"cell_type": "code",
"execution_count": 2,
"id": "6057ac63",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"=== TD(0) 经过 500 局学习后的最终估值 ===\n",
"状态 1-5 学习到的 V 值: [0.123 0.311 0.416 0.688 0.864]\n",
"真实的理论 V 值 : [0.167, 0.333, 0.5 , 0.667, 0.833]\n",
"结论:TD(0) 完美地学会了评估这个策略的真实价值!\n"
]
}
],
"source": [
"# 重新初始化\n",
"V_td_500 = np.zeros(7)\n",
"# 为了让它在终点附近也能学到正确的平均值,通常设定非终点状态价值初始值为 0.5 会收敛更快\n",
"V_td_500[1:6] = 0.5 \n",
"alpha = 0.1\n",
"gamma = 1.0\n",
"\n",
"for episode in range(500):\n",
" state = 3\n",
" while True:\n",
" next_state, reward, done = step(state)\n",
" # TD(0) 更新\n",
" td_target = reward + gamma * V_td_500[next_state]\n",
" V_td_500[state] = V_td_500[state] + alpha * (td_target - V_td_500[state])\n",
" \n",
" state = next_state\n",
" if done:\n",
" break\n",
"\n",
"print(\"=== TD(0) 经过 500 局学习后的最终估值 ===\")\n",
"print(\"状态 1-5 学习到的 V 值: \", np.round(V_td_500[1:6], 3))\n",
"print(\"真实的理论 V 值 : [0.167, 0.333, 0.5 , 0.667, 0.833]\")\n",
"print(\"结论:TD(0) 完美地学会了评估这个策略的真实价值!\")"
]
},
{
"cell_type": "markdown",
"id": "ab8ef2b0",
"metadata": {},
"source": [
"## 3. Sarsa:稳扎稳打的“老实人”\n",
"Sarsa 是一个 **同策略(On-policy)**算法。所谓“同策略”,也就是 **“知行合一”**:智能体用来在环境中瞎逛生成数据的策略(行为策略 Behavior Policy),和它在内心中不断更新优化的策略(目标策略 Target Policy)是**同一个策略**。\n",
"\n",
"**为什么叫 Sarsa**\n",
"因为更新一次 Q 值,刚好需要这 5 个元素连在一起:当前状态 $S_t$、当前动作 $A_t$、得到的奖励 $R_{t+1}$、下一个状态 $S_{t+1}$、以及在下一个状态**实际采取**的动作 $A_{t+1}$ 。连起来就是 S-A-R-S-A。\n",
"\n",
"**核心更新公式:**\n",
"$$Q(S_t, A_t) \\leftarrow Q(S_t, A_t) + \\alpha [R_{t+1} + \\gamma Q(S_{t+1}, A_{t+1}) - Q(S_t, A_t)]$$\n",
"注意看 TD 目标:它是用下一个状态**实际走的那一步**的 $Q(S_{t+1}, A_{t+1})$ 来更新当前的。"
]
},
{
"cell_type": "markdown",
"id": "4e6e3d04",
"metadata": {},
"source": [
"## 4. Q-learning:开天辟地的“聪明人”\n",
"Q-learning 也是基于时间差分的,但它极其特殊,它是一个**异策略(Off-policy**算法。\n",
"所谓“异策略”,就是 **“看别人踩坑,自己学经验”**。它用来收集数据的策略(比如带有随机探索的 $\\epsilon$-greedy),和它真正在心里学到的“终极通关秘籍”(绝对贪心策略)是 **分离的**。\n",
"\n",
"**核心更新公式:**\n",
"$$Q(S_t, A_t) \\leftarrow Q(S_t, A_t) + \\alpha [R_{t+1} + \\gamma \\max_a Q(S_{t+1}, a) - Q(S_t, A_t)]$$\n",
"注意看 TD 目标:不管智能体在 $S_{t+1}$ 实际采取了什么动作,它在更新心里那本账的时候,**永远假设下一步会采取最优的动作(即 $\\max_a Q$** 。\n",
"它直接去求解了第三章的“贝尔曼最优方程(Bellman Optimality Equation)”。"
]
},
{
"cell_type": "code",
"execution_count": 3,
"id": "c632f093",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"--- Sarsa 的更新方式 (老实人) ---\n",
"Sarsa 的 TD 目标: 10 + 0.9 * 3.0 (状态1动作0的价值) = 12.7\n",
"\n",
"--- Q-learning 的更新方式 (聪明人) ---\n",
"Q-learning 的 TD 目标: 10 + 0.9 * 5.0 (状态1里动作1的最大价值) = 14.5\n",
"\n",
"核心结论:\n",
"Sarsa 评估的是当前的探索策略,如果当前策略经常犯错跳崖,Sarsa 就会学得非常保守 (宁愿绕远路也不靠近悬崖)。\n",
"Q-learning 默认未来一定会做最优选择,所以它学到的一定是理论上最短的通关路线,哪怕当前还在跌跌撞撞。\n"
]
}
],
"source": [
"import numpy as np\n",
"\n",
"# 假设我们有一个 Q 表 (状态数目为2,动作数目为2)\n",
"Q = np.array([\n",
" [1.0, 2.0], # 状态 0 的动作价值: 动作0价值=1, 动作1价值=2\n",
" [3.0, 5.0] # 状态 1 的动作价值: 动作0价值=3, 动作1价值=5\n",
"])\n",
"\n",
"# 假设智能体经历了这样一步:\n",
"# 当前在 状态0 (S_t = 0),采取了 动作0 (A_t = 0)\n",
"# 得到了奖励 10 (R_{t+1} = 10)\n",
"# 进入了 状态1 (S_{t+1} = 1)\n",
"S_t, A_t, R, S_next = 0, 0, 10, 1\n",
"\n",
"# 因为有探索率 epsilon 的存在,智能体在 状态1 脑子一抽,\n",
"# 没有选价值最高的动作1(价值5),而是“实际”采取了动作0 (A_{t+1} = 0,价值为3)\n",
"A_next = 0 \n",
"\n",
"alpha = 0.1\n",
"gamma = 0.9\n",
"\n",
"print(\"--- Sarsa 的更新方式 (老实人) ---\")\n",
"# Sarsa 说:“我不管别人怎么选,反正我下一步实际手贱选了动作 0,我就得为我实际的行动买单!”\n",
"# Sarsa 使用的是 Q(S_{t+1}, A_{t+1})\n",
"sarsa_target = R + gamma * Q[S_next, A_next] \n",
"print(f\"Sarsa 的 TD 目标: {R} + {gamma} * {Q[S_next, A_next]} (状态1动作0的价值) = {sarsa_target}\")\n",
"\n",
"\n",
"print(\"\\n--- Q-learning 的更新方式 (聪明人) ---\")\n",
"# Q-learning 说:“虽然我这一步瞎选了动作 0,但我心里门儿清,状态 1 里的最优解其实是动作 1!我是要当海贼王的男人,我的认知必须基于最优选择!”\n",
"# Q-learning 使用的是 max_a Q(S_{t+1}, a) \n",
"max_q_next = np.max(Q[S_next]) \n",
"q_learning_target = R + gamma * max_q_next\n",
"print(f\"Q-learning 的 TD 目标: {R} + {gamma} * {max_q_next} (状态1里动作1的最大价值) = {q_learning_target}\")\n",
"\n",
"print(\"\\n核心结论:\")\n",
"print(\"Sarsa 评估的是当前的探索策略,如果当前策略经常犯错跳崖,Sarsa 就会学得非常保守 (宁愿绕远路也不靠近悬崖)。\")\n",
"print(\"Q-learning 默认未来一定会做最优选择,所以它学到的一定是理论上最短的通关路线,哪怕当前还在跌跌撞撞。\")"
]
}
],
"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
}