一句话定位
给 muzero 的 MCTS 装一个”采样 + 重要性加权”的补丁:节点展开时不再枚举全部动作,而是从策略网络采样 K 个动作,用修正后的先验 π̂β 做 PUCT 搜索,使 MuZero 第一次能用到围棋(362)、Atari(18)之外的高维/连续动作空间——DM Control 最高 56 维的人形机器人——上,且 K 远小于动作空间大小时(围棋 K=50/362、Atari K=3/18、56 维人形 K=3~5)性能已能逼近穷举全动作空间的基线。
背景与定位
muzero 的价值等价隐模型 + MCTS 在围棋、国际象棋、将棋、Atari 等小/中等动作空间上验证有效,但其 MCTS 核心的 PUCT 公式默认每个节点能枚举全部动作;真实世界的动作空间(连续控制的关节角度、机器人多维扭矩等)往往无法枚举。此前处理连续/复杂动作空间的主流是 model-free 路线:DDPG/D4PG(确定性策略梯度+分布式)、TRPO/PPO(信赖域/裁剪的策略优化)、SAC/MPO/AWR((相对)熵正则化的 actor-critic),它们都不具备 MuZero 式的前瞻规划能力。零星把 AlphaZero/MuZero 搬到连续动作的尝试(A0C 用 REINFORCE 估计梯度,只验证 1D Pendulum;Yang et al. 2020 只验证 1-2 维动作)尚未在高维任务上证明有效;(Tang & Agrawal, 2020) 的因子化策略表示(每个动作维度用独立的类别分布)能避免离散化带来的指数爆炸,本文直接复用了这一表示。
本文提出一个更一般的”sample-based policy iteration”框架:只要策略改进算子是 action-independent 的(MuZero 的 MCTS 改进、MPO/PPO/AWR 皆属此类),就可以把策略改进/评估写成对完整改进策略 Iπ 的期望,用采样到的 K 个动作及重要性加权来无偏估计这个期望。论文将该框架具体实例化到 MuZero 上得到 Sampled MuZero。同期工作 muzero-unplugged 的连续动作实验直接依赖本文提出的采样式搜索。
模型架构
网络骨干与标准 muzero 一致——表示函数 h / 动力学函数 g / 预测函数 f 三件套,联合训练;本文唯一的结构性改动集中在 MCTS 的先验替换上,网络架构本身只是针对连续控制输入做了卷积→全连接的适配:
- 表示函数(1 维状态输入,非 Atari/棋类的 2 维图像输入):输入块 = 线性层 + LayerNorm + tanh;随后接 ResNet v2 风格 pre-activation 残差塔,10 个残差块,每块 2 层,隐藏宽度 512,LayerNorm + ReLU。
- 动力学函数:动作块 = 线性层 + LayerNorm + ReLU,动作嵌入与表示函数输出的嵌入相加后,输入同架构(10 块×2 层×512)的残差塔。
- value/reward 预测:沿用 MuZero 的类别化(categorical)表示,51 个 bin,value 覆盖 [-150, 150],reward 覆盖 [-1, 1]。
- policy 预测:主结果用 (Tang & Agrawal, 2020) 的因子化类别分布——每个动作维度独立表示为 B=7 个 bin 的类别分布;附录额外测了高斯参数化(在 hard/manipulator 任务上表现相近,但优化更难,需要 5e-3 系数的熵正则化才能训好,作者未对高斯版本做正则化调优)。
- Real-World RL (RWRL):在表示函数的输入块和残差塔之间插入 LSTM 处理部分可观测性,8 步截断 BPTT,每步拼接最近 4 个观测,有效展开步数 32。
核心算法修改(相对标准 MuZero 的搜索):
- 节点展开时不再返回全部 N=|A| 个动作,而是从提议分布 β 采样 K 个动作,同时返回每个动作对应的 π(s,a) 和 β(s,a)。
- PUCT 公式里把先验 π 换成 π̂β ∝ (β̂/β)·π(β̂ 是 K 个采样动作的经验分布,即 Kronecker/Dirac delta 的均值),其余照搬标准 MuZero 的 PUCT:
argmax_a Q(s,a) + c(s)·π̂β(s,a)·√(ΣN(s,b))/(1+N(s,a)),c(s)=c1+log((1+ΣN(s,b)+c2)/c2),c1=1.25, c2=19652(与标准 MuZero 相同)。 - 论文证明(Theorem,附录 F):用 π̂β 做先验搜索得到的 visit-count 分布 I_π̂β 近似等于对”完整动作空间搜索得到的 visit-count 分布 Iπ”做 sample-based 改进算子后的结果,即 I_π̂β ≈ Î_β π;换言之搜索之外(训练策略网络、n-step bootstrap 训练价值网络、acting)的整套 MuZero 流程都不需要改动。
- 采样分布取 β=π(可加温度调节),并在 β 和 π 两处都加 Dirichlet 噪声,保证低先验动作仍有机会被采到和搜到。
数据
纯自对弈/自交互 RL,无预先采集的数据集,动作空间大小是核心自变量:
- 围棋:19×19 棋盘 + 停一手,动作空间大小 362;测试 K ∈ {15, 25, 50, 100}。
- Atari (Ms. Pacman):动作空间大小 18;测试不同 K,K=2 不足以有效学习,K=3 起性能迅速逼近穷举基线。
- DM Control Suite:沿用 Acme (Hoffman et al., 2020) 的 easy/medium/hard 任务分级和数据预算,额外评测 manipulator 任务;主结果用 K=20 训练。
- DM Control(像素输入):与状态输入相同任务集,数据预算 25M 帧;与 Dreamer 对比时用 (Hafner et al., 2019) 定义的 20 个任务、500 万环境步设置。
- Real-World RL (RWRL) Suite:cartpole/walker/quadruped/humanoid 四类任务各设 easy/medium/hard 三档难度。
- dm_control Locomotion(CMU 人形):56 维动作的人形机器人控制(forage / go-to-target / run-walls / run-gaps 等目标)。
- Replay buffer:保留最近 2000 条序列,episode 切成长度最多 500 的子序列;采样按 MuZero 同款 prioritized replay(相同超参)。
训练方法
目标函数与标准 MuZero 完全一致(policy 用 visit-count 分布做交叉熵目标,value/reward 用 n-step return 回归),改动只集中在 MCTS 内部的先验替换(π→π̂β),训练流程按域分别配置:
- 围棋:值目标改用 n=25 步 bootstrap(不同于 MuZero 原论文直接 bootstrap 到终局),且对目标网络在 n+i, i∈[0,3] 的 4 个连续时间步预测取平均,以降低双人博弈中因视角交替导致的 value 过拟合;训练时搜索预算降到 400 次模拟/步(标准 MuZero 用 800,为省算力),评测时恢复到 800 次模拟/步;Elo 标尺锚定到 MuZero 基线的最终表现为 2000 Elo。
- Atari:架构、优化器、超参与标准 MuZero 完全相同;评测搜索预算 50 次模拟/步。
- DM Control / RWRL:Adam 优化器 + decoupled weight decay(权重衰减系数 2e-5),batch size 1024,初始学习率 1e-4,用 cosine 退火在 100 万个训练 batch 内衰减到 0;训练时 K=20 个采样动作,搜索预算 50 次模拟/步;根节点额外优化——在搜索开始前先对全部 K 个采样动作评估一次,用其结果初始化 PUCT 公式里的 Q(s,a);评测时用 100 局、搜索预算 50 次模拟/步,取访问数最高的动作。
- 基线对比公平性:DMPO/D4PG 基线复现自 Acme,但本文为公平对比把网络加大(policy 网络层 (512,512,256,128),critic 网络层 (1024,1024,512,256))并统一 batch size 为 1024,每个任务跑 3 个随机种子。
- 消融 1(采样数 K):在 humanoid.run(21 维动作)上测 K∈{3,5,10,20,40},K=3 已足以学出好策略,K>10 后性能不再明显提升。
- 消融 2(π̂β vs π):在 humanoid.run 上对比 K=5/K=20、以及是否在根节点预评估 Q 值(no-Q)。结果:用 π̂β 明显比直接用 π 更稳健,尤其在采样数小、且不做根节点 Q 值预评估时差距最大;即使做了根节点 Q 值预评估,π̂β 依然更优。
Infra(训练 / 推理工程)
- 网络用 Haiku(JAX) 实现。
- 训练用加速器数量、总 GPU/TPU-hours、混合精度设置:论文正文与附录均未披露(与 MuZero 原论文披露 TPU 数量不同,本文未给出对应数字)。
- 推理侧只披露了搜索预算(围棋评测 800 次模拟/步,Atari/DM Control/RWRL 评测 50 次模拟/步),未换算为墙钟延迟、控制频率(Hz)或边缘硬件指标——未披露。
评测 benchmark
围棋(Figure 2,1 个随机种子/实验):MuZero 基线用全部 362 个动作搜索,Elo 锚定该基线最终表现为 2000。Sampled MuZero 随 K 增大单调逼近基线,K=50 已经很接近全动作基线。
Atari Ms. Pacman(Figure 3,1 个随机种子/实验):动作空间 18,K=2 学习效率不足,K=3 起迅速逼近全动作基线。
DM Control Suite(状态输入):3 个随机种子/实验,对比 DMPO (Hoffman et al., 2020) 与 D4PG (Barth-Maron et al., 2018);在 hard 和 manipulator 类任务(如 humanoid.run)上表现尤其突出(Figure 4,完整结果见附录 Figure 7)。
DM Control Suite(像素输入):同任务集、同数据预算(25M 帧)、同超参,学习曲线与状态输入版本相近(Figure 5);对比 Dreamer (Hafner et al., 2019),用其定义的 20 任务 / 500 万环境步设置,Sampled MuZero 在全部 20 个任务上持平或超过 Dreamer,且未使用 action repeat(Dreamer 用 action repeat=2)、未做观测重建、未针对每个任务重调超参(Table 2 逐任务数值对比)。
56 维 CMU 人形 Locomotion(Figure 6,1 个随机种子/实验):Sampled MuZero 在 forage、go-to-target、run-walls(对比 Merel et al., 2019)以及 run-gaps(对比 Song et al., 2020)上均超过此前报告的最优结果,且使用的环境交互次数比对比结果少一个数量级以上。
Real-World RL (RWRL) Challenge(Table 1,3 个随机种子/实验;DMPO/D4PG 结果取自 (Dulac-Arnold et al., 2020),STACX 取自 (Zahavy et al., 2020)):
| 任务/难度 | DMPO | D4PG | STACX | SMuZero |
|---|---|---|---|---|
| Cartpole Easy | 464.05 | 482.32 | 734.40 | 861.05 |
| Walker Easy | 474.44 | 512.44 | 487.75 | 959.83 |
| Quadruped Easy | 567.53 | 787.73 | 865.80 | 987.20 |
| Humanoid Easy | 1.33 | 102.92 | 1.21 | 289.36 |
| Cartpole Medium | 155.63 | 175.47 | 398.71 | 516.69 |
| Walker Medium | 64.63 | 75.49 | 94.01 | 448.51 |
| Quadruped Medium | 180.30 | 268.01 | 466.43 | 946.21 |
| Humanoid Medium | 1.27 | 1.28 | 1.18 | 108.56 |
| Cartpole Hard | 138.06 | 108.20 | 135.26 | 244.71 |
| Walker Hard | 63.05 | 59.85 | 58.11 | 71.16 |
| Quadruped Hard | 144.69 | 280.75 | 351.56 | 348.09 |
| Humanoid Hard | 1.40 | 1.27 | 1.26 | 1.19 |
Sampled MuZero 在绝大多数任务/难度组合上领先(论文正文总结为”在三档难度上都显著优于基线”),但需要如实指出:Humanoid Hard 这一格是唯一例外——四种方法的分数都在 1.2~1.4 附近(该任务本身接近所有方法都失败的水平),Sampled MuZero(1.19) 反而略低于 DMPO(1.40),属于奖励接近失败下限时的噪声区间,而非该方法的普遍劣势。
创新点与影响
- 核心贡献:提出 sample-based policy iteration 框架——把策略改进/评估写成对完整改进策略 Iπ 的期望,用 K 个采样动作 + 重要性加权(β̂/β)无偏估计该期望;证明该估计量 Î_β π 随 K→∞ 依分布收敛到 Iπ,且渐近正态、方差 ∝ 1/K(定理见 Sec 4.4,附录 E)。
- 具体化为 Sampled MuZero:证明把 MuZero PUCT 公式里的先验 π 换成 π̂β 之后,搜索得到的 visit-count 分布近似等价于对完整动作空间搜索结果做 sample-based 改进(Sec 5.1 定理,附录 F),因此除了采样这一步,MuZero 的训练/acting/价值学习流程无需任何其他改动——是一个算法上很轻量的扩展。
- 改变了什么:首次让 MuZero 式”学到的模型 + MCTS 前瞻规划”能够应用到穷举不可行的高维/连续动作空间(DM Control hard/manipulator、56 维 CMU 人形 Locomotion),且在这些任务上打平或超过专门为连续控制设计的 model-free 方法(DMPO、D4PG)以及同为 model-based 的 Dreamer;证明了即便 K 远小于动作空间大小,性能损失也很小(围棋 K=50/362、Atari K=3/18、56 维人形 K=3~5 起即可)。
- 作者自陈局限:框架仅覆盖 action-independent 的策略改进算子(MuZero 的 MCTS 改进、MPO/PPO/AWR 一类),未声称覆盖所有策略迭代算法;根节点之外的其他节点仍固定采样 K 个动作,论文明确指出对高访问量路径动态采样更多动作(如 progressive widening)是可行的未来方向而非本文已做的工作;连续动作的高斯参数化比因子化类别分布更难优化,需要额外熵正则化系数(5e-3)才能训练稳定。
原始链接
一手源存档(sources/)
- 官方未发布配套代码仓库、博客或项目页(github_url / hf_url / project_url 均核实为空,与同期 muzero-unplugged 一致)。
- 论文全文见上方 arXiv 链接(arXiv 原文 PDF,不入 git)。