转载自公众号:敢敢AUTOHUB
0. 简介
上一篇文章从论文视角拆解了 LaWAM 的动机与核心机制。这篇文章换一个维度:从代码仓库出发,逐层剖析这套系统的工程骨架是怎么搭建的、数据在各模块之间怎么流转的、以及 Teacher-Student 蒸馏范式在具身智能场景下究竟是怎么落地的。写这篇文章的目标不是复述论文摘要,而是让读者读完之后能回答三个问题:第一,如果我想在自己的项目里实现类似的"世界模型增强策略",我需要哪些组件、它们之间的接口长什么样;第二,Teacher-Student 蒸馏在这里为什么不是简单的"大模型教小模型",而是"后验教先验";第三,从训练到推理的切换点在哪里,推理时哪些模块被移除了、哪些被保留了。
论文地址:https://arxiv.org/abs/2606.15768
GitHub项目地址:https://github.com/RLinf/LaWAM
项目地址:https://rlinf.github.io/LaWAM/
1. 两阶段训练架构:先练基础件,再组装系统
1.1 为什么不是端到端一步训完
LaWAM 的训练并非一个脚本从头跑到底的单一流程,而是拆成了两个物理上独立的阶段——各自有不同的训练框架、不同的配置系统、不同的产出物。进一步看这种两阶段的设计思路,它在大模型时代并不罕见:GPT 系列先做预训练再做对齐,扩散模型先练 VAE 再练去噪网络,本质上都是"先把基础件练好,再在基础件上组装端到端系统"。LaWAM 的特殊之处在于,它的基础件不仅仅是一个特征提取器或编码器——阶段一训练出的 LAM(Latent Action Model)同时扮演了两个后续角色:encoder 路径在阶段二充当蒸馏的 Teacher,decoder 路径在阶段二充当世界模型。这种"一次训练,双重复用"的设计,使得阶段一的质量直接决定了阶段二整个系统的上限。
1.2 阶段一:LAM 预训练——在视频中学会"动作意图"
阶段一的入口文件是 latent_action_model/core/lam_lightinng.py,基于 PyTorch Lightning 框架。它的目标是训练一个能从连续视频帧中无监督地抽取"动作意图"的模型。这里的"无监督"指的是不需要机器人动作标签——模型只看视频帧的前后变化,就要学会"这两帧之间发生了什么"。具体而言,LAM 包含三个核心组件:一个冻结的视觉编码器(基于 VJEPA2 / DINOv2 的 ViT-B/16)把像素帧映射到 patch token 特征空间;一个 Q-Former 风格的编码器把"当前帧特征 + 未来帧特征"压缩成极低维的潜动作向量(通常 num_queries=1, code_dim=32);一个 VQ 量化瓶颈把连续向量离散化为码本中的码字索引。> 直觉理解:可以把 LAM 的训练类比成一个"看图说话"测试——给你两张前后照片,你要说出"中间发生了什么"(encoder 的工作),然后另一个人只听你的描述、只看第一张照片,就要画出第二张(decoder 的工作)。如果第二张画得像,说明你的描述确实抓住了关键变化。
训练损失由三项构成,各自承担不同的监督职责:解码器重建损失检验 decoder 能否从"当前帧 + 潜动作"恢复出未来帧的视觉特征——这是核心自监督信号;VQ 量化损失约束码本的利用率与多样性,防止码本坍缩到少数几个码字上;状态辅助损失则是一个轻量的正则项,逼迫潜动作偏向编码具身运动信息(如末端位移方向)而非纯视觉外观变化(如光照变化、背景晃动)。
这个设计的灵感源流可以追溯到 2024 年 LAPA(Latent Action Pretraining from Videos, ICLR 2025)和 DULA(Distilling Universal Latent Actions)等工作。它们都证明了一件事:用 VQ-VAE 目标在大量视频上训练出的离散潜动作码本,能够跨本体、跨场景地编码"帧间发生了什么",而不依赖于具体的机械臂关节定义。LaWAM 的阶段一继承了这个思路,但多走了一步——它不仅保留了编码器,还保留了解码器。
1.3 阶段二:LaWAM 策略微调——把 LAM 组装进 VLA 系统
回到阶段二。阶段二的入口文件是 starVLA/training/train_starvla.py,基于原生 PyTorch + HuggingFace Accelerate,作者在代码注释中明确表示选择手写训练循环而非使用高层 Trainer 是为了"保持显式、易于 hack"——这在需要频繁实验不同训练策略(如 Scheduled Sampling、repeated diffusion steps 等非标准技巧)的研究代码中是一个合理的选择。阶段二把四个大模块组装在一起:Qwen-VL 多模态主干(处理图像 + 语言指令)、阶段一预训练好的 LAM(提供 Teacher 信号和世界模型能力)、一个轻量的 VLMToLAMQFormer(充当 Student)、以及一个基于 Flow Matching 的动作生成头(输出最终动作序列)。这里的关键是两个阶段的衔接。衔接点是一行配置:
# starVLA/model/framework/vlas/lawam.py
lam_ckpt_path: str = str(DEFAULT_LAM_ROOT / "logs/dino_base_ae_bridge/version_0/checkpoints/epoch=39.ckpt")
lam_yaml_path: str = str(DEFAULT_LAM_ROOT / "logs/dino_base_ae_bridge/version_0/dino_base_ae.yaml")
阶段二通过这两个路径加载阶段一的权重和结构定义,把 LAM 的 encoder 路径冻结、用 torch.no_grad() 包裹(充当产出蒸馏标签的 Teacher),把 decoder 路径保留并允许梯度流过(充当世界模型 LaWM)。从这一刻起,阶段一的训练产物分裂成了两种截然不同的身份——同一个模型的两半,在阶段二的训练和推理中各自扮演不同角色,彼此之间只通过一个 32 维的潜动作向量发生耦合。
2. 系统拓扑:四大子模块的协作关系
2.1 模块分工与数据接口
整个 LaWAM 策略网络的核心类是 LatentWorldPolicyBackend,定义在 starVLA/model/framework/vlas/lawam.py,总计约 920 行代码。它把四个子模块组装成一条完整的"观测→意图→未来→动作"链路,每个子模块各司其职,通过明确的张量接口互相传递信息——下面这张表列出了它们各自的输入输出契约和训练时行为:
| 模块 | 代码字段 | 输入 | 输出 | 训练时行为 |
|---|---|---|---|---|
| Qwen-VL | self.vlm |
图像 + 语言指令 + 占位符 token | 多模态上下文 hidden states | 参数可微调 |
| VLMToLAMQFormer | self.vlm_to_lam |
VLM 中 <ACT_PH> 位置的 hidden states |
潜动作预测 pred_action_emb [B,1,32] |
Student,全梯度 |
| LAM | self.lam |
视频帧(encoder)/ 当前帧+潜动作(decoder) | Teacher 信号 / 未来视觉子目标 | encoder 冻结;decoder 视配置 |
| Flow Head | self.flow |
当前视觉 + 未来子目标 + VLM 上下文 + 状态 | 动作序列 [B, H, Da] | 全梯度 |
这四个模块的关系不是简单的串联。训练时存在两条并行的数据通路:一条是 Student 路径(VLM → QFormer → decoder → 子目标),另一条是 Teacher 路径(视频帧 → LAM encoder → 量化潜动作)。两条路径在蒸馏损失处交汇——Student 的输出要逼近 Teacher 的输出。推理时,Teacher 路径被完全移除,只剩 Student 路径驱动整个系统。
2.2 占位符注入机制:VLM 如何与动作空间对接
一个值得注意的工程细节是 LaWAM 如何让 Qwen-VL(一个本来只做语言和视觉理解的模型)产生动作相关的表示。答案是占位符注入:在 tokenize 阶段,模型在输入序列中插入若干 <ACT_PH>(Action Placeholder)特殊 token。这些 token 在 VLM 前向传播时会像普通 token 一样参与自注意力计算,但在 VLM 输出侧,它们所在位置的 hidden states 会被单独提取出来,送入 VLMToLAMQFormer 做进一步翻译。这种设计避免了修改 VLM 的内部结构,只通过"在输入侧插 token、在输出侧抽特征"就完成了语言模型与动作空间的对接。代码中的关键步骤如下:
# starVLA/model/framework/vlas/lawam.py
def _run_vlm_stage(self, input_ids, attention_mask, pixel_values, image_grid_thw,
act_placeholder_mask, flow_placeholder_mask, act_query, flow_query):
# 1. 把 VLM embedding 层中 <ACT_PH> 位置的向量替换为可训练的 query
hidden = self.vlm.model.embed_tokens(input_ids)
hidden[act_placeholder_mask] = act_query.expand_as(hidden[act_placeholder_mask])
hidden[flow_placeholder_mask] = flow_query.expand_as(hidden[flow_placeholder_mask])
# 2. 正常跑 VLM 前向
vlm_out = self.vlm(inputs_embeds=hidden, attention_mask=attention_mask,
pixel_values=pixel_values, image_grid_thw=image_grid_thw)
h_vlm = vlm_out.last_hidden_state
# 3. 提取 <ACT_PH> 位置的 hidden state,送入 QFormer
h_act = h_vlm[act_placeholder_mask] # [B * num_queries, vlm_hidden_dim]
pred_latent = self.vlm_to_lam(h_act.view(bsz, q, -1)) # → [B, 1, code_dim]
return {"h_vlm": h_vlm, "pred_latent": pred_latent}
这段代码揭示了一个设计选择:act_query 和 flow_query 是可训练的参数,而非固定的 embedding。这意味着训练过程中,模型会自动学会"在这些占位位置应该关注什么样的上下文信息",从而让 VLM 在这些位置产生最有利于动作预测的表示。这种做法比在 VLM 输出之后加一个固定的投影头更灵活——可训练 query 实际上允许 VLM 的自注意力层"看到"动作任务的需求,从而在内部就为动作预测做好准备,而不是事后再从通用表示中费力提取动作相关信号。
3. Teacher-Student 蒸馏:后验教先验的工程实现
3.1 蒸馏的本质:弥合训练与推理之间的信息鸿沟
在传统的知识蒸馏中(如 Hinton et al., 2015),Teacher 是一个参数更多、能力更强的大模型,Student 是一个参数更少但推理更快的小模型,蒸馏的目标是让小模型模仿大模型的输出分布。LaWAM 中的蒸馏与此截然不同——Teacher 和 Student 的参数量差异并不大,真正的差异在于它们能看到的信息量。
Teacher(LAM encoder)在训练时能同时看到当前帧和未来帧,它拥有后验信息:知道"未来实际发生了什么",因此能精确地推断出"这段时间内的动作意图是什么"。Student(VLMToLAMQFormer)只能看到当前观测和语言指令,它必须在没有未来信息的条件下猜出动作意图——这是一个先验估计。蒸馏的目标,是让先验尽可能逼近后验。
一句话理解:Teacher 是那个已经看完电影结局的人,Student 是那个只看了开头就要猜结局的人。蒸馏就是让"只看开头的人"学会像"看完全片的人"一样准确地描述"接下来会发生什么"。
这种"后验→先验"的蒸馏范式在机器学习中有更广泛的根基。VAE 中的 encoder 和 decoder 就存在类似的不对称:encoder 看到完整数据后给出后验分布 ,而 decoder 在生成时只能从先验 采样。DAgger(Ross et al., 2011)中的专家策略能看到全局状态,而学生策略只能看到局部观测。LaWAM 把这种范式搬到了具身智能的语境下:Teacher 看到了"任务实际如何完成"(后验),Student 必须在"任务还没执行"时就预测出意图(先验)。
3.2 Teacher 的实现:冻结的 LAM Encoder
先看 Teacher 这一侧。Teacher 的具体实现在 _run_lam_teacher 方法中,它的工作流程直截了当:接收训练视频帧(包含当前帧和未来帧),经过冻结的 LAM encoder 抽取潜动作,再经过 VQ 量化得到离散码字,整个过程在 torch.no_grad() 保护下执行,不接受任何梯度回传——Teacher 的职责只是产出一个稳定的蒸馏标签,它自身不需要在阶段二中继续学习:
# starVLA/model/framework/vlas/lawam.py
def _run_lam_teacher(self, *, primary_video: torch.Tensor, embodiment_id: torch.Tensor) -> torch.Tensor:
# 对齐时间轴:训练视频帧数可能多于 LAM 预训练时的设定,均匀采样到 Teacher 期望长度
primary_video_t = self._build_lam_teacher_inputs_for_distill(primary_video)
with torch.no_grad():
# Teacher 路径完全冻结——不接受梯度,保证蒸馏标签的稳定性
lam_out = self.lam.get_latent_action(
videos=primary_video_t,
states=None,
dec_videos=primary_video_t,
predict_future_frame=False,
embodiment_ids=embodiment_id,
)
# 取量化后的潜动作(码本空间中的离散点),detach+clone 双重保险
return lam_out["quantized"].detach().clone()
这里有三个值得注意的工程决策。第一,时间轴对齐:_build_lam_teacher_inputs_for_distill 使用 torch.linspace 均匀采样帧索引,确保无论训练视频有多少帧,Teacher 始终在它预训练时"舒适"的时间分辨率上工作。第二,取 quantized 而非连续向量:量化后的码字位于有限码本空间中,相当于给蒸馏目标加了一层离散正则化——换句话说,Student 不需要精确匹配一个可能在连续空间中漂移的向量,只需要对齐到码本中最近的锚点。第三,detach().clone() 的双重保险:detach 切断计算图防止意外梯度回传,clone 避免 tensor 内存共享导致的潜在 bug(某些 LAM 实现在推理模式下可能返回 view 而非独立 tensor)。
3.3 Student 的实现:VLMToLAMQFormer
再看 Student 这一侧。Student 的实现是一个极其轻量的 Cross-Attention 模块——只有一层 cross-attention 加一个 FFN,参数量不到 LAM encoder 的百分之一。但它承担的任务却很重:把 VLM 产生的高维上下文表示(3584 维的 Qwen-VL hidden states,包含了图像理解、语言指令解析、多模态对齐的全部信息)压缩到与 LAM 潜动作完全相同的低维空间(32 维的 code_dim),且压缩后的向量要能被 LAM decoder 正确消费、展开为有意义的未来子目标:
# starVLA/model/framework/vlas/lawam.py
class VLMToLAMQFormer(nn.Module):
"""把 VLM 的 placeholder hidden states 翻译成 LAM 能消费的 latent action。"""
def __init__(self, *, vlm_hidden_dim: int, lam_code_dim: int,
num_layers: int = 1, num_heads: int = 8):
super().__init__()
# 只有 1 个可学习 query:极端瓶颈,强制只保留"动作意图"这一维信息
self.query = nn.Parameter(torch.randn(1, 1, lam_code_dim) * 0.02)
# Cross-attention:query 在 code_dim 空间,KV 在 vlm_hidden_dim 空间
self.cross_attns = nn.ModuleList([
nn.MultiheadAttention(embed_dim=lam_code_dim, kdim=vlm_hidden_dim,
vdim=vlm_hidden_dim, num_heads=num_heads, batch_first=True)
for _ in range(num_layers)
])
# 每层 cross-attention 后接 FFN
self.ffns = nn.ModuleList([...])
self.final_norm = nn.LayerNorm(lam_code_dim)
def forward(self, context: torch.Tensor) -> torch.Tensor:
# context: [B, num_act_queries, vlm_hidden_dim]
queries = self.query.expand(context.shape[0], -1, -1)
for norm_q, norm_kv, xattn, ffn in zip(...):
q = norm_q(queries)
kv = norm_kv(context)
attn_out, _ = xattn(q, kv, kv)
queries = queries + attn_out
queries = queries + ffn(queries)
return self.final_norm(queries) # → [B, 1, code_dim]
这里的设计哲学是信息瓶颈最大化。VLM 输出的 hidden states 可能有数千维、数十个 token 位置,但 QFormer 只用 1 个 32 维的可学习 query 去"问"所有这些上下文。换句话说,这个极端的压缩比(从 num_queries × vlm_hidden_dim 到 1 × 32)强制网络只保留对动作预测最关键的那一丝信息——"接下来应该执行什么样的动作意图"。任何与动作无关的视觉细节、语言修饰、背景信息,都会在这个瓶颈处被丢弃。
还有一个容易被忽略的实现细节是 nn.MultiheadAttention 的 kdim/vdim 参数。普通的自注意力要求 Q、K、V 同维度,但这里 Q 在 32 维空间,K 和 V 在 3584 维空间。PyTorch 的 MultiheadAttention 内部会自动为 K 和 V 分配从 kdim/vdim 到 embed_dim 的投影矩阵,完成跨维度的注意力计算。这避免了手写线性投影层的冗余代码。
3.4 蒸馏损失的度量选择
接着看连接 Teacher 和 Student 的桥梁——蒸馏损失本身是怎么算的。蒸馏损失的计算逻辑实现在 _compute_latent_loss 方法中,代码只有短短几行,但背后的度量选择反映了作者对潜动作空间几何结构的深入理解。这里提供了 MSE 距离和 cosine 距离两种可切换的度量方式,默认配置选用的是前者,原因与 VQ 量化后的码字空间性质有关:
def _compute_latent_loss(self, *, pred_latent, teacher_latent, latent_loss_type):
if pred_latent.shape != teacher_latent.shape:
raise ValueError(f"latent shape mismatch: pred={tuple(pred_latent.shape)}, "
f"teacher={tuple(teacher_latent.shape)}")
if latent_loss_type == "mse":
return F.mse_loss(pred_latent, teacher_latent)
# cosine 距离:只关注方向对齐,忽略模长差异
return 1 - F.cosine_similarity(pred_latent, teacher_latent, dim=-1).mean()
默认配置使用 MSE。选择 MSE 而非 cosine 的原因是:Teacher 输出的是经过 VQ 量化后的码字向量,这些向量在码本空间中的绝对位置是有意义的(每个码字对应一种特定的动作模式)。MSE 同时约束方向和模长,确保 Student 不仅预测出正确的"动作类别",还预测出正确的"动作强度"。cosine 距离是备选方案——当码本向量的模长信息不重要、只需方向对齐时更合适,但在 LaWAM 的默认设定下未被启用。
4. 世界模型的角色复用:decoder 从辅助件到核心组件
4.1 传统做法的浪费与 LaWAM 的反转
在 LAPA、VQ-BeT(Behavior Generation with Latent Actions)等前序工作中,LAM 训练完成后,decoder 通常被视为一个训练时的辅助工具——它的唯一作用是验证 encoder 是否学好了(进一步看它的逻辑——如果 decoder 能从潜动作重建未来帧,说明潜动作确实编码了足够的变化信息),训练完就被丢弃。LaWAM 的关键洞察是:这个 decoder 已经具备了"给定当前视觉状态和一个动作意图,预测未来视觉状态"的能力——这恰恰就是一个世界模型的定义。把它保留下来直接当世界模型用,比从头训练一个新的世界模型既省参数又省数据。
4.2 decoder 在阶段二的使用方式
顺着 Teacher-Student 蒸馏的线索往下走。在阶段二中,decoder(论文称为 LaWM,Latent World Model)接收两个输入:当前观测的视觉 token h_t,以及 Student 预测的潜动作 pred_action_emb。它的输出是一组预测的未来视觉 token h_t1_pred,即"如果执行了这个动作意图,未来的视觉状态在特征空间里会长什么样":
# starVLA/model/framework/vlas/lawam.py
def _decode_future_tokens_strict_single_query(self, *, h_t, pred_action_emb, source):
"""用 LAM decoder(世界模型)把潜动作展开为未来视觉子目标"""
if pred_action_emb.shape[1] != 1:
raise ValueError(f"[{source}] future_prediction requires single-query latent action, "
f"got query_dim={pred_action_emb.shape[1]}.")
# h_t: [B, N_vis, D_vis] — 当前观测的 256 个视觉 patch token
# pred_action_emb: [B, 1, code_dim] — Student 预测的 32 维潜动作
decoded = self.lam.decoder(h_t, pred_action_emb)
# decoded: [B, N_vis, D_vis] — 预测的未来视觉 token(隐空间视觉子目标)
if decoded.dim() == 4:
decoded = decoded[:, -1, :, :] # 多步时取最后一帧
return decoded
这段代码的输出 decoded 就是论文所说的"隐空间视觉子目标"(latent visual subgoal)。它不是一张图片,不是一段视频,而是一组 256 个 patch token 构成的特征序列。直观理解是,这些 token 描述的是"未来场景中哪些区域会发生变化、交互点在哪里、物体会移动到什么位置"——所有这些信息都以与动作决策最相关的形式被编码,省略了纹理、光照、背景这些对控制无用的细节。
这正是 LaWAM 相比像素空间世界模型(如 LingBot-VA 等基于视频生成的方案)的核心优势所在。像素 WAM 需要逐像素迭代去噪来生成一段"看起来合理"的未来视频,然后再让策略从视频中提取有用信息;LaWAM 直接在特征空间里做一次前向传播就得到了"对动作有用的未来状态表示",中间跳过了所有与控制无关的像素重建开销。
工程价值:这个设计选择的经济账很清楚——像素 WAM 花 4 秒生成一段"给人看的未来视频",然后策略还得从视频里提取有用信息;LaWAM 花 30 毫秒直接产出"给策略看的未来特征",跳过了中间所有冗余环节。省掉的不只是计算量,还有整个视频生成模型的参数预算(5B vs 230M)。根据论文实验数据,LingBot-VA 每次策略推理需要 4482 毫秒,而 LaWAM 只需 187 毫秒——24 倍的差距主要就来自于这个"在哪个空间做未来预测"的设计选择。
4.3 decoder 内部结构:AdaLN-DiT 条件注入
这里值得单独拆开看的是 LAM decoder 的内部结构值得单独说一下。它是一个 24 层的 Transformer,采用 AdaLN(Adaptive Layer Normalization)机制把潜动作注入到每一层的计算中。这种设计来自 DiT(Scalable Diffusion Models with Transformers)的实践:与其把条件信息拼接到输入序列(会增加序列长度和计算量),不如让条件信息通过 LayerNorm 的 scale 和 shift 参数来调制每一层的特征。潜动作先经过一个 MLP 映射到与 hidden dim 匹配的调制向量 ,然后在每层 Transformer 的 LayerNorm 中做:
其中 是潜动作向量, 是当前帧的视觉 patch token 序列。这种 AdaLN 注入方式的核心优势在于:潜动作的影响不是集中在网络的某一层,而是均匀地渗透到了 decoder 全部 24 层 Transformer 的每一次归一化操作中,使得 decoder 能在从浅层到深层的不同抽象粒度上,根据不同的动作意图调整对未来视觉状态的预测——浅层可能调整局部纹理的变化方向,深层可能调整全局物体位置的偏移趋势。
5. 三条损失的协同:谁主导、谁辅助、谁塑形
5.1 总损失的装配公式
阶段二的总损失函数把三个互相配合的组件按各自的权重装配在一起,形成一个联合优化目标。三个组件各司其职、缺一不可——动作流匹配损失直接监督最终输出的动作轨迹质量,蒸馏损失负责把 Student 的潜动作预测对齐到 Teacher 的输出空间上,感知损失则约束世界模型解码出的子目标不偏离真实未来太远。它们之间的权重配比体现了"主力 + 辅助"的层次关系:
其中感知损失的权重 把子目标监督压到辅助地位——它的作用是"确保世界模型不跑偏"而非主导训练方向。蒸馏权重 看起来数值上与流匹配主损失持平,但实际梯度贡献要小得多,原因是蒸馏回归的对象是一个仅有 32 维的潜动作向量,而流匹配回归的是一个形状为 [B, 8, Da] 的完整动作速度场张量。对应到代码中的实现如下:
# starVLA/model/framework/vlas/lawam.py — forward 方法
loss_total = (
loss_flow # 动作流匹配主损失
+ self.model_cfg.perceptual_weight * loss_perceptual # 子目标监督,权重 0.1
+ self.model_cfg.lam_encoder_distill_weight * loss_distill # 蒸馏,权重 1.0
)
5.2 各损失的角色与量级分析
换一个角度来拆解这三条损失。三条损失在训练中扮演不同角色。loss_flow 是 Flow Matching 的速度场回归损失,它直接监督动作生成的质量——网络预测的速度场 要匹配真实的"噪声到数据"方向 。这是主力损失,决定了最终输出动作的精度。loss_distill 让 Student 的潜动作预测对齐 Teacher 的输出,它的作用是"塑形"——把 Student 拉到正确的动作意图空间上,但不直接约束最终动作的质量。loss_perceptual 约束世界模型预测的子目标要贴近真实的未来帧特征,它的作用是"校准"——确保 Student 驱动 decoder 展开后的结果不会偏离真实未来太远。
蒸馏权重给了 1.0 看起来很大,但实际梯度贡献是温和的。原因是 loss_distill 回归的是一个 32 维向量的 MSE,其数值量级天然比 loss_flow(回归一个 [B, 8, Da] 动作张量的速度场 MSE)小得多。loss_perceptual 给 0.1 的原因更直接:它回归的是 256 个 patch token 每个维度的 MSE,数值量级比 flow loss 大,需要用小权重压下来,让它只起到"别跑太偏"的约束作用而非主导训练方向。
5.3 perceptual loss 的计算
loss_perceptual 的计算非常简洁,它直接在冻结视觉编码器的特征空间中对比两个张量——Student 驱动 decoder 产生的未来子目标 h_t1_pred(模型认为未来会是什么样),与真实未来帧经过同一个冻结编码器后的特征 h_t1_gt(未来实际是什么样),不回归任何像素级别的内容:
if self.model_cfg.future_prediction:
# 在 LAM 视觉 token 空间中直接对比,不回归 RGB 像素
loss_perceptual = F.mse_loss(shared.h_t1_pred, shared.h_t1_gt)
else:
loss_perceptual = torch.tensor(0.0, device=device, dtype=lam_stage_dtype)
这里 h_t1_gt 来自训练数据中真实的未来帧经过冻结视觉编码器的编码结果(代码中 features[:, -1, :, :])。这里要注意,这个 loss 只在 future_prediction=True 时生效——这是配置文件中的一个布尔开关,决定了整个世界模型增强机制是否被激活。当这个开关关闭时,h_t1_pred 会退化为 h_t(即不预测未来,直接把当前视觉特征当作"未来"),整个系统就变成了一个不使用未来预测的普通 VLA 基线——论文消融实验中的 baseline 正是这样配置的。
6. Scheduled Sampling:从 Teacher 条件到 Student 条件的平滑过渡
6.1 问题的核心矛盾
回过头看 Flow Head 训练时面对的一个核心矛盾。训练 Flow Head 时,它需要一个"未来视觉子目标"作为条件来学习动作生成。这里存在一个经典的 train-test mismatch 问题(在序列生成领域被称为 exposure bias):训练时可以把真实的未来帧特征 h_t1_gt 当作条件——这是完美无噪声的;但推理时只能拿到 Student 预测 → decoder 展开后的 h_t1_pred——这是有误差的。核心问题在于:如果 Flow Head 只在完美条件下训练,推理时面对 Student 的"噪声"预测就会表现失常。
反过来,如果从训练初期就只用 Student 的预测当条件输入给 Flow Head,问题同样严重:这意味着训练早期 Student 的预测质量很差——它还没学会对齐 Teacher 的潜动作空间,输出的向量几乎是随机的——这些随机向量经过 decoder 展开后产生的子目标同样毫无意义,Flow Head 被迫在这种垃圾条件下训练,学到的速度场也是垃圾,形成一个"垃圾进、垃圾出"的恶性循环,最终可能导致整个训练坍缩。
这正是 Scheduled Sampling 这一经典技术(最早由 Bengio et al., 2015 提出用于序列模型的 exposure bias 问题)要解决的核心矛盾。在 LaWAM 的代码中,这个机制的实现集中在 _build_flow_future_condition 方法里——它决定了训练过程中每个样本使用 GT 子目标还是 Student 预测的子目标作为 Flow Head 的条件输入:
6.2 实现:sample-level 混合与 straight-through 桥接
# starVLA/model/framework/vlas/lawam.py
def _build_flow_future_condition(self, *, h_t1_pred, h_t1_gt):
if not self.training:
return h_t1_pred # 推理时只用 Student 预测,Teacher 已不可用
# 概率随训练步数线性增长:从 0.0(完全用 GT)到 1.0(完全用 Student 预测)
prob_pred = self._flow_h_t1_pred_prob()
bsz = h_t1_pred.shape[0]
# 每个样本独立抛硬币,决定用 GT 还是 Student 预测
pred_mask = (torch.rand(bsz, 1, 1, device=h_t1_pred.device) < prob_pred)
cond_future = torch.where(pred_mask, h_t1_pred, h_t1_gt)
if not self.model_cfg.detach_future_feature:
# Straight-through bridge:数值上不改变前向结果,但让梯度流回 Student
gt_mask = (~pred_mask).to(dtype=h_t1_pred.dtype)
cond_future = cond_future + gt_mask * (h_t1_pred - h_t1_pred.detach())
else:
cond_future = cond_future.detach()
return cond_future
def _flow_h_t1_pred_prob(self):
"""线性退火:从 start 到 end,经过 ramp_steps 步完成过渡"""
if not self.model_cfg.enable_flow_h_t1_scheduled_sampling:
return 1.0 # 不启用 scheduled sampling 时,始终用 Student 预测
start = self.model_cfg.flow_h_t1_pred_prob_start # 默认 0.0
end = self.model_cfg.flow_h_t1_pred_prob_end # 默认 1.0
ramp_steps = max(1, self.model_cfg.flow_h_t1_pred_ramp_steps) # 默认 20000
ratio = min(1.0, self._flow_train_step / ramp_steps)
return start + (end - start) * ratio
6.3 三个关键设计决策的解读
这段代码虽然只有十几行,但包含了三个精心设计的技术决策,每一个都有明确的工程动机,分别对应"如何在单个 batch 内混合两种条件来源"、"如何在选择了 GT 条件的样本上仍然让 Student 接收到梯度信号"和"如何随训练进程平滑地调节两种来源的混合比例"三个关键子问题。下面逐一展开。
第一个决策是 sample-level 混合而非 batch-level 切换。代码用 torch.rand(bsz, 1, 1) 为每个样本独立抛硬币,而不是在某个训练步数之后整个 batch 统一切换。这避免了"某一步之后 Flow Head 突然从完美条件切到噪声条件"带来的训练震荡。每个 batch 中总有一部分样本在用 GT、一部分在用 Student 预测,Flow Head 始终在"混合信号"环境下训练,过渡更平滑。
第二个决策是 straight-through 梯度桥接。cond_future + gt_mask * (h_t1_pred - h_t1_pred.detach()) 这一行数学上等于零(前向值不变),但在反向传播时为 Student 路径提供了额外的梯度通道。即使某个样本在本次前向时选了 GT 作为 Flow Head 的条件,Student 仍然能通过这条"幽灵梯度"接收到 Flow Head 的反馈。如果没有这条桥接,Student 只有在"轮到自己"时才能学到 Flow Head 的信号,学习效率会打折扣。这个技巧与 Gumbel-Softmax 中的 straight-through estimator 同源,核心思想都是"前向走离散/选择路径,反向走连续梯度"。
第三个决策是线性退火 schedule。概率从 0 匀速增长到 1,在 20000 步(约整个训练的前 1/3 到 1/2)内完成过渡。没有使用更复杂的 cosine、exponential 或 warmup-plateau-decay schedule,因为在实验中线性已经足够,而过于复杂的 schedule 反而增加了超参数调优的负担。这体现了一种"简单够用就不加复杂"的工程哲学。
难点提示:straight-through 桥接是整段代码中最不直观的一行。
cond_future + gt_mask * (h_t1_pred - h_t1_pred.detach())前向计算结果等于cond_future本身(因为h_t1_pred - h_t1_pred.detach()数值为零),但反向传播时h_t1_pred的梯度不为零,会传回 Student。可以把它理解为一条"隐形电线"——不改变电路的输出,但让电流(梯度)有了一条额外的回路。
7. 推理时的完整数据流:Teacher 消失,系统自持
7.1 从训练到推理的角色切换
先看推理时系统变成了什么样。训练完成后推理路径极其干净——整个 Teacher 路径(LAM encoder + VQ 量化 + 蒸馏损失计算)被完全移除,不再参与任何计算,也不占用推理时的显存。推理时的数据流是一条纯粹的、没有分支的前向链路,从观测一路走到动作输出中间不做任何迭代:
# starVLA/model/framework/vlas/lawam.py
@torch.inference_mode()
def predict_action(self, *, batch, guidance_scale=None, num_inference_steps=None, ...):
# Step 1-3: 共享编码(与训练时相同的 _run_shared_encoding_core)
shared = self._run_shared_encoding_infer(
prepared_batch=prepared_batch,
source="LatentWorldPolicyBackend.predict_action",
lam_features_with_no_grad=False,
)
# shared 内部做了什么:
# VLM 前向 → 提取 <ACT_PH> hidden → QFormer(Student) → pred_action_emb
# LAM 视觉编码器提取当前帧 h_t
# LAM decoder(h_t, pred_action_emb) → h_t1_pred(未来子目标)
# Step 4: Flow Head 从噪声积分到动作
actions = self.flow.sample_actions_cfg(
h_t=shared.h_t, # 当前视觉 token
h_t1_star=shared.h_t1_pred, # Student 驱动的未来子目标
h_vlm=shared.h_vlm, # VLM 上下文
state=batch["state"], # 机器人本体状态
cfg_scale=guidance_scale, # CFG 引导强度
num_inference_steps=num_inference_steps, # 默认 4 步 Euler 积分
)
return actions
7.2 推理延迟的来源分析
整条推理链路不依赖任何"未来帧"输入——系统不需要等待未来发生就能做出决策。根据论文公布的延迟测试数据,187 毫秒的端到端延迟分布在四个阶段:Qwen-VL 的前向传播是主要耗时来源(约 2.3B 参数的 Transformer 推理占据了大部分时间预算),QFormer 只做一次轻量的 cross-attention 计算(几乎可忽略不计),LAM decoder 做一次 230M 参数的 Transformer 前向传播就完成了世界模型推理,Flow Head 做 4 步 Euler 积分(每步执行一次 DiT 前向)完成从噪声到动作的生成。
对比像素空间世界模型的推理流程——先跑一个视频扩散模型生成未来帧序列(数十到数百步迭代去噪,每步都在 megapixel 分辨率上做 U-Net 前向),再用策略网络从生成的视频中提取信息——LaWAM 的世界模型部分只做了一次 Transformer 前向传播,没有任何迭代过程。这不是"把扩散做快了",而是从根本上改变了"未来以什么形式进入策略"的问题定义。
7.3 CFG 在速度空间的应用
推理时 Flow Head 还使用了 Classifier-Free Guidance(简称 CFG)这一在图像生成领域已经被广泛验证的技术来增强未来子目标的条件影响力。与图像扩散模型中在样本空间做 CFG 不同的是,LaWAM 选择在速度空间做 CFG——训练阶段以一定概率 cfg_drop_prob 随机丢弃未来子目标条件输入,让网络同时学会"有条件下的速度"和"无条件下的速度"两种模式,推理时再把两者按引导系数做线性外推:
当引导系数 时,条件与无条件之间的差异被放大,生成的动作轨迹会更贴合给定的未来子目标所指示的运动方向——子目标的"引导力"变强了。这种选择在速度空间而非样本空间做 CFG 的实践,与 NVIDIA 的 GR00T 框架和 Isaac Lab 中工业级机器人策略系统的实现方式保持一致,已被证明在连续动作生成场景下数值更稳定、收敛更快。
8. Flow Matching 动作头:为什么用整流流而非扩散
8.1 从扩散到整流流的范式迁移
动作生成头的选择是 LaWAM 的另一个关键设计点。近两年机器人策略领域经历了一轮从 Diffusion Policy 到 Flow Matching Policy 的范式迁移,代表工作包括 pi0(Physical Intelligence, 2024)和 Conditional Flow Matching for Robotic Manipulation。迁移的核心动机很直接:Rectified Flow(整流流)把噪声到数据之间的路径拉成一条直线,使得推理时只需要 4 步 Euler 积分就能得到高质量的动作轨迹,而经典扩散模型通常需要 20-100 步去噪。对于要求低延迟的实时机器人控制场景,这个差距足以决定一套方案能否上机部署。
整流流的训练目标极其简洁。定义噪声 和真实动作 之间的线性插值路径 ,其中 是生成进度。沿这条直线的瞬时速度是常量 ,> 解读:整流流之所以只需 4 步就够,是因为"沿直线走的速度是常量"——网络只要在任意一点学会预测这个常量,就等于学会了整条路径。扩散模型的路径是弯曲的,网络必须在很多点上分别学不同的"瞬时方向",所以需要更多步积分才能准确走完全程。
正因为整流流的插值路径是一条直线、沿路径的瞬时速度是常量,网络的学习目标变得极其清晰而友好——它只需要在训练时被随机采样到的任意中间点 处,输出那个与当前位置和当前时间步都无关的常量速度向量 即可,完全无需像经典扩散模型那样学习一个随时间非线性变化的复杂分数函数。训练损失因此可以写成如下简洁的形式:
8.2 物理时间与 Flow 时间的双时间轴
LaWAM 的 Flow Head 实现中有一个容易混淆的设计:它同时维护两套完全不同的"时间"概念。第一套是 flow time ,描述"从噪声到数据的生成进度",通过 DiT 内部的 TimestepEncoder 编码为调制向量,注入到 AdaLN 中。第二套是物理时间(秒),描述"每个动作 token 发生在真实时间轴的哪个时刻",通过 ContinuousTimeEncoder 编码为正弦位置嵌入,加到动作 token 上。
物理时间编码的存在是为了解决混频训练问题。真实机器人数据来自不同平台,控制频率从 5 Hz 到 20 Hz 不等。同一个"第 3 个动作 token",在 5 Hz 平台上对应 0.6 秒后的动作,在 20 Hz 平台上只对应 0.15 秒后的动作。如果只用 token 索引做位置编码,模型在混合数据上训练时会彻底混乱。物理时间编码消除了这个歧义:两个不同频率平台上"0.3 秒后的动作"会拿到相同的时间编码,即便它们的 token 索引不同。
8.3 repeated_diffusion_steps:同一 batch 的多重采样
代码中有一个初看起来"多余"甚至"浪费算力"的操作——在计算 flow loss 之前,先把整个 batch 的所有条件张量(包括视觉 token、VLM 上下文、状态向量、动作标签等)沿 batch 维度重复 repeated_diffusion_steps 次(默认配置为 4 次),然后再把这个扩展了 4 倍的 batch 送入 Flow Head 做前向和损失计算。这不是数据增强——图像和文本没有被变换过——而是一种专门针对 flow matching 训练特点的监督密度提升技巧,下面的代码展示了这个过程:
repeat_steps = int(self.model_cfg.repeated_diffusion_steps)
h_t_rep = shared.h_t.repeat(repeat_steps, *([1] * (shared.h_t.ndim - 1)))
# ... 所有条件张量都做同样的 repeat ...
loss_flow = self.flow(h_t=h_t_rep, h_t1_star=h_t1_cond_rep, ...)
这不是数据增强,而是增加 flow matching 的监督密度。每次 repeat 会独立采样不同的 flow time 和噪声 ,相当于对同一组条件(同一张图片、同一条指令、同一个子目标)生成 4 组不同的"噪声→数据"配对。这使得 Flow Head 在单个 gradient step 中看到了更丰富的 组合,加速了速度场的学习收敛,代价是约 4 倍的 Flow Head 前向开销(但不增加 VLM 和 LAM 的开销,因为它们的输出被复用)。
9. 总结
LaWAM 的代码设计证明了一件事:在具身智能中引入世界模型不必然意味着引入一个昂贵的视频生成器。通过把阶段一 LAM 的 decoder 直接复用为隐空间世界模型、用 Teacher-Student 蒸馏弥合训练时后验与推理时先验的信息鸿沟、再用 Scheduled Sampling 平滑条件切换,LaWAM 在 LIBERO 上达到 98.6% 成功率、在 RoboTwin 上达到 91.22% 成功率、推理延迟压到 187 毫秒——这三个数据点合在一起,足以让这套"隐空间世界模型 + 蒸馏"范式进入下一代高效具身策略的候选清单。但细粒度可形变物体的动态覆盖(如绳索、布料操作)和移动相机下的自运动解耦,这两个缺口依然悬着——前者决定它能否啃下高难度精细操作,后者决定它能否从固定相机走向人形和移动平台。
178