一句话定位

UniZero 用一个基于 Transformer 的模块化隐世界模型替换 muzero 系(MuZero-style)算法里”编码器+递归动力学网络”的三件套骨架,把隐状态 z_t 和隐式历史 h_t 显式解耦,从而同时解决 MuZero-style 架构的两个缺陷——隐状态与历史信息纠缠(entanglement)、训练时只用初始观测导致轨迹数据利用不足——并借助 KV Cache 在推理期保留完整近期上下文;在需要长期记忆的 VisualMatch 基准、8→26 个 Atari 游戏的在线多任务学习、以及 DMControl 18 个连续控制任务上均取得优于 MuZero/DreamerV3 的表现,2025.06 被 Transactions on Machine Learning Research(TMLR 2025)接收,代码合入 lightzero 开源工具箱。

背景与定位

muzero(Schrittwieser et al., 2019)把 MCTS 规划与”编码器 h_θ / 动力学网络 g_θ / 预测网络 f_θ”三件套隐模型结合,在棋类和 Atari 上取得超人类表现,但其训练时只用初始观测 o_t(可能是堆叠帧)加整段动作序列作为输入,第 k 步隐状态 s_t^k 靠动力学网络递归 rollout 得到——这使得 s_t^k 与历史信息紧密纠缠,且与后续 EfficientZero(de Vries et al., 2021;Ye et al., 2021)提出的自监督正则化损失(对齐 s_t^k 与观测编码 z_t^k)在 POMDP 场景下互相冲突:MuZero w/ SSL 在 frame_stack=4(近似 MDP)下样本效率不错,但在 frame_stack=1(近似 POMDP,参考 Hausknecht & Stone, 2017 的 DRQN 框架)下 500k 步内不收敛。论文称这是”纠缠”问题;另一个直接把 k 步递归预测的隐状态当 MCTS 根节点(MuZero w/ Context)的朴素修复方案,则因递归预测的累积误差导致”不完整上下文”(incomplete context),同样表现不佳。同期用 GRU 替代 Transformer 骨架的 UniZero (RNN) 变体也因 GRU 有限记忆长度存在同样问题,论文未将其纳入正式 baseline。

同一研究团队(Shanghai AI Lab / SenseTime)在 MCTS+RL 方向还有 lightzero(统一开源 benchmark,UniZero 的官方实现宿主)与 rezero-mcts(用 backward-view + entire-buffer reanalyze 给 reanalyze 阶段提速,是与 UniZero 正交的训练流水线优化,不改动网络架构)。UniZero 与 TD-MPC 系(td-mpctd-mpc2)、Dreamer 系(dreamer-v3)、以及同为 Transformer 世界模型的 irisstorm 的根本区别在训练范式:后几者都是”先学世界模型、再用想象/MPC 学策略”的两阶段流水线,而 UniZero 延续 MuZero 的模型-策略联合优化(model-policy joint training)+ MCTS 策略提升,只是把骨架从 MLP/GRU 换成 Transformer(Appendix F 对比表,详见”创新点与影响”)。

版本说明:本页六个维度以 arXiv v2(2025-01-03,https://arxiv.org/html/2406.10667v2)全文为主要依据——相对 v1(2024-06-15)新增了 DMControl 连续控制评测、把多任务学习从 4 个游戏扩展到 26 个游戏、补齐了连续动作空间的 MCTS 扩展方案;v2 对应被 TMLR 2025 接收的版本(信息来源:lightzero 官方 GitHub README 2025.06 news 条目,非 arXiv 本身披露)。个别超参数 v1→v2 有变动(如默认 Transformer 层数 N:v1 为 4,v2 为 2),本页统一按 v2 为准并在正文标注。

模型架构

UniZero 是基于 Transformer 的模块化隐世界模型,四个组件(Eq. 2):

  • 编码器 h_θ:Atari/DMC 用 LightZero 同款卷积网络 + 末端线性层,把三维卷积特征映射为长度 D=768 的一维隐状态;VisualMatch 用更小的卷积网络,D=64(输入 3×5×5 → Conv1/BN1/LeakyReLU(16×5×5) → Conv2(32×5×5) → Conv3(64×5×5) → AdaptiveAvgPool2d(64×1×1) → Linear(64) → SimNorm(64))。编码器末端统一接 SimNorm(借鉴 td-mpc2):把 D 维隐状态划分成 L=D/V 个 simplex,每个 simplex 维度 V=8,做温度 τ=1 的 softmax 归一化,对隐空间做 L1 范数约束——消融显示这是训练稳定性的关键(Softmax 次之,Sigmoid 不收敛)。
  • 动力学头 g_θ:ẑ_{t+1}, r̂_t = g_θ(z_{≤t}, a_{≤t}),条件是到当前步为止的完整隐状态+动作序列(经 Transformer 骨架处理),而非 MuZero 式的单步递归;两层线性网络 + GELU,预测下一隐状态(维度=D,再过一次 SimNorm)和奖励(101-bin 离散回归,Bellemare et al. 2017 distributional RL 风格)。
  • 决策头 f_θ:p_t, v_t = f_θ(z_{≤t}, a_{≤t-1}),同样两层线性+GELU,价值/奖励用 101 个 bin 的离散回归,策略头输出维度=动作空间大小(离散)或高斯分布参数(连续,见下)。
  • Transformer 骨架:基于 nanoGPT 实现,每个时间步拆成两个 token——隐状态 token(SimNorm 归一化后)和动作 token(离散动作查表 nn.Embedding,v2 新增连续动作用两层 MLP 编码),加可学习位置编码(nn.Embedding)。默认配置:8 个注意力头(Atari/DMC)/4 个头(VisualMatch),层数 N=2(v2 Table 8 默认值,v1 为 4;消融另外扫过 H=5/10/20/40 时的层数-上下文长度联合影响),dropout 0.1,FFN 用 GELU。不使用 decoder 重建观测——消融证实加 decode regularization(L1 重建 loss + perceptual loss,系数 0.05)对 Pong 和 VisualMatch 性能均无明显影响,支持”隐状态只需保留决策相关信息”的设计假设。
  • 推理期 KV Cache:维护 KV_M = {KV(z_{t-H_infer}, a_{t-H_infer}, …, z_t, a_t)},H_infer=4(Atari/DMC 默认)或 memory_length+16(VisualMatch,需要覆盖整个 episode 才能记住探索阶段看到的目标颜色);新观测编码后作为 MCTS 根节点,动力学头在该 KV Cache 上递归展开内部节点。
  • 连续动作扩展(v2 新增,借鉴 Sampled MuZero/Sampled Policy Iteration, Hubert et al. 2021):策略头输出高斯分布参数 μ_θ(s)、σ_θ(s)(Eq. 9);MCTS 节点展开时从该高斯提议分布采样 K=20 个动作(而非枚举全部连续动作,Eq. 10);PUCT 公式里的先验 P(ẑ,a) 换成均匀分布 u(ẑ,a) 以避免对采样动作引入额外偏置;策略从 MCTS 访问计数分布 π̂_β 蒸馏进网络时改用 KL 散度损失 ℒ_policy = -Σ_i π̂_β(a_i|s)·log π_θ(a_i|s)(Eq. 11-13),而非离散动作用的交叉熵。

目标网络:维护一个 EMA(软更新,momentum 0.05)的 target world model W̄ = (h̄_θ, ḡ_θ, f̄_θ),用于产生 target 隐状态和 n-step TD value target;消融证实软目标最稳定,硬拷贝(每 100 次迭代)稳定性稍差,完全不用目标网络会导致 Pong 不收敛、VisualMatch 出现 NaN 梯度。

数据

在线自我博弈式强化学习,无预采集数据集,“数据规模”体现为交互步数与环境覆盖面:

  • Atari 100k(SimPLe 引入,Kaiser et al. 2024):26 个游戏,每局交互 100,000 步(跳帧 4,对应 400,000 环境帧)。观测格式 (3,64,64) 单帧 RGB(stack=1,UniZero 默认)或 (4,64,64) 灰度四帧堆叠(stack=4,MuZero baseline),区别于先前工作常用的 (4,96,96) 格式。
  • DMControl Proprio Control Suite(v2 新增):18 个连续控制任务(经典控制/运动/机械臂操作),预算 500,000 环境步,动作重复(action repeat)=2,固定 episode 长度 1000 步(有效决策步数 500),无终止条件。
  • VisualMatch 长期依赖基准(改编自 Ni et al. 2024 的 transformers_rl 设置):网格世界,探索阶段(1 步,随机 RGB 颜色房间)→干扰阶段(步数=memory_length,随机出现苹果、拾取无奖励)→奖励阶段(15 步,需选出与探索阶段颜色匹配的方块);目标色从 blue/red/green 三色中随机;智能体视野限制在 5×5 网格;奖励完全稀疏——只有奖励阶段成功才给 +1。与原始设置的区别:探索阶段从 15 步压缩到 1 步、干扰阶段拾取苹果不给分。
  • 多任务学习:先在 8 个 Atari 游戏(Alien/Boxing/ChopperCommand/Hero/MsPacman/Pong/RoadRunner/Seaquest)上验证,再扩展评测到全部 26 个 Atari 游戏;full_action_space=True 统一成 18 维离散动作空间;每个任务独立的数据 collector 与 replay buffer,训练时每任务采样 task_batch_size=32 拼成大 minibatch,per-task loss 取平均反传;共享 Transformer 骨架和编码器,仅解码头(决策头/动力学头)per-task 独立(借鉴 Kumar et al. 2022),编码器用 LayerNorm 而非 BatchNorm(v1 报告:因各任务早期内部状态统计差异大,BatchNorm 导致性能下降)。
  • Replay buffer 容量 1,000,000 条 transition,均匀采样;UniZero 默认不做数据增强(MuZero w/ SSL baseline 用数据增强=True)。
  • 无人工动作标注、无 sim-to-real、无跨模态数据混合——全部由环境模拟器在线生成轨迹并即时用于训练。

训练方法

目标函数(Eq. 3,联合优化模型与策略): next-latent 预测项(β_z·‖ẑ_{t+1} - sg(z̄_{t+1})‖²₂,target 由 EMA 编码器产生,stop-grad)+ reward 预测项(β_r·CE(r̂_t, r_t))+ policy 预测项(β_p·CE(p_t, π_t),π_t 是 MCTS 改进后的访问计数分布)+ value 预测项(β_v·CE(v_t, v̂_t),v̂_t 是 n-step TD bootstrap target)。系数取值(v2 Table 8):β_z=10;β_r/β_p 在 Atari/VisualMatch 为 1、DMC 为 0.1;β_v 在 Atari/VisualMatch 为 0.5、DMC 为 0.1。价值/奖励按离散回归(101 个 bin)建模。

训练循环(Algorithm 1):collect_experience(用 MCTS 改进策略与环境交互,存入 buffer)与 update_world_model(从 buffer 采样长度 H 的序列,联合梯度更新模型+策略+价值)交替进行;replay ratio=0.25(环境步与训练步之比);每次训练迭代后清空 KV Cache(因为世界模型参数已更新,旧缓存失效)。

MCTS 细节:PUCT 公式 a* = argmax_a[Q(ẑ,a) + P(ẑ,a)·√(ΣN(ẑ,b))/(1+N(ẑ,a))·(c₁+log((ΣN(ẑ,b)+c₂+1)/c₂))],c₁=1.25,c₂=19652;每次决策 50 次模拟;Dirichlet 噪声 α=0.3、权重 0.25;温度 T=0.25(访问计数 π_t ∝ N(z_t,a_t)^{1/T} 归一化)。

超参数(v2 Table 8,Atari/DMC/VisualMatch 统一配置,除标注外):训练上下文长度 H=10(消融扫过 5/10/20/40);推理上下文长度 H_infer=4(Atari/DMC)或 memory_length+16(VisualMatch);batch size 64;AdamW,学习率 1×10⁻⁴;weight decay 1×10⁻⁴;max grad norm 5;折扣因子 γ=0.997;policy entropy 系数 1×10⁻⁴;软目标更新动量 0.05;硬目标更新频率 100(仅用于”硬目标”消融变体,非默认);TD steps=5;Buffer Reanalyze Frequency:0(DMC/VisualMatch)、1/50(Atari)。DMC 连续动作额外:采样动作数 K=20,动作重复=2。MuZero baseline(w/ SSL)用 SGD、学习率按 schedule 0.2→0.02→0.002 衰减、batch size 256、SSL loss 系数 2、数据增强=True,其余超参与 UniZero 基本一致(Table 9)。

消融关键结论(Pong + VisualMatch,Section 4.5 / Appendix E.1):

  1. H_infer=4 全面优于 H_infer=8(不论 Transformer 层数),说明 Atari 短期任务不需要很长的推理上下文;VisualMatch 则要求训练上下文长度等于 memory_length+16 才能记住探索阶段的目标色。
  2. 更长的训练上下文 H 不总是带来更好效果——可能是 MCTS 推理误差随之增大所致;但配合更深的 Transformer,长上下文对表征学习有帮助,与”预测更远的未来有助于表征学习”的既有结论一致。
  3. SimNorm > Softmax > Sigmoid(Sigmoid 训练不收敛),凸显隐空间正则化对训练稳定性的重要性。
  4. Decode Regularization 对 Pong、VisualMatch 均无明显增益,支持”隐状态只需决策相关信息、无需重建观测”的假设。
  5. 目标网络消融:软目标(默认)最稳定;硬拷贝目标(每 100 迭代)出现一定不稳定;完全去掉目标网络导致 Pong 不收敛、VisualMatch 出现 NaN 梯度(与 DQN 里目标网络的作用类比一致)。
  6. 多任务梯度校正方法 PCGrad(Yu et al. 2020)、CAGrad(Liu et al. 2021)在初步实验中收益甚微,未纳入最终结果;反比于任务平均 episode 长度的采样策略、任务专属可学习嵌入两种尝试同样未见显著提升。

Infra(训练 / 推理工程)

  • 硬件:所有实验跑在 Kubernetes 集群上,单个实验实例配置为单张 NVIDIA A100 80GB GPU + 24 CPU 核心 + 100GB 内存;论文未披露任何多卡/分布式训练配置,也未讨论 UniZero 本身的多 worker 并行方案(这点与同实验室的 rezero-mcts 论文明确”留待未来工作”不同——UniZero 论文里未见对应讨论)。
  • 训练时长(wall-clock):Atari 单游戏训练到 100k 环境步约需 4 小时;VisualMatch(memory_length=500)跑 1M 训练步约需 30 小时。未披露 DMControl、多任务(8/26 游戏)设置下的具体 wall-clock 时长。
  • 精度/并行策略:未披露(论文未提及 fp16/bf16/混合精度或数据/模型并行细节)。
  • 推理 FPS / 控制 Hz / 边缘硬件:未披露——论文只给出 MCTS 每次决策 50 次模拟这一规划开销的代理指标,没有给出单步决策延迟(ms)或部署侧的实测帧率。
  • 总 GPU 卡时聚合数字:未披露(只有上述两个具体场景的 wall-clock 小时数,没有跨全部 26 游戏 + 18 DMC 任务 + 多任务实验的总卡时统计)。

评测 benchmark

VisualMatch(长期依赖,Figure 4):定性结果——MuZero 因缺乏上下文在所有 memory length 下表现都差;SAC-GPT(Ni et al. 2024 的 Transformer+SAC-Discrete 基线,训练 3M 环境步后的最终成功率)随 memory length 增加明显退化;UniZero 在 memory length 增加时保持稳定的高成功率。论文正文只以学习曲线图(Figure 4)呈现,未给出具体的成功率数值表

多任务学习,8 个 Atari 游戏(Table 2,400K 环境步,Normalized Mean/Median 为人类归一化分数):

算法AlienBoxingChopperCommandHeroMsPacmanPongRoadRunnerSeaquestNormed MeanNormed Median
UniZero (多任务)10035350130039891963007130.45540.4085
MuZero (多任务)590119891999999-158036000.21920.0895
UniZero (单任务)58032802299110121855037500.32230.1739

UniZero(MT) 在归一化均值和中位数上同时超过 UniZero(ST) 和 MuZero(MT)。进一步把 Transformer 骨架大小从 nlayer=4/8/12 扫描,8 个游戏上样本效率随模型增大一致提升(Appendix Figure 11)。扩展到全部 26 个 Atari 游戏的多任务训练(Appendix Figure 13)在归一化均值上与单任务训练基本持平(论文未给出具体数值表,只有学习曲线图)。

Atari 100k 单任务,26 games 完整对比(v2 Table 10,节选代表性行;Normalized Mean/Median 为全部 26 games 汇总):

GameRandomHumanMuZero (原始, Schrittwieser 2019)MuZero (Reproduced, stack4)UniZero (Ours, stack1)
Alien227.87127.7530.0300600
Boxing0.112.115207
BattleZone2360.037187.52688758711410
Kangaroo52.03035.0632001885
PrivateEye24.969571.356100500
Seaquest68.442054.7208466620
Normalized Mean (↑)0.0001.0000.560.440.39
Normalized Median (↑)0.0001.0000.230.130.22

UniZero(stack=1,单帧输入)在 15/26 个游戏上超过同框架下的 MuZero(Reproduced, stack=4),归一化中位数(0.22)高于 MuZero(Reproduced) 的 0.13(但低于原始 MuZero 论文报告的 0.23),验证”单帧输入即可同时建模短期和长期依赖”的核心论点;作者强调 UniZero 与 MuZero(Reproduced) 用完全相同的 LightZero 框架、相同超参、无逐游戏调参,因此二者可直接公平比较。

DMControl Proprio Control Suite,18 tasks vs dreamer-v3(v2 Table 3,人类归一化分数):

任务UniZeroDreamerV3
acrobot-swingup400.3154.5
cartpole-balance_sparse1000.0996.8
finger-turn_easy1000.0745.4
hopper-hop120.5111.0
Mean787.2743.7
Median875.1845.5

UniZero(借助 Sampled Policy Iteration 式连续动作扩展)在 18 个任务的均值与中位数上均超过 DreamerV3。

创新点与影响

  • 核心贡献:首次系统指出 MuZero-style 架构的两个根因性缺陷——隐状态与历史信息纠缠、训练轨迹数据利用不足——并用”Transformer 骨架显式分离隐状态 z_t 与隐式历史 h_t”这一单一架构改动同时解决两者,同时保留 MuZero 的 MCTS + 模型-策略联合优化训练范式(不像 Dreamer/TD-MPC/IRIS/STORM 系那样退化为两阶段流水线,见 Appendix F 对比表:TWM/IRIS/DreamerV3/STORM/TD-MPC2 均为 two-stage,只有 MuZero 与 UniZero 是 model-policy joint training)。
  • 不使用观测重建:消融证实 decode regularization 无收益,支持”隐状态只需保留决策相关信息”的设计选择,这与 Dreamer 系依赖重建损失塑造表征的做法形成对照。
  • 单一架构横跨异构场景:同一套 UniZero 架构和训练流程在短期依赖(Atari)、长期依赖(VisualMatch)、离散动作(Atari)、连续动作(DMControl,需扩展 MCTS 节点展开与策略蒸馏方式)、单任务与多任务(8→26 个 Atari 游戏共享一个模型)之间均无需结构性改动即可迁移,论文以此论证其作为”可扩展决策基础模型(scalable foundational model for decision-making)“的潜力。
  • 开源影响:代码合入 lightzero 官方 benchmark,成为该工具箱里 MuZero/EfficientZero 之外的标准算法变体(atari_unizero_segment_config.py),并衍生出 Sampled UniZero 变体(WIP,覆盖 DMControl/连续控制场景);2025.06 被 TMLR 2025 正式接收。
  • 作者自陈的局限(Section 6 / Appendix C “Future Directions”):(1) 多任务学习中任务平衡策略、MCTS 里的信息复用、多模态多任务统一框架、大规模预训练-微调方法均留待未来工作;(2) 尝试过的梯度校正(PCGrad/CAGrad)、反比例任务采样、任务专属嵌入均未见显著收益,原因未深入分析;(3) 未讨论多 worker/分布式训练场景下的扩展性(与同团队 rezero-mcts 论文形成对照,后者明确把这一点列为局限);(4) 论文没有给出跨全部评测场景的推理延迟/FPS 或总训练算力的聚合统计。

原始链接

一手源存档(sources/)