VLA-JEPA 方法详解

VLA-JEPA 方法详解

1. 核心直觉

VLA-JEPA 把动作分成两层:latent action 不是电机指令,而是“执行后世界会怎样变化”的隐式变量;embodied action 才是机器人真正执行的连续控制量。

flowchart LR A[当前多视角图像+指令] --> B[Qwen3-VL] B --> C[latent actions] --> D[潜空间世界模型] --> E[未来语义状态预测] F[未来视频帧] --> G[冻结 V-JEPA2] --> E B --> H[embodied action tokens] --> I[Flow Matching DiT] --> J[未来动作轨迹]

像素中有光照、纹理、相机运动等干扰。模型不重建像素,而是预测 V-JEPA2 的 latent state。只有当目标编码器对某些外观变化具有不变性时,这种训练才会相应弱化这些变化的影响;JEPA 本身不自动保证“所有 embedding 都是语义”。

flowchart TD P[未来像素] --> Q[像素重建: 易学到背景/光照] P --> R[冻结 V-JEPA2] --> S[语义世界状态] --> T[latent 对齐: 更关注交互]

防泄漏原则:未来帧只能构造监督目标,不能进入产生 latent action 的 VLM。

flowchart LR O[当前图像+语言] --> V[Qwen3-VL] --> Z[latent action] F[未来帧] --> E[冻结 V-JEPA2] --> Y[target state] Z --> P[predictor] --> YH[predicted state] Y & YH --> L[world loss] F -.不能进入 VLM.-> O

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

2.1 JEPA 是什么?

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

先看三种学习目标的区别:

flowchart TB X[被遮挡/目标内容] A[像素生成模型] --> A1[预测每个 RGB 像素] C[JEPA] --> C1[预测目标内容的 embedding] V1[两个增强视图/样本对] --> B[对比学习] --> B1[拉近正对, 推远其他对] X --> A X --> C

例如,看到“一只手正推向杯子”,若 target encoder 的表示保留物体运动而弱化纹理,JEPA 不需要画出下一帧中杯子的每个像素,而要预测“杯子将向右移动”对应的特征。背景纹理变化的惩罚会较小;运动方向错误则会被更明显地惩罚。

2.2 JEPA 的四个基本组件

flowchart LR C[可见上下文 x] --> CE[Context Encoder] CE --> CX[上下文表示 h_x] T[目标区域/未来片段 y] --> TE[Target Encoder] TE --> TY[目标表示 h_y] CX --> P[Predictor] 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 是要预测的目标 token;m 描述目标的位置或时间;
  • f_theta 是 context encoder,f_bar_theta 是 target encoder;
  • g_phi 是 predictor;
  • sg 表示 stop-gradient,目标分支不接收当前损失的梯度;
  • d 可以是 L1、L2、cosine distance 等。

上式是便于理解的抽象写法。在 I-JEPA/V-JEPA 中,target encoder 常常先编码完整、未遮挡的图像或视频,再按 target mask 取出目标 token;不一定是将裁出的 y 单独送入 target encoder。

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 participant L as Loss/Optimizer D->>M: 采样上下文 x 与隐藏目标 y M->>C: 只给可见上下文 x M->>T: 给完整样本并选择目标 token C->>P: 上下文表示 h_x M->>P: 目标位置 m P-->>L: 预测表示 h_y_hat T-->>L: stop-gradient 目标 h_y L-->>C: 更新 context encoder 和 predictor

伪代码:

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 = select_target_tokens(target_encoder(sample), target_position)

    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 CE[Context Encoder theta] -->|正常反向传播| U[梯度更新] CE -->|EMA 参数更新| TE[Target Encoder bar theta] TE -.stop-gradient.-> L[Loss] P[Predictor phi] -->|正常反向传播| U

EMA 的完整更新公式放在图外书写,避免 Mermaid 将箭头符号误解析:

\[\bar\theta \leftarrow \tau\bar\theta+(1-\tau)\theta. \]

与当前 VLA-JEPA 代码的区别:这里没有从零训练一对 JEPA encoder。项目直接加载已经预训练好的 V-JEPA2 encoder,并将它冻结;world loss 更新 Qwen3-VL 的 latent-action 路径和 world predictor,机器人 batch 的 action loss 还会更新动作头。冻结 V-JEPA2 既不接收梯度,也不做 EMA,因此 target space 不会随这次训练漂移。若该预训练表示本身非坍塌且随输入变化,常量预测无法把对齐损失降到零。

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

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

可以把它们记成:

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

JEPA 的优势是不用显式训练像素解码器去拟合所有细节;代价是表示最终保留什么,取决于 target encoder 和训练任务是否保留了任务所需信息。

2.6 I-JEPA 的完整训练过程

I-JEPA 的训练样本只有一张图像。它通过遮住同一张图中的若干区域,构造“根据可见区域预测隐藏区域表示”的自监督任务。

2.6.1 第一步:图像切成 patch,并采样 context 与 target

设输入图像为 I。Vision Transformer 先把它划分成规则 patch。Mask sampler 从这些 patch 中采样:

  • 一个较大的 context 区域;
  • 一个或多个连续的 target block;
  • 从 context 中删除与 target 重叠的 patch。
flowchart LR I[一张完整图像] --> P[切分为图像 patch] P --> M[Mask Sampler] M --> C[可见 context patch] M --> T1[目标块一] M --> T2[目标块二] C -.不包含目标内容.-> T1

这里的 target 不是类别标签,而是图像中被隐藏区域的位置集合。每次训练都可重新采样,因此同一张图能产生许多不同的预测任务。

2.6.2 第二步:两条 encoder 路径

Context Encoder 只看到可见 patch;Target Encoder 通常编码完整图像,然后按 target mask 取出目标位置的 token。

flowchart TB I[完整图像 I] --> M[Context 和 Target masks] M --> VC[只保留可见 patch] VC --> CE[Context Encoder] CE --> HC[context representations] I --> TE[Target Encoder] TE --> HA[完整图像 representations] M --> SEL[按 target mask 选择] HA --> SEL SEL --> HT[target representations]

这不构成信息泄漏,因为 HT 只作为 loss 的标签,不会作为 predictor 的输入。Target Encoder 也不接收本次 loss 的梯度。

数学表示为:

\[H_c=f_\theta(I_{context}), \]

\[H_t=\operatorname{Select}\left(f_{\bar\theta}(I),M_{target}\right). \]

2.6.3 第三步:Predictor 根据什么预测?

Predictor 的输入只有:

  1. 可见区域的 context representations;
  2. target block 的二维位置编码;
  3. 放在 target 位置上的可学习 mask token。

位置编码只告诉模型“需要预测哪里”,不会告诉它“那里是什么”。

flowchart LR HC[可见区域 representations] --> PR[Predictor] MP[Target 位置编码] --> PR MT[可学习 mask tokens] --> PR PR --> PH[预测 target representations] HT[Target Encoder 标签] --> L[Latent Loss] PH --> L

例如,目标区域位于猫身体上方。Predictor 可以看到猫的身体、耳朵边缘和沙发,并知道目标位置,于是尝试预测该位置应具有“猫头相关”的表示;它不需要重建每根毛发的 RGB 值。

2.6.4 第四步:计算 latent loss

对所有 target token 比较预测表示与目标表示:

\[\mathcal L_{I\text{-}JEPA}=\frac{1}{N_t}\sum_{i=1}^{N_t}d\left(\hat h_i,h_i^{target}\right). \]

其中 d 是表示距离,不同实现可以使用 L1、Smooth L1 或 L2。核心要求是比较表示,不是比较像素。

2.6.5 第五步:参数更新

sequenceDiagram participant I as 图像 participant C as Context Encoder participant T as Target Encoder participant P as Predictor participant L as Latent Loss I->>C: 仅可见 context patch I->>T: 完整图像 C->>P: context representations P->>L: predicted target representations T->>L: stop-gradient target representations L-->>C: 梯度更新 L-->>P: 梯度更新 C-->>T: EMA 参数更新
  • Context Encoder:正常反向传播;
  • Predictor:正常反向传播;
  • Target Encoder:不反向传播,通过 Context Encoder 参数的 EMA 缓慢更新。

\[\bar\theta \leftarrow \tau\bar\theta+(1-\tau)\theta. \]

I-JEPA 完整伪代码:

def train_ijepa(image):
    # 1. 从同一张图采样可见区域与多个目标块
    context_mask, target_masks = sample_image_blocks(image)

    # 2. 学生分支只能看到 context,不能看到 target 内容
    context_repr = context_encoder(
        image,
        visible_mask=context_mask,
    )

    with no_grad():
        # 3. 教师分支编码完整图像,再取目标位置作为标签
        full_target_repr = target_encoder(image)
        target_repr = select_tokens(full_target_repr, target_masks)

    # 4. Predictor 只拿 context 特征和目标位置
    predicted_repr = predictor(
        context_repr,
        target_positions=target_masks,
    )

    # 5. 在表示空间对齐,而不是重建 RGB
    loss = latent_regression_loss(predicted_repr, target_repr)
    update_by_gradient(context_encoder, predictor, loss)

    # 6. Target Encoder 通过 EMA 缓慢跟随
    ema_update(target_encoder, context_encoder)

训练完成后,一般保留 Context Encoder 作为图像表征模型;Target Encoder 和 Predictor 主要用于自监督预训练。

2.7 V-JEPA 的完整训练过程

V-JEPA 把 I-JEPA 从二维图像扩展到视频。核心机制不变,但 patch 变成了同时具有时间和空间位置的 tubelet / spatiotemporal token

2.7.1 第一步:视频切成时空 token

输入视频可写成 V=[I_0,I_1,...,I_{T-1}]。视频编码器把连续若干帧和一个空间 patch 组合成 tubelet。

flowchart LR V[一段 T 帧视频] --> TB[Tubelet Embedding] TB --> ST[时间 x 高度 x 宽度 的 token 网格] ST --> MS[时空 Mask Sampler] MS --> C[可见时空 context] MS --> TG[隐藏时空 target blocks]

Target block 可以覆盖连续空间区域和多个时间步。标准 V-JEPA 预测的是被遮挡时空块,它不要求 target 一定全在未来;只用过去预测未来是 VLA-JEPA 采用的更强时序约束。

2.7.2 第二步:Context 与 Target 视频路径

flowchart TB V[完整视频] --> M[时空 masks] M --> VC[只保留可见 tubelets] VC --> CE[Video Context Encoder] CE --> HC[context video representations] V --> TE[Video Target Encoder] TE --> HF[完整视频 representations] M --> SEL[选择 target 时空位置] HF --> SEL --> HT[target video representations]

与 I-JEPA 一样,Target Encoder 可以看到完整视频,但它的输出只用于监督;Context Encoder 和 Predictor 看不到 target 内容。

2.7.3 第三步:加入时空位置进行预测

Predictor 接收可见视频 representations,以及 target tubelet 的时间、高度、宽度位置编码:

\[\hat H_t=g_\phi\left(H_c,M_{time},M_{height},M_{width}\right). \]

flowchart LR HC[可见视频 representations] --> P[Video Predictor] TP[时间位置编码] --> P SP[空间位置编码] --> P MT[Target mask tokens] --> P P --> PH[预测隐藏时空 representations] HT[Target Encoder 标签] --> L[Video latent loss] PH --> L

时间位置很重要:同一个空间位置在 t0t5 可能对应完全不同的运动阶段。Predictor 必须结合可见帧中的物体、运动方向和时序关系预测目标 token。

2.7.4 第四步:视频 latent loss 与 EMA

\[\mathcal L_{V\text{-}JEPA}=\frac{1}{N_{st}}\sum_{i=1}^{N_{st}}d\left(\hat h_i,h_i^{target}\right). \]

反向传播仍只更新 Video Context Encoder 和 Predictor,Video Target Encoder 通过 EMA 更新。

sequenceDiagram participant V as 视频 participant C as Video Context Encoder participant T as Video Target Encoder participant P as Video Predictor participant L as Latent Loss V->>C: 可见 tubelets V->>T: 完整视频 C->>P: context video representations P->>L: predicted target representations T->>L: stop-gradient target representations L-->>C: 梯度更新 L-->>P: 梯度更新 C-->>T: EMA 参数更新

V-JEPA 完整伪代码:

def train_vjepa(video):
    # 1. 在时间和空间两个维度采样隐藏块
    context_mask, target_masks = sample_spatiotemporal_blocks(video)

    # 2. Context Encoder 只处理未被遮挡的 tubelets
    context_repr = video_context_encoder(
        video,
        visible_mask=context_mask,
    )

    with no_grad():
        # 3. Target Encoder 编码完整视频,目标分支不反向传播
        full_target_repr = video_target_encoder(video)
        target_repr = select_tokens(full_target_repr, target_masks)

    # 4. 时空位置告诉 Predictor 预测哪些 tubelets
    predicted_repr = video_predictor(
        context_repr,
        target_spatiotemporal_positions=target_masks,
    )

    # 5. 只对齐视频 latent,不生成未来像素
    loss = latent_regression_loss(predicted_repr, target_repr)
    update_by_gradient(video_context_encoder, video_predictor, loss)
    ema_update(video_target_encoder, video_context_encoder)

通过预测跨帧的隐藏表示,V-JEPA 被迫学习物体身份、运动方向、遮挡关系和动作阶段等时间信息;但这些性质仍取决于数据、mask 策略和 target representation。

2.8 I-JEPA 与 V-JEPA 对照

对比项 I-JEPA V-JEPA
输入 单张图像 多帧视频
基础 token 2D image patch 3D spatiotemporal tubelet
Target 被遮挡空间块 被遮挡时空块
Predictor 位置信息 高度、宽度 时间、高度、宽度
主要学习内容 物体与空间结构 物体、空间结构与运动变化
预测目标 target image embedding target video embedding
是否重建像素
Target Encoder 更新 EMA EMA
flowchart LR I[单张图像] --> IJ[I-JEPA] IJ --> IS[空间 target representations] V[多帧视频] --> VJ[V-JEPA] VJ --> VS[时空 target representations] IS --> J[共同点: latent prediction] VS --> J

最重要的连续关系是:

I-JEPA: 可见空间区域 + 目标空间位置 -> 目标图像表示
V-JEPA: 可见时空区域 + 目标时空位置 -> 目标视频表示

2.9 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] --> VP[Predictor] --> VF[隐藏目标 latent] end subgraph VLA[VLA-JEPA] SC[历史 world states] --> WP[World Predictor] ZA[Qwen 根据当前图像+语言产生 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,应该足以帮助世界模型预测未来语义状态。 单靠这条约束,它只保证 latent action 对状态转移有用,不保证它是唯一可辨识的真实机器人动作;机器人 batch 的 flow-matching 动作损失进一步把这套表征对齐到可执行的 7D 控制轨迹。

2.10 用一个最小数值例子理解 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] --> L[L1=0.25] B[目标 0.8,0.2] --> L L --> P[更新 world predictor] L --> Q[更新 latent-action 产生路径] L -.不更新.-> V[V-JEPA2]

到这里,可以把 JEPA 记为一句话:不要求模型复原目标长什么样,而要求它在目标表示空间里预测目标是什么。 在 VLA-JEPA 这个具体方法中,目标恰好是未来世界状态,因此可以把“目标”理解成“未来”。

3. 数据

3.1 两类数据

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

统一样本接口:

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[复制/选择两路视角] --> W[world video] S --> I[第0帧缩放224x224] --> Q[VLM image] L[CSV文本] --> Q

3.3 机器人数据

LeRobot 数据由 modality.json 和 robot type 配置对齐相机、状态、动作字段。论文说明连续位置/轴角动作 min-max 到 [0,1],夹爪二值化为 {0,1}

flowchart TD E[LeRobot episode] --> I[当前多视角图像] E --> V[连续视频片段] E --> A[未来7步动作, 每步7维] E --> S[可选8维状态] I & V & A & S --> B[统一样本字典]

4. 模型总览

flowchart TB I[当前图像] & L[指令] --> Q[Qwen3-VL-2B] Q --> Z[latent token hidden states] Q --> ZA[embodied token hidden states] V[未来视频] --> F[V-JEPA2 frozen encoder] --> Y[目标 world states] Z --> P[action-conditioned predictor] --> YH[预测未来 states] Y & YH --> LW[L_WM] ZA & S[可选 state] & A[真实动作] --> H[Flow Matching DiT] --> LA[L_FM]

4.1 Qwen3-VL 与特殊 token

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

数学上:

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视频] --> F1[Frozen V-JEPA2] V2[视角2视频] --> F2[Frozen V-JEPA2] F1 & F2 --> C[embedding维拼接] --> 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] z0[latent] <--> x0[state patches] end subgraph t1[时间 t1] z1[latent] <--> x1[state patches] end subgraph t2[时间 t2] z2[latent] <--> x2[state patches] end t0 --> t1 --> t2 t1 -.禁止看 t2.-> t2

伪代码:

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)。当前实现明确使用:

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[动作+时间编码] ZA[z_a] & S[state] --> D[条件 DiT] AE --> D --> V[预测速度] V & A & E --> L[速度 MSE]

伪代码:

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[随机动作] --> D1[DiT预测速度] --> U1[Euler更新] --> D2[DiT预测速度] --> U2[Euler更新] U2 --> D3[DiT预测速度] --> U3[Euler更新] --> D4[DiT预测速度] --> O[输出7步动作]

7. 训练 pipeline

7.1 人类视频 batch

无动作标签,所以只产生 wm_loss

flowchart LR B[视频batch] --> Q[Qwen: 当前帧+语言] --> Z[latent] B --> F[V-JEPA2: 未来帧] --> Y[target state] Z & F --> P[predictor] --> YH[prediction] Y & YH --> L[wm_loss]

7.2 机器人 batch

flowchart TB R[机器人batch] --> Q[Qwen3-VL] Q --> Z[latent] --> W[world predictor] Q --> ZA[embodied tokens] --> H[Flow DiT] R --> F[V-JEPA2 target] --> W W --> LW[wm_loss] R --> H --> LA[action_loss] LW & LA --> Total[total loss]

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

联合训练 step:

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 --> D[latent world model] C --> E[flow action head] D --> F[时序理解/鲁棒性] E --> G[机器人执行]

只记住三点:

  1. 未来帧只做 target,不进入 VLM
  2. latent action 通过预测未来 V-JEPA 状态获得语义
  3. 真正的机器人控制由 embodied token + flow-matching DiT 生成
posted @ 2026-07-22 15:25  S-X-Q  阅读(37)  评论(0)    收藏  举报