强化学习实战:用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. 实际应用与优化建议

通过可视化分析,我们可以得出一些实用结论:

  1. 策略效率评估

    • 激进策略在无障碍路径上表现高效,但容易陷入障碍物附近的局部区域
    • 保守策略访问状态更均匀,但到达目标需要更长时间
  2. 超参数调优

    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()
    
  3. 混合策略设计

    • 初期可采用保守策略探索环境
    • 后期切换为激进策略提高效率
    • 关键是在代码中实现策略切换条件:
      def adaptive_policy(state, episode_idx):
          if episode_idx < 500:  # 前500回合用保守策略
              return conservative_policy(state)
          else:  # 之后用激进策略
              return aggressive_policy(state)
      

完整代码已上传至GitHub仓库,包含更多可视化功能和交互式演示。实际项目中,这种分析方法可以帮助我们快速识别策略缺陷,特别是在机器人路径规划、游戏AI等场景中。

更多推荐