• 正文
  • 相关推荐
申请入驻 产业图谱

具身智能 | 用世界模型替代师生学习:WMP 是如何让四足机器人「看清地形」

07/09 10:52
2422
加入交流群
扫码加入
获取工程师必备礼包
参与热点资讯讨论

转载自公众号:敢敢AUTOHUB

0. 简介

笔者第一次看到 WMP 这个标题的反应是「视觉腿足又出新名词了吗」。把论文从头读一遍以后才发现,WMP 不是一种新的视觉编码器,也不是把 Dreamer 直接搬到机器人上的横向移植,而是在「师生学习」这套主流范式里抠了一个真实的痛点:师生之间的信息不对称是天然的,scandots 永远没法描述 Tilt 这种需要精确距离的薄壁通道,也没法表示 Crawl 这种悬空障碍。换句话说,特权信息本身就有上限。这个痛点不是用「再训练一个更大的 student」就能糊掉的,它要么继续凑合,要么换一条路线。

进一步看现有视觉腿足的主流做法,可以归到两类:一类是 Extreme Parkour 与 Legged Locomotion in Challenging Terrains 用的 scandots 教师,一类是 Robot Parkour Learning 把障碍几何当作 privileged 信息分地形训练。前一类的问题在于「scandots 表达力天花板低」,后一类则要为每种地形单独训一个 teacher,工程成本巨大。

WMP 想做的是绕开这道围墙:不再让 student 模仿一个不完美的 teacher,而是让 policy 直接在一个统一的世界模型上长出来。这里的关键是世界模型把高维视觉历史压成了带预测能力的隐状态,policy 拿到的就不是「scandots 的伪 ground truth」,而是「未来一段时间会发生什么」的语义级摘要。

1.1 输入输出接口

WMP 把腿足控制建模成 POMDP,机器人每个时刻能拿到的观测分两块:本体感知  与第一视角深度图 。底层状态里还有仅在仿真里可见的特权信息 ,包含 scandots、足底接触力、随机化的物理参数。policy 的输出是 12 维关节目标位置,通过 PD 控制器换算成扭矩。

这里要厘清的是:policy 实际接受的输入并不是原始深度图,而是世界模型在每隔  步采一次深度图后吐出的循环隐状态 ,所以 policy 频率仍然是 50Hz,世界模型只在 10Hz 上更新。

1.2 关键不对称:世界模型与策略的频率错位

直观理解是把世界模型当成「慢一点的眼睛+大脑」,policy 当成「快一点的肌肉」。深度图采集本身就有 40ms 左右的真机延迟,再叠加 RSSM 计算开销,如果让世界模型和 policy 同频跑,板载算力扛不住。WMP 的工程取舍是让 RSSM 每  步运行一次,即每 100ms 更新一次 ,policy 在两次更新之间复用同一个 。论文里把  从 2 扫到 30,发现仿真里  越小奖励越高,但真机由于延迟不可忽略, 是甜点。这意味着 RSSM 的离散时间步不是简单的 0.02s,而是 0.1s,这件事在阅读代码时尤其要注意,否则会把 reward sum 和 batch_length 都算错。

公式上四个 RSSM 组件可以写成下面这一组方程,对应 Recurrent / Encoder / Dynamic predictor / Decoder 四个角色,缺一不可。这套写法把世界模型拆成「记忆 + 推断 + 预测 + 重建」四步流水线,是 Dreamer 系列在 Atari 与 DMC 上反复迭代后留下来的稳定结构,WMP 的关键改造在于让循环步长跨越 5 个 policy 时间步:

其中  是 GRU 实现的循环模型, 是带 CNN 与 MLP 的多模态编码器, 既预测下一时刻潜变量分布,也负责把  解码回原观测。这里要厘清的是 prior 与 posterior 共享同一个解码 head,但训练时通过 KL 项让两者互相靠近,从而把「带观测推断」与「无观测想象」捆在同一个网络里。这种共享是 Dreamer 系列在 Atari 与 DMC 上反复验证有效的设计,WMP 把它直接搬到了腿足任务,并通过把 num_actions 扩到 5 倍来吃 100ms 时间步。

2.1 为什么不直接用 ConvNet-RNN

传统 student 用的是 ConvNet-RNN:每帧深度图过 CNN,再串联 GRU 输出 latent,再用模仿损失对齐 teacher。核心问题在于 ConvNet-RNN 只是被动总结过去的视觉信号,它没有「预测未来」的内在压力,因此当 student 走到没见过的地形时,latent 容易塌缩到「平均地形」上。RSSM 与之相反,它有显式的 prior 分布  与 posterior 分布 ,并通过 KL 散度逼迫两者靠近。这意味着即便没有当前观测,模型也能凭  一路展开未来若干步的「想象」,足够支撑机器人走过没看见的脚下区域。

2.2 RSSM 在 WMP 仓库里的真实形态

WMP 直接复用 dreamerv3-torch 的 RSSM,但把 num_actions 改成 num_actions × update_interval,让 dynamic 在 100ms 时间步内吃下 5 个动作。这件改造看似只是参数缩放,实际上把整个模型的「动作时间分辨率」从 50Hz 压到 10Hz,与 policy 的运行频率显式解耦。下面这段是 WorldModel 的训练 step,注意它并不靠 imagined rollout 做 RL,只用真实仿真数据训世界模型,这与 Dreamer 原始范式有本质差别:

# dreamer/models.py
def _train(self, data):
    data = self.preprocess(data)
    with tools.RequiresGrad(self):
        with torch.cuda.amp.autocast(self._use_amp):
            embed = self.encoder(data)
            post, prior = self.dynamics.observe(
                embed, data["action"], data["is_first"]
            )
            kl_loss, kl_value, dyn_loss, rep_loss = self.dynamics.kl_loss(
                post, prior, self._config.kl_free,
                self._config.dyn_scale, self._config.rep_scale,
            )
            preds = {}
            for name, head in self.heads.items():
                grad_head = name in self._config.grad_heads
                feat = self.dynamics.get_feat(post)
                feat = feat if grad_head else feat.detach()
                preds[name] = head(feat)
            losses = {n: -p.log_prob(data[n]) for n, p in preds.items()}
            scaled = {k: v * self._scales.get(k, 1.0) for k, v in losses.items()}
            model_loss = sum(scaled.values()) + kl_loss
        metrics = self._model_opt(torch.mean(model_loss), self.parameters())
    return {k: v.detach() for k, v in post.items()}, _, metrics

这段代码做了三件事:先用 MultiEncoder 把 prop 与 image 编码成 embed,再让 RSSM.observe 在整段轨迹上一次性算出 posterior post 与 prior prior,最后在 decoder/reward 两个 head 上算 NLL,加上 KL 项就得到模型损失。这里的关键是 grad_heads 默认只含 decoder 和 reward,意味着 cont head 在 WMP 里是禁用状态——四足任务不需要 termination 概率,因为 episode 结束由仿真器判定,无需模型预测。

难点提示(KL_free 是怎么回事):训练 RSSM 容易走两个极端,要么 prior 与 posterior 对齐过度导致 posterior 塌缩,要么放任不管导致开环预测发散。kl_free=1.0 给 KL 项设了一个「免罚区」,即每个时间步如果 KL 散度小于 1 nat 就不再施加梯度,这相当于告诉模型「分布间距足够近就够了,别再硬追完美对齐」。直觉上类似驾校教练只看你别压线,不会管你方向盘晃没晃。

3.1 为什么策略要拿  而不是

WMP 在论文里反复强调 policy 拿到的是 deterministic 部分 ,不是 stochastic 部分 。这里值得展开: 是带噪声采样的,每次走 forward 都会抖动,policy 跟它对齐就成了「对着随机数学开车」;而  是 GRU 的 deterministic state,包含了到目前为止所有压缩的视觉与本体历史,是真正稳定可用的「世界摘要」。论文里 t-SNE 把  在六种地形下投到二维平面上,会看到 Slope/Stair/Gap/Climb/Crawl/Tilt 之间有清晰边界,只有 Slope 与 Climb 微微重叠,这正好对应「Climb 可以看作 90 度 Slope」这个直观事实。

3.2 ActorCriticWMP 的真实拼装方式

下面这段是 WMP 仓库里 ActorCriticWMP 的 act 路径,可以看到 policy 的输入由三部分拼出来:本体历史 latent、3 维 command、世界模型 latent。这种「三合一」设计让策略既能利用 RSSM 的视觉语义,又能保留 RMA 风格的 explicit context,是把两条路线的优点都吃进同一个 actor 的关键写法。注意 actor 与 critic 的 wm_feature_encoder 参数是独立的,这一点对 PPO 收敛非常重要:

# rsl_rl/modules/actor_critic_wmp.py
def act(self, observations, history, wm_feature, **kwargs):
    latent_vector = self.history_encoder(history)
    command = observations[:, self.privileged_dim + 6:self.privileged_dim + 9]
    wm_latent_vector = self.wm_feature_encoder(wm_feature)
    concat_observations = torch.concat(
        (latent_vector, command, wm_latent_vector), dim=-1,
    )
    self.update_distribution(concat_observations)
    return self.distribution.sample()

def evaluate(self, critic_observations, wm_feature, **kwargs):
    wm_latent_vector = self.critic_wm_feature_encoder(wm_feature)
    concat_observations = torch.concat(
        (critic_observations, wm_latent_vector), dim=-1,
    )
    return self.critic(concat_observations)

这段代码的工程价值在于它把三个不同来源的信息分别压缩再拼接:history_encoder 是一个 MLP,吃进 5 步本体感知历史并输出 latent,对应 RMA/HIM 那条线的 explicit context;wm_feature_encoder 是另一个 MLP,把 RSSM 输出的 1536 维 deterministic state 压到 32 维供 actor 使用;critic_wm_feature_encoder 是 critic 专用的同结构 MLP,但参数独立。这三路 latent 加上 3 维 command 就是 actor 的最终输入。

3.3 不对称 Actor-Critic 的细节

注意 evaluate 里 critic 拿到的是 critic_observations,也就是 privileged_obs,里面包含 scandots 和 base linear vel,这种 asymmetric AC 方案沿用自 [Pinto et al., 2017]。这里的关键是论文强调即使有了 scandots,critic 在 Tilt/Crawl 这种 scandots 表达不完整的地形里依然要靠  才能给出准确的 value 估计,否则 baseline 会偏,PPO 的 advantage 也会跟着偏。换句话说, 不只是 actor 的眼睛,也是 critic 的「补盲器」。

4. 训练时的监督模块:仿真里同时学世界模型与策略

4.1 单阶段训练的工程设计

WMP 与 student-teacher 最大的差别就是只有一个训练阶段。仓库里的 WMPRunner.learn 在每个迭代里同时做三件事:用 PPO+AMP 训 actor-critic、用监督回归训 depth predictor、用 ELBO 训世界模型。这种「同 episode 内三模型并训」的写法非常 ralph,工程上反而稳,因为 RSSM 不参与 RL 反传(policy 的梯度通过 stop-gradient 截断),所以三个优化器之间不会互相破坏。

工程价值(为什么单阶段比两阶段强):student-teacher 两阶段训练的根本痛点是「上限被 teacher 卡死」,因为 student 模仿误差会累积。WMP 把世界模型嵌入到 PPO 的感知通路里,让世界模型的训练目标(重建过去观测、预测未来)与 policy 的训练目标(最大化 return)正交,但又共享同一份仿真数据。这意味着两边互相不抢数据,还能并行收敛。换句话说,把感知与决策放进同一个训练循环,比串行的两阶段流水线更接近动物的学习方式。

# rsl_rl/runners/wmp_runner.py
for it in range(self.current_learning_iteration, tot_iter):
    with torch.inference_mode():
        for i in range(self.num_steps_per_env):
            if (self.env.global_counter % self.wm_update_interval == 0):
                wm_embed = self._world_model.encoder(wm_obs)
                wm_latent, _ = self._world_model.dynamics.obs_step(
                    wm_latent, wm_action, wm_embed, wm_obs["is_first"],
                )
                wm_feature = self._world_model.dynamics.get_deter_feat(wm_latent)
                wm_is_first[:] = 0
            history = self.trajectory_history.flatten(1).to(self.device)
            actions = self.alg.act(obs, critic_obs, amp_obs, history,
                                   wm_feature.to(self.env.device))
            obs, privileged_obs, rewards, dones, infos, *_ = self.env.step(actions)
        self.alg.compute_returns(critic_obs, wm_feature.to(self.env.device))
    mean_value_loss, mean_surrogate_loss, *_ = self.alg.update()

    if (sum_wm_dataset_size > self.wm_config.train_start_steps):
        if (it % self.depth_predictor_cfg["training_interval"] == 0):
            depth_mse_loss = self.train_depth_predictor()
        wm_metrics = self.train_world_model()

4.2 Depth Predictor 的角色

DepthPredictor 是 WMP 工程链路里非常关键但容易被忽略的一环:仿真里 4096 个并行 A1 实例如果每个都跑深度相机,IsaacGym 的 GPU 会炸。核心问题在于,WMP 给一个折中方案——只让 1024 个 env 真的跑深度相机(camera_num_envs=1024),剩下的 env 用 forward_height_map 加一个 ConvTranspose 解码器近似生成深度图,并通过监督学习让两路对齐。下面这段是 DepthPredictor.forward

# rsl_rl/modules/depth_predictor.py
def forward(self, forward_heightmap, prop):
    x = torch.concat((forward_heightmap, prop), dim=-1)
    x = self.encoder(x)
    x = x.reshape(
        [-1, self.h_list[0], self.w_list[0],
         self._embed_size // (self.h_list[0] * self.w_list[0])]
    )
    x = x.permute(0, 3, 1, 2)
    x = self.layers(x)
    mean = x.permute(0, 2, 3, 1)
    if self._cnn_sigmoid:
        mean = F.sigmoid(mean)
    return mean

这段做的是「把 525 维 forward heightmap + 33 维 prop 解码为 64×64×1 深度图」,监督信号就是真相机渲染的深度图与该预测结果的 MSE。进一步看,这是一个聪明的工程取舍:仿真里 forward heightmap 是几乎免费的(terrain mesh 直接采样),相机渲染是昂贵的,所以让大部分 env 用便宜信号训 RSSM,关键 env 用真深度图监督 predictor,最后所有 env 在 RSSM 阶段拿到的都是「视觉化」的输入。

直觉理解:可以把 depth predictor 想象成一个「假摄像头」,它本身没有 RGB-D 的物理传感能力,但因为知道地形高度场和机器人姿态,它可以伪造出几何正确的深度图。RSSM 不在乎深度图是真是假,只要分布一致就行。

5. 推理时的执行模块:板载推理的延迟约束

5.1 板载流水线

真机上的 WMP 部署在 Unitree A1 板载的 Jetson NX 上,没有外接计算盒。这意味着整个 RSSM + actor 都要在 NX 的 GPU 上跑,policy 输入数据流大致是:Intel D435i 60Hz 输出 424×240 深度图 → spatial+temporal 滤波 → 中心裁剪+下采样到 64×64 → 送入 RSSM 编码器,整段链路含 100ms 延迟。这个延迟在仿真训练时是显式建模的:env 每隔 5 个 timestep 才把深度图灌给 RSSM,并刻意附加 100ms 滞后,对应论文 Section IV-C「Depth images are computed every k timesteps and sent to the policy with 100ms latency」。

5.2 Sim-to-real 的关键超参

仓库 a1_amp_config.py 里 class depth 的几行配置直接写进了真机部署的契约,每一行都对应一个不可或缺的 sim-to-real 假设。如果跨平台部署时不重扫这几个数字,整个 WMP 的延迟与视场假设都会失效,这也是开源用户最容易踩的坑之一。这里要厘清的是配置里的每一项都是仿真 / 真机协商出来的结果,不是随便填的默认值:

# legged_gym/envs/a1/a1_amp_config.py
class depth:
    use_camera = True
    camera_num_envs = 1024
    update_interval = 5  # 5 works without retraining, 8 worse
    original = (64, 64)
    resized = (64, 64)
    horizontal_fov = 58
    near_clip = 0
    far_clip = 2
    dis_noise = 0.0
    scale = 1
    invert = True

这里值得注意几个数字:horizontal_fov=58 与 D435i 的 default 视场吻合;far_clip=2 表示超过 2m 的深度被截断,因为对 A1 这种 0.25m 高的小型四足来说,2m 之外的远景对 1 秒内的步态决策没有用;update_interval=5 直接对应 RSSM 频率,注释里特地标了「5 不需要重训,8 就明显变差」,是论文 ablation 的工程结论。

6. 训练目标:三套损失并存

6.1 世界模型与 PPO 的损失装配

WMP 的总损失可以写成三块叠加:世界模型 ELBO、策略 PPO+AMP、深度预测 MSE。三块损失分别由三个独立的优化器更新各自参数,互不交叉。这里的关键是 RSSM 的梯度只在 ELBO 内部流动,不会被 PPO 的 surrogate 推着乱跑,反过来 PPO 对 wm_feature 也是 stop-gradient,避免 RL 信号污染世界模型。三组公式如下:

第一块是 RSSM 的 ELBO,第二块是策略侧的 PPO + AMP 风格奖励(外加 vel-predict 辅助损失),第三块是 depth predictor 的 MSE。三者通过共享数据缓冲区交换信息,但梯度互不相通。这里的关键是 PPO 的 advantage 通过 compute_returns(critic_obs, wm_feature) 拿到 critic 的 value,而 critic 同样吃 ,于是世界模型间接影响了所有三个损失,但梯度只在 ELBO 里反传。这种「数据共享、梯度隔离」的写法,是 WMP 单阶段训练能稳定收敛的根本工程保证。

6.2 关键超参表

# dreamer/configs.yaml + a1_amp_config.py 节选
dyn_deter:        512
dyn_stoch:        32
dyn_discrete:     32
units:            512
kl_free:          1.0
dyn_scale:        0.5
rep_scale:        0.1
batch_size:       16
batch_length:     64
train_steps_per_iter: 10
train_start_steps:    10000
model_lr:         1e-4
update_interval:  5     # WMP / 100ms
camera_num_envs:  1024  # depth-truth envs
amp_reward_coef:  0.01  # 0.5 * 0.02
entropy_coef:     0.01
vel_predict_coef: 1.0

这里要厘清的是几个有反差的取舍:batch_length=64 对应 6.4 秒训练片段,论文 Figure 4 显示 6.4s 是性能拐点,再短记忆不够、再长 RSSM 反传不稳;amp_reward_coef=0.5×0.02=0.01 是 AMP 风格奖励的实际权重,足够小到不会 overshadow tracking reward 但又足够大让步态自然;train_start_steps=10000 让 RSSM 在策略开始训练前先离线吃一波数据,避免冷启动期 policy 拿到无意义 latent。

7. 可选模式:t-SNE 可视化与开环预测

7.1 循环状态的可视化

论文 Section V-B 的 t-SNE 实验把六种地形下的  投到二维平面,可以看到不同地形的 cluster 几乎完美分开。这意味着  不仅压缩了视觉历史,还隐式编码了「我现在是在爬楼还是在过缝」这种语义信息。核心问题在于,policy 不需要任何额外的地形分类头,只要拿到  就能在 actor MLP 里学会「不同地形不同步态」的策略分支,这一点跟 [Robot Parkour Learning] 必须为每种地形单独训 teacher 形成鲜明对比。

解读(t-SNE cluster 是怎么自然分开的):RSSM 训练时唯一的目标是「重建观测 + 预测未来」,并没有任何显式的地形分类损失。但因为不同地形的视觉模式与本体响应序列差异巨大(爬楼时的关节角度模式与穿缝时完全不同),世界模型必须把它们编码到 latent 空间的不同区域才能有效压缩信息。换句话说,地形分离是 ELBO 的副产物,不是设计目标。这种「无监督学到结构」的现象,正是世界模型路线最优雅的地方。

WMP 论文的另一个亮眼实验是把仿真训练的 RSSM 拿到真机上做开环 rollout:给定真实第一帧观测和真实动作序列,让 RSSM 想象后续若干秒的深度图。结果显示,在 Crawl 任务里,模型预测的悬空横梁形状与真实形状不完全一致(仿真里没见过这种形状),但「机器人能穿过的缝隙」位置和角度高度吻合。换句话说,RSSM 没有学到「物体识别」,但学到了「可通行性」这一关键属性,这才是 sim-to-real 平滑迁移的根本原因。

公式上即便没有 reward head 主导训练(reward.loss_scale=0),decoder 重建出的深度图本身已是足够丰富的监督信号。这种「不靠 reward 也能学」的特性是 Dreamer 系列的标志,WMP 把它发挥得更彻底——既然 RSSM 不参与策略梯度反传,那么 reward head 在 WMP 里实际上只是一个观测变量,留着只是为了和原始 dreamerv3-torch 保持兼容。这意味着如果有人想把 WMP 移植到没有 reward 信号的纯模仿场景里,直接把 reward head 拆掉对训练效果几乎没有影响。

8. 总结

WMP 是 Dreamer 路线在视觉腿足里的落点,不是 student-teacher 路线的横向变体。它的真正贡献不是「用世界模型替代 student」,而是「让世界模型同时承担视觉感知与未来预测两职」,从而绕过 scandots 表达力天花板,又避开 ConvNet-RNN 没有预测压力的缺陷。

相关推荐