231 lines
8.7 KiB
Plaintext
231 lines
8.7 KiB
Plaintext
{
|
|
"cells": [
|
|
{
|
|
"cell_type": "markdown",
|
|
"id": "13468a42",
|
|
"metadata": {},
|
|
"source": [
|
|
"# 第 4 章:价值迭代与策略迭代 (Value Iteration and Policy Iteration)\n",
|
|
"\n",
|
|
"本章介绍的算法属于**动态规划 (Dynamic Programming)**,因为它们需要提前知道环境的完整模型(状态转移概率和奖励)。\n",
|
|
"\n",
|
|
"## 1. 价值迭代 (Value Iteration)\n",
|
|
"价值迭代是直接求解第三章“贝尔曼最优方程 (BOE)”的迭代算法。\n",
|
|
"它在每次迭代中包含两步:\n",
|
|
"* **策略更新 (Policy Update)**:算出 Q 值,找到当前最贪心(收益最大)的动作。\n",
|
|
"* **价值更新 (Value Update)**:把状态的价值直接更新为最大的 Q 值:$v_{k+1}(s) = \\max_a q_k(s,a)$。\n",
|
|
"* **注意**:价值迭代中间产生的 $v_k$ 仅仅是计算过程中的临时数值,它们并不代表任何特定策略的真实状态价值!\n",
|
|
"\n",
|
|
"## 2. 策略迭代 (Policy Iteration)\n",
|
|
"策略迭代的过程更像是一个稳扎稳打的“评估-改进”循环。\n",
|
|
"* **策略评估 (Policy Evaluation)**:对当前策略,死磕到底,算出它真实的、收敛的状态价值 $v_{\\pi_k}$(这通常需要内部再套一个循环)。\n",
|
|
"* **策略改进 (Policy Improvement)**:根据算出来的真实价值,选择 Q 值最大的动作,更新策略。这保证了新策略一定比老策略更好(或者一样好)。"
|
|
]
|
|
},
|
|
{
|
|
"cell_type": "code",
|
|
"execution_count": 1,
|
|
"id": "9ae8b2ba",
|
|
"metadata": {},
|
|
"outputs": [],
|
|
"source": [
|
|
"import numpy as np\n",
|
|
"\n",
|
|
"# --- 1. 定义 2x2 环境 ---\n",
|
|
"states = [\"s1\", \"s2\", \"s3\", \"s4\"]\n",
|
|
"actions = [\"up\", \"right\", \"down\", \"left\", \"stay\"]\n",
|
|
"gamma = 0.9\n",
|
|
"\n",
|
|
"# 状态转移规则 (依据 Table 4.1 逆向推导的简单网格规律)\n",
|
|
"def get_transition(state, action):\n",
|
|
" # s4 是目标(吸收态),到了就停在原地\n",
|
|
" if state == \"s4\": return \"s4\"\n",
|
|
" \n",
|
|
" if state == \"s1\":\n",
|
|
" if action == \"right\": return \"s2\"\n",
|
|
" if action == \"down\": return \"s3\"\n",
|
|
" if action == \"stay\": return \"s1\"\n",
|
|
" return \"s1\" # 撞墙反弹\n",
|
|
" elif state == \"s2\":\n",
|
|
" if action == \"left\": return \"s1\"\n",
|
|
" if action == \"down\": return \"s4\"\n",
|
|
" if action == \"stay\": return \"s2\"\n",
|
|
" return \"s2\" # 撞墙反弹\n",
|
|
" elif state == \"s3\":\n",
|
|
" if action == \"up\": return \"s1\"\n",
|
|
" if action == \"right\": return \"s4\"\n",
|
|
" if action == \"stay\": return \"s3\"\n",
|
|
" return \"s3\" # 撞墙反弹\n",
|
|
" return state\n",
|
|
"\n",
|
|
"# 奖励规则 (依据 Table 4.1 提取)\n",
|
|
"def get_reward(state, action, next_state):\n",
|
|
" if state == \"s4\": return 1 # 目标奖励\n",
|
|
" \n",
|
|
" # 判断是否撞墙 (尝试移动但留在原地)\n",
|
|
" if state == next_state and action != \"stay\":\n",
|
|
" return -1\n",
|
|
" \n",
|
|
" # 正常移动的奖励\n",
|
|
" if next_state == \"s2\": return -1 # 禁区\n",
|
|
" if next_state == \"s4\": return 1 # 目标\n",
|
|
" return 0 # 其他移动 (书中这题为 0)\n",
|
|
"\n",
|
|
"# 辅助函数:计算 Q 值 (这里假设转移是 100% 确定性的)\n",
|
|
"def compute_q_value(state, action, v_values):\n",
|
|
" next_s = get_transition(state, action)\n",
|
|
" reward = get_reward(state, action, next_s)\n",
|
|
" return reward + gamma * v_values[states.index(next_s)]"
|
|
]
|
|
},
|
|
{
|
|
"cell_type": "code",
|
|
"execution_count": 7,
|
|
"id": "32dc5456",
|
|
"metadata": {},
|
|
"outputs": [
|
|
{
|
|
"name": "stdout",
|
|
"output_type": "stream",
|
|
"text": [
|
|
"=== 开始价值迭代 (Value Iteration) ===\n",
|
|
"迭代 1: V值 = [0. 1. 1. 1.], 策略 = ['down', 'down', 'right', 'up']\n",
|
|
"迭代 2: V值 = [0.9 1.9 1.9 1.9], 策略 = ['down', 'down', 'right', 'up']\n",
|
|
"迭代 3: V值 = [1.71 2.71 2.71 2.71], 策略 = ['down', 'down', 'right', 'up']\n",
|
|
"迭代 4: V值 = [2.44 3.44 3.44 3.44], 策略 = ['down', 'down', 'right', 'up']\n",
|
|
"迭代 5: V值 = [3.1 4.1 4.1 4.1], 策略 = ['down', 'down', 'right', 'up']\n"
|
|
]
|
|
}
|
|
],
|
|
"source": [
|
|
"print(\"=== 开始价值迭代 (Value Iteration) ===\")\n",
|
|
"\n",
|
|
"# 初始化 V 值为 0\n",
|
|
"V_vi = np.zeros(len(states))\n",
|
|
"policy_vi = [\"stay\"] * len(states) \n",
|
|
"iterations = 5\n",
|
|
"\n",
|
|
"for k in range(iterations):\n",
|
|
" new_V = np.zeros(len(states))\n",
|
|
" \n",
|
|
" # 遍历所有状态\n",
|
|
" for i, s in enumerate(states):\n",
|
|
" q_values = []\n",
|
|
" # 遍历所有动作,计算 Q 值\n",
|
|
" for a in actions:\n",
|
|
" q = compute_q_value(s, a, V_vi)\n",
|
|
" q_values.append(q)\n",
|
|
" \n",
|
|
" # 核心:价值更新 (直接取最大的 Q 值)\n",
|
|
" best_q = max(q_values)\n",
|
|
" new_V[i] = best_q\n",
|
|
" \n",
|
|
" # 核心:策略更新 (记录最大 Q 值对应的动作)\n",
|
|
" best_action_idx = np.argmax(q_values)\n",
|
|
" policy_vi[i] = actions[best_action_idx]\n",
|
|
" \n",
|
|
" print(f\"迭代 {k+1}: V值 = {np.round(new_V, 2)}, 策略 = {policy_vi}\")\n",
|
|
" \n",
|
|
" # 如果价值不再变化,说明收敛了\n",
|
|
" if np.max(np.abs(new_V - V_vi)) < 1e-5:\n",
|
|
" print(f\"-> 价值迭代在第 {k+1} 步提前收敛!\")\n",
|
|
" V_vi = new_V\n",
|
|
" break\n",
|
|
" \n",
|
|
" V_vi = new_V"
|
|
]
|
|
},
|
|
{
|
|
"cell_type": "code",
|
|
"execution_count": 9,
|
|
"id": "9e6c5d0d",
|
|
"metadata": {},
|
|
"outputs": [
|
|
{
|
|
"name": "stdout",
|
|
"output_type": "stream",
|
|
"text": [
|
|
"=== 开始策略迭代 (Policy Iteration) ===\n",
|
|
"\n",
|
|
"第 1 轮主迭代,当前策略: ['stay', 'stay', 'stay', 'stay']\n",
|
|
" 评估完成,当前策略的真实 V值 = [ 0. -10. 0. 10.]\n",
|
|
"\n",
|
|
"第 2 轮主迭代,当前策略: ['down', 'down', 'right', 'up']\n",
|
|
" 评估完成,当前策略的真实 V值 = [ 9. 10. 10. 10.]\n",
|
|
"-> 策略不再改变,策略迭代收敛!最优策略已找到。\n"
|
|
]
|
|
}
|
|
],
|
|
"source": [
|
|
"print(\"=== 开始策略迭代 (Policy Iteration) ===\")\n",
|
|
"\n",
|
|
"V_pi = np.zeros(len(states))\n",
|
|
"# 初始给一个极差的策略:全部原地不动 (类似书中图 4.3 的烂策略)\n",
|
|
"current_policy = [\"stay\", \"stay\", \"stay\", \"stay\"] \n",
|
|
"\n",
|
|
"for k in range(5):\n",
|
|
" print(f\"\\n第 {k+1} 轮主迭代,当前策略: {current_policy}\")\n",
|
|
" \n",
|
|
" # --- 步骤 1: 策略评估 (Policy Evaluation) ---\n",
|
|
" # 死磕到底,一直循环直到 V 值收敛,算出当前策略的真实价值\n",
|
|
" while True:\n",
|
|
" new_V = np.zeros(len(states))\n",
|
|
" for i, s in enumerate(states):\n",
|
|
" action = current_policy[i] # 只看当前策略指定的动作\n",
|
|
" new_V[i] = compute_q_value(s, action, V_pi)\n",
|
|
" \n",
|
|
" if np.max(np.abs(new_V - V_pi)) < 1e-5:\n",
|
|
" break\n",
|
|
" V_pi = new_V\n",
|
|
" print(f\" 评估完成,当前策略的真实 V值 = {np.round(V_pi, 2)}\")\n",
|
|
" \n",
|
|
" # --- 步骤 2: 策略改进 (Policy Improvement) ---\n",
|
|
" policy_stable = True\n",
|
|
" for i, s in enumerate(states):\n",
|
|
" old_action = current_policy[i]\n",
|
|
" \n",
|
|
" # 看看有没有更好的动作\n",
|
|
" q_values = [compute_q_value(s, a, V_pi) for a in actions]\n",
|
|
" best_action = actions[np.argmax(q_values)]\n",
|
|
" \n",
|
|
" current_policy[i] = best_action\n",
|
|
" if old_action != best_action:\n",
|
|
" policy_stable = False # 策略发生了改变\n",
|
|
" \n",
|
|
" if policy_stable:\n",
|
|
" print(\"-> 策略不再改变,策略迭代收敛!最优策略已找到。\")\n",
|
|
" break"
|
|
]
|
|
},
|
|
{
|
|
"cell_type": "code",
|
|
"execution_count": null,
|
|
"id": "5282df09",
|
|
"metadata": {},
|
|
"outputs": [],
|
|
"source": []
|
|
}
|
|
],
|
|
"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
|
|
}
|