一句话定位

STORM 用 GPT 式 causal Transformer 替换 Dreamer 系列的 GRU 序列模型,同时保留 DreamerV2/V3 的 categorical-VAE 随机隐变量与”在想象中训 actor-critic”的整套训练法,在 Atari 100k 上拿到 126.7% mean human-normalized(不用 lookahead 搜索的方法中的新纪录),且单卡 RTX 3090 上 4.3 小时训完对应 1.85 小时真实游戏时长的样本量,NeurIPS 2023。

背景与定位

Model-based RL 靠世界模型的”想象”训练策略,但自回归展开的世界模型会累积预测误差,导致 agent 在想象里追逐现实中不存在的虚假目标;给隐状态引入随机性(VAE 式采样噪声)被证明能缓解这个问题。Dreamer 谱系(dreamer-v1 用 LSTM+连续想象、dreamer-v2 换成 GRU+categorical-VAE 离散隐状态、dreamer-v3 加一整套鲁棒性变换做到跨领域零调参)一直用 RNN 做序列模型;simple-atari(SimPLe)更早用 LSTM 世界模型直接在 Atari 上试过这条路线但采样效率一般。RNN 的循环结构限制并行训练速度。此前把 Transformer 接入世界模型的尝试——iris(VQ-VAE 把每帧编码成 4×4=16 个 token,spatial-temporal Transformer 建模,token 数太多拖慢训练)、TWM(Robine et al. 2023,用 Transformer-XL,把 observation/action/reward 当作等价的三种 token 类比 Decision Transformer)、transdreamer(直接把 Dreamer 的 GRU 换成 Transformer,但缺乏在标准 benchmark 上的验证)——要么没有真正发挥 Transformer 的效率优势训练更慢,要么性能没有超过 GRU 版的 DreamerV3。STORM 的定位就是把”强序列建模能力的 Transformer”和”VAE 式随机隐变量”结合到一起,同时保持效率。

模型架构

STORM 整体是 categorical-VAE 编码器 + GPT 式因果 Transformer 序列模型 + DreamerV2/V3 风格 actor-critic,三部分端到端联合训练。

图像编码器/解码器(categorical VAE)

  • 输入 oₜ 是 3×64×64 RGB 图像(未转灰度、未做帧堆叠),编码器:4 层卷积(kernel=4, stride=2, padding=1)+ BatchNorm + ReLU,通道 3→32→64→128→256,空间 64→32→16→8→4,flatten 得 4096 维,接一层 Linear 到 1024 维,reshape 成隐分布 Zₜ32 个类别变量、每个 32 类(与 DreamerV2/V3、TWM 一致的配置)。
  • Zₜ 采样得到 zₜ(32×32 one-hot),用 straight-through 直通梯度保留反传路径。
  • 解码器结构对称:Linear 1024→4096→reshape 256×4×4,4 层转置卷积(DeConv,kernel=4/stride=2/pad=1)+ BatchNorm + ReLU 还原到 3×64×64。

Action mixer(zₜ 与 aₜ 融合成单 token)eₜ = mφ(zₜ, aₜ),把 32×32 隐采样展平后与动作 one-hot(维度 A,Atari 各游戏 3~18 不等)拼接,过 Linear+LN+ReLU→Linear+LN 输出到 Transformer 特征维 D=512 的单个 token eₜ。这是与 TWM 的关键区别之一:TWM 把 observation、action、reward 当作三个独立同等地位的 token,STORM 把观测和动作先融合成一个 token 再喂进 Transformer,序列长度更短。

序列模型(GPT-like causal Transformer)h₁:T = fφ(e₁:T),标准 vanilla Transformer(非 Transformer-XL),自注意力用后续掩码(subsequent mask)保证 eₜ 只能看到 e₁,...,eₜ。默认配置仅 2 层(远小于 IRIS/TWM 的 10 层)、隐藏维 D=5128 个注意力头、dropout=0.1;位置编码是可学习的加性参数矩阵 w₁:T(非正弦),每个 Transformer block 是后置 LN 结构:MHSA→Linear+Dropout→残差→LN→FFN(隐藏维 2D,ReLU)→Linear+Dropout→残差→LN。推理(想象展开)阶段用 KV cache 加速自回归采样。

预测头(均为 MLP,输入取 hₜ:Dynamics predictor Ẑₜ₊₁=g_φᴰ(ẑₜ₊₁|hₜ)(1 层,输出 1024 维对应 32×32 隐分布);Reward predictor r̂ₜ=g_φᴿ(hₜ)(3 层,输出 255-bin symlog two-hot,沿用 DreamerV3 的做法);Continuation predictor ĉₜ=g_φᶜ(hₜ)(3 层,输出 1 维伯努利)。

Agent(在想象中训练的 actor-critic):agent 状态 sₜ=[zₜ,hₜ](隐采样与序列模型隐状态拼接,而非只用其一,消融见下);Actor πθ(aₜ|sₜ) 与 Critic Vψ(sₜ) 均为 3 层 MLP,Critic 输出 255-bin symlog two-hot 分布(distributional,做法与 DreamerV3 一致)。训练/推理时不是从后验 Zₜ 而是从先验 Ẑₜ 采样 zₜ 来推进想象。

与同类方法的架构对比(论文 Table 1):SimPLe 用 LSTM+Binary-VAE+PPO;TWM 用 Transformer-XL、三 token、Categorical-VAE、无历史信息、agent 用重建图像;IRIS 用 vanilla Transformer、4×4 latent、VQ-VAE、agent 也用重建图像;DreamerV3 用 GRU+Categorical-VAE、agent state 是 [latent, hidden];STORM 用 vanilla Transformer+Categorical-VAE、单 latent token、无历史信息(重建不依赖 hₜ)、agent state 同样是 [latent, hidden]。

数据

STORM 是纯在线 RL,无任何离线数据集或预训练,全部数据来自 agent 与环境的实时交互:

  • 交互 3 步循环:S1)用当前策略与真实环境交互采样,存入 FIFO replay buffer;S2)从 buffer 采轨迹更新世界模型;S3)从 buffer 采起点、用世界模型生成想象轨迹改进 actor-critic。
  • Benchmark 为 Atari 100k:26 款游戏,动作维度最多 18;100k “样本步”经过 4 倍帧跳(frame skip,取最后 2 帧的 max)对应 400k 实际游戏帧,约合 1.85 小时真实游戏时长
  • 世界模型训练:每次采 B1=16 条长度 T=64 的轨迹;想象展开:B2=1024 条轨迹,用 C=8 步上下文起始、想象展望 L=16 步。
  • 环境细节:生命值信号(life info)纳入 done 信号但环境持续跑到真正 reset 才停(对齐 IRIS 设置);不做灰度化、不做帧堆叠。
  • 每采 1 个环境步,世界模型和 agent 各更新 1 次(TrainDynamicsEverySteps=1, TrainAgentEverySteps=1)。
  • 训练用 5 个随机种子,每 2500 个样本步存一次 checkpoint,每个 checkpoint 跑 20 个评估回合取平均,论文报告的是各 seed 最终 checkpoint 分数的均值(Table 2)。
  • 示范轨迹(demonstration trajectory)消融:对 MsPacman/Pong/Freeway 三个探索困难的游戏,各加入 1 条用预训练 DQN agent(Gogianu et al., “Atari agents”)采集的轨迹放进 replay buffer(Table 11:MsPacman return 5860 / 1612×4 帧;Pong return 18 / 2079×4 帧;Freeway return 27 / 2048×4 帧)。主结果(126.7% mean、58.4% median)包含 Freeway 的示范轨迹;论文同时报告不用该轨迹的公平对比结果——mean human-normalized 降到 122.3%(Table 2 同时给出 “Freeway w/o traj” 一行)。

训练方法

世界模型损失(式 3,端到端自监督,β1=0.5, β2=0.1):

L(φ) = (1/BT) Σ [ L_rec + L_rew + L_con + β1·L_dyn + β2·L_rep ]
  • L_rec = ||ôₜ-oₜ||²:图像重建 MSE(直接对编码器输出 zₜ 而非序列模型输出的 ẑₜ 做重建——消融证实这一点很关键,“Decoder at rear” 变体若改用 ẑₜ 重建会掉点);
  • L_rew:symlog two-hot 分类损失(继承自 DreamerV3,把回归转成分类避免不同环境间损失尺度不一致);
  • L_con:连续标志的二元交叉熵;
  • L_dyn = max(1, KL(sg(qφ(zₜ₊₁|oₜ₊₁)) ‖ g_φᴰ(ẑₜ₊₁|hₜ))):动力学损失(KL 平衡里带 free-bits 的 max(1,·)),只更新序列模型侧;
  • L_rep = max(1, KL(qφ(zₜ₊₁|oₜ₊₁) ‖ sg(g_φᴰ(ẑₜ₊₁|hₜ)))):表示损失,让编码器输出被序列模型预测弱引导,二者通过 stop-gradient 分离更新方向(与 DreamerV2/V3 的 KL balancing 思路一致)。

Actor-Critic 损失(沿用 DreamerV3 的 actor 训练设置,式 7-10):

  • λ-return 递归定义 Gλₜ = rₜ + γcₜ[(1-λ)Vψ(sₜ₊₁) + λGλₜ₊₁]γ=0.985, λ=0.95
  • Actor 损失用百分位归一化 S = percentile(Gλₜ,95) - percentile(Gλₜ,5) 缩放 advantage,外加熵正则(系数 η=3×10⁻⁴);
  • Critic 损失 = 对 λ-return 的回归 + 对 critic 自身 EMA 副本的正则(EMA 衰减 σ=0.98),稳定训练防止过拟合。

优化器与关键超参(论文 Table 10,与 GitHub config_files/STORM.yaml 核对一致):Adam;世界模型学习率 1×10⁻⁴、梯度裁剪 1000;actor-critic 学习率 3×10⁻⁵、梯度裁剪 100;Transformer 2 层、D=512、8 头、dropout 0.1。

消融研究要点(Section 5):

  • 模块顺序:把重建损失挪到序列模型输出后(“Decoder at rear”)或把 reward/continuation 头挪到 zₜ 之前(“Predictor at front”)都会掉点,尤其后者在需要多帧上下文推断 reward 的游戏(如 Ms. Pacman)上明显下降,而单帧可判断 reward 的游戏(如 Pong)影响很小。
  • Transformer 层数:从默认 2 层增到 4/6 层没有带来性能提升——作者归因于三点:(1) 相邻帧差异小+残差连接使预测本身不需要复杂模型;(2) Atari 100k 数据量和领域多样性不足以喂饱更大模型;(3) 端到端训练下 L_rep 会让编码器被过大的序列模型过度牵引。
  • Agent state 选择:sₜ=[zₜ,hₜ] 优于只用 hₜ(在 Ms. Pacman 这类需要长上下文的环境里 hₜ 帮助明显)或只用 zₜ(在 Pong 这种非平稳、模型不准的环境里,纯 hₜ 会出现类似灾难性遗忘的行为,引入 zₜ 的随机性有帮助)。

Infra(训练 / 推理工程)

  • 训练硬件:单张 NVIDIA GeForce RTX 3090(论文强调”仅需单卡”),高频 CPU(作者用 Intel i9-11900K)配合,避免 GPU idle;训练全程使用 bfloat16 混合精度加速(README:V100 等不支持 bf16 的设备需手动改回 float16 并调整 mask 填充值防止溢出;A100 上 bf16 反而可能更慢,需切换 use_amp)。
  • 训练时间:完成对应 1.85 小时真实游戏时长(100k 样本步)的训练,单卡 RTX 3090 上耗时 4.3 小时;论文另外把 STORM 直接在 V100 上实测,耗时 9.3 V100-小时(Table 12 标注该数字无星号=实测,非外推)。
  • 横向对比(Table 12,V100-小时,标 * 为按 DreamerV3 论文的外推方法换算——P100 记 2 倍慢、A100 记 2 倍快):SimPLe(P100 20 天 → 外推 240*);TWM(A100 10 小时 或 RTX 3090 12.5 小时 → 外推 20*);IRIS(A100 两次跑共 7 天 → 外推 168*);DreamerV3(V100 实测 12 小时);STORM(V100 实测 9.3 小时,或 RTX 3090 实测 4.3 小时)——是五者中训练成本最低的。
  • 推理/想象加速:想象展开时用 KV cache 加速 Transformer 的自回归采样。
  • 论文 Figure 1 另给出”V100 上训练 FPS”条形图对比(含 STORM/DreamerV3/IRIS/TWM/SimPLe),但受 PDF 图表文本抽取顺序影响,本页未能把具体数值与方法逐一对应,故不在此引用该图读数,仅采信 Table 12 有明确表头标注的小时数。

评测 benchmark

Atari 100k,26 款游戏,Table 2(论文核心结果表)

指标RandomHumanSimPLeTWMIRISDreamerV3STORM
Human Mean0%100%33%96%105%112%126.7%
Human Median0%100%13%51%29%49%58.4%
  • STORM 是”不使用 lookahead 搜索(MCTS 等)“的方法中的新纪录(论文明确不与 MuZero/EfficientZero/SpeedyZero 等带搜索的方法直接比较,因为搜索可以叠加在世界模型之上,不是本文研究对象)。
  • 去掉 Freeway 示范轨迹后的公平对比(fair-comparison ablation):mean human-normalized 降为 122.3%(对应表中 “Freeway w/o traj” 行),仍全面超过此前方法。
  • 具体游戏层面的定性观察:STORM 在目标物体大或多个目标同时存在的游戏(Amidar、Ms Pacman、Chopper Command、Gopher)上明显优于此前方法——归因于注意力机制能显式保留移动物体的历史轨迹,便于推断速度和方向,这是 RNN 方法难以做到的;但在单个小物体的游戏(Pong、Breakout)上表现不如预期,作者认为原因在于自编码器本身对小物体的重建能力有限,且随机采样噪声可能在这类场景下过度干扰注意力权重。
  • 附录 F 给出多步想象可视化(context 8 帧 + 自回归想象 56 帧)在 Boxing / ChopperCommand / MsPacman / Pong / RoadRunner 上与真实轨迹的定性对比。

创新点与影响

  • 首次证明 Transformer 序列模型能在效率和性能上同时超过 GRU 版 Dreamer:此前 IRIS、TWM、TransDreamer 引入 Transformer 都未能同时兼顾效率与效果,STORM 通过”单 latent token(而非 IRIS 的 4×4=16 token 或 TWM 的三类型 token)+ 仅 2 层 vanilla Transformer + 因果掩码”的极简设计做到了两者兼得。
  • 验证了小 Transformer 在低数据领域的合理边界:层数消融表明堆叠更多 Transformer 层在 Atari 100k 这种数据量小、领域单一的场景下不但无益反而可能有害,这与 Transformer 在 NLP/CV 领域”越大越好”的直觉相反,是对”Transformer scaling 是否普适”的一个反例数据点。
  • 示范轨迹作为探索问题的轻量解法:用单条离线示范轨迹注入 replay buffer 就能显著改善 Freeway 这类探索困难环境的表现,为稀疏奖励探索提供了一个比专门设计好奇心驱动探索更简单的备选方案。
  • 论文自陈的局限性:(1) 世界模型端到端联合训练,编码器需要预测自己序列模型的输出,引入额外的非平稳性,可能限制模型的可扩展性;(2) 想象展开的起点从 replay buffer 均匀采样,而 agent 用 on-policy actor-critic 训练,策略梯度公式中理论要求的 on-policy 状态分布 μ(s) 并未被显式建模。
  • 后续官方维护的实现迁移到了新仓库 OC-STORM(作者在 README 中声明本仓库已停止维护)。

原始链接

一手源存档(sources/)

  • storm—github-readme(GitHub README 全文 + config_files/STORM.yaml 超参数配置)
  • arXiv 原文 PDF(2310.09615,NeurIPS 2023),不入 git