#!/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. """Tests for the VLA-JEPA image-prep processor step and its back-compat contract with the model. The step moves image resize + 1->3 channel-expand out of the model into the (serialized) preprocessor. The model keeps the same ops as idempotent guards, so: - old checkpoints (JSON without the step) are unaffected — the model still does the prep; - new checkpoints (JSON with the step) get it done in the step, and the model guards no-op. These tests pin the step's numerics (bit-identical to the model's F.interpolate(area)) and the equivalence of the two paths on the Qwen image path. """ from __future__ import annotations import logging from copy import deepcopy import pytest import torch import torch.nn.functional as F # noqa: N812 pytest.importorskip("transformers") pytest.importorskip("diffusers") from conftest import ( # noqa: E402 BATCH_SIZE, IMAGE_SIZE, make_config, make_inference_batch, make_train_batch, ) from lerobot.configs.types import ( # noqa: E402 FeatureType, NormalizationMode, PipelineFeatureType, PolicyFeature, ) from lerobot.policies.vla_jepa.modeling_vla_jepa import VLAJEPAPolicy # noqa: E402 from lerobot.policies.vla_jepa.processor_vla_jepa import ( # noqa: E402 ImagePrepProcessorStep, make_vla_jepa_pre_post_processors, ) from lerobot.processor import PolicyProcessorPipeline, ProcessorStepRegistry # noqa: E402 from lerobot.utils.constants import ACTION, OBS_IMAGES, OBS_STATE # noqa: E402 RESIZE = (IMAGE_SIZE // 2, IMAGE_SIZE // 2) # (4, 4) IMG_KEY = f"{OBS_IMAGES}.laptop" # --------------------------------------------------------------------------- # Step numerics / shape handling # --------------------------------------------------------------------------- @pytest.mark.parametrize( "shape", [ (3, IMAGE_SIZE, IMAGE_SIZE), # [C, H, W] (raw single-sample inference) (BATCH_SIZE, 3, IMAGE_SIZE, IMAGE_SIZE), # [B, C, H, W] (BATCH_SIZE, 2, 3, IMAGE_SIZE, IMAGE_SIZE), # [B, T, C, H, W] (video stack) ], ) def test_image_prep_resize_shapes_and_area_numerics(shape: tuple[int, ...]) -> None: step = ImagePrepProcessorStep(resize_to=RESIZE) x = torch.rand(*shape) out = step.observation({IMG_KEY: x})[IMG_KEY] assert out.shape[:-2] == x.shape[:-2] # leading + channel dims unchanged assert tuple(out.shape[-2:]) == RESIZE assert out.dtype == torch.float32 # bit-identical to the model-side F.interpolate(mode="area"), no clamp ref = F.interpolate(x.float().reshape(-1, *x.shape[-3:]), size=RESIZE, mode="area").reshape( *x.shape[:-2], *RESIZE ) assert torch.equal(out, ref) def test_image_prep_channel_expand() -> None: step = ImagePrepProcessorStep(resize_to=None, expand_channels=True) x = torch.rand(BATCH_SIZE, 1, IMAGE_SIZE, IMAGE_SIZE) out = step.observation({IMG_KEY: x})[IMG_KEY] assert out.shape[1] == 3 # all three channels are copies of the single input channel assert torch.equal(out[:, 0], x[:, 0]) and torch.equal(out[:, 1], x[:, 0]) def test_image_prep_resize_skip_when_already_target_size() -> None: step = ImagePrepProcessorStep(resize_to=RESIZE) x = torch.rand(BATCH_SIZE, 3, *RESIZE) out = step.observation({IMG_KEY: x})[IMG_KEY] # size already matches -> only the float cast happens, values preserved exactly assert torch.equal(out, x) def test_image_prep_leaves_non_image_keys_untouched() -> None: step = ImagePrepProcessorStep(resize_to=RESIZE) state = torch.randn(BATCH_SIZE, 4) out = step.observation({IMG_KEY: torch.rand(BATCH_SIZE, 3, IMAGE_SIZE, IMAGE_SIZE), OBS_STATE: state}) assert torch.equal(out[OBS_STATE], state) def test_image_prep_config_roundtrip_via_registry() -> None: step = ImagePrepProcessorStep(resize_to=RESIZE, expand_channels=True) cfg = step.get_config() assert cfg == {"resize_to": [RESIZE[0], RESIZE[1]], "expand_channels": True} rebuilt = ProcessorStepRegistry.get("vla_jepa_image_prep")(**cfg) assert rebuilt.resize_to == RESIZE assert rebuilt.expand_channels is True def test_image_prep_transform_features() -> None: step = ImagePrepProcessorStep(resize_to=RESIZE, expand_channels=True) features = { PipelineFeatureType.OBSERVATION: { IMG_KEY: PolicyFeature(type=FeatureType.VISUAL, shape=(3, IMAGE_SIZE, IMAGE_SIZE)), "observation.images.depth": PolicyFeature( type=FeatureType.VISUAL, shape=(1, IMAGE_SIZE, IMAGE_SIZE) ), OBS_STATE: PolicyFeature(type=FeatureType.STATE, shape=(4,)), } } out = step.transform_features(features)[PipelineFeatureType.OBSERVATION] assert out[IMG_KEY].shape == (3, *RESIZE) # already 3-channel, only resized assert out["observation.images.depth"].shape == (3, *RESIZE) # 1->3 expanded assert out[OBS_STATE].shape == (4,) # non-image untouched # --------------------------------------------------------------------------- # Pipeline wiring + back-compat with the model # --------------------------------------------------------------------------- def test_image_prep_step_wired_into_preprocessor() -> None: cfg = make_config() cfg.resize_images_to = RESIZE preprocessor, _ = make_vla_jepa_pre_post_processors(cfg, dataset_stats=None) prep_steps = [s for s in preprocessor.steps if isinstance(s, ImagePrepProcessorStep)] assert len(prep_steps) == 1 assert prep_steps[0].resize_to == RESIZE @torch.no_grad() @pytest.mark.parametrize("batch_fn", [make_inference_batch, make_train_batch]) def test_image_prep_matches_model_qwen_path(patch_vla_jepa_external_models: None, batch_fn) -> None: """The Qwen image path is identical whether the step resized (new ckpt) or the model does (old ckpt). Both use F.interpolate(mode="area"), so pre-resizing in the step then letting the model's size guard no-op yields byte-identical Qwen inputs to the pure model path. This is the contract that keeps already-uploaded checkpoints correct. """ cfg = make_config() cfg.resize_images_to = RESIZE policy = VLAJEPAPolicy(cfg) policy.eval() training = batch_fn is make_train_batch batch = batch_fn() # Path A (old checkpoint, no processor step): the model resizes internally. imgs_a = policy._prepare_model_inputs(deepcopy(batch), training=training)["images"] # Path B (new checkpoint): the step resizes first; the model's guard becomes a no-op. step = ImagePrepProcessorStep(resize_to=RESIZE) resized = step.observation({IMG_KEY: batch[IMG_KEY]}) batch_b = deepcopy(batch) batch_b[IMG_KEY] = resized[IMG_KEY] imgs_b = policy._prepare_model_inputs(batch_b, training=training)["images"] assert len(imgs_a) == len(imgs_b) == BATCH_SIZE for views_a, views_b in zip(imgs_a, imgs_b, strict=True): for a, b in zip(views_a, views_b, strict=True): assert torch.equal(a, b) # --------------------------------------------------------------------------- # Gripper post-step serialization and the normalization-mode coupling # --------------------------------------------------------------------------- def _stats_for(cfg, gripper_min: float = 0.0, gripper_max: float = 1.0): """Dataset stats matching `cfg`'s features, with a settable gripper range.""" action_min = torch.zeros(cfg.action_dim) action_max = torch.ones(cfg.action_dim) gripper = cfg.resolved_gripper_dim if gripper < cfg.action_dim: action_min[gripper] = gripper_min action_max[gripper] = gripper_max stats = { ACTION: { "min": action_min, "max": action_max, "mean": torch.zeros(cfg.action_dim), "std": torch.ones(cfg.action_dim), } } for key, feat in cfg.input_features.items(): stats[key] = { "min": torch.zeros(feat.shape), "max": torch.ones(feat.shape), "mean": torch.zeros(feat.shape), "std": torch.ones(feat.shape), } return stats def _find(pipeline, class_name: str): return next((s for s in pipeline.steps if type(s).__name__ == class_name), None) def test_gripper_steps_survive_a_save_load_round_trip(tmp_path): """gripper_dim/gripper_threshold must not silently revert to the class defaults. Both steps used to inherit `get_config() -> {}`, so a reloaded pipeline came back at `gripper_dim=6, threshold=0.5` no matter what the training config said. """ cfg = make_config(action_dim=7) cfg.gripper_dim = 5 cfg.gripper_threshold = 0.25 cfg.pre_snap_gripper_action = True cfg.binarize_gripper_action = True _, post = make_vla_jepa_pre_post_processors(cfg, _stats_for(cfg)) post.save_pretrained(str(tmp_path)) reloaded = PolicyProcessorPipeline.from_pretrained( str(tmp_path), config_filename="policy_postprocessor.json", overrides={} ) for class_name in ("PreSnapGripperProcessorStep", "BinarizeGripperProcessorStep"): step = _find(reloaded, class_name) assert step is not None, class_name assert (step.gripper_dim, step.threshold) == (5, 0.25), class_name def test_gripper_steps_default_off(): """The LIBERO-specific gripper steps are opt-in, since they assume LIBERO's units.""" cfg = make_config(action_dim=7) assert not cfg.pre_snap_gripper_action assert not cfg.binarize_gripper_action _, post = make_vla_jepa_pre_post_processors(cfg, _stats_for(cfg)) assert _find(post, "PreSnapGripperProcessorStep") is None assert _find(post, "BinarizeGripperProcessorStep") is None def test_gripper_dim_out_of_range_raises(): """An out-of-range gripper index used to make both steps no-op silently.""" cfg = make_config(action_dim=7) cfg.gripper_dim = 7 cfg.pre_snap_gripper_action = True with pytest.raises(ValueError, match="out of range"): cfg.validate_features() def test_gripper_dim_resolved_from_action_names(): cfg = make_config(action_dim=4) cfg.gripper_dim = 6 cfg.action_feature_names = ["shoulder.pos", "elbow.pos", "gripper.pos", "wrist.pos"] assert cfg.resolved_gripper_dim == 2 def test_physical_range_mismatch_warns(caplog): """A gripper in degrees would be pinned to a constant by the 0.5 threshold.""" cfg = make_config(action_dim=7) cfg.gripper_dim = 6 cfg.pre_snap_gripper_action = True cfg.binarize_gripper_action = True with caplog.at_level(logging.WARNING): make_vla_jepa_pre_post_processors(cfg, _stats_for(cfg, gripper_min=0.0, gripper_max=90.0)) assert "looks misconfigured" in caplog.text def test_libero_style_gripper_range_does_not_warn(caplog): cfg = make_config(action_dim=7) cfg.gripper_dim = 6 cfg.pre_snap_gripper_action = True cfg.binarize_gripper_action = True with caplog.at_level(logging.WARNING): make_vla_jepa_pre_post_processors(cfg, _stats_for(cfg, gripper_min=-1.0, gripper_max=1.0)) assert "looks misconfigured" not in caplog.text def test_action_clipping_is_skipped_unless_min_max(caplog): """Clipping to [-1, 1] is a range bound under MIN_MAX but a 1-sigma truncation under MEAN_STD.""" cfg = make_config(action_dim=7) assert cfg.clip_normalized_actions _, post = make_vla_jepa_pre_post_processors(cfg, _stats_for(cfg)) assert _find(post, "ClipActionsProcessorStep") is not None cfg.normalization_mapping = {**cfg.normalization_mapping, "ACTION": NormalizationMode.MEAN_STD} with caplog.at_level(logging.WARNING): _, post = make_vla_jepa_pre_post_processors(cfg, _stats_for(cfg)) assert _find(post, "ClipActionsProcessorStep") is None assert "clip_normalized_actions" in caplog.text def test_observation_delta_indices_collapse_without_world_model(): """Only the world model reads frames past index 0; asking for more decodes video for nothing.""" assert make_config(num_video_frames=8).observation_delta_indices == list(range(8)) cfg = make_config(num_video_frames=8) cfg.enable_world_model = False assert cfg.observation_delta_indices == [0]