Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions src/physicalai/inference/component_factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -110,13 +110,15 @@ 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")
component_registry.register("to_float_tensor", "physicalai.inference.preprocessors.ToFloatTensorPreprocessor")

# 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")


Expand Down
2 changes: 2 additions & 0 deletions src/physicalai/inference/constants.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
ACTION = "action"
TASK = "task"
STATE = "state"
STATE_PASSTHROUGH = "state_passthrough"

TOKENIZED_PROMPT = "tokenized_prompt"
TOKENIZED_PROMPT_MASK = "tokenized_prompt_mask"
Expand All @@ -31,6 +32,7 @@
"RTC_INFERENCE_DELAY",
"RTC_MAX_GUIDANCE_WEIGHT",
"STATE",
"STATE_PASSTHROUGH",
"TASK",
"TOKENIZED_PROMPT",
"TOKENIZED_PROMPT_MASK",
Expand Down
2 changes: 2 additions & 0 deletions src/physicalai/inference/postprocessors/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
]
133 changes: 133 additions & 0 deletions src/physicalai/inference/postprocessors/xr0.py
Original file line number Diff line number Diff line change
@@ -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
2 changes: 2 additions & 0 deletions src/physicalai/inference/preprocessors/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand All @@ -36,4 +37,5 @@
"Rldx1TokenComposer",
"StatsNormalizer",
"ToFloatTensorPreprocessor",
"XR0Preprocessor",
]
Loading
Loading