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表示空间对齐损失
- Context Encoder:编码模型允许看到的上下文。
- Target Encoder:把被遮挡区域或未来片段编码成监督目标。
- Predictor:根据上下文表示和目标位置,预测目标表示。
- 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
这里发生了三个关键变化:
- 目标不是任意未来视频特征,而是冻结 V-JEPA2 给出的世界状态 token;
- predictor 除历史状态外还必须读取 Qwen3-VL 产生的 latent action;
- 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. 推理:当前代码实际执行的路径
- 输入当前多视角图像、语言和可选 state;
- Qwen3-VL 取最后一层 hidden states;
- 只抽取
<|embodied_action|>的 32 个 hidden states; - Flow DiT 从噪声动作块开始,做 4 步 Euler;
- 返回归一化动作数组。
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. 常见误解
- latent action 不是动作标签压缩版,而是为未来语义状态预测服务的中间变量。
- 未来帧不是完全不用,而是只进入冻结 target encoder。
- 预测的是 V-JEPA2 embedding,不是 RGB。
- predictor 训练使用真实历史 state 做 teacher forcing,不是完全开环。
- latent action 不是相邻帧差分,而是当前观测+语言条件下的 learnable query。
- 论文的
L_WM距离写得抽象,当前代码是 mean L1。 - 论文写
L_FM + beta L_WM,当前代码约为action_loss + 0.1*wm_loss。 - 人类视频主要提升扰动场景的稳定性;对 SimplerEnv 不保证提升。
- 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机器人执行
只记住三点:
- 未来帧只做 target,不进入 VLM;
- latent action 通过预测未来 V-JEPA 状态获得语义;
- 真正的机器人控制由 embodied token + flow-matching DiT 生成。