VLA-JEPA

这是 VLA-JEPA 的 LeRobot 移植版。VLA-JEPA 是一个视觉-语言-action(Vision-Language-Action)模型,它将 Qwen3-VL 语言主干、自监督视频世界模型(V-JEPA2)以及基于流匹配的 DiT action 头相结合。


架构概览

VLA-JEPA 有三个主要组成部分:

组件模块作用
Qwen3-VL 主干Qwen3VLInterface将图像与语言指令融合为上下文 token
DiT-B action 头VLAJEPAActionHead在 action chunk 上进行流匹配扩散
V-JEPA2 世界模型ActionConditionedVideoPredictor自监督视频预测损失(仅训练时)

数据流

训练:

  1. 一个由 num_video_frames 帧组成的视频片段由 V-JEPA2 编码为逐帧的 patch token。
  2. Qwen3-VL 主干处理多视角图像与任务指令,并生成一系列上下文 token,其中包含特殊的 action token(用于世界模型条件)和具身 token。
  3. action 头将这些上下文 token 作为交叉注意力的键/值,并通过流匹配预测去噪后的 action chunk。
  4. 世界模型预测器使用从 Qwen 提取的 action token 来预测未来的 V-JEPA2 帧嵌入;这些预测的回归损失会加到 action 损失上。

inference: 仅使用 Qwen 和 action 头。inference 时不需要世界模型。

action 头详解

可通过 action_model_type 使用的预设:

预设头数头维度
DiT-B1264
DiT-L3248

预设仅设置注意力几何形状,每个条目都可以通过 action_num_heads / action_attention_head_dim 覆盖。由此得到两个宽度:

  • DiT 的内部宽度为 heads x head_dimDiT-B 为 768),是推导而来而非配置所得;
  • DiT 的输出宽度,以及 action 解码器和 state 编码器 MLP 的宽度,为 action_hidden_size(默认 1024)。

因此 DiT-B 运行一个宽度为 768 的 transformer,并将其投影到 1024。二者相互独立。

世界模型详解

视频预测器是一个 ViT 风格的 transformer(ActionConditionedVideoPredictor),其输入包括:

  • 帧 token:V-JEPA2 patch 嵌入,投影到 predictor_embed_dim
  • action token:Qwen action token 嵌入,投影到 predictor_embed_dim

它使用块因果注意力,使每个时间步都能关注所有之前的步骤。预测器的输入 embed_dim 等于 num_views × video_encoder_hidden_size(例如预训练 checkpoint 为 2 视角 × 1024 = 2048)。


预训练 checkpoint

LeRobot 组织内直接提供了三个 checkpoint:lerobot/VLA-JEPA,由 ginwind/VLA-JEPA 转换而来:

checkpointdataset相机世界模型action 维度
lerobot/VLA-JEPA-LIBEROLIBERO-102(agentview + wrist)启用7
lerobot/VLA-JEPA-PretrainDROID 1.0.12(外部左侧视角)启用7
lerobot/VLA-JEPA-SimplerEnvOXE Bridge / RT-11(视角复制 ×2)启用7

所有 checkpoint 都使用 Qwen/Qwen3-VL-2B-Instruct 作为语言主干。


配置

VLAJEPAConfig 中的关键参数:

参数默认值描述
chunk_size7每次 inference 调用预测的 action 数
n_action_steps7重新规划前从预测块中执行的步数
num_video_frames8输入世界模型的视频片段长度
enable_world_modelTrue是否加载并训练 V-JEPA2 预测器
world_model_loss_weight0.1JEPA 预测损失相对于 action 损失的权重
causal_world_model_contextFalse以因果方式编码世界模型上下文(每个上下文位置执行一次仅前缀的 V-JEPA2 前向传播),使双向注意力无法将未来帧泄漏到预测器输入中。代价是额外的 t_enc_ctx 次编码器调用
num_inference_timesteps4action 去噪的欧拉积分步数
freeze_qwenFalse冻结 Qwen3-VL 主干,只训练 action 头
reinit_modulesNone允许在加载时随机重新初始化的键前缀(用于跨具身迁移,参见在不同具身上 fine-tune
resize_images_toNone每个相机帧在进入 Qwen3-VL 视觉塔之前都会被调整到 (height, width) 大小。None 保留原始分辨率,而 Qwen3-VL 的 patch 数量会随之增长,因此 720x1280 的相机可能耗尽 GPU 内存。已发布的 checkpoint 使用 [224, 224]
gripper_dim6action 向量中 gripper 维度的索引。当 gripper_joint_names 与 dataset 的某个 action 名称匹配时忽略
gripper_joint_names["gripper"]识别 gripper 的 action 维度名称;匹配到的索引优先于 gripper_dim
gripper_threshold0.5pre_snap_gripper_actionbinarize_gripper_action 使用的阈值。注意 binarize 在反归一化之后运行,因此它比较的是 gripper 的物理值
pre_snap_gripper_actionFalse在反归一化之前将 gripper 维度吸附到 {0, 1}。LIBERO 专属设置,见下文
binarize_gripper_actionFalse在反归一化之后将 gripper 维度二值化为 {-1, 1}。LIBERO 专属设置,见下文
clip_normalized_actionsTrue在反归一化之前将归一化 action 裁剪到 [-1, 1]。仅在 ACTION 使用 MIN_MAX 时应用;在 MEAN_STD 下会被忽略(并给出警告),因为在那里会在 1 个标准差处截断
world_model_num_viewsNone世界模型预测器所针对的相机视角数。已固化在 checkpoint 形状中。None 时回退到 jepa_tubelet_size,已发布的 checkpoint 编码的正是该值

pre_snap_gripper_actionbinarize_gripper_action 移植自 starVLA 的 LIBERO 评估 循环,仅对 LIBERO 的 action 约定成立。pre_snap 将 {0, 1} 写入 归一化空间,反归一化器将其映射到中点和最大值,binarize 随后将 该物理值与 gripper_threshold(0.5)进行比较。对于以度、 毫米或 [0, 100] 为单位的 gripper,两个值都落在阈值之上,导致命令的 gripper 变成一个 常量。因此它们默认设为 False;仅在 LIBERO 风格的设置中启用它们,并且 如果启用,请以 gripper 自身的单位设置 gripper_threshold。当 dataset 统计显示该范围不可行时,处理器工厂会发出警告。


训练

训练步数可能因 dataset 大小和计算预算而异。原论文在 ssv2 + droid 上联合预训练了 50k 步,随后针对 LIBERO 额外训练了 30k 步;但从提供的预训练 checkpoint fine-tune 时,较少的步数仍可能获得良好的性能。

从零开始完整训练

lerobot-train \
  policy.type=vla_jepa \
  policy.repo_id=your_org/your_repo \
  dataset.repo_id=your_org/your_dataset

从预训练 checkpoint fine-tune

lerobot-train \
  --policy.path=lerobot/VLA-JEPA-Pretrain \
  --policy.repo_id=your_org/your_repo \
  --dataset.repo_id=your_org/your_dataset

如果你想冻结 Qwen 主干,只训练 action 头,请设置 policy.freeze_qwen=True

lerobot-train \
  --policy.path=lerobot/VLA-JEPA-Pretrain \
  --policy.repo_id=your_org/your_repo \
  --policy.freeze_qwen=true \
  --dataset.repo_id=your_org/your_dataset

在不同具身上 fine-tune

当目标机器人与预训练 checkpoint 具有不同的 action 或 state 维度时,action 头的输入/输出投影层会出现形状不匹配,无法直接加载。reinit_modules 允许你列出可以允许不匹配的键前缀——这些层会被随机重新初始化,而所有其他权重则从 checkpoint 复用。列出的前缀之外的任何形状不匹配都会引发错误。

依赖于 action_dimstate_dim 的层有:

键前缀
action 编码器(action_dim → inner_dim)model.action_model.action_encoder
action 解码器(hidden_size → action_dim)model.action_model.action_decoder
state 编码器(state_dim → inner_dim)model.action_model.state_encoder
lerobot-train \
  --policy.path=lerobot/VLA-JEPA-Pretrain \
  --policy.repo_id=your_org/your_repo \
  --policy.freeze_qwen=true \
  --policy.reinit_modules='["model.action_model.action_encoder", "model.action_model.action_decoder", "model.action_model.state_encoder"]' \
  --dataset.repo_id=your_org/your_dataset

如果你的机器人没有本体感觉 state,请从列表中省略 model.action_model.state_encoder

复现 LIBERO 结果

在 LIBERO 上训练: 从 Pretrain checkpoint 开始训练,在 LIBERO dataset 上训练 30k 步。 原论文提到在 8 个 GPU 上以 32 的批次大小训练,即全局批次大小为 256。

lerobot-train \
  --policy.path=lerobot/VLA-JEPA-Pretrain \
  --policy.repo_id=your_org/your_repo \
  --dataset.repo_id=HuggingFaceVLA/libero \
  --steps=30000

评估预训练的 LIBERO-10 checkpoint:

lerobot-eval \
  --policy.path=lerobot/VLA-JEPA-LIBERO \
  --env.type=libero \
  --env.task=libero_spatial,libero_object,libero_goal,libero_10 \
  --eval.n_episodes=10 \
  --eval.batch_size=5

只评估任务子集:

lerobot-eval \
  --policy.path=lerobot/VLA-JEPA-LIBERO \
  --env.type=libero \
  --env.task=libero_10 \
  --env.task_ids='[0,1,2]' \
  --eval.n_episodes=10 \
  --eval.batch_size=5

预期结果:

测试套件episode 数成功次数成功率
libero_spatial1009395.0%
libero_object100100100.0%
libero_goal1009898.0%
libero_101009693.0%
总计40038796.5%

在相机数量不同的 dataset 上 fine-tune

预训练的世界模型预测器使用 embed_dim = world_model_num_views × 1024 训练,即两个相机视角。

这个视角数此前是从 jepa_tubelet_size 读取的,而该字段还命名了 JEPA 编码器的时间 tubelet 大小。现在 world_model_num_views 是这个视角数的字段;将其保留为 None 会回退到 jepa_tubelet_size,因此已发布的 checkpoint 可以保持不变地加载。

默认行为——视角填充 / 裁剪(无需任何操作)

VLA-JEPA-Pretrain fine-tune 时,模型会自动调整输入世界模型的视角数,以匹配 world_model_num_views

  • 单视角 dataset(例如 BridgeV2): 单视角隐向量会被复制以产生双视角的世界模型输入,从而在没有任何权重不匹配的情况下保留 JEPA 自监督信号。
  • 超过 2 个视角的 dataset(例如有 3 个视角的 DROID): 所有视角都会传入 Qwen 主干(以获得更丰富的上下文),但只有前 world_model_num_views 个视角(按照配置的视角顺序取一个手腕视角和一个第三人称视角)用于世界模型。

选项 1——禁用世界模型

设置 enable_world_model=False 以完全跳过 JEPA 损失。只加载并训练 Qwen 主干和 action 头。这足以获得良好的 action 性能。

lerobot-train \
  --policy.path=lerobot/VLA-JEPA-Pretrain \
  --policy.enable_world_model=false \
  --policy.repo_id=your_org/your_repo \
  --dataset.repo_id=your_org/single_camera_dataset

选项 2——重新初始化预测器输入投影

如果你想将 world_model_num_views 改为 2 以外的值,请使用 strict=False 加载 checkpoint,并针对新的 embed_dim 重新初始化 model.video_predictor.predictor_embed。所有其他预测器块权重(注意力、MLP、归一化、输出投影)都与相机数量无关,可以复用预训练 checkpoint 中的权重。


引用

@misc{sun2026vlajepaenhancingvisionlanguageactionmodel,
  title         = {VLA-JEPA: Enhancing Vision-Language-Action Model with Latent World Model},
  author        = {Jingwen Sun and Wenyao Zhang and Zekun Qi and Shaojie Ren and Zezhi Liu and Hanxin Zhu and Guangzhong Sun and Xin Jin and Zhibo Chen},
  year          = {2026},
  eprint        = {2602.10098},
  archivePrefix = {arXiv},
  primaryClass  = {cs.RO},
  url           = {https://arxiv.org/abs/2602.10098},
}

许可证

权重按照原始 ginwind/VLA-JEPA 仓库的许可条款分发(Apache 2.0 许可证)。LeRobot 集成代码遵循 Apache 2.0 许可证

在 GitHub 上更新