一句话定位
TWM 把世界模型的骨架从 RNN/RSSM 换成一个 Transformer-XL:把隐状态、动作、奖励三种模态各自线性嵌入后拼成 token 序列喂给 Transformer-XL 做自回归预测,但策略只吃隐状态 z(不吃 Transformer 的隐藏输出 h),因此推理时完全不需要跑 Transformer;在 Atari 100k(每局仅 10 万步交互)上以 5 组随机种子取得 人类归一化均值 0.956 / 中位数 0.505,全面超过 SimPLe、DER、CURL、DrQ(ε)、SPR 五个基线;ICLR 2023。
背景与定位
「在想象中学策略」这条 model-based RL 路线由 world-models-ha-schmidhuber 开创,simple-atari 首次把它做到 Atari 100k 上有竞争力,dreamer-v2 用 RSSM(确定性 GRU 状态 + 随机类别隐变量)在非样本高效的 2 亿帧设定下取得当时最强效果。RSSM 的根本限制是历史只能通过压缩后的确定性状态 hₜ 间接访问,且必须逐步串行更新。同期 transdreamer(Chen et al. 2022)已尝试把 RSSM 换成 Transformer State-Space Model,但其 Transformer 输出仍要反馈进策略网络,推理期仍需要跑 Transformer,且未在 Atari 100k 上评测。
TWM 的定位是:用 Transformer-XL 直接对隐状态/动作/奖励序列做自回归建模,同时让策略只依赖压缩的隐状态 z 而非 Transformer 的输出 h,从而在保留「直接访问历史」这一 Transformer 优势的同时,把推理开销压回model-free 方法的水平。论文强调这与 transdreamer 的关键区别:TransDreamer 的策略依赖 Transformer 输出、推理期计算成本更高;TWM 的世界模型只在训练/想象阶段用到 Transformer。与 TWM 同期独立的 IRIS(iris,Micheli et al. 2022)则走另一条路——把整帧图像离散 token 化后用 GPT 式 Transformer 建模,在 Atari 100k 上取得比 TWM 更高的分数(人类归一化均值 1.046),是当时的新 SOTA;TWM 论文正文承认这一点但未与 IRIS 做直接对比实验。范式上,TWM 与 muzero / efficientzero 这类基于 MCTS 的 lookahead search 路线正交——TWM 决策时不做搜索。
模型架构
观测模型(Observation Model):沿用 dreamer-v2 的 CNN 架构(略作修改)实现一个变分自编码器,把观测 oₜ 编码成离散隐状态 zₜ(32 个类别变量 × 32 个类别),解码器只用于给 zₜ 提供学习信号。单个观测 oₜ 本身就是 4 帧堆叠(frame stacking),因此观测模型隐含短程时序信息,这一点与 DreamerV2 不同。
自回归动力学模型(Autoregressive Dynamics Model):核心骨架是一个因果掩码的 Transformer-XL(Dai et al. 2019),聚合函数 hₜ = f(z_{t-ℓ:t}, a_{t-ℓ:t}, r_{t-ℓ:t-1}) 直接基于过去 ℓ 步的隐状态、动作、奖励历史计算确定性隐藏状态 hₜ。三种模态(隐状态/动作/奖励)先各自过 modality-specific 线性嵌入再送入 Transformer;输入 token 数为 3ℓ-1(因为最后一步奖励不作为输入)。只取动作模态位置的 Transformer 输出作为 hₜ,隐状态与奖励模态的输出被丢弃。奖励预测器 r̂ₜ、折扣预测器 γ̂ₜ、下一隐状态预测器 ẑ_{t+1} 均为条件于 hₜ 的 MLP。关键设计:预测出的奖励会被回填进下一步的 Transformer 输入(“fed back into the transformer”),让模型能感知自己已经产生过的奖励,消融实验证明这一设计能显著提升性能。
策略与推理解耦:策略 πθ(aₜ|zₜ) 与 critic vξ(zₜ) 只吃隐状态 z,不吃 Transformer 输出 h;训练想象阶段用预测隐状态 ẑₜ(不做 reconstruction、不做 frame stacking),真实环境推理阶段用编码隐状态 zₜ + frame stacking 补足短程信息。这使得策略在真实环境部署时完全不需要运行 Transformer。
Transformer-XL 配置数字(Table 4/5):嵌入维度 256,10 层,4 个注意力头(每头 64 维,4×64=256),前馈层维度 1024。隐状态预测器 / 奖励预测器 / 折扣预测器 / actor / critic 的 MLP 隐藏层分别为 4×512 / 4×256 / 4×256 / 4×512 / 4×512,激活函数 SiLU。参数量(Table 5):观测模型 8.2M,动力学模型 10.8M,世界模型合计 19M;actor 1.3M + critic 1.3M = actor-critic 合计 2.6M;总计 21.6M;真实环境推理时只需要 encoder + actor,共 4.4M 参数。
数据
TWM 是纯在线 model-based RL,无离线预训练语料,「数据」即 agent 与环境交互产生并存入回放缓冲区 D 的轨迹 (o,a,r,d):
- Atari 100k 基准(Kaiser et al. 2020 提出):Arcade Learning Environment 中 26 个游戏子集,每局限制 10 万次交互,对应帧跳(frame skip=4)后 40 万帧,约合 2 小时人类游戏时间——是常规 Atari 训练(2 亿帧,Mnih et al. 2015 / Hafner et al. 2021)的 1/500。
- 环境预处理(Table 4):帧下采样至 64×64、灰度化、帧堆叠 4、Terminate-on-life-loss=是、单局最大帧数 108K、随机 no-op 最多 30 步。
- 平衡数据集采样(Balanced Dataset Sampling):因为数据集随训练缓慢增长,均匀采样会过度偏向早期经验、易过拟合。TWM 维护每条数据的采样计数 v₁…v_T,用 softmax(-v/τ) 转成采样概率(τ=20),让新数据被更频繁采样;τ→∞ 退化为均匀采样。消融(Figure 7,BankHeist/Breakout/Boxing/KungFuMaster/MsPacman/Pong)证明该策略显著提升人类归一化分数。
- 每次实验:世界模型批大小 N=100 条长度 ℓ=16 的序列;想象阶段从 N×ℓ 个观测中选出 M=400 个作为想象起点,生成长度 H=15 的想象轨迹用于策略训练。
- 5 组随机种子/游戏,训练结束后每个 run 评测 100 个 episode 取平均分。
训练方法
目标函数:把 DreamerV2 的平衡 KL 散度损失重新推导为平衡交叉熵损失(附录 A.2 给出与原 KL 目标梯度等价的证明),分离出观测模型损失 L_Obs(解码器负对数似然 + 熵正则化 α₁·H(qφ) + 一致性损失 α₂·H(qφ,pψ))与动力学模型损失 L_Dyn(隐状态交叉熵 + β₁·奖励负对数似然 + β₂·折扣负对数似然)。这样做的好处是可以独立调节交叉熵/熵项的相对权重,而不像原始平衡 KL 只有一个 λ。超参(Table 4):α₁=5.0, α₂=0.01, β₁=10.0, β₂=50.0。
策略学习:标准 advantage actor-critic(Mnih et al. 2016),用 Generalized Advantage Estimation(λ=0.95)计算优势,但用世界模型预测的折扣 γ̂ₜ(而非固定 γ)逐步加权,并按折扣因子累积乘积对 actor/critic 损失加权以软性处理 episode 终止(沿用 DreamerV2 做法)。
阈值化熵损失(Thresholded Entropy Loss):对策略熵做归一化并设阈值 Γ,只有当归一化熵 H(πθ)/ln(m) 低于阈值 Γ=0.1 时才惩罚(η=0.01),使得可以用同一套超参在不同动作数的游戏间保持稳定探索率,无需 ε-greedy 或调温度。消融(附录 Figure 15)显示不设阈值时熵容易崩溃或发散,得分更低。
其余超参(Table 4):折扣 γ=0.99;观测/动力学/actor 学习率均为 1e-4,critic 学习率 1e-5;均用 Adam。训练循环(Algorithm 1)为标准三段式:采集真实经验 → 用回放数据更新世界模型 → 用世界模型生成的想象轨迹更新策略,不涉及蒸馏或额外的离线预训练/微调阶段。
Infra(训练 / 推理工程)
- 训练硬件:单卡训练,单个 NVIDIA A100 GPU 约 10 小时完成一次完整训练+评测预算(预算按更新步数固定,故耗时略有波动);同等预算在 NVIDIA GeForce RTX 3090 上需 12–13 小时;若把 Transformer-XL 换成不带记忆机制的 vanilla Transformer,在 A100 上耗时升至 15.5 小时(约 1.5 倍),体现 Transformer-XL 缓存机制的必要性。
- 与其他方法的运行时对比(Table 3,统一在单张 NVIDIA P100 GPU 上测量,数字取自 Schwarzer et al. 2021):SimPLe 500 小时,TWM 23.3 小时(比 SimPLe 快 20 倍以上),SPR(带数据增强)4.6 小时 / (不带)3.0 小时,DER/DrQ(带增强)2.1 小时 / (不带)1.4 小时——TWM 比 model-free 方法慢,但比 SimPLe 快很多。
- 吞吐量(A100,论文自身 batch size 下测得,约等于值):世界模型训练 16,800 samples/s;世界模型想象——Transformer-XL 版 39,000 samples/s,vanilla Transformer 版 19,900 samples/s(XL 版快近 2 倍,验证记忆机制对速度的贡献);策略训练 700,000 samples/s。
- 真实环境推理速度(CPU,batch size=1):策略仅接受隐状态 z 输入时 653 帧/秒;若改为接受 [z,h](需要跑 Transformer)则降到 213 帧/秒,慢约 3 倍——这也是论文选择「策略只吃 z」这一设计的直接工程动机。
- 精度(fp16/bf16/混合精度)、并行策略(是否用了数据/模型并行)论文未披露。
评测 benchmark
Atari 100k 主结果(Table 1,5 组种子,每 run 结束后 100 episode 均值;人类归一化分数 = (agent-random)/(human-random);基线 DER/CURL/DrQ(ε)/SPR 分数取自 Agarwal et al. 2021 的 100-run 重新评测,SimPLe 为 5-run):
| 方法 | 人类归一化 Mean | 人类归一化 Median |
|---|---|---|
| DER | 0.350 | 0.189 |
| CURL | 0.261 | 0.092 |
| DrQ(ε) | 0.465 | 0.313 |
| SPR | 0.616 | 0.396 |
| SimPLe | 0.332 | 0.134 |
| TWM(本文) | 0.956 | 0.505 |
论文按 Agarwal et al. (2021) 的建议额外报告了 median / IQM / mean / optimality gap 四个带 95% 分层 bootstrap 置信区间的聚合指标(Figure 3)和性能分布曲线(Figure 4),结论一致:TWM 在全部四个聚合指标上显著优于五个基线,optimality gap 更接近零(具体区间值绘制在图中,正文未给出可直接摘录的数字表)。
样本效率(Table 2,TWM 自身在不同交互步数下的均值人类归一化分数):5K=0.007,10K=0.133,25K=0.408,50K=0.624,75K=0.832,100K=0.956。对照基线终值:SimPLe=0.332(TWM 在 10K25K 之间即超过),SPR=0.616 即当时最强基线(TWM 在 25K50K 之间即超过)。即 TWM 只用 25% 的交互预算就已超过 SimPLe 的最终成绩,约 50% 的交互预算就超过此前最强基线 SPR 的最终成绩。
消融研究:① 平衡数据集采样(τ=20)显著优于均匀采样(τ=∞)(Figure 7,BankHeist/Breakout/Boxing/KungFuMaster/MsPacman/Pong);② 奖励回填入 Transformer 输入相比不回填,在部分游戏(如 CrazyClimber、Pong)显著提升性能,另一些游戏(不依赖奖励反馈即可预测)效果持平(Figure 8);③ 阈值化熵损失优于普通熵惩罚(Figure 15);④ 历史长度 ℓ=16 优于 ℓ=4(Figure 16);⑤ 策略输入用 z 优于用 [z,h](Figure 17)。
论文未与同期的 IRIS(iris,人类归一化均值 1.046)做直接对比实验。
创新点与影响
贡献(论文自述六点):① 提出基于 Transformer-XL 的自回归世界模型,配合只吃隐状态 z 的 model-free 策略,策略推理阶段完全不需要跑 Transformer(与需要在推理期运行完整世界模型的 dreamer-v2 / transdreamer 形成对比);② 把预测的奖励回填进 Transformer 输入,让模型感知自身已产生的奖励,消融证明有效;③ 把 dreamer-v2 的平衡 KL 散度损失重写为平衡交叉熵损失,可独立调节交叉熵与熵正则化项的权重;④ 提出阈值化熵损失,稳定策略熵、简化跨游戏超参数选择;⑤ 提出平衡数据集采样,用 softmax 温度控制的可变采样概率把训练重心向最新经验倾斜;⑥ 在 Atari 100k 上用 Agarwal et al. (2021) 的统计严谨评测方法,在全部四个聚合指标上超过 SimPLe/DER/CURL/DrQ(ε)/SPR。
影响:TWM 证明了「Transformer 世界模型」可以在保留自回归、直接访问历史等优势的同时,通过让策略只依赖压缩隐状态而非 Transformer 输出,把推理开销压回 model-free 方法水平——为后续 Transformer/序列建模类世界模型(如同期 IRIS 的离散 token 自回归路线)提供了另一种「训练用 Transformer、推理不用」的设计参照。
局限(论文未设独立 Limitations 小节,以下为正文/附录中明确讨论的权衡):① 消融显示策略若改吃 [z,h] 反而分数更低,作者推测策略网络难以跟上 h 分布在训练中的持续变化,这一负面结果也说明「Transformer 输出对策略是否有用」并非显然;② 历史长度 ℓ=16 好于 ℓ=4,意味着更短历史会明显损害性能,而更长历史的计算成本未做进一步探索;③ 论文承认同期独立工作 IRIS 在 Atari 100k 上取得比 TWM 更高的分数,两者未做直接对比实验;④ 平衡数据集采样、阈值化熵损失都引入了新的超参数(τ、Γ),论文未给出这些超参在其他环境/任务上的敏感性分析。
原始链接
- arXiv abstract:https://arxiv.org/abs/2303.07109
- arXiv PDF:https://arxiv.org/pdf/2303.07109
- GitHub:https://github.com/jrobine/twm
一手源存档(sources/)
- twm-transformer-world-model—github-readme — GitHub README 快照
- arXiv 原文 PDF(2303.07109,arXiv 原文 PDF,不入 git)——见上方 arXiv 链接