zhicao's picture
Scheduled Commit
84c418f verified
Raw
History Blame Contribute Delete
15.5 kB
"""Model-side transform helpers for fixed-size track/force columns."""
from __future__ import annotations
from typing import Any
import numpy as np
from pydantic import Field
from groot.vla.model.dreamzero.transform.dreamzero_cotrain import (
DefaultDataCollator,
DreamTransform,
)
from .force import (
ACTION_HORIZON,
FORCE_HISTORY_FRAMES,
FORCE_OFFSETS,
FORCE_SENSOR_COUNT,
FORCE_SENSOR_DIM,
)
from .track import NUM_TRACK_POINTS, TRACK_HORIZON
TRACK_XY_ALIASES = (
"track_xy",
"tracks",
"track.xy",
"observation.track_xy",
"observation.tracks_xy",
)
TRACK_VISIBILITY_ALIASES = (
"track_visibility",
"visibility",
"vis",
"track.visibility",
"observation.track_visibility",
"observation.tracks_visibility",
)
FORCE_CURRENT_ALIASES = (
"current_force",
"force_current",
"observation.force_current",
)
FORCE_HISTORY_ALIASES = (
"force_history",
"tactile_force_history",
"force_history_raw",
"observation.tactile_force",
"force_history_vq",
"force_vq_codes",
"observation.force_history_vq",
)
FORCE_HISTORY_PADDING_MASK_ALIASES = (
"force_history_padding_mask",
"tactile_force_history_padding_mask",
)
def _first_present(data: dict[str, Any], names: tuple[str, ...]) -> Any | None:
for name in names:
if name in data:
return data[name]
return None
def _reshape_track_xy(values: Any) -> np.ndarray:
xy = np.asarray(values, dtype=np.float32)
if xy.ndim == 2 and xy.shape[-1] == NUM_TRACK_POINTS * 2:
xy = xy.reshape(xy.shape[0], NUM_TRACK_POINTS, 2)
if xy.ndim != 3 or xy.shape[1:] != (NUM_TRACK_POINTS, 2):
raise ValueError(
"track xy must be [T,250,2] or flattened [T,500], got "
f"{xy.shape}"
)
if not np.isfinite(xy).all():
raise ValueError("track xy contains NaN or infinity")
if xy.min(initial=0.0) < -1e-4 or xy.max(initial=1.0) > 1.0001:
raise ValueError("track xy must be normalized to [0,1]")
return xy
def _reshape_track_visibility(values: Any, time: int) -> np.ndarray:
visibility = np.asarray(values)
if visibility.ndim == 1 and visibility.size == time * NUM_TRACK_POINTS:
visibility = visibility.reshape(time, NUM_TRACK_POINTS)
if visibility.shape != (time, NUM_TRACK_POINTS):
raise ValueError(
f"track visibility must be [T,250], got {visibility.shape}"
)
return visibility.astype(np.float32)
def prepare_track_windows(
track_xy: Any,
track_visibility: Any,
*,
num_blocks: int,
history_frames: int = FORCE_HISTORY_FRAMES,
future_horizon: int = TRACK_HORIZON,
) -> dict[str, np.ndarray]:
"""Split a contiguous track timeline into clean-past and future windows.
The first future step immediately follows the final clean-past step.
Boundary indices are clamped, which is deterministic for online prefixes.
"""
if num_blocks <= 0 or history_frames <= 0:
raise ValueError("num_blocks and history_frames must be positive")
xy = _reshape_track_xy(track_xy)
visibility = _reshape_track_visibility(track_visibility, xy.shape[0])
total_time = xy.shape[0]
minimum = history_frames + num_blocks * future_horizon
if total_time < future_horizon:
raise ValueError(
f"track timeline has {total_time} frames; need at least {future_horizon}"
)
# A full training window is [H past][K*16 future]. For a shorter online
# prefix, anchor at its first frame and use deterministic edge padding.
first_anchor = history_frames - 1 if total_time >= minimum else 0
past_xy, past_vis, future_xy, future_vis = [], [], [], []
for block_index in range(num_blocks):
anchor = first_anchor + block_index * future_horizon
past_indices = np.arange(
anchor - history_frames + 1, anchor + 1, dtype=np.int64
)
future_indices = np.arange(
anchor + 1, anchor + 1 + future_horizon, dtype=np.int64
)
past_indices = np.clip(past_indices, 0, total_time - 1)
future_indices = np.clip(future_indices, 0, total_time - 1)
past_xy.append(xy[past_indices])
past_vis.append(visibility[past_indices])
future_xy.append(xy[future_indices])
future_vis.append(visibility[future_indices])
return {
"track_past_xy": np.stack(past_xy),
"track_past_visibility": np.stack(past_vis),
"track_future_xy": np.stack(future_xy),
"track_future_visibility": np.stack(future_vis),
}
class TrexTrackForceTransform(DreamTransform):
"""DreamTransform extension that preserves track and force model inputs."""
track_history_frames: int = FORCE_HISTORY_FRAMES
track_horizon: int = TRACK_HORIZON
track_points: int = NUM_TRACK_POINTS
require_track_inputs: bool = True
require_force_inputs: bool = True
use_deform_tactile: bool = False
track_xy_aliases: tuple[str, ...] = Field(default=TRACK_XY_ALIASES)
track_visibility_aliases: tuple[str, ...] = Field(
default=TRACK_VISIBILITY_ALIASES
)
force_current_aliases: tuple[str, ...] = Field(default=FORCE_CURRENT_ALIASES)
force_history_aliases: tuple[str, ...] = Field(default=FORCE_HISTORY_ALIASES)
force_history_padding_mask_aliases: tuple[str, ...] = Field(
default=FORCE_HISTORY_PADDING_MASK_ALIASES
)
def _num_blocks(self, data: dict[str, Any]) -> int:
if "action" not in data:
return 1
action = np.asarray(data["action"])
if action.shape[0] % ACTION_HORIZON:
raise ValueError("action timeline must contain complete 16-step blocks")
return max(action.shape[0] // ACTION_HORIZON, 1)
def _prepare_tracks(self, data: dict[str, Any], blocks: int) -> dict[str, np.ndarray]:
explicit_keys = (
"track_past_xy",
"track_past_visibility",
"track_future_xy",
"track_future_visibility",
)
if all(key in data for key in explicit_keys):
result = {
key: np.asarray(data[key], dtype=np.float32) for key in explicit_keys
}
expected = {
"track_past_xy": (
blocks,
self.track_history_frames,
self.track_points,
2,
),
"track_past_visibility": (
blocks,
self.track_history_frames,
self.track_points,
),
"track_future_xy": (
blocks,
self.track_horizon,
self.track_points,
2,
),
"track_future_visibility": (
blocks,
self.track_horizon,
self.track_points,
),
}
for key, shape in expected.items():
if result[key].shape != shape:
raise ValueError(f"{key} is {result[key].shape}, expected {shape}")
return result
xy = _first_present(data, self.track_xy_aliases)
visibility = _first_present(data, self.track_visibility_aliases)
if xy is None or visibility is None:
if self.require_track_inputs:
raise KeyError(
"250-point track xy/visibility columns were not provided"
)
return {}
return prepare_track_windows(
xy,
visibility,
num_blocks=blocks,
history_frames=self.track_history_frames,
future_horizon=self.track_horizon,
)
def _prepare_force(self, data: dict[str, Any], blocks: int) -> dict[str, np.ndarray]:
current = _first_present(data, self.force_current_aliases)
history = _first_present(data, self.force_history_aliases)
history_padding_mask = _first_present(
data, self.force_history_padding_mask_aliases
)
if history is None:
if self.require_force_inputs:
raise KeyError(
"raw 16x10x6 tactile history or 10 precomputed codes are required"
)
return {}
history_array = np.asarray(history)
is_raw = np.issubdtype(history_array.dtype, np.floating)
force_updates = len(FORCE_OFFSETS)
if is_raw:
if history_array.shape == (
FORCE_HISTORY_FRAMES,
FORCE_SENSOR_COUNT * FORCE_SENSOR_DIM,
):
history_array = history_array.reshape(
FORCE_HISTORY_FRAMES, FORCE_SENSOR_COUNT, FORCE_SENSOR_DIM
)
if history_array.shape == (
blocks,
FORCE_HISTORY_FRAMES,
FORCE_SENSOR_COUNT * FORCE_SENSOR_DIM,
):
history_array = history_array.reshape(
blocks,
FORCE_HISTORY_FRAMES,
FORCE_SENSOR_COUNT,
FORCE_SENSOR_DIM,
)
if history_array.shape == (
FORCE_HISTORY_FRAMES,
FORCE_SENSOR_COUNT,
FORCE_SENSOR_DIM,
):
history_array = np.broadcast_to(
history_array,
(
blocks,
FORCE_HISTORY_FRAMES,
FORCE_SENSOR_COUNT,
FORCE_SENSOR_DIM,
),
).copy()
expected_history = (
blocks,
FORCE_HISTORY_FRAMES,
FORCE_SENSOR_COUNT,
FORCE_SENSOR_DIM,
)
expected_update_history = (
blocks,
force_updates,
FORCE_HISTORY_FRAMES,
FORCE_SENSOR_COUNT,
FORCE_SENSOR_DIM,
)
if history_array.shape not in (expected_history, expected_update_history):
raise ValueError(
"raw tactile history must be [blocks,16,10,6] or "
f"[blocks,{force_updates},16,10,6], "
f"got {history_array.shape}"
)
history_array = history_array.astype(np.float32)
expected_padding_shape = history_array.shape[:-2]
if history_padding_mask is None:
padding_array = np.zeros(expected_padding_shape, dtype=bool)
else:
padding_array = np.asarray(history_padding_mask, dtype=bool)
if padding_array.shape == (FORCE_HISTORY_FRAMES,):
padding_array = np.broadcast_to(
padding_array, expected_padding_shape
).copy()
if padding_array.shape != expected_padding_shape:
raise ValueError(
"force history padding mask must match the raw history "
f"time axes {expected_padding_shape}, got "
f"{padding_array.shape}"
)
else:
if history_array.shape == (FORCE_SENSOR_COUNT,):
history_array = np.broadcast_to(
history_array, (blocks, FORCE_SENSOR_COUNT)
).copy()
expected_codes = (blocks, FORCE_SENSOR_COUNT)
expected_update_codes = (blocks, force_updates, FORCE_SENSOR_COUNT)
if history_array.shape not in (expected_codes, expected_update_codes):
raise ValueError(
"precomputed tactile codes must be [blocks,10] or "
f"[blocks,{force_updates},10], "
f"got {history_array.shape}"
)
history_array = history_array.astype(np.int64)
if current is None and is_raw:
current = (
history_array[:, :, -1]
if history_array.ndim == 6
else history_array[:, -1]
)
if current is None:
if self.require_force_inputs:
raise KeyError("current 10x6 force is required with precomputed codes")
return {}
current_array = np.asarray(current, dtype=np.float32)
if current_array.shape == (FORCE_SENSOR_COUNT * FORCE_SENSOR_DIM,):
current_array = current_array.reshape(FORCE_SENSOR_COUNT, FORCE_SENSOR_DIM)
if current_array.shape == (FORCE_SENSOR_COUNT, FORCE_SENSOR_DIM):
current_array = np.broadcast_to(
current_array, (blocks, FORCE_SENSOR_COUNT, FORCE_SENSOR_DIM)
).copy()
expected_current = (
blocks,
FORCE_SENSOR_COUNT,
FORCE_SENSOR_DIM,
)
expected_update_current = (
blocks,
force_updates,
FORCE_SENSOR_COUNT,
FORCE_SENSOR_DIM,
)
if current_array.shape not in (expected_current, expected_update_current):
raise ValueError(
"current force must be [blocks,10,6] or "
f"[blocks,{force_updates},10,6], got {current_array.shape}"
)
result = {"current_force": current_array}
result[
"tactile_force_history" if is_raw else "force_history_vq"
] = history_array
if is_raw:
result["tactile_force_history_padding_mask"] = padding_array
return result
def apply_single(self, data: dict) -> dict:
blocks = self._num_blocks(data)
track_data = self._prepare_tracks(data, blocks)
force_data = self._prepare_force(data, blocks)
transformed = super().apply_single(data)
images = np.asarray(transformed["images"])
expected_training_frames = 1 + blocks * 8
if self.training and images.shape[0] == expected_training_frames:
# Preserve DreamZero's 4-block/33-frame autoregressive geometry:
# one clean conditioning frame plus eight 10 Hz targets per block.
transformed["history_images"] = images[:1]
transformed["future_images"] = images[1:]
elif images.shape[0] >= 1 and not self.training:
transformed["history_images"] = images
else:
raise ValueError(
"T-Rex Track-Force expects one conditioning frame plus "
f"eight frames per autoregressive block ({expected_training_frames} "
"training frames), or one or more history frames for inference; "
f"got {images.shape[0]}"
)
transformed.update(track_data)
transformed.update(force_data)
if self.use_deform_tactile:
if "deform_current" not in data:
raise KeyError(
"use_deform_tactile=True but the dataset provided no "
"deform_current frames (enable load_deform_tactile)"
)
transformed["deform_current"] = np.asarray(
data["deform_current"], dtype=np.uint8
)
return transformed
TrackForceTransform = TrexTrackForceTransform
TrackForceDataCollator = DefaultDataCollator
__all__ = [
"TrackForceDataCollator",
"TrackForceTransform",
"TrexTrackForceTransform",
"prepare_track_windows",
]