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