Files
worldmodel/research/multiply/V-JEPA2源码逐模块剖析.md

27 KiB
Raw Permalink Blame History

title, date, draft, tags, categories, description
title date draft tags categories description
V-JEPA 2 源码逐模块剖析(代码篇) 2026-05-24 false
V-JEPA 2
JEPA
源码
code-walkthrough
self-supervised
RoPE
masking
world-model
robotics
JEPA
基于 facebookresearch/vjepa2 本地源码的逐模块剖析:multi-block tube 掩码采样、编码器(3D-RoPE/apply_masks)、动作无关预测器、V-JEPA 2-AC 块因果预测器,以及自监督训练全链路闭环。是《V-JEPA 2 深度技术剖析》的代码篇配套。

本文是 《V-JEPA 2 深度技术剖析》代码篇配套,全部基于本地 clone 的官方仓库 facebookresearch/vjepa2MIT License)逐行阅读整理。行号与文件路径以仓库 main 分支为准,升级版本可能略有变动。

阅读主线(数据流顺序):掩码采样 → 编码器 → 动作无关预测器 → AC 预测器 → 训练闭环


〇、文件地图与真实配置速查

模块 文件 关键类/函数
掩码采样 src/masks/multiseq_multiblock3d.py MaskCollator / _MaskGenerator
取 token 工具 src/masks/utils.py apply_masks
编码器 src/models/vision_transformer.py VisionTransformer / vit_giant_xformers_rope
tubelet 切分 src/models/utils/patch_embed.py PatchEmbed3D
注意力 / RoPE src/models/utils/modules.py RoPEAttention / ACRoPEAttention / rotate_queries_or_keys / build_action_block_causal_attention_mask
动作无关预测器 src/models/predictor.py VisionTransformerPredictor
AC 预测器 src/models/ac_predictor.py VisionTransformerPredictorAC
预训练循环 app/vjepa/train.py train_step / forward_context / forward_target / loss_fn
AC 训练循环 app/vjepa_droid/train.py teacher-forcing + rollout
AC 推理/规划桥接 notebooks/utils/world_model_wrapper.py WorldModelencode / infer_next_action
CEM + 能量函数 notebooks/utils/mpc_utils.py cem / l1 / compute_new_pose

编码器 ViT-gvit_giant_xformers_rope):embed_dim=1408, depth=40, num_heads=22, mlp_ratio=48/11, use_rope=True

动作无关预测器pretrain-256px-16f.yaml):pred_embed_dim=384, pred_depth=12, pred_num_heads=12, use_mask_tokens=True(≈22M,刻意做小)。

AC 预测器configs/train/vitg16/droid-256px-8f.yaml):pred_embed_dim=1024, pred_depth=24, pred_num_heads=16, pred_is_frame_causal=True, use_extrinsics=False, use_rope=True(≈300M);context_encoder_key=target_encoder(直接拿 EMA target encoder 当冻结编码器)。

掩码pretrain-256px-16f.yaml,两套并行):

mask:
- num_blocks: 8        # 短程:8 个小块
  spatial_scale: [0.15, 0.15]   # 每块覆盖 15% 空间
  temporal_scale: [1.0, 1.0]    # 时间贯穿整段(tube
  aspect_ratio: [0.75, 1.5]
- num_blocks: 2        # 长程:2 个大块
  spatial_scale: [0.7, 0.7]
  temporal_scale: [1.0, 1.0]

一、掩码采样:multiseq_multiblock3d.py

整条 JEPA 管线的起点。masks_enc(上下文下标)/ masks_pred(目标下标)都在这里生产。

1.1 两层结构:调度 vs 采样

MaskCollator 是 DataLoader 的 collate_fn,本身不采样,只负责调度:按 frames-per-clip(fpc) 分桶(不同帧数 token 网格大小不同,不能混),再对每套掩码配置各调一次生成器:

for i, mask_generator in enumerate(self.mask_generators[fpc]):
    masks_enc, masks_pred = mask_generator(batch_size)
    collated_masks_enc.append(masks_enc)
    collated_masks_pred.append(masks_pred)

所以 masks_enc / masks_predlist(每套配置一项)——这解释了编码器里 for m in masks:torch.cat(masks, dim=0) 为什么要遍历。

1.2 种子心机:块"大小"同步,块"位置"随机

seed = self.step()                       # 跨 worker 共享计数器
g = torch.Generator(); g.manual_seed(seed)
p_size = self._sample_block_size(generator=g, ...)   # 用种子 → 块大小确定

尺寸 (t,h,w) 用种子采样,保证所有数据并行 worker/GPU 采到相同大小(token 序列等长,可凑规整 batch);块位置用不带种子的 torch.randint,每个样本各自随机。同步大小 + 随机位置 = 既能 batch 又有多样性。

1.3 一个块怎么采

定大小(scale + 长宽比反解):

t = max(1, int(self.duration * temporal_mask_scale))   # temporal_scale=1.0 → t=duration(贯穿全部帧)
spatial_num_keep = int(self.height * self.width * spatial_mask_scale)
h = round(sqrt(spatial_num_keep * aspect_ratio))
w = round(sqrt(spatial_num_keep / aspect_ratio))

temporal_scale=1.0t=duration,块在时间上贯穿整段 = 一根"时空管(tube)"。这是要点:把同一空间位置在所有帧上一起遮掉,模型无法靠抄相邻帧作弊,只能真正学运动

定位置并标 00=待预测,1=上下文):

mask = torch.ones((duration, height, width))
mask[start:start+t, top:top+h, left:left+w] = 0

1.4 多块取并集,切出两组下标

mask_e = torch.ones((duration, height, width))
for _ in range(self.npred):                 # npred = num_blocks,如 8
    mask_e *= self._sample_block_mask(p_size)   # 连乘 = 并集
mask_e = mask_e.flatten()

mask_p = torch.argwhere(mask_e == 0).squeeze()   # 目标(被遮)下标
mask_e = torch.nonzero(mask_e).squeeze()         # 上下文(保留)下标

mask_e *= block 连乘是精髓:patch 只有落在所有块之外才保持 1,落进任意块就被乘成 0。npred 个块取并集,两套配置(8×15% + 2×70%)叠加后被遮区逼近 ~90%。argwhere(==0) 出目标下标、nonzero 出上下文下标,二者是同一 token 网格上互补的索引集合。

empty_context while 守卫:若某次把所有 patch 都遮了就重采,保证编码器有输入。

1.5 截断到 min_keep(为了 stack

每样本被遮数不同 → 下标长度不齐 → 没法堆张量。于是全 batch 取最短并截断:

collated_masks_enc  = [cm[:min_keep_enc]  for cm in collated_masks_enc]
collated_masks_pred = [cm[:min_keep_pred] for cm in collated_masks_pred]
collated_masks_enc  = torch.utils.data.default_collate(collated_masks_enc)

这就是 apply_masks 拿到的 masks 为规整 [B, K] 张量的原因。

1.6 几个开关

  • full_complement:预测掩码 = 编码掩码的精确补集(集合差)。
  • pred_full_complement:反向,编码掩码 = 预测掩码补集。
  • inv_block:交换返回 return masks_pred, masks_enc(用块预测上下文的反向任务)。默认 return masks_enc, masks_pred

二、取 token 工具:apply_maskssrc/masks/utils.py

def apply_masks(x, masks, concat=True):
    # x: [B, N, D]; masks: list of [B, K] 下标
    all_x = []
    for m in masks:
        mask_keep = m.unsqueeze(-1).repeat(1, 1, x.size(-1))  # [B, K, D]
        all_x += [torch.gather(x, dim=1, index=mask_keep)]    # 按下标取 token
    return torch.cat(all_x, dim=0)

核心是 torch.gather:沿 token 维按下标取出需要的 token。多块掩码沿 batch 维(dim=0)堆叠,相当于把一个样本的多种掩码当成更大 batch 一起算。


三、编码器:VisionTransformer.forwardvision_transformer.py:161

关键认知:代码里没有"裁一块图"。整段视频先切成全部 token,再用下标 gather 出上下文 token。被遮 token 不是置零,而是直接删除、不参与计算

def forward(self, x, masks=None):
    # (a) 切 token3D 卷积;use_rope 时不加绝对位置
    if not self.use_rope:
        pos_embed = self.interpolate_pos_encoding(x, self.pos_embed)
        x = self.patch_embed(x); x += pos_embed
    else:
        x = self.patch_embed(x)              # ViT-g 走这条

    # (b) 挑出上下文 token
    if masks is not None:
        x = apply_masks(x, masks)            # 丢掉被遮 token
        masks = torch.cat(masks, dim=0)      # 下标继续往下传(给 RoPE 还原位置)

    # (c) Transformer 编码
    for blk in self.blocks:
        x = blk(x, mask=masks, attn_mask=None, T=T, H_patches=H_patches, W_patches=W_patches)
    x = self.norm(x)
    return x

apply_masks 在 tokenize 之后、进 Transformer 之前 → ~90% token 在进注意力前就丢了,这是高掩码率下省算力的根因。

3.1 tubelet 切分:PatchEmbed3D

self.proj = nn.Conv3d(
    in_channels=3, out_channels=embed_dim,
    kernel_size=(tubelet_size, patch_size, patch_size),  # (2, 16, 16)
    stride=(tubelet_size, patch_size, patch_size),       # 不重叠
)
def forward(self, x):                # x: [B, C, T, H, W]
    return self.proj(x).flatten(2).transpose(1, 2)   # -> [B, N, embed_dim]

一个 2×16×16 时空管元 → 一个向量。

3.2 删了 token,位置靠"下标"续命:RoPEAttentionmodules.py:266

token 被删一大半、顺序也乱,模型怎么知道每个上下文块原来的位置?答案在传下来的下标:

if mask is not None:
    mask = mask.unsqueeze(1).repeat(1, self.num_heads, 1)
    d_mask, h_mask, w_mask = self.separate_positions(mask, H_patches, W_patches)

separate_positions 把一维下标反解成 (帧, 高, 宽) 三个坐标:

frame_ids  = ids // tokens_per_frame
height_ids = (ids - 帧分量) // tokens_per_row
width_ids  =  ids - 帧分量 - 高分量

再把每个 head 的特征维三等分,按 t/h/w 各自旋转:

qd = rotate_queries_or_keys(q[..., s:s+self.d_dim], pos=d_mask)   # 时间
qh = rotate_queries_or_keys(q[..., s:s+self.h_dim], pos=h_mask)   # 高
qw = rotate_queries_or_keys(q[..., s:s+self.w_dim], pos=w_mask)   # 宽
q = torch.cat([qd, qh, qw, qr], dim=-1)   # qr 是不旋转的余数维

要点:位置不靠 token 在序列里的顺序,而靠它的原始下标算出来——所以哪怕 token 删得七零八落,每个块仍带着真实 (t,h,w)。

3.3 RoPE 旋转真身 + 一个官方标注的 bug:rotate_queries_or_keysmodules.py:26

def rotate_queries_or_keys(x, pos):
    B, num_heads, N, D = x.size()
    omega = torch.arange(D // 2) / (D / 2.0)
    omega = 1.0 / 10000**omega                       # 角速度 ω_i
    freq = torch.einsum("..., f -> ... f", pos, omega)   # 角度 = 位置 × ω
    emb_sin, emb_cos = freq.sin(), freq.cos()

    # ↓↓↓ 官方注释:这里有个 subtle bug,频率在配对维上被错误复制
    emb_sin = emb_sin.squeeze(-1).repeat(1, 1, 1, 2)     # 实际用的(有 bug
    emb_cos = emb_cos.squeeze(-1).repeat(1, 1, 1, 2)
    # emb_sin = emb_sin.repeat_interleave(2, dim=-1)      # 正确写法(被注释掉)
    # emb_cos = emb_cos.repeat_interleave(2, dim=-1)

    y = x.unflatten(-1, (-1, 2))
    y1, y2 = y.unbind(dim=-1)
    y = torch.stack((-y2, y1), dim=-1).flatten(-2)        # (y1,y2)->(-y2,y1)
    return (x * emb_cos) + (y * emb_sin)

RoPE 思想:不是"加"位置向量,而是按位置把特征在每个二维平面里旋转一个角度;两 token 点积时旋转角之差只取决于相对位置 → 天然编码相对距离、可外推到没见过的分辨率/长度(这正是论文用 3D-RoPE 稳住大模型 + 支持渐进分辨率的底层依据)。

那个 bug 的工程教训:标准 RoPE 频率应与相邻配对 (y1,y2)repeat_interleave(2)[a,b,c]→[a,a,b,b,c,c])对齐;代码用了 repeat(...,2)[a,b,c]→[a,b,c,a,b,c]),频率与配对维错位。官方明说:发布的预训练权重就是在这个"错"的实现上训出来的,改对反而和 checkpoint 不兼容,故保留原样、正确版注释旁边。加载官方权重做下游必须沿用 bug 版;只有从头自训才该切正确版。


四、动作无关预测器:VisionTransformerPredictor.forwardpredictor.py:174

任务:拿编码器给的上下文表示 z_ctx,在目标位置把表示补出来。

def forward(self, x, masks_x, masks_y, ...):
    # x: 上下文表示(z_ctx)masks_x: 上下文下标;masks_y: 目标下标

4.1 先降维("预测器很轻"的代码证据)

x = self.predictor_embed(x)   # Linear: 1408 -> 384

进门就压到 384 维,整个预测器在窄空间算,算完再投影回去。

4.2 给目标位置造可学习占位符

pred_tokens = self.mask_tokens[mask_index]          # 一个可学习 [MASK] 向量
pred_tokens = pred_tokens.repeat(B, self.num_patches, 1)
pred_tokens = apply_masks(pred_tokens, masks_y)     # 只在目标位置放占位符

目标位置不喂内容(内容正是要预测的),只放一个所有目标位置共享的可学习 mask token,信息全靠位置编码注入。

4.3 拼接 + 排序

x = torch.cat([x, pred_tokens], dim=1)        # [上下文表示 | 目标占位符]
masks = torch.cat([masks_x, masks_y], dim=1)  # 每个 token 的原始下标
argsort = torch.argsort(masks, dim=1)         # 按原始位置排序
x = ...[argsort]; masks = ...[argsort]

拼起来的序列空间上是乱的,按原始下标排序恢复成 (t,h,w) 顺序,再把 masks 传进 blk(x, mask=masks) 让 RoPE 还原坐标。

4.4 过 Transformer,抠回目标、升回维

for blk in self.predictor_blocks:
    x = blk(x, mask=masks, attn_mask=None)
x = self.predictor_norm(x)
reverse_argsort = torch.argsort(argsort, dim=1)
x = ...[reverse_argsort]; x = x[:, N_ctxt:]   # 复原排序,只取目标位置
x = self.predictor_proj(x)                     # Linear: 384 -> 1408
return x                                        # ẑ_tgt

整条链路 1408 → 压到 384 算 → 升回 1408 就是"轻量预测器"的全部秘密。


五、AC 预测器:VisionTransformerPredictorACac_predictor.py

把"被遮 patch"换成"未来帧""mask token"换成"动作/状态 token + 块因果掩码",复用同一套 latent 预测骨架做规划。

5.1 四个编码器:把"看的"和"做的"统一进 1024 维

self.predictor_embed    = nn.Linear(embed_dim, 1024)          # 1408 -> 1024(图像 patch
self.action_encoder     = nn.Linear(action_embed_dim, 1024)   # 7 -> 1024(动作)
self.state_encoder      = nn.Linear(action_embed_dim, 1024)   # 7 -> 1024(末端位姿/本体感知)
self.extrinsics_encoder = nn.Linear(action_embed_dim-1, 1024) # 6 -> 1024(相机外参,droid 关闭)

action_embed_dim=7:Franka 动作/状态正好 7 维(末端 6-DoF 位姿增量 + 1 维夹爪)。动作/位姿各过一个线性层变成与 patch 同维的 token——这是"注入"的第一层含义。

5.2 注入方式:把动作/状态插进每帧开头

s = self.state_encoder(states).unsqueeze(2)    # [B, T, 1, D]
a = self.action_encoder(actions).unsqueeze(2)  # [B, T, 1, D]
x = x.view(B, T, H*W, D)
x = torch.cat([a, s, x], dim=2).flatten(1, 2)  # [B, T*(H*W+2), D]

序列按帧分块:[ a₁ s₁ p₁... | a₂ s₂ p₂... | ... ]cond_tokens=2(开 extrinsics 则 3)。这对应论文"每个 patch 能注意到同一时间步的动作、末端状态和其他 patch"。

5.3 块因果掩码:build_action_block_causal_attention_mask

def build_action_block_causal_attention_mask(T, H, W, add_tokens=1):
    N_T = add_tokens + (H * W)        # 每帧 block 的 token 数
    mask = torch.zeros(T*N_T, T*N_T).bool()
    mask_block = torch.ones(N_T, N_T).bool()   # 帧内:全连接
    local_window_time = T                       # 看全部历史
    for t1 in range(T):
        for t2 in range(max(0, t1 - local_window_time + 1), t1 + 1):  # t2 <= t1
            mask[t1*N_T:(t1+1)*N_T, t2*N_T:(t2+1)*N_T] = mask_block
    return mask
  • mask_block=ones帧内部全连接(同一时刻动作/状态/所有 patch 互看)。
  • t2 只到 t1只能 attend 过去帧,禁看未来
  • 即"因果":帧粒度因果、块内全连接(不希望同帧后一个 patch 看不到前一个,只希望整帧看不到下一帧)。
  • 布尔约定:直接喂 F.scaled_dot_product_attention(attn_mask=mask)True=允许 attend

is_causal 怎么生效:真正的因果性靠这个显式 attn_mask 实现,不是 SDPA 的 is_causal flag(那是逐 token 下三角,对块结构是错的)。config 的 pred_is_frame_causal: true 控制的是"造不造并用这个块因果矩阵"(构造函数里 if self.is_frame_causal: attn_mask = build_...),而非把 SDPA 的 is_causal 设 True。

5.4 动作 token 与 patch 的 RoPE 不同(ACRoPEAttention

# 动作/状态 token:只在时间维(depth)旋转,位置=帧号
qd = rotate_queries_or_keys(q[..., :self.d_dim], pos=torch.arange(T))
qr = q[..., self.d_dim:]          # 其余维不旋转

patch token 走完整 3D-RoPE。对应论文"动作和位姿用时间维 rotary,视频 patch 用 3D-RoPE"——动作没有空间行列,只有"第几帧"。另有分辨率归一化:h_mask *= grid_size/H; w_mask *= grid_size/W(换分辨率时位置编码一致)。

5.5 出口:只要 patch 预测

x = x.view(B, T, cond_tokens + H*W, D)
x = x[:, :, cond_tokens:, :].flatten(1, 2)   # 丢掉动作/状态 token
x = self.predictor_norm(x)
x = self.predictor_proj(x)                    # 1024 -> 1408
return x

动作/状态只是条件输入、不需预测;输出只取 patch、投影回 1408,得到"给定历史帧 + 当前动作,下一帧每个 patch 的预测表示"。

注:ac_predictor.pyfrom src.models.utils.modules import ACBlock as Block,块内部用的是 ACBlock(其 forward 接收 action_tokens 参数)。


六、预训练训练循环:app/vjepa/train.py:424

def forward_target(c):
    with torch.no_grad():                 # stop-grad
        h = target_encoder(c)             # 喂【完整视频】,不传 masks
        h = [F.layer_norm(hi, (hi.size(-1),)) for hi in h]
        return h

def forward_context(c):
    z = encoder(c, masks_enc)             # 只看上下文 → z_ctx
    z = predictor(z, masks_enc, masks_pred)   # 补目标位置 → ẑ_tgt
    return z

def loss_fn(z, h):
    h = [apply_masks(hi, mi, concat=False) for hi, mi in zip(h, masks_pred)]  # 此处才取目标块
    loss = torch.mean(torch.abs(zij - hij) ** loss_exp) / loss_exp            # loss_exp=1 → L1

两条编码路径的对比是关键:

  • 上下文分支encoder(c, masks_enc) —— 只看上下文、可训练。
  • 目标分支target_encoder(c) —— 看完整视频、是 encoder 的 EMA 副本no_grad;直到 loss_fn 才用 masks_pred 抠出目标块当真值。

损失 |z h|loss_exp=1 即 L1),只在被遮 token 上算——与论文"表示空间 L1、仅在 masked patch"一致。EMA + stop-grad 即防坍塌机制。


七、AC 推理:CEM + 能量函数规划(notebooks/utils/

定位提醒:app/vjepa_droid/训练侧;真正"跑起来"的规划在 notebooks/utils/mpc_utils.pyCEM)与 world_model_wrapper.py(桥接)。二者与训练侧的 rollout 是同一机制的"推理版/训练版"。

7.1 桥接层 WorldModelworld_model_wrapper.py

把"编码器 + AC 预测器"包成可被规划器调用的世界模型。默认 MPC 超参:

mpc_args = {"rollout": 2, "samples": 400, "topk": 10, "cem_steps": 10,
            "momentum_mean": 0.15, "momentum_std": 0.15, "maxnorm": 0.05}

encode(image):单帧得先复制成 2 帧(tubelet 时间核=2)再编码,且必须 layer_norm(与训练目标同一归一化空间,能量尺度才对齐):

clip = clip.permute(...).flatten(0,1).unsqueeze(2).repeat(1, 1, 2, 1, 1)  # 单帧 → 2 帧
h = self.encoder(clip)
if self.normalize_reps: h = F.layer_norm(h, (h.size(-1),))

infer_next_action:定义"走一步"的闭包交给 cem。要点——视觉表示靠 AC 预测器"想象",本体位姿靠运动学硬算

def step_predictor(reps, actions, poses):
    next_rep = self.predictor(reps, actions, poses)[:, -self.tokens_per_frame:]  # 取最后一帧预测
    if self.normalize_reps: next_rep = F.layer_norm(next_rep, (next_rep.size(-1),))
    next_pose = compute_new_pose(poses[:, -1:], actions[:, -1:])                 # 运动学积分,非网络预测
    return next_rep, next_pose
mpc_action = cem(context_frame=rep, context_pose=pose, goal_frame=goal_rep,
                 world_model=step_predictor, **self.mpc_args)[0]

7.2 能量函数就是 latent L1mpc_utils.py

def l1(a, b):
    return torch.mean(torch.abs(a - b), dim=-1)   # 想象状态 vs 目标状态,逐元素 L1

论文的 goal-conditioned energy function,落到代码就这一行。

7.3 只优化"平移 + 夹爪",旋转冻结为 0(读源码才见的简化)

动作名义 7 维,分布只建在 4 维(xyz + gripper)上;采样时旋转 3 维恒填 0:

std = cat([ones((rollout,3)) * maxnorm, ones((rollout,1))], dim=-1)   # xyz 的 std=maxnorm
action_samples = randn(samples, 4) * std[h] + mean[h]
action_samples[:, :3]  = clip(action_samples[:, :3], -maxnorm, maxnorm)  # ← L1 球动作约束
action_samples = cat([action_samples[:, :3], zeros((len,3)), action_samples[:, -1:]], -1)  # 旋转恒 0

clip(±maxnorm) 即论文"每步动作约束在 L1 球内"(默认 maxnorm=0.05;论文桌面实验报半径 0.075≈13cm,数量级一致,实验脚本可覆盖默认值),理由是大动作 out-of-distribution。

7.4 想象 rollout + 选精英 + 动量更新

for h in range(rollout):                       # 400 条候选同时在 latent 里播放 rollout 步
    next_frame, next_pose = world_model(frame_traj, action_traj, pose_traj)
    ...
sims = l1(final_state.flatten(1), goal_state.flatten(1))   # 只比【最终想象帧】vs 目标
idx  = sims.topk(topk, largest=False).indices             # 取能量最小的 topk 精英
mean = mean_selected * (1 - momentum_mean) + mean * momentum_mean   # 带动量更新分布
std  = std_selected  * (1 - momentum_std)  + std  * momentum_std

全程不解码回像素——预测/比较/打分都在表示空间,这是比扩散式 Cosmos 快一个数量级(16s vs 4min)的根因。迭代 cem_steps 轮后返回分布均值作动作。标准 CEM:采样 → 评估 → 留精英 → 更新分布,纯采样无梯度,单卡可跑。

7.5 受控时域(receding horizon

cem 返回 [1, rollout, 7] 整条轨迹,但外层执行循环只落地第一个动作,再重新编码当前观测、重新规划rollout=2 = 每次做 2 步前瞻、只执行 1 步,用前瞻避免短视。

7.6 位姿用确定性运动学:compute_new_pose

new_xyz        = pose[:, :3] + action[:, :3]                     # 平移相加
angle_diff     = [dm @ m for dm, m in zip(delta_matrices, matrices)]  # 旋转矩阵相乘复合
new_closedness = np.clip(pose[:, -1:] + action[:, -1:], 0, 1)    # 夹爪裁剪

位姿用真实运动学递推(平移加、旋转复合、夹爪裁剪),不让网络在低维本体状态上犯错。poses_to_diff 是逆运算(从当前位姿→目标位姿反解动作)。

7.7 训练侧为何配套:teacher-forcing + rollout 双损失(app/vjepa_droid/train.py:417

推理要多步 rollout,训练就必须同时练单步和多步,否则误差累积漂移:

z_tf = _step_predictor(z[:, :-tokens_per_frame], actions, states[:,:-1], extrinsics[:,:-1])  # teacher forcing
_z = cat([z[:, :tokens_per_frame], z_tf[:, :tokens_per_frame]], dim=1)
for n in range(1, auto_steps):                                   # 用自己的预测往后滚
    _z_nxt = _step_predictor(_z, actions[:, :n+1], states[:, :n+1], extrinsics[:, :n+1])[:, -tokens_per_frame:]
    _z = cat([_z, _z_nxt], dim=1)
z_ar = _z[:, tokens_per_frame:]
loss = loss_fn(z_tf, h) + loss_fn(z_ar, h)                       # 单步 L1 + 多步 rollout L1

slossrollout 项)让预测器在"喂自己预测"时也不漂,推理时 CEM 的 latent rollout 才可信——训练 rollout 与推理 CEM rollout 是配套设计。


八、全链路闭环

MaskCollator(_MaskGenerator):多块连乘取并集
   → (masks_enc 上下文下标, masks_pred 目标下标)                         §一
        │
        ▼
encoder(clips, masks_enc)            → 只编码上下文 → z_ctx              §三
predictor(z_ctx, masks_enc, masks_pred) → 在目标位置补出 ẑ_tgt          §四
target_encoder(clips) + apply_masks(masks_pred) → 目标真值 hEMA, stop-grad
loss = |ẑ_tgt  h|(仅 masks_pred 上, L1                              §六

AC 阶段(§五):把"被遮 patch"换成"未来帧""mask token"换成"动作/状态 token + 块因果掩码",复用同一 latent 预测骨架;推理时反过来用——给目标图像编码成 z_goal,枚举动作让模型在 latent 里 rollout,以 ‖ẑ_future z_goal‖₁ 为能量,CEM 搜最优动作、执行第一步再重规划(receding horizon)。


九、两条值得记住的设计要点

  1. multi-block tube 掩码:① 时间贯穿的 tubetemporal_scale=1.0)逼模型学运动而非抄帧;② 块大小同步、位置随机 + 多块并集,在可批处理前提下做出 ~90% 高难度掩码。
  2. 下标即位置:掩码通过"删 token + 传下标"实现,位置信息全程由下标 → 3D-RoPE 还原,而非依赖序列顺序。这套机制同时服务于编码器、动作无关预测器和 AC 预测器。

十、工程备忘

  • RoPE bug 不可贸然修rotate_queries_or_keys 的频率展开有已知 bug,但权重与之绑定。加载官方 checkpoint 必须沿用 bug 版;从头自训可切注释里的 repeat_interleave(2) 正确版。
  • 块因果 ≠ is_causalAC 的时间因果由显式 attn_mask 实现,pred_is_frame_causal 控制是否构造该矩阵;勿与 SDPA 的逐 token is_causal 混淆。
  • min_keep 截断有信息损失:为可批处理,全 batch 下标截到最短,会丢掉少量 patch;高 batch 多样性下影响小。
  • AC 复用 EMA 编码器context_encoder_key=target_encoder,AC 阶段冻结的是预训练的 EMA target encoder,不是在线 encoder。
  • 规划默认不搜旋转cem 默认只优化平移 + 夹爪,旋转 3 维恒为 0(axis 可手动钉死其他维);大幅缩小搜索空间。
  • 能量尺度靠 layer_norm 对齐:编码、预测、目标三处都 layer_norm,规划能量(latent L1)才在同一尺度;推理侧 normalize_reps 必须与训练一致。
  • 训练 rollout ↔ 推理 CEM rollout 配套:缺了训练侧的 slossautoregressive rollout 损失),推理时多步 latent rollout 会漂移、CEM 失准。

配套阅读:方法/数字/机器人结果见 《V-JEPA 2 深度技术剖析》;系列脉络见 《Meta JEPA 系列全面调研》