一句话定位
DIAMOND(DIffusion As a Model Of eNvironment Dreams)是把扩散模型直接当世界模型、并让 RL 智能体完全在这个扩散世界模型的”想象”里训练的工作:在竞争激烈的 Atari 100k 基准上拿到 1.46 的平均人类归一化分(mean HNS),创下”完全在世界模型内训练的智能体”的新纪录;同时证明同一个扩散世界模型可以脱离 RL 单独作为可玩的神经游戏引擎——在 87 小时静态 CS:GO 对局数据上训练后,得到一个键鼠可实时操控、在 RTX 3090 上以 10Hz 运行的 Dust II 地图神经模拟器。NeurIPS 2024 Spotlight。
背景与定位
DIAMOND 属于基于模型的强化学习(model-based RL)+ 世界模型在想象中训练这一范式(Ha & Schmidhuber 2018 提出、SimPLe 引入 Atari 100k)。它要挑战的是当时主流世界模型的一个共同选择:把观测压成离散潜变量序列再建模动力学——DreamerV2/V3 的 RSSM 离散潜、iris(IRIS,同组 Micheli 等)的离散自编码器 + 自回归 transformer token、STORM/TWM 的 transformer 化 tokenization。论文的核心论点是:这种离散压缩会丢掉对 RL 关键的视觉细节(远处的红点奖励、敌人 vs 奖励的几像素差别),从而损害策略学习。
作者把当时图像生成领域”扩散模型正在取代离散 token 建模”的范式迁移(Ho et al. DDPM、Karras et al. EDM、Stable Diffusion)搬到世界模型里:用连续图像空间的条件扩散模型直接预测下一帧观测,绕开离散化的信息瓶颈。扩散模型天然易于条件化、能无模式坍缩地建模多模态分布,这两点对世界模型都很关键(更贴合动作条件 → 更可靠的信用分配;更丰富的多模态想象 → 更多样的训练场景)。
在”神经游戏引擎”这条并行脉络上,DIAMOND 与同期的 deepmind-genie(Genie,潜动作 2D)、GameNGen(扩散版 DOOM)、Decart/decart-oasis(Minecraft)同处一个浪潮;DIAMOND 的特色是同一套扩散世界模型既服务 Atari 的 RL 想象训练、又能 scale 到 3D 的 CS:GO 实时可玩引擎。范式命名:diffusion world model / RL in imagination(扩散世界模型 + 想象中强化学习)。
模型架构
DIAMOND 的世界模型 = 一个动力学扩散模型 Dθ + 一个奖励/终止模型 Rψ,再配一个在想象中训练的 actor-critic 智能体 (π,V)ϕ。三者分别是独立网络。
动力学扩散模型 Dθ(backbone = 标准 2D U-Net,走 diffusion-forcing 式逐帧自回归)
- 扩散范式:EDM(Karras et al. 2022),而非 DDPM。 这是全文最关键的架构决策。采用 EDM 的网络预处理(preconditioning):Dθ 参数化为噪声观测与网络 Fθ 输出的加权和
Dθ = c_skip·x^τ + c_out·Fθ(c_in·x^τ, y),其中 c_skip = σ_data² / (σ_data² + σ²(τ))。这种信号/噪声自适应混合让模型在高噪声时被训练去预测干净图像,从而用极少去噪步(甚至 1 步)就能稳定生成,避免 DDPM 在低步数下的复合误差(compounding error)。 - 动作条件(action-conditioning):保留一个长度 L=4 的历史缓冲区(过去 4 帧观测 + 4 个动作)。**过去观测按通道维拼接(frame stacking)**到待去噪的下一帧噪声观测上;动作通过残差块里的 adaptive group normalization(AdaGN)注入;扩散时间 τ 也通过 AdaGN 注入。
- latent / tokenizer:无。DIAMOND 直接在 64×64×3 像素图像空间做扩散,不用 VAE/离散码本——这正是它保住视觉细节的来源。
- memory / 一致性:靠 frame-stacking(4 帧)提供短时记忆;作者在 Limitations 明确指出这是”最小记忆机制”,未来应换成沿环境时间的自回归 transformer(如 DiT 式)以获得长时记忆。附录 M 试过 cross-attention 架构,但早期实验里 frame-stacking 更有效。
- 采样器:Euler 方法(一阶),n=3 去噪步(NFE=3/帧),未用高阶采样器或随机采样。
- 配置:残差块 layers=[2,2,2,2]、channels=[64,64,64,64]、条件维度 256。Atari 版扩散模型 约 4M 参数。
奖励/终止模型 Rψ:独立的 CNN + LSTM(处理部分可观测性),残差块 channels=[32,32,32,32]、LSTM 维度 512;输入帧+动作序列,想象前先 burn-in 4 帧初始化 LSTM 隐状态。奖励用符号三分类交叉熵(reward ∈ {−1,0,1})。
actor-critic (π,V)ϕ:共享 CNN-trunk(4 个残差块 + 2×2 max-pool stride 2)+ LSTM(维度 512),策略头与价值头分开;策略用 REINFORCE + 价值基线,价值网络用 λ-returns 的 Bellman 误差训练(类 IRIS)。整套系统合计约 13M 参数(扩散 4M + 奖励/终止 + actor-critic)。
数据
分两条实验线,数据完全不同:
Atari 100k(RL 想象训练)
- 遵循 Atari 100k 协议:每个游戏智能体只允许在真实环境里采 100k 个动作(考虑 frameskip=4,约等于 2 小时人类游戏经验)。对照:常规无约束 Atari 智能体训 50M 步(约 500 倍经验)。
- 数据是 on-policy 在线采集:智能体在真实环境采数据 → 用至今收集的全部数据更新世界模型 → 在更新后的世界模型里用 RL 训智能体 → 循环。每 epoch 采 100 个环境步、训练 400 步、batch 32、共 1000 epoch。采集用 ε-greedy(ε=0.01)。
- 评测覆盖标准 26 个 Atari 游戏,每个 5 个随机种子。
CS:GO(神经游戏引擎,纯世界模型,无 RL)
- 用 Pearce & Zhu (2022) 的 Online 数据集:5.5M 帧(95 小时)在线人类对局,16Hz 采集,全部来自 Dust II 地图。
- 随机留出 0.5M 帧(500 局 / 8 小时)做测试,其余 5M 帧(87 小时)训练。无 RL 智能体、无在线数据采集——纯离线静态数据上的世界模型学习。
- 分辨率处理:世界建模时把原始 280×150 降到 56×30,再用一个更小的扩散上采样模型恢复到原分辨率。
数据配比 / 清洗 / 过滤细节:两条线都未涉及复杂 mixture(Atari 是单游戏 on-policy,CS:GO 是单一数据集)。
训练方法
- 世界模型目标:EDM 的去噪重建 L2 损失(条件在过去观测+动作+扩散时间上):
L(θ)=E[‖Dθ(x_{t+1}^τ, τ, x_{≤t}^0, a_{≤t}) − x_{t+1}^0‖²]。噪声等级 σ(τ) 从经验选定的 log-normal 分布采样(集中在中噪声区)。 - 三网络分别优化、想象中训练 RL:
- 价值网络用 λ-returns(λ=0.95)做回归目标的 Bellman 误差;
- 策略用 REINFORCE + 价值基线 + 熵正则(熵权 η=0.001);
- 想象 horizon H=15,折扣 γ=0.985。智能体只在想象里训练,只在真实环境采数据。
- 关键设计消融:
- EDM vs DDPM:在 Breakout 100k 帧静态数据上同架构对比。DDPM 在 ≤10 去噪步时自回归到 t=1000 步会严重漂移出分布;EDM 即使 n=1 步也长时程稳定。这是选 EDM 的核心依据。
- 去噪步数选择:Breakout 这类确定性转移单步即可;但 Boxing 等部分可观测游戏观测分布多模态,单步预测会取期望→模糊/OOD,需迭代求解器锁定某个模态,故全实验统一 n=3。
- 优化超参:AdamW,学习率 1e-4,ε=1e-8;扩散/奖励模型 weight decay=1e-2,actor-critic weight decay=0。图像 64×64×3,frameskip=4,max noop=30,life loss 触发终止,奖励裁剪到 {−1,0,1}。
- CS:GO 训练:把 U-Net 通道数放大,参数从 Atari 的 4M 提到 381M(含 51M 上采样器);动力学模型仍只用 3 去噪步,但上采样器引入随机采样并加到 10 去噪步以提升画质,在画质与推理成本间取平衡。
Infra(训练 / 推理工程)
- 硬件:全部实验用单张消费级 GPU。Atari 每个 run 约占 12GB VRAM,在单张 Nvidia RTX 4090 上约 2.9 天训完;26 游戏 ×5 种子合计 约 1.03 GPU-年。
- 参数量/训练时间对比(附录 H, Table 4):DIAMOND 13M 参数 / 2.9 天 / mean HNS 1.459;IRIS 30M / 4.1 天 / 1.046;DreamerV3 18M / <1 天 / 1.097。即 DIAMOND 参数比 IRIS、DreamerV3 都少,训练比 IRIS 快、比 DreamerV3 慢,但分数最高。
- 单次更新耗时(RTX 4090,附录 I):一次完整更新 543ms = 扩散模型 88ms + 奖励/终止 115ms + actor-critic 340ms;其中 15 步想象每步 20.4ms(下一帧预测 12.7ms = 3×4.2ms 去噪步 + 奖励/终止 7.0ms + 动作 0.7ms)。一个 epoch 约 217s。
- 推理成本:Atari 每帧 3 NFE(对比 IRIS 每帧 16 NFE),同为 64×64 分辨率。
- CS:GO 推理:组合模型 381M(含 51M 上采样器,即动力学约 330M) 在 RTX 4090 上训 12 天;推理时在 RTX 3090 上以 10Hz 实时运行,可键鼠操控。
评测 benchmark
Atari 100k(26 游戏、5 种子,Table 1 与 Figure 2)——完全在世界模型内训练的智能体横向对比:
| 指标 | SimPLe | TWM | IRIS | DreamerV3 | STORM | DIAMOND |
|---|---|---|---|---|---|---|
| Mean HNS ↑ | 0.332 | 0.956 | 1.046 | 1.097 | 1.266 | 1.459 |
| IQM ↑ | 0.130 | 0.459 | 0.501 | 0.497 | 0.636 | 0.641 |
| 超人游戏 ↑ | 1 | 8 | 10 | 9 | 10 | 11 |
- DIAMOND mean HNS = 1.46(新 SOTA),IQM 0.64(与 STORM 持平且高于其余全部基线),在 11 个游戏超越人类。
- 在”捕捉小细节很重要”的游戏上尤其突出:Asterix 3698.5(STORM 1028、IRIS 854)、RoadRunner 20673.2(STORM 17564)、Breakout、Pong 20.4。也有明显短板:BankHeist 19.7(远低于 DreamerV3 649)、BattleZone 4702(IRIS/DreamerV3/STORM 均 1.2–1.35 万)。
- 与 IRIS 的定性对比(Fig 5):同静态数据集训两模型,IRIS 的想象轨迹在帧间出现视觉不一致(Asterix 里敌人↔奖励反复互变、Breakout 砖块/分数跳变、RoadRunner 奖励点闪烁),DIAMOND 无这些不一致——Breakout 里打碎红砖分数还能可靠 +7。作者强调这不是靠更多算力:同分辨率下 DIAMOND 每帧仅 3 NFE vs IRIS 16 NFE、参数更少、训练更快。
- 未与 model-free/搜索类 SOTA(BBF、EfficientZero)直接比——附录 J 说明它们用了正交技术(周期性重置+超参调度、昂贵的 MCTS lookahead),不直接可比。
- CS:GO:无定量 benchmark,以可玩性 + 视频质量(键鼠实时游玩 Dust II、10Hz)作定性证据。
创新点与影响
- 核心贡献:首次系统论证图像空间条件扩散模型可以做稳定、高效的世界模型,并给出让它可行的关键设计——用 EDM 而非 DDPM(自适应信号/噪声混合 → 低去噪步下长时程稳定,破解自回归复合误差)、n=3 去噪步在画质与成本间取平衡、frame-stacking + AdaGN 的动作条件。由此拿下 Atari 100k 世界模型内训练的新 SOTA(1.46 HNS)。
- “视觉细节 matters”的实证:把离散潜世界模型(IRIS 等)丢失的几像素细节,与 RL 策略学习的成败直接挂钩——奖励/敌人的一致渲染、分数的可靠更新,都是像素级扩散带来的可测收益。
- 同一模型 = 可玩神经游戏引擎:CS:GO Dust II 的实时(10Hz、键鼠)神经模拟器,是当年”fully neural playable game engine”浪潮的代表作之一,与 GameNGen(DOOM)、Genie、Oasis 并列,把世界模型从”RL 内部工具”推向”可交互媒介”。全部代码/预训练权重/可玩世界模型开源(GitHub + HF
eloialonso/diamond)。 - 作者自陈局限:① 主评测局限于离散控制环境,连续控制域待验证;② frame-stacking 是最小记忆机制,长时记忆与可扩展性应换成沿环境时间的自回归 transformer(DiT 式);③ 奖励/终止预测目前是独立网络,把它整合进扩散模型很不平凡(从扩散模型抽表征困难),留作未来工作。
原始链接
- arXiv abstract:https://arxiv.org/abs/2405.12399
- arXiv PDF:https://arxiv.org/pdf/2405.12399
- 项目页(视频/demo):https://diamond-wm.github.io/
- GitHub 代码:https://github.com/eloialonso/diamond (CS:GO 在
csgo分支、DDPM 对比在ddpm分支) - Hugging Face 预训练权重:https://huggingface.co/eloialonso/diamond
- 发表:NeurIPS 2024 Spotlight
一手源存档(sources/)
- diamond—github-readme — GitHub README 快照(sources/world-model/2024/diamond—github-readme.md)
- diamond—project-page — 项目页快照(sources/world-model/2024/diamond—project-page.md)
- arXiv 全文(HTML v2)已通读,未入 git;引用见上方 arXiv URL(arXiv 原文 PDF/HTML,不入 git)