一句话定位

EfficientZero 在 muzero Reanalyze 之上加三件套——自监督时序一致性损失(SimSiam 风格)、端到端预测 value prefix(LSTM 预测「奖励前缀和」)、基于模型的 off-policy 校正——把 MCTS 类基于模型的世界模型 RL 首次推进到样本高效区间:在 Atari 100k(每个游戏仅 10 万环境步 ≈ 2 小时真人游戏时长)上取得 194.3% mean / 109.0% median 人类归一化分,第一个在如此少数据下同时在均值与中位数上超过人类,并把开源实现从 MuZero 的「64 TPU × 12 小时」压到「4 张 3090 × 7 小时」一个 agent。

背景与定位

样本效率是 RL 落地物理世界(机器人、医疗、推荐)的核心瓶颈:DQN 需 2 亿帧、AlphaZero 训练时要自对弈 2100 万局。基于模型的方法理论上更省样本(真实数据 + 模型「想象」数据都能训练策略),但在图像输入的低数据区间一直打不过 model-free:muzero、DreamerV2 虽超人却极不省样本,SimPLe 省样本但性能差(median 仅 0.144),而当时的 SoTA 是把自监督/数据增强套到 model-free 上的 SPR(median 0.415)。很多前人(van Hasselt 等)甚至怀疑图像输入下基于模型的方法到底能不能带来数据效率。

作者给出肯定答案,并把 MuZero 在低数据下失效归因于三个具体问题,一一对应三个改动:

  1. 环境模型缺少监督——MuZero 的动态模型只靠 reward/value/policy 三个标量信号训练,信号稀疏、value 有 bootstrap 噪声,不足以学好几百维的隐状态转移。→ 自监督一致性损失。
  2. 难以处理偶然不确定性(aleatoric uncertainty)导致的状态混叠(state aliasing)——递归 rollout 越深,reward 预测误差越累积;预测「哪一步恰好丢分」本身就是病态问题。→ 端到端 value prefix。
  3. 多步 value 的 off-policy 问题——MuZero Reanalyze 的多步 value target 用旧策略采的轨迹,数据受限时必须反复重用旧数据,off-policy 偏差被放大。→ 基于模型的 off-policy 校正。

技术脉络上它同时继承两条线:MCTS-as-policy-improvement 的 AlphaGo/AlphaZero/muzero 一脉,以及 SimCLR/MoCo/SimSiam/BYOL 的自监督表征一脉(与 dreamer-v1、PlaNet=planet 的隐空间世界模型属同代但不同范式)。作者自陈其一致性损失与 SPR 高度相似,区别是 SPR 用 BYOL 且只在 model-free 的 Rainbow 上做表征,而本文用 SimSiam 且把学到的模型真正喂给 MCTS 做探索与策略提升。NeurIPS 2021 收录。

模型架构

沿用 MuZero 的三网络流水线:表征网络 $\mathcal{H}$($s_t=\mathcal{H}(o_t)$)、动态网络 $\mathcal{G}$($\hat{s}_{t+1}=\mathcal{G}(s_t,a_t)$)、预测网络(reward/value/policy)。整体网络比 MuZero 更小(作者发现低数据下小网无容量瓶颈)。

输入:堆叠 4 帧历史(帧间隔 frame-skip 4,等效覆盖 16 帧游戏历史),按通道拼接为 $96\times96\times12$ 张量(4 帧 × RGB 3 通道)。

表征网络 $\mathcal{H}$(kernel 均 $3\times3$):conv stride2 32 通道(48×48,BN+ReLU)→ 1 residual block(32) → residual downsample stride2 64 通道(24×24)→ 1 resblock(64) → avgpool stride2(12×12)→ 1 resblock(64) → avgpool stride2(6×6)→ 1 resblock(64)。输出隐状态 $6\times6\times64$。

动态网络 $\mathcal{G}$:沿用 MuZero 结构但 residual block 从 16 个减到 1 个,并额外加一条 residual link 让递归推理时保留历史隐状态信息。结构:state 与 action 拼成 65 平面 → conv stride2 64 通道(BN) → 残差相加(ReLU) → 1 resblock(64)。

Action-conditioning:action 作为平面与隐状态拼接输入动态网络(MuZero 惯例);连续动作在 DMControl 上把每维离散成 5 档再喂 MCTS。

value prefix(奖励前缀)预测头:这是本文关键改动之一。不再逐步预测单步 reward 再求和,而是用 LSTM(hidden 512) 把展开状态序列 $(s_t,\hat{s}{t+1},\dots,\hat{s}{t+k-1})$ 端到端映到「奖励前缀和 $\sum_{i=0}^{k-1}\gamma^i r_{t+i}$」的 601 维分类 support(1×1 conv 16 通道 → flatten → LSTM(512) → FC(32) → FC(601))。训练时 LSTM 每步都被监督(每来一个新状态就能算 value prefix),故低数据也训得动;MCTS 内展开更深时 LSTM 隐状态在 $\zeta=5$ 步后重置。

reward / policy / value 头:resblock(64) → 1×1 conv 16 通道 → flatten → FC(32) → FC(D),value 用 601 维分类 support(MuZero 式 categorical value),policy 的 $D=$ 动作数。预测头末层权重/偏置置零以稳训练。

自监督一致性(SimSiam 风格):动态网络输出的 $\hat{s}{t+1}$ 应与真实下一观测的表征 $s{t+1}=\mathcal{H}(o_{t+1})$ 一致。做法:$\hat{s}{t+1}$ 过 projector $P_1$ 再过 predictor $P_2$,与 $s{t+1}$ 过 $P_1$(stop-gradient,作 target 分支)算负余弦相似度: $$\mathcal{L}{\text{similarity}}(s{t+1},\hat{s}{t+1})=\mathcal{L}2\big(\text{sg}(P_1(s{t+1})),\ P_2(P_1(\hat{s}{t+1}))\big)$$ $P_1$ 为 3 层 MLP、$P_2$ 为 2 层 MLP,隐层 512、输出 1024,层间加 BN(末层除外)。动态网络递归展开 5 步、对 $k=1..5$ 都拉近 $\hat{s}{t+k}$ 与 $s{t+k}$。这一致性直接经动态网络成形,无需额外解码器/模型。

总损失(单步 rollout 示意,展开 $l_{\text{unroll}}=5$ 步平均): $$\mathcal{L}t=\mathcal{L}(u_t,r_t)+\lambda_1\mathcal{L}(\pi_t,p_t)+\lambda_2\mathcal{L}(z_t,v_t)+\lambda_3\mathcal{L}{\text{similarity}}+c\lVert\theta\rVert^2$$ 其中 $\lambda_1=1,\ \lambda_2=0.25,\ \lambda_3=2$,reward/policy 交叉熵、一致性用负余弦相似度。

数据

纯在线 RL 自对弈采集,无任何外部/离线数据集:

  • Atari 100k(SimPLe 提出,26 个游戏):每个游戏允许 10 万环境步 = 40 万帧(frame-skip 4),约等于 2 小时真人游戏时长;作为参照 DQN 用 2 亿帧 ≈ 925 小时。人类基线也是在同样 2 小时熟悉游戏后测得。指标为人类归一化分 $(\text{score}{\text{agent}}-\text{score}{\text{random}})/(\text{score}{\text{human}}-\text{score}{\text{random}})$。
  • DMControl 100k(3 个低维连续控制任务:Cartpole Swingup、Reacher Easy、Ball-in-cup Catch),同样 10 万环境步。因 MCTS 不能直接处理连续动作,把每维离散成 5 档;为避免维度爆炸只选低维任务。
  • 自对弈时按 400 步为一段收集中间序列;replay buffer 用优先级采样($\alpha=0.6$,$\beta$ 从 0.4 退火到 1.0,$P(i)\propto p_i^\alpha$,$p_i$=训练时 value 的 L1 误差),但作者指出低数据下优先级只带来微弱提升。
  • reward 做 clipping,terminal-on-loss-of-life=True,单局最长 108K 帧。

训练方法

建立在 MuZero Reanalyze 之上,三大改动如上。三个核心机制的训练细节:

① 自监督一致性:同步展开 5 步、逐步拉近预测隐状态与真实观测表征(见架构)。消融显示这是三件套里最关键的一件(去掉后 Normed Mean 从 1.943 跌到 0.881);且改进主要来自自监督损失本身而非数据增强(去掉小幅度随机平移 0–4 像素 + 强度扰动的增强,性能几乎不变)。

② 端到端 value prefix:把 UCT 里 $Q(s,a)=\sum_{i=0}^{k-1}\gamma^i r_{t+i}+\gamma^k v_{t+k}$ 的奖励求和项整体交给 LSTM 端到端预测,绕开「精确预测哪一步丢分」的病态子问题。在半训练 Pong 模型 rollout 出的静态 100k 数据集上对照:直接逐步 reward 预测训练误差更低,但 value prefix 在展开 5 步时验证误差明显更小,即避免过拟合硬 reward 预测、缓解状态混叠。

③ 基于模型的 off-policy 校正:用动态时域 $l\le k$ 的旧轨迹奖励 + 在末状态 $s_{t+l}$ 用当前策略重跑 MCTS 取根节点均值 value: $$z_t=\sum_{i=0}^{l-1}\gamma^i u_{t+i}+\gamma^l,\nu^{\text{MCTS}}{t+l},\qquad l=\Big(k-\big\lfloor\tfrac{T{\text{current}}-T_{s_t}}{\tau T_{\text{total}}}\big\rfloor\Big).\text{clip}(1,k)$$ 轨迹越旧 $l$ 越小(少 rollout 几步以减少策略发散),$k=5,\ \tau=0.3,\ T_{\text{total}}=100\text{k}$。消融(以 UpNDown 为例):带校正后 target value 对真值的 L1 误差全面下降(全状态 0.657→0.569),且轨迹越旧误差越大、校正收益越明显(20k 阶段 0.657→0.569,100k 阶段 0.441→0.397);进一步拆解发现「动态时域」比「MCTS 根值重搜」更重要。

MCTS 细节:$N_{\text{sim}}=50$ 次模拟;UCT 常数 $c_1=1.25,\ c_2=19652$;对未访问节点用 ELF OpenGo 式 mean-Q 机制给初值(而非默认 0)以改善探索;backup 用带阈值 $\epsilon=0.01$ 的 soft min-max 归一化避免低数据下 min-max 区间过窄导致的过度自信;根节点加 Dirichlet 噪声($\rho=0.25,\ \xi=0.3$);输出访问计数分布温度在训练 50%/75% 处退火到 0.5/0.25。

关键超参(Atari,Table 6):discount 0.9974、minibatch 256、SGD lr 0.2(100k 处降到 0.02)momentum 0.9 weight-decay 1e-4、max grad norm 5、unroll 步 =5、TD 步 $k=5$、训练 120K 步(仅前 100K 步采数据,让后段轨迹被充分利用)、评估 32 seeds、min replay 2000、self-play 网络更新间隔 100 / target 网络 200、reanalyze policy 比例 0.99(value 100%)。

Infra(训练 / 推理工程)

  • 训练算力:一个 Atari agent 训 100k 步只需 4 张 GPU(3090)× 约 7 小时;作为对照 MuZero 训一个 Atari agent 要 64 TPU × 12 小时。作者把「高质量、算力友好的开源实现」本身列为贡献。
  • 分布式框架:基于 Ray + 双缓冲(double buffering)。四类 actor 并行——self-play 数据 worker(用 600 训练步内的模型自对弈,把轨迹送 replay buffer)、CPU rollout worker(在 CPU 上准备 batch 上下文)、GPU batch worker(用 target 模型 reanalyze 过去数据、跑 MCTS)、learner。CPU/GPU worker 数按吞吐匹配,主瓶颈在 reanalyze 模块。
  • MCTS 工程:用 C++ 实现 selection/expansion/backup、Python 做神经网络推理、Cython 桥接两侧上下文;实现 batch MCTS 并行搜一批树;单 GPU 上并置多个 batch 计算线程(仿 ELF OpenGo)。用 torch AMP 加速。20G 显存不足时可给每个 GPU worker 分配 0.25 卡。
  • 推理 / 控制频率、边端部署:论文未披露(面向 Atari/DMControl 基准,非实机部署)。

评测 benchmark

Atari 100k(26 游戏,3 runs × 32 eval seeds),人类归一化 mean / median:

方法Normed MeanNormed Median
Random0.0000.000
Human1.0001.000
SimPLe0.4430.144
OTRainbow0.2640.204
CURL0.3810.175
DrQ0.3570.268
SPR(前 SoTA)0.7040.415
MuZero(作者复现,同超参)0.5620.227
EfficientZero1.9431.090
  • 头条:194.3% mean / 109.0% median,相对前 SoTA(SPR)分别领先 176% / 163%(Fig.1、Table1 口径;正文 §5.2 另给 170% / 180%)。第一个用仅 2 小时游戏数据在均值与中位数上双超人类。
  • 与 DQN 参照:DQN 用 500× 数据(2 亿帧)才 220% mean / 96% median;EfficientZero 已逼近其性能。
  • 26 游戏中 14 个超过人类;部分单局分远超人类(如 Asterix 25557.8 vs 人类 8503.3、CrazyClimber 83940 vs 35829、KungFuMaster 30944 vs 22736),个别仍逊于人类(如 Seaquest 1100.2 vs 42054.7、Frostbite 296.3 vs 4334.7)。
  • 用 Agarwal 等的 rliable 稳健聚合指标(mean/median/IQM/optimality-gap,95% CI)复核,EfficientZero 在四项指标上均显著领先。

DMControl 100k(10 seeds,episode return)

任务CURL(前 SoTA)DreamerMuZeroState SAC(用真值状态,oracle)EfficientZero
Cartpole Swingup582±146326±27218.5±122835±22813±19
Reacher Easy538±233314±155493±145746±25952±34
Ball-in-cup Catch769±43246±174542±270746±91942±17

即便与直接读取真值状态的 State SAC(视作 oracle)相比也可比甚至反超,而同样离散化的 MuZero 表现差。

三件套消融(26 游戏,Normed Mean / Median):Full 1.943 / 1.090;去一致性 0.881 / 0.340(跌最多);去 value prefix 1.482 / 0.552;去 off-policy 校正 1.475 / 0.836。作者结论:一致性提供的丰富学习信号是 MuZero 低数据下最缺的;value prefix 在早期学习更有用;off-policy 校正是专为低数据设计、数据充足时非必需。用带解码器重建可视化:无一致性时展开预测状态 $\hat{s}_{t+k}$ 无法重建回观测,有一致性时可以,佐证一致性缩小了表征网络与动态网络输出之间的分布漂移。

创新点与影响

  • 贡献:在 MuZero 上的三处外科手术式改动(SimSiam 时序一致性 / LSTM 端到端 value prefix / 基于模型的 off-policy 校正),首次让 MCTS 类基于模型的世界模型 RL 在图像输入 + 极限低数据下超过人类,正面回答了「基于模型方法能否带来图像输入数据效率」这一长期质疑。
  • 改变了什么:把 Atari 100k 的天花板从 SPR 的 0.415 median 抬到 1.090,并把「MuZero 太贵」的门槛从 64 TPU 降到 4 张消费级 3090,开源了含 C++/Cython batch-MCTS 的可复现框架,推动了后续 MCTS-RL 研究的普及。value prefix 与 mean-Q / soft min-max 等 MCTS 稳定化技巧被后续工作沿用。
  • 作者自陈局限:连续动作只能靠「每维离散成 5 档」处理,维度易爆炸,故 DMControl 只测 3 个低维任务,缺乏更好的连续动作设计;MCTS 仍较慢;未涉及终身学习。作者把这些列为 future work。
  • 后续:同组 EfficientZero V2(ICML 2024 Spotlight)把方法扩展到离散与连续控制统一框架,正是对本文连续动作局限的回应。

原始链接

一手源存档(sources/)

  • efficientzero—github-readme.md — 官方 GitHub README 快照(4 GPU 训练命令、C++/Cython + Ray 架构、依赖)
  • arXiv 2111.00210 全文 PDF(arXiv 原文 PDF,不入 git;正文所有数字取自此 PDF 全文含附录 A.1–A.6)