Files
RL-Study/Notebooks/C4.ipynb
T
2026-02-27 21:17:25 +08:00

8.7 KiB

第 4 章:价值迭代与策略迭代 (Value Iteration and Policy Iteration)

本章介绍的算法属于动态规划 (Dynamic Programming),因为它们需要提前知道环境的完整模型(状态转移概率和奖励)。

1. 价值迭代 (Value Iteration)

价值迭代是直接求解第三章“贝尔曼最优方程 (BOE)”的迭代算法。 它在每次迭代中包含两步:

  • 策略更新 (Policy Update):算出 Q 值,找到当前最贪心(收益最大)的动作。
  • 价值更新 (Value Update):把状态的价值直接更新为最大的 Q 值:$v_{k+1}(s) = \max_a q_k(s,a)$。
  • 注意:价值迭代中间产生的 v_k 仅仅是计算过程中的临时数值,它们并不代表任何特定策略的真实状态价值!

2. 策略迭代 (Policy Iteration)

策略迭代的过程更像是一个稳扎稳打的“评估-改进”循环。

  • 策略评估 (Policy Evaluation):对当前策略,死磕到底,算出它真实的、收敛的状态价值 $v_{\pi_k}$(这通常需要内部再套一个循环)。
  • 策略改进 (Policy Improvement):根据算出来的真实价值,选择 Q 值最大的动作,更新策略。这保证了新策略一定比老策略更好(或者一样好)。
In [1]:
import numpy as np

# --- 1. 定义 2x2 环境 ---
states = ["s1", "s2", "s3", "s4"]
actions = ["up", "right", "down", "left", "stay"]
gamma = 0.9

# 状态转移规则 (依据 Table 4.1 逆向推导的简单网格规律)
def get_transition(state, action):
    # s4 是目标(吸收态),到了就停在原地
    if state == "s4": return "s4"
    
    if state == "s1":
        if action == "right": return "s2"
        if action == "down": return "s3"
        if action == "stay": return "s1"
        return "s1" # 撞墙反弹
    elif state == "s2":
        if action == "left": return "s1"
        if action == "down": return "s4"
        if action == "stay": return "s2"
        return "s2" # 撞墙反弹
    elif state == "s3":
        if action == "up": return "s1"
        if action == "right": return "s4"
        if action == "stay": return "s3"
        return "s3" # 撞墙反弹
    return state

# 奖励规则 (依据 Table 4.1 提取)
def get_reward(state, action, next_state):
    if state == "s4": return 1 # 目标奖励
    
    # 判断是否撞墙 (尝试移动但留在原地)
    if state == next_state and action != "stay":
        return -1
    
    # 正常移动的奖励
    if next_state == "s2": return -1 # 禁区
    if next_state == "s4": return 1  # 目标
    return 0 # 其他移动 (书中这题为 0)

# 辅助函数:计算 Q 值 (这里假设转移是 100% 确定性的)
def compute_q_value(state, action, v_values):
    next_s = get_transition(state, action)
    reward = get_reward(state, action, next_s)
    return reward + gamma * v_values[states.index(next_s)]
In [7]:
print("=== 开始价值迭代 (Value Iteration) ===")

# 初始化 V 值为 0
V_vi = np.zeros(len(states))
policy_vi = ["stay"] * len(states) 
iterations = 5

for k in range(iterations):
    new_V = np.zeros(len(states))
    
    # 遍历所有状态
    for i, s in enumerate(states):
        q_values = []
        # 遍历所有动作,计算 Q 值
        for a in actions:
            q = compute_q_value(s, a, V_vi)
            q_values.append(q)
            
        # 核心:价值更新 (直接取最大的 Q 值)
        best_q = max(q_values)
        new_V[i] = best_q
        
        # 核心:策略更新 (记录最大 Q 值对应的动作)
        best_action_idx = np.argmax(q_values)
        policy_vi[i] = actions[best_action_idx]
        
    print(f"迭代 {k+1}: V值 = {np.round(new_V, 2)}, 策略 = {policy_vi}")
    
    # 如果价值不再变化,说明收敛了
    if np.max(np.abs(new_V - V_vi)) < 1e-5:
        print(f"-> 价值迭代在第 {k+1} 步提前收敛!")
        V_vi = new_V
        break
        
    V_vi = new_V
=== 开始价值迭代 (Value Iteration) ===
迭代 1: V值 = [0. 1. 1. 1.], 策略 = ['down', 'down', 'right', 'up']
迭代 2: V值 = [0.9 1.9 1.9 1.9], 策略 = ['down', 'down', 'right', 'up']
迭代 3: V值 = [1.71 2.71 2.71 2.71], 策略 = ['down', 'down', 'right', 'up']
迭代 4: V值 = [2.44 3.44 3.44 3.44], 策略 = ['down', 'down', 'right', 'up']
迭代 5: V值 = [3.1 4.1 4.1 4.1], 策略 = ['down', 'down', 'right', 'up']
In [9]:
print("=== 开始策略迭代 (Policy Iteration) ===")

V_pi = np.zeros(len(states))
# 初始给一个极差的策略:全部原地不动 (类似书中图 4.3 的烂策略)
current_policy = ["stay", "stay", "stay", "stay"] 

for k in range(5):
    print(f"\n{k+1} 轮主迭代,当前策略: {current_policy}")
    
    # --- 步骤 1: 策略评估 (Policy Evaluation) ---
    # 死磕到底,一直循环直到 V 值收敛,算出当前策略的真实价值
    while True:
        new_V = np.zeros(len(states))
        for i, s in enumerate(states):
            action = current_policy[i] # 只看当前策略指定的动作
            new_V[i] = compute_q_value(s, action, V_pi)
            
        if np.max(np.abs(new_V - V_pi)) < 1e-5:
            break
        V_pi = new_V
    print(f"  评估完成,当前策略的真实 V值 = {np.round(V_pi, 2)}")
    
    # --- 步骤 2: 策略改进 (Policy Improvement) ---
    policy_stable = True
    for i, s in enumerate(states):
        old_action = current_policy[i]
        
        # 看看有没有更好的动作
        q_values = [compute_q_value(s, a, V_pi) for a in actions]
        best_action = actions[np.argmax(q_values)]
        
        current_policy[i] = best_action
        if old_action != best_action:
            policy_stable = False # 策略发生了改变
            
    if policy_stable:
        print("-> 策略不再改变,策略迭代收敛!最优策略已找到。")
        break
=== 开始策略迭代 (Policy Iteration) ===

第 1 轮主迭代,当前策略: ['stay', 'stay', 'stay', 'stay']
  评估完成,当前策略的真实 V值 = [  0. -10.   0.  10.]

第 2 轮主迭代,当前策略: ['down', 'down', 'right', 'up']
  评估完成,当前策略的真实 V值 = [ 9. 10. 10. 10.]
-> 策略不再改变,策略迭代收敛!最优策略已找到。
In [ ]: