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

价值评估的两种路径
-
蒙特卡洛 (MC) 方法:等 episode 结束后用实际 return 作为价值估计,无偏但方差大。在 CartPole 中表现为某些 episode 因运气好获得高回报,导致策略过度偏向偶然成功的动作
-
时序差分 (TD) 方法 :用当前奖励加下一状态估计值作为目标,偏差较大但方差低。实验发现用 TD(0) 训练时,前期进步快但后期容易陷入局部最优
-
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)
同步更新策略
-
采集完一个 batch 的数据后,先计算所有状态的价值估计
-
用 GAE 计算每个状态的 Advantage 时,注意对 batch 做标准化处理
-
Actor 的梯度计算:
# probs 是 Actor 输出的动作概率 policy_loss = -(log_probs * advantages.detach()).mean() -
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 导致早期训练不稳定
生产环境注意事项
- 共享网络层处理:
- 当 Actor 和 Critic 共享底层特征提取层时,建议对 policy_loss 和 value_loss 设置不同权重
-
实验发现 0.8:1.2 的比例在多数任务表现良好
-
异步训练技巧:
- 使用 torch 的 DistributedDataParallel 时,注意 sync_parameters 的调用频率
- 推荐每 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()
这能防止单次更新时策略变化过大,但需要配合经验回放使用。
