312 lines
15 KiB
Plaintext
312 lines
15 KiB
Plaintext
{
|
||
"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": 3,
|
||
"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 决定 向左,进入了状态 3。拿到真实奖励 0.0。\n",
|
||
" -> 我猜状态 3 的价值是 0.000。\n",
|
||
" -> 所以我的 TD 目标是 0.0 + 1.0 * 0.000 = 0.000。\n",
|
||
" -> 我把状态 4 的价值从 0.000 更新为了 0.000。\n",
|
||
"\n",
|
||
"[第 3 步] 在状态 3 决定 向左,进入了状态 2。拿到真实奖励 0.0。\n",
|
||
" -> 我猜状态 2 的价值是 0.000。\n",
|
||
" -> 所以我的 TD 目标是 0.0 + 1.0 * 0.000 = 0.000。\n",
|
||
" -> 我把状态 3 的价值从 0.000 更新为了 0.000。\n",
|
||
"\n",
|
||
"[第 4 步] 在状态 2 决定 向右,进入了状态 3。拿到真实奖励 0.0。\n",
|
||
" -> 我猜状态 3 的价值是 0.000。\n",
|
||
" -> 所以我的 TD 目标是 0.0 + 1.0 * 0.000 = 0.000。\n",
|
||
" -> 我把状态 2 的价值从 0.000 更新为了 0.000。\n",
|
||
"\n",
|
||
"[第 5 步] 在状态 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",
|
||
"[第 6 步] 在状态 4 决定 向左,进入了状态 3。拿到真实奖励 0.0。\n",
|
||
" -> 我猜状态 3 的价值是 0.000。\n",
|
||
" -> 所以我的 TD 目标是 0.0 + 1.0 * 0.000 = 0.000。\n",
|
||
" -> 我把状态 4 的价值从 0.000 更新为了 0.000。\n",
|
||
"\n",
|
||
"[第 7 步] 在状态 3 决定 向左,进入了状态 2。拿到真实奖励 0.0。\n",
|
||
" -> 我猜状态 2 的价值是 0.000。\n",
|
||
" -> 所以我的 TD 目标是 0.0 + 1.0 * 0.000 = 0.000。\n",
|
||
" -> 我把状态 3 的价值从 0.000 更新为了 0.000。\n",
|
||
"\n",
|
||
"[第 8 步] 在状态 2 决定 向左,进入了状态 1。拿到真实奖励 0.0。\n",
|
||
" -> 我猜状态 1 的价值是 0.000。\n",
|
||
" -> 所以我的 TD 目标是 0.0 + 1.0 * 0.000 = 0.000。\n",
|
||
" -> 我把状态 2 的价值从 0.000 更新为了 0.000。\n",
|
||
"\n",
|
||
"[第 9 步] 在状态 1 决定 向左,进入了状态 0。拿到真实奖励 0.0。\n",
|
||
" -> 我猜状态 0 的价值是 0.000。\n",
|
||
" -> 所以我的 TD 目标是 0.0 + 1.0 * 0.000 = 0.000。\n",
|
||
" -> 我把状态 1 的价值从 0.000 更新为了 0.000。\n",
|
||
"\n",
|
||
"游戏结束!最终到达终点 0。\n",
|
||
"\n",
|
||
"跑完这 1 局后的最新状态 V 表 (状态1到5): [0. 0. 0. 0. 0.]\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": 4,
|
||
"id": "6057ac63",
|
||
"metadata": {},
|
||
"outputs": [
|
||
{
|
||
"name": "stdout",
|
||
"output_type": "stream",
|
||
"text": [
|
||
"=== TD(0) 经过 500 局学习后的最终估值 ===\n",
|
||
"状态 1-5 学习到的 V 值: [0.153 0.291 0.513 0.666 0.919]\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": 5,
|
||
"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
|
||
}
|