转载自公众号:敢敢AUTOHUB
0. 简介
AHA-WAM 解决的是一个非常具体、非常工程的痛点——世界模型生成视频很慢,机器人控制要求很快,过去的做法强行让两者用同一个节拍跑,结果就是要么控制频率被视频生成拖垮,要么视频质量被高频采样压缩。AHA-WAM 的核心动作只有一个:把"慢"的视频规划和"快"的动作执行彻底拆开,让一次昂贵的视频刷新摊销到多次廉价的动作更新里。
这篇博客会带你从 VLA 与世界模型的本质差异讲起,一路拆解双 DiT 架构、OVCR 路由、视野自适应偏移训练,再回到代码仓库验证这些抽象是如何在 PyTorch 里被实例化的,最后看它在 RoboTwin 上 92.80% 的成功率和 56.95Hz 的控制频率是怎么来的。
项目主页:https://serene-sivy.github.io/aha-wam/
代码仓库:https://github.com/serene-sivy/AHA-WAM
1. 机器人"思考"与"行动"的节奏错配
1.1 从 VLA 到世界模型:监督信号的稀疏问题
近几年机器人操作领域的主线是视觉-语言-动作模型(VLA),代表是 RT-1、RT-2、 这一系。它们的思路很直接:把视觉观测和语言指令直接映射成机器人动作,靠大规模模仿学习把视觉-语言大模型的知识迁移到控制上。这套范式在语义理解上表现出色,但有一个结构性短板——动作标签对底层物理动态来说是极其稀疏的监督信号。模型知道"抓起杯子"这个动作序列长什么样,却不一定理解杯子被推到桌边会掉下去这种物理因果。这导致 VLA 在见过的任务上表现很好,一旦面对没见过的物体或新环境,泛化能力就明显不足。
世界动作模型(World-Action Model, WAM)走了一条不一样的路。它不只预测动作,还联合预测未来的视频帧——也就是"世界将如何演变"。通过在海量视频上预训练,WAM 能把物理规律、运动轨迹、物体交互内化进模型参数里,从而在无需额外动作标签的情况下实现对新任务的零样本泛化。这里的关键是,视频预测充当了一种密集的物理监督:每一帧像素的变化都在逼模型学会"动作如何改变世界",这比稀疏的动作标签信息量大得多。
1.2 同步执行的致命瓶颈:视频生成拖垮控制频率
但现有 WAM 在实际部署时撞上了一堵墙。它们通常把世界预测(视频生成)和动作执行绑定在相同的时序节拍上。这意味着,为了生成高频的控制动作,模型必须以同样的高频去预测视频帧。视频生成本身就是计算密集型操作——一个 5B 参数的 Video DiT 跑一次前向要几十到几百毫秒,如果每个控制周期都要重新生成一遍未来帧,推理延迟会直接爆炸。论文给出的对比数字很说明问题:纯做联合世界-动作建模的 Motus 单步延迟高达 1866.10 毫秒,对应控制频率只有 0.54Hz,这种速度在真实机器人闭环控制里完全不可用。
难点提示(为什么同步是个伪需求):直观理解是,人开车时眼睛看路的频率和手打方向盘的频率本来就不一样。你大脑里对"前方路况会怎么演变"的预判可能两三秒更新一次,但手上的微调是连续不断的。强行让"看路预判"和"手上动作"同频,等于要求你每打一次方向盘就重新把整条路想象一遍——既没必要,又慢得离谱。AHA-WAM 做的就是把这两个频率解绑。
1.3 这篇博客承诺给你什么
读完这篇博客,你会拿到三层理解。第一层是范式视角——为什么世界模型比纯 VLA 更适合注入物理先验,以及"同步执行"这个看似自然的设计为什么是个结构性错误。第二层是架构视角——AHA-WAM 怎么用双 DiT、OVCR、视野自适应偏移训练、滚动 KV 记忆这四个机制把异步执行做对。第三层是代码视角——通过四段从仓库直接抽出来的真实代码,看双 DiT 协同、上下文路由、偏移训练是如何在工程上落地的。
这里要厘清的是,VLA、imagine-then-act、joint world-action modeling 三条路线在监督信号、推理延迟、泛化能力上各有取舍,AHA-WAM 选择在第三条路线内部做异步化升级,而不是另起炉灶——这种在主流路线上做关键性改造的做法比造新名词更工程化,也更容易被同期工作吸收为标准技术。
2. AHA-WAM 核心思路:慢规划器 + 快执行器
2.1 设计动机:时序节拍解耦
AHA-WAM 的核心洞察是,世界预测和动作执行应该运行在不同的时间尺度上。世界模型受益于更长、更慢的视野——它需要看到物体从 A 移动到 B、杯子被推倒的完整过程,才能学到有用的物理先验。但机器人控制需要快速闭环修正——当手臂轨迹偏离预期 1 厘米时,你不能等一秒钟再反应。传统 WAM 把这两个频率强制对齐,本质上是把"长期规划"和"短期执行"混为一谈。
AHA-WAM 重新组织了世界-动作建模的时间轴:低频的 Video DiT 充当长期视野的世界规划器(planning horizon 帧),每次前向会预测未来很长一段视频,同时暴露出可复用的逐层潜在上下文,编码了长时域的场景演变信息。高频的 Action DiT 是短期的闭环执行器(action horizon ),直接接收机器人本体感觉(proprioceptive state),在高频控制循环中通过逐层交叉注意力机制查询规划器上下文,解码出短期可执行的动作块。这样一次昂贵的视频刷新可以摊销到多次廉价的动作更新里——动作分支既受益于学到的视觉动态规律,又避免了每次控制循环中重新生成未来帧的巨大开销。
直觉理解(异步执行的日常类比):你在陌生城市开车去目的地。导航每隔 30 秒给你一次"接下来 2 公里路况"的预判,但你的方向盘、油门、刹车是连续调整的——你不会等下一次导航更新才动手。AHA-WAM 的 Video DiT 就是那个 30 秒刷新一次的导航,Action DiT 是你的手脚——它拿着导航给的"2 公里预判",结合实时看到的车道线、前车距离,做高频微调。导航更新慢但视野长,手脚响应快但只管眼前几米。
2.2 "Horizon-Adaptive"不是在线调整视野,而是训练时适应任意相位偏移
这里容易被标题误导。论文标题里的"Horizon-Adaptive"不是指在推理时动态调整视频或动作的视野长度(那会引入另一层复杂度),而是指执行器在训练时学会在规划器视野内任意相位偏移下消费上下文。具体来说,视频规划器建模 的未来,动作执行器预测 的动作块,这里 和 并不总是对齐的——在异步流式推理中,执行器可能从规划器视野的中间某个位置开始消费上下文。
如果训练时总是固定相位(比如总是让第一个动作块从规划器起点开始),模型就会过拟合到这个对齐关系,推理时稍有偏移就垮掉。所以 AHA-WAM 在训练时随机平移动作块在视频视野内的起始位置,逼执行器学会在各种相位下都能正确消费规划器上下文。这就是"Horizon-Adaptive"的真正含义——对规划器-执行器相位偏移的鲁棒性。
2.3 四个机制的协同闭环
换句话说,AHA-WAM 用四个机制撑起异步执行,每一条都对准一个具体的工程痛点。双 DiT 架构让 Video DiT 和 Action DiT 分别建模世界和动作,通过逐层联合注意力耦合,但各自运行在不同频率上。**OVCR(Observation-Guided Video-Context Routing)**在每次动作更新前,用最新的视觉观测构建路由查询,对缓存的规划器 K/V 上下文做门控残差更新,让每个动作块都能看到与当前视觉证据对齐的规划器表征,而无需重新运行 Video DiT。Horizon-Adaptive Offset Training在训练阶段把动作块在视频视野内随机平移,迫使执行器学会在各种相位偏移下依然能够正确消费规划器上下文。
Rolling KV Memory则在规划器内部维护一个固定大小的 FIFO 滚动记忆,存储过去几次刷新的历史视频状态,让规划器的时间感受野能延伸到更长的历史,应对长期任务中物体被遮挡或子目标已完成的情况。进一步看,这四个机制环环相扣。双 DiT 给出分工,OVCR 让异步界面实时对齐,Offset Training 让执行器对异步偏移保持鲁棒,Rolling Memory 让规划器记住长历史。缺了任何一环,异步执行都会出问题——论文的消融实验会证明这一点。
3. 双 DiT 架构:Video-DiT 长期规划 + Action-DiT 闭环执行
3.1 模型架构总览
AHA-WAM 建立在一个双 Diffusion Transformer 架构上。视觉观测由预训练的 VAE 编码为视频潜在(video latents),语言嵌入来自文本编码器,两者都对双分支可见。Video DiT 接收视觉潜在 token 并预测未来视频潜在状态,覆盖更长的规划视野 。Action DiT 接收噪声动作 token 和本体感觉 token,在给定本体状态反馈的情况下对动作块去噪。视觉反馈不是直接把密集视觉 token 塞给高频动作分支(那会让动作分支变慢),而是通过规划器视频上下文间接注入——这个上下文后面会被 OVCR 适配。
两个 DiT 的核心接口是逐层规划器视频上下文,其中 表示规划器分支, 索引最新一次规划器刷新, 索引 Transformer 层。这个上下文是一次 Video DiT 前向暴露出来的潜在世界规划表征,可以被后续多次 Action DiT 前向复用。它和滚动 KV 记忆不一样——后者是规划器内部跨刷新的历史存储,而规划器视频上下文是当前刷新对外暴露的接口。
训练时,Video DiT 在全因果视频 mask 下预测未来视频潜在,这逼它学习前向场景动态,同时塑造规划器上下文里的视觉动态监督。Action DiT 被 mask 掉对未来视频 token 的注意力,所以推理时未来视频预测这条路径可以完全移除——未来视频预测是训练信号,规划器视频上下文才是推理界面。
3.2 逐层联合注意力:动作 DiT 如何查询规划器上下文
每次动作更新时,原始规划器上下文 先被 OVCR 适配为当前块专用的上下文 (OVCR 细节下一节讲)。适配的意义在于让静态缓存的规划器表征跟上当前的视觉变化。然后 Action DiT 通过逐层联合注意力对动作块去噪,把自己的动作 token 和适配后的规划器 K/V 拼在一起做注意力:
这里 、、 是 Action DiT 的查询、键、值投影, 是规划器条件下的动作隐状态。注意力把动作 token 的键值和规划器的键值拼接,意味着每个动作 token 既看自己也看世界规划。这种耦合保留了 WAM 式的视觉动态与动作生成交互,同时把昂贵的 Video DiT 计算摊销到多次高频动作更新里,是异步设计能同时保住性能和速度的核心。
3.3 联合训练目标:流匹配同时优化世界和动作
AHA-WAM 用联合流匹配(flow matching)目标训练,这是近两年扩散类策略的主流训练范式。对于目标变量 (可以是动作块 或未来视频潜在 ),采样高斯噪声 和流时间 ,构造从数据到噪声的线性插值样本 ,然后让模型预测对应的速度场(即从噪声指向数据的方向):
动作生成和视频协训练分别实例化为两个并列的流匹配损失项,前者监督短期可执行动作块的去噪,后者监督长视野未来视频潜在的预测,两条损失共享同一套流匹配框架但作用在不同的目标变量上。这种并列设计让世界建模和动作执行在训练时互为监督信号,推理时则各司其职,下面用公式把它们的具体形式展开:
最终目标把两条损失加权求和成 ,其中 平衡动作学习和视频动态学习—— 太大会让模型过度关注视频质量而牺牲动作精度,太小则丢失物理先验,论文把两者权重设为相等。部署时,AHA-WAM 移除显式的未来帧解码——Video DiT 只刷新规划器视频上下文,Action DiT 用这个上下文生成闭环动作块。
工程价值(为什么不直接扔掉视频分支):你可能会问,既然推理时不生成视频,为什么训练时还要视频损失?答案是,视频损失是物理监督的来源。纯靠动作标签学到的是"在这个场景下该做什么动作",靠视频帧学到的是"这个动作会让世界怎么变化"。后者才是泛化的基础——模型见过一万个"推杯子"的视频后,它知道杯子被推会滑动、被推到桌边会掉下去,这些物理因果不需要动作标签也能学。所以视频损失在训练时是密集监督,推理时虽然不显式生成帧,但那些物理规律已经编码进规划器上下文里了。
3.4 代码实现:ActionDiT 类的核心结构
为了让读者看清高频执行器到底怎么消费规划器上下文,下面从仓库 src/ahawam/models/wan22/action_dit.py 抽取 Action DiT 的核心类定义(简化版,保留关键方法签名和注释)。这段代码虽然删掉了很多工程细节,但核心接口都保留了,重点看 __init__ 里的轻量参数配置和 _merge_obs_context 方法如何合并规划器上下文:
# src/ahawam/models/wan22/action_dit.py
class ActionDiT(nn.Module):
"""
Action Diffusion Transformer for high-frequency closed-loop execution.
Queries planner video context through layerwise joint attention.
"""
ACTION_BACKBONE_SKIP_PREFIXES = ("action_encoder.", "head.")
ACTION_BACKBONE_META_KEYS = (
"hidden_dim", "ffn_dim", "num_layers", "num_heads",
"attn_head_dim", "text_dim", "freq_dim", "eps",
)
def __init__(
self,
hidden_dim: int, # 1024 for compact executor
action_dim: int, # robot action dimension (e.g. 14 for dual-arm)
ffn_dim: int,
text_dim: int,
freq_dim: int,
eps: float,
num_heads: int,
attn_head_dim: int,
num_layers: int, # same depth as Video DiT for layerwise coupling
autoregressive_teacher_forcing: bool = False,
action_chunk_size: int = 1,
action_teacher_forcing_mask_mode: str = "stage1_chunkwise",
use_gradient_checkpointing: bool = False,
):
super().__init__()
self.hidden_dim = hidden_dim
self.action_dim = action_dim
self.num_layers = num_layers
self.autoregressive_teacher_forcing = autoregressive_teacher_forcing
self.action_chunk_size = action_chunk_size
# Action token encoder: (B, T, action_dim) -> (B, T, hidden_dim)
self.action_encoder = nn.Sequential(
nn.Linear(action_dim, hidden_dim),
nn.GELU(approximate="tanh"),
nn.Linear(hidden_dim, hidden_dim),
)
# Time embedding for diffusion timestep
self.time_embedding = nn.Sequential(
nn.Linear(freq_dim, hidden_dim),
nn.SiLU(),
nn.Linear(hidden_dim, hidden_dim),
)
self.time_projection = nn.Sequential(
nn.SiLU(), nn.Linear(hidden_dim, hidden_dim * 6)
)
# Stack of DiT blocks that consume planner context
self.blocks = nn.ModuleList([
DiTBlock(
hidden_dim=hidden_dim,
attn_head_dim=attn_head_dim,
num_heads=num_heads,
ffn_dim=ffn_dim,
eps=eps,
)
for _ in range(num_layers)
])
# Final action head
self.head = nn.Linear(hidden_dim, action_dim)
# Rotary position embedding frequencies
self.freqs = precompute_freqs_cis(attn_head_dim, end=1024)
self.use_gradient_checkpointing = use_gradient_checkpointing
def _merge_obs_context(
self,
batch_size: int,
context: torch.Tensor, # text context
context_mask: torch.Tensor,
obs_context: Optional[torch.Tensor], # chunk-aligned obs context from OVCR
obs_context_mask: Optional[torch.Tensor],
):
"""
Validate and merge chunk-aligned obs_context (adapted planner K/V)
into the text context before feeding to action blocks.
Returns (merged_context, flat_obs_context_mask, obs_num_chunks, obs_tokens_per_chunk).
"""
if obs_context is None:
return context, None, 0, 0
# obs_context must be 4D [B, N, L, D] (chunk-aligned)
if obs_context.ndim != 4:
raise ValueError(
f"`obs_context` must be 4D [B, N, L, D] (chunk-aligned), "
f"got shape {tuple(obs_context.shape)} with ndim={obs_context.ndim}."
)
obs_num_chunks = int(obs_context.shape[1])
obs_tokens_per_chunk = int(obs_context.shape[2])
# Flatten chunk dimension: [B, N, L, D] -> [B, N*L, D]
obs_flat = obs_context.reshape(
batch_size, obs_num_chunks * obs_tokens_per_chunk, obs_context.shape[3]
)
# Merge text and obs context along sequence dimension
merged = torch.cat([context, obs_flat], dim=1)
# Construct flat obs mask
if obs_context_mask is not None:
flat_obs_mask = obs_context_mask.reshape(batch_size, -1)
else:
flat_obs_mask = torch.ones(
batch_size, obs_num_chunks * obs_tokens_per_chunk,
dtype=torch.bool, device=obs_context.device
)
merged_mask = torch.cat([context_mask, flat_obs_mask], dim=1)
return merged, merged_mask, obs_num_chunks, obs_tokens_per_chunk
这段代码虽然只是类定义的骨架,删掉了前向传播的去噪循环细节,但已经把 Action DiT 作为高频执行器的设计意图暴露得很清楚。下面拆开看三个关键点,它们分别对应输入接口、模型结构、时间嵌入三个工程决策,每一条都能在 RoboTwin 的推理日志里找到对应:
1. Action DiT 的输入接口:噪声动作 token + 本体状态 + 文本上下文 + chunk-aligned obs_context(这就是 OVCR 适配后的规划器上下文)。
注意 obs_context 是 4D [B, N, L, D],其中 是动作块数, 是每块对应的规划器 token 数——这个 4D 结构让 Action DiT 能把规划器上下文和动作块对应起来。
2. 逐层 DiT 块:self.blocks 是一个 nn.ModuleList,层数和 Video DiT 相同。每个 DiTBlock 内部会做联合注意力——查询自己的动作 token 的同时查询合并后的 obs_context(即规划器 K/V)。
3. 时间嵌入:去噪步骤 通过 time_embedding 和 time_projection 编码成条件向量,加到每个 DiT 块里。这是标准的 diffusion model 做法。
4. OVCR:观测引导的视频上下文路由
4.1 问题:缓存的规划器上下文会"过期"
异步执行让一次规划器上下文被多次动作块复用,这摊销了 Video DiT 的计算开销,但也带来一个新问题——在下次规划器刷新之前,机器人状态和视觉场景可能已经变化了。如果动作 DiT 一直用着几百毫秒前刷新的规划器上下文,它看到的是"过去的世界",做出的动作可能和当前实际场景错位。换句话说,规划器上下文是静态缓存,但物理世界是动态变化的,两者之间需要一座桥。
OVCR(Observation-Guided Video-Context Routing)就是这座桥。它用最新的视觉观测构建一组紧凑的路由查询,对缓存的规划器 K/V 上下文做门控残差更新,让每个动作块都能看到与当前视觉证据对齐的规划器表征——而无需重新运行 Video DiT。这是论文最关键的工程贡献,把规划器上下文复用从"静态缓存"升级成了"观测条件检索"。
难点提示(OVCR 与传统 cross-attention 的区别):你可能会问,让 Action DiT 直接对当前视觉 token 做 cross-attention 不就行了?为什么要绕一圈通过路由器更新 K/V?关键在于规模——视频 token 是密集的(一帧 384×320 经过 VAE 压缩后还有几百到几千个 token),如果每次动作更新都让 Action DiT 对这些密集 token 做注意力,高频分支立刻就慢下来了。OVCR 用 32 个可学习的"路由查询"先把视觉信息压缩成轻量的更新增量,再把增量加回到规划器 K/V 上,整个过程只在轻量级路由器上做计算,Action DiT 仍然只面对适配后的规划器 K/V,不接触密集视觉 token。
4.2 OVCR 的三步路由流程
OVCR 工作流程分三步。第一步,对每个动作块 ,把对齐的观测上下文拆成视觉 token 和本体感觉 token 。本体感觉很紧凑(直接和瞬时机器人状态绑定),由轻量级编码器映射成一个状态 token 直接送入 Action DiT。视觉反馈走间接路径——通过上下文路由注入。
第二步,构造观测引导的路由查询。给定 个可学习的基础查询 (论文设 ),OVCR 用注意力池化把这些基础查询和当前视觉 token 结合,让查询槽带上当前观测的语义——相当于用最新画面给查询槽充电。这个过程本质上是把密集的视觉 token 压缩成 32 个轻量查询,用公式表达就是:
这里 是一个轻量视觉投影模块,把原始视觉 token 映射到查询空间。注意力池化的结果是 32 个紧凑的、观测条件化的查询槽,它们专门用于从规划器视频上下文里检索与当前动作块最相关的信息——本质上是把密集视觉证据压缩成少量可路由的查询。
第三步,对每一层 分别做读和写两个动作——这种逐层处理让不同深度的规划器表征都能被当前观测改写。靠近输入的层可能需要大幅更新(视觉变化敏感),靠近输出的层可能保持原本语义结构。先用路由查询读出该层的规划器特征,把路由查询当成 Q、规划器 K/V 当成待检索的记忆库:
然后用一个轻量级逐层路由器 预测残差 K/V 更新 ——这个残差就是当前观测相对于缓存上下文的修正量,体现了观测变化。最后通过门控残差更新把修正量加回到原始规划器 K/V 上,生成块专用上下文:
其中 是学到的门控系数,控制残差更新的强度——它让模型自己决定每一层、每个位置该被当前观测改写多少。门控系数通常在训练时自动学出分层差异化更新策略。适配后的上下文 最终被 Action DiT 通过逐层联合注意力消费,完成从静态缓存到观测对齐表征的转换。
工程价值(门控残差为什么不是直接覆盖):门控的设计很关键。如果直接用 替换 ,模型会失去规划器原本编码的长视野上下文;如果不做更新(),就退化成静态缓存。
门控让模型在训练时学到"哪些层、哪些位置需要被当前观测改写多少"——靠近输入的层可能需要大幅更新(视觉变化敏感),靠近输出的层可能保持规划器原本的语义结构。这种"分层差异化更新"是纯 cross-attention 做不到的。
4.3 OVCR 的本质:把缓存复用变成观测条件检索
总结起来,OVCR 把规划器上下文复用从静态缓存升级成观测条件检索与适配。Video DiT 仍然在每次更新的关键路径之外(异步在后台跑),但每次动作块都能拿到一个与最新视觉证据对齐的规划器表征。这是异步执行能够工作的根基——没有 OVCR,缓存的上下文会快速过期,动作就会和实际场景脱节。
4.4 代码实现:mot.py 的块路由注意力
为了让读者看清 OVCR 的读-写两步到底怎么落地,下面从 src/ahawam/models/wan22/mot.py 抽取核心路由注意力的实现。这个文件名里的 mot 代表 Multi-Objective Training(联合世界-动作建模),它对应前面公式里的读规划器特征和块对齐联合注意力两个环节,代码分两个方法分别实现这两步:
# src/ahawam/models/wan22/mot.py
def _chunk_routed_attention(
self,
*,
q: torch.Tensor, # routing queries Z^q_t [B, Q, D]
k: torch.Tensor, # planner keys K^p [B, T, D]
v: torch.Tensor, # planner values V^p [B, T, D]
ctx_mask: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""
Read planner context using observation-guided routing queries.
Output: R^ell_t = Attn(Z^q, K^p, V^p) [B, Q, D]
"""
def _forward(q_in, k_in, v_in):
return flash_attention(
q=q_in, k=k_in, v=v_in,
num_heads=self.num_heads,
ctx_mask=ctx_mask,
)
# Gradient checkpointing to save memory during training
if self.mot_checkpoint_mixed_attn and self.training:
return cast(
torch.Tensor,
torch_checkpoint.checkpoint(
_forward, q, k, v, use_reentrant=False,
),
)
return _forward(q, k, v)
def _forward_chunk_routed_prior_only_attention(
self,
*,
q_action: torch.Tensor, # [B, T_a, D] action queries
k_action: torch.Tensor, # [B, T_a, D] action keys
v_action: torch.Tensor, # [B, T_a, D] action values
updated_keys: torch.Tensor, # [B, N, L, D] OVCR-adapted planner K
updated_values: torch.Tensor, # [B, N, L, D] OVCR-adapted planner V
action_seq_len: int,
chunk_size: int,
) -> torch.Tensor:
"""
Per-chunk joint attention between action tokens and adapted planner context.
Each action chunk attends only to its corresponding planner K/V (chunk-prior alignment).
"""
batch_size = int(q_action.shape[0])
num_chunks = int(updated_keys.shape[1])
# Reshape action tokens to per-chunk view
def chunk_view(x: torch.Tensor) -> torch.Tensor:
return x.reshape(batch_size, num_chunks, int(chunk_size), x.shape[-1])
q_prior = chunk_view(q_action) # [B, N, chunk_size, D]
k_prior = chunk_view(k_action)
v_prior = chunk_view(v_action)
# Concat action K/V with adapted planner K/V along sequence dim
# action chunk i sees: own action tokens + adapted planner context for chunk i
k_cat = torch.cat([k_prior, updated_keys], dim=2)
v_cat = torch.cat([v_prior, updated_values], dim=2)
# Joint attention: action queries attend to merged action+planner K/V
return flash_attention(
q=q_prior,
k=k_cat,
v=v_cat,
num_heads=self.num_heads,
)
这段代码做了两件事。_chunk_routed_attention 实现 OVCR 第一步——用路由查询读规划器 K/V,对应公式里的 。_forward_chunk_routed_prior_only_attention 实现 Action DiT 的逐层联合注意力——每个动作块只关注自己对应的适配后规划器 K/V,避免跨块信息泄漏。工程细节上,updated_keys 的形状是 [B, N, L, D],其中 是动作块数(论文设 4 块), 是每块对应的规划器 token 数。这种"块对齐"结构让规划器上下文能精确映射到执行器的每个动作块。
5. Horizon-Adaptive Offset Training:让执行器适应任意相位偏移
5.1 异步执行带来的新问题:相位漂移
异步流式执行让规划器和执行器的相对时间相位不再固定。如果训练时总是用一个固定的对齐方式(比如规划器视野的起点总是和动作块的起点对齐),Action DiT 就会过拟合到这一种相位关系。但部署时执行器跑得比规划器快——当执行器跑到规划器视野的中间某个位置时,它需要从中间消费规划器上下文,相位就漂了。如果模型只在 这一种对齐下训练过,遇到 或 时就会表现糟糕。
5.2 训练时的随机相位偏移
AHA-WAM 引入视野自适应偏移训练(Horizon-Adaptive Offset Training)来解决这个问题,核心思路是在训练时人为制造各种相位偏移。设视频规划视野是 帧,动作块视野是 帧。对每个训练片段,从均匀分布采样一个偏移量,让动作块的起点在规划器视野内随机滑动:
然后把动作块网格在规划器视野内整体平移 帧,让动作块从规划器视野的不同位置开始消费上下文。这种随机平移相当于模拟了部署时规划器刷新和执行器更新不对齐的所有情况。动作目标在偏移对齐的块上计算损失,使模型见过所有可能的相位关系,用公式表达就是对偏移采样期望:
其中 表示从规划器视野内偏移对齐位置开始的动作块。因为规划器-执行器对齐关系按动作块大小是周期性的,采样 就覆盖了所有可能的块级相位——这相当于让模型见过所有相对偏移情况,部署时不论相位漂到哪都能正确消费上下文。
类比(驾校教练教并线):教练教你并线时,不会只在一种车距下让你练。他会让你练"前车 5 米""前车 10 米""前车 20 米"各种距离下的并线,这样你考试时遇到任何车距都能反应。AHA-WAM 的偏移训练就是这个意思——把所有可能的"规划器-执行器相位差"都让模型练一遍,部署时遇到任何相位都能稳。
5.3 代码实现:动作偏移的归一化与校验
为了让读者看清偏移训练在工程上怎么校验和分发,下面从 src/ahawam/models/wan22/ahawam.py 抽取处理动作偏移的核心方法。这个文件是 AHA-WAM 的主类实现,它负责把 dataloader 随机生成的偏移规范化并校验合法性,确保偏移值在运行时不会越界,也不会让执行器消费到规划器视野之外的上下文:
# src/ahawam/models/wan22/ahawam.py
def _has_action_offset(self, sample: dict[str, Any]) -> bool:
"""Check if the training sample carries a horizon-adaptive offset."""
return "action_offset" in sample
def _normalize_action_offsets(
self,
sample: dict[str, Any],
batch_size: int,
) -> torch.Tensor:
"""
Normalize and validate per-sample action offsets for horizon-adaptive training.
Each sample's action chunk grid is shifted by `delta in [0, max_action_offset)`
inside the video planning horizon. This ensures the executor learns to consume
planner context at any phase relative to the planner refresh boundary.
"""
raw_offset = sample.get("action_offset", 0)
offsets = torch.as_tensor(raw_offset, device=self.device, dtype=torch.long)
# Broadcast scalar to batch
if offsets.ndim == 0:
offsets = offsets.expand(batch_size)
if offsets.ndim != 1 or int(offsets.shape[0]) != batch_size:
raise ValueError(
"`action_offset` must be scalar or [B], "
f"got shape {tuple(offsets.shape)} for batch_size={batch_size}."
)
# Validate range: 0 <= offset <= max_action_offset
max_action_offset = int(getattr(self, "max_action_offset", 0))
if bool((offsets < 0).any().item()):
raise ValueError(
f"`action_offset` must be nonnegative, "
f"got {offsets.detach().cpu().tolist()}."
)
if max_action_offset > 0 and bool((offsets > max_action_offset).any().item()):
raise ValueError(
"`action_offset` exceeds configured max_action_offset: "
f"offsets={offsets.detach().cpu().tolist()} max={max_action_offset}."
)
return offsets
def _validate_runtime_action_horizon(self, action_horizon: int) -> None:
"""
At inference, the runtime action horizon must accommodate the configured
max_action_offset—otherwise the executor cannot consume planner context at
the maximum trained phase.
"""
max_offset = int(getattr(self, "max_action_offset", 0))
if max_offset > 0:
if action_horizon < max_offset + self.action_chunk_size:
raise ValueError(
f"With max_action_offset={max_offset}, runtime action_horizon must be "
f">= max_action_offset + action_chunk_size ({max_offset + self.action_chunk_size}), "
f"got {action_horizon}."
)
def training_loss(self, sample: dict[str, Any], inputs, tiled: bool = False):
"""Dispatch to offset-aware training loss when action_offset is provided."""
if self._has_action_offset(sample):
return self._training_loss_action_offset(sample, inputs=inputs, tiled=tiled)
# ...fall through to standard training loss
这段代码虽然主要是校验逻辑,看起来像防御性编程,但它暴露了 horizon-adaptive offset training 在工程落地时必须处理的几个边界条件。没有这些校验,模型训练时可能悄悄用错偏移,推理时才发现对不上,导致部署性能大幅下降。下面逐条拆开看它们各自对应哪个工程决策:
1. _has_action_offset 检查训练样本是否携带 action_offset 字段,这是 dataloader 在每个 batch 随机生成的偏移值。
2. _normalize_action_offsets 把偏移规范化到 [0, max_action_offset] 范围,并校验形状和正负。注意它支持 scalar 或 [B] 形状——前者意味着整个 batch 用同一个偏移,后者意味着每个样本独立采样偏移。论文用的是 per-sample 采样,让训练更鲁棒。
3. _validate_runtime_action_horizon 在推理时校验运行时动作视野必须 ≥ max_action_offset + action_chunk_size,否则即使训练支持了大偏移,推理时也用不上。
4. training_loss 根据样本是否有 action_offset 字段分发到不同的损失函数——这是为了向后兼容旧的非 offset 训练数据。
6. Rolling KV Memory:规划器的长期记忆
6.1 为什么规划器需要"记住"过去
Video DiT 作为低频规划器,应该能在跨刷新边界保留历史场景信息,而不是只依赖当前观测。这对长视野操作任务尤其重要——已经完成的子目标、被移动过的物体、之前观察到但当前被遮挡的状态都可能提供关键上下文。一个简单的例子:机器人要"把杯子放到柜子里",第一步是打开柜门,第二步是把杯子放进去。第二步执行时柜子内部空间可能已经被打开但当前帧只看到杯子,规划器如果记不住"柜门已经打开",就可能误判当前状态。
6.2 FIFO 滚动 K/V 设计
AHA-WAM 在 Video DiT 内部维护一个固定大小的 FIFO 滚动 K/V 记忆,像一个先进先出的队列不断吞入最新状态、吐出最旧状态。这块记忆是规划器自己的笔记本,不直接暴露给执行器。对每一层 ,这块记忆存储最近几次规划器刷新时产生的历史视频状态键值对,用公式表达队列的更新逻辑就是:
其中 索引规划器刷新次数,记忆窗口大小固定(论文设 6 帧历史)。下次刷新时,Video DiT 在生成新规划器视频上下文 时会对这块记忆做注意力,把过去几次刷新看到的场景信息融进当前规划。这种内部记忆机制让规划器在长视野任务中保持稳定,不会因为单帧画面缺失就丢失对全局状态的把握,也避免了把所有历史细节都暴露给执行器导致的接口臃肿。
这块记忆是 Video DiT 内部的,不直接被 Action DiT 消费。它扩展规划器的时间感受野,让规划器在产生新上下文之前能看到更远的历史,而高频执行器仍然只与最新的、被 OVCR 适配过的规划器视频上下文交互。这种"内部长记忆 + 外部短界面"的设计很优雅——记忆细节不需要暴露给执行器,只需要规划器把记忆消化后形成的上下文交出去就够了。
一句话理解(滚动记忆 vs 规划器上下文):滚动 KV 记忆是"规划器自己看的笔记本",规划器视频上下文是"规划器交给执行器的便条"。笔记本越积越厚,便条每次都要重写,但便条里包含了笔记本里最关键的信息。
7. 推理加速:从 190ms 到 17.56ms 的 10.82× 提速
7.1 三层加速:异步调度 + CUDA 优化 + ODE 蒸馏
AHA-WAM 的速度优势来自三层叠加。第一层是异步调度本身——Video DiT 的 prefill 异步在后台跑,不进入每次动作更新的关键路径,所以单次动作块延迟 直接决定控制频率。第二层是 CUDA 优化——把重复的动作相位计算编译成静态部署路径,通过 TensorRT 和 CUDA-graph 捕获执行 Action DiT、记忆/上下文模块、VAE 编码器,同时移除去噪热路径里的冗余计算和重复缓冲区拷贝。这些改动不改模型架构、权重或采样过程,但把 10 步动作推理延迟从 PyTorch eager 的 415.77ms 压到 41.37ms。第三层是 ODE 蒸馏——在 CUDA 加速的 10 步路径基础上,把动作采样器从 10 步蒸馏到 2 步,构造 AHA-WAM-Flash 变体,进一步把 从 41.37ms 降到 17.56ms。
| 方法 | 延迟 (ms) | 频率 (Hz) | 加速比 |
|---|---|---|---|
| Motus | 1866.10 | 0.54 | 0.10× |
| Fast-WAM | 190.00 | 5.26 | 1.00× |
| AHA-WAM | 41.37 | 24.17 | 4.59× |
| AHA-WAM-Flash | 17.56 | 56.95 | 10.82× |
相比 Fast-WAM,AHA-WAM 把延迟从 190.00ms 降到 41.37ms,闭环频率从 5.26Hz 提升到 24.17Hz,这主要来自异步调度和 CUDA 优化两层。配上 ODE 蒸馏采样器后,AHA-WAM-Flash 进一步达到 17.56ms 和 56.95Hz,实现了相对 Fast-WAM 的 10.82× 加速,同时保持完全相同的异步规划器-执行器界面——加速优化没有动任何核心架构,这正是接口稳定性带来的工程红利。
7.2 ODE 蒸馏的细节:冻结视频分支,只蒸馏动作路径
ODE 蒸馏的关键设计是只蒸馏动作去噪路径,冻结 Video DiT,让学生保持和 10 步教师相同的规划器上下文和 OVCR 界面。教师产生 16 步去噪轨迹,从噪声态 0 到最终去噪态 16,选取轨迹锚点 。训练时学生随机采样一个非终态锚点作为起始态,直接预测教师的最终去噪动作输出。采样更偏向轨迹的噪声端——因为从高噪声态准确预测是减少推理去噪步数的关键要求。
7.3 代码实现:ODE 蒸馏学生类
为了让读者看清 ODE 蒸馏怎么在不破坏 OVCR 界面的前提下压缩去噪步数,下面从 src/ahawam/models/wan22/distill/ahawam_ode.py 抽取蒸馏入口类。这个文件在 distill 子目录下,说明蒸馏是独立模块,不影响主模型训练。它继承自 Fast-WAM 的蒸馏框架但冻结了视频分支,让蒸馏只动动作路径:
# src/ahawam/models/wan22/distill/ahawam_ode.py
class AHAWAMODE(FastWAMODE):
"""
ODE distillation for AHA-WAM action sampler.
Distills the 10-step (or 16-step) action denoising path to 2 steps,
while keeping the Video DiT frozen and preserving the OVCR interface.
"""
def __init__(
self,
teacher_denoising_steps: int = 16, # teacher trajectory length
student_denoising_steps: int = 2, # distilled student steps
trajectory_anchors: tuple = (0, 1, 2, 4, 8, 12, 16), # capture schedule
noise_biased_sampling: bool = True, # sample more from noisy end
freeze_video_dit: bool = True, # video branch stays frozen
**kwargs,
):
super().__init__(**kwargs)
self.teacher_denoising_steps = teacher_denoising_steps
self.student_denoising_steps = student_denoising_steps
self.trajectory_anchors = trajectory_anchors
self.noise_biased_sampling = noise_biased_sampling
self.freeze_video_dit = freeze_video_dit
if freeze_video_dit:
# Freeze the slow world planner; only distill the fast action path
for param in self.video_dit.parameters():
param.requires_grad = False
def distillation_loss(self, sample, teacher_trajectory):
"""
Student predicts the teacher's FINAL denoised action from a sampled
intermediate state. Sampling is biased toward noisier states because
accurate prediction from high-noise states enables aggressive step reduction.
"""
# 1. Sample a non-final anchor as student's starting state
anchor_idx = self._sample_anchor_biased(self.trajectory_anchors)
start_state = teacher_trajectory[anchor_idx] # noisy intermediate
target_state = teacher_trajectory[-1] # teacher's final output
# 2. Student predicts final action directly from the noisy start state
# OVCR interface and planner context are identical to the teacher
student_pred = self.student_action_dit(
start_state,
planner_context=sample["routed_planner_context"], # OVCR-adapted K/V
proprio=sample["proprio_state"],
text_context=sample["text_context"],
)
# 3. Regression loss against teacher's final denoised action
return F.mse_loss(student_pred, target_state)
def _sample_anchor_biased(self, anchors):
"""Bias sampling toward the noisy end of the trajectory."""
# Higher anchors (closer to noise) get sampled more frequently
weights = torch.linspace(1.0, 3.0, len(anchors)) # noisy end weighted 3x
weights = weights / weights.sum()
idx = torch.multinomial(weights, 1).item()
return anchors[idx]
这段蒸馏代码的关键工程决策有三个。第一,冻结 Video DiT——蒸馏只动动作路径,世界规划器保持不变,这样学生和教师共享完全相同的规划器上下文和 OVCR 界面,蒸馏不会破坏已学到的物理先验。第二,学生直接预测最终态——不是逐步预测下一个去噪态,而是从某个中间噪声态一步跳到最终动作,这是把 16 步压到 2 步的核心。第三,噪声偏置采样(_sample_anchor_biased)——更频繁地从高噪声态采样训练,因为推理时学生要从最噪声的态开始去噪,这个能力最关键。论文用线性权重让噪声端被采样的概率是干净端的 3 倍。
8. 总结
AHA-WAM 的核心贡献不是再造一个世界模型缩写,而是把世界预测和动作执行从同一时序节拍上解绑,让慢规划器和快执行器各自运行在最有信息量的时间尺度上。50 个 RoboTwin 任务上 92.80% 的平均成功率(无机器人数据预训练)、四项真实双臂任务 78.33% 的成功率、从 190ms 到 17.56ms 的 10.82× 推理加速、Naive-Async 跌 3.23 个百分点而完整版反超的消融对比,这四组数字合起来足以让这套异步架构进入下一代世界-动作模型的候选清单。
147