跳转至

02.1 JEPA 技术细节深挖

flowchart LR
    x["输入 x"]
    ctx["Context Encoder<br/>(可见 patches)"]
    tgt["Target Encoder<br/>(EMA, 完整 x)"]
    pred["Predictor<br/>窄 ViT"]
    sx["s_x (context repr)"]
    sy["s_y (target repr)"]
    sy_hat["ŝ_y (predicted)"]
    loss["L1 / cosine loss<br/>在 latent 空间"]

    x --> ctx --> sx --> pred --> sy_hat
    x --> tgt --> sy
    sy --> loss
    sy_hat --> loss

    classDef stage fill:#fff,stroke:#cc785c,color:#1a1a1a;
    class x,ctx,tgt,pred,sx,sy,sy_hat,loss stage

核心:在 latent 空间预测,而非像素重建——避免高频细节噪声,强制模型学习语义抽象。

V-JEPA 2 总览图(原论文 arXiv:2506.09985 Figure 1)

I-JEPA 架构图(原论文 arXiv:2301.08243 Figure 3)

最后更新: 2026-04-14 | 深度调研

核心发现

I-JEPA 架构

  • 三模块: 目标编码器(EMA) + 上下文编码器 + 预测器(窄ViT)

  • 掩码: 4个目标块(15-20%面积) + 1个上下文块(85-100%)

  • 损失: L2(表示空间), 非像素空间

  • ViT-H/14 ImageNet线性探测: 79.3% (vs MAE 77.2%)

  • 比MAE快10x+, 无需数据增强

V-JEPA 2 (2025.06)

  • 1.2B参数 ViT-g, 100万+小时视频训练

  • L1损失(比L2更鲁棒), 3D-RoPE位置编码

  • 动作条件训练: 300M参数Transformer预测器, AdaLN注入

  • 机器人杯子操控: 80%成功率(vs Octo 15%)

  • 规划速度: 16秒/动作(vs Cosmos 4分钟)

LeWorldModel (2026.03) — 端到端JEPA

  • SIGReg正则化: 基于Cramér-Wold定理, 通过Epps-Pulley检验强制等向高斯

  • 仅2项损失(从7项简化), 无需EMA/停止梯度

  • 规划速度比DINO-WM快48x

  • 局限: 低复杂度环境表现差

LeCun反对生成模型的技术论证

  • H(像素下一帧|当前帧) >> H(表示下一帧|当前帧)

  • 像素级预测浪费计算在不可预测的高频细节上

  • JEPA用低维表示空间(192-768维) vs 像素空间(150K+维)

AMI Labs (2025.12成立)

  • LeCun执行主席, 2026.03完成\(10.3亿种子轮, 估值\)45亿

  • 总部巴黎, 目标: 将JEPA工业化为通用世界模型

JEPA vs 其他SSL方法

  • vs SimCLR: 无需负样本和数据增强

  • vs BYOL: 预测不同区域而非增强不变性

  • vs DINO/DINOv2: 无需精选数据(142M)

  • vs MAE: 表示空间预测 >> 像素预测(40.7% vs 66.9% @1%标签)

参考


工程实现细节

I-JEPA 多块掩码算法

这是 I-JEPA 成功的关键:目标块必须足够大(迫使学习语义),上下文必须去除与目标重叠(避免信息泄漏)。

import numpy as np

def sample_target_context_blocks(
    patch_h=14, patch_w=14,    # ViT-H/14: 14x14 patches
    n_targets=4,
    target_scale=(0.15, 0.2),   # 15-20% 图像面积
    target_ar=(0.75, 1.5),      # 长宽比
    context_scale=(0.85, 1.0),  # 85-100% 图像面积
):
    """
    I-JEPA 块采样算法 (基于 github.com/facebookresearch/ijepa/src/masks/multiblock.py)

    张量形状约定:
    - patch_mask: (patch_h, patch_w) bool 掩码
    - 返回: (target_mask, context_mask, target_block_coords)
    """

    # ============ 1. 采样 4 个目标块 ============
    target_mask = np.zeros((patch_h, patch_w), dtype=bool)
    target_blocks = []

    for _ in range(n_targets):
        # 尺度 s ~ U(0.15, 0.2), 实际面积
        s = np.random.uniform(*target_scale)
        total_patches = patch_h  patch_w  s  # 期望 patch 数

        # 长宽比 a ~ U(0.75, 1.5)
        a = np.random.uniform(*target_ar)

        # 推出 h, w (满足 h*w = total_patches 且 h/w = a)
        h = int(np.round(np.sqrt(total_patches * a)))
        w = int(np.round(np.sqrt(total_patches / a)))
        h = min(max(h, 1), patch_h)
        w = min(max(w, 1), patch_w)

        # 随机位置
        top = np.random.randint(0, patch_h - h + 1)
        left = np.random.randint(0, patch_w - w + 1)

        # 标记为目标 (可重叠)
        target_mask[top:top+h, left:left+w] = True
        target_blocks.append((top, left, h, w))

    # ============ 2. 采样 1 个上下文块 ============
    s_ctx = np.random.uniform(*context_scale)
    total_ctx = patch_h  patch_w  s_ctx
    h_ctx = int(np.round(np.sqrt(total_ctx)))
    w_ctx = int(np.round(np.sqrt(total_ctx)))

    top_ctx = np.random.randint(0, patch_h - h_ctx + 1)
    left_ctx = np.random.randint(0, patch_w - w_ctx + 1)

    context_mask = np.zeros((patch_h, patch_w), dtype=bool)
    context_mask[top_ctx:top_ctx+h_ctx, left_ctx:left_ctx+w_ctx] = True

    # ============ 3. 关键:从 context 中移除 target 重叠 ============
    context_mask = context_mask & ~target_mask

    return target_mask, context_mask, target_blocks

# ============ 张量形状 (ViT-H/14, 224×224) ============
# Patch grid: 14 × 14 = 196 patches
# 每 patch 特征维度 D = 1280 (ViT-H 的 embed_dim)
# Patch tokens: (B, 196, 1280)
#
# 经过 sample_target_context_blocks:
#   target_mask:  (14, 14) - 约 59-78 个 True (15-20%)
#   context_mask: (14, 14) - 约 70-190 个 True (去除 target 后)
#
# Encoder 输入:
#   patch_tokens[context_mask]  → (B, N_ctx, 1280)    N_ctx ≈ 70-190
#   patch_tokens[target_mask]   → (B, N_tgt, 1280)    N_tgt ≈ 59-78

EMA Target Encoder 与 Stop-Gradient 精确位置

import torch
import copy

class IJEPA:
    def __init__(self, encoder_online, total_steps=300_000):
        self.encoder = encoder_online              # 在线编码器 (有梯度)
        self.target_encoder = copy.deepcopy(encoder_online)  # EMA 目标编码器

        # 冻结 target encoder 的参数
        for p in self.target_encoder.parameters():
            p.requires_grad = False

        self.total_steps = total_steps
        self.step = 0

        # EMA 余弦调度: 0.996 → 1.0
        self.tau_base = 0.996
        self.tau_final = 1.0

    def get_ema_tau(self):
        """EMA 系数随步数余弦增长"""
        progress = min(self.step / self.total_steps, 1.0)
        # cos(pi * progress) 从 1 → -1
        # (1-cos)/2 从 0 → 1
        tau = self.tau_base + (self.tau_final - self.tau_base) * (
            0.5  (1 - np.cos(np.pi  progress))
        )
        return tau

    @torch.no_grad()
    def update_target_encoder(self):
        """EMA 参数更新 (无梯度)"""
        tau = self.get_ema_tau()
        for p_online, p_target in zip(
            self.encoder.parameters(),
            self.target_encoder.parameters()
        ):
            p_target.data.mul_(tau).add_(p_online.data, alpha=1 - tau)
        self.step += 1

    def training_step(self, image):
        """
        image: (B, 3, 224, 224)

        Stop-gradient 精确位置 (CRUCIAL):
        1. target encoder 的 forward 整个在 no_grad 中
        2. predictor 输出与 target 计算 loss 时, target 已经 detach
        """
        # ========== 1. 采样掩码 ==========
        target_mask, context_mask, _ = sample_target_context_blocks()

        # Patch 化
        patches = patchify(image)  # (B, 196, D_in)

        # ========== 2. Online Encoder 处理 context (有梯度) ==========
        context_patches = patches[:, context_mask.flatten()]  # (B, N_ctx, D)
        z_ctx = self.encoder(context_patches)                 # (B, N_ctx, 1280)

        # ========== 3. Target Encoder 处理 FULL image (stop gradient!) ==========
        with torch.no_grad():
            z_full_target = self.target_encoder(patches)      # (B, 196, 1280)
            z_target = z_full_target[:, target_mask.flatten()]  # (B, N_tgt, 1280)

        # ========== 4. Predictor 预测 target 位置的表征 ==========
        # mask tokens with target positional embeddings
        target_pos_emb = self.pos_embed[target_mask.flatten()]  # (N_tgt, 384)
        mask_tokens = self.mask_token.expand(B, N_tgt, 384) + target_pos_emb

        # Predictor 输入: [context_tokens | mask_tokens]
        pred_input = torch.cat([
            self.ctx_proj(z_ctx),   # (B, N_ctx, 384) 投影到 predictor 维度
            mask_tokens             # (B, N_tgt, 384)
        ], dim=1)

        pred_output = self.predictor(pred_input)                # (B, N_ctx+N_tgt, 384)
        z_pred = pred_output[:, -N_tgt:]                        # (B, N_tgt, 384)
        z_pred = self.pred_proj(z_pred)                         # (B, N_tgt, 1280)

        # ========== 5. L2 Loss (target 已被 no_grad 保护) ==========
        # 无需再 .detach(), 因为 z_target 来自 no_grad context
        loss = torch.nn.functional.smooth_l1_loss(z_pred, z_target)

        loss.backward()
        self.optimizer.step()
        self.update_target_encoder()  # 反传后 EMA 更新

        return loss.item()

V-JEPA 2: 3D Tubelet 编码与 3D-RoPE

# ============ 3D Tubelet 提取 ============
# 视频: (B, T=64, C=3, H=384, W=384)
# Tubelet 大小: (t=2, h=16, w=16)
# → 张量形状: (B, 32, 24, 24, embed_dim) → 展平为 18,432 tokens

class TubeletEmbed3D(nn.Module):
    def __init__(self, tubelet=(2, 16, 16), embed_dim=1408, in_chans=3):
        # embed_dim=1408 for ViT-g (V-JEPA 2 标准配置)
        # ViT-g config: depth=40, num_heads=16, mlp_ratio=48/11
        super().__init__()
        self.tubelet = tubelet
        # 用 3D Conv 实现 tubelet 提取 + linear projection
        self.proj = nn.Conv3d(
            in_chans, embed_dim,
            kernel_size=tubelet,
            stride=tubelet
        )
        # 3D-RoPE 位置编码(将在后续应用)
        self.pos_embed = nn.Parameter(
            torch.randn(1, 321616, embed_dim)
        )

    def forward(self, video):
        # video: (B, T=64, C=3, H=384, W=384)
        B, T, C, H, W = video.shape
        video = video.permute(0, 2, 1, 3, 4)  # (B, C, T, H, W)
        tubelets = self.proj(video)
        # (B, 1408, 32, 24, 24)
        B, D, nt, nh, nw = tubelets.shape
        tubelets = tubelets.permute(0, 2, 3, 4, 1).reshape(B, ntnhnw, D)
        # (B, 18432, 1408) ← 注意: 最后一维是 1408,非 1024
        tubelets = tubelets + self.pos_embed[:, :tubelets.size(1), :]
        encoded = self.vit(tubelets)  # (B, 18432, 1408)
        return encoded

V-JEPA 2 渐进分辨率训练(8.4× 加速来源)

每阶段 batch size 自动调整以保持显存使用恒定。早期低分辨率快速训练让网络学会基础特征,再逐步升级让网络适应高分辨率精细信息。整体比固定 64×384² 快 8.4×。

LeWM SIGReg 完整实现

import torch
import torch.nn.functional as F

def sigreg_loss(z, n_projections=256, temperature=1.0):
    """
    SIGReg: Sketched Isotropic Gaussian Regularization
    强制 z 在所有一维投影上服从标准正态

    z: (B, D) - embedding batch
    n_projections: M, 投影方向数

    时间复杂度: O(B  M  D), 内存 O(B * M)
    """
    B, D = z.shape

    # ========== 1. 标准化 z ==========
    # 先减去 batch 均值 (让 mean ≈ 0)
    z_centered = z - z.mean(dim=0, keepdim=True)

    # ========== 2. 采样 M 个球面均匀方向 ==========
    # 从 N(0,I) 采样后归一化得到 S^{D-1} 上均匀分布
    u = torch.randn(n_projections, D, device=z.device)
    u = F.normalize(u, dim=1)  # (M, D)

    # ========== 3. 投影 ==========
    projections = z_centered @ u.T  # (B, M)
    # 每列 projections[:, m] 是 B 个 z 在方向 u_m 上的一维投影

    # 再标准化每列 (使 var ≈ 1)
    projections = projections / (projections.std(dim=0, keepdim=True) + 1e-8)

    # ========== 4. Epps-Pulley 检验统计量 ==========
    # 理论: 如果 z ~ N(0,I), 则 projections[:, m] 都 ~ N(0,1)
    # E-P 检验基于经验特征函数 vs 标准正态特征函数的平方差

    # 选择 K 个节点 t_k (一维投影空间的特征函数评估点)
    K = 50
    t_knots = torch.linspace(0.1, 5.0, K, device=z.device)

    # 权重: 高斯核, 越远权重越小
    w = torch.exp(-0.5  t_knots*2)
    w = w / w.sum()  # (K,)

    # 经验特征函数: φ_n(t) = (1/B) Σ exp(itx_b)
    # 实部 = cos, 虚部 = sin
    tx = torch.einsum('bm,k->bmk', projections, t_knots)  # (B, M, K)
    phi_emp_real = torch.cos(tx).mean(dim=0)               # (M, K)
    phi_emp_imag = torch.sin(tx).mean(dim=0)               # (M, K)

    # 理论特征函数 (标准正态): φ(t) = exp(-t²/2), 虚部=0
    phi_true_real = torch.exp(-0.5  t_knots*2)           # (K,)
    # phi_true_imag = 0

    # 平方差
    diff_sq = (phi_emp_real - phi_true_real)2 + phi_emp_imag2  # (M, K)

    # 加权求和
    ep_stats = (diff_sq * w).sum(dim=1)  # (M,) - 每个投影方向一个检验统计量

    # ========== 5. 聚合成单一 loss ==========
    sigreg = ep_stats.mean()

    return sigreg

# ============ 完整训练循环 ============
def train_step_lewm(batch, model, optimizer, lambda_sig=0.05):
    images = batch['images']      # (B, T, C, H, W)
    actions = batch['actions']    # (B, T, A)

    # Encode all frames
    z_all = model.encoder(images.flatten(0, 1))  # (B*T, 192)
    z_all = z_all.reshape(B, T, 192)

    # Predict next z given action
    z_pred = model.predictor(z_all[:, :-1], actions[:, :-1])  # (B, T-1, 192)

    # ========== Loss 1: Prediction ==========
    L_pred = F.mse_loss(z_pred, z_all[:, 1:].detach())

    # ========== Loss 2: SIGReg (only on z, not z_pred) ==========
    L_sig = sigreg_loss(z_all.flatten(0, 1))  # (B*T, 192)

    # ========== Total ==========
    L = L_pred + lambda_sig * L_sig

    optimizer.zero_grad()
    L.backward()
    optimizer.step()

    return {'L_pred': L_pred.item(), 'L_sig': L_sig.item()}

SIGReg 理论优势对比



附录:关键数学公式

I-JEPA 训练目标

I-JEPA 在表示空间预测被遮挡目标块:

\(\mathcal{L}_{\text{I-JEPA}} = \frac{1}{M}\sum_{i=1}^{M}\sum_{j\in\mathcal{B}_i}\left\|\hat{s}_{y_j} - \operatorname{sg}(s_{y_j})\right\|_2^2\)

其中 M=4 个目标块,B_i 是第 i 个目标块中的 patch 集合:

\(\hat{s}_{y_j} = g_\phi(s_x,\operatorname{pos}(j)),\quad s_{y_j} = f_{\bar\theta}(y_j)\)

EMA 更新规则(余弦调度)

\(\bar\theta \leftarrow \tau \cdot \bar\theta + (1-\tau)\cdot\theta\)

\(\tau\) 在训练过程中按余弦从 \(\tau_{\text{base}}=0.996\) 增长到 \(\tau_{\text{final}}=1.0\)

\(\tau(t) = \tau_{\text{base}} + (\tau_{\text{final}} - \tau_{\text{base}})\cdot\frac{1-\cos(\pi t/T)}{2}\)

V-JEPA 2 训练目标(L1 更鲁棒)

\(\mathcal{L}_{\text{V-JEPA}2} = \left\|P_\phi\left(\Delta_y,\ E_\theta(x)\right) - \operatorname{sg}(E_{\bar\theta}(y))\right\|_1\)

V-JEPA 2-AC 动作条件损失

教师强制:\(\mathcal{L}_{\text{tf}} = \frac{1}{T}\sum_{k=1}^{T}\|\hat z_{k+1} - z_{k+1}\|_1\)

多步展开:\(\mathcal{L}_{\text{roll}} = \|P_\phi(a_{1:T},s_1,z_1) - z_{T+1}\|_1\)

LeWM SIGReg(基于 Cramér-Wold 定理)

总损失:

\(\mathcal{L}_{\text{LeWM}} = \underbrace{\|\hat z_{t+1} - z_{t+1}\|_2^2}_{\text{prediction}} + \lambda\cdot\underbrace{\operatorname{SIGReg}(Z)}_{\text{Gaussianity}}\)

SIGReg 是 M 个随机方向上一维投影的 Epps-Pulley 检验统计量平均:

\(\operatorname{SIGReg}(Z) = \frac{1}{M}\sum_{m=1}^{M} T^{(m)},\quad u^{(m)}\in\mathcal{S}^{d-1}\)

对每个方向的投影 h^{(m)} = Z u^{(m)},Epps-Pulley 统计量为经验特征函数与目标特征函数之差的加权积分:

\(T^{(m)} = \int w(t)\left|\hat\varphi(t;h^{(m)}) - \varphi_0(t)\right|^2 dt\)

其中:

经验特征函数:\(\hat\varphi(t) = \frac{1}{N}\sum_{b=1}^N e^{it h_b^{(m)}}\)

标准正态特征函数:\(\varphi_0(t) = e^{-t^2/2}\)

JEPA vs 生成模型的信息论论证

LeCun 的核心论点:像素空间条件熵远大于表示空间:

\(H(X_{t+1}\mid X_t) \gg H(s_{t+1}\mid s_t)\)

因此在表示空间预测可以过滤不可预测的高频细节,把算力集中在可预测结构上。

防坍缩的互信息下界

\(I(s_x;s_y) \ge \frac{1}{2}\log\left(1+\frac{\|\hat s_y - s_y\|^2}{\operatorname{Var}(s_y)}\right)^{-1}\)

最小化预测误差同时最大化互信息,等效于信息瓶颈。


Code 引用索引与可信度

本节所有伪代码的出处与可信度。✓=官方代码直接移植, △=基于论文描述重构, ⚠=推测填充 ——

注: 所有 "✓" 级代码的完整版本请参考对应 GitHub repo;本文的 PyTorch 伪代码经过简化以突出核心逻辑。


上级 · 02 技术架构与核心方法