Actor-Critic框架实战:解决强化学习中的高方差与偏差平衡问题

1次阅读
没有评论

共计 2038 个字符,预计需要花费 6 分钟才能阅读完成。

image.webp

从 CartPole 看 REINFORCE 的效率瓶颈

最近用 REINFORCE 算法训练 CartPole-v1 时遇到典型问题:虽然最终能学会,但需要超过 2000 个 episode 才能稳定。通过记录每个 episode 的 reward,发现曲线像过山车一样剧烈波动。这其实是蒙特卡洛梯度估计的高方差特性导致的——用整条轨迹的累计回报作为策略评估,单个 episode 的偶然性会直接影响参数更新方向。

Actor-Critic 框架实战:解决强化学习中的高方差与偏差平衡问题

价值评估的两种路径

  1. 蒙特卡洛 (MC) 方法:等 episode 结束后用实际 return 作为价值估计,无偏但方差大。在 CartPole 中表现为某些 episode 因运气好获得高回报,导致策略过度偏向偶然成功的动作

  2. 时序差分 (TD) 方法 :用当前奖励加下一状态估计值作为目标,偏差较大但方差低。实验发现用 TD(0) 训练时,前期进步快但后期容易陷入局部最优

  3. Actor-Critic 的折中方案:Critic 网络用 TD 方法学习价值函数,为 Actor 提供低方差梯度信号;同时通过多步回报或 GAE 保持一定无偏性。在 CartPole 中测试显示,A2C 算法只需约 500episode 就能稳定

核心实现细节

Actor 网络设计

class Actor(nn.Module):
    def __init__(self, state_dim, action_dim):
        super().__init__()
        self.fc1 = nn.Linear(state_dim, 64)
        self.fc2 = nn.Linear(64, action_dim)  # 输出各动作 logit

    def forward(self, x):
        x = F.relu(self.fc1(x))
        return F.softmax(self.fc2(x), dim=-1)  # Softmax 转换为概率

Critic 与 GAE 实现

def compute_gae(rewards, values, gamma=0.99, lam=0.95):
    """
    rewards: 轨迹中的即时奖励序列
    values: Critic 对每个状态的价值估计
    lam: GAE 的超参数,控制偏差 - 方差权衡
    """
    deltas = rewards[:-1] + gamma * values[1:] - values[:-1]
    gae = 0
    returns = []
    for delta in reversed(deltas):
        gae = delta + gamma * lam * gae
        returns.insert(0, gae + values[:-1])
    return torch.stack(returns)

同步更新策略

  1. 采集完一个 batch 的数据后,先计算所有状态的价值估计

  2. 用 GAE 计算每个状态的 Advantage 时,注意对 batch 做标准化处理

  3. Actor 的梯度计算:

    # probs 是 Actor 输出的动作概率
    policy_loss = -(log_probs * advantages.detach()).mean()

  4. Critic 的梯度计算:

    value_loss = F.mse_loss(returns, predicted_values)

性能对比实验

使用相同超参数(lr=3e-4, γ=0.99)在 CartPole 上的对比:

  • REINFORCE:
  • 收敛所需 episode:2100±300
  • 最后 100episode 平均 reward:195±15

  • A2C:

  • 收敛所需 episode:480±50
  • 最后 100episode 平均 reward:198±5

调整 γ 值的发现:
– γ=0.9 时训练更快但最终性能下降约 10%
– γ=0.999 导致早期训练不稳定

生产环境注意事项

  1. 共享网络层处理
  2. 当 Actor 和 Critic 共享底层特征提取层时,建议对 policy_loss 和 value_loss 设置不同权重
  3. 实验发现 0.8:1.2 的比例在多数任务表现良好

  4. 异步训练技巧

  5. 使用 torch 的 DistributedDataParallel 时,注意 sync_parameters 的调用频率
  6. 推荐每 10 个 batch 同步一次参数,而非每个 step

完整训练模板

包含以下关键组件:
1. 带 tensorboard 日志的 Trainer 类
2. 支持多环境并行采样的 RolloutWorker
3. 自动调整学习率的 Scheduler
4. 关键超参数注释示例:

config = {
    'gamma': 0.99,     # 折扣因子
    'gae_lambda': 0.95, # GAE 参数
    'entropy_coef': 0.01, # 熵正则化系数
    'max_grad_norm': 0.5  # 梯度裁剪阈值
}

延伸思考

在实现 PPO 的 clip 机制时,可以修改 policy_loss 的计算:

ratio = torch.exp(new_log_probs - old_log_probs)
surr1 = ratio * advantages
surr2 = torch.clamp(ratio, 1-eps, 1+eps) * advantages
policy_loss = -torch.min(surr1, surr2).mean()

这能防止单次更新时策略变化过大,但需要配合经验回放使用。

正文完
 0
评论(没有评论)