27 KiB
title, date, draft, tags, categories, description
| title | date | draft | tags | categories | description | ||||||||||
|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|---|
| V-JEPA 2 源码逐模块剖析(代码篇) | 2026-05-24 | false |
|
|
基于 facebookresearch/vjepa2 本地源码的逐模块剖析:multi-block tube 掩码采样、编码器(3D-RoPE/apply_masks)、动作无关预测器、V-JEPA 2-AC 块因果预测器,以及自监督训练全链路闭环。是《V-JEPA 2 深度技术剖析》的代码篇配套。 |
本文是 《V-JEPA 2 深度技术剖析》 的代码篇配套,全部基于本地 clone 的官方仓库
facebookresearch/vjepa2(MIT 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 |
WorldModel(encode / infer_next_action) |
| CEM + 能量函数 | notebooks/utils/mpc_utils.py |
cem / l1 / compute_new_pose |
编码器 ViT-g(vit_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_pred 是 list(每套配置一项)——这解释了编码器里 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.0 ⇒ t=duration,块在时间上贯穿整段 = 一根"时空管(tube)"。这是要点:把同一空间位置在所有帧上一起遮掉,模型无法靠抄相邻帧作弊,只能真正学运动。
定位置并标 0(0=待预测,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_masks(src/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.forward(vision_transformer.py:161)
关键认知:代码里没有"裁一块图"。整段视频先切成全部 token,再用下标
gather出上下文 token。被遮 token 不是置零,而是直接删除、不参与计算。
def forward(self, x, masks=None):
# (a) 切 token:3D 卷积;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,位置靠"下标"续命:RoPEAttention(modules.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_keys(modules.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.forward(predictor.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 预测器:VisionTransformerPredictorAC(ac_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.py里from 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.py(CEM)与world_model_wrapper.py(桥接)。二者与训练侧的 rollout 是同一机制的"推理版/训练版"。
7.1 桥接层 WorldModel(world_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 L1(mpc_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
sloss(rollout 项)让预测器在"喂自己预测"时也不漂,推理时 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) → 目标真值 h(EMA, stop-grad)
loss = |ẑ_tgt − h|(仅 masks_pred 上, L1) §六
AC 阶段(§五):把"被遮 patch"换成"未来帧","mask token"换成"动作/状态 token + 块因果掩码",复用同一 latent 预测骨架;推理时反过来用——给目标图像编码成 z_goal,枚举动作让模型在 latent 里 rollout,以 ‖ẑ_future − z_goal‖₁ 为能量,CEM 搜最优动作、执行第一步再重规划(receding horizon)。
九、两条值得记住的设计要点
- multi-block tube 掩码:① 时间贯穿的 tube(
temporal_scale=1.0)逼模型学运动而非抄帧;② 块大小同步、位置随机 + 多块并集,在可批处理前提下做出 ~90% 高难度掩码。 - 下标即位置:掩码通过"删 token + 传下标"实现,位置信息全程由下标 → 3D-RoPE 还原,而非依赖序列顺序。这套机制同时服务于编码器、动作无关预测器和 AC 预测器。
十、工程备忘
- RoPE bug 不可贸然修:
rotate_queries_or_keys的频率展开有已知 bug,但权重与之绑定。加载官方 checkpoint 必须沿用 bug 版;从头自训可切注释里的repeat_interleave(2)正确版。 - 块因果 ≠ is_causal:AC 的时间因果由显式
attn_mask实现,pred_is_frame_causal控制是否构造该矩阵;勿与 SDPA 的逐 tokenis_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 配套:缺了训练侧的
sloss(autoregressive rollout 损失),推理时多步 latent rollout 会漂移、CEM 失准。
配套阅读:方法/数字/机器人结果见 《V-JEPA 2 深度技术剖析》;系列脉络见 《Meta JEPA 系列全面调研》。