一句话定位

IRIS(Imagination with auto-Regression over an Inner Speech)把世界模型拆成一个离散自编码器(VQVAE,把每帧压成 16 个 image token)+ 一个 GPT 式自回归 Transformer(在 token 序列上逐 token 预测下一帧、奖励、终止),策略完全在这个世界模型「想象」出的轨迹里用 actor-critic 学。在只允许 10 万步 real-env 交互(≈2 小时人类游戏)的 Atari 100k 上,IRIS 取得 1.046 的人类归一化均值、26 个游戏中 10 个超越人类,成为无 lookahead search 类方法的新 SOTA(还顺带超过了非样本高效设定下训练的 MuZero)。ICLR 2023。

背景与定位

RL 长期的痛点是样本效率极低——DreamerV2 在 Atari 上要几个月的游戏量、OpenAI Five 在 Dota2 要几千年。Model-based / world model 是通往数据高效的主要方向,其中「在世界模型的想象里学策略」这一支尤其吸引人:一旦世界模型够准,策略就完全脱离真实交互的样本约束。这条线由 world-models-ha-schmidhuber(2018,autoencoder + RNN,玩具环境)开创,SimPLe(2020,像素空间视频预测 + PPO)首次在 Atari 100k 上显现潜力,而当时「在想象中学」的最强 Atari agent 是 dreamer-v2——但它是在 2 亿帧的非样本高效设定下开发评测的。此前的世界模型骨架几乎都是 RSSM(planet / dreamer-v1 / dreamer-v2 这类卷积自编码器 + 循环状态空间模型。

IRIS 的立意:把 Transformer 从 NLP / CV / offline-RL(decision-transformer、Trajectory Transformer)引进世界模型。Transformer 在离散 token 序列上最擅长,但图像不像文本天然有词表——朴素地把像素当 token 会让序列长度爆炸、注意力二次方开销不可行。VQGAN(Esser 2021)与 DALL·E(Ramesh 2021)给出的答案是用离散自编码器(VQVAE, Van Den Oord 2017)把一帧压成极少量 image token,再让 Transformer 自回归建模——IRIS 把这套「图像 token 化 + 自回归 Transformer」的组合直接改造成世界模型,把动力学学习整体转写成序列建模问题:自编码器造出一套 image token 的「语言」,Transformer 在时间上「作曲」这套语言。范式上属于 model-based RL / learning-in-imagination,与 muzero / efficientzero 这类 lookahead search(MCTS) 路线正交——IRIS 决策时不做任何搜索

模型架构

IRIS = 一个 model-based agent 的三件套循环:collect_experience(真实环境采数据)→ update_world_model(学奖励/终止/下一帧预测)→ update_behavior(在想象里学策略与价值)。策略只在想象 MDP 里学,真实经验只用来学环境动力学。世界模型由两个部件组成:

部件一:离散自编码器 (E, D)——把图像 ↔ token。

  • 编码器 E: R^{h×w×3} → {1,…,N}^K 把一帧图像经 CNN 产出 y_t ∈ R^{K×d},再对每个位置取嵌入表 E={e_i} 中最近邻的下标得到 K 个 token z_t^k = argmin_i ||y_t^k − e_i||(VQ 量化,straight-through 估计器反传)。
  • 解码器 D: {1,…,N}^K → R^{h×w×3} 是 CNN,把 K 个 token 还原成图像。
  • 实现上直接改自 VQGAN:去掉判别器,退化成带感知损失的 vanilla VQVAE
  • 量化配置数字:词表大小 N = 512每帧 K = 16 个 token、token 嵌入维度 d = 512;输入帧 64×64;编解码器各 4 层、每层 2 个残差块、卷积 64 通道,在 8/16 分辨率上加 self-attention 层。

部件二:GPT 式自回归 Transformer G——建模动力学。

  • G 在交错的帧-动作 token 序列 (z_0^1,…,z_0^K, a_0, z_1^1,…,z_1^K, a_1, …) 上运行,输入序列由原始 (x_0,a_0,x_1,a_1,…) 经 E 编码得到。
  • 动作条件化(action-conditioning):动作作为独立 token 插在每帧 K 个 image token 之后,用 A×D 的动作嵌入表、帧 token 用 N×D 的嵌入表,一并送进 M 个 Transformer block。
  • 每个时间步 G 建模三个分布:Transition(下一帧 token 逐 token 自回归 ẑ_{t+1}^k ~ p_G(· | z_{≤t}, a_{≤t}, z_{t+1}^{<k})——自回归发生在 token 级,第 k 个 token 的条件里也包含已预测的前 k−1 个 token)、Reward r̂_tTermination d̂_t ∈ {0,1}
  • block 是 GPT2 式(minGPT 实现改来):pre-LN self-attention + 残差,接 pre-LN 的逐位置 MLP + 残差。作者特别指出:Parisotto 2020 发现标准 Transformer 用 RL 目标难优化、需 gating,而 IRIS 的世界模型无需此类改造,很可能因为它的训练目标是自监督的(而非直接吃 RL 梯度)。
  • Transformer 配置数字:时序长度 L = 20 步、嵌入维 D = 256、层数 M = 10、注意力头 4、weight decay 0.01、embedding/attention/residual dropout 均 0.1。输入序列长度 = L×(K+1) = 20×17 = 340 token。

部件三:Actor-Critic(策略与价值)。

  • actor 与 critic 共享权重、仅最后一层分开;输入是 64×64×3 的重建帧 x̂_t(注意:策略吃的是解码器重建的图像,不是 token)。
  • 结构 = 一个卷积块(3×3 conv stride1 + ReLU + 2×2 maxpool,重复 4 次)接一个 LSTM cell(隐状态 512);从某帧开始想象前,先 burn-in 前 20 帧初始化 LSTM 隐状态。
  • 想象在真实观测 x_0 处初始化,向前推 H 步(imagination horizon),若中途预测到 episode 结束则停。因为想象步数固定不能用 Monte Carlo,故用价值网络 V(x̂_t) 对超出时程的回报做 bootstrap。

数据

IRIS 是在线 RL,无外部预训练语料——「数据」就是 agent 自己在 Atari 里交互、不断增长的经验回放:

  • 样本预算 = Atari 100k:26 个 Atari 游戏,每个只允许 10 万个动作(≈2 小时人类游戏);作为对照,无约束 Atari agent 通常训 5000 万步,是其 500 倍经验。
  • 三个部件各自的采样/批(在存储的过往经验上采):自编码器批大小 256、Transformer 批大小 64、actor-critic 批大小 64;Transformer 在长度 L = 20 步的片段上自监督训练。
  • 采数据细节:collection 用固定 ε-greedy = 0.01 并从策略采样(默认采样温度 1);评测采样温度 0.5Freeway 这个稀疏奖励游戏特殊:把采数据/探索时的采样温度从 1 降到 0.01(Appendix H,属探索采数据环节,非评测),以避免早期训练阶段不利于学习的随机游走。注意即使在真实环境采数据时,帧也要过一遍自编码器(x̂ = D(E(x)))再喂策略,以保持策略输入分布一致。
  • sim-vs-real / 想象放大:策略「精确地模拟了数百万条轨迹」来学——真实交互只有 10 万步,想象轨迹是它的成千上万倍。
  • 无 co-training、无跨域数据:纯单游戏、单环境,从零学。

训练方法

总循环(Algorithm 1):每个 epoch 先 collect_experience 采若干步,再做若干步 update_world_model(更新 E、D 和 G),再做若干步 update_behavior(更新 π 和 V)。三个部件错峰启动:自编码器从第 5 epoch、Transformer 从第 25 epoch、actor-critic 从第 50 epoch 开始训。共 600 epoch(其中 500 个采数据 epoch),每 epoch 真实环境 200 步、每 epoch 200 个训练步

世界模型目标:

  • 自编码器:等权重组合 L1 重建损失 + commitment 损失(VQ)+ 感知损失(perceptual, Johnson 2016),straight-through 估计器反传。
  • Transformer G:自监督地在过往经验的片段上训——transition 与 termination 用交叉熵reward 用 MSE 或交叉熵(取决于奖励函数形式)。

行为学习(沿用 dreamer-v2 的 actor-critic 目标与超参,作者称”为简单起见”):

  • Critic:回归 λ-return(λ=0.95),平方损失、stop-gradient 目标(Dreamer 式)。
  • Actor:用 REINFORCE(想象里轨迹充裕,可以直接用),以 V(x̂_t) 作 baseline 降方差,加熵正则 η=0.001 维持探索。
  • RL 超参:imagination horizon H = 20、γ = 0.995、λ = 0.95、熵系数 η = 0.001
  • 共享优化超参:学习率 1e-4、Adam(β1=0.9, β2=0.999)、max grad norm 10.0。

无蒸馏、无 lookahead search;「minimal tuning」是作者反复强调的卖点——相比 battle-hardened 的 model-free baseline(带 prioritized replay、ε 调度、数据增强等技巧),IRIS 几乎不做额外调参。

Infra(训练 / 推理工程)

  • 训练硬件8 张 Nvidia A100 40GB;每个 Atari 环境用 5 个随机种子重复训练。
  • 训练时长两个 Atari 环境跑在同一张 GPU 上,一轮训练约 7 天,摊到每个环境平均 3.5 天
  • 并行/精度:IRIS 走直白的单 GPU / 单 CPU 实现(与无搜索类 baseline 一样),不像 EfficientZero 那样需要 CPU/GPU 多线程分布式 + C++/Cython 版 MCTS。论文未披露混合精度等细节(未披露)。
  • 同类工程对照(Appendix G):SimPLe(唯一另一个在想象中学的 baseline)用单张 P100 在单环境训 3 周;最强无搜索 baseline SPR 用 P100 只需 4.6 小时(IRIS 慢很多,因为要训生成式世界模型);MuZero 原版用 40 TPU × 12 小时训单环境,EfficientZero / MuZero 复现用 4×RTX 3090 × 7 小时
  • 推理:论文未给部署 FPS / 控制频率 / 延迟数字(未披露);GitHub 提供交互脚本可让用户用键盘直接在世界模型里游玩(为加速交互,Transformer 记忆每 20 帧刷新一次)。

评测 benchmark

Atari 100k,26 游戏,2 小时真实经验后的人类归一化聚合指标(Table 1)。 IRIS 评测为每游戏训练结束后收集 100 个 episode 取均值、5 个 seed。粗体为无搜索类最优、下划线为总体最优:

指标RandomHumanSimPLeCURLDrQSPRIRISMuZero*EfficientZero*
Superhuman0N/A123610514
Mean ↑0.0001.0000.3320.2610.4650.6161.0460.5621.943
Median ↑0.0001.0000.1340.0920.3130.3960.2890.2271.090
IQM ↑0.0001.0000.1300.1130.2800.3370.501N/AN/A
Optimality Gap ↓1.0000.0000.7290.7680.6310.5770.512N/AN/A

* MuZero / EfficientZero 属 lookahead search 类(用 MCTS),列出作对照;IRIS 无搜索。

  • 相对最强无搜索 baseline SPR:mean +70%、IQM +49%、optimality gap +11%、超人游戏数 +67%(6→10)。这是无 lookahead search 类的新 SOTA,且 IRIS 也超过了 MuZero(后者非为样本高效设定设计)。
  • rliable 统计口径(Agarwal 2021,分层 bootstrap 置信区间):IRIS 在其后 50% 游戏上与最强 baseline 持平,之后**随机占优(stochastically dominates)**其他方法;对所有 baseline 的 improvement probability 都 > 0.5。median 上与其他方法重叠(median 只被少数决定性游戏影响)。
  • 世界模型定性能力:Pong 中训练 120 局后即逐像素精确预测球轨迹与记分板;KungFuMaster 中能在不确定性下生成不同数量/种类的敌人、并复现「敌人被击中后消失」的机制;Breakout / Gopher 中准确预测正奖励帧与 episode 终止。

关键消融:

  • token 数(Appendix E):默认每帧 16 token;增到 64 token 在 3 个视觉复杂游戏上——Alien +36%、Asterix +121%、BankHeist +432%(Table 7),代价是算力/显存上升。说明视觉细节多的游戏受益于更多 token。
  • 数据规模(Appendix F, Table 8):把环境步从 100k 提到 10M(为控训练时间把优化:环境步比从 1:1 降到 1:50),mean 从 1.046 飙到 7.488、超人游戏 10→15、IQM 0.501→2.239——证明 IRIS 可扩展到样本高效设定之外。

作者指出的失败模式:① double exploration problem——当新关卡/机制要靠一个低概率事件触发时(如 Frostbite 造冰屋),世界模型没见过就学不到,策略也就无法在想象里发现它(Krull 因关卡切换频繁反而 IRIS 拿到该游戏 SOTA);② 视觉复杂、小物体重要的游戏,16 token 不够。

创新点与影响

贡献:① 提出 IRIS——第一个把「离散自编码器(image token)+ GPT 式自回归 Transformer」组合系统性地用作 RL 世界模型的 agent,把动力学学习整体转写成离散 token 序列建模问题;② 在 Atari 100k 上以 mean 1.046、10/26 超人成为无 lookahead search 类新 SOTA,甚至超过 MuZero;③ 证明自监督训练目标让标准 Transformer 无需 gating 等改造即可稳定学世界模型(对比 Parisotto 2020 直接用 RL 目标的困难);④ token 数、数据规模两个方向都给了清晰的 scaling 证据(64 token、10M 步);⑤ 全程 minimal tuning,开源代码与预训练模型。

影响:IRIS 把 Transformer + 离散 image token 立为 RSSM 之外世界模型 RL 的另一条主干范式,直接启发后续 Transformer 世界模型(TransDreamer、TWM、storm、Δ-IRIS 等)与更大规模的自回归/token 化世界模型。它与 dreamer-v2(提供 actor-critic 目标)、muzero / efficientzero(搜索路线对照)共同界定了 2022-2023 样本高效 RL 世界模型的版图。

作者自述局限:① double exploration——依赖低概率事件解锁的新机制学不进世界模型(Frostbite);② 策略目前从重建帧学,其实可以直接复用世界模型的内部表征(更省、更准);③ 视觉复杂游戏需更多 token(算力代价);④ 未与 MCTS 结合——作者认为「想象中学 + MCTS」两者贡献可能互补,是未来方向。

原始链接

一手源存档(sources/)

  • iris—github-readme — GitHub README 快照(含 tl;dr、训练/可视化命令、预训练模型说明)
  • arXiv 原文 PDF(2209.00588,arXiv 原文 PDF,不入 git)——见上方 arXiv 链接