一句话定位
REM(Retentive Environment Model)给 token 化世界模型换上 RetNet 骨干,并设计并行观测预测(POP)机制,把想象阶段生成下一帧全部 token 的世界模型调用次数从 K·H 降到 2H(K=每帧 token 数,H=想象步数),带来想象阶段 15.4 倍加速,在 Atari 100k 的 26 个游戏上 12 个超越人类、训练耗时不到 12 小时;ICML 2024 收录。
背景与定位
iris(ICLR 2023)开创了 token 化世界模型(token-based world model, TBWM)路线:用 VQ-VAE 把每帧压成一串离散 token,再用 GPT 式自回归 Transformer 在 token 序列上建模动力学。但 IRIS 的致命瓶颈是想象(imagination)阶段必须逐 token 串行生成下一帧——预测第 k 个 token 依赖已生成的前 k-1 个 token,生成 H 步观测需要 K·H 次串行世界模型调用,GPU 利用率极低、训练耗时随 token 数线性增长,实际上限制了 IRIS 只能用很粗的 token 网格(4×4=16 个 token/帧)。
REM 直接针对这个瓶颈:把 IRIS 的 causal Transformer 骨干换成 RetNet(Retentive Network, Sun et al. 2023)——一种具有”chunkwise”并行/循环对偶形式的序列模型,用一个 d×d 的循环状态 S 汇总历史信息。基于 RetNet 的这一特性,REM 设计了 POP 机制,让下一帧的全部 K 个 token 在一次前向里并行生成,从而摆脱对 K 的依赖。同期的非 token 化 Transformer 世界模型 twm-transformer-world-model(Transformer-XL 骨干)与 storm(小型 causal Transformer + categorical VAE,把每帧当作单一 token 处理)代表了绕开这一瓶颈的另一条路径;REM 则证明了 token 化路线本身也可以通过序列模型架构创新解决瓶颈,而不必放弃 token 化带来的更细粒度视觉表征。论文归类为想象中学习(learning in imagination)范式下的 V-M-C(Ha & Schmidhuber 2018)结构:视觉感知(Tokenizer)+ 动力学模型(World Model)+ 控制器(Controller)。
模型架构
REM 遵循 V-M-C 三段式结构,三部分独立训练(真实经验采集 → tokenizer 训练 → world model 训练 → actor-critic 在想象中训练,见”训练方法”)。
𝒱 Tokenizer(VQ-VAE / VQGAN 风格离散自编码器)
- 编码器把 64×64×3 的输入帧映射为 K=64 个(8×8 网格)d 维潜向量,每个向量最近邻量化到码本 E∈ℝ^(N×d) 的索引(token),码本大小 N=512;IRIS 只用 K=4×4=16 个 token/帧,REM 把网格分辨率提到 8×8,是 IRIS 的 4 倍。
- 架构:编码器 3 个下采样 EncoderBlock(每个含 GroupNorm(8 groups)+SiLU+非对称 padding 卷积,通道 32→64→128→256,空间 64→32→16→8),末尾接 GN+SiLU+3×3 卷积;解码器对称使用最近邻插值上采样代替反卷积。
- 训练目标(Eqn. 8):L1 重建损失 + 两项 commitment loss(VQ-VAE 标准做法,停梯度算子分别应用于编码器输出和码本)+ 感知损失(perceptual loss),无判别器(不同于原始 VQGAN)。
ℳ Retentive World Model(POP + RetNet)
- 骨干:5 层 RetNet,4 个 Retention head,embedding 维度 d=256,dropout 0.1,前馈维度 1024,LayerNorm epsilon 1e-6;用开源实现 yet-another-retnet。
- 输入是观测-动作交替的 token 轨迹,按”block”(K 个观测 token + 1 个动作 token)组织,训练时切成大小 B=c(K+1) 的 chunk(c=3,即每 chunk 含 3 个 block),chunk 间靠 RetNet 的循环状态 S 传递历史信息。
- 世界模型上下文为 2 帧(IRIS 原始实现只用单帧上下文做预测,REM 扩展为 2 帧)。
- POP 核心机制:额外引入 K 个专用”预测 token” u=(u₁,…,u_K) 及其独立的可学习 embedding 表 E_u;在想象阶段,从当前时刻的循环状态 S_t 出发,把 u 整体喂给 RetNet 一次前向即可并行输出全部 K 个下一观测 token 的分布 p(ẑ_{t+1}^k | 历史, u_{≤k})——u 只用于生成观测预测,从不进入循环状态的更新。这样将想象阶段的世界模型调用次数从 K·H(IRIS 逐 token 串行)降到 2H(默认模式:一次调用消费上一 block 算出当前状态,代价 K+1;另一次调用用 u 生成观测,代价 K);论文还给出一种”合并调用”的替代模式,把串行调用次数进一步压到 H,但每次调用代价升到 2K+1,总计算量两种模式相同((2K+1)H),只是分布在不同数量的串行步里,具体选哪种取决于硬件和 batch size(REM 实验里发现默认模式更快)。
- 代价权衡:POP 把总计算量从 (K+1)H 提高到 (2K+1)H(因为要多算 K 个预测 token 的输出),但换来的是把串行深度从 K·H 砍到 2H(或 H),GPU 利用率大幅提升,净壁钟时间反而更快——这与 Transformer 用更高计算成本换更好可扩展性的历史经验一致。
- 训练侧的核心挑战(对应 Section 2.4):要让世界模型学会”在每个时间步 t 都能用 u 做出有效预测”,需要计算每个前缀的循环状态 S_[i,j](chunk i 内第 j 个 block 后的状态),因为直接把 z_t 替换成 u 会破坏”预测未来还需要真实 z_t”的因果链。论文设计了 Algorithm 1/2 的两阶段并行方案:先并行算出所有中间状态 S̃(Eqn. 6),再顺序累积成 chunk 内各 block 的完整状态 S(Eqn. 7,该步计算量小,对总加速影响可忽略),最后用这些状态批量算出所有 (S_[i,j], u) 元组的观测预测输出——这需要扩展 RetNet 使其支持”同一 batch 内不同时间步状态+共享位置编码”的批量前向(标准 RetNet chunkwise 前向只支持同一时间步的状态)。
𝒞 Controller(actor-critic,在想象中训练)
- 共享骨干:观测 token 先映射回 tokenizer 的码本向量、按空间排布,经 2 层卷积(256×8×8 → 128×8×8 → 64×8×8,SiLU 激活)+ 展平 + 线性层得到 512 维向量;动作 token 用独立可学习 embedding 表;再经一个 512 维 LSTM 得到历史相关的隐向量,接线性 actor/critic 头。
- 与 IRIS 的关键差异:IRIS 的 controller 直接处理重建后的像素帧且不输入历史动作(π(a_t|ô_{≤t}));REM 的 controller 直接吃离散 token 的码本向量(不重建像素)且输入已采样的历史动作(π(a_t|ẑ_{≤t},a_{<t},ẑ_t)),消融证明这两点都对最终表现有贡献(见”评测”)。
- 训练目标:λ-return(γ=0.995, λ=0.95)+ REINFORCE 策略梯度(以 value 为 baseline)+ 熵正则(权重 α=0.001),与 IRIS/DreamerV2 一脉相承。
数据
REM 完全在真实环境交互中在线学习(无离线数据集、无跨环境预训练):
- Atari 100k 基准(26 款游戏):每个游戏严格限制 100,000 次交互步(因标准 frame-skip=4,对应 400,000 游戏帧),约等于 2 小时人类游戏时长——相比原始 Atari 基准的 2 亿步(多 500 倍),是极端样本受限设定。
- 经验采集节奏(Table 2):600 个总 epoch 中,前 500 个 epoch 每 epoch 采集 200 个真实环境步(500×200=100,000,精确对应 Atari 100k 预算),采集阶段用 ε-greedy=0.01 探索;剩余 100 个 epoch 不再采集新数据,只用已有回放缓冲区继续训练三个组件收敛。
- 消融实验用 26 款游戏中的子集 8 款(Assault, Asterix, ChopperCommand, CrazyClimber, DemonAttack, Gopher, Krull, RoadRunner——挑选 IRIS 与 REM 分差最大的游戏),每个设置 5 个随机种子,受限于计算资源。
- 主实验每个游戏 5 个随机种子,每个种子训练结束后跑 100 个 episode 取平均作为最终得分;帧分辨率统一 64×64,最大 no-op 数(训练,测试)=(30,1),最大 episode 步数(训练,测试)=(20K,108K),是否因失去一条命而终止 episode(训练,测试)=(No,Yes)。
训练方法
REM 每个 epoch 循环 4 步(Figure 2 / Algorithm 3):① 用当前策略与真实环境交互采集经验存入回放缓冲区 → ② 从缓冲区均匀采样帧训练 tokenizer → ③ 采样轨迹片段训练 RetNet 世界模型 → ④ 在世界模型想象出的轨迹里训练 actor-critic 策略。
- Tokenizer 训练:损失见”模型架构”(重建 L1 + 两项 commitment + 感知损失);学习率 1e-4,batch size 128,梯度裁剪阈值 10,权重衰减 0.01,从第 5 个 epoch 开始训练,每 epoch 训练 200 步。
- 世界模型训练(POP 训练模式):输入是从缓冲区均匀采样的 H=10 步轨迹片段,切成 c=3 个 block 一组的 chunk;对转移和终止输出用交叉熵损失,奖励损失依任务用 MSE(连续)或交叉熵(离散);学习率 2e-4,batch size 64,梯度裁剪阈值 100,权重衰减 0.05,从第 25 个 epoch 开始,每 epoch 训练 200 步。
- Actor-critic 训练(想象展开):从缓冲区采样的短轨迹片段初始化世界模型和 controller 的状态后,在想象里展开 H=10 步;学习率 1e-4,batch size 128,梯度裁剪阈值 3,权重衰减 0.01,从第 50 个 epoch 开始,每 epoch 训练 100 步;评估时采样温度 0.5。
- 优化器统一 AdamW(β1=0.9, β2=0.999)。
Infra(训练 / 推理工程)
- 计时基准硬件:单张 Nvidia RTX 4090 工作站,用于 REM 与 IRIS 的运行时间对比实验(Figure 1/10/11)。
- 主实验训练硬件:其余(性能)实验均在 Nvidia V100 GPU 上完成;论文未披露具体 GPU 数量、是否多卡并行、混合精度设置。
- 速度结果:POP 使想象阶段获得 15.4 倍加速(相对 IRIS 式逐 token 串行生成);REM 整体训练在不到 12 小时内完成(对应 100k 交互步的 Atari 100k 预算)。论文额外做了”用 REM 配置改造后的 IRIS”消融(Figure 11)以证明 POP 本身(而非仅仅是配置差异)带来的加速。
- 世界模型调用次数对比(Table 6,仅计观测预测部分,忽略动作 token 处理开销):
算法 训练总计算量 想象串行调用次数 想象单次调用代价 POP(默认模式) 2KH 2H K POP(合并模式) 2KH H 2K No POP(IRIS 式) KH KH 1 - 推理侧 FPS / 控制频率:论文未披露(未给出具体 fps 数值,只给出上表的调用次数对比及图示的壁钟时间对比)。
评测 benchmark
主评测:Atari 100k 全部 26 款游戏(Table 1,每游戏 5 seed,每 seed 训练后跑 100 episode 取平均),对照 Random / Human / SimPLe / DreamerV3 / TWM / STORM(非 token 化)与 IRIS(token 化):
| 指标 | SimPLe | DreamerV3 | TWM | STORM | IRIS | REM(本文) |
|---|---|---|---|---|---|---|
| Superhuman(26 局,↑) | 1 | 9 | 8 | 9 | 10 | 12 |
| Mean HNS(↑) | 0.332 | 1.124 | 0.956 | 1.222 | 1.046 | 1.222 |
| Median HNS(↑) | 0.134 | 0.485 | 0.505 | 0.425 | 0.289 | 0.280 |
| IQM HNS(↑) | 0.130 | 0.487 | 0.459 | 0.561 | 0.501 | 0.673 |
| Optimality Gap(↓) | 0.729 | 0.510 | 0.513 | 0.472 | 0.512 | 0.482 |
REM 在 IQM 上全面领先所有 baseline(0.673),#Superhuman 数(12/26)也是所有方法中最高;相对 IRIS,REM 在 mean、optimality gap、IQM 三项指标上都更优,median 大致持平(略低)。逐游戏亮点:Assault 1764.2(IRIS 1524.4)、Asterix 1637.5(IRIS 853.6)、Boxing 87.5(IRIS 70.1)、Breakout 90.7(IRIS 83.7)、DemonAttack 5738.6(IRIS 2034.4,STORM 仅 164.6)、Gopher 5452.4(IRIS 2236.1)、RoadRunner 14060.2(STORM 17564.0 略高但 REM 仍超人类)。也有明显劣势项:BankHeist REM 仅 19.2(远低于 IRIS 53.1、DreamerV3 648.7)、Kangaroo 467.6(IRIS 838.2、DreamerV3 4098.3)、Hero 6484.8、Qbert 743.0 均低于多数 baseline——这也是 REM median 分数偏低的主因(少数游戏拖累中位数,但均值/IQM 受益于其余游戏的大幅超越)。
消融(8 局子集,5 seed,Table 9):REM 完整版 Mean=1.947 / IQM=2.201 / 6 个超人类;对照 IRIS(原配置)Mean=1.564 / IQM=1.191 / 5 个超人类;“No POP”(把 POP 换回逐 token 串行生成)Mean=2.357 / IQM=2.068 / 6 个超人类——证明 POP 基本不损失(甚至略微超过)智能体表现,同时大幅缩短总训练时间;世界模型用独立 embedding 表(而非复用 tokenizer 码本)版本 Mean=1.778;tokenizer 分辨率降到 4×4(等同 IRIS 分辨率)版本 Mean=1.341/IQM=1.026(全消融里表现最差之一),证明高分辨率 token 网格对性能有显著贡献;controller 换成 IRIS 架构(吃像素、无动作输入)版本 Mean=1.340/IQM=1.234;仅去掉 controller 动作输入版本 Mean=1.571/IQM=1.535。附录 A.4 另给出观测预测的交叉熵损失对比,证实 POP、复用 tokenizer embedding 两项设计均能改善世界模型自身的预测质量(而不仅是下游策略表现)。
创新点与影响
- 核心贡献:提出 POP 机制,把 RetNet 的 chunkwise 循环/并行对偶形式扩展出”用专用预测 token 批量查询循环状态”的新前向模式,首次让 token 化世界模型在想象阶段摆脱”逐 token 串行生成下一帧”的复杂度依赖(从 O(KH) 串行调用降到 O(H)),代价是总计算量翻倍但换来数量级的壁钟加速与更高 GPU 利用率。
- REM 是首个采用 RetNet 架构的世界模型智能体,论文称这是 RetNet 在强化学习场景下有效性的”首个证据”。
- 影响:POP 使 token 化路线可以负担得起远高于 IRIS(4×4=16)的 token 网格分辨率(8×8=64),消融证明分辨率提升是 REM 相对 IRIS 性能提升的重要来源之一;论文结尾提出的未来方向(世界模型/controller 用能覆盖完整历史的循环状态、tokenizer 独立优化以便利用大规模预训练视觉模型、把 POP 用于视频生成任务的整帧并行生成)之后被同一作者的后续工作 Simulus(GitHub 仓库 leor-c/Simulus)承接改进。
- 作者自陈局限(论文正文与结论明确提及):(1) POP 默认将总计算量提高到 (2K+1)H(较 No-POP 的 (K+1)H 更贵),是用更多 FLOPs 换取更好的可扩展性和更低串行深度;(2) REM 在 median HNS 上不如 DreamerV3/TWM/STORM,个别游戏(BankHeist、Kangaroo、Hero、Qbert 等)明显弱于 baseline,论文未深入分析具体原因;(3) 论文本身承认其消融实验受限于计算资源,只能在 8 局子集上做,而非全部 26 局;(4) 结论部分明确把”让世界模型/controller 的循环状态覆盖整段历史”列为未解决的开放方向,当前实现仍受限于按 H 步片段初始化状态的训练方式。
原始链接
- arXiv: https://arxiv.org/abs/2402.05643
- PDF: https://arxiv.org/pdf/2402.05643
- GitHub: https://github.com/leor-c/REM
- OpenReview(ICML 2024 收录页): https://openreview.net/forum?id=Lfp5Dk1xb6
一手源存档(sources/)
- rem-parallel-observation-prediction—github-readme — GitHub README 快照(sources/world-model/2024/rem-parallel-observation-prediction—github-readme.md)
- arXiv 全文(HTML, v5 终版)已通读,未入 git;引用见上方 arXiv URL(arXiv 原文 PDF,不入 git)