diff --git a/src/lerobot/envs/utils.py b/src/lerobot/envs/utils.py index 8b9c4f94b14..995d6adf814 100644 --- a/src/lerobot/envs/utils.py +++ b/src/lerobot/envs/utils.py @@ -28,6 +28,7 @@ from torch import Tensor from lerobot.configs import FeatureType, PolicyFeature +from lerobot.types import ImageInputFormat from lerobot.utils.constants import OBS_ENV_STATE, OBS_IMAGE, OBS_IMAGES, OBS_STATE, OBS_STR from lerobot.utils.utils import get_channel_first_image_shape @@ -65,7 +66,10 @@ def _convert_nested_dict(d): return result -def preprocess_observation(observations: dict[str, np.ndarray]) -> dict[str, Tensor]: +def preprocess_observation( + observations: dict[str, np.ndarray], + image_input_format: ImageInputFormat = ImageInputFormat.FLOAT32_0_1, +) -> dict[str, Tensor]: # TODO(jadechoghari, imstevenpmwork): refactor this to use features from the environment (no hardcoding) """Convert environment observation to LeRobot format observation. Args: @@ -96,10 +100,13 @@ def preprocess_observation(observations: dict[str, np.ndarray]) -> dict[str, Ten # sanity check that images are uint8 assert img_tensor.dtype == torch.uint8, f"expect torch.uint8, but instead {img_tensor.dtype=}" - # convert to channel first of type float32 in range [0,1] + # convert to channel first, keeping the policy's expected image range img_tensor = einops.rearrange(img_tensor, "b h w c -> b c h w").contiguous() - img_tensor = img_tensor.type(torch.float32) - img_tensor /= 255 + if image_input_format is ImageInputFormat.FLOAT32_0_1: + img_tensor = img_tensor.type(torch.float32) + img_tensor /= 255 + elif image_input_format is not ImageInputFormat.UINT8_0_255: + raise ValueError(f"Unsupported policy image input format: {image_input_format}") return_observations[imgkey] = img_tensor diff --git a/src/lerobot/policies/__init__.py b/src/lerobot/policies/__init__.py index 7f0bed2e051..dbb26d09b46 100644 --- a/src/lerobot/policies/__init__.py +++ b/src/lerobot/policies/__init__.py @@ -28,7 +28,7 @@ from .pi0.configuration_pi0 import PI0Config as PI0Config from .pi0_fast.configuration_pi0_fast import PI0FastConfig as PI0FastConfig from .pi05.configuration_pi05 import PI05Config as PI05Config -from .pretrained import PreTrainedPolicy as PreTrainedPolicy +from .pretrained import ImageInputFormat as ImageInputFormat, PreTrainedPolicy as PreTrainedPolicy from .smolvla.configuration_smolvla import SmolVLAConfig as SmolVLAConfig from .tdmpc.configuration_tdmpc import TDMPCConfig as TDMPCConfig from .utils import make_robot_action, prepare_observation_for_inference @@ -61,6 +61,7 @@ "WallXConfig", "XVLAConfig", # Base class + "ImageInputFormat", "PreTrainedPolicy", # RTC utilities "ActionInterpolator", diff --git a/src/lerobot/policies/groot/modeling_groot.py b/src/lerobot/policies/groot/modeling_groot.py index 415af930990..8a470e3f3dd 100644 --- a/src/lerobot/policies/groot/modeling_groot.py +++ b/src/lerobot/policies/groot/modeling_groot.py @@ -39,7 +39,7 @@ from lerobot.utils.constants import ACTION, OBS_IMAGES from lerobot.utils.import_utils import _transformers_available, require_package -from ..pretrained import PreTrainedPolicy +from ..pretrained import ImageInputFormat, PreTrainedPolicy from ..utils import get_device_from_parameters from .configuration_groot import ( GROOT_N1_5, @@ -67,6 +67,7 @@ class GrootPolicy(PreTrainedPolicy): name = "groot" config_class = GrootConfig + input_image_format = ImageInputFormat.UINT8_0_255 def __init__(self, config: GrootConfig, **kwargs): """Initialize Groot policy wrapper.""" diff --git a/src/lerobot/policies/groot/processor_groot.py b/src/lerobot/policies/groot/processor_groot.py index 20b3518a306..365a3894001 100644 --- a/src/lerobot/policies/groot/processor_groot.py +++ b/src/lerobot/policies/groot/processor_groot.py @@ -25,7 +25,6 @@ import numpy as np import torch import torchvision.transforms.v2.functional as tv_functional -from einops import rearrange from torchvision.transforms import InterpolationMode from lerobot.utils.import_utils import _datasets_available, _transformers_available, require_package @@ -104,6 +103,8 @@ # action chunks, so processor-side horizons are capped at this value. N1_7_NATIVE_ACTION_HORIZON = 40 +type _GrootN17CameraBatch = tuple[torch.Tensor, ...] + N1_7_EMBODIMENT_MAPPING = { "oxe_droid_relative_eef_relative_joint": 24, "xdof_relative_eef_relative_joint": 27, @@ -1336,19 +1337,17 @@ def make_groot_pre_post_processors( # GR00T specific processor steps -def _to_uint8_np_bthwc(img_t: torch.Tensor) -> np.ndarray: - # img_t: (B, C, H, W) or (B, T, C, H, W), float in [0,1] or uint8 - if img_t.dtype.is_floating_point: - img_t = (img_t.clamp(0, 1) * 255.0).to(torch.uint8) - if img_t.dim() == 4: - return rearrange(img_t.cpu().numpy(), "b c h w -> b 1 h w c") - if img_t.dim() == 5: - return rearrange(img_t.cpu().numpy(), "b t c h w -> b t h w c") - raise ValueError(f"Expected image tensor shape (B, C, H, W) or (B, T, C, H, W), got {tuple(img_t.shape)}") +def _as_uint8_video_tensor_btchw(image: torch.Tensor) -> torch.Tensor: + if image.ndim not in (4, 5): + raise ValueError( + f"Expected image tensor shape (B, C, H, W) or (B, T, C, H, W), got {tuple(image.shape)}." + ) + image = tv_functional.to_dtype(image, torch.uint8, scale=True) + return image.unsqueeze(1) if image.ndim == 4 else image -def _align_video_horizon(video: np.ndarray, horizon: int | None) -> np.ndarray: - """Match the checkpoint video horizon by truncating or left-padding frames.""" +def _align_video_horizon_tensor(video: torch.Tensor, horizon: int | None) -> torch.Tensor: + """Match the checkpoint video horizon without changing dtype or tensor layout.""" if horizon is None or horizon <= 0: return video @@ -1357,8 +1356,12 @@ def _align_video_horizon(video: np.ndarray, horizon: int | None) -> np.ndarray: return video if current > horizon: return video[:, -horizon:] - pad = np.repeat(video[:, :1], horizon - current, axis=1) - return np.concatenate([pad, video], axis=1) + pad = video[:, :1].expand(-1, horizon - current, -1, -1, -1) + return torch.cat([pad, video], dim=1) + + +def _uint8_image_numpy_hwc(image: torch.Tensor) -> np.ndarray: + return image.detach().cpu().permute(1, 2, 0).contiguous().numpy() def _build_n1_7_processor(model_name: str = GROOT_N1_7_BACKBONE_MODEL) -> ProcessorMixin: @@ -1844,21 +1847,19 @@ def _cache_raw_state(state: torch.Tensor) -> None: if grouped: self._last_raw_state = grouped + batch_size, batch_device = infer_n1_7_batch_size_and_device(obs, transition.get(TransitionKey.ACTION)) img_keys = self._ordered_image_keys(obs) if img_keys: - cams = [_align_video_horizon(_to_uint8_np_bthwc(obs[k]), self.video_horizon) for k in img_keys] - video = np.stack(cams, axis=2) # (B, T, V, H, W, C) - obs["video"] = video - image_keys_to_remove = [key for key in obs if key.startswith(OBS_IMAGES)] - if OBS_IMAGE in obs: - image_keys_to_remove.append(OBS_IMAGE) - for k in image_keys_to_remove: - obs.pop(k, None) - - bsz, _device = infer_n1_7_batch_size_and_device(obs, transition.get(TransitionKey.ACTION)) + obs["video"] = tuple( + _align_video_horizon_tensor(_as_uint8_video_tensor_btchw(obs[key]), self.video_horizon) + for key in img_keys + ) + # Preserve channels-first tensors until VLM preprocessing. + for key in [key for key in obs if key.startswith(OBS_IMAGES) or key == OBS_IMAGE]: + obs.pop(key) comp["language"] = prepare_n1_7_language_batch( comp.get(self.language_key), - bsz, + batch_size, formalize_language=self.formalize_language, ) @@ -1967,13 +1968,14 @@ def _cache_raw_state(state: torch.Tensor) -> None: comp["action_mask"] = action_mask emb_id = self.embodiment_mapping.get(self.embodiment_tag, 0) - bsz, device = infer_n1_7_batch_size_and_device(obs, transition.get(TransitionKey.ACTION)) if "action_mask" not in comp: - action_mask = torch.zeros(bsz, self.action_horizon, dtype=torch.float32, device=device) + action_mask = torch.zeros( + batch_size, self.action_horizon, dtype=torch.float32, device=batch_device + ) valid_horizon = min(self.valid_action_horizon, self.action_horizon) action_mask[:, :valid_horizon] = 1.0 comp["action_mask"] = action_mask - comp["embodiment_id"] = torch.full((bsz,), emb_id, dtype=torch.int32, device=device) + comp["embodiment_id"] = torch.full((batch_size,), emb_id, dtype=torch.int32, device=batch_device) transition[TransitionKey.OBSERVATION] = obs transition[TransitionKey.COMPLEMENTARY_DATA] = comp @@ -2035,9 +2037,9 @@ def load_state_dict(self, state: dict[str, torch.Tensor]) -> None: class GrootN17VLMEncodeStep(ProcessorStep): """Tokenize N1.7's packed video-language prompt with the Qwen3-VL processor. - The packed video has shape ``(B, T, V, H, W, C)``. Each frame/view becomes - an image item in the same chat message so the resulting image tokens match - the temporal VLM packing used by Isaac-GR00T. + The packed video is an ordered tuple of ``(B, T, C, H, W)`` camera tensors. + Each frame/view becomes an image item in the same chat message so the + resulting image tokens match the temporal VLM packing used by Isaac-GR00T. Images are handed to the torchvision-backed Qwen3-VL processor as ``(C, H, W)`` uint8 tensors (no per-frame PIL roundtrip), and, when ``device`` resolves to a @@ -2081,7 +2083,9 @@ def _target_device(self) -> torch.device | None: return None def _build_sample_images( - self, video: Any, batch_size: int, target_device: torch.device | None + self, + cameras: _GrootN17CameraBatch, + target_device: torch.device | None, ) -> list[list[Any]]: """Return, per batch item, its ordered ``(timestep, view)`` frames. @@ -2089,8 +2093,10 @@ def _build_sample_images( otherwise frames are ``(C, H, W)`` uint8 tensors (moved to ``target_device`` when set) for the torchvision-backed Qwen processor. """ + batch_size = cameras[0].shape[0] + horizon = cameras[0].shape[1] + if self.use_albumentations: - video_np = np.asarray(video) train_crop = self.training and torch.is_grad_enabled() sample_images: list[list[Any]] = [] for batch_idx in range(batch_size): @@ -2101,7 +2107,7 @@ def _build_sample_images( sample_images.append( [ _transform_n1_7_image_for_vlm_albumentations( - video_np[batch_idx, timestep, view_idx], + _uint8_image_numpy_hwc(cameras[view_idx][batch_idx, timestep]), image_crop_size=self.image_crop_size, image_target_size=self.image_target_size, shortest_image_edge=self.shortest_image_edge, @@ -2109,54 +2115,51 @@ def _build_sample_images( letter_box_transform=self.letter_box_transform, crop_position=crop_position, ) - for timestep in range(video_np.shape[1]) - for view_idx in range(video_np.shape[2]) + for timestep in range(horizon) + for view_idx in range(len(cameras)) ] ) return sample_images - video_t = video if torch.is_tensor(video) else torch.from_numpy(np.ascontiguousarray(video)) - # (B, T, V, H, W, C) uint8 -> (B, T, V, C, H, W) - video_t = video_t.permute(0, 1, 2, 5, 3, 4).contiguous() - if target_device is not None and video_t.device != target_device: - video_t = video_t.to(target_device, non_blocking=(target_device.type == "cuda")) - - frames_per_sample: list[list[Any]] = [] - for batch_idx in range(batch_size): - sample = video_t[batch_idx] # (T, V, C, H, W) - frames_per_sample.append( - [ - _transform_n1_7_image_for_vlm_torch( - sample[timestep, view_idx], - image_crop_size=self.image_crop_size, - image_target_size=self.image_target_size, - shortest_image_edge=self.shortest_image_edge, - crop_fraction=self.crop_fraction, - letter_box_transform=self.letter_box_transform, - ) - for timestep in range(sample.shape[0]) - for view_idx in range(sample.shape[1]) - ] - ) - return frames_per_sample + prepared_cameras = [ + camera.to(target_device, non_blocking=(target_device.type == "cuda")) + if target_device is not None and camera.device != target_device + else camera + for camera in cameras + ] + + return [ + [ + _transform_n1_7_image_for_vlm_torch( + prepared_cameras[view_idx][batch_idx, timestep], + image_crop_size=self.image_crop_size, + image_target_size=self.image_target_size, + shortest_image_edge=self.shortest_image_edge, + crop_fraction=self.crop_fraction, + letter_box_transform=self.letter_box_transform, + ) + for timestep in range(horizon) + for view_idx in range(len(prepared_cameras)) + ] + for batch_idx in range(batch_size) + ] def __call__(self, transition: EnvTransition) -> EnvTransition: obs = transition.get(TransitionKey.OBSERVATION, {}) or {} comp = transition.get(TransitionKey.COMPLEMENTARY_DATA, {}) or {} - video = obs.get("video") - if video is None: + cameras: _GrootN17CameraBatch | None = obs.get("video") + if cameras is None: return transition - batch_size = int(video.shape[0]) + target_device = self._target_device() + sample_images = self._build_sample_images(cameras, target_device) + batch_size = len(sample_images) languages = prepare_n1_7_language_batch( comp.get("language"), batch_size, formalize_language=False, ) - target_device = self._target_device() - sample_images = self._build_sample_images(video, batch_size, target_device) - texts: list[str] = [] images: list[Any] = [] for batch_idx in range(batch_size): diff --git a/src/lerobot/policies/groot/utils.py b/src/lerobot/policies/groot/utils.py index 9a65404fe85..e332835bf52 100644 --- a/src/lerobot/policies/groot/utils.py +++ b/src/lerobot/policies/groot/utils.py @@ -224,9 +224,6 @@ def infer_n1_7_batch_size_and_device( for value in list(obs.values()) + [action]: if isinstance(value, torch.Tensor): return value.shape[0], value.device - video = obs.get("video") - if isinstance(video, np.ndarray): - return video.shape[0], torch.device("cpu") return 1, torch.device("cpu") diff --git a/src/lerobot/policies/pretrained.py b/src/lerobot/policies/pretrained.py index 702569b8c7b..cdd04c6fd11 100644 --- a/src/lerobot/policies/pretrained.py +++ b/src/lerobot/policies/pretrained.py @@ -21,7 +21,7 @@ from importlib.resources import files from pathlib import Path from tempfile import TemporaryDirectory -from typing import TYPE_CHECKING, TypedDict, TypeVar, Unpack +from typing import TYPE_CHECKING, ClassVar, TypedDict, TypeVar, Unpack import packaging import safetensors @@ -34,6 +34,7 @@ from lerobot.__version__ import __version__ from lerobot.configs import PreTrainedConfig from lerobot.configs.train import TrainPipelineConfig +from lerobot.types import ImageInputFormat from lerobot.utils.hub import HubMixin from .utils import log_model_loading_keys @@ -102,6 +103,7 @@ class PreTrainedPolicy(nn.Module, HubMixin, abc.ABC): config_class: None name: None + input_image_format: ClassVar[ImageInputFormat] = ImageInputFormat.FLOAT32_0_1 def __init__(self, config: PreTrainedConfig, *inputs, **kwargs): super().__init__() diff --git a/src/lerobot/scripts/lerobot_eval.py b/src/lerobot/scripts/lerobot_eval.py index 722763d6e79..dc8b7575073 100644 --- a/src/lerobot/scripts/lerobot_eval.py +++ b/src/lerobot/scripts/lerobot_eval.py @@ -82,7 +82,7 @@ make_env_pre_post_processors, preprocess_observation, ) -from lerobot.policies import PreTrainedPolicy, make_policy, make_pre_post_processors +from lerobot.policies import ImageInputFormat, PreTrainedPolicy, make_policy, make_pre_post_processors from lerobot.processor import PolicyProcessorPipeline from lerobot.types import PolicyAction from lerobot.utils.constants import ACTION, DONE, OBS_IMAGE, OBS_IMAGES, OBS_STR, REWARD @@ -260,7 +260,10 @@ def rollout( try: while not np.all(done) and step < max_steps: # Numpy array to tensor and changing dictionary keys to LeRobot policy format. - observation = preprocess_observation(observation) + observation = preprocess_observation( + observation, + image_input_format=getattr(policy, "input_image_format", ImageInputFormat.FLOAT32_0_1), + ) if return_observations: all_observations.append(deepcopy(observation)) @@ -378,7 +381,10 @@ def rollout( # Track the final observation. if return_observations: - observation = preprocess_observation(observation) + observation = preprocess_observation( + observation, + image_input_format=getattr(policy, "input_image_format", ImageInputFormat.FLOAT32_0_1), + ) all_observations.append(deepcopy(observation)) # Stack the sequence along the first dimension so that we have (batch, sequence, *) tensors. diff --git a/src/lerobot/scripts/lerobot_train.py b/src/lerobot/scripts/lerobot_train.py index 6e845852306..f47ffd18a05 100644 --- a/src/lerobot/scripts/lerobot_train.py +++ b/src/lerobot/scripts/lerobot_train.py @@ -54,7 +54,7 @@ from lerobot.envs import close_envs, make_env, make_env_pre_post_processors from lerobot.jobs import submit_to_hf from lerobot.optim.factory import make_optimizer_and_scheduler -from lerobot.policies import PreTrainedPolicy, make_policy, make_pre_post_processors +from lerobot.policies import ImageInputFormat, PreTrainedPolicy, make_policy, make_pre_post_processors from lerobot.rewards import make_reward_pre_post_processors from lerobot.utils.collate import lerobot_collate_fn from lerobot.utils.import_utils import register_third_party_plugins @@ -71,6 +71,23 @@ from .lerobot_eval import eval_policy_all +def prep_camera_images_for_policy( + batch: dict[str, Any], + camera_keys: list[str], + input_image_format: ImageInputFormat, +) -> None: + """Adapt worker-produced camera images to the policy's raw input contract.""" + + if input_image_format is ImageInputFormat.UINT8_0_255: + return + if input_image_format is not ImageInputFormat.FLOAT32_0_1: + raise ValueError(f"Unsupported policy image input format: {input_image_format}") + + for cam_key in camera_keys: + if cam_key in batch and batch[cam_key].dtype == torch.uint8: + batch[cam_key] = batch[cam_key].to(dtype=torch.float32) / 255.0 + + def update_policy( train_metrics: MetricsTracker, policy: PreTrainedPolicy, @@ -296,6 +313,10 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None): rename_map=cfg.rename_map, ) + # Capture this before Accelerate or PEFT wraps the policy. Reward models do + # not currently declare an image format and retain the historical default. + input_image_format = getattr(policy, "input_image_format", ImageInputFormat.FLOAT32_0_1) + if cfg.peft is not None: if cfg.is_reward_model_training: raise ValueError("PEFT is only supported for policy training. ") @@ -566,9 +587,7 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None): for _ in range(step, cfg.steps): start_time = time.perf_counter() batch = next(dl_iter) - for cam_key in dataset.meta.camera_keys: - if cam_key in batch and batch[cam_key].dtype == torch.uint8: - batch[cam_key] = batch[cam_key].to(dtype=torch.float32) / 255.0 + prep_camera_images_for_policy(batch, dataset.meta.camera_keys, input_image_format) batch = preprocessor(batch) train_tracker.dataloading_s = time.perf_counter() - start_time @@ -621,9 +640,7 @@ def train(cfg: TrainPipelineConfig, accelerator: "Accelerator | None" = None): n_eval_batches = 0 with torch.no_grad(), accelerator.autocast(): for eval_batch in eval_dataloader: - for cam_key in dataset.meta.camera_keys: - if cam_key in eval_batch and eval_batch[cam_key].dtype == torch.uint8: - eval_batch[cam_key] = eval_batch[cam_key].to(dtype=torch.float32) / 255.0 + prep_camera_images_for_policy(eval_batch, dataset.meta.camera_keys, input_image_format) eval_batch = preprocessor(eval_batch) loss, _ = policy.forward(eval_batch) eval_loss_sum += loss.item() diff --git a/src/lerobot/types.py b/src/lerobot/types.py index 9de504870b0..7b81215a818 100644 --- a/src/lerobot/types.py +++ b/src/lerobot/types.py @@ -16,7 +16,7 @@ from __future__ import annotations -from enum import Enum +from enum import Enum, StrEnum from typing import Any, TypedDict import numpy as np @@ -36,6 +36,18 @@ class TransitionKey(str, Enum): COMPLEMENTARY_DATA = "complementary_data" +class ImageInputFormat(StrEnum): + """Raw image dtype/range expected by a policy before preprocessing. + + This only describes the tensor dtype and numeric range. It does not encode + channel count, channel order, or layout; those are defined by dataset/env + feature conventions and processor steps. + """ + + FLOAT32_0_1 = "float32_0_1" + UINT8_0_255 = "uint8_0_255" + + PolicyAction = torch.Tensor RobotAction = dict[str, Any] EnvAction = np.ndarray diff --git a/tests/envs/test_envs.py b/tests/envs/test_envs.py index c6a0b077dd8..67b4e8ac6ee 100644 --- a/tests/envs/test_envs.py +++ b/tests/envs/test_envs.py @@ -31,6 +31,7 @@ _parse_hub_url, preprocess_observation, ) +from lerobot.policies import ImageInputFormat from tests.utils import require_env OBS_TYPES = ["state", "pixels", "pixels_agent_pos"] @@ -43,6 +44,26 @@ AVAILABLE_ENVS = ["aloha", "pusht"] +def test_preprocess_observation_preserves_uint8_for_uint8_policy_contract(): + image = np.zeros((4, 5, 3), dtype=np.uint8) + image[0, 0] = [0, 127, 255] + image[1, 2] = [255, 127, 0] + + obs = preprocess_observation( + {"pixels": image}, + image_input_format=ImageInputFormat.UINT8_0_255, + ) + + assert obs["observation.image"].dtype == torch.uint8 + assert obs["observation.image"].shape == (1, 3, 4, 5) + torch.testing.assert_close( + obs["observation.image"][0, :, 0, 0], torch.tensor([0, 127, 255], dtype=torch.uint8) + ) + torch.testing.assert_close( + obs["observation.image"][0, :, 1, 2], torch.tensor([255, 127, 0], dtype=torch.uint8) + ) + + @pytest.mark.parametrize("obs_type", OBS_TYPES) @pytest.mark.parametrize("env_name, env_task", ENV_TASK_PAIRS) @require_env diff --git a/tests/policies/groot/test_groot_n1_7.py b/tests/policies/groot/test_groot_n1_7.py index 8b74e4664f4..b639dc9331e 100644 --- a/tests/policies/groot/test_groot_n1_7.py +++ b/tests/policies/groot/test_groot_n1_7.py @@ -27,6 +27,7 @@ from torch import nn from lerobot.configs import FeatureType, PolicyFeature +from lerobot.policies import ImageInputFormat from lerobot.policies.factory import make_policy_config, make_pre_post_processors from lerobot.policies.groot.configuration_groot import ( GROOT_ACTION_DECODE_TRANSFORM_LIBERO, @@ -1234,11 +1235,13 @@ def test_groot_n1_7_pack_inputs_orders_video_by_checkpoint_modality_keys(): normalize_min_max=False, video_modality_keys=["image", "wrist_image"], ) + wrist = torch.full((1, 3, 2, 2), 22, dtype=torch.uint8) + front = torch.full((1, 3, 2, 2), 11, dtype=torch.uint8) transition = { TransitionKey.OBSERVATION: { f"{OBS_IMAGES}.zz_extra": torch.full((1, 3, 2, 2), 33, dtype=torch.uint8), - f"{OBS_IMAGES}.image2": torch.full((1, 3, 2, 2), 22, dtype=torch.uint8), - f"{OBS_IMAGES}.image": torch.full((1, 3, 2, 2), 11, dtype=torch.uint8), + f"{OBS_IMAGES}.image2": wrist, + f"{OBS_IMAGES}.image": front, OBS_STATE: torch.zeros(1, 8), }, TransitionKey.COMPLEMENTARY_DATA: {"task": ["Move"]}, @@ -1247,9 +1250,10 @@ def test_groot_n1_7_pack_inputs_orders_video_by_checkpoint_modality_keys(): output = step(transition) video = output[TransitionKey.OBSERVATION]["video"] - assert video.shape == (1, 1, 2, 2, 2, 3) - assert np.unique(video[0, 0, 0]).tolist() == [11] - assert np.unique(video[0, 0, 1]).tolist() == [22] + assert isinstance(video, tuple) and len(video) == 2 + assert video[0].shape == (1, 1, 3, 2, 2) + assert video[0].data_ptr() == front.data_ptr() + assert video[1].data_ptr() == wrist.data_ptr() assert f"{OBS_IMAGES}.zz_extra" not in output[TransitionKey.OBSERVATION] assert f"{OBS_IMAGES}.image" not in output[TransitionKey.OBSERVATION] assert f"{OBS_IMAGES}.image2" not in output[TransitionKey.OBSERVATION] @@ -1664,6 +1668,7 @@ def test_groot_n1_7_processors_are_registered_lazily_without_external_gr00t(): preprocessor, _ = make_groot_pre_post_processors(config) step_types = {type(step) for step in preprocessor.steps} + assert GrootPolicy.input_image_format is ImageInputFormat.UINT8_0_255 assert GrootN17PackInputsStep in step_types assert GrootN17VLMEncodeStep in step_types assert "gr00t" not in sys.modules @@ -1752,7 +1757,7 @@ def __call__(self, text, images, return_tensors, padding): step._proc = fake_proc transition = { TransitionKey.OBSERVATION: { - "video": np.zeros((2, 1, 1, 2, 2, 3), dtype=np.uint8), + "video": (torch.zeros((2, 1, 3, 2, 2), dtype=torch.uint8),), }, TransitionKey.COMPLEMENTARY_DATA: { "language": ["first task", "second task"], @@ -1807,15 +1812,18 @@ def __call__(self, text, images, return_tensors, padding): fake_proc = FakeProcessor() step = GrootN17VLMEncodeStep() step._proc = fake_proc - video = np.zeros((2, 2, 2, 2, 2, 3), dtype=np.uint8) + cameras = ( + torch.zeros((2, 2, 3, 2, 2), dtype=torch.uint8), + torch.zeros((2, 2, 3, 2, 2), dtype=torch.uint8), + ) image_id = 1 for batch_idx in range(2): for timestep in range(2): for view_idx in range(2): - video[batch_idx, timestep, view_idx, :, :, :] = image_id + cameras[view_idx][batch_idx, timestep] = image_id image_id += 1 transition = { - TransitionKey.OBSERVATION: {"video": video}, + TransitionKey.OBSERVATION: {"video": cameras}, TransitionKey.COMPLEMENTARY_DATA: {"language": ["task a", "task b"]}, } @@ -1936,7 +1944,9 @@ def __call__(self, text, images, return_tensors, padding): camera_a = np.arange(3 * 5 * 3, dtype=np.uint8).reshape(3, 5, 3) camera_b = (np.arange(3 * 5 * 3, dtype=np.uint16).reshape(3, 5, 3) * 3 % 251).astype(np.uint8) - video = np.stack([camera_a, camera_b], axis=0).reshape(1, 1, 2, 3, 5, 3) + cameras = tuple( + torch.from_numpy(camera).permute(2, 0, 1).unsqueeze(0).unsqueeze(0) for camera in (camera_a, camera_b) + ) fake_proc = FakeProcessor() step = GrootN17VLMEncodeStep( image_target_size=[8, 8], @@ -1948,7 +1958,7 @@ def __call__(self, text, images, return_tensors, padding): step( { - TransitionKey.OBSERVATION: {"video": video}, + TransitionKey.OBSERVATION: {"video": cameras}, TransitionKey.COMPLEMENTARY_DATA: {"language": ["move"]}, } ) diff --git a/tests/policies/groot/test_groot_n1_7_oss_parity.py b/tests/policies/groot/test_groot_n1_7_oss_parity.py index 3fced5909a8..e39c885519f 100644 --- a/tests/policies/groot/test_groot_n1_7_oss_parity.py +++ b/tests/policies/groot/test_groot_n1_7_oss_parity.py @@ -96,7 +96,10 @@ def __call__(self, **kwargs): step._proc = processor transition = { TransitionKey.OBSERVATION: { - "video": np.zeros((1, 1, 2, 480, 640, 3), dtype=np.uint8), + "video": ( + torch.zeros((1, 1, 3, 480, 640), dtype=torch.uint8), + torch.zeros((1, 1, 3, 480, 640), dtype=torch.uint8), + ), }, TransitionKey.COMPLEMENTARY_DATA: {"language": ["pick up the vial"]}, } diff --git a/tests/policies/groot/test_groot_train_random_crop.py b/tests/policies/groot/test_groot_train_random_crop.py index adef2958b59..50499c34161 100644 --- a/tests/policies/groot/test_groot_train_random_crop.py +++ b/tests/policies/groot/test_groot_train_random_crop.py @@ -106,7 +106,8 @@ def crop_at(position): def _video(img, views=2): - return np.stack([img] * views, axis=0).reshape(1, 1, views, *img.shape) + frame = torch.from_numpy(img).permute(2, 0, 1).unsqueeze(0).unsqueeze(0) + return tuple(frame.clone() for _ in range(views)) def _step(training): @@ -121,31 +122,23 @@ def _step(training): def test_training_crop_replays_one_window_across_views(): video = _video(_structured_image()) - frames = _step(training=True)._build_sample_images(video, batch_size=1, target_device=None)[0] + frames = _step(training=True)._build_sample_images(video, target_device=None)[0] np.testing.assert_array_equal(np.asarray(frames[0]), np.asarray(frames[1])) def test_training_crop_differs_from_eval_center_crop(): video = _video(_structured_image()) random.seed(3) # a draw that is not the exact center - train_frame = np.asarray( - _step(training=True)._build_sample_images(video, batch_size=1, target_device=None)[0][0] - ) - eval_frame = np.asarray( - _step(training=False)._build_sample_images(video, batch_size=1, target_device=None)[0][0] - ) + train_frame = np.asarray(_step(training=True)._build_sample_images(video, target_device=None)[0][0]) + eval_frame = np.asarray(_step(training=False)._build_sample_images(video, target_device=None)[0][0]) assert not np.array_equal(train_frame, eval_frame) def test_training_crop_is_disabled_under_no_grad(): video = _video(_structured_image()) with torch.no_grad(): - no_grad_frame = np.asarray( - _step(training=True)._build_sample_images(video, batch_size=1, target_device=None)[0][0] - ) - eval_frame = np.asarray( - _step(training=False)._build_sample_images(video, batch_size=1, target_device=None)[0][0] - ) + no_grad_frame = np.asarray(_step(training=True)._build_sample_images(video, target_device=None)[0][0]) + eval_frame = np.asarray(_step(training=False)._build_sample_images(video, target_device=None)[0][0]) np.testing.assert_array_equal(no_grad_frame, eval_frame) @@ -162,8 +155,6 @@ def test_training_crop_respects_global_seed(): def draw(): random.seed(11) - return np.asarray( - _step(training=True)._build_sample_images(video, batch_size=1, target_device=None)[0][0] - ) + return np.asarray(_step(training=True)._build_sample_images(video, target_device=None)[0][0]) np.testing.assert_array_equal(draw(), draw()) diff --git a/tests/training/test_visual_validation.py b/tests/training/test_visual_validation.py index 1df8006b273..efc86553070 100644 --- a/tests/training/test_visual_validation.py +++ b/tests/training/test_visual_validation.py @@ -30,6 +30,7 @@ import numpy as np import pytest +import torch pytest.importorskip("datasets", reason="datasets is required (install lerobot[dataset])") @@ -37,8 +38,9 @@ from lerobot.configs.policies import PreTrainedConfig from lerobot.configs.train import TrainPipelineConfig from lerobot.datasets.lerobot_dataset import LeRobotDataset +from lerobot.policies import ImageInputFormat from lerobot.policies.factory import make_policy_config -from lerobot.scripts.lerobot_train import train +from lerobot.scripts.lerobot_train import prep_camera_images_for_policy, train from lerobot.utils.device_utils import auto_select_torch_device pytest.importorskip("transformers") @@ -57,6 +59,32 @@ def temp_dir(tmp_path): DEVICE = auto_select_torch_device() +def test_prep_camera_images_for_policy_obeys_policy_contract(): + image = torch.tensor([0, 127, 255], dtype=torch.uint8) + + float_batch = {"observation.images.front": image.clone()} + prep_camera_images_for_policy( + float_batch, + ["observation.images.front"], + ImageInputFormat.FLOAT32_0_1, + ) + assert float_batch["observation.images.front"].dtype == torch.float32 + torch.testing.assert_close( + float_batch["observation.images.front"], + torch.tensor([0.0, 127.0 / 255.0, 1.0]), + ) + + uint8_image = image.clone() + uint8_batch = {"observation.images.front": uint8_image} + prep_camera_images_for_policy( + uint8_batch, + ["observation.images.front"], + ImageInputFormat.UINT8_0_255, + ) + assert uint8_batch["observation.images.front"].data_ptr() == uint8_image.data_ptr() + assert uint8_batch["observation.images.front"].dtype == torch.uint8 + + def make_dummy_dataset(camera_keys, tmp_path): """Creates a minimal dummy dataset for testing rename_mapping logic.""" features = {