一句话定位
TWISTER(University of Würzburg,Burchi & Timofte,ICLR 2025)指出 Transformer 世界模型此前”预测下一步隐状态”的目标不足以发挥 Transformer 的表征能力,于是给 Transformer 状态空间模型(TSSM)加上一条**动作条件对比预测编码(AC-CPC)**辅助损失,让模型学会预测未来 10 步的高层特征表示,在 Atari 100k 上取得非前视搜索(look-ahead search)类方法中的新纪录(人类归一化均值 162%)。
背景与定位
DreamerV3(DreamerV3)用 RNN(RSSM)世界模型统治了多个领域,但其后一批工作尝试把 RSSM 换成 Transformer 以获得更好的训练效率和可扩展性——TransDreamer(TransDreamer)、TWM(TWM)、STORM(STORM)都属于这条”Transformer 状态空间模型”路线,但它们的性能提升相对 DreamerV3 有限。另一条路线(IRIS IRIS、Δ-IRIS Δ-IRIS)则用 VQ-VAE 离散化图像 token 后训练自回归 Transformer 直接在像素空间重建轨迹。TWISTER 沿用第一条路线的骨干(TSSM,直接借鉴 TransDreamer 的设计),但指出问题根源在于训练目标太弱:论文观察到世界模型相邻隐状态的余弦相似度非常高(图 2),预测”下一步”对 Transformer 而言太容易,无法逼迫它学到高质量的长程表示——这与 Zhang et al.(STORM 作者)此前的猜测一致。TWISTER 的方案是引入 CPC(Contrastive Predictive Coding,Oord et al. 2018,此前多用于语音/图像/文本预训练及 DeepMind Lab 上作为 A3C 的辅助损失),并把预测扩展到未来动作序列条件下的 K=10 步,构造出”action-conditioned CPC”(AC-CPC)。
模型架构
范式:TSSM(Transformer State-Space Model,masked self-attention 自回归 Transformer)+ 分类变分自编码器(Categorical VAE)离散隐状态 + 动作条件对比预测编码(AC-CPC)辅助表示学习头;世界模型 + actor + critic 三网络联合训练(沿用 DreamerV3 的行为学习设置)。
- 编码器:卷积 VAE,输入图像 64×64×3 → 4 层 Conv+LN+SiLU(32/64/128/256 通道,逐层下采样到 4×4)→ flatten 4096 → Linear 1024 → reshape 为 32 个特征×32 类别的 categorical 分布,采样得到离散随机状态 $z_t$(32×32,即 32 类特征各 32 类,与 DreamerV3 一致)。
- 解码器:对称的转置卷积结构,从 $z_t$(32×32=1024 维)重建 64×64×3 图像。
- Transformer 网络(TSSM 核心):先用”Action Mixer”把 $z_{0:T-1}$ flatten 后与动作 $a_{0:T-1}$ concat,经 Linear+LN+SiLU → Linear+LN 投影到 512 维,再送入 4 层 Transformer block(8 个注意力头,512 通道,dropout 0.1),使用 相对位置编码(Transformer-XL 式,Dai et al. 2019)而非绝对位置编码——这样在 imagination/评估阶段可以直接缓存 K/V 特征做上下文扩展,无需在超出训练长度时重新计算位置编码。注意力上下文长度(Attention Context Length)设为 8。输出隐状态 $h_t$ 与随机状态 $z_t$ 拼接构成 agent 状态 $s_t={h_t,z_t}$。
- 动作条件(Action-conditioning):动作直接在 Action Mixer 阶段与隐状态 concat 后进入 Transformer;此外 AC-CPC 预测头也显式以未来动作序列 $a_{t:t+k}$ 为条件(见下)。
- 预测头(均为简单 MLP,隐藏维度 512):Reward predictor(3 层,输出 255 维 symlog 离散分布)、Continue predictor(3 层,Bernoulli)、Representation network(2 层,把增强后的未来随机状态 $z’{t+k}$ 投影为 512 维对比特征 $e_t^k$)、AC-CPC predictor(2 层,以 $s_t$ 与未来动作序列 $a{t:t+k}$ 为输入,输出 512 维预测特征 $\hat e_t^k$)、Critic network(3 层,255 维 symlog 离散)、Actor network(3 层,输出 one-hot categorical,维度=动作数 A)。
- AC-CPC 机制:对比学习目标最大化模型状态 $s_t$ 与未来(增强视图下的)随机状态 $z’{t:t+K}$ 之间的互信息,K=10 步;负样本采用batch 内其余 B×T−1 个样本;相似度用点积 $\mathrm{sim}(z_j’, s_t)=q\phi^k(z_j’)^\top p_\phi^k(s_t,a_{t:t+k})$,对每个步长 k 各学一对 MLP($q_\phi^k$、$p_\phi^k$)。与原始 CPC(连续特征)不同,TWISTER 处理的是离散隐状态,因此需要额外学一个投影网络把离散 $z’$ 映射为对比特征。
- 与同族方法的架构对比(论文 Table 1):TWM/IRIS/STORM/Δ-IRIS 均只预测”下一状态”,唯有 TWISTER 把预测视野扩展到 K=10 步;IRIS/Δ-IRIS 用 VQ-VAE 离散 token(分别 4×4、2×2 空间 token),TWISTER 与 TWM/DreamerV3/STORM 一样用 categorical-VAE 单一 latent;agent 状态上,TWISTER 与 TWM/STORM 一样用 {latent, hidden},IRIS/Δ-IRIS 用原始图像。
数据
论文在 Atari 100k 与 DeepMind Control Suite(DMC)两个仿真基准上做在线强化学习,不涉及预训练用的离线大规模数据集:
- Atari 100k(Kaiser et al. 2020):26 款 Atari 游戏,每个游戏预算 100k 次环境交互(400k 环境帧,默认 action-repeat=4),约合 1.85 小时真实游戏时长;数据来自 agent 与环境在线交互产生的经验,存入 replay buffer 采样训练(batch size 16、序列长度 64)。
- DeepMind Control Suite:20 项连续控制任务,仅用高维图像观测,预算 100 万环境步(1M steps)。
- 数据增强:AC-CPC 的正/负样本用随机裁剪+缩放(crop scale 0.25–1.0,aspect ratio 0.75–1.33)做图像增强,构造”增强视图”下的未来状态 $z’$;论文也测试了随机平移(±4 像素,Yarats et al. 2021)增强,但发现效果不如随机裁剪+缩放。
- 没有额外的离线数据混合、跨具身或 sim-to-real 环节;训练/评测数据都直接来自 Atari 模拟器与 MuJoCo/DMC 的在线 rollout。
训练方法
世界模型损失(式 2):$L(\phi)=L_{rew}+L_{con}+L_{rec}+L_{dyn}+L_{cpc}$,五项联合训练:
- $L_{rew}$、$L_{con}$:沿用 DreamerV3 的 symlog cross-entropy / binary cross-entropy 预测奖励与 episode continuation。
- $L_{rec}$:像素重建 L2 损失,从随机状态 $z_t$(而非含时序信息的 $s_t$)重建图像,与前作一致地防止解码器”作弊”利用时序信息。
- $L_{dyn}$:动态预测器与编码器分布间的 KL(含 stop-gradient 的双向正则项,$\beta_{dyn}=0.5$,$\beta_{reg}=0.1$,两项 KL 均用 $\max(1,\cdot)$ 设一个下限/free-bits 阈值以避免 KL 塌缩为 0)。
- $L_{cpc}$(式 5):K=10 步 InfoNCE 损失,动作条件由 AC-CPC predictor 显式接收未来动作序列。
行为学习(actor-critic):完全在想象(latent imagination)中进行,沿用 DreamerV3 设置——把采样序列的 batch×time 展平成 $B^{img}=B\times T$ 条轨迹,world model 用缓存的 K/V 继续想象 H=15 步;critic 用离散化 λ-return 回归(symlog cross-entropy)+ 对自身 EMA 网络的正则项(无独立 target network);actor 用 Reinforce + 熵正则,advantage 按 5/95 分位数的 EMA 统计归一化(momentum 0.99)。
关键超参(Appendix A.3):Batch size 16,序列长度 64,图像分辨率 64×64 RGB;Transformer 4 层,8 头,dropout 0.1,注意力上下文长度 8;随机状态 32 类特征×32 类别;世界模型学习率 1e-4(Adam,β=(0.9,0.999),ε=1e-8,梯度裁剪 1000);actor-critic 学习率 3e-5(Adam,ε=1e-5,梯度裁剪 100),回报折扣 γ=0.997,λ=0.95,critic EMA decay 0.98,actor 熵系数 η=3e-4。消融确认 K=10 是最优 CPC 步数(K=15 时性能反而下降,均值/中位数曲线见图 6a),对应约 0.67 秒游戏时长的展望。
Infra(训练 / 推理工程)
- 训练硬件、GPU 数量、GPU-hours、并行策略:论文正文与附录均未披露具体训练硬件配置(未提及 GPU 型号/数量/训练时长),仅在致谢中说明工作由 Alexander von Humboldt Foundation 资助——未披露。
- 精度(fp16/bf16/fp32):未披露。
- 推理 FPS / control-Hz / 延迟:论文未报告推理速度或控制频率的量化数字——未披露。GitHub 仓库(burchim/TWISTER)仅提供训练/评估脚本入口和超参覆盖方式,未给出硬件或速度基准。
- 边缘设备部署:不适用,论文为纯研究性 RL 基准评测,未讨论边缘部署。
评测 benchmark
Atari 100k(Table 2,26 款游戏,主指标为人类归一化均值/中位数):
| 方法 | Normed Mean (%) | Normed Median (%) | # Superhuman (/26) |
|---|---|---|---|
| Random | 0 | 0 | 0 |
| SimPLe | 33 | 13 | 1 |
| TWM | 96 | 51 | 8 |
| IRIS | 105 | 29 | 10 |
| DreamerV3 | 112 | 49 | 9 |
| STORM | 127 | 58 | 10 |
| Δ-IRIS | 139 | 53 | 11 |
| TWISTER(本作) | 162 | 77 | 12 |
TWISTER 在 26 局中的多项游戏上大幅领先,最突出的是 Gopher(22234 分,DreamerV3 为 3730、STORM 为 8240)与 Private Eye(1608 分,远超 DreamerV3 的 882,但仍远低于 STORM 的 7781);Kangaroo(6016 分)也明显领先其他方法。论文指出 TWISTER 在”关键奖励物体数量多”(Amidar、Bank Heist、Gopher、Ms Pacman)以及”小型移动物体”(Breakout、Pong、Asterix)类游戏上收益明显,归因于 AC-CPC 迫使模型关注运动物体的位置来完成未来预测,缓解了重建损失容易忽略小物体的问题。
消融研究(Table 3,主消融,26 局聚合分数):
| 配置 | Normed Mean (%) | Normed Median (%) |
|---|---|---|
| TWISTER(完整) | 162 | 77 |
| 去掉 AC-CPC(No AC-CPC) | 112 | 44 |
| 世界模型换成 DreamerV3 的 RSSM | 121 | 69 |
| CPC 预测头去掉未来动作条件(No Action Conditioning) | 111 | 42 |
| 去掉数据增强(No Data Augmentation) | 120 | 68 |
三项关键消融结论:① AC-CPC 本身贡献巨大(162%→112%,约 50 个百分点);② 用 Transformer(TSSM) 替代 RNN(RSSM) 本身收益有限(两者不加 AC-CPC 时相近),但 AC-CPC 只在 Transformer 上才充分发挥作用(RSSM+AC-CPC 121% vs TSSM+AC-CPC 162%),论文将其归因于自注意力比 RNN 更擅长学习长程特征表示;③ 去掉未来动作条件后 CPC 目标几乎失去作用(111% 与不加 CPC 的 112% 接近),说明”动作条件”是 AC-CPC 有效的关键,因为不知道未来动作时预测遥远状态本质上是不可解的任务。
DeepMind Control Suite(Table 11,20 项任务,1M 环境步):TWISTER 均值 801.8,中位数 907.6,对比 DreamerPro 770.9/858.1、DreamerV3 739.6/808.5、TD-MPC2 720.9/795.9;TWISTER(No CPC)消融为 728.7/802.4——AC-CPC 在 Acrobot Swingup(239.4 vs 无 CPC 的 81.6)、Quadruped Run/Walk(652.1/904.9 vs 503.8/741.6)、Walker Run(711.2 vs 566.2)等复杂任务上收益尤其明显。
创新点与影响
- 指出并解决了 Transformer 世界模型的表示学习瓶颈:此前 TransDreamer/TWM/STORM 一系列 Transformer 化世界模型性能提升有限,TWISTER 首次系统论证根因在于”预测下一状态”目标过弱(相邻隐状态余弦相似度过高),并给出简洁的解法——把预测视野拉长并引入对比学习。
- 提出 action-conditioned CPC(AC-CPC):把经典 CPC(此前主要用于语音/图像/文本预训练及 DeepMind Lab 辅助损失)迁移到 model-based RL,创新点是用未来动作序列作为预测条件,解决了”不知道未来动作则长程预测本质不可解”的问题。
- 刷新非前视搜索类方法在 Atari 100k 上的纪录(162% 均值 / 77% 中位数,12/26 项达到超人类水平),同时在 DMC 连续控制上也取得 SOTA(均值 801.8)。
- 作者自述局限/未竞争范围:论文明确将比较范围限定在”不使用前视搜索(look-ahead search)“的方法,未与 EfficientZero V2(当前 Atari 100k 总榜第一,用 MCTS)或 BBF(使用周期性网络重置等正交技巧的 model-free 方法)直接竞争,作者认为将 AC-CPC 与这些正交技术结合是未来方向;此外论文未讨论 real-time 推理速度、大规模/跨环境泛化等问题。
原始链接
- arXiv abs:https://arxiv.org/abs/2503.04416
- arXiv PDF:https://arxiv.org/pdf/2503.04416
- OpenReview(ICLR 2025):https://openreview.net/forum?id=YK9G4Htdew
- GitHub:https://github.com/burchim/TWISTER
一手源存档(sources/)
- twister-contrastive-world-model—github-readme — GitHub README(fetched 2026-07-16)
- arXiv 2503.04416 全文(arXiv 原文 PDF,不入 git,见上方链接)