一句话定位
SimPLe(Simulated Policy Learning)直接在像素级学一个 Atari 的逐帧视频预测模型(4 帧堆叠 + 动作 → 下一帧 + 奖励,核心是一个新提出的离散隐变量随机模型),然后完全在这个学到的「模拟器」里用 PPO 训策略,只跟真实环境交互 100K 步(≈2 小时游戏时间)。在 26 个 Atari 游戏上,SimPLe 在样本效率上大多超过精调过的 Rainbow / PPO 基线(Freeway 上最多 10 倍),是第一个在 ALE 上成功证明「学到的模型 + model-based 控制」可以打过 model-free 方法的工作,正面回应了 Machado et al. (2018) 综述提出的公开挑战。论文 2019-03 挂 arXiv 首个预印本,最终发表于 ICLR 2020。
背景与定位
人类几分钟内就能学会 Atari 游戏,而当时最好的 model-free 深度 RL(DQN 系、Rainbow)要几千万到上亿步、相当于几周实时训练。作者假设人类的优势部分来自「对物理过程的直觉预测」,于是探索:能否学一个视频预测模型,让 agent 在脑内模拟器里训练,从而大幅降低与真实环境的交互次数?
此前 Oh et al. (2015)、Chiappa et al. (2017) 已证明 Atari 的像素级预测模型本身可以做到低误差、长时程稳定,Leibfried et al. (2016) 进一步加了奖励预测,但没有一篇把学到的模型真正用来训出能打游戏的策略。Machado et al. (2018) 综述(Section 7.2)因此把它列为公开挑战:「目前尚无在 ALE 上用学到的模型成功做规划的清晰证明」。SimPLe 是对这个挑战的正面回答。
与后续同年出现的 planet / dreamer-v1 / muzero 相比,SimPLe 走的是不同的技术路线:PlaNet/Dreamer 把观测压缩进紧凑连续隐状态(RSSM)后在隐空间想象/规划,MuZero 学一个只对规划有用、从不重建观测的抽象隐模型配 MCTS;而 SimPLe 直接建模逐帧像素(video prediction),用一个自回归离散隐变量描述随机性,模型输出就是下一帧真实分辨率的图像和奖励,再用标准 model-free 算法(PPO)在这个像素级模拟器里训练。这条路线继承自 world-models-ha-schmidhuber(VAE+RNN+进化控制器)和 Kaiser & Bengio (2018) 的离散自编码器,也直接对比、改进了 Babaeizadeh et al. (2017a) 的连续 VAE 式随机视频预测基线(SV2P)。SimPLe 论证了这条「像素级视频预测 + model-free 策略优化」路线在低数据、Atari 场景下的可行性,其数据聚合(每轮用当前策略在真实环境采数据、重训模型)思路直接借鉴经典 Dyna-Q,但此前在 Atari 领域从未取得可比结果。
模型架构
世界模型(Figure 2):输入为 4 帧堆叠的游戏画面(105×80×3,由 210×160 下采样 2 倍)+ 动作 one-hot 编码,输出下一帧预测图像(per-pixel 256 色 softmax)和奖励。
- 确定性主干:仿 Oh et al. (2015) 的卷积前馈网络。典型配置为 4 层 64-filter 卷积 + 3 层全连接(前两层 1024 单元);全连接层与一个 64 维可学习动作 embedding 拼接后逐通道相乘,实现动作条件化;随后 3 层 64-filter 反卷积,末层反卷积输出 105×80 原尺寸图像;下采样-上采样层间有 skip connection;卷积/反卷积前接 dropout 0.15,后接 layer normalization。
- 随机性——离散隐变量(本文核心创新):一个卷积推断网络(仅训练时用,2 层,kernel 8×8/stride 4×4,输出通道 128→512)近似给定未来帧的后验分布,采样得到隐变量后离散化为 bit(训练时反传绕过离散化,遵循 Kaiser & Bengio 2018),再训一个 LSTM 自回归地按 8 bit 一组预测这些 bit;推理时不再从先验采样,而是由该 LSTM 自回归生成隐 bit。作者认为这解决了此前 VAE 式随机模型(如 SV2P)的两个问题:KL 权重需按游戏调参、后验容易偏离先验导致推理时出现未见过的隐值。
- 损失函数:像素预测用 clipped L2(阈值 C=10)或 clipped 256 维 softmax 交叉熵(阈值 C=0.03,对应置信度超过 97% 时不再产生梯度);奖励预测为在最后一层全连接上接 softmax。
- Scheduled sampling:训练中按线性递增概率(到第一轮训练中段升到 100%)用模型自身前一步预测替换部分输入帧,缓解自回归复合误差。
- 确定性 / 确定性+循环 / 随机离散 三种变体都做了消融(见「评测」),随机离散(SD)综合表现最好。
- 规模:整个模型约 7400 万参数;隐 bit 预测器每步输出 128 bit(8 bit 一组)。
数据
- 交互预算:与真实 Atari 环境仅 100K agent steps(帧跳过 4 → 400K 帧,60 FPS 下约 114 分钟游戏时间)。因为循环开始前已采集部分数据,实际总交互数为
6400 × 16 = 102,400。 - 预处理:标准 Atari 流程——frame skip = 4(每个动作重复 4 次),画面下采样 2 倍到 105×80,4 帧堆叠作为观测以缓解部分可观测性。
- 游戏集合:从 ALE 中选取 26 个游戏,选择标准是「用 SimPLe 或 Rainbow 在 100K 交互下能取得非随机结果」。
- 数据采集是迭代式的(Algorithm 1):初始数据来自随机策略 rollout;此后循环 15 次,每次用当前策略在真实环境采 6400 步新数据加入缓冲区 D,再重训世界模型(自监督于观测、监督于奖励),再重训策略。真实数据除了训练世界模型外,也被直接混入 PPO 的训练(但由于模拟环境交互量高达 1500 万步而真实只有 10 万步,真实数据对策略训练的直接影响可忽略)。
- 随机性验证:额外做了 sticky-actions 实验(Machado et al. 2018 建议的协议),证明随机离散世界模型能在动作粘滞(stochastic)设定下学出与确定性设定接近的效果,且无需额外调参。
- 论文未涉及仿真到真实(sim-to-real)迁移问题——这里的「模拟器」本身就是从真实 Atari 观测中学出来的,不存在跨域 gap 的讨论。
训练方法
SimPLe 主循环(Algorithm 1),共迭代 15 次:
- 用当前策略 π 在真实环境 env 中采集数据,加入缓冲区 D;
- 用 D 有监督训练世界模型 env′(
TRAIN_SUPERVISED)——首轮训练 45K steps,此后每轮 15K steps(后续轮次更短是因为模型已捕获大部分游戏动态,只需扩展到新情形); - 完全在世界模型 env′ 内用 PPO(
TRAIN_RL)更新策略 π(γ=0.95,Schulman et al. 2017)。
PPO 策略训练细节:
- 为缓解模型复合误差,用短 rollout:每 N = 50 步(默认;消融了 N=25/100)从真实数据缓冲区 D 中均匀采样起始状态重启模拟环境(random starts,消融证明去掉这一步在 Seaquest 上效果显著变差);rollout 末步额外加入 value function 估计作为奖励,弥补短 rollout 无法反映更远期收益的问题。
- 每个 PPO epoch 用 16 个并行 agent,每个从模拟环境采 25/50/100 步(默认 50)。
- PPO epoch 数 =
z·1000,除最后一轮 z=3、第 8/12 轮 z=2 外均为 z=1,即每轮在模拟环境中产生800K·z步交互;整个训练过程中策略在模拟环境里累计交互约 1520 万步(15.2M),远超真实环境的 10 万步。 - 评测策略温度:用
softmax(logits(π)/T)采样动作,经验发现 T = 0.5 效果最好(比原策略更确定但不完全贪心)。 - 折扣因子/rollout 长度消融(Section 6.4):γ=0.95 略优于 γ=0.99;N=25 与 N=50 相当,N=100 因复合误差略差;世界模型训练时长加长 5 倍能进一步提升效果,但受算力限制其余消融都用短训练配置。
Infra(训练 / 推理工程)
- 世界模型规模与速度:约 7400 万参数;在单张 NVIDIA Tesla P100 上,batch=16 推理约 0.5s、batch=2 反传约 0.7s,折算约 32ms / 模拟帧;作为对比,真实 ALE 模拟器单步约 0.4ms(即学到的模拟器比真实模拟器慢约 80 倍)。
- 总训练规模(GPU 数、总 GPU-小时):论文未披露(未给出多游戏并行训练所用 GPU 总数或总训练小时数;仅在致谢中提及使用了波兰 PLGrid 高性能计算基础设施 ACK Cyfronet AGH / PCSS 的计算资源)。官方 GitHub 仓库(tensor2tensor/rl)说明「完整的 model-based 训练流程需要显著更长时间(几天到一周,取决于硬件和所用模型)」。
- 框架:基于 TensorFlow 的 Tensor2Tensor 库,开源在
tensorflow/tensor2tensor的tensor2tensor/rl子目录,提供trainer_model_based.py/trainer_model_free.py/evaluator.py/player.py等脚本,并提供约 180 组预训练策略+世界模型 checkpoint(Google Cloud Storagegs://tensor2tensor-checkpoints/modelrl_experiments/train_sd/,每游戏 5 个 run)。 - 推理/控制频率:论文未给出线上部署意义的 FPS/控制 Hz 指标(这是研究性 world model,非实时控制系统)。
评测 benchmark
主结果(26 个 Atari 游戏,100K 交互,Figure 3/4):SimPLe 在几乎所有游戏上比精调过的 Rainbow(基于 Dopamine 实现调参)更样本高效;在超过一半的游戏上,达到同等分数 Rainbow 所需交互数是 SimPLe 的 2 倍以上;最佳情形 Freeway 上超过 10 倍。与 PPO 基线相比优势更大。即便让 Rainbow/PPO 用 2 倍交互(200K),SimPLe 仍占优(Figure 4)。模型在除 Bank Heist 外的所有游戏上都优于随机策略;在 6 个游戏上 5 次运行中的最佳成绩超过人类平均水平(Avg. Human,引自 Pohlen et al. 2018 Table 3)。
与其他 model-based 基线的对比(论文原文给出的近似比值,因原作者未提供逐表数据):
- vs. Dyna-DQN(Holland et al. 2018):Asterix 约 330%、Q-Bert 约 120%、Seaquest 约 150%、Ms. Pac-Man 约 80%(random-normalized score @100K)。
- vs. GATS(Azizzadenesheli et al. 2018):Pong 约 64 倍、Breakout 约 10 倍。
交互预算扫描(Section 6.2,Figure 5):20K 交互下效果差;50K 已接近 100K 水平;此后持续提升到 500K(此时与 model-free PPO 打平)——即模型基方法的优势在低数据区间明显,随数据增多逐渐消失。用 SimPLe@100K 得到的策略作为 model-free PPO 的初始化也被验证有效(Figure 5b)。
架构消融(Table 1,7 种配置 × 26 游戏,按「最佳次数 / 至少达到中位数次数」计):
| 配置 | best(/26) | at least median(/26) |
|---|---|---|
| deterministic | 0 | 7 |
| det. recurrent | 3 | 13 |
| SD(随机离散) | 8 | 16 |
| SD, γ=0.9 | 1 | 14 |
| default | 10 | 21 |
| SD, N=100 | 0 | 14 |
| SD, N=25 | 4 | 19 |
即随机离散隐变量模型显著优于确定性/循环确定性变体,其中 default 行综合表现最好;论文另在同一节说明「用 5 倍长时间训练世界模型能取得本文最好结果,受算力所限其余消融都用较短训练」,但原文未明确逐行标注 Table 1 中 default 与 SD 两行具体在哪些超参上不同,此处不做过度解读。
后续复核的重要限定:论文正文(Introduction)说明,首个预印本发布后,van Hasselt et al. (2019) 与 Kielak (2020) 证明经过针对低数据区间调参的 Rainbow 也能达到相近水平——用改进后的 model-free 基线重新比较,两者在 26 个游戏中打平(各 13 胜)。
创新点与影响
贡献:
- 提出并验证了一种新的随机离散隐变量视频预测架构,在 Atari 逐帧预测任务上显著优于此前的确定性模型和连续 VAE 式随机模型(如 SV2P)。
- 首次在 ALE 上成功证明「用学到的(像素级)模型训练出的策略可以在真实环境中打过精调 model-free 基线」,正面回应 Machado et al. (2018) 提出的公开挑战。
- 把经典 Dyna-Q 式「真实数据采集 ↔ 模型训练 ↔ 模拟环境策略训练」交替循环,首次在 Atari 这一大规模视觉域证明可行、且样本效率显著优于当时 SOTA。
- 完整开源(Tensor2Tensor
rl子模块),提供约 180 组预训练 checkpoint,成为后续 model-based Atari 研究的常用基线与代码起点。
改变了什么:证明了低数据(100K 交互,约 2 小时游戏时间)Atari 学习在 model-based 范式下可行,把「model-based RL 在像素级视觉域打不过 model-free」这一长期认知推翻,直接推动了同年稍晚出现的 dreamer-v1、muzero 等 model-based 工作。
作者自陈的局限:
- 最终渐近分数总体仍低于最好的 model-free 方法(这是 model-based RL 的通病);
- 同一游戏不同训练 run 之间方差很大,原因可能是模型-策略-数据采集三者的复杂交互;
- 训练成本(在世界模型内训练策略)在计算和时间上都相当可观,呼吁未来研究更轻量的模型;
- 定性分析(Appendix B)指出模型难以捕捉大范围全局场景切换,例如 Private Eye 中角色在不同场景间瞬移的情形;
- 论文明确把「用模型做规划而非仅作为学习到的模拟器」「利用模型可微性将梯度信息注入 RL」「把预测模型学到的表征直接喂给策略」都列为未来方向。
原始链接
- 论文(arXiv 1903.00374,v1 2019-03,ICLR 2020):https://arxiv.org/abs/1903.00374
- PDF:https://arxiv.org/pdf/1903.00374
- 官方代码(Tensor2Tensor
rl子模块):https://github.com/tensorflow/tensor2tensor/tree/master/tensor2tensor/rl - 项目视频页(README 中引用):https://sites.google.com/corp/view/modelbasedrlatari/home
一手源存档(sources/)
- simple-atari—github-readme — Tensor2Tensor
rl子模块 README 快照(sources/world-model/2019/simple-atari—github-readme.md) - arXiv 原文 PDF(不入 git):https://arxiv.org/pdf/1903.00374