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 技术架构与核心方法