
从Token Mask到完整训练闭环:彻底理解V-JEPA的x-Encoder、y-Encoder、Mask Token、Predictor如何协同工作并预测被遮挡视频的目标特征与完整数据流和训练逻辑前面已经知道,一段视频经过时空块嵌入(Tubelet Embedding)之后,会从原始的视频帧转换成一个三维时空 Token 网格(Spatio-temporal Token Grid)。例如 V-JEPA 中典型的输入:[16,224,224,3][16,224,224,3][16,224,224,3]经过大小为:2×16×162\times16\times162×16×16的三维 Patch Embedding 后,会得到:[8,14,14,de][8,14,14,d_e][8,14,14,de]再 Flatten 成:[1568,de][1568,d_e][1568,de]因此完整视频一共有:L=1568L=1568L=1568个 Token。接下来真正进入 V-JEPA 最核心的训练流程:哪些 Token 会被删除?x-Encoder(上下文编码器,Context Encoder)到底看到什么?被删除的位置怎样重新告诉Predictor(预测器)?Mask Token(掩码 Token)里面有没有真实的视频内容?为什么y-Encoder(目标编码器,Target Encoder)反而要看到完整视频?Predictor 最终预测的是 RGB,还是 Feature?理解这些问题之后,V-JEPA Figure 3 中最核心的数据流就完全可以串起来。1. 先正式定义LLL、MMM和NNN完整 Token 序列:xL=(x1,x2,…,xL)x_L=(x_1,x_2,\ldots,x_L)xL=(x1,x2,…,xL)在 V-JEPA 的真实例子中:L=1568L=1568L=1568假设其中有MMM个 Token 被掩码(Masked Tokens)。Mask 索引为:i1,i2,…,iMi_1,i_2,\ldots,i_Mi1,i2,…,iM剩下的未被 Mask 的 Token 数量(Visible Token Number)为:N=L−MN=L-MN=L−M可见 Token 的索引为:j1,j2,…,jNj_1,j_2,\ldots,j_Nj1,j2,…,jN因此:L=N+ML=N+ML=N+M可以理解为:完整视频 Token ───────────────────────────────── Visible Tokens Target Tokens N M N + M = L所以这里三个变量一定要先牢牢记住:LLL:完整视频的 Token 总数;MMM:被 Mask、需要预测的 Target Token 数;NNN:Mask 后剩余的可见 Context Token 数。因此:L=N+M\boxed{L=N+M}L=N+M后面几乎所有 Shape 的变化,都围绕这三个量展开。2.xxx到底表示什么?在 V-JEPA 中,可以把xxx理解成:视频经过 Mask 以后剩下的可见上下文(Visible Context)。也就是:xN=(xj1,xj2,…,xjN)x_N=(x_{j_1},x_{j_2},\ldots,x_{j_N})xN=(xj1,xj2,…,xjN)而 Target 对应的则是被 Mask 的那些位置:i1,…,iMi_1,\ldots,i_Mi1,…,iM这些 Target 位置最终需要产生:si1,…,siMs_{i_1},\ldots,s_{i_M}si1,…,siM作为yyy分支给出的目标表示。所以整个问题可以先粗略理解成:完整视频 Token │ ├───────────────┐ │ │ ▼ ▼ Visible Context Target Positions N M │ │ ▼ │ x-Encoder │ │ │ ▼ │ Context Features │ │ │ └──── Predictor ◄──┘ │ ▼ Predicted Targets这里真正特别的地方在于:x-Encoder 本身并不会看到 Target Token。3. x-Encoder 最关键的一点:不是放[MASK],而是真的删除假设现在使用一个玩具例子:L=12L=12L=12Mask 为:2,6,10{2,6,10}2,6,10完整 Token 是:1 2 3 4 5 6 7 8 9 10 11 12一个非常容易产生的错误理解是:1 MASK 3 4 5 MASK 7 8 9 MASK 11 12然后把这121212个位置全部交给x-Encoder(上下文编码器,Context Encoder)。V-JEPA 实际不是这样。它会直接删除(Drop)被 Mask 的 Target Token:1 3 4 5 7 8 9 11 12因此:N=9N=9N=9真正输入 x-Encoder 的是:xN=(x1,x3,x4,x5,x7,x8,x9,x11,x12)x_N=(x_1,x_3,x_4,x_5,x_7,x_8,x_9,x_{11},x_{12})xN=(x1,x3,x4,x5,x7,x8,x9,x11,x12)形状为:[9,de][9,d_e][9,de]经过 x-Encoder:zN=Eθ(xN)z_N=E_\theta(x_N)zN=Eθ(xN)得到:zN=(z1,z3,z4,z5,z7,z8,z9,z11,z12)z_N=(z_1,z_3,z_4,z_5,z_7,z_8,z_9,z_{11},z_{12})zN=(z1,z3,z4,z5,z7,z8,z9,z11,z12)形状仍然为:[9,de][9,d_e][9,de]论文 Figure 3 明确把第一步画成:Remove masked tokens也就是:[L,de]→[N,de][L,d_e]\rightarrow[N,d_e][L,de]→[N,de]而不是:[L,de]→[L,de][L,d_e]\rightarrow[L,d_e][L,de]→[L,de]并在目标位置塞入[MASK]。这一点非常重要。4. 为什么要先删除,再进入大型 Encoder?这里除了一个显而易见的原因:不让 x-Encoder 直接看到需要预测的 Target 内容。还有一个非常现实的原因:节省计算量(Computational Efficiency)。Transformer 中自注意力(Self-Attention)的计算复杂度大致随 Token 数量平方增长:O(N2)O(N^2)O(N2)如果完整视频原本有:L=1568L=1568L=1568个 Token,而其中很大一部分已经被 Mask,那么根本没有必要让昂贵的大型 ViT Encoder 继续处理那些 Target Token。所以整个思路是:完整 Token │ ▼ 先删除大量 Target Token │ ▼ 只保留 Visible Tokens │ ▼ 大型 x-Encoder也就是:昂贵的大型 Encoder 只处理真正可见的 Context Token。这种设计也延续了 MAE 和 I-JEPA 中:Encoder 只处理可见 Patch。的高效思想。5. 但是这马上产生一个问题:Predictor 怎么知道哪里缺了?现在 x-Encoder 输出的是:z1,z3,z4,z5,z7,z8,z9,z11,z12z_1,z_3,z_4,z_5,z_7,z_8,z_9,z_{11},z_{12}z1,z3,z4,z5,z7,z8,