一句话定位
Recall to Imagine(R2I)把状态空间模型(SSM,S4/S5 一脉)塞进 dreamer-v3 的 RSSM 世界模型里,做出一个叫 S3M(Structured State-Space Model)的新序列骨干:用可并行扫描(parallel scan)+ 可重置隐状态的对角 MIMO SSM 替代 GRU 循环核,在保持 DreamerV3”一套超参走天下”的通用性前提下,把长程记忆与长程 credit assignment 能力大幅拉升——在 BSuite、POPGym 上刷新 SOTA,在复杂 3D 记忆迷宫 Memory Maze 上首次做到超越人类,同时训练速度比 DreamerV3 快最多 9 倍。ICLR 2024 口头报告(作者自述 top-1.2%)。
背景与定位
Model-based RL(MBRL)里的世界模型骨干长期在 RNN 与 Transformer 之间二选一:RNN(如 dreamer-v3 用的 GRU-RSSM)有梯度消失问题,长程依赖学不动;Transformer 虽然语言建模强,但自注意力是序列长度的平方复杂度,训练长序列时不稳定,也扩不到 RL 可能需要的长上下文。同期,SSM(S4 系列)在监督/自监督序列任务上证明能建模上万步的依赖且训练/推理都是次二次复杂度,于是很自然地被想到搬进世界模型。
R2I 是 DreamerV3 谱系(dreamer-v3 的直接后继)的记忆增强分支:保留 DreamerV3 的整体 RSSM 结构(表征模型 + 动力学模型 + 序列模型 + 观测/奖励/续episode 预测头、actor-critic 在想象轨迹里训练),只是把序列模型从 GRU 换成一叠 SSM 层(S3M),并把表征模型改造成非循环形式以便时间维完全并行计算。论文点名的并发工作是 S4WM(s4wm-world-model-backbones,Deng et al. 2023)——同样把 S4 塞进 Dreamer 式世界模型,但 S4WM 只做世界模型本身的预测准确度(图像重建 MSE),不训练 RL 策略、也不报告 reward;R2I 论文用附录整节(Appendix P/P.1,含逐项对比表 Table 5)说明”世界模型似然更高不等于 RL 表现更好”,并列出与 S4WM 不同的多项设计选择(见下方评测部分)。
模型架构
S3M(Structured State-Space Model)序列骨干:在 DreamerV3 RSSM 的基础上,用 SSM 替换 GRU 作为序列模型 fθ。
- 非循环表征模型(non-recurrent representation model):DreamerV3 里 zₜ ~ qθ(zₜ|hₜ,oₜ) 依赖上一步的 hₜ,导致必须逐步串行计算;R2I 去掉这个依赖,变成 zₜ ~ qθ(zₜ|oₜ),使所有时间步的后验样本能独立、并行地一次性算出,这是让整段序列的 h₁:T 能被并行扫描算出来的前提(附录 M 的消融显示这个改动不损失甚至提升性能)。
- SSM 层内部结构:每层 SSM 按 Eq.2 的离散线性递推算完之后接 GeLU → 全连接 GLU(gated linear unit)→ LayerNorm(沿用 Smith et al. 2023 S5 的架构),采用 post-norm(DreamerV3/在线 RL 场景下比 pre-norm 更利于泛化,论文初步实验证实)。最后一层 SSM 的输出即确定性状态 hₜ;所有 SSM 层的隐状态集合记为 xₜ。
- SSM 设计选择(对照 S4WM,Table 5):矩阵参数化用 对角(Diagonal)而非 DPLR(性能相近但 DPLR 想象步计算慢 2–3 倍);维度上用 MIMO(多输入多输出,而非 SISO)——省去混合层、参数更省;离散化用 双线性(bilinear)(S4 常用的 Woodbury 恒等式离散化每步要矩阵求逆,R2I 一早就放弃了 S4 走 bilinear);计算模式用 parallel scan(而非 S4 原始的全局卷积模式)——因为只有 parallel scan 才能同时吐出策略需要的隐状态 x₁:T,也只有它能支持 batch 内多 episode 的状态重置(修改 Lu et al. 2024 的 associative 二元算子,加入 done 标志位项,clip 到 [0,1]),且能把序列长度维度跨设备并行扩展。
- 三档模型尺寸(Table 2,配置名称照论文原文,注意 Atari/DMC 用的这一档论文本身就叫”Small”,和 BSuite/POPGym 那档的”Small Memory”是两个不同名字,不是”中”配置):
| 配置(论文原名) | 用于 | hₜ 维度 | xₜ 维度/层 | SSM 层数 | SSM units |
|---|---|---|---|---|---|
| Small Memory | BSuite / POPGym | 512 | 512 | 3 | 1024 |
| Small | Atari / DMC | 512 | 192 | 5 | 512 |
| Medium Memory | Memory Maze | 2048 | 512 | 5 | 1024 |
- 动作条件化:与 RSSM 一致,序列模型每步吃 (aₜ₋₁, zₜ₋₁) 作为输入更新 (hₜ, xₜ):ht,xt = fθ((at−1,zt−1), xt−1),只是 fθ 内部换成了多层 SSM。
- 隐变量:与 DreamerV3 相同的多类别(multi-categorical)随机隐变量——32 个类别变量 × 每个 32 类,unimix 概率 0.01。
- 图像/表格环境编解码:图像用 CNN encoder/decoder;表格(向量)观测用 MLP。
- 策略输入的三种变体:output-state policy π(â|z,h)、hidden-state policy π(â|z,x)、full-state policy π(â|z,h,x)。三者在不同 domain 里表现不同——非记忆环境(Atari、DMC-proprio、BSuite)用 output-state 更好;POPGym 用 hidden-state 更好;Memory Maze 用 full-state 更好;DMC-vision 用 hidden-state。论文强调这不是因为哪种输入信息量更大,而是因为 hₜ、zₜ、xₜ 三者在训练过程中分布会漂移,把全部塞进策略反而因非平稳性伤害稳定性(附录 N 有逐环境消融)。
数据
R2I 是纯在线 RL 设定,没有预训练数据集,“数据”即环境交互产生的经验,存入 FIFO replay buffer(大小 10⁷ = 1000 万步),在全部 5 个 domain 上统一使用这个偏大的 buffer(DreamerV3 默认是 100 万步)以稳定 SSM 世界模型、防止在小 buffer 上过拟合。
评测覆盖 5 大 RL domain:
- BSuite:Memory Length(须记住初始观测直到 episode 结束)、Discounting Chain(首个动作触发的奖励延迟给出,考验 credit assignment)。
- POPGym(Table 4,三种任务 × Easy/Medium/Hard):RepeatPrevious 需同时记住 k 个类别值(Easy 4 / Medium 32 / Hard 64);Autoencode 分”观察-复现”两阶段,需记住 episode 前半段全部观测(episode 长度 Easy 104 / Medium 208 / Hard 312,即需记住 T/2 = 52/104/156 个类别值);Concentration(模拟翻牌配对游戏,每步观察含多个类别、每类别 N 个取值——类别数 Easy/Hard 52、Medium 104,取值数 N Easy/Medium 3、Hard 14;episode 长度取 Morad et al. 2023 给出的”最优解所需最小步数”,本文未给具体步数;正文另指出 Concentration Easy/Medium 需要同时记住的信息量最多约 208 个(Table 4 给出的具体值为 Easy 104、Medium 208),且该任务可被无记忆策略部分求解)。
- Atari 100K、DMC(proprio + vision):作为泛化性 sanity check,本身不特别需要长程记忆。
- Memory Maze(Pasukonis et al. 2022):4 种迷宫尺寸 9x9/11x11/13x13/15x15,每 episode 最长 4000 环境步,agent 需记住已探索过的墙体布局、物体位置和自身位置。
数据构成不涉及离线数据集混合、也不涉及 sim-to-real,全部是在线策略与环境交互产生的经验回放。POPGym 的超参搜索(附录 L)里额外确认了:把 replay buffer 从 DreamerV3 默认 100 万步增大到 1000 万步这个改动,对 R2I 有效但对 DreamerV3 反而更不稳定(方差更大)。
训练方法
沿用 DreamerV3 的整体训练范式:collect–train world model–train actor-critic in imagination 的在线循环。
- 世界模型目标(Eq.3–7):与 DreamerV3 相同的 ELBO 变体——预测损失(观测/奖励/续episode 对数似然)+ 动力学损失 + 表征损失,两个 KL 项各自 clip 在阈值 1(free bits)并做 KL balancing,用缩放系数 β_pred=1、β_dyn=0.5、β_rep=0.1 加权(Table 3)。SSM 状态矩阵没有像部分 SSM 文献建议的那样用更小学习率单独调,实验发现和世界模型其余部分共用同一学习率(1e-4)效果最好。
- actor-critic 训练:完全在想象轨迹里训练(imagination horizon H=15),离散动作空间用 REINFORCE,连续控制(DMC)用穿过学到的动力学反传梯度;λ-return(λ=0.95,折扣 γ=0.997)、critic 用 twohot 回归 + EMA(衰减 0.98)、return 按 95/5 百分位差归一化(衰减 0.99)、固定熵奖励系数 3e-4——这套流程完全照搬 DreamerV3(附录 D)。
- 关键工程决策:(1) 打断表征模型对 hₜ 的依赖,换成非循环 qθ(zₜ|oₜ),才能让 parallel scan 一次性算出整段 h₁:T、x₁:T,避免像 RNN/Transformer 那样需要多步 burn-in(否则想象阶段复杂度会退化成 O((L+H)²));(2) 选 parallel scan 而非卷积模式,核心原因是只有前者才输出策略需要的隐状态 x₁:T,且支持 batch 内多 episode 的状态重置;(3) 除 SSM 尺寸(Table 2)外,其余超参数在全部 5 个 domain 上固定不变,延续 DreamerV3”零调参”的通用性主张。
- 超参数细节(Table 3):Batch length L=1024,batch size=4(即每个训练 batch 覆盖 4×1024=4096 个时间步);世界模型梯度裁剪 1000,Adam epsilon 1e-8;actor-critic 梯度裁剪 100,Adam epsilon 1e-5;SSM 离散化范围 (1e-3, 1e-1),HiPPO 矩阵分块数 8。
Infra(训练 / 推理工程)
- 硬件:Memory Maze 实验用 2 张 NVIDIA A100(40GB),通过 batch-wise 数据并行分布训练;另配 40 个环境 worker 进程做数据采集,其中一个 rollout worker 与训练进程共享一张训练 GPU。
- 吞吐:系统总吞吐约 350 FPS;训练强度为”每 1 个采样步对应 51 个回放步”,论文原话为”82 environment steps per one gradient update of the world model and actor-critic (since the batch size is 4096)”。
- 达到超人类水平所需算力(Table 6,Appendix T):
| 迷宫尺寸 | 原始分数 | 人类归一化 | Oracle 归一化 | 环境步数 | 墙钟天数 | A100 天数 |
|---|---|---|---|---|---|---|
| 9x9 | 33.55 | 127% | 96% | 17M | 0.6 | 1.2 |
| 11x11 | 51.96 | 117% | 89% | 66M | 2.2 | 4.4 |
| 13x13 | 58.14 | 104.7% | 78% | 206M | 6.9 | 13.8 |
| 15x15 | 40.86 | 60% | 46% | 未达到 | 未达到 | 未达到 |
| 均值 | 46.13 | 102% | 77% | — | — | — |
- 主实验预算:Memory Maze 主结果(Figure 5)在 400M 环境步或两周训练后截止评测;R2I 因计算效率更高,在同样墙钟时间内比 DreamerV3 积累更多环境步,实测训练速度提升最多 9 倍(Figure 2/26/27)。
- 软件栈:JAX 实现(dm-haiku + flax + optax),直接基于 Danijar Hafner 的 DreamerV3 JAX 代码库扩展;SSM 并行扫描实现参考 the Annotated S4 博客与 Smith et al. 2023 的 S5 代码库。GitHub 仓库要求单张支持 CUDA 的 GPU 起步,Memory Maze 默认 1000 万步的图像 replay buffer 需要约 130GB 磁盘空间。
- 推理侧 FPS/延迟未单独披露(论文只报告训练吞吐,未区分推理阶段单独指标)。
评测 benchmark
- BSuite(Figure 3,10 个随机种子,中位数 + 25/75 百分位):此前 SOTA DreamerV3 在 Memory Length 和 Discounting Chain 上只能应对最长约 30 步的奖励延迟;R2I 把这个上限推到约 100 步,且在更宽的延迟范围内保持高成功率。
- POPGym(Figure 4):R2I 在 Autoencode-Easy/Medium、RepeatPrevious-Medium/Hard 上刷新 SOTA,超过 DreamerV3 及 POPGym 全部 13 个 model-free baseline(其中 PPO+GRU 是最强 model-free baseline,PPO+LSTM 次之);Concentration-Easy 上与 DreamerV3 持平,Concentration-Medium 略优于 DreamerV3(该任务本身可被无记忆 MLP 策略部分解决)。
- Memory Maze(Figure 5 + Table 6):400M 步 / 两周训练后,R2I 全面超过 DreamerV3 与 model-free 的 IMPALA(Pasukonis et al. 2022 中该 domain 最强 model-free 方法);9x9 与 Dreamer 相近但明显强于 IMPALA,11x11/13x13/15x15 上明显强于两者;人类归一化分数在 9x9/11x11/13x13 上均超过 100%(超越人类),15x15 未达到人类水平(60%)。
- Atari 100K / DMC(Figure 6,Figure 25 用 RLiable 库画的 performance profile):R2I 与 DreamerV3 表现几乎一致——Atari 100K 上 profile 高度重合;DMC 的 proprio/vision 上 R2I 在 return 500–900 区间的比例略低,其余区间差异很小。结论是记忆能力的提升没有牺牲在非记忆任务上的通用性。
- 与 S4WM 的架构对比(Table 5,Appendix P.1):论文用表格逐项列出与并发 SSM 世界模型工作 S4WM(s4wm-world-model-backbones)的 11 项设计差异,并在正文中逐条展开了其中 9 项——计算模式(parallel scan vs 卷积)、归一化位置(post-norm vs pre-norm)、先验/后验 SSM 是否共享参数(共享 vs 不共享)、SSM 后的变换(GLU vs 线性层)、SSM 维度(MIMO vs SISO)、参数化(对角 vs DPLR)、离散化方法(bilinear vs 论文未说明)、世界模型训练目标(与 DreamerV3 相同 vs KL 项无 free info)、episode 重置处理(有 vs 无实现);表格还额外列出隐状态是否对策略可见(R2I 有、S4WM 无实现)与策略训练模式(R2I 用 DreamerV3 objective、S4WM 为 None)两项。最核心的一条:S4WM 只报告世界模型的图像预测准确度(MSE 更低),不训练 RL 策略、不报告 reward;R2I 用 DreamerV3 同款 objective 完整训练策略并给出 RL 分数,附录 P 专门论证”世界模型似然更高不代表 RL 表现更好”。
- 消融:非循环表征模型不损失甚至提升 Memory Maze 表现(附录 M);策略输入选择高度依赖 domain——记忆密集型环境几乎都需要把 xₜ(SSM 隐状态)喂给策略,非记忆环境反而用 output-state 更稳(附录 N);POPGym 上对 DreamerV3 做了单独的网络尺寸超参搜索以保证对比公平(附录 L)。
创新点与影响
R2I 是第一个把 SSM(S4 家族)用作世界模型骨干并同时训练出可用 RL 策略的 MBRL 方法(并发的 S4WM 只做世界模型预测,没有 RL 结果)。核心工程贡献是把 RSSM 的表征模型改造成非循环形式,从而让整段序列可以用 parallel scan 一次性并行算完,同时仍能把 SSM 的原始递归隐状态 xₜ 暴露给策略——论文的消融证明这个”策略能看到 xₜ”恰恰是解决长程记忆任务的关键,而这一点是 S4WM 的卷积模式做不到的。最终效果:在 Memory Maze 上首次做到超越人类,同时训练速度比 DreamerV3 快最多 9 倍,并且在全部 5 个 domain 上沿用同一套(除网络尺寸外)固定超参数,继承了 DreamerV3”通用智能体”的定位。ICLR 2024 接收为 oral(作者/仓库自述 top-1.2%,未见第三方复核)。
论文自陈的局限(Conclusion):世界模型训练所用的 batch 序列长度、以及 actor-critic 想象阶段的 horizon,目前都还不算”极长”,未来可以在这两个维度上进一步拉长;另外作者认为 SSM 与注意力机制存在互补性(如语言建模里已出现的混合架构),把注意力机制引入 R2I 是可探索的方向。
原始链接
- arXiv: https://arxiv.org/abs/2403.04253
- PDF: https://arxiv.org/pdf/2403.04253
- GitHub: https://github.com/chandar-lab/Recall2Imagine
- 项目主页: https://recall2imagine.github.io/
- OpenReview: https://openreview.net/forum?id=1vDArHJ68h
一手源存档(sources/)
- recall-to-imagine-r2i—github-readme — GitHub README 存档(sources/world-model/2024/recall-to-imagine-r2i—github-readme.md)
- recall-to-imagine-r2i—project — 项目主页存档(sources/world-model/2024/recall-to-imagine-r2i—project.md)
- arXiv 原文 PDF,不入 git(https://arxiv.org/pdf/2403.04253)