一句话定位
把 muzero 里一笔带过的 Reanalyse 技巧做深做全:用学到的模型 + MCTS 重新分析(而非重新交互)历史数据来生成更好的策略/价值训练目标,通过调节 Reanalyse fraction(0%=纯在线交互 → 100%=纯离线、零环境交互)把在线 RL、数据高效 RL、离线 RL 统一成同一套不做任何离策略修正的算法——100% 档即为 MuZero Unplugged,在 RL Unplugged 离线 benchmark 上刷新 SOTA,同时顺带把 200M 帧标准设定下的 Atari 在线 SOTA 也刷新了。NeurIPS 2021。
背景与定位
离线 RL 此前的主流思路是约束/正则化:CRR 用 critic 过滤掉坏动作只从好动作学策略、REM 用 ensemble 随机凸组合正则化 Q 值、CQL 学一个保守 Q 函数下界、POPO 用悲观价值函数——这些方法共享的问题是都要专门处理离策略偏差。另一条路是model-based 离线 RL:MOReL 先学一个悲观 MDP 再在其中做策略优化、MOPO 用模型不确定性惩罚奖励、LOMPO 把 MOPO 扩展到图像输入——但这些工作普遍局限于低维状态/动作空间,且都只把学到的模型用于不确定性估计去训策略,从未直接用模型做前瞻搜索,也没有谁把这套方法用到 Atari 这种视觉复杂域。
本文的路径不同:直接沿用 muzero 的价值等价隐模型 + MCTS,把 MuZero 论文里一笔带过、仅用于离散动作数据效率提升的 “Reanalyse” 预览版做深——反复用最新网络参数对存量数据重跑 MCTS 生成新目标,构成一个自我强化的改进循环。当 Reanalyse fraction 推到 100%,学习完全基于存量轨迹、零环境交互,即离线 RL;由于目标生成方式(MCTS 重规划)和策略/价值本身的训练方式与在线情形完全一致,作者强调不需要任何针对离策略或离线场景的特殊改动——这是与上述所有正则化类方法的核心区别。连续动作空间的实验则依赖同期工作 sampled-muzero(Hubert et al., 2021)的采样式搜索扩展来处理高维/连续动作。
模型架构
复用 muzero 的三件套(表示函数 h、动力学函数 g、预测函数 f,MCTS 全程在隐空间跑),但做了两处更新:
- 骨干换代:从原版 MuZero 的普通残差块换成 ResNet v2 风格 pre-activation 残差块 + Layer Normalisation,优化器从 SGD 换成 Adam + decoupled weight decay(这一改动本身就显著提升了 Atari 分数,见 Infra/评测节 Table 7)。表示函数与动力学函数均为 10 个残差块,每块 2 层;图像输入每层 3×3 卷积、256 个平面;状态输入每层全连接、隐藏维 512。
- 连续动作支持:借用 sampled-muzero 的采样式搜索——策略头产出一组候选动作供 MCTS 搜索,而非枚举整个动作空间;策略只在被采样到的动作上被更新。连续动作用对角协方差的单高斯表示(作者未观察到混合高斯带来收益),训练目标是对数似然。
- Reanalyse 时的动作注入:离线数据/示教轨迹的行为策略往往与当前学到的策略差异很大,尤其在高维连续动作空间中,MCTS 采样式搜索很可能永远采不到数据集里出现过的动作。解决办法是把该轨迹的真实动作显式注入根节点的候选采样集合,先验权重与 Dirichlet 探索噪声的比例一致,取 25%(离散动作如 Atari 不受影响,因策略天然覆盖所有动作)。消融(Table 10)显示这一步对高维任务是决定性的:如 humanoid.run(21 维动作)不注入时仅 3.3 分,注入后跳到 633.4;manipulator.insert ball(5 维)从 21.2 → 557.0。
数据
- 数据效率实验(在线,Atari):不新增交互数据,只调 Reanalyse fraction,验证同一套算法能否覆盖跨数量级的数据预算(Table 1):fraction 50.0% / 95.0% / 99.5% 分别对应 2000M / 200M / 20M 帧的等效预算,median 人类归一化分 1331.7% / 1006.4% / 126.6%。
- 离线 benchmark:RL Unplugged(Gulcehre et al., 2020):
- Atari:46 个游戏,每局 200M 帧,离散动作、像素观测,数据由一个在线 DQN agent 生成,含 sticky actions(25% 概率重复上一动作,Machado et al., 2017 协议)。
- DM Control Suite:9 个任务,连续动作维度 1–21,状态观测,episode 数按任务从 40 到 3000 相差近百倍(如 cartpole.swingup 40 episodes、finger.turn hard 500、manipulator.insert ball/peg 各 1500、humanoid.run 3000)。
- 低数据消融:Atari 数据集裁到 **1%(2M 帧)**与 **10%(20M 帧)**子集,验证数据越多性能越好(Table 3)。
- Replay buffer:Atari 保留最近 50,000 条子序列,DM Control 保留最近 2,000 条,episode 切成长度至多 500 的子序列;采样用 prioritized replay,优先级
p_i=|νᵢ−zᵢ|(搜索价值与 n 步回报之差),重要性采样权重修偏,全部实验 α=β=1。 - 网络规模按数据量缩放:DM Control 各任务数据量相差百倍,为防止小数据集上模型过参数化导致记忆/过拟合,按
channels=√(datapoints/layers)缩放隐藏层宽度(附录 Table 11 消融了 64/128/256/512 四档隐藏维,验证过大的网络在 cartpole 这类小数据集上确实过拟合更严重)。 - 无独立的动作标注环节——数据本身就是 RL 轨迹(在线 agent 或专家策略生成),动作已内含。
训练方法
- Reanalyse 算法(Algorithm 1):从存量轨迹里采一个时间步 t,用当前参数对该状态重跑 MCTS 产生新的策略目标 πₜ、价值 νₜ,不选择/不执行任何动作——纯粹用于产生更优的训练标签,标签质量随参数迭代不断变好,形成正反馈。
- 损失形式:与原版 MuZero 一致,K 步展开求和
l_p(policy) + l_v(value) + l_r(reward);策略目标 = MCTS 搜索出的访问分布 πₜ₊ₖ,reward 目标 = 真实观测奖励 uₜ₊ₖ,value 目标 zₜ₊ₖ 是 5 步 TD(对 target network,每 100 步同步一次权重)在 Atari 上的选择;DM Control 因数据集普遍很小,改为直接回归到 MCTS 搜索价值估计以防止过拟合、让价值学习独立于轨迹——但消融(Table 8)显示数据更多时(20M 帧、99.5% Reanalyse fraction)5-step TD 反而优于 0-step TD(median 126.6% vs 115.3%,mean 450.6% vs 385.8%),即便此时学习几乎全离策略。 - CRR 复现作为强基线:critic 直接用 MuZero 的价值头(5-step TD 训练),CRR loss 训练策略头,结果显著优于其他离线基线但仍不及 MuZero Unplugged(Table 2)。
- 关键超参:Adam + decoupled weight decay,weight decay scale 1e-4;初始学习率 1e-4,按 cosine schedule 在 1,000,000 个训练 batch 内退火到 0;batch size 1024(所有实验统一);折扣 0.997(Atari)/ 0.99(DM Control)。
- 评估协议:训练结束(1M mini-batch)后用最终 checkpoint 评估,取 300 个评估 episode 的均值;连续动作评估时把高斯策略尺度压到接近确定性但仍可采样(
scale_eval=min(scale_predicted, 0.05)),以便 MCTS 仍能拿到一组不同候选动作。
Infra(训练 / 推理工程)
- 实现基于 JAX(复现自 Schrittwieser et al., 2020 的 MuZero),网络用 Haiku 库搭建。
- GPU/TPU 数量、并行方式、GPU-hours、推理 FPS / 控制频率 / 边缘硬件延迟:本文均未披露——通读全文(含附录 A–F)未出现任何算力配置或时延数字(对照原版 muzero 论文披露的”棋类 16 TPU 训练+1000 TPU 自对弈、Atari 8 TPU 训练+32 TPU 自对弈、12 小时/1M 步”,本篇后续工作完全没有重复给出这类工程细节)。
评测 benchmark
Atari 在线 RL,200M 帧标准设定(Table 7,57 游戏人类归一化分):
| 版本 | Median | Mean |
|---|---|---|
| LASER(此前 SOTA) | 431.0% | – |
| MuZero(原版) | 741.7% | 2183.6% |
| MuZero sticky actions | 692.9% | 2188.4% |
| MuZero Res2 Adam(本文架构改动) | 1006.4% | 2856.2% |
即仅换成 ResNet v2 + LayerNorm + Adam 就把 median 从 741.7% 推到 1006.4%,刷新 200M 帧设定下的 Atari SOTA。
Reanalyse 数据效率扫描(Table 1,57 游戏):Reanalyse fraction 50.0%/95.0%/99.5% → median 1331.7%/1006.4%/126.6%,mean 4094.4%/2856.2%/450.6%,对应等效帧预算 2000M/200M/20M——同一算法覆盖两个数量级的数据预算。
RL Unplugged Atari 离线 benchmark(Table 2,46 游戏,median/mean 归一化分):
| 方法 | Median | Mean |
|---|---|---|
| BC | 53.3% | 48.5% |
| DQN | 86.2% | 89.5% |
| IQN | 100.8% | 96.1% |
| BCQ | 107.5% | 120.0% |
| REM | 107.9% | 113.5% |
| CRR(本文复现) | 155.6% | 271.2% |
| MuZero BC | 54.0% | 46.9% |
| MuZero Unplugged | 265.3% | 595.5% |
逐局对比数据生成策略(在线 DQN),MuZero Unplugged 在 44/46 局达到或超过训练数据水平,仅 2 局略降(Figure 2)。
消融:动作选择方式 × 训练损失(Table 4,median 分,MCTS 评估行):监督模仿(unroll-5)169.7 → CRR 损失 172.5 → Reanalyse 损失 265.3,即完整 MuZero Unplugged(Reanalyse 损失 + MCTS 评估选动作)在同一评估方式下大幅领先;反过来固定 Reanalyse 损失、只换评估时的动作选择方式:采样动作 203.2 < 取最大 value 239.9 < MCTS 265.3,即 MCTS 评估动作选择在任一损失下都是最优选择。
低数据 Atari(Table 3,1%=2M 帧 / 10%=20M 帧,5 个游戏):MuZero Unplugged 在多数游戏上大幅领先 QR-DQN/REM/CQL(H)(如 1% 档 asterix 达 27220.5 vs 三个基线 166.3/386.5/592.4;10% 档 asterix 达 40554.0),但并非全面碾压——pong 上两档均落后于 CQL(H)(1% 档 MZ −16.2 vs CQL(H) 19.3;10% 档 MZ 15.6 vs CQL(H) 18.5),1% 档 qbert 上也不及 CQL(H)(MZ 6953.2 vs CQL(H) 14012.0)。
RL Unplugged DM Control(Table 5,9 任务均值):原始 Gulcehre BC 基线均值 365.3,D4PG 415.9、BRAC 398.5、RABM 486.8;用 MuZero 网络复现的 BC 均值 406.7(与原始 BC 基线大致吻合,验证了离线数据接入无误);MuZero Unplugged 601.9,在 humanoid.run(633.4 vs 各基线 1.7–408.5)、manipulator 系列等高难任务上领先明显;cartpole 等小数据集任务(仅 40 episodes)上 MuZero Unplugged(343.3)反而不及 D4PG/BRAC/RABM(856.0/869.0/798.0),作者归因于小样本导致训练后期过拟合(图 5/6 显示最优 checkpoint 明显好于训练末期)。与另一强基线 CRR 单独对比(Table 6,用其评测口径按训练中最高均值选 checkpoint):CRR 640.9 vs MuZero Unplugged 733.3。
创新点与影响
- 核心贡献:把 MuZero 论文中简略提及的 Reanalyse 技巧系统化、推向极限(fraction 100%),证明”用学到的模型 + MCTS 对历史数据重新规划以生成训练目标”这一单一机制,只需调节一个标量比例,就能无缝跨越在线 RL、任意数据预算的数据高效 RL、完全离线 RL 三种此前分别研究的设定,且离散/连续动作、像素/状态观测通吃,不需要任何离策略修正。
- 改变了什么:在 RL Unplugged 离线 benchmark(Atari + DM Control)上同时刷新 SOTA,是第一个把”模型直接用于规划”(而非仅用于不确定性估计)的方法成功推广到 Atari 这种视觉复杂离线域的工作;同时顺带把 200M 帧在线 Atari SOTA 也刷新了(ResNet v2 + LayerNorm + Adam 架构改动)。论文本身也观察到 Reanalyse fraction 数据效率曲线呈对数关系,呼应语言模型的 scaling law(Kaplan et al., 2020)。
- 作者自陈局限:花了”意外多”的时间排查离线数据集的动作空间不匹配、数据差异和压缩伪影问题,建议任何离线 RL 论文应先复现基线结果再改算法;cartpole 等极小数据集任务上仍有过拟合问题,作者提出 dropout 等正则化可能缓解但留作未来工作;连续动作 MCTS 需要额外的动作注入技巧才能在离线场景下工作,并非开箱即用;论文未探讨与已有正则化类离线 RL 方法(CQL/MOPO 等)的组合,留作未来工作。
原始链接
- arXiv(v1 2021-04-13):https://arxiv.org/abs/2104.06294
- PDF:https://arxiv.org/pdf/2104.06294
一手源存档(sources/)
- 未发现独立于 arXiv 论文的官方博客 / GitHub / HuggingFace / project page(DeepMind 未为本文单独发布代码或博客;已用 Bing 搜索确认,检索到的仅为原版 MuZero 博客与第三方解读,均非本文一手源)。
- arXiv 全文见上方链接(arXiv 原文 PDF,不入 git)。