PPO实现异常:仅单epoch无裁剪损失可训练,多epoch带裁剪损失失效
PPO在CartPole-v2中多epoch训练失效的问题修复
核心问题分析
你的代码在A2C模式(单epoch、无裁剪)下能正常训练,但启用PPO裁剪损失+多epoch后失效,主要是因为优势函数计算逻辑错误、Critic更新方式错误,以及部分超参数和训练流程的问题。
具体修复点
1. 修正优势函数的归一化逻辑
你当前先对returns做归一化,再计算优势函数,这会破坏回报与价值的相对关系,导致优势信号失真。正确的做法是先计算原始优势,再对优势函数做归一化:
# 原错误代码 returns = (returns - returns.mean()) / (returns.std() + eps) advantages = returns - state_values.squeeze() # 修正后代码 returns = torch.tensor(returns, device=device) # 先计算原始优势 advantages = returns - state_values.squeeze() # 对优势函数做归一化 advantages = (advantages - advantages.mean()) / (advantages.std() + eps)
2. 修正Critic的更新逻辑
你当前仅在最后一个epoch用收集轨迹时的旧价值计算Critic损失,这导致Critic无法学习到新的价值信息。正确的做法是在每个epoch用当前Critic网络重新计算状态价值:
# 原错误代码 for epoch in range(num_epochs): # Actor训练逻辑... if epoch == num_epochs - 1: critic_loss = F.smooth_l1_loss(state_values.squeeze(), returns) critic_optimizer.zero_grad() critic_loss.backward(retain_graph=False) critic_optimizer.step() # 修正后代码 for epoch in range(num_epochs): # Actor训练逻辑 new_probs = actor(states).gather(1, actions.unsqueeze(-1)).squeeze() ratios = new_probs / old_probs surr1 = ratios * advantages surr2 = torch.clamp(ratios, 1 - epsilon, 1 + epsilon) * advantages actor_loss = -torch.min(surr1, surr2).mean() # 启用裁剪损失 actor_optimizer.zero_grad() actor_loss.backward() actor_optimizer.step() # Critic训练:用当前网络计算状态价值 current_state_values = critic(states).squeeze() # 若Actor与Critic共享网络,需同时输出概率与价值 critic_loss = F.smooth_l1_loss(current_state_values, returns) critic_optimizer.zero_grad() critic_loss.backward() critic_optimizer.step()
注:如果你的Actor网络同时输出动作概率和状态价值(共享主干),需调整前向传播逻辑,避免重复计算:
new_probs, current_state_values = actor(states) new_probs = new_probs.gather(1, actions.unsqueeze(-1)).squeeze()
3. 调整超参数
- 裁剪系数
epsilon从0.3降至0.2:CartPole是简单环境,过大的裁剪范围会限制策略更新幅度,导致学习停滞。 - 设
num_epochs为3-5:简单环境无需过多迭代,过多epoch易引发过拟合。
4. 移除不必要的retain_graph=True
多epoch训练中,每次迭代的计算图相互独立,保留计算图会占用额外内存且可能引发梯度计算错误:
# 原错误代码 actor_loss.backward(retain_graph=True) # 修正后代码 actor_loss.backward()
5. 确认旧策略概率的固定性
你从saved_actions中提取old_probs的逻辑是正确的——PPO要求旧策略概率在多epoch训练中保持固定,不能随策略更新而改变,此部分无需修改。
额外建议
- 尝试用GAE(广义优势估计)替代蒙特卡洛回报:GAE能生成更稳定的优势信号,提升多epoch场景下的训练效果。
- 添加梯度裁剪:对Actor和Critic的梯度进行裁剪(如
torch.nn.utils.clip_grad_norm_(actor.parameters(), max_norm=0.5)),避免梯度爆炸。
内容的提问来源于stack exchange,提问作者yanis-falaki
相关产品推荐
相关产品推荐

