一句话定位
TransDreamer 把 dreamer-v2 的世界模型骨架从 RNN(RSSM)换成第一个基于 Transformer 的随机状态空间模型——Transformer State-Space Model(TSSM)——并让这个世界模型与策略网络共享,用于需要长程记忆的部分可观察任务;在自建的 2D/3D「Hidden Order Discovery」记忆任务上明显超过 DreamerV2(如 2D 4-Ball 成功率 23% vs 7%),在不需要长程记忆的 DMC/Atari 短任务上则与 DreamerV2 打平(收敛更慢)。
背景与定位
Dreamer 系列(planet → dreamer-v1 → dreamer-v2)证明了「在 RSSM 学到的隐空间里做想象式 actor-critic 训练」这条 model-based RL 路线的有效性,但 RSSM 的确定性状态更新 hₜ = GRU(hₜ₋₁, zₜ₋₁, aₜ₋₁) 天生继承 RNN 的两个限制:只能通过压缩后的 hₜ 间接访问历史,且更新必须逐步串行、无法并行训练。与此同时 Transformer 在 NLP/CV 中已证明在长程依赖和「直接访问历史、做复杂交互」的记忆式推理任务上优于 RNN(Ritter et al. 2020; Banino et al. 2020),但先前工作(GTrXL, Parisotto et al. 2020)也表明「把 Transformer 直接套进 RL 策略网络」训练很不稳定,需要额外的 GRU 门控层才能收敛。
TransDreamer 要回答的问题是:能否设计一个同时满足随机性、可直接访问历史状态、可并行训练、又能在测试时逐步 rollout 想象这四个条件的 Transformer 世界模型,并证明它在需要长程记忆推理的任务上确实优于 RSSM。论文强调,简单地把 RNN 换成 Transformer 并不可行——RSSM 的后验表征模型 q(zₜ|hₜ,xₜ) 会把 Transformer 的输出(hₜ)又喂回作为下一步输入,破坏了 Transformer 训练所需的并行性,这是本文要解决的核心架构难题。
模型架构
TransDreamer 沿用 Dreamer 的三段式框架(世界模型学习 → 想象轨迹上的策略学习 → 环境交互采集数据),核心改动是把 RSSM 换成 Transformer State-Space Model(TSSM)。
TSSM 各部件(对照 RSSM,论文 Table 1):
- 确定性状态更新:RSSM 是
hₜ = GRU(hₜ₋₁, zₜ₋₁, aₜ₋₁)(逐步串行);TSSM 改为hₜ = Transformer(z₁:ₜ₋₁, a₁:ₜ₋₁)——Transformer 一次前向直接吃全部历史随机状态与动作序列,输出全部时间步的 hₜ,可并行计算。 - 表征模型(后验,关键改动——“myopic representation model”):RSSM 是
zₜ ~ q(zₜ|hₜ,xₜ)(依赖 hₜ);TSSM 去掉对 hₜ 的依赖,改为zₜ ~ q(zₜ|xₜ)——因为 Transformer 训练不允许把自己的输出(hₜ)当作下一步输入,否则破坏并行性。论文论证:完整模型状态是随机 zₜ 与确定性 hₜ 的拼接,时序信息已经由确定性通路携带,因此表征模型不编码历史信息的假设「可能」不影响性能;论文用一个简化版 Dreamer(同样把q(zₜ|xₜ,hₜ)换成q(zₜ|xₜ))做对照实验,验证了这一假设。 - 随机状态先验 / 图像 / 奖励 / 折扣预测器:形式与 RSSM 一致(
p(ẑₜ|hₜ)、p(x̂ₜ|hₜ,zₜ)、p(r̂ₜ|hₜ,zₜ)、p(γ̂ₜ|hₜ,zₜ)),只是条件变量 hₜ 来自 Transformer 而非 GRU。 - 想象(imagination):测试 / 训练想象阶段用先验
ẑₜ ~ p(ẑ|hₜ)作为 Transformer 的输入自回归生成未来状态序列,与真实 rollout 保持一致的接口。
策略学习:完全沿用 Dreamer 的 actor-critic-on-imagined-trajectories 框架(REINFORCE + 通过可微世界模型的动力学反传的混合梯度),世界模型参数在训练 agent 时冻结。论文特别指出,由于 TSSM 只训练来预测图像/奖励/折扣(而非直接用策略梯度训练 Transformer 本身),TransDreamer 不需要 GTrXL 那样的额外门控层就没有遇到 Transformer-in-RL 常见的训练不稳定问题。
训练稳定性 trick:可选优先经验回放(prioritized replay)——以概率 α 只从有非零奖励的轨迹采样、其余 (1−α) 均匀采样全缓冲区,用于在奖励稀疏时更快学好奖励预测器(Hidden Order Discovery 实验中 α=0.5)。
想象轨迹数(K):因 Transformer 显存开销远大于 RNN,无法像 Dreamer 那样对回放批次里的每个采样状态都生成想象轨迹,而是随机挑选一个更小的子集(大小 K)来生成想象——DMC/Atari 上 K=3;Hidden Order Discovery 任务上每次只随机采样 1 个起始状态、想象到该 episode 的最大步数。
配置数字(Appendix A.3, A.4):
- DMC/Atari:2 层 Transformer,无 dropout / 门控 / identity map reordering(Atari Pong 例外,用 identity map reordering 效果更好);Multihead Attention 10 个头;隐藏状态维度 DMC=200,Atari=600(与对应 DreamerV2 确定性状态维度对齐)。
- Hidden Order Discovery(2D/3D):6 层 Transformer,带 identity map reordering;发现把 attention 各层的中间输出拼接起来作为 hₜ 能加速收敛。
- 图像分辨率、CNN encoder/decoder 结构等继承 DreamerV2 的实现(未在正文单独列出新数字)。
数据
TransDreamer 是在线 model-based RL,无预训练语料,「数据」即 agent 与环境交互产生、存入回放缓冲区的轨迹:
- 自建 Hidden Order Discovery 任务(受 NumPad 任务 [Humplik et al. 2019; Parisotto et al. 2020] 启发):
- 2D 版(Minigrid 框架,俯视图):8×8 网格,4/5/6 个彩色球,每球最小间距 2 格;每 episode 最多 100 步;隐藏的球收集顺序随机化,收错则重置球位置(但不重置 agent 位置和隐藏顺序)。
- 3D 版(Unity 框架,第一人称视角,更强的部分可观察性):4-Ball Dense / 4-Ball Sparse / 5-Ball Dense 三种配置;sparse 设置球间距为 dense 的约 2 倍(球间欧氏距离 ≥4 个球径 vs ≥2 个球径)。
- 奖励设计:按正确顺序收集一球得 +3;收错则该轮内已收集球的奖励清零并重置地图(防止 agent 靠反复碰第一个球刷分)。
- 短程记忆对照任务:DeepMind Control Suite(dreamer-v1 的配置)与 Atari(DreamerV2 的配置),选取几个不需要长程记忆的子任务作为「不应该输」的 sanity check。
- 世界模型独立评测数据:为公平比较图像生成 MSE 和奖励预测准确率,TSSM 与 RSSM 分别用各自 agent 收集的轨迹单独训练(不联合策略训练),在 3D 5-Ball Dense 配置上用长度 100 步的轨迹、以不同 context 长度(60/70/80 步)评估剩余步的生成质量。
- 论文未披露总环境步数 / episode 数等聚合数据规模的具体数字(只在 Appendix A.4 给出每个任务 100 步上限、DMC/Atari 沿用各自原论文的训练步数)。
训练方法
- 目标函数:与 RSSM 相同形式的 ELBO(论文 Appendix A.2 给出完整推导),把表征后验从
∏_t qφ(zₜ|z₁:ₜ₋₁,xₜ)简化为∏_t qφ(zₜ|xₜ);损失包含图像/奖励/折扣的负对数似然项,加权系数 ηₓ、η_r、η_γ,以及后验-先验的 KL 项。 - 策略学习:与 Dreamer 相同,用想象轨迹的价值估计 V(sₜ) 构造 actor 目标,通过 REINFORCE 和/或对可微世界模型的动力学反传取梯度;critic 用时序差分学习拟合 V(sₜ)。
- DMC/Atari 配置几乎照搬 dreamer-v1/dreamer-v2 官方超参(action repeat、每 5 步训练一次世界模型和策略等),唯一必须修改的是想象轨迹数 K(因显存限制,3 条/样本)。
- Hidden Order Discovery 配置沿用 DreamerV2 的 Crafter 超参,额外加入优先回放(α=0.5)。
- 未见蒸馏、离散动作 tokenization 或额外的多阶段预训练/微调流程——训练范式即标准 Dreamer 循环(世界模型学习 ↔ 策略学习 ↔ 环境交互)。
Infra(训练 / 推理工程)
论文未披露 GPU 型号、GPU 数量、GPU-hours、并行策略或训练精度(无 DreamerV2 式的 accelerator-days 对比表)。也未披露推理 FPS / control-Hz / 延迟等部署指标。可确定的工程约束是:Transformer 相比 RNN 显存开销更大,导致必须限制并行想象轨迹数 K(DMC/Atari 上 K=3,而 Dreamer 对回放批次中每个采样状态都想象)——这是论文唯一明确讨论的资源权衡,其余基础设施细节均未披露。
评测 benchmark
Hidden Order Discovery 成功率(Table 3,完整收集顺序至少一次的轨迹占比,1000 条轨迹统计):
| 环境 | TransDreamer | DreamerV2 |
|---|---|---|
| 2D 4-Ball | 23% | 7% |
| 2D 5-Ball | 5% | 0% |
| 2D 6-Ball | 1% | 0% |
| 3D 4-Ball Dense | 18% | 10% |
| 3D 4-Ball Sparse | 11% | 1% |
| 3D 5-Ball Dense | 4% | 0% |
2D 4-Ball 上 TransDreamer 平均 episode 奖励约 7(对应平均正确收集 2 球以上),DreamerV2 约 4(约 1 球)。
世界模型图像生成前景 MSE(Table 2a,越低越好,3D 任务):
| 任务 / Context 步数 | TransDreamer (60/70/80) | DreamerV2 (60/70/80) |
|---|---|---|
| 4-Ball Dense | 211.2 / 133.1 / 69.8 | 281.9 / 194.2 / 110.8 |
| 4-Ball Sparse | 195.5 / 115.2 / 56.8 | 215.8 / 138.6 / 72.4 |
| 5-Ball Dense | 245.2 / 163.1 / 85.0 | 300.9 / 217.0 / 124.9 |
前景 MSE 占整体 MSE 差距的 60% 以上(尽管球只占图像很小比例)。
世界模型非零奖励(+3)预测准确率(Table 2b,越高越好):
| 任务 / Context 步数 | TransDreamer (60/70/80) | DreamerV2 (60/70/80) |
|---|---|---|
| 4-Ball Dense | 46.9 / 53.2 / 73.2 | 28.2 / 34.6 / 50.5 |
| 4-Ball Sparse | 32.4 / 36.5 / 48.6 | 32.0 / 33.3 / 42.3 |
| 5-Ball Dense | 17.7 / 18.1 / 32.35 | 9.8 / 6.2 / 15.3 |
零奖励预测准确率两者都很高(≈91.1–96.6%,差距小,见 Appendix Table 5)。
短程记忆对照(DMC + Atari 子集):TransDreamer 最终收敛到与 DreamerV2 相近的回报,但收敛速度普遍更慢(符合论文预期,因为这些任务不需要长程记忆,RNN 的近期偏置反而是优势),例外是 DMC Cheetah Run 上 TransDreamer 收敛略快且性能略优。
论文没有报告与非 Dreamer 系基线(如 Decision Transformer、GTrXL 策略、MuZero)的直接对比,只对比 TransDreamer vs DreamerV2。
创新点与影响
贡献:① 提出 Transformer State-Space Model(TSSM)——第一个满足「直接访问历史 + 可并行训练 + 可自回归 rollout + 保留随机隐变量」四个要求的 Transformer 世界模型,核心技巧是把表征后验从 q(zₜ|hₜ,xₜ) 简化为 q(zₜ|xₜ)(myopic representation model)以打破训练时的循环依赖;② 提出 TransDreamer——把 TSSM 世界模型与策略网络共享的完整 MBRL agent,是第一个基于 Transformer 的 model-based RL agent;③ 发现 Transformer 世界模型不需要 GTrXL 式的门控层就能稳定训练(因为世界模型只用图像/奖励/折扣信号训练,而非直接用策略梯度训练 Transformer);④ 定量证明 Transformer 世界模型在长程记忆任务上比 RSSM 生成更准的图像和奖励预测,且这一优势随想象步数增加而扩大。
影响:TransDreamer 是把 Transformer 引入 Dreamer 式隐空间想象框架的早期尝试之一,验证了「世界模型的历史访问方式」是长程记忆任务的关键瓶颈,为后续将 Transformer/attention 机制引入世界模型骨架(如后续基于 Transformer 或 diffusion 的世界模型工作)提供了一个直接的 RSSM→Transformer 替换范式和自建的长程记忆评测任务(Hidden Order Discovery)。
作者自述局限:① 不确定「表征模型忽略历史」这一假设在更复杂任务上是否依然成立;② 由于显存限制无法对回放批次中每个状态都生成想象轨迹,只能采样子集,这削弱了每次迭代的样本利用率;③ 短程记忆任务上收敛慢于 DreamerV2;④ 未在更复杂任务(如作者提到的 Crafter)上验证;⑤ 未处理需要长程记忆同时需要良好探索策略的 Atari 游戏(论文明确避开了这类游戏,因为不解决探索问题)。
原始链接
- arXiv abstract:https://arxiv.org/abs/2202.09481
- arXiv PDF:https://arxiv.org/pdf/2202.09481
一手源存档(sources/)
- 未发现官方 GitHub / 项目主页 / 博客(论文正文与附录均未提供代码或项目链接,仅引用了 danijar/dreamerv2 与 danijar/crafter 作为基线配置参考,非本工作自身发布)。
- arXiv 原文 PDF(2202.09481,arXiv 原文 PDF,不入 git)——见上方 arXiv 链接