diff --git a/src/lerobot/policies/vla_jepa/processor_vla_jepa.py b/src/lerobot/policies/vla_jepa/processor_vla_jepa.py index e9cca0527..585fe427b 100644 --- a/src/lerobot/policies/vla_jepa/processor_vla_jepa.py +++ b/src/lerobot/policies/vla_jepa/processor_vla_jepa.py @@ -17,11 +17,14 @@ from __future__ import annotations from typing import Any import torch +import torch.nn.functional as F # noqa: N812 +from lerobot.configs import PipelineFeatureType, PolicyFeature from lerobot.policies.vla_jepa.configuration_vla_jepa import VLAJEPAConfig from lerobot.processor import ( AbsoluteActionsProcessorStep, EnvTransition, + ObservationProcessorStep, PolicyAction, PolicyProcessorPipeline, ProcessorStep, @@ -34,6 +37,69 @@ from lerobot.processor import ( ) +@ProcessorStepRegistry.register(name="vla_jepa_image_prep") +class ImagePrepProcessorStep(ObservationProcessorStep): + """Prepares image observations for the VLA-JEPA model: float cast, 1->3 channel expand, resize. + + This makes explicit (in the serialized pipeline) the image prep the model used to do + internally. The model keeps the same operations as idempotent guards, so: + - checkpoints saved WITHOUT this step (older uploads) are unaffected — the model still + does the prep; + - checkpoints saved WITH this step get it done here, and the model-side guards no-op. + + Mirrors `Qwen3VLInterface.to_pixel_values` + the `F.interpolate(mode="area")` resize in + `VLAJEPAPolicy._prepare_model_inputs`/`predict_action`. Deliberately does NOT clamp (the + model path doesn't), so values stay bit-identical. Handles [C,H,W], [B,C,H,W]/[T,C,H,W] + and [B,T,C,H,W] image tensors. + """ + + def __init__(self, resize_to: tuple[int, int] | None = None, expand_channels: bool = True): + self.resize_to = tuple(resize_to) if resize_to is not None else None + self.expand_channels = expand_channels + + def observation(self, observation: dict) -> dict: + new_observation = dict(observation) + for key in observation: + if "image" not in key: + continue + image = observation[key].float() + if self.expand_channels and image.shape[-3] == 1: + repeats = [1] * image.ndim + repeats[-3] = 3 + image = image.repeat(*repeats) + if self.resize_to is not None and tuple(image.shape[-2:]) != self.resize_to: + device = image.device + # NOTE: no "area" kernel on mps; resize on cpu then move back. + if device.type == "mps": + image = image.cpu() + lead = image.shape[:-3] + c, h, w = image.shape[-3:] + flat = image.reshape(-1, c, h, w) + flat = F.interpolate(flat, size=self.resize_to, mode="area") + image = flat.reshape(*lead, c, *self.resize_to).to(device) + new_observation[key] = image + return new_observation + + def get_config(self) -> dict[str, Any]: + return { + "resize_to": list(self.resize_to) if self.resize_to is not None else None, + "expand_channels": self.expand_channels, + } + + def transform_features(self, features): + for key in features[PipelineFeatureType.OBSERVATION]: + if "image" not in key: + continue + feat = features[PipelineFeatureType.OBSERVATION][key] + # Match `to_pixel_values`: only a single channel is expanded to 3. + nb_channel = 3 if (self.expand_channels and feat.shape[0] == 1) else feat.shape[0] + spatial = self.resize_to if self.resize_to is not None else tuple(feat.shape[1:]) + features[PipelineFeatureType.OBSERVATION][key] = PolicyFeature( + type=feat.type, shape=(nb_channel, *spatial) + ) + return features + + @ProcessorStepRegistry.register(name="vla_jepa_clip_actions") class ClipActionsProcessorStep(ProcessorStep): """Clips action tensor to [-1, 1] before unnormalization.""" @@ -125,6 +191,9 @@ def make_vla_jepa_pre_post_processors( steps.rename_observations, steps.add_batch_dim, steps.to_device, + ImagePrepProcessorStep( + resize_to=tuple(config.resize_images_to) if config.resize_images_to else None, + ), relative_step, steps.normalize, ] diff --git a/tests/policies/vla_jepa/test_processor_vla_jepa.py b/tests/policies/vla_jepa/test_processor_vla_jepa.py new file mode 100644 index 000000000..7ce31ea76 --- /dev/null +++ b/tests/policies/vla_jepa/test_processor_vla_jepa.py @@ -0,0 +1,182 @@ +#!/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 + +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 FeatureType, PipelineFeatureType, PolicyFeature # noqa: E402 +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 ProcessorStepRegistry # noqa: E402 +from lerobot.utils.constants import 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)