8.7 KiB
8.7 KiB
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 [ ]: