mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-24 10:16:09 +00:00
2e9cd87bbd
* first commit * feat(policies): add VLA-JEPA * feat(policies): add VLA-JEPA * support vla_jepa * (feat)policies: add VLA-JEPA * linting * adding deps to pyproject.toml * updating uv lock * adding guards to avoid needing transformers and diffusers for type checking and basic tests * fixing action and state dim * fix warnings with qwen processor kwargs * fixing wm_loss not propagating * adjusting obs steps, tublets size to match original implementation * some more fixes to be closer to the original implem * adding more tests to ensure good coverage * align VLA-JEPA architecture with original checkpoint - Remove stale `action_num_heads` / `action_attention_head_dim` config fields; DiT head dimensions are now always derived from the preset (DiT-B/L/test). - Add `num_target_vision_tokens` and `action_max_seq_len` config fields required by the action head's future-token embedding and positional embedding tables. - Fix default `qwen_model_name` to 2B (matches all released checkpoints). - Rename `ActionEncoder` attrs w1/w2/w3 → layer1/layer2/layer3 to match checkpoint key names; replace `nn.Sequential` decoder/state-encoder with `_MLP2` (layer1/layer2 naming). - Fix `VLAJEPAActionHead` to size ActionEncoder and StateEncoder at `inner_dim` (DiT input width) rather than `action_hidden_size` (DiT output width). - Rename `DiT.blocks` → `transformer_blocks` and `attn` → `attn1` to match checkpoint; add alternating cross/self attention (even blocks cross-attend to Qwen context, odd blocks self-attend). - Add `DiT-test` preset for unit tests. - Rewrite `ActionConditionedVideoPredictor` with explicit ViT-style blocks (`_PredictorBlock` with fused qkv) to match checkpoint structure; rename `encoder`/`norm`/`proj` → `predictor_blocks`/`predictor_norm`/`predictor_proj`. * propagate action_is_pad masking through VLA-JEPA policy pipeline Pass the `action_is_pad` tensor from the batch through to the action head so padded timesteps are excluded from the flow-matching loss. * update VLA-JEPA tests for arch changes and action_is_pad - Switch conftest to use `action_model_type="DiT-test"` now that `action_num_heads` / `action_attention_head_dim` have been removed. - Add action_head tests covering fully-padded loss (zero) and equivalence of action_is_pad=None vs all-zeros mask. - Remove obsolete `test_native_to_lerobot_wm_only` test. * add VLA-JEPA documentation Covers architecture overview, pretrained checkpoints, config reference, training/eval commands for LIBERO-10, and guidance on fine-tuning for single-camera datasets. * add one-shot script to convert ginwind/VLA-JEPA checkpoints to safetensors (will remove once migrated) * make default params more aligned with paper and pretrained models - adding possibility of freezing qwen backbone and world model - added tests for weight loading * trying out to re-init the action head to avoid pretraining dimension mismatch * allow different state dim and action dim * removing missleading future_action_window_size to just use chunk_size * lots of changes to make existing weights work, need to massively refactor the pre and post processing * refactoring into using pre and post processor * pre-commit cleanup * fixing doc defaults args Signed-off-by: Maxime Ellerbach <maxime@ellerbach.net> * adressing dtype zeros issue * adding guard for diffusers * fixing training and exal examples * trying to close success rate gap * fix qwen norm layer output libero eval is now as expected * adding instructions for different embodiement + fixing some tests * smol fix to avoid having default CPU device when training * fixing misconception about multiview / singleview handling * removing conversion script * adding licences * adding .mdx docs and shortening polivy_vla_jepa_README.md * removing useless pre-processor * cleanup * removing swish in favor of silu * adding configuration gripper index and threshold * fixing simlink --------- Signed-off-by: Maxime Ellerbach <maxime@ellerbach.net> Co-authored-by: ginwind <ginwind@mail.ustc.edu.cn>
118 lines
4.7 KiB
Python
118 lines
4.7 KiB
Python
# Copyright 2026 The HuggingFace Inc. team. All rights reserved.
|
|
#
|
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
|
# you may not use this file except in compliance with the License.
|
|
# You may obtain a copy of the License at
|
|
#
|
|
# http://www.apache.org/licenses/LICENSE-2.0
|
|
#
|
|
# Unless required by applicable law or agreed to in writing, software
|
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
# See the License for the specific language governing permissions and
|
|
# limitations under the License.
|
|
|
|
from __future__ import annotations
|
|
|
|
from collections.abc import Sequence
|
|
from typing import TYPE_CHECKING
|
|
|
|
import numpy as np
|
|
import torch
|
|
from PIL import Image
|
|
|
|
from lerobot.utils.import_utils import _transformers_available
|
|
|
|
if TYPE_CHECKING or _transformers_available:
|
|
from transformers import AutoProcessor, Qwen3VLForConditionalGeneration
|
|
else:
|
|
AutoProcessor = None
|
|
Qwen3VLForConditionalGeneration = None
|
|
|
|
from .configuration_vla_jepa import VLAJEPAConfig
|
|
|
|
|
|
class Qwen3VLInterface(torch.nn.Module):
|
|
def __init__(self, config: VLAJEPAConfig) -> None:
|
|
super().__init__()
|
|
self.config = config
|
|
self.model = Qwen3VLForConditionalGeneration.from_pretrained(
|
|
config.qwen_model_name,
|
|
torch_dtype=self._get_torch_dtype(config.torch_dtype),
|
|
)
|
|
self.processor = AutoProcessor.from_pretrained(config.qwen_model_name)
|
|
self.processor.tokenizer.padding_side = config.tokenizer_padding_side
|
|
self.model.config.hidden_size = self.model.config.text_config.hidden_size
|
|
|
|
@staticmethod
|
|
def _get_torch_dtype(dtype_name: str) -> torch.dtype:
|
|
if dtype_name == "float32":
|
|
return torch.float32
|
|
if dtype_name == "float16":
|
|
return torch.float16
|
|
return torch.bfloat16
|
|
|
|
def expand_tokenizer(self) -> tuple[list[str], list[int], int]:
|
|
# starVLA/JEVLA checkpoints expand action tokens as action_horizon * 4,
|
|
# independent of vj2 num_action_tokens_per_timestep. Keeping this count
|
|
# is required for Qwen embedding/lm_head checkpoint shapes to match.
|
|
max_action_tokens = self.config.chunk_size * 4
|
|
tokenizer = self.processor.tokenizer
|
|
action_tokens = []
|
|
action_token_ids = []
|
|
for idx in range(max_action_tokens):
|
|
token = self.config.special_action_token.format(idx)
|
|
action_tokens.append(token)
|
|
if token not in tokenizer.get_vocab():
|
|
tokenizer.add_tokens([token], special_tokens=True)
|
|
action_token_ids.append(tokenizer.convert_tokens_to_ids(token))
|
|
|
|
embodied_action_token = self.config.embodied_action_token
|
|
if embodied_action_token not in tokenizer.get_vocab():
|
|
tokenizer.add_tokens([embodied_action_token], special_tokens=True)
|
|
embodied_action_token_id = tokenizer.convert_tokens_to_ids(embodied_action_token)
|
|
|
|
if self.model.get_input_embeddings().weight.size(0) < len(tokenizer):
|
|
self.model.resize_token_embeddings(len(tokenizer))
|
|
return action_tokens, action_token_ids, embodied_action_token_id
|
|
|
|
def build_inputs(
|
|
self,
|
|
images: Sequence[Sequence[Image.Image]],
|
|
instructions: Sequence[str],
|
|
action_prompt: str,
|
|
embodied_prompt: str,
|
|
) -> dict[str, torch.Tensor]:
|
|
messages = []
|
|
for sample_images, instruction in zip(images, instructions, strict=True):
|
|
prompt = self.config.prompt_template.format(
|
|
instruction=instruction,
|
|
actions=action_prompt,
|
|
e_actions=embodied_prompt,
|
|
)
|
|
content = [{"type": "image", "image": img} for img in sample_images]
|
|
content.append({"type": "text", "text": prompt})
|
|
messages.append([{"role": "user", "content": content}])
|
|
|
|
batch_inputs = self.processor.apply_chat_template(
|
|
messages,
|
|
tokenize=True,
|
|
add_generation_prompt=True,
|
|
return_dict=True,
|
|
processor_kwargs={"padding": True, "return_tensors": "pt"},
|
|
)
|
|
return batch_inputs.to(self.model.device)
|
|
|
|
@staticmethod
|
|
def tensor_to_pil(image_tensor: torch.Tensor) -> Image.Image:
|
|
image = image_tensor.detach().cpu()
|
|
if image.ndim == 3 and image.shape[0] in (1, 3):
|
|
image = image.permute(1, 2, 0)
|
|
image = image.float()
|
|
if image.max() <= 1.0:
|
|
image = image * 255.0
|
|
image = image.clamp(0, 255).round().to(torch.uint8).numpy()
|
|
if image.shape[-1] == 1:
|
|
image = np.repeat(image, 3, axis=-1)
|
|
return Image.fromarray(image)
|