外观
渐进式教程:Gymnasium 三版跑起来
一句话定位:这一页用同一个 CartPole 问题,带你把三种代表算法——表格 Q-learning、REINFORCE、PPO——从零写出来并跑通。它解决"原理都懂、一写代码就报错"的断层,适合刚看完概念页、准备第一次动手的读者。读完你会有三份能跑的代码、一张三算法对比表,以及一套最小可复现实验模板。
本页代码基于 Gymnasium(OpenAI Gym 的社区维护继任者)。三个版本共享同一个环境接口,但学习机制完全不同:
CartPole 三版进化路线
┌──────────────────────────────────────────────────────────────┐
│ 版本一 表格 Q-learning 价值学习:查表,观测要离散化 │
│ │ │
│ 版本二 REINFORCE 策略梯度:学一个网络 π(a|s) │
│ │ │
│ 版本三 极简 PPO 价值+策略:Actor-Critic + GAE │
│ │ │
│ 共同点:同一个 env API,同一个"能上 475 分"验收标准 │
└──────────────────────────────────────────────────────────────┘先说明:这不是三选一,而是三座里程碑。你亲手写一遍 Q-learning,才能真懂"价值高估"和离散化的代价;写一遍 REINFORCE,才能真懂什么叫高方差;写一遍 PPO,才算真正踏入现代深度 RL 的门槛。三版的原理分别对应价值学习、策略梯度方法、Actor-Critic 家族三页,建议边写边对照。
一、安装与 Gymnasium API 速览
1. 安装
bash
# 基础环境
pip install gymnasium numpy
# PyTorch(版本二、三需要,CPU 版即可跑 CartPole)
pip install torch --index-url https://download.pytorch.org/whl/cpu # CPU 版示例
# 可选:画学习曲线
pip install matplotlib版本锁定建议
RL 生态对版本敏感。本文以 gymnasium >= 0.28、numpy >= 1.24 为准。装完跑一下 python -c "import gymnasium; print(gymnasium.__version__)"。环境库的完整索引见数据集与工具档案。
2. 最小 API 探针
先跑通下面这段"API 探针"——它一次性验证你环境中每个 API 的名字和返回结构:
python
import gymnasium as gym
# 创建环境;render_mode 设为 "human" 会弹出窗口
env = gym.make("CartPole-v1", render_mode="rgb_array")
# 观测空间与动作空间
print("observation_space:", env.observation_space)
print("action_space:", env.action_space)
# reset:新版返回 (obs, info) 两元组
obs, info = env.reset(seed=42)
print("reset 后 obs 的 shape:", obs.shape, "dtype:", obs.dtype)
print("reset 后 info:", info)
# 随机采样一个动作,执行一步
action = env.action_space.sample()
obs, reward, terminated, truncated, info = env.step(action)
print(f"step 返回: obs={obs}, reward={reward}")
print(f"terminated={terminated}, truncated={truncated}, info={info}")
env.close()输出会告诉你几件重要的事:
| 事实 | 内容 | 为什么重要 |
|---|---|---|
observation_space | Box([-4.8 -inf -0.418 -inf], [4.8 inf 0.418 inf], (4,), float32) | 连续观测、4 维、float32 |
action_space | Discrete(2) | 离散动作:左推/右推 |
reset | 返回 (obs, info) 两元组 | 版本不同,返回结构不同 |
step | 返回 5 元组,含 terminated 与 truncated | 终止与截断必须分开处理 |
Gymnasium API 的两个历史坑
- 老教程里
env.reset()只返回obs;从 Gymnasium 0.26 起返回(obs, info)。 - 老教程里只有 4 元组
(obs, reward, done, info);Gymnasium 把done拆成了terminated(任务达成)与truncated(外力截断,如超时)。在 CartPole 里,truncated恰恰是"站到了 500 步"这种好事,把它当失败处理会毁掉学习信号——这是本文档最要强调的坑之一,后面代码里你会看到正确的处理方式。
3. 关于 render
gym.make("CartPole-v1", render_mode="human") 弹窗可视化;不想弹窗就用 "rgb_array"(可保存成图或 GIF)。训练时不要用 "human",它会大幅拖慢速度。想亲眼看效果,在训练脚本最后跑几集带渲染的评估即可。
二、CartPole 问题介绍
CartPole 是最经典的"第一个 RL 问题":一根杆子立在可移动的小车上,你要左右推车,让杆子尽可能久地保持竖直。
- 状态(观测):4 个实数——车位置
x、车速度、杆子与竖直的夹角θ、角速度。 - 动作:离散 2 个——向左推 / 向右推。
- 奖励:每存活一个时间步
+1。 - 终止:杆子倾角超过 ±12°、车滑出轨道(±2.4)、或坚持到 500 步(
truncated)。 - 目标:最大化累计奖励,理论上限 500。
杆子(角度 θ 越接近 0 越好)
│
┌────────┴───┐
│ 小车 │← 推力(左 / 右)
└────┬───────┘
└─────── 轨道(x 超界即终止)为什么拿它做教程:状态只有 4 维(表格化可行)、动作只有 2 个、一个 episode 只要几十秒、CPU 上分钟级出结果。它是验证"我的训练管线对不对"的极廉价探针——任何算法在 CartPole 上都应该能轻松学到 400+ 分,学不到,说明你的实现或环境处理有 bug。
三个版本的统一验收标准:连续 N 集的滑动平均回报 ≥ 475(即接近 500 上限)。达不到这个标准,先排查代码,不要怀疑"任务太难"。
三、版本一:表格 Q-learning(价值学习)
1. 思路
Q-learning 把"状态-动作对 → 价值"存进一张表。但 CartPole 的观测是 4 维连续值,先做离散化:把每一维切成若干区间,把一个连续观测映射成桶索引,Q 表就变成了"桶组合 × 动作"的矩阵。
更新规则(off-policy TD 控制):
text
Q(s,a) ← Q(s,a) + α · [ r + γ · max_a' Q(s',a') − Q(s,a) ]直觉:TD 误差 r + γ·max Q(s') − Q(s,a) 表示"这步之后价值比我以为的更好/更差",用学习率 α 挪一小步。细节推导见价值学习:从动态规划到 DQN。
2. 完整代码
python
"""
版本一:表格 Q-learning 求解 CartPole
跑通标准:滑动平均回报 ≥ 475(接近 500 上限)
运行:python q_learning_cartpole.py
"""
import numpy as np
import gymnasium as gym
# ---------- 1. 状态离散化 ----------
# CartPole 观测 4 维的大致取值范围
OBS_BOUNDS = np.array([
[-4.8, 4.8], # 车位置 x
[-3.0, 3.0], # 速度(截断)
[-0.5, 0.5], # 杆子角度 θ(超范围必终止)
[-3.0, 3.0], # 角速度(截断)
])
# 每维切成多少个桶:[位置, 速度, 角度, 角速度]
NUM_BINS = np.array([10, 10, 10, 10])
def discretize(obs):
"""把连续观测映射成一个离散索引(范围 0 .. prod(NUM_BINS)-1)。"""
bins = []
for i, (low, high) in enumerate(OBS_BOUNDS):
# 截断到边界内,再按桶宽划分
clipped = np.clip(obs[i], low, high)
idx = int((clipped - low) / (high - low) * NUM_BINS[i])
idx = min(idx, NUM_BINS[i] - 1)
bins.append(idx)
# 编码成单一索引(行主序)
index = 0
for i, b in enumerate(bins):
index = index * NUM_BINS[i] + b
return index
# ---------- 2. 超参数 ----------
ALPHA = 0.1 # 学习率
GAMMA = 0.99 # 折扣因子
EPS_START = 1.0 # 初始探索率
EPS_END = 0.01 # 最小探索率
EPS_DECAY = 0.9995 # 每集衰减
NUM_EPISODES = 8000
MAX_STEPS = 500 # 与 CartPole-v1 的截断步数一致
env = gym.make("CartPole-v1")
n_state = int(np.prod(NUM_BINS))
n_action = env.action_space.n
Q = np.zeros((n_state, n_action))
# ---------- 3. 训练循环 ----------
episode_returns = []
for ep in range(NUM_EPISODES):
obs, _ = env.reset(seed=ep) # 每集播种,保证可复现
state = discretize(obs)
eps = max(EPS_END, EPS_START * (EPS_DECAY ** ep))
total_reward = 0
for t in range(MAX_STEPS):
# ε-greedy 选动作
if np.random.rand() < eps:
action = env.action_space.sample()
else:
action = int(np.argmax(Q[state]))
obs, reward, terminated, truncated, _ = env.step(action)
next_state = discretize(obs)
# Q-learning 更新(off-policy:直接用 max 下一状态价值)
best_next = np.max(Q[next_state])
Q[state, action] += ALPHA * (reward + GAMMA * best_next - Q[state, action])
state = next_state
total_reward += reward
if terminated or truncated:
break
episode_returns.append(total_reward)
if (ep + 1) % 500 == 0:
mean = float(np.mean(episode_returns[-100:]))
print(f"episode {ep+1:>5d} | mean_return(last_100) = {mean:.1f} | eps = {eps:.3f}")
env.close()
print("Q-learning 训练完成。")3. 学习曲线预期与"跑通标准"
把这 8000 集的回报画出来,你会看到一条典型的阶梯式上升曲线:
text
回报
500 ┤ ┌─────────────────── 平台期(逼近上限)
│ ┌─────┘
400 ┤ ┌───┘
│ ┌──┘
300 ┤ ┌──┘
│ ┌─┘
200 ┤ ┌──┘
│ ┌─┘
100 ┤ ┌──┘
│ ┌──┘
0 ┤─┘
└───────────────────────────────────────────────▶ 回合数
0 1000 2000 3000 4000 5000 6000- 前几百集:探索主导(ε 很大),回报徘徊在 10~50。
- 几百到两千集:Q 表开始有用,回报阶梯式爬升——注意是"阶梯"不是"斜坡",因为学到一定水平后,episode 会突然长一截。
- 两三千集后:逼近 500 上限,滑动平均稳定在 475+。
如果想看它"失败"一次
把 EPS_DECAY 改成 0.9999(ε 衰减极慢),前 2000 集会一直停留在低分——这就是"探索太多"压住性能的直观教材,与探索与利用那页讲的现象完全一致。
代码级坑位(版本一专属):
| 坑 | 现象 | 修法 |
|---|---|---|
discretize 越界 | 某些观测落在桶索引之外,Q 索引报错 | 先 clip 再算索引(代码里已处理) |
把 truncated 当 terminated | 每次 500 步时被误判为"失败",学习被反向惩罚 | 两分支都 break 但不区分信号即可(本版不区分对学习无碍) |
| ε 衰减太快 | 还没探索完就贪心,卡在低分 | 从 1.0 起步、指数衰减到 0.01 |
观测 dtype 是 float32,discretize 里直接比较 | 精度问题导致桶边界偏移 | 统一 np.clip 处理(已内置) |
4. 版本一的三个领悟
- 离散化即信息损失:10×10×10×10 = 一万个格子,CartPole 勉强够用;高维问题(如 Atari 画面)格子数指数爆炸——这是维度诅咒的雏形。
- 查表法不能泛化:没见过的格子永远是 0 价值,必须靠足够密的离散覆盖。
- 价值学习天然适配离散动作:
max_a Q(s,a)在动作少时很便宜;动作变多就难了,这是 DQN 家族的硬约束(价值学习)。
四、版本二:REINFORCE(策略梯度)
1. 思路
不学价值表,直接学策略网络 π_θ(a|s):输入 4 维观测,输出 2 个动作的概率。参数更新方向是策略梯度定理:
text
∇J(θ) = E_τ [ Σ_t G_t · ∇log π_θ(a_t|s_t) ]直觉:哪个动作带来了更高的回报 G_t,就把它选中的概率往上提。REINFORCE 用整集蒙特卡洛回报估计 G_t,不用 bootstrap,因此无偏但方差大。推导与直觉见策略梯度方法。
2. 完整代码
python
"""
版本二:REINFORCE(带折扣回报归一化 baseline)求解 CartPole
跑通标准:滑动平均回报 ≥ 475
运行:python reinforce_cartpole.py
"""
import gymnasium as gym
import numpy as np
import torch
import torch.nn as nn
import torch.optim as optim
from torch.distributions import Categorical
class PolicyNet(nn.Module):
"""策略网络:观测 → 动作概率分布。"""
def __init__(self, obs_dim, act_dim, hidden=128):
super().__init__()
self.net = nn.Sequential(
nn.Linear(obs_dim, hidden), nn.Tanh(),
nn.Linear(hidden, hidden), nn.Tanh(),
nn.Linear(hidden, act_dim), # 输出 logits(不经过 softmax)
)
def forward(self, obs):
return Categorical(logits=self.net(obs)) # 内部自带 softmax
def discount_returns(rewards, gamma=0.99):
"""计算每个时间步的折扣回报 G_t(蒙特卡洛,无 bootstrap)。"""
returns = np.zeros(len(rewards), dtype=np.float32)
g = 0.0
for t in reversed(range(len(rewards))):
g = rewards[t] + gamma * g
returns[t] = g
return returns
def main(seed=0, num_episodes=3000, lr=1e-2, gamma=0.99):
torch.manual_seed(seed)
np.random.seed(seed)
env = gym.make("CartPole-v1")
policy = PolicyNet(env.observation_space.shape[0], env.action_space.n)
optimizer = optim.Adam(policy.parameters(), lr=lr)
episode_returns = []
for ep in range(num_episodes):
obs, _ = env.reset(seed=seed + ep)
log_probs, rewards = [], []
# 1) 采样一整集
for _ in range(500):
obs_t = torch.as_tensor(obs, dtype=torch.float32)
dist = policy(obs_t)
action = dist.sample()
log_probs.append(dist.log_prob(action))
obs, reward, terminated, truncated, _ = env.step(action.item())
rewards.append(reward)
if terminated or truncated:
break
# 2) 计算回报,并做标准化(减去均值除以标准差——这是"baseline"的蒙特卡洛版)
returns = torch.as_tensor(discount_returns(rewards, gamma), dtype=torch.float32)
returns = (returns - returns.mean()) / (returns.std() + 1e-8)
# 3) 策略梯度更新(REINFORCE 损失 = - Σ G_t · log π(a_t|s_t))
loss = -torch.stack([
lp * ret for lp, ret in zip(log_probs, returns)
]).sum()
optimizer.zero_grad()
loss.backward()
optimizer.step()
episode_returns.append(sum(rewards))
if (ep + 1) % 100 == 0:
mean = float(np.mean(episode_returns[-100:]))
print(f"episode {ep+1:>5d} | mean_return(last_100) = {mean:.1f}")
env.close()
return episode_returns
if __name__ == "__main__":
main()3. 学习曲线预期:方差之王
REINFORCE 的曲线锯齿极大,这是它的本性,不是 bug:
text
回报
500 ┤ ┌──┐ ┌──┐
│ ┌──┘ └──┐ ┌──┘ └──┐
400 ┤──┘ └─┘ └────┐
│ └──┐ ┌────
300 ┤ └──┘
│
200 ┤
│
100 ┤
│
0 ┤
└──────────────────────────────────────────▶ 回合数
0 1000 2000 3000- 前几百集回报在 10~60 之间剧烈振荡。
- 大约 1000~2000 集开始出现"突然上 400"的尖峰,随后又掉下来——梯度方差大的直接表现。
- 3000 集左右,滑动平均(窗口 100)多数时间能站上 475。
高方差是你这一版要"看见"的东西
同样是 CartPole,Q-learning 是平滑爬升,REINFORCE 是震荡爬升。同样的任务、同样的目标,差的只是估计器——这就是"REINFORCE 的蒙特卡洛回报方差高"这句话的实体形态。后面 PPO 用 GAE + 裁剪把方差压下来,你会直观看到曲线变稳。想更深刻地体会方差,把 lr 从 1e-2 提到 5e-2,观察曲线从"震荡"变成"发疯"。
4. 版本二的坑位
| 坑 | 现象 | 修法 |
|---|---|---|
| 回报不标准化 | 梯度量级不稳定,收敛极慢 | returns = (returns - mean)/(std+eps) |
loss 忘了 sum() 用 mean() | 梯度量级随集长变化 | 用 sum 或先归一化再 mean,保持量级一致 |
| 学习率太大(>1e-2) | 一集好成绩被"过冲"回滚,训练崩 | 降到 1e-2 或更小;这是策略梯度最敏感的坑 |
torch.as_tensor(obs) 忘了指定 dtype | 偶发 dtype 不匹配报错 | 显式 dtype=torch.float32 |
把 step 的返回值顺序记错 | 解包报错 | 牢记五元组 (obs, reward, terminated, truncated, info) |
5. 版本二的三个领悟
- 连续动作不再是难题:策略网络直接输出分布参数,REINFORCE 天然支持连续动作(高斯分布)——而表格 Q-learning 在连续动作面前寸步难行。
- 方差是策略梯度的阿喀琉斯之踵:整集回报的方差随 episode 长度线性增长,这也是 TRPO/PPO 存在的原因(Actor-Critic 家族)。
- baseline 思想:回报标准化就是最简单的 baseline。价值网络做 advantage baseline 是更优解——这正是版本三。
五、版本三:极简 PPO(Actor-Critic + GAE + Clip)
1. 思路
PPO 在策略梯度上加了三个东西,对应三个动机:
| 组件 | 作用 | 直觉 |
|---|---|---|
| Actor-Critic | 价值网络提供优势估计 A(s,a) 做 baseline | 用"这步比平均好多少"替代"这步绝对回报多少",降方差 |
| GAE(广义优势估计) | 在方差与偏差间做可调权衡 | λ 大→偏 MC(高方差低偏差);λ 小→偏 TD(低方差高偏差) |
| Clip 裁剪 | 限制每轮更新步长,防止一次过冲 | "信任域"的廉价近似,让训练稳 |
PPO 的裁剪目标(surrogate loss):
text
L = min( r(θ)·A , clip(r(θ), 1−ε, 1+ε)·A )
其中 r(θ) = π_θ(a|s) / π_θ_old(a|s) (新旧策略概率比)直觉一句话:只有"旧策略选了它"的动作才有资格被强化,且单轮更新把新旧概率比限制在 [1−ε, 1+ε] 内。公式与机制详解见策略梯度方法与Actor-Critic 家族。
2. 完整代码(约 150 行,单环境版)
python
"""
版本三:极简 PPO 求解 CartPole(单环境、无向量化,突出可读性)
跑通标准:滑动平均回报 ≥ 475
运行:python ppo_cartpole.py
"""
import gymnasium as gym
import numpy as np
import torch
import torch.nn as nn
import torch.optim as optim
from torch.distributions import Categorical
class ActorCritic(nn.Module):
"""共享特征层 + 策略头 + 价值头。"""
def __init__(self, obs_dim, act_dim, hidden=128):
super().__init__()
self.shared = nn.Sequential(
nn.Linear(obs_dim, hidden), nn.Tanh(),
nn.Linear(hidden, hidden), nn.Tanh(),
)
self.policy_head = nn.Linear(hidden, act_dim) # logits
self.value_head = nn.Linear(hidden, 1) # 标量价值 V(s)
def dist(self, obs):
return Categorical(logits=self.policy_head(self.shared(obs)))
def value(self, obs):
return self.value_head(self.shared(obs)).squeeze(-1)
def compute_gae(rewards, values, next_value, dones, gamma, lam):
"""GAE:返回 (advantages, returns)。values 不含 next_value。"""
T = len(rewards)
advantages = torch.zeros(T)
last_gae = 0.0
for t in reversed(range(T)):
# 若 t 步 truncated/terminated,则后续无自举
mask = 1.0 - float(dones[t])
delta = rewards[t] + gamma * next_value * mask - values[t]
last_gae = delta + gamma * lam * mask * last_gae
advantages[t] = last_gae
next_value = values[t]
returns = advantages + values
return advantages, returns
def main(seed=0, total_steps=100_000, n_steps=256, gamma=0.99, lam=0.95,
lr=3e-4, clip_eps=0.2, update_epochs=10, minibatch_size=64,
ent_coef=0.01, vf_coef=0.5):
torch.manual_seed(seed)
np.random.seed(seed)
env = gym.make("CartPole-v1")
ac = ActorCritic(env.observation_space.shape[0], env.action_space.n)
optimizer = optim.Adam(ac.parameters(), lr=lr)
obs, _ = env.reset(seed=seed)
num_updates = total_steps // n_steps
for update in range(num_updates):
# ---------- 1) 采样一段 rollout ----------
# 收集阶段只存 Python 标量,结束再一次性转张量(清晰且不易踩 dtype 坑)
obs_list, act_list, rew_list, dones_list, val_list = [], [], [], [], []
for _ in range(n_steps):
obs_t = torch.as_tensor(obs, dtype=torch.float32)
with torch.no_grad():
dist = ac.dist(obs_t)
value = ac.value(obs_t)
action = dist.sample()
obs_list.append(obs_t)
act_list.append(action)
val_list.append(value)
obs, reward, terminated, truncated, _ = env.step(action.item())
done = terminated or truncated
rew_list.append(float(reward))
dones_list.append(float(done))
if done:
obs, _ = env.reset()
# ---------- 2) GAE 计算优势 ----------
obs_t = torch.as_tensor(obs, dtype=torch.float32)
with torch.no_grad():
next_value = ac.value(obs_t)
obs_tensor = torch.stack(obs_list)
with torch.no_grad():
old_dist = ac.dist(obs_tensor)
old_log_probs = old_dist.log_prob(torch.stack(act_list))
values_tensor = torch.stack(val_list)
rewards_tensor = torch.tensor(rew_list) # [T] 一次性转张量
dones_tensor = torch.tensor(dones_list) # [T]
advantages, returns = compute_gae(
rewards_tensor, values_tensor, next_value, dones_tensor, gamma, lam
)
advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)
# ---------- 3) 多轮小批量更新 ----------
n_samples = n_steps
for _ in range(update_epochs):
perm = np.random.permutation(n_samples)
for start in range(0, n_samples, minibatch_size):
idx = perm[start:start + minibatch_size]
batch_obs = obs_tensor[idx]
batch_act = torch.stack(act_list)[idx]
batch_old = old_log_probs[idx]
batch_adv = advantages[idx]
batch_ret = returns[idx]
dist = ac.dist(batch_obs)
log_probs = dist.log_prob(batch_act)
entropy = dist.entropy().mean()
# 概率比 r(θ)
ratio = (log_probs - batch_old).exp()
# 裁剪目标
pg_loss = -torch.min(
ratio * batch_adv,
torch.clamp(ratio, 1.0 - clip_eps, 1.0 + clip_eps) * batch_adv,
).mean()
vf_loss = nn.functional.mse_loss(ac.value(batch_obs), batch_ret)
loss = pg_loss + vf_coef * vf_loss - ent_coef * entropy
optimizer.zero_grad()
loss.backward()
optimizer.step()
# ---------- 4) 日志 ----------
total_reward = 0
eval_obs, _ = env.reset(seed=seed + 9999)
for _ in range(500):
with torch.no_grad():
eval_act = ac.dist(torch.as_tensor(eval_obs, dtype=torch.float32)).sample()
eval_obs, r, term, trunc, _ = env.step(eval_act.item())
total_reward += r
if term or trunc:
break
print(f"update {update+1:>4d}/{num_updates} | eval_return = {total_reward}")
env.close()
if __name__ == "__main__":
main()张量收集的两种姿势
上面用了"收集期存 Python 标量、结束一次性 torch.tensor 转张量"的干净姿势。另一种常见写法是直接收集张量再 torch.stack——两者都行,关键是全程 dtype 一致。如果你在调试中看到 torch.stack 报 size/dtype 不匹配,多半是列表里混进了不同形状的张量(比如把标量 reward 和 (1,) 张量混存),对照上面 GAE 段的转换方式逐行检查即可。
一个更干净、推荐正式使用的做法:
python
# 收集阶段只存 Python 标量
obs_list, act_list, rew_list, dones_list, val_list = [], [], [], [], []
for _ in range(n_steps):
...
rew_list.append(float(reward))
dones_list.append(float(done))
...
# 结束后一次性转张量
obs_tensor = torch.stack(obs_list) # [T, obs_dim]
act_tensor = torch.stack(act_list) # [T]
rew_tensor = torch.tensor(rew_list) # [T]
done_tensor = torch.tensor(dones_list) # [T]
val_tensor = torch.stack(val_list) # [T]3. 学习曲线预期:又快又稳
text
回报
500 ┤ ┌─────────────── 平台期
│ ┌─┘
400 ┤ ┌─┘
│ ┌──┘
300 ┤ ┌──┘
│ ┌──┘
200 ┤ ┌───┘
│ ┌──┘
100 ┤─┘
│
0 ┤
└──────────────────────────────────────────▶ 更新次数
0 20 40 60 80 100 120注意横轴单位从"集"变成了"更新次数"(每次更新 256 步)。在 CartPole 上,约 50~100 次更新(对应 1.3~2.5 万环境步)就能稳定上 475——样本效率比 REINFORCE 高一个量级,且曲线几乎没有 REINFORCE 那种骤降。
| 对比维度 | Q-learning | REINFORCE | 极简 PPO |
|---|---|---|---|
| 达到 475 所需环境步数 | ~2~3 万 | ~10~20 万 | ~1.5~3 万 |
| 曲线平稳性 | 平滑 | 锯齿剧烈 | 平稳 |
| 代码量(行) | ~70 | ~80 | ~150 |
| 理论基础 | 价值学习 | 策略梯度 | Actor-Critic + GAE |
| 能否上连续动作 | 不能 | 能 | 能 |
4. 版本三的坑位(比前两版多一倍)
| 坑 | 现象 | 修法 |
|---|---|---|
old_log_probs 用的不是采样时的策略 | 概率比失真,训练不稳定 | 必须在更新前用冻结的旧策略算好,且用 no_grad |
GAE 的 dones 忘加 mask | 跨 episode 边界错误自举,价值估计被污染 | delta = r + γ·V'·mask − V |
| 优势不归一化 | 不同集优势量级不同,学习率难调 | (A−mean)/(std+eps) |
单环境连续跑,reset 漏调用 | 观测沿用终止帧,value 预测失真 | done 后立即 reset |
| 熵系数设 0 | CartPole 尚可,复杂环境容易过早收敛到次优 | 保留小熵正则 ent_coef=0.01 |
env.step 里 action 是 torch.Tensor | Gymnasium 报 TypeError | .item() 转成 Python int |
想跳过手写,用 SB3 对照?
把框架对比里的 Stable-Baselines3 拿出来,一行 PPO("MlpPolicy", "CartPole-v1").learn(50000) 就能得到同样的结果。手写版本的价值在于:你能读懂 SB3 日志里每个字段。建议先手写跑通,再上框架,两条腿走路。
六、三版对比表与选择建议
| 维度 | 版本一 Q-learning | 版本二 REINFORCE | 版本三 PPO |
|---|---|---|---|
| 学习对象 | Q 表 | 策略网络 | 策略 + 价值网络 |
| 是否 bootstrap | 是(TD) | 否(纯 MC) | 是(GAE) |
| 方差 | 低 | 很高 | 低 |
| 偏差 | 有(离散化 + 自举) | 无(无偏) | 有(自举可控) |
| 观测要求 | 必须离散化 | 连续直用 | 连续直用 |
| 动作空间 | 离散 | 任意 | 任意 |
| 样本效率 | 中 | 低 | 高 |
| 离线更新 | 支持(off-policy) | 不支持(on-policy) | 支持(小批量 on-policy) |
| 典型用途 | 教学、小规模离散问题 | 教学、理解策略梯度 | 现代 RL 默认选择 |
工程师怎么选:
- 教学/学习:按本页顺序三个都写一遍。
- 小规模离散任务(棋类、简单调度):Q-learning 家族够用且极稳。
- 真实业务(连续控制、机器人、组合优化):直接 PPO/SAC 系(Actor-Critic 家族)。
- 先跑通基线验证管线:用 SB3 的 PPO 五分钟出基线,再决定要不要自研。
七、常见坑位总表(三版通用)
这些坑在 CartPole 上就能复现,但往往要到真实项目里才吃大亏:
| # | 坑 | 症状 | 检测/修法 |
|---|---|---|---|
| 1 | 观测 dtype 不统一(float32/float64 混用) | 偶发类型报错或精度漂移 | 统一 torch.float32 |
| 2 | terminated/truncated 混为一谈 | 超时被当失败,学习信号反向 | 该区分的场景显式区分(CartPole 的 500 步是好结局) |
| 3 | 训练用 render_mode="human" | 训练慢 10 倍 | 训练 rgb_array/不渲染,评估时再渲染 |
| 4 | seed 只播一个随机源 | 环境随机与算法随机混在一起 | 三层播种(Python/numpy/torch + env),见从零构建一个 RL 项目 |
| 5 | 只看单次运行曲线 | 被方差骗 | 多 seed 运行,画 IQR 区间,见从零搭一套 RL 评估 |
| 6 | env 不 close() | 多进程环境句柄泄漏 | try/finally 或 with gym.make(...) as env |
八、最小可复现实验模板
三个版本跑通后,把下面这个模板存成 run_experiment.py,它会成为你所有后续实验的起点——把"训练 + 评估 + 存档"固化成一个命令:
python
"""最小可复现实验模板:跑一次 PPO-CartPole 并保存完整产物。"""
import argparse
import json
import numpy as np
import gymnasium as gym
import torch
from ppo_cartpole import ActorCritic # 复用版本三的模型
def evaluate(ac, seed, episodes=10):
env = gym.make("CartPole-v1")
returns = []
for i in range(episodes):
obs, _ = env.reset(seed=seed * 100 + i)
total = 0
for _ in range(500):
with torch.no_grad():
act = ac.dist(torch.as_tensor(obs, dtype=torch.float32)).sample()
obs, r, term, trunc, _ = env.step(act.item())
total += r
if term or trunc:
break
returns.append(total)
env.close()
return np.mean(returns), np.std(returns)
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--seed", type=int, default=0)
parser.add_argument("--lr", type=float, default=3e-4)
parser.add_argument("--total_steps", type=int, default=50_000)
args = parser.parse_args()
# 记录完整配置(实验身份证)
config = vars(args)
torch.manual_seed(args.seed); np.random.seed(args.seed)
ac = ActorCritic(4, 2)
# ... 这里接版本三的训练循环 ...
mean, std = evaluate(ac, args.seed)
# 产物:模型 + 配置 + 指标,全部落盘
torch.save(ac.state_dict(), f"model_seed{args.seed}.pt")
with open(f"result_seed{args.seed}.json", "w", encoding="utf-8") as f:
json.dump({**config, "eval_mean": mean, "eval_std": std}, f, ensure_ascii=False, indent=2)
print(f"seed={args.seed} eval_return={mean:.1f} ± {std:.1f}")
if __name__ == "__main__":
main()跑三次不同 seed 并合并结果:
bash
for s in 0 1 2; do python run_experiment.py --seed $s; done这样你就有了"多 seed、可复现、产物落盘"的最小闭环——它正是评估实践里实验矩阵的最小区块。在此基础上,再加 Hydra 配置管理 和 W&B/TensorBoard 日志,就是一个工程级的实验系统了。
九、扩展到其他环境
CartPole 跑通后,把 gym.make 的环境名换掉即可试水更难的任务:
| 环境 | 动作空间 | 难度跃迁 | 需要改什么 |
|---|---|---|---|
LunarLander-v2 | 离散 4 | 奖励稀疏、接触动力学 | 几乎不改代码(离散动作三版通用) |
MountainCar-v0 | 离散 3 | 奖励极稀疏,纯 Q-learning 会很难 | Q-learning 需加奖励塑形或课程 |
Acrobot-v1 | 离散 3 | 连续摆动控制 | 代码不变,调超参 |
Pendulum-v1 | 连续 1 | 首次接触连续动作 | 版本三把 Categorical 换成高斯分布 |
HalfCheetah-v4/Humanoid-v4(MuJoCo) | 连续 6/17 | 高维连续控制 | 需要向量化采样 + GPU;上框架(SB3/Brax)更实际 |
每次换环境,重复同一套验收动作
新环境跑通 = 随机基线有分数 → 简单策略能学到东西 → 三版对比表能重现"PPO 又快又稳"。这是评估实践强调的"先建基线再换算法"最小实践。
环境库的更多选择(Brax、MuJoCo、Procgen、Atari 等)见数据集与工具档案;三版代码对应的算法家族详解见价值学习、策略梯度方法、Actor-Critic 家族。
延伸阅读
- 价值学习:从动态规划到 DQN —— 版本一的理论完整版:从 TD 到 DQN 家族
- 策略梯度方法 —— 版本二的理论完整版:REINFORCE、baseline、PPO 目标
- Actor-Critic 家族 —— 版本三的理论完整版:GAE、SAC、DDPG 谱系
- 从零构建一个 RL 项目 —— 把本页的"单实验"升级成"完整项目流水线"
- 从零搭一套 RL 评估 —— 本页验收标准的多 seed 严格版
- 框架与工具怎么选 —— 手写跑通后,怎么接 SB3/RLlib/CleanRL
参考资料
- Gymnasium 官方文档(Farama Foundation):gymnasium.farama.org;
Env接口见 gymnasium.farama.org/api/env/ - Watkins, C. J. C. H. & Dayan, P. (1992). Q-Learning. Machine Learning, 8, 279–292. link.springer.com/article/10.1007/BF00992698
- Williams, R. J. (1992). Simple Statistical Gradient-Following Algorithms for Connectionist Reinforcement Learning. Machine Learning, 8, 229–256. link.springer.com/article/10.1007/BF00992696
- Schulman, J. et al. (2017). Proximal Policy Optimization Algorithms. arXiv:1707.06347
- Schulman, J. et al. (2016). High-Dimensional Continuous Control Using Generalized Advantage Estimation. ICLR 2016. arXiv:1506.02438
- CleanRL 官方仓库与文档:github.com/vwxyzjn/cleanrl、docs.cleanrl.dev(本页 PPO 结构参考了 CleanRL 的极简风格)
- Stable-Baselines3 官方文档:stable-baselines3.readthedocs.io(PPO 的成熟参考实现)