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

543 lines
27 KiB
Markdown
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
---
title: "V-JEPA 2 源码逐模块剖析(代码篇)"
date: 2026-05-24
draft: false
tags: ["V-JEPA 2", "JEPA", "源码", "code-walkthrough", "self-supervised", "RoPE", "masking", "world-model", "robotics"]
categories: ["JEPA"]
description: "基于 facebookresearch/vjepa2 本地源码的逐模块剖析:multi-block tube 掩码采样、编码器(3D-RoPE/apply_masks)、动作无关预测器、V-JEPA 2-AC 块因果预测器,以及自监督训练全链路闭环。是《V-JEPA 2 深度技术剖析》的代码篇配套。"
---
> 本文是 [《V-JEPA 2 深度技术剖析》](V-JEPA2深度技术剖析.md) 的**代码篇配套**,全部基于本地 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`,两套并行):
```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 网格大小不同,不能混),再对每套掩码配置各调一次生成器:
```python
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 种子心机:块"大小"同步,块"位置"随机
```python
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 + 长宽比反解):
```python
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=上下文`):
```python
mask = torch.ones((duration, height, width))
mask[start:start+t, top:top+h, left:left+w] = 0
```
### 1.4 多块取并集,切出两组下标
```python
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 取最短并截断:
```python
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`
```python
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 不是置零,而是**直接删除、不参与计算**。
```python
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`
```python
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 被删一大半、顺序也乱,模型怎么知道每个上下文块原来的位置?答案在传下来的下标:
```python
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` 把一维下标反解成 (帧, 高, 宽) 三个坐标:
```python
frame_ids = ids // tokens_per_frame
height_ids = (ids - 帧分量) // tokens_per_row
width_ids = ids - 帧分量 - 高分量
```
再把每个 head 的特征维三等分,按 t/h/w 各自旋转:
```python
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`
```python
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`,在**目标位置**把表示补出来。
```python
def forward(self, x, masks_x, masks_y, ...):
# x: 上下文表示(z_ctx)masks_x: 上下文下标;masks_y: 目标下标
```
### 4.1 先降维("预测器很轻"的代码证据)
```python
x = self.predictor_embed(x) # Linear: 1408 -> 384
```
进门就压到 384 维,整个预测器在窄空间算,算完再投影回去。
### 4.2 给目标位置造可学习占位符
```python
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 拼接 + 排序
```python
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,抠回目标、升回维
```python
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 维
```python
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 注入方式:把动作/状态插进每帧开头
```python
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`
```python
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`
```python
# 动作/状态 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 预测
```python
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`
```python
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 超参:
```python
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**(与训练目标同一归一化空间,能量尺度才对齐):
```python
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 预测器"想象",本体位姿靠运动学硬算**:
```python
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`
```python
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:
```python
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 + 选精英 + 动量更新
```python
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`
```python
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,训练就必须同时练单步和多步,否则误差累积漂移:
```python
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) → 目标真值 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 掩码**:① 时间贯穿的 tube`temporal_scale=1.0`)逼模型学运动而非抄帧;② 块大小同步、位置随机 + 多块并集,在可批处理前提下做出 ~90% 高难度掩码。
2. **下标即位置**:掩码通过"删 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 的逐 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 配套**:缺了训练侧的 `sloss`autoregressive rollout 损失),推理时多步 latent rollout 会漂移、CEM 失准。
---
> 配套阅读:方法/数字/机器人结果见 [《V-JEPA 2 深度技术剖析》](V-JEPA2深度技术剖析.md);系列脉络见 [《Meta JEPA 系列全面调研》](index.md)。