PPO 算法详解与 CartPole-v1 的 PyTorch 实现
本文从策略梯度开始推导 Proximal Policy Optimization(PPO,近端策略优化),并逐步对应本目录中的 train.py 和 infer.py。阅读本文不要求预先掌握强化学习,但默认读者熟悉 Python、PyTorch 和基础概率知识。
目录
- 1. 项目概览
- 2. CartPole-v1 问题
- 3. 强化学习基础
- 4. 从策略梯度到 Actor-Critic
- 5. GAE:如何估计优势函数
- 6. PPO 的核心:限制策略更新幅度
- 7. 完整损失函数
- 8. 当前代码的网络结构
- 9. 训练数据流与代码对应
- 10. terminated 与 truncated
- 11. 推理与可视化
- 12. 安装和使用
- 13. 超参数详解
- 14. 常见问题和调试方法
- 15. 参考资料
1. 项目概览
本实现具有以下特点:
- 不依赖 Stable-Baselines3,PPO 的 rollout、GAE 和优化过程全部由 PyTorch 实现。
- 使用单个 Gymnasium
CartPole-v1环境收集 on-policy 数据。 - 使用共享主干的 Actor-Critic 网络。
- 使用离散动作的
torch.distributions.Categorical策略。 - 使用 GAE、PPO clipped objective、价值损失、熵奖励、优势标准化和梯度裁剪。
- 正确区分 Gymnasium 的
terminated与truncated。 - 训练后保存自描述 checkpoint,推理脚本可直接恢复网络并实时渲染。
整体流程如下:
环境 observation
│
▼
Actor-Critic 网络 ───────► V(s),估计当前状态价值
│
▼
Categorical 策略 π(a|s)
│ sample
▼
执行 action ─► 环境返回 next_observation、reward、terminated、truncated
│
▼
收集一段 rollout ─► GAE/returns ─► 多轮 minibatch PPO 更新
│
└────────────────────────────► 继续采样,最终保存 checkpoint
PPO 是 on-policy 算法:每批数据由更新前的策略采集,完成若干轮优化后便不再复用。旧数据来自不同策略,继续反复训练会使重要性采样比率失真,破坏 PPO 的近端更新假设。
2. CartPole-v1 问题
CartPole 的目标是通过左右推动小车,使铰接在小车上的杆尽可能长时间保持直立。
2.1 状态空间
每个 observation 是长度为 4 的 float32 向量:
| 索引 | 含义 | 直观作用 |
|---|---|---|
| 0 | 小车位置 | 判断小车是否接近轨道边界 |
| 1 | 小车速度 | 判断小车正向哪个方向运动 |
| 2 | 杆的角度 | 判断杆向左还是向右倾斜 |
| 3 | 杆的角速度 | 判断倾斜趋势和纠正力度 |
2.2 动作空间
动作空间是 Discrete(2):
0:向左推动小车。1:向右推动小车。
Actor 因此只需输出两个 logits,经过 Categorical(logits=logits) 后得到两个动作的概率。
2.3 奖励和回合结束
每存活一步获得 +1 奖励,因此回合回报等于杆保持平衡的步数。CartPole-v1 最长为 500 步:
- 杆角度或小车位置越界时,环境产生
terminated=True。 - 达到 500 步时间限制时,外层
TimeLimit产生truncated=True。
最大回报为 500。官方环境定义参见 Gymnasium CartPole 文档。
3. 强化学习基础
3.1 交互序列
智能体与环境交互产生轨迹:
其中:
- \(s_t\):时刻 \(t\) 的状态;在代码中是
observation。 - \(a_t\):策略选择的动作;在代码中是
action。 - \(r_t\):执行动作后的即时奖励;在代码中是
reward。 - \(\pi_\theta(a_t\mid s_t)\):参数为 \(\theta\) 的策略网络给动作 \(a_t\) 的概率。
3.2 折扣回报
折扣回报表示:从当前时刻开始,将未来所有奖励按照“距离当前有多远”逐步降低权重。定义为:
其中:
- \(G_t\):时刻 \(t\) 的折扣回报。
- \(r_t\):执行当前动作后立即得到的奖励。
- \(r_{t+k}\):距离当前 \(k\) 步的奖励。
- \(\gamma\in[0,1]\):折扣因子。
- 距离当前 \(k\) 步的奖励,其权重为 \(\gamma^k\)。
例如未来四步的奖励都是 1,且 \(\gamma=0.9\):
虽然四步奖励的原始总和是 4,但距离当前越远的奖励权重越低。
折扣因子 \(\gamma\) 的影响
当 \(\gamma=0\) 时:
智能体只关心当前一步,完全忽略未来奖励。当 \(\gamma\) 接近 1 时:
智能体更加重视长期收益。本项目默认使用 \(\gamma=0.99\),未来第 \(k\) 步奖励的权重为 \(0.99^k\):
| 距离当前的步数 | 奖励权重 |
|---|---|
| 0 | \(1\) |
| 10 | \(0.99^{10}\approx0.904\) |
| 50 | \(0.99^{50}\approx0.605\) |
| 100 | \(0.99^{100}\approx0.366\) |
| 500 | \(0.99^{500}\approx0.0066\) |
为什么需要折扣
折扣有三个主要作用:
- 鼓励尽早获得奖励:同样大小的奖励,越早获得,其当前价值越高。
- 降低远期不确定性的影响:越远的未来通常越难预测,因此赋予较低权重。
- 让无限奖励序列保持有限:当每步奖励都是 1 且 \(0\le\gamma<1\) 时,折扣回报是收敛的几何级数:
例如 \(\gamma=0.99\) 时,无限奖励序列的折扣回报上限为:
在 CartPole 中的含义
CartPole 每存活一步获得奖励 1。假设从当前时刻开始还能保持平衡 \(n\) 步,则:
能让杆保持更久的动作具有更大的折扣回报,因此智能体仍然有动力学习长期平衡策略。
需要区分文档和代码中的两个“回报”概念:
- 训练日志的
return是一个 episode 内的未折扣奖励总和。CartPole 存活 500 步时,日志回报就是 500。 - PPO 的 TD error、GAE 和 value target 使用
--gamma 0.99对未来价值进行折扣,这是网络实际优化的训练信号。
3.3 状态价值、动作价值和优势
状态价值函数表示从状态 \(s\) 出发,继续使用策略 \(\pi\) 时的期望回报:
动作价值函数还指定了当前动作:
优势函数衡量某动作相对该状态平均水平好多少:
- \(A>0\):这个动作比策略在该状态下的平均动作好,应提高其概率。
- \(A<0\):这个动作较差,应降低其概率。
- \(A\approx0\):动作与平均水平接近,不需要明显改变概率。
PPO 的核心工作可以概括为:在限制新旧策略差异的前提下,提高正优势动作的概率、降低负优势动作的概率。
4. 从策略梯度到 Actor-Critic
4.1 优化目标
策略的目标是最大化期望累计回报:
环境转移通常不可微,不能直接对环境求梯度。策略梯度定理利用 log-derivative trick,将梯度写成:
直观理解:
log_prob表示策略对已执行动作的对数概率。- 回报或优势为正时,梯度提升该动作概率。
- 回报或优势为负时,梯度降低该动作概率。
4.2 为什么使用 baseline
直接用 \(G_t\) 或 \(Q(s_t,a_t)\) 估计策略梯度通常方差很大。减去一个只依赖状态、不依赖动作的 baseline 不改变梯度期望:
令 baseline 为价值函数 \(V(s_t)\),便得到优势函数:
这不会改变理想情况下梯度的期望,但可以显著降低方差。
4.3 Actor-Critic
Actor-Critic 同时学习两个目标:
- Actor:策略 \(\pi_\theta(a\mid s)\),负责选择动作。
- Critic:价值 \(V_\phi(s)\),负责评估状态和构造优势估计。
本实现让二者共享两层 MLP 主干,然后分别连接 actor head 和 critic head。共享主干减少参数量,并让两项任务共同学习状态表示。
5. GAE:如何估计优势函数
精确计算 \(A^\pi\) 需要知道未来的完整分布,实际训练中只能从有限样本估计。Generalized Advantage Estimation(GAE)通过参数 \(\lambda\) 在偏差和方差之间折中。
5.1 一步 TD 误差
先定义一步 Temporal Difference(TD)误差:
它可看作一步优势估计:即时奖励加上下个状态的预测价值,再减去当前状态的预测价值。
- 只使用一步 bootstrap,方差低,但 Critic 不准确时偏差较大。
- 使用完整 Monte Carlo 回报,偏差较低,但方差大,而且必须等待回合结束。
5.2 GAE 公式
GAE 将未来的 TD 误差指数加权:
等价的反向递推形式是:
代码从 rollout 最后一步向前循环,因此只需维护一个标量 gae,不必显式构造所有多步回报。
5.3 \(\lambda\) 的含义
- \(\lambda=0\):接近一步 TD,方差较低、偏差较高。
- \(\lambda\rightarrow1\):接近 Monte Carlo,偏差较低、方差较高。
- 默认
0.95:PPO 中常见的折中值。
5.4 Critic 的训练目标
先用未标准化的优势构造 value target:
在代码中对应:
returns = advantages + value_batch
这里的 value_batch 是收集 rollout 时旧网络的预测。随后 Critic 回归到 \(\hat R_t\)。
6. PPO 的核心:限制策略更新幅度
普通策略梯度如果一次更新过大,可能让新策略突然远离生成数据的旧策略,导致性能崩溃。PPO 使用新旧策略对已采样动作的概率比率衡量变化。
6.1 新旧策略概率比率
定义:
- \(r_t=1\):新旧策略对该动作给出相同概率。
- \(r_t>1\):新策略提高了该动作概率。
- \(r_t<1\):新策略降低了该动作概率。
直接计算概率比可能产生数值问题,因此代码存储旧 log_prob,更新时计算:
log_ratio = new_log_probs - old_log_prob_batch[indices]
ratio = log_ratio.exp()
因为:
6.2 未裁剪的 surrogate objective
重要性采样修正后的目标为:
同一批 rollout 可以因此进行多轮更新。但是如果 \(r_t\) 偏离 1 太远,新策略就会过度利用这批旧数据。
6.3 Clipped objective
PPO-Clip 使用以下目标:
默认 \(\epsilon=0.2\),也就是把有利方向上的比率收益限制在大致 \([0.8,1.2]\) 范围。
当优势为正
\(\hat A_t>0\) 表示动作好,算法希望提高动作概率。但当 \(r_t>1+\epsilon\) 时,继续提高概率不再增加 clipped objective,从而抑制过大的正向更新。
当优势为负
\(\hat A_t<0\) 表示动作差,算法希望降低动作概率。但当 \(r_t<1-\epsilon\) 时,继续降低概率不再带来收益,从而抑制过大的负向更新。
clip 不是直接把网络参数或最终策略概率截断,而是让超出范围后的优化收益进入平台区。它也不是严格的 KL trust region 保证,但实现简单,允许对同一 rollout 做多轮 minibatch 更新。
6.4 代码为何使用 maximum
论文公式最大化 min(...),PyTorch 优化器默认最小化 loss,因此代码先取负号:
policy_loss_unclipped = -advantages[indices] * ratio
policy_loss_clipped = -advantages[indices] * torch.clamp(
ratio, 1.0 - args.clip_coef, 1.0 + args.clip_coef
)
policy_loss = torch.maximum(
policy_loss_unclipped, policy_loss_clipped
).mean()
数学上:
因此该实现与最大化论文中的 clipped objective 等价。
7. 完整损失函数
7.1 策略损失
代码最小化:
7.2 价值损失
Critic 使用均方误差拟合 return target:
代码中的系数 \(1/2\) 只改变梯度尺度,不改变最优点。
7.3 熵奖励
离散策略的熵为:
高熵表示动作概率更均匀、探索更多;低熵表示策略更确定。因为训练过程最小化 loss,所以熵以负号加入:
当前默认:
7.4 优势标准化
每个 rollout 内执行:
标准化让策略梯度的尺度更稳定,降低不同 rollout 回报尺度变化对学习率的影响。注意代码先用原始优势构造 returns,再标准化只用于 policy loss 的优势。
7.5 梯度裁剪
反向传播后执行:
nn.utils.clip_grad_norm_(model.parameters(), args.max_grad_norm)
默认最大梯度范数为 0.5,用于缓解异常 minibatch 造成的梯度爆炸。它与 PPO 的 probability-ratio clipping 作用不同:前者限制梯度范数,后者限制策略目标中的更新收益。
8. 当前代码的网络结构
ActorCritic 接受形状为 [B, 4] 的 observation batch:
observation [B, 4]
│
Linear(4, 64) + Tanh
│
Linear(64, 64) + Tanh
│
├── actor: Linear(64, 2) ──► logits [B, 2]
│ │
│ └──► Categorical policy
│
└── critic: Linear(64, 1) ──► value [B]
其中 \(B\) 是 batch size,默认隐藏层大小为 64。
8.1 为什么 actor 输出 logits
logits 是未归一化的对数概率。Categorical(logits=logits) 内部会完成数值稳定的归一化,并提供:
sample():按概率随机采样训练动作。log_prob(action):计算策略梯度和概率比率所需的对数概率。entropy():计算探索奖励。
8.2 初始化
所有线性层使用正交权重初始化和零 bias。正交初始化有助于深度强化学习初期保持较稳定的激活与梯度尺度。
9. 训练数据流与代码对应
9.1 Rollout 中保存什么
单次 rollout 默认收集 2048 个环境 step。每一步保存:
| 数据 | 单步形状 | rollout batch 形状 | 用途 |
|---|---|---|---|
observation |
[4] |
[T, 4] |
重新计算新策略和新价值 |
action |
标量 | [T] |
计算动作的新 log probability |
old_log_prob |
标量 | [T] |
构造新旧策略概率比率 |
reward |
标量 | [T] |
TD error 和 return |
value |
标量 | [T] |
GAE 和 value target |
next_value |
标量 | [T] |
TD bootstrap |
terminated |
标量 | [T] |
控制是否允许 bootstrap |
done |
标量 | [T] |
阻止 GAE 穿过回合边界 |
这里 \(T=\min(\text{rollout steps},\text{remaining total steps})\),因此最后一个 rollout 可以短于 2048 步。
9.2 一次完整更新
一次 PPO 迭代可概括为:
1. 用当前策略与环境交互 T 步
2. 保存 observation、action、old_log_prob、reward、value 和边界标记
3. 从后向前计算 GAE
4. 计算 return = raw_advantage + old_value
5. 标准化 policy 使用的 advantage
6. 打乱 T 个样本
7. 切分 minibatch
8. 对同一 rollout 重复 update_epochs 轮:
a. 重新计算 new_log_prob、new_value 和 entropy
b. 计算 ratio 和 clipped policy loss
c. 计算 value loss 和 total loss
d. 反向传播、梯度裁剪、Adam 更新
9. 丢弃旧 rollout,使用更新后的策略收集新数据
9.3 伪代码
初始化 Actor-Critic 参数 θ 和 Adam optimizer
重置环境得到 s
while global_step < total_timesteps:
rollout = []
for t in 0 ... T-1:
dist, V(s) = model(s)
a ~ dist # 训练时随机采样
old_log_prob = log π_old(a|s)
s_next, r, terminated, truncated = env.step(a)
next_value = V(s_next)
保存 transition
if terminated or truncated:
重置环境
else:
s = s_next
从后向前计算 δ_t 和 GAE advantage
returns = raw_advantages + old_values
标准化 advantages
repeat update_epochs times:
随机打乱 rollout
for each minibatch:
计算新策略的 log_prob、value、entropy
ratio = exp(new_log_prob - old_log_prob)
计算 clipped policy loss
计算 value loss 和 total loss
optimizer.step()
保存模型权重和网络元数据
9.4 为什么能重复使用一个 rollout
第一次 minibatch 更新前,ratio 接近 1。多轮更新后,新策略逐渐偏离采样策略。PPO clipping 会抑制继续扩大这种偏离的收益,使同一批样本能够进行有限次数的复用。
update_epochs 不能无限增大:clip 并不保证策略完全不会漂移。过多轮更新仍可能对 rollout 过拟合,并降低下一批 on-policy 数据的质量。
10. terminated 与 truncated
Gymnasium 的 step 返回:
next_observation, reward, terminated, truncated, info = env.step(action)
二者都表示当前 episode 需要结束,但数学意义不同:
terminated=True:MDP 内部终止,例如杆倒下。终止状态之后没有未来回报,因此不 bootstrap。truncated=True:外部时间限制结束。底层任务本身未必终止,因此应使用 \(V(s_{t+1})\) bootstrap。
当前实现定义两个 mask。令:
TD error 使用 terminated mask:
- 真正终止时 \(u_t=1\),未来价值被置零。
- 时间截断时 \(u_t=0\),保留下个状态价值。
GAE 递推使用 done mask:
不论真正终止还是时间截断,环境接下来都会 reset。因此 GAE 不能从新 episode 传播回旧 episode。这个双 mask 设计同时满足“截断时 bootstrap”和“不跨 reset 传播”两个要求。
在 rollout 自身的最后一步,即使 episode 尚未结束,也会用已计算的 next_value 完成一步 bootstrap;但更远期的 GAE 链在本批边界停止。
11. 推理与可视化
训练和推理的动作选择不同:
11.1 训练时随机采样
action = distribution.sample()
随机采样保留探索。如果训练时始终选择概率最大的动作,策略可能过早锁定在次优行为上。
11.2 推理时确定性选择
action = logits.argmax(dim=-1).item()
推理选择概率最大的动作,使结果更稳定,并展示网络当前认为最优的策略。
11.3 Checkpoint 格式
训练完成后默认保存 ppo_cartpole.pt,其中包含:
| 字段 | 含义 |
|---|---|
model_state_dict |
Actor-Critic 参数 |
env_id |
训练环境 ID |
obs_dim |
observation 维度 |
action_dim |
动作数量 |
hidden_size |
隐藏层宽度 |
training_config |
训练 CLI 配置 |
global_step |
实际训练环境步数 |
infer.py 根据这些元数据重建网络,不需要手动重复指定结构。加载后调用 model.eval() 并在 torch.no_grad() 中推理。
可视化要求在创建环境时指定:
gym.make(checkpoint["env_id"], render_mode="human")
这是 Gymnasium 新渲染 API 的要求,不能在环境创建后临时切换渲染模式。
12. 安装和使用
以下命令默认从仓库根目录执行。
12.1 安装依赖
项目的 environment.yaml 已声明 PyTorch 和 Gymnasium Classic Control。使用 Conda/Mamba 创建环境:
mamba env create -f environment.yaml
mamba activate sketch
如果已有 PyTorch 环境,只安装 Gymnasium Classic Control:
pip install "gymnasium[classic-control]"
12.2 使用默认参数训练
python tests/ppo/train.py
默认训练 100,000 个环境 step,并保存:
tests/ppo/ppo_cartpole.pt
日志示例:
step= 12000 episode= 230 return= 145.0 length=145 mean_20= 118.4
return:当前 episode 总奖励。length:当前 episode 步数;CartPole 中与 return 相同。mean_20:最近 20 个完整 episode 的平均奖励。
12.3 自定义训练
python tests/ppo/train.py \
--total-timesteps 200000 \
--rollout-steps 2048 \
--learning-rate 3e-4 \
--device auto \
--checkpoint /tmp/my_cartpole.pt
查看全部参数:
python tests/ppo/train.py --help
12.4 可视化推理
python tests/ppo/infer.py --episodes 5
加载自定义 checkpoint:
python tests/ppo/infer.py \
--checkpoint /tmp/my_cartpole.pt \
--episodes 10 \
--device auto
推理会打开 pygame 窗口,并在终端打印每回合奖励。如果当前机器没有图形桌面,render_mode="human" 可能无法创建窗口,应在具有显示服务的本机运行。
13. 超参数详解
| CLI 参数 | 默认值 | 作用 | 调整建议 |
|---|---|---|---|
--total-timesteps |
100000 | 总环境交互步数 | 奖励仍上升时可增加 |
--rollout-steps |
2048 | 每次更新前收集的样本数 | 太小会使优势估计噪声大,太大则更新不频繁 |
--learning-rate |
3e-4 | Adam 学习率 | 不稳定时降低,学习过慢时谨慎提高 |
--gamma |
0.99 | 奖励折扣因子 | 越大越重视长期奖励 |
--gae-lambda |
0.95 | GAE 偏差-方差折中 | 越大通常偏差低、方差高 |
--clip-coef |
0.2 | PPO ratio 裁剪范围 | 太大可能更新剧烈,太小可能学习缓慢 |
--update-epochs |
10 | 每个 rollout 的重复优化轮数 | 太大可能过拟合旧 rollout |
--minibatch-size |
64 | 每次梯度更新样本数 | 小 batch 噪声更大,大 batch 更新更平滑 |
--value-coef |
0.5 | 价值损失权重 | Critic 欠拟合时可提高,但会影响共享主干 |
--entropy-coef |
0.01 | 探索奖励权重 | 过早确定化时提高,长期随机时降低 |
--max-grad-norm |
0.5 | 梯度范数上限 | 梯度爆炸时降低 |
--hidden-size |
64 | 共享隐藏层宽度 | CartPole 通常无需很大网络 |
--seed |
1 | Python、NumPy、PyTorch 和环境种子 | 用多个 seed 评估稳定性 |
--device |
cpu | cpu、cuda 或 auto |
小型 CartPole 使用 CPU 通常已经足够快 |
--checkpoint |
脚本目录下默认路径 | 模型保存位置 | 实验间使用不同路径避免覆盖 |
13.1 参数之间的关系
每个 rollout 大致产生:
默认完整 rollout 对应:
次 optimizer step。增大 rollout-steps、增大 update-epochs 或减小 minibatch-size 都会增加每批数据的计算量,但含义并不相同。
13.2 如何判断学会了
CartPole 的回报上限为 500。比单次达到 500 更可靠的判断方式是观察多个 episode 的平均回报。由于初始化状态和动作采样具有随机性,应使用多个随机种子重复训练,而不是只报告最好的一次结果。
14. 常见问题和调试方法
14.1 奖励一直停留在 10~30
检查:
- 是否误把推理的
argmax用在训练阶段,导致探索不足。 old_log_prob是否在 rollout 时保存,而不是更新期间重新计算。- PPO ratio 是否为
exp(new_log_prob - old_log_prob)。 - policy loss 的正负号是否正确。
- 优势是否在构造 returns 之后再标准化。
可尝试降低学习率到 1e-4,或增加总训练步数。单个 seed 的短期波动不一定表示实现错误。
14.2 奖励先提高后突然下降
可能原因:
- 学习率或
clip-coef太大,策略更新过猛。 update-epochs太大,对单批 rollout 过拟合。entropy-coef太低,策略过早变得确定。- Critic 误差过大,导致优势估计不可靠。
当前实现没有 KL early stopping。如果扩展到更困难的任务,可记录 approximate KL、clip fraction、value loss 和 entropy,以判断策略是否更新过度。
14.3 为什么 CPU 可能比 CUDA 更快
网络只有两层、batch 也很小,环境交互本身发生在 CPU。将小张量频繁传到 GPU 的调度开销可能超过计算收益,因此 CartPole 默认使用 CPU。CUDA 支持主要用于接口完整性和后续扩展。
14.4 无法打开推理窗口
Gymnasium Classic Control 的 human rendering 需要 pygame 和有效的显示服务。确认:
- 安装的是
gymnasium[classic-control],而不只是最小版gymnasium。 - Linux 环境存在有效的
DISPLAY或 Wayland 会话。 - 没有在纯 SSH、Docker 或 CI 无显示环境中直接请求 human rendering。
14.5 checkpoint 找不到
默认训练和推理都使用 tests/ppo/ppo_cartpole.pt。如果训练时传入了自定义 --checkpoint,推理时必须传入相同路径。
14.6 checkpoint 结构不兼容
infer.py 会检查 model_state_dict、env_id、obs_dim、action_dim 和 hidden_size。缺少这些字段时会明确报错。不要把只有裸 state_dict 的其他模型直接当作本项目 checkpoint。
14.7 CUDA 请求失败
--device cuda 在 CUDA 不可用时会直接抛出错误。希望自动回退 CPU 时使用:
python tests/ppo/train.py --device auto
15. 参考资料
- John Schulman et al., Proximal Policy Optimization Algorithms, 2017.
- John Schulman et al., High-Dimensional Continuous Control Using Generalized Advantage Estimation, 2015.
- Gymnasium, CartPole-v1.
- Gymnasium, 旧 Gym 到新 step API 的迁移说明.
- PyTorch,
torch.distributions与Categorical.
最简运行顺序:
python tests/ppo/train.py
python tests/ppo/infer.py --episodes 5