强化学习实战:用Python代码可视化不同策略下的状态访问分布(附GitHub源码)
·
强化学习实战:用Python代码可视化不同策略下的状态访问分布
在强化学习领域,理解智能体如何与环境交互是核心挑战之一。当我们设计不同策略时,智能体访问状态的方式会显著不同,这直接影响学习效果和最终性能。本文将带你用Python实现一个网格世界环境,通过可视化手段直观展示两种典型策略(激进型与保守型)导致的状态访问分布差异。
1. 环境搭建与基础概念
我们先构建一个5x5的网格世界环境,这是强化学习教学中最常用的实验场景之一。这个环境中,智能体从左上角出发,目标是到达右下角的终止状态。某些格子被设置为障碍物,智能体无法通过。
import numpy as np
import matplotlib.pyplot as plt
class GridWorld:
def __init__(self, size=5):
self.size = size
self.obstacles = [(1,1), (2,3), (3,1)] # 障碍物位置
self.goal = (size-1, size-1) # 终止状态
self.state = (0, 0) # 初始状态
def reset(self):
self.state = (0, 0)
return self.state
def step(self, action):
x, y = self.state
if action == 0: # 上
x = max(x-1, 0)
elif action == 1: # 右
y = min(y+1, self.size-1)
elif action == 2: # 下
x = min(x+1, self.size-1)
elif action == 3: # 左
y = max(y-1, 0)
# 检查是否碰到障碍物
if (x, y) not in self.obstacles:
self.state = (x, y)
done = (self.state == self.goal)
reward = 10 if done else -0.1 # 稀疏奖励设置
return self.state, reward, done
状态访问分布 (State Visitation Distribution)描述了智能体在长期交互中访问各个状态的概率。它与策略直接相关,计算公式为:
vπ(s) = (1-γ)∑γᵗPₜπ(s)
其中γ是折扣因子,Pₜπ(s)表示在策略π下时刻t处于状态s的概率。
2. 策略设计与实现
我们将实现两种对比鲜明的策略:激进型策略倾向于朝着目标直线前进,而保守型策略则采取更谨慎的移动方式。
def aggressive_policy(state):
"""激进策略:优先向右和向下移动"""
x, y = state
if y < 4 and (x, y+1) not in env.obstacles: # 优先向右
return 1
elif x < 4 and (x+1, y) not in env.obstacles: # 其次向下
return 2
elif y > 0 and (x, y-1) not in env.obstacles: # 然后向左
return 3
else:
return 0 # 最后向上
def conservative_policy(state):
"""保守策略:避免靠近障碍物和边界"""
x, y = state
possible_actions = []
if x > 0 and (x-1, y) not in env.obstacles: # 上
possible_actions.append(0)
if y < 4 and (x, y+1) not in env.obstacles and x != 3: # 右(避开特定列)
possible_actions.append(1)
if x < 4 and (x+1, y) not in env.obstacles and y != 3: # 下(避开特定行)
possible_actions.append(2)
if y > 0 and (x, y-1) not in env.obstacles: # 左
possible_actions.append(3)
return np.random.choice(possible_actions) if possible_actions else 0
3. 状态访问统计与可视化
我们通过模拟多个回合的运行轨迹,统计每个状态被访问的次数,并将其转化为概率分布。
def run_episodes(env, policy, num_episodes=1000, max_steps=100):
visitation_counts = np.zeros((env.size, env.size))
for _ in range(num_episodes):
state = env.reset()
for _ in range(max_steps):
action = policy(state)
state, _, done = env.step(action)
visitation_counts[state] += 1
if done:
break
# 归一化为概率分布
visitation_dist = visitation_counts / np.sum(visitation_counts)
return visitation_dist
def plot_visitation_distribution(dist, title):
plt.figure(figsize=(8, 6))
plt.imshow(dist, cmap='YlOrRd')
plt.colorbar(label='访问概率')
plt.title(title)
plt.xticks([])
plt.yticks([])
# 标注障碍物和目标位置
for obs in env.obstacles:
plt.text(obs[1], obs[0], 'X', ha='center', va='center', fontsize=14)
plt.text(env.goal[1], env.goal[0], 'G', ha='center', va='center', fontsize=14)
plt.show()
执行模拟并可视化结果:
env = GridWorld()
# 运行激进策略
agg_dist = run_episodes(env, aggressive_policy)
plot_visitation_distribution(agg_dist, "激进策略状态访问分布")
# 运行保守策略
cons_dist = run_episodes(env, conservative_policy)
plot_visitation_distribution(cons_dist, "保守策略状态访问分布")
4. 占用度量的计算与分析
占用度量(Occupancy Measure)进一步考虑了动作选择,表示状态-动作对被访问的概率。它与状态访问分布的关系为:
ρπ(s,a) = vπ(s)π(a|s)
我们可以扩展之前的统计代码来计算占用度量:
def run_episodes_with_actions(env, policy, num_episodes=1000, max_steps=100):
occupancy = np.zeros((env.size, env.size, 4)) # 4个动作
for _ in range(num_episodes):
state = env.reset()
for _ in range(max_steps):
action = policy(state)
next_state, _, done = env.step(action)
occupancy[state][action] += 1
state = next_state
if done:
break
# 归一化
occupancy /= np.sum(occupancy)
return occupancy
def plot_action_distribution(occupancy, state):
actions = ['上', '右', '下', '左']
dist = occupancy[state] / np.sum(occupancy[state])
plt.figure(figsize=(6, 4))
plt.bar(actions, dist)
plt.title(f"状态{state}下的动作选择分布")
plt.ylabel("概率")
plt.show()
比较两种策略在关键状态下的动作选择:
# 激进策略的占用度量
agg_occupancy = run_episodes_with_actions(env, aggressive_policy)
plot_action_distribution(agg_occupancy, (0, 0)) # 初始状态
plot_action_distribution(agg_occupancy, (2, 2)) # 中心状态
# 保守策略的占用度量
cons_occupancy = run_episodes_with_actions(env, conservative_policy)
plot_action_distribution(cons_occupancy, (0, 0))
plot_action_distribution(cons_occupancy, (2, 2))
5. 实际应用与优化建议
通过可视化分析,我们可以得出一些实用结论:
-
策略效率评估 :
- 激进策略在无障碍路径上表现高效,但容易陷入障碍物附近的局部区域
- 保守策略访问状态更均匀,但到达目标需要更长时间
-
超参数调优 :
def sensitivity_analysis(): gammas = [0.5, 0.7, 0.9, 0.99] fig, axes = plt.subplots(2, 2, figsize=(10, 8)) for gamma, ax in zip(gammas, axes.flatten()): dist = compute_discounted_visitation(env, aggressive_policy, gamma) ax.imshow(dist) ax.set_title(f"γ={gamma}") plt.tight_layout() plt.show() -
混合策略设计 :
- 初期可采用保守策略探索环境
- 后期切换为激进策略提高效率
- 关键是在代码中实现策略切换条件:
def adaptive_policy(state, episode_idx): if episode_idx < 500: # 前500回合用保守策略 return conservative_policy(state) else: # 之后用激进策略 return aggressive_policy(state)
完整代码已上传至GitHub仓库,包含更多可视化功能和交互式演示。实际项目中,这种分析方法可以帮助我们快速识别策略缺陷,特别是在机器人路径规划、游戏AI等场景中。
更多推荐


所有评论(0)