一句话定位

PWM 复用同一个预训练 TD-MPC2 世界模型,把策略抽取方式从「在线 MPC 采样规划」换成「对世界模型直接做一阶梯度(First-order Gradient, FoG)反传」的离线策略提取:每个任务 <10 分钟内学出的策略,不仅超过在线规划的 TD-MPC2、无模型的 PPO/SAC,甚至超过直接用真实仿真器梯度的 SHAC——核心反直觉发现是:世界模型对策略学习的价值不在预测精度,而在于它诱导出的优化地形是否足够光滑、最优性 gap 是否够小。

背景与定位

  • ICLR 2025 会议论文(所有页眉标注 “Published as a conference paper at ICLR 2025”),作者 Ignat Georgiev、Varun Giridhar、Animesh Garg(Georgia Institute of Technology)与 Nicklas Hansen(UC San Diego,td-mpc2 一作)。论文全文未出现 NVIDIA 字样,仅用 Nvidia RTX6000/RTX3090 GPU 做实验——作者机构应为 Georgia Tech + UC San Diego 两家。
  • 直接建立在 td-mpc2 之上:复用其世界模型架构(encoder/dynamics/reward 三件套 + SimNorm 隐归一化 + NormedLinear + 离散回归奖励)与其公开的多任务离线数据集/checkpoint,但改变了策略如何从世界模型里被抽取出来。TD-MPC2 本身用 DDPG 式梯度 ∇θJ(θ)≈E[∇θQ(s,a)] 学一个策略先验,再在推理时跑 MPPI 采样规划(零阶梯度,ZoG);PWM 反其道而行——把预训练好的世界模型当作纯粹的可微分物理模拟器,用类似 differentiable-simulation 流派(Xu et al. 2022 SHAC/dflex;一作后续 AHAC, ICML 2024)的 on-policy actor-critic,直接对 H 步隐空间 rollout 做一阶梯度反传,训练结束后策略前向推理即可用,无需在线规划。
  • dreamer-v3 对比:论文把 DreamerV3 也归为 FoG 阵营(其 actor 训练同样对想象 rollout 反传梯度),但 DreamerV3 专注单任务在线学习,未处理多任务场景;PWM 明确聚焦「大规模多任务预训练世界模型 + 每任务快速抽取专家策略」这一新框架。
  • 提出的核心命题反直觉:更精确的世界模型不等于更好的策略。通过 ball-wall(接触不连续)和双摆(混沌动力学)两个教学式玩具例子(Appendix A/B),论文论证 FoG 优化真正在意的是模型的光滑性最优性 gap,而非拟合误差本身——这与 Suh et al. (2022)、Lambert et al. (2020) 的 objective-mismatch 观察一脉相承,但把结论落到了「应该怎么设计/正则化世界模型」上。
  • 多任务框架层面的转变:不再像 TD-MPC2/GATO 那样训练一个统一的多任务策略,而是「先预训练一个多任务世界模型,再为每个任务单独抽取一个专家策略」——论文用消融证明这个解耦设计是必要的(单一多任务策略在 MT30 上失败,见下)。

模型架构

World model(WM)组件与配置直接沿用 td-mpc2 的架构(细节见该页),PWM 论文只强调其针对差分策略学习所需的关键设定:

  • 三件套 Eφ(s,a→e latent), Fφ(z,a,e), Rφ(z,a,e):全连接 MLP,LayerNorm + Mish 激活(NormedLinear),/ 输出层用 SimNorm(维度 V=8)隐归一化。奖励头是离散回归(SymLog 空间两热编码,101 个 bin)。
  • 48M 参数配置(MT30/MT80 主结果所用):latent 维 z=768,encoder 隐层 [1792,1792,1792],dynamics 隐层 [1792,1792],reward 隐层 [1792,1792],任务嵌入维 96。世界模型训练 horizon H=16、batch 1024、学习率 αφ=3×10⁻⁴、梯度裁剪范数 20(五档规模 1M/5M/19M/48M/317M 全部直接继承 TD-MPC2,论文只在 48M 上跑主实验)。
  • 单任务 dflex 实验(Section 4.1,5 个高维 locomotion 环境)所用世界模型具体参数规模论文正文未明确标注是哪一档(仅说明训练自 20,480 条 episode、100k 梯度步);唯一给出精确配置的是 Appendix B 双摆玩具例子——5M 档:latent 512、encoder 隐层 256、dynamics/reward 隐层各 [512,512]。
  • PWM 自己新增的”策略组件”(TD-MPC2 没有对应设计,替代其 Q-ensemble + policy-prior + MPPI 规划器):
    • Actor πθ:MLP 隐层 [400,200,100]。
    • Critic :MLP 隐层 [400,200],3 个 critic 的 ensemble(降方差)。
    • 策略训练 horizon H=16、actor batch 64;critic 训练把同一份 rollout 数据切成 4 个 mini-batch、跑 8 个梯度步(等效 critic batch = 256)。
    • 学习率 αθ=αψ=5×10⁻⁴;actor 梯度裁剪范数=1,critic=100;TD(λ) 的 λ=0.95,discount γ=0.99。
    • 动作噪声下限 0.24:FoG 下策略容易在大量梯度步后过快坍缩为确定性策略,论文没加熵正则,而是直接给动作分布标准差设一个下限来维持随机性。
  • 奖励头的差分化改造(Appendix C,是 PWM 相对 TD-MPC2 世界模型”为可微分化做的轻微修改”的核心):TD-MPC2 原版两热编码求逆需要 SymExp(x)=sign(x)(e^|x|-1),但 sign(x) 不可导;PWM 直接丢弃 SymExp 反变换,只做 softmax(x) 加权求和得到”伪奖励”(pseudo-reward)——数值上不再是精确反归一化的标量奖励,但保留了梯度、且实测足够支撑策略学习。

数据

  • MT30 / MT80 世界模型预训练数据完全复用 TD-MPC2 的离线数据集与 checkpoint,未采集新数据:“we harness the same data and world model architecture as TD-MPC2”。
    • MT30(30 个 dm_control 任务,11 种具身,动作维 m=1~6):每任务 120k 条 trajectory,由 3 个随机种子的 TD-MPC2 训练 run 生成。
    • MT80(MT30 的 30 个 dm_control 任务 + 额外 50 个 MetaWorld 操作任务,MetaWorld 侧 n=39, m=4):MetaWorld 每任务 40k 条 trajectory(同样来自 3-seed TD-MPC2 采集)。
    • 世界模型用 H=16、γ=0.99 重新预训练(而非 TD-MPC2 原始设定),论文指出这是为了给后续 FoG 提供更好梯度(呼应 3.2 节的 ESNR 分析)。
  • 单任务 dflex 数据(Hopper n=11,m=3;Ant n=37,m=8;Anymal n=49,m=12;Humanoid n=76,m=21;SNU Humanoid m=152 肌肉驱动):每任务 20,480 条 episode,由 SHAC 算法(Xu et al. 2022)采集,质量谱系覆盖从接近 0 奖励到 SHAC 收敛后的最高奖励(非专家纯净数据),世界模型在此基础上跑 100k 梯度步预训练。
  • 双摆玩具实验(Appendix B):3 个 SHAC run 共采集 24,000 条 episode(每条 240 步),采集期间最高 episode 奖励 -942.95;世界模型训练 100k batch 样本、batch size 1024,对比 H=3 与 H=16 两种训练 horizon。
  • ball-wall 玩具实验(Appendix A):对闭式解析奖励函数 f(θ) 均匀采样 1000 个点(θ∈[-π,π]),纯监督拟合,不涉及环境交互。
  • 全部数据均来自仿真器(MuJoCo/dm_control、MetaWorld、dflex),无真实机器人数据。

训练方法

Algorithm 1 的两阶段流程:

  1. 世界模型预训练(一次性):在多任务/单任务离线 buffer B 上按 TD-MPC2 式损失训练 Eφ,Fφ,Rφ——L_wm(φ) = Σ_t γ^t [‖z_{t+1} − sg(Eφ(s_{t+1},e))‖² + CE(r̂_t, r_t)](stop-gradient 的 joint-embedding 一致性项 + 奖励交叉熵)。
  2. 逐任务策略抽取(<10 分钟/任务):冻结/继续微调世界模型,针对每个任务单独训练 actor πθ + critic :
    • Actor 损失(Eq.6):在隐空间展开 H=16 步 rollout z_{h+1}=Fφ(z_h,a_h),动作采样自 a_h~πθ(·|z_h),累加折扣奖励 Rφ(z_h,a_h) 并用 critic Vψ(z_H) 做终端 bootstrap——整个表达式对世界模型做一阶/重参数化梯度反传,而非 TD-MPC2 的 DDPG-Q-梯度,也非 REINFORCE 式零阶梯度。
    • Critic 损失(Eq.7-9):同一 H 步 rollout 上用 TD(λ)(λ=0.95)学值函数,3-critic ensemble 降方差。
    • 训练时把 rollout 数据切 4 个 mini-batch、跑 8 个梯度步,得到等效 critic batch 256。
  • 世界模型在线微调:对于数据难采的高维任务(如 SNU Humanoid),在策略训练同时用默认超参 + replay buffer(容量 1024)在线微调世界模型。
  • 多任务实验中每任务训练 10k 梯度步,耗时 9.3 分钟(RTX6000)——这是论文 “<10 分钟/任务” 的具体来源。
  • 消融揭示的训练层面发现:
    • 策略 batch size(Ant 任务,32 / 512 / 2048):FoG 优化里更小的 batch(32)在单位时间内学到更好策略,与无模型 RL “batch 越大越好” 的直觉相反。
    • 世界模型正则化强度(Ant 任务,ReLU→Mish→SimNorm 依次加强正则):正则越强、世界模型损失(拟合误差)越大,但策略奖励反而越高——弱正则模型让策略前期(<1M 步)学得更快,但最终收敛到更差的次优解,直接印证论文的核心反直觉命题;该消融在 Appendix F.1 扩展到 Hopper/Anymal/Humanoid,结论一致。
    • 世界模型训练 horizon(H=3 vs H=16,同样跑 50k 步策略训练):H=16 只带来边际提升,且增益主要来自更难的 dm_control 任务而非 MetaWorld。
    • 世界模型预训练步数对策略样本效率的影响(50k/100k/250k 梯度步预训练,5 个 dm_control 任务,3 seed):固定 50k 步策略训练后,PWM 策略组件比 TD-MPC2(关闭规划)明显更样本高效,但代价是需要更充分训练的世界模型才能拿到高奖励。
    • TD(λ) actor 消融(Hopper/Ant/Anymal):把 actor 目标换成类 Dreamer 的 TD(λ),整体奖励低于 PWM 默认的 TD(N) 式目标,且多耗约 10% 计算量(在 Anymal 等高维任务上学习曲线更稳,但未被采纳为默认)。
    • 框架消融(MT30):去掉逐任务专属策略、改用「单一多任务策略」(PWM-single-policy)或「关闭在线规划的 TD-MPC2」(TD-MPC2-no-planning)都在 MT30 上明显失败,验证了”预训练一个多任务世界模型 + 逐任务抽取专家策略”这一框架设计的必要性。

Infra(训练 / 推理工程)

  • 世界模型预训练成本(48M 参数,MT30/MT80 复用 TD-MPC2 训练脚本):GitHub README 明确给出 单卡 Nvidia RTX 3090(README 原文写作 “RTX 3900”,应为笔误)训练约 2 周,并强调 horizon=16rho=0.99 是复现的关键设置。
  • 逐任务策略抽取成本(核心卖点):10k 梯度步 = 9.3 分钟 / 任务,单张 Nvidia RTX6000 GPU。
  • 单任务 dflex 实验(Section 4.1,5 个高维 locomotion 环境):除 PPO 用 1024 个并行环境外,其余方法均用 64 个并行向量化环境;每个任务/种子的完整训练(含高达 10-20M 仿真步)在单张 Nvidia RTX6000 GPU 上 ≤2 小时完成。
  • 本地复现要求 >24GB 显存的 Nvidia GPU(GitHub README 安装说明)。
  • 推理:PWM 去掉了 TD-MPC2 每决策步 6 次 MPPI 迭代的在线规划,部署时只是 actor MLP 的一次前向;论文 Figure 6 中间面板定性展示 PWM 推理时间”显著低于”TD-MPC2,但未给出具体 ms/FPS 数字
  • 并行策略、混合精度(fp16/bf16)等分布式训练细节论文未披露(全部实验为单卡)。

评测 benchmark

一手结果(均来自论文正文/附录表格,含均值±标准差,多为 50% IQM + 95% CI over 10 seeds,除非另注):

contact-rich 单任务(5 个 dflex 环境,10 seed,Table 3 为 PPO-normalized 值、Table 4 为原始 episode reward):

环境 (m=动作维)PPOSACDreamerV3TD-MPC2SHACPWM
Hopper (m=3)1.00±0.110.87±0.161.15±0.470.85±0.371.02±0.031.20±0.29
Ant (m=8)1.00±0.120.95±0.081.12±0.451.07±0.441.16±0.131.46±0.31
Anymal (m=12)1.00±0.030.98±0.061.18±0.470.98±0.481.26±0.041.16±0.24
Humanoid (m=21)1.00±0.051.04±0.041.03±0.431.05±0.461.15±0.041.19±0.025
SNU Humanoid (m=152)1.00±0.090.88±0.110.48±0.210.26±0.121.44±0.081.36±0.56

(数值为 PPO-normalized 50% IQM ± std;原始 reward 见 Table 4,如 Ant: PWM 9672±2012 vs SHAC 7662±859 vs TD-MPC2 7080±2885)。PWM 在 4/5 任务上渐进奖励超过 SHAC(直接用真实仿真器梯度的方法),且在全部 5 个任务上都超过用同一世界模型做在线 MPC 规划的 TD-MPC2;但在动作维最高的 SNU Humanoid(152 维)上不敌 SHAC——论文明确承认 “PWM does not scale well to the highest-dimensional task”。

多任务(MT30/MT80,10 seed,50% IQM + 95% CI,Figure 6/15/16 给出聚合与逐任务分数):

  • headline 数字:在 48M 参数世界模型的 80 任务(MT80)设定下,PWM 相比 TD-MPC2(同一世界模型、但依赖在线规划)取得最高 27% 的奖励提升,且无需在线规划。
  • MT30/MT80 逐任务分数(Figure 15/16)显示 PWM 的优势主要集中在更难的 dm_control 任务上,在 MetaWorld 操作任务上与 TD-MPC2 大致相当。
  • 与在线训练的单任务专家 SAC、DreamerV3 相比(MT30 子集),多任务 PWM(仅用离线数据、每任务 ≤10 分钟训练)“能够匹配”两者的表现(定性结论,图中未给出精确数字)。

消融(见”训练方法”节数字):

  • 接触刚度消融(仅 Hopper):刚度加大后 SHAC 渐进表现下降 48%,PPO/PWM 基本不受影响,PWM 渐进仍比 PPO 多 17% 奖励
  • 其余消融(batch size、WM 正则化强度、WM 训练 horizon、WM 预训练步数、TD(λ) actor、单一多任务策略框架)均为图示定性结论,论文正文未给出可摘录的精确数值(除上述已列出者)。

基线:PPO、SAC、SHAC(Xu et al. 2022,直接用 dflex 真实仿真器梯度)、DreamerV3、TD-MPC2(同一世界模型,在线 MPPI 规划)。

创新点与影响

  • 贡献 1(核心反直觉发现):世界模型的准确度与策略表现存在反相关——通过教学玩具例子和消融证明,FoG 优化真正需要的是模型的光滑性与低最优性 gap,而非拟合精度;更精确/更弱正则的模型反而给出更差的最终策略。
  • 贡献 2(FoG 优化效率):把预训练世界模型当作纯粹可微分模拟器直接反传一阶梯度,不仅比零阶梯度(PPO/在线 MPC 规划的 TD-MPC2)更高效,渐进表现甚至能超过直接拿到真实仿真器梯度的方法(SHAC)——说明”经过恰当正则的学习模型”可以是比”真实但不连续的物理”更好的优化对象。
  • 贡献 3(可扩展多任务框架):提出”先预训练一个多任务世界模型、再逐任务在分钟级抽取专家策略”的新范式,替代 TD-MPC2/GATO 式训练单一统一多任务策略的路线;论文用框架消融证明该解耦设计是必要的(单一多任务策略、或关闭规划的 TD-MPC2 均在 MT30 上失败)。
  • 作者自陈局限(Conclusion 节原文):(i) 效果高度依赖预先存在的大规模数据来训练世界模型,在新颖/低数据环境下未必可行;(ii) 尽管单任务训练很快,但每个新任务都需要重新训练,难以支持需要快速适应的场景;(iii) 目前复用的 TD-MPC2 世界模型是自回归形式,难以进一步规模化。此外实证上也观察到:PWM 在动作维最高的 SNU Humanoid(152 维)上不如 SHAC;FoG 训练偏好小 batch,与无模型 RL 的规模化直觉相反,可能限制其吞吐扩展性。

原始链接

一手源存档(sources/)