Add multiple research papers in PDF format to the repository, including recent works on AI and physics, with file sizes ranging from 1.7 MB to 32.3 MB.
Sync to site1 / sync (push) Has been cancelled

This commit is contained in:
gaojie
2026-06-02 04:12:38 +08:00
parent 36fe037bc1
commit c55f47b287
57 changed files with 278820 additions and 3 deletions
@@ -0,0 +1,542 @@
---
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)。