mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-20 16:31:55 +00:00
4dfa8cea65
Port the LingBot-VA policy (Wan2.2 dual-stream video+action world model) into LeRobot, following the EO-1 / VLA-JEPA conventions. Covers inference, checkpoint conversion, and predicted-video saving (training is deferred to a follow-up PR). - Vendored Wan transformer/attention/flex/VAE/scheduler modules (key names preserved for near-identity conversion); torch SDPA default, flashattn/flex lazy-guarded. - LingBotVAConfig (registered "lingbot_va") + processor with fixed-quantile action unnormalization; full dual-stream sampling loop with CFG, two flow-matching schedulers and KV cache, mapped onto select_action with observed-keyframe feedback. - convert_lingbot_va_checkpoints.py (libero/robotwin variants): bundles the ~5B transformer, lazy-pulls the frozen VAE+UMT5 from the source repo. - Predicted-video plumbing in lerobot_eval (predicted_frames_callback; opt-in via --policy.save_predicted_video) and ConstantWithWarmupSchedulerConfig. - pyproject: widen diffusers-dep to <0.37, add lingbot_va + imageio-dep extras, add lingbot_va and (missing) eo1 to `all`. - Factory + policies/__init__ wiring, docs page + toctree, and tests. Note: the LIBERO success-rate correctness gate must be validated on a CUDA GPU with the converted checkpoint. Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
84 lines
2.9 KiB
Python
84 lines
2.9 KiB
Python
#!/usr/bin/env 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
|
|
|
|
import pytest
|
|
|
|
from lerobot.configs.policies import PreTrainedConfig
|
|
from lerobot.configs.types import FeatureType, PolicyFeature
|
|
from lerobot.policies.lingbot_va.configuration_lingbot_va import LingBotVAConfig
|
|
from lerobot.utils.constants import ACTION, OBS_IMAGES
|
|
|
|
|
|
def make_config(**overrides) -> LingBotVAConfig:
|
|
kwargs = {"device": "cpu"}
|
|
kwargs.update(overrides)
|
|
return LingBotVAConfig(**kwargs)
|
|
|
|
|
|
def test_registered_in_choice_registry() -> None:
|
|
assert "lingbot_va" in PreTrainedConfig.get_known_choices()
|
|
assert PreTrainedConfig.get_choice_class("lingbot_va") is LingBotVAConfig
|
|
|
|
|
|
def test_type_property() -> None:
|
|
assert make_config().type == "lingbot_va"
|
|
|
|
|
|
def test_chunk_size_and_action_steps() -> None:
|
|
cfg = make_config(frame_chunk_size=4, action_per_frame=4)
|
|
assert cfg.chunk_size == 16
|
|
assert cfg.n_action_steps == 16
|
|
assert cfg.action_delta_indices == list(range(16))
|
|
assert cfg.observation_delta_indices is None
|
|
assert cfg.reward_delta_indices is None
|
|
|
|
|
|
def test_optimizer_and_scheduler_presets() -> None:
|
|
cfg = make_config()
|
|
opt = cfg.get_optimizer_preset()
|
|
assert opt.lr == cfg.optimizer_lr
|
|
sched = cfg.get_scheduler_preset()
|
|
assert sched.num_warmup_steps == cfg.scheduler_warmup_steps
|
|
|
|
|
|
def test_validate_features_sets_action_feature() -> None:
|
|
cfg = make_config()
|
|
cfg.input_features = {f"{OBS_IMAGES}.image": PolicyFeature(type=FeatureType.VISUAL, shape=(3, 128, 128))}
|
|
cfg.output_features = {}
|
|
cfg.validate_features()
|
|
assert ACTION in cfg.output_features
|
|
assert cfg.output_features[ACTION].shape == (len(cfg.used_action_channel_ids),)
|
|
|
|
|
|
def test_validate_features_no_visual_raises() -> None:
|
|
cfg = make_config()
|
|
cfg.input_features = {}
|
|
cfg.output_features = {}
|
|
with pytest.raises(ValueError, match="at least one visual input feature"):
|
|
cfg.validate_features()
|
|
|
|
|
|
def test_invalid_attn_mode_raises() -> None:
|
|
with pytest.raises(ValueError, match="attn_mode"):
|
|
make_config(attn_mode="banana")
|
|
|
|
|
|
def test_quantile_length_mismatch_raises() -> None:
|
|
with pytest.raises(ValueError, match="action_q01"):
|
|
make_config(used_action_channel_ids=[0, 1, 2], action_q01=[0.0, 0.0], action_q99=[1.0, 1.0, 1.0])
|