VLA-JEPA 方法详解

VLA-JEPA 方法详解

1. 核心直觉

VLA-JEPA 把动作分成两层:latent action 不是电机指令,而是"执行后世界会怎样变化"的隐式变量;embodied action 才是机器人真正执行的连续控制量。
flowchart LR A当前多视角图像+指令 --> BQwen3-VL B --> Clatent actions --> D潜空间世界模型 --> E未来语义状态预测 F未来视频帧 --> G冻结 V-JEPA2 --> E B --> Hembodied action tokens --> IFlow Matching DiT --> J未来动作轨迹

像素中有光照、纹理、相机运动等干扰。模型不重建像素,而是预测 V-JEPA2 的 latent state。
flowchart TD P未来像素 --> Q像素重建: 易学到背景/光照 P --> R冻结 V-JEPA2 --> S语义世界状态 --> Tlatent 对齐: 更关注交互

防泄漏原则:未来帧只能构造监督目标,不能进入产生 latent action 的 VLM。
flowchart LR O当前图像+语言 --> VQwen3-VL --> Zlatent action F未来帧 --> E冻结 V-JEPA2 --> Ytarget state Z --> Ppredictor --> YHpredicted state Y & YH --> Lworld loss F -.不能进入 VLM.-> O

2. JEPA 基础:先理解"在表示空间预测"

2.1 JEPA 是什么?

JEPA 全称 Joint-Embedding Predictive Architecture(联合嵌入预测架构)。它的核心不是生成原始数据,而是预测原始数据在特征空间里的表示。

先看三种学习目标的区别:
flowchart TB X被遮挡/未来的内容 A像素生成模型 --> A1预测每个 RGB 像素 B对比学习 --> B1拉近正样本, 推远负样本 CJEPA --> C1预测目标内容的语义 embedding X --> A X --> B X --> C

例如,看到"一只手正推向杯子",JEPA 不需要画出下一帧中杯子的每个像素;它只需要预测"杯子将向右移动"对应的高层特征。背景纹理稍有变化不会造成很大惩罚,但物体运动方向预测错误会造成明显误差。

2.2 JEPA 的四个基本组件

flowchart LR C可见上下文 x --> CEContext Encoder CE --> CX上下文表示 h_x T目标区域/未来片段 y --> TETarget Encoder TE --> TY目标表示 h_y CX --> PPredictor POS目标位置/掩码信息 --> P P --> PY预测表示 h_y_hat TY & PY --> L表示空间对齐损失

  1. Context Encoder:编码模型允许看到的上下文。
  2. Target Encoder:把被遮挡区域或未来片段编码成监督目标。
  3. Predictor:根据上下文表示和目标位置,预测目标表示。
  4. Latent Loss:比较预测表示与目标表示,而不是比较像素。

最简数学形式:

\h_x=f_\\theta(x),\\qquad h_y=f_{\\bar\\theta}(y),\\qquad \\hat h_y=g_\\phi(h_x,m), \\

\\\mathcal L_{JEPA}=d(\\hat h_y,\\operatorname{sg}(h_y)). \\

其中:

  • x 是可见上下文;y 是隐藏目标;m 描述目标的位置或时间;
  • f_theta 是 context encoder,f_bar_theta 是 target encoder;
  • g_phi 是 predictor;
  • sg 表示 stop-gradient,目标分支不接收当前损失的梯度;
  • d 可以是 L1、L2、cosine distance 等。

2.3 一次 JEPA 训练到底发生了什么?

sequenceDiagram participant D as 一张图/一个视频 participant M as Mask/Sampler participant C as Context Encoder participant T as Target Encoder participant P as Predictor D->>M: 采样上下文 x 与隐藏目标 y M->>C: 只给可见上下文 x M->>T: 给目标区域 y C->>P: 上下文表示 h_x M->>P: 目标位置 m P-->>P: 预测目标表示 h_y_hat T-->>P: stop-gradient 目标 h_y P-->>D: 最小化 latent distance

伪代码:

python 复制代码
def jepa_train_step(sample):
    context, target, target_position = sample_context_and_target(sample)

    context_repr = context_encoder(context)

    with no_grad():
        # 目标只提供学习方向,不被 predictor 反向修改
        target_repr = target_encoder(target)

    pred_repr = predictor(context_repr, target_position)
    loss = latent_distance(pred_repr, target_repr)
    update(context_encoder, predictor, loss)
    update_target_encoder_slowly()  # 通用 JEPA 常见做法;本项目不执行此步

2.4 为什么不会退化成"所有输入都输出同一个向量"?

如果 context encoder 和 target encoder 都能被同一损失随意更新,它们可能一起输出常数,此时损失虽然为零,却什么也没学到,这叫 representation collapse。

通用 JEPA 通常结合以下机制避免坍塌:

  • 目标分支 stop-gradient;
  • target encoder 不直接反向更新,而是 context encoder 参数的指数移动平均(EMA);
  • context 和 target 输入不对称:context 看不到被遮挡目标;
  • predictor 只能依据上下文和位置信息完成非平凡预测。

flowchart LR CEContext Encoder theta -->|正常反向传播| U梯度更新 CE -->|EMA: bar_theta = tau*bar_theta + ...| TETarget Encoder bar_theta TE -.stop-gradient.-> LLoss PPredictor phi -->|正常反向传播| U

与当前 VLA-JEPA 代码的区别:这里没有从零训练一对 JEPA encoder。项目直接加载已经预训练好的 V-JEPA2 encoder,并将它冻结;只训练 Qwen3-VL 相关路径和 world predictor。因此它用"固定、已有语义的目标空间"进一步避免目标随训练漂移。

2.5 JEPA、Autoencoder/MAE、对比学习的区别

方法 预测目标 是否需要像素解码器 是否依赖负样本 更关注什么
Autoencoder / MAE 原始像素或低层 patch 细节、纹理、颜色
对比学习 样本间相似性 通常需要大 batch/负样本或其他约束 全局不变性
JEPA 被遮挡/未来内容的 embedding 可预测的高层结构

可以把它们记成:

text 复制代码
MAE  : 根据上下文,把缺失部分"画出来"
对比学习: 判断两份数据是不是同一个语义
JEPA : 根据上下文,说出缺失部分在语义空间中"应该是什么"

JEPA 的优势是避免把容量浪费在不可预测且不重要的像素细节上;代价是结果质量依赖目标表示空间是否真的保留了任务所需信息。

2.6 从 I-JEPA 到 V-JEPA

I-JEPA 面向静态图像,通常遮挡若干图像区域;V-JEPA 面向视频,把目标变成被遮挡的时空块,因此模型必须利用外观、物体运动和时间连续性来预测目标表示。
flowchart LR I一张图像 --> IJI-JEPA IJ --> IP预测隐藏空间区域的 embedding V一段视频 --> VJV-JEPA VJ --> VP预测隐藏时空块/未来片段的 embedding

举例:视频前几帧显示手正把杯子向右推。
flowchart LR F0t0: 手接触杯子 --> F1t1: 杯子开始右移 --> F2t2: 杯子继续右移 F0 & F1 --> C可见时空上下文 F2 --> T隐藏 target C --> P预测器 P --> Z应包含杯子右移的语义 T --> Etarget encoder embedding Z & E --> Llatent 对齐

模型不必确定 t2 的桌面纹理每个像素是什么,但需要保留"哪个物体、怎样运动、空间关系怎样变化"。这正是 VLA-JEPA 想借给机器人策略的知识。

2.7 VLA-JEPA 如何改造 JEPA

标准视频 JEPA 的 predictor 根据视频上下文预测目标 embedding;VLA-JEPA 在 predictor 中额外加入 latent action,并让它负责解释状态转移:

\\\hat s_{t+1}=g_\\phi(s_{\\le t},z_{\\le t}). \\

flowchart TB subgraph Generic普通 V-JEPA VC视频上下文 state --> VPPredictor --> VF未来 latent end subgraph VLAVLA-JEPA SC历史 world states --> WPWorld Predictor ZAQwen 根据当前图像+语言产生 latent action --> WP WP --> SF未来 world state end

这里发生了三个关键变化:

  1. 目标不是任意未来视频特征,而是冻结 V-JEPA2 给出的世界状态 token;
  2. predictor 除历史状态外还必须读取 Qwen3-VL 产生的 latent action;
  3. latent action 路径看不到未来,所以只能从当前场景、指令和时间 token 推断可能的状态变化。

因此,VLA-JEPA 中 JEPA 的作用不是直接输出机器人动作,而是提供一条自监督训练信号:一个好的 latent action,应该足以帮助世界模型预测未来语义状态。

2.8 用一个最小数值例子理解 latent loss

假设 V-JEPA2 把"杯子向右移动后的状态"编码为二维向量 h_y=[0.8, 0.2]。world predictor 根据当前状态和 latent action 得到 h_hat=[0.5, 0.4]

使用当前代码的 mean L1:

\\\mathcal L=\\frac{\|0.5-0.8\|+\|0.4-0.2\|}{2}=0.25. \\

反向传播会更新 predictor,并继续更新产生 latent action 的 Qwen 路径,使下次预测更接近 [0.8,0.2]。V-JEPA2 目标编码器保持冻结。
flowchart LR A预测 0.5,0.4 --> LL1=0.25 B目标 0.8,0.2 --> L L --> P更新 world predictor L --> Q更新 latent-action 产生路径 L -.不更新.-> VV-JEPA2

到这里,可以把 JEPA 记为一句话:不要求模型复原未来长什么样,而要求它在一个有语义的特征空间里预测未来是什么。

3. 数据

3.1 两类数据

数据 动作标签 代表数据集 用途
人类视频 Something-Something-v2,约 220K 学状态转移和时序语义
机器人示范 Droid、LIBERO、BridgeV2、Fractal 联合学世界模型和控制

统一样本接口:

text 复制代码
image : 当前多视角图像,给 VLM,通常 224x224
video : [V, T, H, W, 3],给 V-JEPA2,通常 256x256
lang  : 语言指令
action: [H, 7],机器人样本才有,H=7
state : [1, 8],可选本体状态

3.2 人类视频

代码从目录读取视频,从 CSV 读取文件编号和文本描述,随机截取连续帧;单视角不足两路时复制一份。
flowchart LR V视频文件 --> R随机取 T 帧 --> S缩放 256x256 --> M复制/选择两路视角 --> Wworld video S --> I第0帧缩放224x224 --> QVLM image LCSV文本 --> Q

3.3 机器人数据

LeRobot 数据由 modality.json 和 robot type 配置对齐相机、状态、动作字段。论文说明连续位置/轴角动作 min-max 到 [0,1],夹爪二值化为 {0,1}
flowchart TD ELeRobot episode --> I当前多视角图像 E --> V连续视频片段 E --> A未来7步动作, 每步7维 E --> S可选8维状态 I & V & A & S --> B统一样本字典

4. 模型总览

flowchart TB I当前图像 & L指令 --> QQwen3-VL-2B Q --> Zlatent token hidden states Q --> ZAembodied token hidden states V未来视频 --> FV-JEPA2 frozen encoder --> Y目标 world states Z --> Paction-conditioned predictor --> YH预测未来 states Y & YH --> LWL_WM ZA & S可选 state & A真实动作 --> HFlow Matching DiT --> LAL_FM

4.1 Qwen3-VL 与特殊 token

词表新增 <|action_i|>(第 i 个时间间隔的 latent action)和 <|embodied_action|>(动作头条件)。默认 T=8,每步重复 K=24/T=3 个 latent token,总数 24;具身 token 默认重复 32 次。

数学上:

text 复制代码
z_i = VLM(<latent_i> | 当前图像, 指令)
z_a = VLM(<embodied_action> | 当前图像, 指令, latent tokens)

未来帧不参与上述两个式子的输入。

5. V-JEPA2 世界状态与 latent world model

每个视角单独编码,再拼接:

\s_t = F(I_t\^{(1)}) \\Vert F(I_t\^{(2)}). \\

F 是冻结的 V-JEPA2;s_t 是多视角语义 world state,而不是 RGB。当前代码通过 get_vision_features 提取特征,目标路径使用 no_grad
flowchart LR V1视角1视频 --> F1Frozen V-JEPA2 V2视角2视频 --> F2Frozen V-JEPA2 F1 & F2 --> Cembedding维拼接 --> S统一 world state

预测器接收历史状态和对应 latent action:

\\\hat{s}_{1:T}=p\^{WM}_\\theta(s_{0:T-1},z_{0:T-1}). \\

同一时间步内双向 attention;跨时间只看过去和当前;训练使用 teacher forcing(真实历史 state 作为输入)。
flowchart LR subgraph t0时间 t0 z0latent <--> x0state patches end subgraph t1时间 t1 z1latent <--> x1state patches end subgraph t2时间 t2 z2latent <--> x2state patches end t0 --> t1 --> t2 t1 -.禁止看 t2.-> t2

伪代码:

python 复制代码
def world_predict(states, latent_tokens):
    # states: [B,T-1,N,Dv],历史 V-JEPA 状态
    # latent_tokens: [B,T-1,K,Dq],由 VLM 产生
    x = project_state(states)
    a = project_action(latent_tokens)
    x = interleave(a, x)                 # 每步 [latent, state patches]
    mask = build_time_causal_mask(x)    # 同步双向、跨时间因果
    for block in predictor_blocks:
        x = block(x, attention_mask=mask)
    return output_projection(remove_action_tokens(x))

抽象损失是 L_WM = sum_t d(predicted_state_t, target_state_t)。当前实现明确使用:

python 复制代码
teacher_forcing_wm_loss = F.l1_loss(predicted_states, gt_states)

也就是 mean L1 latent alignment,不是像素 MSE。

6. Flow-matching 动作头

动作头接收 z_a,可选接收本体状态,输出未来 H=7 步、每步 7 维动作。

训练时从噪声到真实动作线性插值:

\a_\\tau=(1-\\tau)\\epsilon+\\tau a,\\qquad v\^\*(a_\\tau)=a-\\epsilon. \\

DiT 学习速度场:

\\\mathcal L_{FM}=\\mathbb E\\\|v_\\theta(a_\\tau,\\tau\\mid z_a,s_0)-(a-\\epsilon)\\\|_2\^2. \\

flowchart LR E高斯噪声 & A真实动作 & T随机时间 --> M线性插值 M --> AE动作+时间编码 ZAz_a & Sstate --> D条件 DiT AE --> D --> V预测速度 V & A & E --> L速度 MSE

伪代码:

python 复制代码
def flow_matching_loss(z_a, gt_action, state=None):
    noise = randn_like(gt_action)
    tau = sample_beta_time(batch_size)       # 每个样本一个时间
    noisy = (1 - tau) * noise + tau * gt_action
    target = gt_action - noise
    action_feat = encode_action_and_time(noisy, tau)
    cond = concat_state_and_future_tokens(action_feat, state)
    hidden = DiT(cond, encoder_hidden_states=z_a, timestep=discretize(tau))
    pred_velocity = action_decoder(hidden[:, -7:])
    return mean_square(pred_velocity, target)

推理从随机动作开始,执行 4 次 Euler 积分:

\a \\leftarrow a+\\Delta t\\,v_\\theta(a,t\\mid z_a,s_0). \\

flowchart TD N随机动作 --> D1DiT预测速度 --> U1Euler更新 --> D2DiT预测速度 --> U2Euler更新 U2 --> D3DiT预测速度 --> U3Euler更新 --> D4DiT预测速度 --> O输出7步动作

7. 训练 pipeline

7.1 人类视频 batch

无动作标签,所以只产生 wm_loss
flowchart LR B视频batch --> QQwen: 当前帧+语言 --> Zlatent B --> FV-JEPA2: 未来帧 --> Ytarget state Z & F --> Ppredictor --> YHprediction Y & YH --> Lwm_loss

7.2 机器人 batch

flowchart TB R机器人batch --> QQwen3-VL Q --> Zlatent --> Wworld predictor Q --> ZAembodied tokens --> HFlow DiT R --> FV-JEPA2 target --> W W --> LWwm_loss R --> H --> LAaction_loss LW & LA --> Totaltotal loss

论文总目标:L_robot = L_FM + beta * L_WM。当前仓库的具体行为:VLA_JEPA.forward 返回两个损失,机器人路径把 wm_loss 乘以 0.1,训练器对返回字典求和;因此当前实现约为 action_loss + 0.1 * wm_loss。人类视频路径只返回 wm_loss

联合训练 step:

python 复制代码
def train_step(robot_batch, video_batch):
    # 机器人:控制 + 世界模型
    robot_out = model(robot_batch)
    optimizer_step(sum(robot_out.values()))
    # 人类视频:继续强化时序语义
    video_out = model(video_batch)
    optimizer_step(sum(video_out.values()))

论文训练阶段:SSV2+Droid 约 50K steps;仿真微调约 30K;真实机器人微调约 20K。常用 AdamW、cosine schedule + warmup、混合精度。

8. 推理:当前代码实际执行的路径

  1. 输入当前多视角图像、语言和可选 state;
  2. Qwen3-VL 取最后一层 hidden states;
  3. 只抽取 <|embodied_action|> 的 32 个 hidden states;
  4. Flow DiT 从噪声动作块开始,做 4 步 Euler;
  5. 返回归一化动作数组。

sequenceDiagram participant E as 机器人环境 participant P as Policy participant Q as Qwen3-VL participant H as Flow DiT E->>P: 图像+指令+state P->>Q: prompt + embodied tokens Q-->>P: z_a P->>H: z_a + state + noise action loop 4 steps H-->>P: velocity P->>P: Euler update end P-->>E: 未来7步动作

当前 predict_action 不调用未来视频编码器和 vj_predictor;世界模型主要通过训练塑造 VLM 表征。动作执行后重新观察,再滚动预测。

9. 例子:把杯子放进抽屉

输入是机械臂、杯子、打开的抽屉和指令"把杯子放入抽屉"。latent action 不需要输出厘米级位移;在 world loss 约束下,它会形成阶段性的转移变量:
flowchart LR s0杯子在桌面 --> z0靠近 --> s1夹爪对准 s1 --> z1抓取 --> s2杯子被抓住 s2 --> z2搬运 --> s3接近抽屉 s3 --> z3放置 --> s4杯子进入抽屉

随后 z_a 让 DiT 把噪声轨迹变成 [dx,dy,dz,drx,dry,drz,gripper] x 7。人类视频还可能帮助模型学到"抓取失败后松开并重新尝试"的时序决策;它增强的是恢复行为,不等于提供机器人关节标签。

10. 常见误解

  1. latent action 不是动作标签压缩版,而是为未来语义状态预测服务的中间变量。
  2. 未来帧不是完全不用,而是只进入冻结 target encoder。
  3. 预测的是 V-JEPA2 embedding,不是 RGB。
  4. predictor 训练使用真实历史 state 做 teacher forcing,不是完全开环。
  5. latent action 不是相邻帧差分,而是当前观测+语言条件下的 learnable query。
  6. 论文的 L_WM 距离写得抽象,当前代码是 mean L1。
  7. 论文写 L_FM + beta L_WM,当前代码约为 action_loss + 0.1*wm_loss
  8. 人类视频主要提升扰动场景的稳定性;对 SimplerEnv 不保证提升。
  9. ELBO 是解释视角,不是完整概率生成模型。

11. 超参数速查

项目
VLM Qwen3-VL-2B
target encoder V-JEPA2
world horizon 8
latent token/step 3(24/8)
embodied tokens 32
predictor 12 layers, 8 heads
action head DiT-B 风格,16 layers
action horizon 7
action dimension 7
inference steps 4
world loss(代码) mean L1
robot loss(代码近似) action + 0.1 × world

12. 最终心智模型

flowchart TD U大量无标签视频 --> A学会未来状态变化 R少量机器人示范 --> B学会具体控制轨迹 A & B --> C共享 VLM 表征 C --> Dlatent world model C --> Eflow action head D --> F时序理解/鲁棒性 E --> G机器人执行

只记住三点:

  1. 未来帧只做 target,不进入 VLM
  2. latent action 通过预测未来 V-JEPA 状态获得语义
  3. 真正的机器人控制由 embodied token + flow-matching DiT 生成