pi0.5 / inference.py
XiangpengYang's picture
feat: add validated UR action inference
5b21b68
Raw
History Blame Contribute Delete
4.13 kB
"""Validated π₀.₅ UR inference and serializable result formatting."""
from __future__ import annotations
import json
import random
import tempfile
from dataclasses import dataclass
import numpy as np
import pandas as pd
from PIL import Image
ACTION_LABELS = ("dx", "dy", "dz", "droll", "dpitch", "dyaw", "gripper")
ACTION_HORIZON = 10
MODEL_ACTION_DIM = 32
BASE_SEED = 42
@dataclass(frozen=True)
class PredictionResult:
actions: pd.DataFrame
json_path: str
status: str
def _as_rgb_uint8(value) -> np.ndarray:
if isinstance(value, Image.Image):
return np.asarray(value.convert("RGB"), dtype=np.uint8)
array = np.asarray(value)
if np.issubdtype(array.dtype, np.floating):
maximum = float(np.nanmax(array)) if array.size else 0.0
if maximum <= 1.0:
array = array * 255.0
array = np.clip(array, 0, 255).astype(np.uint8)
return np.asarray(Image.fromarray(array).convert("RGB"), dtype=np.uint8)
def _validate_state(values) -> np.ndarray:
try:
state = np.asarray(values, dtype=np.float32)
except (TypeError, ValueError) as exc:
raise ValueError("state must contain seven numeric values") from exc
if state.shape != (7,):
raise ValueError("state must contain exactly seven values")
if not np.isfinite(state).all():
raise ValueError("state must contain seven finite values")
return state
def _validate_trial_index(value) -> int:
if isinstance(value, bool) or not isinstance(value, (int, np.integer)):
raise ValueError("trial index must be a non-negative integer")
result = int(value)
if result < 0:
raise ValueError("trial index must be a non-negative integer")
return result
def set_seed(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
try:
import torch
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
except ImportError:
pass
def run_prediction(
policy,
fixed_image,
wrist_image,
instruction,
state_values,
trial_index,
model_id,
checkpoint_path,
) -> PredictionResult:
if fixed_image is None:
raise ValueError("fixed-camera image is required")
if wrist_image is None:
raise ValueError("wrist-camera image is required")
if not isinstance(instruction, str) or not instruction.strip():
raise ValueError("task instruction is required")
prompt = instruction.strip()
state = _validate_state(state_values)
trial = _validate_trial_index(trial_index)
seed = BASE_SEED + trial
set_seed(seed)
noise = np.random.default_rng(seed).standard_normal(
(ACTION_HORIZON, MODEL_ACTION_DIM), dtype=np.float32
)
observation = {
"observation/image": _as_rgb_uint8(fixed_image),
"observation/wrist_image": _as_rgb_uint8(wrist_image),
"observation/state": state,
"prompt": prompt,
}
result = policy.infer(observation, noise=noise)
actions = np.asarray(result.get("actions"))
expected_shape = (ACTION_HORIZON, len(ACTION_LABELS))
if actions.shape != expected_shape:
raise RuntimeError(f"policy returned action shape {actions.shape}; expected {expected_shape}")
if not np.isfinite(actions).all():
raise RuntimeError("policy returned non-finite actions")
table = pd.DataFrame(actions, columns=ACTION_LABELS)
payload = {
"model_id": model_id,
"checkpoint_path": checkpoint_path,
"instruction": prompt,
"state": state.tolist(),
"trial_index": trial,
"seed": seed,
"action_labels": list(ACTION_LABELS),
"actions": actions.tolist(),
}
with tempfile.NamedTemporaryFile(
mode="w", suffix=".json", delete=False, encoding="utf-8"
) as stream:
json.dump(payload, stream, indent=2, ensure_ascii=False)
json_path = stream.name
return PredictionResult(
actions=table,
json_path=json_path,
status=f"Prediction complete (seed={seed}, action_shape={expected_shape}).",
)