diff --git a/src/physicalai/inference/component_factory.py b/src/physicalai/inference/component_factory.py index 8bb0ceef..d4256d07 100644 --- a/src/physicalai/inference/component_factory.py +++ b/src/physicalai/inference/component_factory.py @@ -110,6 +110,7 @@ def __repr__(self) -> str: component_registry.register("hf_tokenizer", "physicalai.inference.preprocessors.HFTokenizer") component_registry.register("ov_tokenizer", "physicalai.inference.preprocessors.OVTokenizer") component_registry.register("pi05", "physicalai.inference.preprocessors.Pi05Preprocessor") +component_registry.register("xr0", "physicalai.inference.preprocessors.XR0Preprocessor") component_registry.register("rldx1", "physicalai.inference.preprocessors.Rldx1Preprocessor") component_registry.register("rldx1_token_composer", "physicalai.inference.preprocessors.Rldx1TokenComposer") component_registry.register("rldx1_rope", "physicalai.inference.preprocessors.Rldx1RopePreprocessor") @@ -117,6 +118,7 @@ def __repr__(self) -> str: # Postprocessors component_registry.register("denormalize", "physicalai.inference.postprocessors.StatsDenormalizer") +component_registry.register("xr0_denormalize", "physicalai.inference.postprocessors.XR0Postprocessor") component_registry.register("action_chunk_trimmer", "physicalai.inference.postprocessors.ActionChunkTrimmer") diff --git a/src/physicalai/inference/constants.py b/src/physicalai/inference/constants.py index e28bc5f6..dceb315b 100644 --- a/src/physicalai/inference/constants.py +++ b/src/physicalai/inference/constants.py @@ -11,6 +11,7 @@ ACTION = "action" TASK = "task" STATE = "state" +STATE_PASSTHROUGH = "state_passthrough" TOKENIZED_PROMPT = "tokenized_prompt" TOKENIZED_PROMPT_MASK = "tokenized_prompt_mask" @@ -31,6 +32,7 @@ "RTC_INFERENCE_DELAY", "RTC_MAX_GUIDANCE_WEIGHT", "STATE", + "STATE_PASSTHROUGH", "TASK", "TOKENIZED_PROMPT", "TOKENIZED_PROMPT_MASK", diff --git a/src/physicalai/inference/postprocessors/__init__.py b/src/physicalai/inference/postprocessors/__init__.py index 64a45d42..69e4391a 100644 --- a/src/physicalai/inference/postprocessors/__init__.py +++ b/src/physicalai/inference/postprocessors/__init__.py @@ -10,10 +10,12 @@ from physicalai.inference.postprocessors.action_normalizer import ActionNormalizer from physicalai.inference.postprocessors.base import Postprocessor from physicalai.inference.postprocessors.stats_denormalizer import StatsDenormalizer +from physicalai.inference.postprocessors.xr0 import XR0Postprocessor __all__ = [ "ActionChunkTrimmer", "ActionNormalizer", "Postprocessor", "StatsDenormalizer", + "XR0Postprocessor", ] diff --git a/src/physicalai/inference/postprocessors/xr0.py b/src/physicalai/inference/postprocessors/xr0.py new file mode 100644 index 00000000..5b59d8cb --- /dev/null +++ b/src/physicalai/inference/postprocessors/xr0.py @@ -0,0 +1,133 @@ +# Copyright (C) 2026 Intel Corporation +# SPDX-License-Identifier: Apache-2.0 + +"""NumPy postprocessor for the exported XR0 model.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +import numpy as np +from typing_extensions import override + +from physicalai.inference.constants import ACTION, STATE_PASSTHROUGH +from physicalai.inference.postprocessors.base import Postprocessor + +if TYPE_CHECKING: + from collections.abc import Sequence + +# Numerical epsilon added to the action std (matches the training convention). +_ACTION_EPS = 1e-6 + + +class XR0Postprocessor(Postprocessor): + """Denormalize the exported XR0 graph's action output. + + Inverts the source action normalization (``action * (std + eps) + mean``) and + slices the padded action back to its real dimension, mirroring the training + ``XR0Postprocessor`` in physicalai-train. + + In ``action_mode="delta"`` the denormalized output is a delta and the + current-frame ``state`` (a second graph output) is re-added on the + overlapping leading channels before slicing to the real action dimension. + + Args: + action_mean: Per-dimension action mean (padded to ``max_action_dim``). + Per-timestep ``(chunk_size, max_action_dim)`` in delta mode. + action_std: Per-dimension action std (padded to ``max_action_dim``). + Per-timestep ``(chunk_size, max_action_dim)`` in delta mode. + action_dim: Real (unpadded) action dimension; when set, the output is + sliced to it. ``None`` keeps the padded width. + action_mode: ``"absolute"`` (default) returns the denormalized action; + ``"delta"`` re-adds the current-frame ``state`` to the denormalized + delta before slicing. + action_eps: Numerical epsilon added to the std (matches training). + + Examples: + Constructed via manifest (type-based resolution):: + + {"type": "xr0_denormalize", "action_mean": [...], + "action_std": [...], "action_dim": 30, "action_mode": "delta"} + """ + + def __init__( + self, + action_mean: Sequence[float] | Sequence[Sequence[float]], + action_std: Sequence[float] | Sequence[Sequence[float]], + action_dim: int | None = None, + action_mode: str = "absolute", + action_eps: float = _ACTION_EPS, + ) -> None: + """Initialize the XR0 inference postprocessor. + + Raises: + ValueError: If ``action_mean`` and ``action_std`` shapes differ, or + if ``action_mode`` is not a recognized mode. + """ + super().__init__() + self._mean = np.asarray(action_mean, dtype=np.float32) + self._std = np.asarray(action_std, dtype=np.float32) + if self._mean.shape != self._std.shape: + msg = f"action_mean {self._mean.shape} and action_std {self._std.shape} must have the same shape" + raise ValueError(msg) + valid_modes = {"absolute", "delta"} + if action_mode not in valid_modes: + msg = f"Unknown action_mode {action_mode!r}. Expected one of {sorted(valid_modes)}" + raise ValueError(msg) + self._action_dim = int(action_dim) if action_dim is not None else None + self._action_mode = str(action_mode) + self._eps = float(action_eps) + + @override + def __call__(self, outputs: dict[str, np.ndarray]) -> dict[str, np.ndarray]: + """Denormalize and unpad the predicted action chunk. + + Args: + outputs: Runner output dict containing an ``action`` array (and, in + delta mode, the current-frame ``state`` echoed by the graph). + + Returns: + The outputs dict with the ``action`` denormalized (and sliced to the + real action dimension when known). In delta mode the current-frame + state is re-added to reconstruct the absolute action. + """ + action = outputs.get(ACTION) + if action is None: + return outputs + action = np.asarray(action, dtype=np.float32) + # Mirror io.denormalize_action: action * (std + eps) + mean. + action = action * (self._std + self._eps) + self._mean + if self._action_mode == "delta": + action = self._add_current_state(action, outputs.get(STATE_PASSTHROUGH)) + if self._action_dim is not None: + action = action[..., : self._action_dim] + result = dict(outputs) + result[ACTION] = action + return result + + @staticmethod + def _add_current_state(action: np.ndarray, state: np.ndarray | None) -> np.ndarray: + """Re-add the current-frame state to a denormalized delta action. + + Mirrors the Studio-side ``XR0Postprocessor.forward`` delta inverse: the + current (last) state frame is added on the overlapping leading channels, + broadcast over the action chunk. + + Returns: + The reconstructed absolute action. + + Raises: + ValueError: If ``state`` is missing. + """ + if state is None: + msg = "action_mode='delta' requires the graph's 'state_passthrough' output to invert the delta prediction." + raise ValueError(msg) + current = np.asarray(state, dtype=np.float32) + if current.ndim == 3: # (B, T, D) -> current (last) frame # noqa: PLR2004 + current = current[:, -1, :] + overlap = min(action.shape[-1], current.shape[-1]) + # Insert a chunk axis so (..., overlap) broadcasts over the action chunk. + current = np.expand_dims(current[..., :overlap], axis=-2) + action = action.copy() + action[..., :overlap] += current + return action diff --git a/src/physicalai/inference/preprocessors/__init__.py b/src/physicalai/inference/preprocessors/__init__.py index c006b908..b5dec018 100644 --- a/src/physicalai/inference/preprocessors/__init__.py +++ b/src/physicalai/inference/preprocessors/__init__.py @@ -20,6 +20,7 @@ from physicalai.inference.preprocessors.smolvla import ResizeSmolVLA from physicalai.inference.preprocessors.stats_normalizer import StatsNormalizer from physicalai.inference.preprocessors.to_tensor import ToFloatTensorPreprocessor +from physicalai.inference.preprocessors.xr0 import XR0Preprocessor __all__ = [ "HFTokenizer", @@ -36,4 +37,5 @@ "Rldx1TokenComposer", "StatsNormalizer", "ToFloatTensorPreprocessor", + "XR0Preprocessor", ] diff --git a/src/physicalai/inference/preprocessors/xr0.py b/src/physicalai/inference/preprocessors/xr0.py new file mode 100644 index 00000000..2632992a --- /dev/null +++ b/src/physicalai/inference/preprocessors/xr0.py @@ -0,0 +1,435 @@ +# Copyright (C) 2026 Intel Corporation +# SPDX-License-Identifier: Apache-2.0 + +"""NumPy preprocessor for the XR0 model.""" + +from __future__ import annotations + +import math +from typing import TYPE_CHECKING + +import cv2 +import numpy as np +from typing_extensions import override + +from physicalai.inference.constants import IMAGES, STATE, TASK +from physicalai.inference.preprocessors.base import Preprocessor + +if TYPE_CHECKING: + from collections.abc import Sequence + +# Qwen3-VL image-normalization + geometry constants. These are baked into the +# manifest ``init_args`` at export time from the source image processor; the +# defaults mirror ``Qwen/Qwen3-VL-4B-Instruct`` so the component is usable +# standalone. +_QWEN3VL_IMAGE_MEAN = (0.5, 0.5, 0.5) +_QWEN3VL_IMAGE_STD = (0.5, 0.5, 0.5) +_QWEN3VL_RESCALE_FACTOR = 1.0 / 255.0 +_QWEN3VL_PATCH_SIZE = 16 +_QWEN3VL_MERGE_SIZE = 2 +_QWEN3VL_TEMPORAL_PATCH_SIZE = 2 + +# Numerical epsilon added to the state std (matches the training convention). +_STATE_EPS = 1e-6 + +_TEMPORAL_IMAGE_NDIM = 5 +_BATCHED_IMAGE_NDIM = 4 +_CHANNELS_FIRST_NDIM = 3 +_TEMPORAL_STATE_NDIM = 3 + +# Reject images whose aspect ratio exceeds this (matches the source preprocessor). +_MAX_ASPECT_RATIO = 200 + +# --- prompt text (mirrors the Qwen3-VL processor chat template) ------------ +_MULTI_VIEW_HEADER = "The following observations are captured from multiple views.\n" +_TASK_TEMPLATE = "Generate robot actions for the task:\n{instruction} /no_cot" +_ASSISTANT_PRIMER = "" + +# View titles the model was trained with, e.g. +# "wrist_left" -> "Left-Wrist" so the prompt reads "# Left-Wrist View". +_VIEW_TITLES = { + "base": "Base", + "wrist_left": "Left-Wrist", + "wrist_right": "Right-Wrist", +} + +# Qwen3-VL chat special tokens rendered as literal text (a bare tokenizer maps +# each to its dedicated special-token id). +_IM_START = "<|im_start|>" +_IM_END = "<|im_end|>" +_VISION_START = "<|vision_start|>" +_VISION_END = "<|vision_end|>" +_IMAGE_PAD = "<|image_pad|>" + + +def _view_title(view: str) -> str: + """Human-readable view title matching the reference prompt. + + Returns: + The human-readable view title. + """ + key = view.replace("-", "_") + if key in _VIEW_TITLES: + return _VIEW_TITLES[key] + return " ".join(word.capitalize() for word in key.split("_")) + + +def _image_pad_count(grid_t: int, grid_h: int, grid_w: int, merge_size: int) -> int: + """Number of ``<|image_pad|>`` tokens the processor expands one image into. + + Returns: + The image-pad token count for a single image. + """ + return (grid_t * grid_h * grid_w) // (merge_size * merge_size) + + +def _render_chat_prompt(views: Sequence[str], pad_counts: Sequence[int], instruction: str) -> str: + """Render the XR0 Qwen3-VL chat prompt as a raw string. + + Reproduces ``processor.apply_chat_template(..., tokenize=False)`` for the XR0 + multi-view message, with each image's ``<|image_pad|>`` already expanded to + ``pad_counts[i]`` copies. + + Returns: + The fully-rendered chat prompt string. + + Raises: + ValueError: If ``views`` and ``pad_counts`` have different lengths. + """ + if len(views) != len(pad_counts): + msg = f"views ({len(views)}) and pad_counts ({len(pad_counts)}) must have the same length" + raise ValueError(msg) + + parts: list[str] = [_MULTI_VIEW_HEADER] + for view, count in zip(views, pad_counts, strict=True): + parts.extend((f"# {_view_title(view)} View\n", _VISION_START + _IMAGE_PAD * count + _VISION_END, "\n")) + parts.append(_TASK_TEMPLATE.format(instruction=instruction)) + user = "".join(parts) + return f"{_IM_START}user\n{user}{_IM_END}\n{_IM_START}assistant\n{_ASSISTANT_PRIMER}{_IM_END}\n" + + +def _resize_image(image: np.ndarray, factor: int, max_pixels: int) -> np.ndarray: + """Resize an ``(H, W, C)`` uint8 image to patch-aligned dims within an area budget. + + Both sides are rounded to multiples of ``factor`` and the area is kept within + ``[factor**2, max_pixels]``, preserving aspect ratio for the VLM vision encoder. + + Returns: + The resized ``(H, W, C)`` uint8 image. + + Raises: + ValueError: If the image aspect ratio exceeds ``_MAX_ASPECT_RATIO``. + """ + min_pixels = factor * factor + height, width = image.shape[:2] + ratio = max(height, width) / min(height, width) + if ratio > _MAX_ASPECT_RATIO: + msg = f"absolute aspect ratio must be smaller than 200, got {ratio}" + raise ValueError(msg) + + new_height = max(factor, round(height / factor) * factor) + new_width = max(factor, round(width / factor) * factor) + + if new_height * new_width > max_pixels: + scale = math.sqrt(height * width / max_pixels) + new_height = max(factor, math.floor(height / scale / factor) * factor) + new_width = max(factor, math.floor(width / scale / factor) * factor) + elif new_height * new_width < min_pixels: + scale = math.sqrt(min_pixels / (height * width)) + new_height = max(factor, math.ceil(height * scale / factor) * factor) + new_width = max(factor, math.ceil(width * scale / factor) * factor) + + # cv2.resize takes dsize as (width, height); INTER_CUBIC is the closest match to + # the reference PIL bicubic resample (exact parity is not required here). + return cv2.resize(image, (new_width, new_height), interpolation=cv2.INTER_CUBIC) + + +def _build_pixel_grid( + images: Sequence[np.ndarray], + image_mean: Sequence[float], + image_std: Sequence[float], + rescale_factor: float, +) -> np.ndarray: + """Rescale + normalize already-resized images into a Qwen3-VL pixel grid. + + Returns: + The normalized image grid of shape ``(num_images, C, H, W)`` as float32. + """ + mean = np.asarray(image_mean, dtype=np.float32) + std = np.asarray(image_std, dtype=np.float32) + grid = [ + np.transpose((np.asarray(image, dtype=np.float32) * np.float32(rescale_factor) - mean) / std, (2, 0, 1)) + for image in images + ] + return np.stack(grid).astype(np.float32) + + +def _patchify_pixel_grid( + grid: np.ndarray, + grid_thw: Sequence[Sequence[int]], + *, + temporal_patch_size: int, + patch_size: int, + merge_size: int, +) -> np.ndarray: + """Patchify a normalized image grid exactly like the Qwen3-VL image processor. + + Reproduces the transformers ``Qwen2VLImageProcessor`` patchify (temporal + duplication + 9-D reshape/transpose) so the exported graph receives the flat + ``pixel_values`` layout directly (patchify happens off-graph here). + + Returns: + The flat ``pixel_values`` array of shape + ``(sum(grid_t * grid_h * grid_w), C * temporal_patch_size * patch_size ** 2)`` + as float32. + """ + flattened: list[np.ndarray] = [] + for index, (grid_t, grid_h, grid_w) in enumerate(grid_thw): + image = grid[index] # (C, H, W) + channel = image.shape[0] + feature = channel * temporal_patch_size * patch_size * patch_size + patches = np.tile(image[np.newaxis], (temporal_patch_size, 1, 1, 1)) # (tp, C, H, W) + patches = patches.reshape( + grid_t, + temporal_patch_size, + channel, + grid_h // merge_size, + merge_size, + patch_size, + grid_w // merge_size, + merge_size, + patch_size, + ) + patches = patches.transpose(0, 3, 6, 4, 7, 2, 1, 5, 8) + flattened.append(patches.reshape(grid_t * grid_h * grid_w, feature)) + return np.concatenate(flattened, axis=0).astype(np.float32) + + +def _to_rgb_frame(array: object) -> np.ndarray: + """Normalize a NumPy image (``(H,W,C)`` / ``(C,H,W)`` / batched / temporal) to uint8 RGB. + + Collapses arbitrary incoming layouts/dtypes to a single canonical ``(H, W, 3)`` + uint8 RGB frame so the subsequent ``cv2.resize`` sees a consistent input. The + exported graph bakes image geometry from the resized dimensions, so only the + output size has to match the reference; the interpolation kernel does not. + + Returns: + The canonical ``(H, W, 3)`` uint8 RGB frame. + """ + arr = np.asarray(array) + if arr.ndim == _TEMPORAL_IMAGE_NDIM: # (B, T, C, H, W) -> last frame of first sample + arr = arr[0, -1] + elif arr.ndim == _BATCHED_IMAGE_NDIM: # (T|B, C, H, W) -> last frame + arr = arr[-1] + if arr.ndim == _CHANNELS_FIRST_NDIM and arr.shape[0] in {1, 3}: # channels-first + arr = np.transpose(arr, (1, 2, 0)) + if arr.dtype != np.uint8: + arr = (np.clip(arr, 0.0, 1.0) * 255.0).round().astype(np.uint8) + if arr.shape[-1] == 1: + arr = np.repeat(arr, 3, axis=-1) + return np.ascontiguousarray(arr) + + +class XR0Preprocessor(Preprocessor): + """Build the exported XR0 graph inputs from a raw observation dict. + + Lightweight, torch-free NumPy preprocessor: resizes the camera views into the + Qwen3-VL ``pixel_values`` grid, pads/normalizes the ``state`` and renders the + multi-view chat prompt as a plain ``task`` string. + + Args: + max_state_dim: State dimension after padding. + image_factor: Patch-alignment factor for image resizing. + image_max_pixels: Maximum image area for image resizing. + image_mean: Per-channel image mean (baked from the source image processor). + image_std: Per-channel image std (baked from the source image processor). + rescale_factor: Pixel rescale factor (``1/255`` for Qwen3-VL). + patch_size: Vision patch size used to derive the ``<|image_pad|>`` count. + merge_size: Spatial merge size used to derive the ``<|image_pad|>`` count. + temporal_patch_size: Number of frames grouped per temporal patch (the + off-graph patchify duplicates a still image to this many frames). + normalize_state: Whether the exported model expects normalized state. + Defaults to False (raw state), matching the training default. + state_mean: Baked ``max_state_dim`` state mean (identity when disabled). + state_std: Baked ``max_state_dim`` state std (identity when disabled). + + Examples: + Constructed via manifest (type-based resolution):: + + {"type": "xr0", "max_state_dim": 32, "patch_size": 16, "merge_size": 2} + """ + + def __init__( + self, + max_state_dim: int = 32, + image_factor: int = 32, + image_max_pixels: int = 90000, + image_mean: Sequence[float] = _QWEN3VL_IMAGE_MEAN, + image_std: Sequence[float] = _QWEN3VL_IMAGE_STD, + rescale_factor: float = _QWEN3VL_RESCALE_FACTOR, + patch_size: int = _QWEN3VL_PATCH_SIZE, + merge_size: int = _QWEN3VL_MERGE_SIZE, + temporal_patch_size: int = _QWEN3VL_TEMPORAL_PATCH_SIZE, + *, + normalize_state: bool = False, + state_mean: Sequence[float] | None = None, + state_std: Sequence[float] | None = None, + ) -> None: + """Initialize the XR0 inference preprocessor. + + Raises: + ValueError: If ``patch_size`` / ``merge_size`` is not positive. + """ + super().__init__() + if int(patch_size) <= 0 or int(merge_size) <= 0: + msg = f"patch_size and merge_size must be positive, got {patch_size!r} / {merge_size!r}" + raise ValueError(msg) + + self._max_state_dim = int(max_state_dim) + self._image_factor = int(image_factor) + self._image_max_pixels = int(image_max_pixels) + self._image_mean = tuple(float(v) for v in image_mean) + self._image_std = tuple(float(v) for v in image_std) + self._rescale_factor = float(rescale_factor) + self._patch_size = int(patch_size) + self._merge_size = int(merge_size) + self._temporal_patch_size = int(temporal_patch_size) + self._normalize_state = bool(normalize_state) + + # State normalization is opt-in; padded dims use identity stats (mean 0, + # std 1) so they stay zero, mirroring the training preprocessor. + if normalize_state and state_mean is not None and state_std is not None: + self._state_mean = self._pad_state_stat(state_mean, 0.0) + self._state_std = self._pad_state_stat(state_std, 1.0) + else: + self._state_mean = np.zeros(self._max_state_dim, dtype=np.float32) + self._state_std = np.ones(self._max_state_dim, dtype=np.float32) + + def _pad_state_stat(self, values: Sequence[float], fill: float) -> np.ndarray: + """Pad/truncate a state stat to ``max_state_dim`` (padded dims use ``fill``). + + Returns: + The ``(max_state_dim,)`` float32 stat array. + """ + arr = np.asarray(values, dtype=np.float32).flatten() + out = np.full(self._max_state_dim, fill, dtype=np.float32) + dim = min(self._max_state_dim, arr.shape[0]) + out[:dim] = arr[:dim] + return out + + def _extract_images(self, inputs: dict[str, object]) -> tuple[list[str], list[np.ndarray]]: + """Return the ordered view names and resized ``(H, W, C)`` uint8 views. + + The view order is taken directly from the observation image keys + (``images.``) in their natural insertion order, so ``pixel_values`` + stays aligned with the per-view prompt sections (title + pad count). + + Returns: + A ``(views, images)`` tuple: the ordered view names and the resized + uint8 RGB images (one per available camera view). + + Raises: + ValueError: If the observation contains no image entry. + """ + images_value = inputs.get(IMAGES) + if isinstance(images_value, dict): + image_items = {f"{IMAGES}.{view}": array for view, array in images_value.items()} + else: + image_items = { + key: value + for key, value in inputs.items() + if isinstance(key, str) and key.startswith(f"{IMAGES}.") and "is_pad" not in key + } + if not image_items: + msg = "XR0 inference requires at least one image observation" + raise ValueError(msg) + views = [key.removeprefix(f"{IMAGES}.") for key in image_items] + images = [ + _resize_image(_to_rgb_frame(value), factor=self._image_factor, max_pixels=self._image_max_pixels) + for value in image_items.values() + ] + return views, images + + def _prepare_state(self, inputs: dict[str, object]) -> np.ndarray: + """Pad the state into ``(B, 1, max_state_dim)`` (optionally normalized). + + Returns: + The padded ``(B, 1, max_state_dim)`` float32 state array. + + Raises: + ValueError: If the observation has no state entry. + """ + state_value = inputs.get(STATE) + if state_value is None: + msg = "XR0 inference requires a 'state' observation" + raise ValueError(msg) + state = np.asarray(state_value, dtype=np.float32) + if state.ndim == 1: # (D,) -> (1, D) + state = state[None, :] + if state.ndim == _TEMPORAL_STATE_NDIM: # (B, T, D) -> last frame + state = state[:, -1, :] + dim = state.shape[-1] + if dim < self._max_state_dim: + state = np.pad(state, ((0, 0), (0, self._max_state_dim - dim))) + state = state[:, : self._max_state_dim] + if self._normalize_state: + state = (state - self._state_mean) / (self._state_std + _STATE_EPS) + return state[:, None, :].astype(np.float32) # (B, 1, max_state_dim) + + @staticmethod + def _instruction(inputs: dict[str, object]) -> str: + """Extract the task instruction string from the observation. + + Returns: + The (first) task instruction as a string (empty when absent). + """ + task = inputs.get(TASK) + if task is None: + return "" + if isinstance(task, str): + return task + if isinstance(task, np.ndarray): + flat = np.atleast_1d(task).tolist() + return str(flat[0]) if flat else "" + if isinstance(task, (list, tuple)): + return str(task[0]) if task else "" + return str(task) + + @override + def __call__(self, inputs: dict[str, object]) -> dict[str, object]: + """Transform a raw observation into the exported graph inputs. + + Args: + inputs: Observation dict with a ``state`` array, ``images`` (nested + dict or flattened ``images.*`` keys) and a ``task`` string. + + Returns: + Dict with ``pixel_values`` / ``state`` (float32 NumPy) and ``task`` + (a single-element list holding the rendered chat prompt string). + ``pixel_values`` is the flat patchified layout + ``(sum(t*h*w), C * temporal_patch_size * patch_size ** 2)`` the exported + graph consumes directly (patchify happens here, off-graph); the sibling + OpenVINO tokenizer turns ``task`` into the graph's ``tokenized_prompt`` + / ``tokenized_prompt_mask`` inputs. + """ + views, images = self._extract_images(inputs) + pixel_grid = _build_pixel_grid(images, self._image_mean, self._image_std, self._rescale_factor) + grid_thw = [ + (1, image.shape[0] // self._patch_size, image.shape[1] // self._patch_size) # image.shape == (H, W, C) + for image in images + ] + pixel_values = _patchify_pixel_grid( + pixel_grid, + grid_thw, + temporal_patch_size=self._temporal_patch_size, + patch_size=self._patch_size, + merge_size=self._merge_size, + ) + pad_counts = [_image_pad_count(grid_t, grid_h, grid_w, self._merge_size) for grid_t, grid_h, grid_w in grid_thw] + prompt = _render_chat_prompt(views, pad_counts, self._instruction(inputs)) + state = self._prepare_state(inputs) + return { + "pixel_values": np.ascontiguousarray(pixel_values.astype(np.float32)), + "state": np.ascontiguousarray(state), + TASK: [prompt], + } diff --git a/tests/unit/inference/postprocessors/test_xr0.py b/tests/unit/inference/postprocessors/test_xr0.py new file mode 100644 index 00000000..70f77cb0 --- /dev/null +++ b/tests/unit/inference/postprocessors/test_xr0.py @@ -0,0 +1,108 @@ +# Copyright (C) 2026 Intel Corporation +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import numpy as np +import pytest + +from physicalai.inference.constants import ACTION, STATE_PASSTHROUGH +from physicalai.inference.postprocessors import Postprocessor, XR0Postprocessor + + +class TestXR0PostprocessorInit: + def test_is_postprocessor(self) -> None: + post = XR0Postprocessor(action_mean=[0.0] * 4, action_std=[1.0] * 4) + assert isinstance(post, Postprocessor) + + def test_mismatched_stats_raise(self) -> None: + with pytest.raises(ValueError, match="same shape"): + XR0Postprocessor(action_mean=[0.0] * 4, action_std=[1.0] * 3) + + +class TestXR0PostprocessorCall: + def test_denormalizes_action(self) -> None: + post = XR0Postprocessor(action_mean=[1.0] * 4, action_std=[2.0] * 4) + out = post({ACTION: np.zeros((1, 3, 4), dtype=np.float32)}) + # 0 * (2 + eps) + 1 == 1.0 + np.testing.assert_allclose(out[ACTION], 1.0, atol=1e-5) + + def test_slices_to_action_dim(self) -> None: + post = XR0Postprocessor(action_mean=[0.0] * 8, action_std=[1.0] * 8, action_dim=6) + out = post({ACTION: np.ones((1, 5, 8), dtype=np.float32)}) + assert out[ACTION].shape == (1, 5, 6) + + def test_none_action_dim_keeps_width(self) -> None: + post = XR0Postprocessor(action_mean=[0.0] * 8, action_std=[1.0] * 8) + out = post({ACTION: np.ones((1, 5, 8), dtype=np.float32)}) + assert out[ACTION].shape == (1, 5, 8) + + def test_missing_action_passthrough(self) -> None: + post = XR0Postprocessor(action_mean=[0.0] * 4, action_std=[1.0] * 4) + payload = {"other": np.zeros((2,), dtype=np.float32)} + out = post(payload) + assert "other" in out + + def test_roundtrip_with_preprocessor_convention(self) -> None: + mean = np.array([0.5, -1.0, 2.0, 0.0], dtype=np.float32) + std = np.array([1.5, 0.5, 3.0, 1.0], dtype=np.float32) + post = XR0Postprocessor(action_mean=mean.tolist(), action_std=std.tolist()) + raw = np.random.rand(1, 4, 4).astype(np.float32) + normalized = (raw - mean) / (std + 1e-6) + out = post({ACTION: normalized}) + np.testing.assert_allclose(out[ACTION], raw, atol=1e-4) + + +class TestXR0PostprocessorDelta: + """Delta mode re-adds the current-frame state to the denormalized delta.""" + + def _delta_post(self, chunk: int = 5, dim: int = 8, action_dim: int = 6) -> XR0Postprocessor: + # Identity delta stats (mean 0 / std 1) so the graph output is the raw delta. + return XR0Postprocessor( + action_mean=np.zeros((chunk, dim), dtype=np.float32).tolist(), + action_std=np.ones((chunk, dim), dtype=np.float32).tolist(), + action_dim=action_dim, + action_mode="delta", + ) + + def test_adds_current_state(self) -> None: + post = self._delta_post() + delta = np.random.rand(1, 5, 8).astype(np.float32) + state = np.random.rand(1, 1, 8).astype(np.float32) # (B, T=1, D) + out = post({ACTION: delta, STATE_PASSTHROUGH: state}) + # delta + state on the first action_dim=6 channels, sliced to 6. + expected = delta[..., :6] + state[:, -1, :][:, None, :6] + assert out[ACTION].shape == (1, 5, 6) + np.testing.assert_allclose(out[ACTION], expected, atol=1e-5) + + def test_roundtrip_reconstructs_absolute(self) -> None: + mean = np.random.rand(5, 8).astype(np.float32) + std = (np.random.rand(5, 8) + 0.5).astype(np.float32) + post = XR0Postprocessor( + action_mean=mean.tolist(), + action_std=std.tolist(), + action_dim=6, + action_mode="delta", + ) + absolute = np.random.rand(1, 5, 6).astype(np.float32) + state = np.random.rand(1, 1, 8).astype(np.float32) + # Studio-side target: normalized delta = ((absolute - state) - mean) / (std + eps), + # padded to width 8. + delta = absolute - state[:, -1, :][:, None, :6] + padded = np.zeros((1, 5, 8), dtype=np.float32) + padded[..., :6] = delta + normalized = (padded - mean) / (std + 1e-6) + out = post({ACTION: normalized, STATE_PASSTHROUGH: state}) + np.testing.assert_allclose(out[ACTION], absolute, atol=1e-4) + + def test_delta_requires_state(self) -> None: + post = self._delta_post() + with pytest.raises(ValueError, match="state"): + post({ACTION: np.zeros((1, 5, 8), dtype=np.float32)}) + + def test_absolute_mode_ignores_state(self) -> None: + post = XR0Postprocessor(action_mean=[0.0] * 8, action_std=[1.0] * 8, action_dim=6) + action = np.ones((1, 5, 8), dtype=np.float32) + with_state = post({ACTION: action, STATE_PASSTHROUGH: np.ones((1, 1, 8), dtype=np.float32)})[ACTION] + without_state = post({ACTION: action})[ACTION] + np.testing.assert_allclose(with_state, without_state, atol=1e-6) diff --git a/tests/unit/inference/preprocessors/test_xr0.py b/tests/unit/inference/preprocessors/test_xr0.py new file mode 100644 index 00000000..ccf9cc33 --- /dev/null +++ b/tests/unit/inference/preprocessors/test_xr0.py @@ -0,0 +1,399 @@ +# Copyright (C) 2026 Intel Corporation +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import numpy as np +import pytest + +from physicalai.inference.constants import IMAGES, STATE, TASK +from physicalai.inference.preprocessors import Preprocessor, XR0Preprocessor +from physicalai.inference.preprocessors.xr0 import _build_pixel_grid, _render_chat_prompt + + +@pytest.fixture() +def preprocessor(): + return XR0Preprocessor( + max_state_dim=32, + image_factor=32, + image_max_pixels=90000, + image_mean=(0.5, 0.5, 0.5), + image_std=(0.5, 0.5, 0.5), + rescale_factor=1.0 / 255.0, + patch_size=16, + merge_size=2, + ) + + +def _make_inputs(h: int = 256, w: int = 256, state_dim: int = 8) -> dict[str, object]: + """Build a minimal XR0 observation dict (raw uint8 images, 1-D state).""" + return { + IMAGES: { + "base": np.zeros((h, w, 3), dtype=np.uint8), + "wrist_left": np.zeros((h, w, 3), dtype=np.uint8), + }, + STATE: np.zeros((state_dim,), dtype=np.float32), + TASK: "pick up the cup", + } + + +class TestXR0PreprocessorInit: + def test_is_preprocessor(self, preprocessor) -> None: + assert isinstance(preprocessor, Preprocessor) + + def test_nonpositive_patch_size_raises(self) -> None: + with pytest.raises(ValueError, match="must be positive"): + XR0Preprocessor(patch_size=0) + + +class TestXR0PreprocessorOutput: + def test_output_keys(self, preprocessor) -> None: + result = preprocessor(_make_inputs()) + assert set(result) == {"pixel_values", "state", TASK} + + def test_pixel_grid_shape_and_dtype(self, preprocessor) -> None: + result = preprocessor(_make_inputs()) + # Flat patchified pixel_values: 2 views at 256x256 -> grid_h=grid_w=16, + # 256 tokens/image, 2 images -> 512 rows; feature = C * tp * patch**2. + assert result["pixel_values"].shape == (512, 3 * 2 * 16 * 16) + assert result["pixel_values"].dtype == np.float32 + + def test_pixel_grid_normalized(self, preprocessor) -> None: + # Zero uint8 pixels -> (0 * 1/255 - 0.5) / 0.5 == -1.0 for every channel. + result = preprocessor(_make_inputs()) + np.testing.assert_allclose(result["pixel_values"], -1.0, atol=1e-5) + + def test_state_padded_shape(self, preprocessor) -> None: + result = preprocessor(_make_inputs(state_dim=8)) + assert result["state"].shape == (1, 1, 32) + assert result["state"].dtype == np.float32 + + def test_state_zero_padding(self, preprocessor) -> None: + inputs = _make_inputs(state_dim=8) + inputs[STATE] = np.arange(8, dtype=np.float32) + result = preprocessor(inputs) + np.testing.assert_allclose(result["state"][0, 0, :8], np.arange(8)) + np.testing.assert_allclose(result["state"][0, 0, 8:], 0.0) + + +class TestXR0PreprocessorPrompt: + def test_task_is_single_element_list(self, preprocessor) -> None: + result = preprocessor(_make_inputs()) + assert isinstance(result[TASK], list) + assert len(result[TASK]) == 1 + + def test_prompt_structure(self, preprocessor) -> None: + prompt = preprocessor(_make_inputs())[TASK][0] + assert prompt.startswith("<|im_start|>user\n") + assert prompt.endswith("<|im_start|>assistant\n<|im_end|>\n") + assert "pick up the cup" in prompt + assert "# Base View" in prompt + assert "# Left-Wrist View" in prompt + + def test_image_pad_token_count(self, preprocessor) -> None: + # 256/16 = 16 patches per side; merged 2x2 -> (16*16)/(2*2) = 64 pads/image. + prompt = preprocessor(_make_inputs())[TASK][0] + assert prompt.count("<|image_pad|>") == 64 * 2 + + def test_missing_task_defaults_empty(self, preprocessor) -> None: + inputs = _make_inputs() + inputs.pop(TASK) + prompt = preprocessor(inputs)[TASK][0] + assert "<|im_start|>user" in prompt + + def test_task_list_uses_first(self, preprocessor) -> None: + inputs = _make_inputs() + inputs[TASK] = ["open the drawer", "ignored"] + prompt = preprocessor(inputs)[TASK][0] + assert "open the drawer" in prompt + assert "ignored" not in prompt + + +class TestXR0PreprocessorState: + def test_3d_state_uses_last_timestep(self, preprocessor) -> None: + inputs = _make_inputs(state_dim=8) + state = np.zeros((1, 5, 8), dtype=np.float32) + state[0, -1] = np.arange(8, dtype=np.float32) + inputs[STATE] = state + result = preprocessor(inputs) + np.testing.assert_allclose(result["state"][0, 0, :8], np.arange(8)) + + def test_normalized_state(self) -> None: + prep = XR0Preprocessor( + max_state_dim=32, + patch_size=16, + merge_size=2, + normalize_state=True, + state_mean=[1.0] * 8, + state_std=[2.0] * 8, + ) + inputs = _make_inputs(state_dim=8) + inputs[STATE] = np.ones(8, dtype=np.float32) * 3.0 + result = prep(inputs) + state = np.asarray(result["state"]) + # (3 - 1) / (2 + 1e-6) ~= 1.0 on the real dims; padded dims stay 0. + np.testing.assert_allclose(state[0, 0, :8], 1.0, atol=1e-5) + np.testing.assert_allclose(state[0, 0, 8:], 0.0, atol=1e-5) + + +class TestXR0PreprocessorExtractImages: + def test_nested_dict_returns_array_per_view(self, preprocessor) -> None: + views, images = preprocessor._extract_images(_make_inputs()) + assert views == ["base", "wrist_left"] + assert len(images) == 2 + assert all(isinstance(image, np.ndarray) for image in images) + assert all(image.dtype == np.uint8 and image.shape[-1] == 3 for image in images) + + def test_resized_to_patch_aligned(self, preprocessor) -> None: + # 256 is a multiple of factor=32 and area 65536 < 90000 -> unchanged. + _, images = preprocessor._extract_images(_make_inputs(h=256, w=256)) + assert all(image.shape == (256, 256, 3) for image in images) + + def test_odd_size_rounded_to_factor(self, preprocessor) -> None: + # 250 rounds to nearest multiple of factor=32 -> 256. + _, images = preprocessor._extract_images(_make_inputs(h=250, w=250)) + assert all(image.shape == (256, 256, 3) for image in images) + + def test_flattened_keys(self, preprocessor) -> None: + inputs = { + f"{IMAGES}.base": np.zeros((256, 256, 3), dtype=np.uint8), + f"{IMAGES}.wrist_left": np.zeros((256, 256, 3), dtype=np.uint8), + } + views, images = preprocessor._extract_images(inputs) + assert views == ["base", "wrist_left"] + assert len(images) == 2 + + def test_flattened_keys_skip_is_pad(self, preprocessor) -> None: + inputs = { + f"{IMAGES}.base": np.zeros((256, 256, 3), dtype=np.uint8), + f"{IMAGES}.base.is_pad": np.zeros((1,), dtype=bool), + } + _, images = preprocessor._extract_images(inputs) + assert len(images) == 1 + + def test_all_views_used(self, preprocessor) -> None: + # Every provided view is used, in the observation's insertion order. + inputs = { + IMAGES: { + "base": np.zeros((256, 256, 3), dtype=np.uint8), + "wrist_left": np.zeros((256, 256, 3), dtype=np.uint8), + "wrist_right": np.zeros((256, 256, 3), dtype=np.uint8), + }, + } + views, images = preprocessor._extract_images(inputs) + assert views == ["base", "wrist_left", "wrist_right"] + assert len(images) == 3 + + def test_selected_in_observation_order(self) -> None: + # Views follow the observation key insertion order, so pixel_values stay + # aligned with the prompt sections. + prep = XR0Preprocessor(max_state_dim=32) + base = np.zeros((256, 256, 3), dtype=np.uint8) + wrist = np.full((256, 256, 3), 255, dtype=np.uint8) + inputs: dict[str, object] = {IMAGES: {"wrist_left": wrist, "base": base}} + views, images = prep._extract_images(inputs) + # First image is wrist_left (all 255), second is base (all 0). + assert views == ["wrist_left", "base"] + assert images[0].max() == 255 + assert images[1].max() == 0 + + def test_arbitrary_keys_used_in_order(self, preprocessor) -> None: + # Non-reference view names are used as-is in insertion order. + inputs = { + f"{IMAGES}.camA": np.zeros((256, 256, 3), dtype=np.uint8), + f"{IMAGES}.camB": np.zeros((256, 256, 3), dtype=np.uint8), + } + views, images = preprocessor._extract_images(inputs) + assert views == ["camA", "camB"] + assert len(images) == 2 + + def test_no_images_raises(self, preprocessor) -> None: + with pytest.raises(ValueError, match="at least one image"): + preprocessor._extract_images({STATE: np.zeros((8,), dtype=np.float32)}) + + +class TestBuildPixelGrid: + def test_shape_and_dtype(self) -> None: + images = [np.zeros((32, 48, 3), dtype=np.uint8) for _ in range(2)] + grid = _build_pixel_grid(images, (0.5, 0.5, 0.5), (0.5, 0.5, 0.5), 1.0 / 255.0) + # (num_images, C, H, W) -- channels-first from (H, W, C) input. + assert grid.shape == (2, 3, 32, 48) + assert grid.dtype == np.float32 + + def test_zero_pixels_normalized(self) -> None: + images = [np.zeros((16, 16, 3), dtype=np.uint8)] + grid = _build_pixel_grid(images, (0.5, 0.5, 0.5), (0.5, 0.5, 0.5), 1.0 / 255.0) + # (0 * 1/255 - 0.5) / 0.5 == -1.0 for every channel. + np.testing.assert_allclose(grid, -1.0, atol=1e-5) + + def test_per_channel_stats(self) -> None: + # Distinct per-channel value + stats verify the mean/std are applied per channel. + arr = np.zeros((2, 2, 3), dtype=np.uint8) + arr[..., 0], arr[..., 1], arr[..., 2] = 10, 20, 30 + grid = _build_pixel_grid([arr], (1.0, 2.0, 3.0), (2.0, 4.0, 6.0), 1.0) + # (value * 1.0 - mean) / std per channel. + np.testing.assert_allclose(grid[0, 0], (10 - 1.0) / 2.0, atol=1e-5) + np.testing.assert_allclose(grid[0, 1], (20 - 2.0) / 4.0, atol=1e-5) + np.testing.assert_allclose(grid[0, 2], (30 - 3.0) / 6.0, atol=1e-5) + + def test_hwc_to_chw_transpose(self) -> None: + # Unique value per pixel/channel so a wrong transpose would be detected. + arr = np.arange(2 * 3 * 3, dtype=np.uint8).reshape(2, 3, 3) # (H, W, C) + grid = _build_pixel_grid([arr], (0.0, 0.0, 0.0), (1.0, 1.0, 1.0), 1.0) + np.testing.assert_allclose(grid[0], np.transpose(arr.astype(np.float32), (2, 0, 1))) + + +class TestRenderChatPrompt: + def test_chat_envelope(self) -> None: + prompt = _render_chat_prompt(["base"], [4], "pick up the cup") + assert prompt.startswith("<|im_start|>user\n") + assert prompt.endswith("<|im_start|>assistant\n<|im_end|>\n") + + def test_multi_view_header_and_titles(self) -> None: + prompt = _render_chat_prompt(["base", "wrist_left"], [1, 1], "task") + assert "The following observations are captured from multiple views.\n" in prompt + assert "# Base View" in prompt + assert "# Left-Wrist View" in prompt + # View order preserved: base section precedes wrist-left section. + assert prompt.index("# Base View") < prompt.index("# Left-Wrist View") + + def test_unknown_view_title_capitalized(self) -> None: + prompt = _render_chat_prompt(["top_down"], [1], "task") + assert "# Top Down View" in prompt + + def test_image_pad_expansion_per_view(self) -> None: + prompt = _render_chat_prompt(["base", "wrist_left"], [3, 5], "task") + assert prompt.count("<|image_pad|>") == 3 + 5 + # Each view wraps its pads in vision-start/end. + assert prompt.count("<|vision_start|>") == 2 + assert prompt.count("<|vision_end|>") == 2 + + def test_instruction_embedded(self) -> None: + prompt = _render_chat_prompt(["base"], [1], "open the drawer") + assert "Generate robot actions for the task:\nopen the drawer /no_cot" in prompt + + def test_length_mismatch_raises(self) -> None: + with pytest.raises(ValueError, match="must have the same length"): + _render_chat_prompt(["base", "wrist_left"], [1], "task") + + +class TestPrepareState: + @pytest.fixture() + def small_preprocessor(self): + # Small max_state_dim keeps the reference tensors tiny. + return XR0Preprocessor(max_state_dim=4, patch_size=16, merge_size=2) + + @pytest.mark.parametrize( + ("state", "expected"), + [ + # 1-D (D,) shorter than max_state_dim -> promoted to (1, D) then zero-padded. + (np.array([1.0, 2.0], dtype=np.float32), np.array([[[1.0, 2.0, 0.0, 0.0]]], dtype=np.float32)), + # 1-D exact length -> no padding. + ( + np.array([1.0, 2.0, 3.0, 4.0], dtype=np.float32), + np.array([[[1.0, 2.0, 3.0, 4.0]]], dtype=np.float32), + ), + # 1-D longer than max_state_dim -> truncated to first 4. + ( + np.array([1.0, 2.0, 3.0, 4.0, 5.0], dtype=np.float32), + np.array([[[1.0, 2.0, 3.0, 4.0]]], dtype=np.float32), + ), + # 2-D (B, D) batch preserved, padded per row. + ( + np.array([[1.0, 2.0], [3.0, 4.0]], dtype=np.float32), + np.array([[[1.0, 2.0, 0.0, 0.0]], [[3.0, 4.0, 0.0, 0.0]]], dtype=np.float32), + ), + # 3-D (B, T, D) -> last timestep only, then padded. + ( + np.array([[[9.0, 9.0], [1.0, 2.0]]], dtype=np.float32), + np.array([[[1.0, 2.0, 0.0, 0.0]]], dtype=np.float32), + ), + # 3-D multi-batch (B=2, T=2, D=2) -> per-sample last timestep, then padded. + ( + np.array([[[9.0, 9.0], [1.0, 2.0]], [[8.0, 8.0], [3.0, 4.0]]], dtype=np.float32), + np.array([[[1.0, 2.0, 0.0, 0.0]], [[3.0, 4.0, 0.0, 0.0]]], dtype=np.float32), + ), + ], + ) + def test_prepare_state(self, small_preprocessor, state, expected) -> None: + result = small_preprocessor._prepare_state({STATE: state}) + assert result.shape == expected.shape + assert result.dtype == np.float32 + np.testing.assert_allclose(result, expected) + + @pytest.mark.parametrize( + ("state", "expected"), + [ + # (value - mean) / (std + eps); real dims use given stats, padded dims stay 0. + # mean=[1,1], std=[2,2]: (3-1)/2=1.0, (5-1)/2=2.0. + ( + np.array([3.0, 5.0], dtype=np.float32), + np.array([[[1.0, 2.0, 0.0, 0.0]]], dtype=np.float32), + ), + # 3-D (B, T, D) last timestep normalized the same way. + ( + np.array([[[9.0, 9.0], [3.0, 5.0]]], dtype=np.float32), + np.array([[[1.0, 2.0, 0.0, 0.0]]], dtype=np.float32), + ), + ], + ) + def test_prepare_state_normalized(self, state, expected) -> None: + prep = XR0Preprocessor( + max_state_dim=4, + patch_size=16, + merge_size=2, + normalize_state=True, + state_mean=[1.0, 1.0], + state_std=[2.0, 2.0], + ) + result = prep._prepare_state({STATE: state}) + assert result.shape == expected.shape + assert result.dtype == np.float32 + np.testing.assert_allclose(result, expected, atol=1e-5) + + def test_missing_state_raises(self, small_preprocessor) -> None: + with pytest.raises(ValueError, match="requires a 'state'"): + small_preprocessor._prepare_state({}) + + +class TestXR0PreprocessorCall: + def test_output_structure_and_contiguity(self, preprocessor) -> None: + result = preprocessor(_make_inputs()) + assert set(result) == {"pixel_values", "state", TASK} + assert isinstance(result["pixel_values"], np.ndarray) + assert result["pixel_values"].dtype == np.float32 + assert result["pixel_values"].flags["C_CONTIGUOUS"] + assert isinstance(result["state"], np.ndarray) + assert result["state"].dtype == np.float32 + assert result["state"].flags["C_CONTIGUOUS"] + assert isinstance(result[TASK], list) + assert len(result[TASK]) == 1 + assert isinstance(result[TASK][0], str) + + def test_pixel_grid_matches_view_count(self, preprocessor) -> None: + # Flat patchified pixel_values: 2 views at 256x256 -> 256 tokens/image, + # 2 images -> 512 rows; feature = C * temporal_patch_size * patch**2. + result = preprocessor(_make_inputs()) + assert result["pixel_values"].shape == (512, 3 * 2 * 16 * 16) + + def test_pad_counts_derive_from_resized_dims(self, preprocessor) -> None: + # Non-square 128x256 stays patch-aligned (multiples of factor=32). + # patch_size=16 -> grid 8x16; merge 2x2 -> (8*16)/(2*2) = 32 pads/image. + prompt = preprocessor(_make_inputs(h=128, w=256))[TASK][0] + assert prompt.count("<|image_pad|>") == 32 * 2 + assert prompt.count("<|vision_start|>") == 2 + + def test_fewer_images_than_views_truncates(self, preprocessor) -> None: + # Only one flattened image -> single view drives pixel grid and prompt. + inputs = { + f"{IMAGES}.base": np.zeros((256, 256, 3), dtype=np.uint8), + STATE: np.zeros((8,), dtype=np.float32), + TASK: "pick up the cup", + } + result = preprocessor(inputs) + # Single 256x256 view -> 256 patch rows; feature = C * tp * patch**2. + assert result["pixel_values"].shape == (256, 3 * 2 * 16 * 16) + prompt = result[TASK][0] + assert "# Base View" in prompt + assert "# Left-Wrist View" not in prompt + assert prompt.count("<|vision_start|>") == 1