dreamzero / scripts /data /build_trex_track_force_v2.py
zhicao's picture
Upload dreamzero source code (no model checkpoints)
fbd9366 verified
Raw
History Blame Contribute Delete
76.7 kB
"""Build the T-Rex track/EEF extension in a LeRobot v2 dataset.
This builder preserves every existing parquet column and adds:
* model-facing ``observation.track_xy`` / ``observation.track_visibility``;
* view-preserving ``observation.tracks.{head_left,left_wrist,right_wrist}``,
each frame stored as fixed-size ``[x, y, visibility]`` values;
* ``observation.state_eef62`` and ``action.eef62_absolute`` using T-Rex FK and
the canonical ``translation + rotation-6D + hand`` representation.
Writes are resumable and atomic. Existing valid episode outputs are skipped,
and the original parquet/metadata files receive one-time ``.trex_track_force.bak``
backups before their first replacement.
"""
from __future__ import annotations
import argparse
import hashlib
import importlib.util
import json
import os
import shutil
import sys
import tempfile
from datetime import datetime, timezone
from functools import lru_cache
from pathlib import Path
from typing import Callable, Iterable, Sequence
import numpy as np
_DATA_SCRIPT_DIR = Path(__file__).resolve().parent
_SCRIPT_DIR = _DATA_SCRIPT_DIR.parent
_DREAMZERO_ROOT = _SCRIPT_DIR.parent
if str(_SCRIPT_DIR) not in sys.path:
sys.path.insert(0, str(_SCRIPT_DIR))
from trex_track.layout import ( # noqa: E402
NUM_COMBINED_POINTS,
POINT_SLICES,
TRACK_LAYOUT_VERSION,
VIEW_ORDER,
VIEW_POINT_COUNTS,
VIEW_SLICES,
identity_metadata,
layout_metadata,
)
SCHEMA_VERSION = "trex_track_force_v2.3"
BACKUP_SUFFIX = ".trex_track_force.bak"
DEFAULT_DATASET_ROOT = _DREAMZERO_ROOT / "data" / "trex_small_force"
DEFAULT_TREX_ROOT = Path("/scratch1/home/zhicao/T-Rex")
PARQUET_SCHEMA_METADATA_KEY = b"trex_track_force_schema_version"
TARGET_RATE_HZ = 20.0
ACTION_CHUNK_STEPS = 16
ACTION_CHUNK_DURATION_SECONDS = ACTION_CHUNK_STEPS / TARGET_RATE_HZ
ACTION_CHUNK_TIMESTAMP_SPAN_SECONDS = (ACTION_CHUNK_STEPS - 1) / TARGET_RATE_HZ
AUTOREGRESSIVE_BLOCKS = 4
VIDEO_FRAMES_PER_BLOCK = 8
TRAINING_VIDEO_FRAMES = 1 + AUTOREGRESSIVE_BLOCKS * VIDEO_FRAMES_PER_BLOCK
FORCE_COLUMN = "observation.tactile_force"
FORCE_FLAT_DIM = 60
FORCE_SENSOR_COUNT = 10
FORCE_SENSOR_DIM = 6
FORCE_HISTORY_FRAMES = 16
RELATIVE_ACTION_STATS_FILENAME = "relative_stats_dreamzero.json"
STATE_EEF_COLUMN = "observation.state_eef62"
ACTION_EEF_COLUMN = "action.eef62_absolute"
TRACK_XY_COLUMN = "observation.track_xy"
TRACK_VISIBILITY_COLUMN = "observation.track_visibility"
TRACK_COLUMNS = {
"head_left": "observation.tracks.head_left",
"left_wrist": "observation.tracks.left_wrist",
"right_wrist": "observation.tracks.right_wrist",
}
NEW_COLUMNS = (
TRACK_XY_COLUMN,
TRACK_VISIBILITY_COLUMN,
*TRACK_COLUMNS.values(),
STATE_EEF_COLUMN,
ACTION_EEF_COLUMN,
)
LEFT_EEF = slice(0, 9)
LEFT_HAND_EEF = slice(9, 31)
RIGHT_EEF = slice(31, 40)
RIGHT_HAND_EEF = slice(40, 62)
EefConverter = Callable[[np.ndarray], np.ndarray]
class DatasetSchemaError(RuntimeError):
"""Raised when an episode cannot satisfy the track/EEF schema."""
def sample_timestamps_nearest(
source_timestamps: np.ndarray | Sequence[float],
target_rate_hz: float = TARGET_RATE_HZ,
*,
anchor_index: int | None = None,
anchor_timestamp: float | None = None,
offsets: Sequence[int] | np.ndarray | None = None,
alignment_tolerance: float = 1e-6,
) -> dict[str, object]:
"""Deterministically align a target-rate grid to nearest source frames.
Ties choose the earlier source frame. Non-padding source indices must be
unique, and every non-padding alignment error is bounded by half the
median source period plus ``alignment_tolerance``. Queries outside the
source interval clamp to an endpoint and are explicitly marked in
``padding_mask``.
"""
source = np.asarray(source_timestamps, dtype=np.float64)
if source.ndim != 1 or source.size < 2:
raise DatasetSchemaError(
f"source_timestamps must be a 1D array with >=2 values, got {source.shape}"
)
if not np.isfinite(source).all():
raise DatasetSchemaError("source_timestamps contain NaN/Inf")
source_deltas = np.diff(source)
if not np.all(source_deltas > 0.0):
raise DatasetSchemaError("source_timestamps must be strictly increasing")
source_period = float(np.median(source_deltas))
if not np.isfinite(source_period) or source_period <= 0.0:
raise DatasetSchemaError("could not infer a positive source period")
max_source_period = float(source_deltas.max())
if max_source_period > 1.5 * source_period:
raise DatasetSchemaError(
"source timestamps contain a dropped-frame gap: "
f"max={max_source_period:.9f}s median={source_period:.9f}s"
)
target_rate = float(target_rate_hz)
if not np.isfinite(target_rate) or target_rate <= 0.0:
raise DatasetSchemaError("target_rate_hz must be finite and positive")
tolerance = float(alignment_tolerance)
if not np.isfinite(tolerance) or tolerance < 0.0:
raise DatasetSchemaError("alignment_tolerance must be finite and non-negative")
if anchor_index is not None and anchor_timestamp is not None:
raise DatasetSchemaError("set only anchor_index or anchor_timestamp")
if anchor_timestamp is None:
index = 0 if anchor_index is None else int(anchor_index)
if index < 0 or index >= source.size:
raise DatasetSchemaError(
f"anchor_index {index} is outside [0,{source.size})"
)
anchor = float(source[index])
else:
anchor = float(anchor_timestamp)
if not np.isfinite(anchor):
raise DatasetSchemaError("anchor_timestamp must be finite")
if offsets is None:
last_offset = int(
np.floor((float(source[-1]) - anchor) * target_rate + tolerance * target_rate)
)
if last_offset < 0:
raise DatasetSchemaError("anchor is after the source timestamp interval")
offset_array = np.arange(last_offset + 1, dtype=np.int64)
else:
raw_offsets = np.asarray(offsets)
if raw_offsets.ndim != 1 or raw_offsets.size == 0:
raise DatasetSchemaError("offsets must be a non-empty 1D sequence")
offset_array = raw_offsets.astype(np.int64)
if not np.array_equal(raw_offsets, offset_array):
raise DatasetSchemaError("offsets must contain integer target steps")
if not np.all(np.diff(offset_array) > 0):
raise DatasetSchemaError("offsets must be strictly increasing and unique")
target = anchor + offset_array.astype(np.float64) / target_rate
if not np.all(np.diff(target) > 0.0):
raise DatasetSchemaError("target timestamps must be strictly increasing")
padding = (target < source[0] - tolerance) | (target > source[-1] + tolerance)
insertion = np.searchsorted(source, target, side="left")
lower = np.clip(insertion - 1, 0, source.size - 1)
upper = np.clip(insertion, 0, source.size - 1)
lower_error = np.abs(target - source[lower])
upper_error = np.abs(source[upper] - target)
# Differences within tolerance count as midpoint ties and choose earlier.
choose_upper = upper_error < (lower_error - tolerance)
indices = np.where(choose_upper, upper, lower).astype(np.int64)
indices[target < source[0]] = 0
indices[target > source[-1]] = source.size - 1
alignment_errors = np.abs(source[indices] - target)
non_padding = ~padding
# Real MP4 timestamps have small per-frame jitter. Nearest-neighbour error
# is bounded by half the local gap, not half the median source period.
max_allowed_error = max_source_period / 2.0 + tolerance
if non_padding.any() and np.any(
alignment_errors[non_padding] > max_allowed_error
):
worst = float(alignment_errors[non_padding].max())
raise DatasetSchemaError(
f"timestamp alignment error {worst:.9f}s exceeds "
f"source_period/2+tolerance={max_allowed_error:.9f}s"
)
selected = indices[non_padding]
if np.unique(selected).size != selected.size:
raise DatasetSchemaError(
"nearest timestamp alignment selected duplicate non-padding source frames"
)
return {
"indices": indices,
"target_timestamps": target,
"offsets": offset_array,
"padding_mask": padding.astype(bool),
"alignment_errors": alignment_errors,
"source_period_seconds": source_period,
"max_source_period_seconds": max_source_period,
"source_rate_hz": 1.0 / source_period,
"target_rate_hz": target_rate,
"max_allowed_alignment_error_seconds": max_allowed_error,
}
def summarize_timestamp_sampling(
source_timestamps: np.ndarray | Sequence[float],
*,
target_rate_hz: float = TARGET_RATE_HZ,
action_chunk_steps: int = ACTION_CHUNK_STEPS,
) -> dict[str, object]:
"""Return a JSON-safe per-episode 20 Hz coverage/chunk validation summary."""
source = np.asarray(source_timestamps, dtype=np.float64)
coverage = sample_timestamps_nearest(
source,
target_rate_hz=target_rate_hz,
anchor_index=0,
)
chunk_steps = int(action_chunk_steps)
if chunk_steps <= 0:
raise DatasetSchemaError("action_chunk_steps must be positive")
chunk = sample_timestamps_nearest(
source,
target_rate_hz=target_rate_hz,
anchor_index=0,
offsets=np.arange(chunk_steps, dtype=np.int64),
)
coverage_indices = np.asarray(coverage["indices"], dtype=np.int64)
coverage_targets = np.asarray(coverage["target_timestamps"], dtype=np.float64)
coverage_errors = np.asarray(coverage["alignment_errors"], dtype=np.float64)
coverage_padding = np.asarray(coverage["padding_mask"], dtype=bool)
chunk_padding = np.asarray(chunk["padding_mask"], dtype=bool)
target_rate = float(target_rate_hz)
chunk_duration = chunk_steps / target_rate
chunk_timestamp_span = (chunk_steps - 1) / target_rate
if chunk_steps == ACTION_CHUNK_STEPS and np.isclose(
target_rate, TARGET_RATE_HZ
) and (
not np.isclose(chunk_duration, ACTION_CHUNK_DURATION_SECONDS, atol=1e-12)
or not np.isclose(
chunk_timestamp_span,
ACTION_CHUNK_TIMESTAMP_SPAN_SECONDS,
atol=1e-12,
)
):
raise DatasetSchemaError("invalid 16-step/20 Hz action chunk definition")
return {
"source_frame_count": int(source.size),
"source_start_timestamp": float(source[0]),
"source_end_timestamp": float(source[-1]),
"source_timestamp_span_seconds": float(source[-1] - source[0]),
"source_period_seconds": float(coverage["source_period_seconds"]),
"source_rate_hz": float(coverage["source_rate_hz"]),
"target_rate_hz": target_rate,
"target_sample_count": int(coverage_indices.size),
"target_start_timestamp": float(coverage_targets[0]),
"target_end_timestamp": float(coverage_targets[-1]),
"first_source_index": int(coverage_indices[0]),
"last_source_index": int(coverage_indices[-1]),
"max_alignment_error_seconds": float(coverage_errors.max(initial=0.0)),
"max_allowed_alignment_error_seconds": float(
coverage["max_allowed_alignment_error_seconds"]
),
"coverage_padding_count": int(coverage_padding.sum()),
"action_chunk_steps": chunk_steps,
# A 16-step control horizon is the half-open interval [t, t+0.8s).
"action_chunk_duration_seconds": chunk_duration,
# The first/last sampled timestamps in that horizon are 15/20=0.75s apart.
"action_chunk_timestamp_span_seconds": chunk_timestamp_span,
"action_chunk_fully_covered": not bool(chunk_padding.any()),
"action_chunk_padding_count": int(chunk_padding.sum()),
"action_chunk_padding_mask": chunk_padding.tolist(),
"complete_action_chunks": int(coverage_indices.size // chunk_steps),
}
def _utc_now() -> str:
return datetime.now(timezone.utc).isoformat()
def _load_json(path: Path) -> dict:
if not path.is_file():
raise FileNotFoundError(path)
with path.open("r", encoding="utf-8") as file:
value = json.load(file)
if not isinstance(value, dict):
raise DatasetSchemaError(f"expected JSON object in {path}")
return value
def _fsync_directory(path: Path) -> None:
try:
fd = os.open(path, os.O_RDONLY)
except OSError:
return
try:
os.fsync(fd)
finally:
os.close(fd)
def _atomic_write_json(path: Path, value: dict) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
fd, tmp_name = tempfile.mkstemp(
prefix=f".{path.name}.",
suffix=".tmp",
dir=path.parent,
)
try:
with os.fdopen(fd, "w", encoding="utf-8") as file:
json.dump(value, file, indent=2, sort_keys=False)
file.write("\n")
file.flush()
os.fsync(file.fileno())
os.replace(tmp_name, path)
_fsync_directory(path.parent)
except BaseException:
try:
os.unlink(tmp_name)
except FileNotFoundError:
pass
raise
def backup_path(path: Path) -> Path:
return path.with_name(path.name + BACKUP_SUFFIX)
def _atomic_backup(path: Path) -> Path | None:
"""Create a one-time atomic backup, never replacing an existing backup."""
if not path.exists():
return None
destination = backup_path(path)
if destination.exists():
return destination
fd, tmp_name = tempfile.mkstemp(
prefix=f".{destination.name}.",
suffix=".tmp",
dir=path.parent,
)
os.close(fd)
try:
shutil.copy2(path, tmp_name)
with open(tmp_name, "rb") as file:
os.fsync(file.fileno())
# A concurrent builder may have completed the backup while we copied.
if destination.exists():
os.unlink(tmp_name)
return destination
os.replace(tmp_name, destination)
_fsync_directory(path.parent)
return destination
except BaseException:
try:
os.unlink(tmp_name)
except FileNotFoundError:
pass
raise
def _sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as file:
for block in iter(lambda: file.read(1024 * 1024), b""):
digest.update(block)
return digest.hexdigest()
def _import_pyarrow():
try:
import pyarrow as pa
import pyarrow.parquet as pq
except ImportError as exc:
raise RuntimeError("pyarrow is required to build LeRobot parquet files") from exc
return pa, pq
def _fixed_size_array(values: np.ndarray):
"""Convert ``[rows, *shape]`` to nested Arrow FixedSizeListArray."""
pa, _ = _import_pyarrow()
array = np.asarray(values, dtype=np.float32)
if array.ndim < 2:
raise ValueError(f"fixed-size feature must have at least 2 dims, got {array.shape}")
result = pa.array(array.reshape(-1), type=pa.float32())
for size in reversed(array.shape[1:]):
result = pa.FixedSizeListArray.from_arrays(result, int(size))
if len(result) != array.shape[0]:
raise AssertionError(f"Arrow rows {len(result)} != numpy rows {array.shape[0]}")
return result
def _set_or_append_column(table, name: str, values: np.ndarray):
column = _fixed_size_array(values)
index = table.schema.get_field_index(name)
if index >= 0:
return table.set_column(index, name, column)
return table.append_column(name, column)
def _column_to_numpy(table, name: str, dtype=np.float32) -> np.ndarray:
if name not in table.column_names:
raise DatasetSchemaError(f"missing parquet column {name!r}")
values = table[name].combine_chunks().to_pylist()
try:
return np.asarray(values, dtype=dtype)
except (TypeError, ValueError) as exc:
raise DatasetSchemaError(f"column {name!r} is not a dense numeric array") from exc
def validate_tactile_force(table) -> dict[str, object]:
"""Validate the existing force-only source without creating VQ-code columns."""
force = _column_to_numpy(table, FORCE_COLUMN, dtype=np.float32)
expected = (int(table.num_rows), FORCE_FLAT_DIM)
if force.shape != expected:
raise DatasetSchemaError(
f"{FORCE_COLUMN} must have shape {expected}, got {force.shape}"
)
if not np.isfinite(force).all():
raise DatasetSchemaError(f"{FORCE_COLUMN} contains NaN/Inf")
return {
"source_column": FORCE_COLUMN,
"stored_shape": [FORCE_FLAT_DIM],
"reshape": [FORCE_SENSOR_COUNT, FORCE_SENSOR_DIM],
"history_frames": FORCE_HISTORY_FRAMES,
"history_encoding": "online_model_encoder",
"vq_codes_on_disk": False,
"finite": True,
}
def _is_fixed_shape(field_type, shape: Sequence[int]) -> bool:
pa, _ = _import_pyarrow()
current = field_type
for size in shape:
if not pa.types.is_fixed_size_list(current) or current.list_size != int(size):
return False
current = current.value_type
return pa.types.is_float32(current)
def _atomic_write_parquet(table, path: Path, *, expected_rows: int) -> None:
_, pq = _import_pyarrow()
path.parent.mkdir(parents=True, exist_ok=True)
fd, tmp_name = tempfile.mkstemp(
prefix=f".{path.name}.",
suffix=".tmp.parquet",
dir=path.parent,
)
os.close(fd)
tmp_path = Path(tmp_name)
try:
pq.write_table(table, tmp_path, compression="zstd")
with tmp_path.open("rb") as file:
os.fsync(file.fileno())
validate_episode_parquet(
tmp_path,
expected_frames=expected_rows,
verify_source_fk=False,
)
_atomic_backup(path)
os.replace(tmp_path, path)
_fsync_directory(path.parent)
except BaseException:
tmp_path.unlink(missing_ok=True)
raise
def _scalar_text(value: np.ndarray) -> str:
scalar = np.asarray(value)
if scalar.shape != ():
raise DatasetSchemaError(f"expected scalar string, got shape {scalar.shape}")
return str(scalar.item())
def _validate_unit_interval(name: str, values: np.ndarray) -> None:
array = np.asarray(values)
if not np.isfinite(array).all():
raise DatasetSchemaError(f"{name} contains NaN/Inf")
if array.size and (float(array.min()) < 0.0 or float(array.max()) > 1.0):
raise DatasetSchemaError(
f"{name} must be in [0,1], got [{array.min()}, {array.max()}]"
)
def load_track_payload(
path: Path,
*,
expected_frames: int | None = None,
episode_index: int | None = None,
) -> dict[str, np.ndarray]:
"""Load and strictly validate a canonical extraction NPZ."""
if not path.is_file():
raise FileNotFoundError(path)
with np.load(path, allow_pickle=False) as archive:
forbidden = [name for name in archive.files if name.startswith("images_")]
if forbidden:
raise DatasetSchemaError(
f"{path} embeds full RGB arrays ({forbidden}); regenerate with the new extractor"
)
payload = {name: np.asarray(archive[name]).copy() for name in archive.files}
required = {
"tracks",
"vis",
"tracks_head_left",
"tracks_left_wrist",
"tracks_right_wrist",
"vis_head_left",
"vis_left_wrist",
"vis_right_wrist",
"episode_index",
"num_steps",
"point_slices",
"track_layout_version",
"point_view_ids",
"point_hand_ids",
"point_role_ids",
"point_local_ids",
"point_global_ids",
"point_names",
}
missing = sorted(required.difference(payload))
if missing:
raise DatasetSchemaError(f"{path} is missing keys: {missing}")
tracks = np.asarray(payload["tracks"], dtype=np.float32)
visibility = np.asarray(payload["vis"], dtype=np.float32)
if tracks.ndim != 3 or tracks.shape[1:] != (NUM_COMBINED_POINTS, 2):
raise DatasetSchemaError(
f"{path}: tracks must be (T,{NUM_COMBINED_POINTS},2), got {tracks.shape}"
)
if visibility.shape != tracks.shape[:2]:
raise DatasetSchemaError(
f"{path}: visibility {visibility.shape} != {tracks.shape[:2]}"
)
num_frames = int(tracks.shape[0])
if int(np.asarray(payload["num_steps"]).item()) != num_frames:
raise DatasetSchemaError(f"{path}: num_steps does not match tracks")
if expected_frames is not None and num_frames != int(expected_frames):
raise DatasetSchemaError(
f"{path}: {num_frames} track frames != {expected_frames} parquet frames"
)
stored_episode = int(np.asarray(payload["episode_index"]).item())
if episode_index is not None and stored_episode != int(episode_index):
raise DatasetSchemaError(
f"{path}: episode_index={stored_episode}, expected {episode_index}"
)
if _scalar_text(payload["track_layout_version"]) != TRACK_LAYOUT_VERSION:
raise DatasetSchemaError(f"{path}: unsupported track layout version")
if not np.array_equal(
np.asarray(payload["point_slices"], dtype=np.int32),
np.asarray(POINT_SLICES, dtype=np.int32),
):
raise DatasetSchemaError(f"{path}: point_slices do not match canonical layout")
expected_ids = identity_metadata()
identity_keys = {
"point_view_ids": "view_ids",
"point_hand_ids": "hand_ids",
"point_role_ids": "role_ids",
"point_local_ids": "local_ids",
"point_global_ids": "global_ids",
}
for stored_key, expected_key in identity_keys.items():
if not np.array_equal(
np.asarray(payload[stored_key], dtype=np.int64),
np.asarray(expected_ids[expected_key], dtype=np.int64),
):
raise DatasetSchemaError(f"{path}: unstable identity metadata in {stored_key}")
if not np.array_equal(
np.asarray(payload["point_names"]).astype(str),
np.asarray(expected_ids["point_names"]).astype(str),
):
raise DatasetSchemaError(f"{path}: unstable identity metadata in point_names")
_validate_unit_interval("tracks", tracks)
_validate_unit_interval("visibility", visibility)
if not np.all((visibility == 0.0) | (visibility == 1.0)):
raise DatasetSchemaError(f"{path}: visibility must be binary")
view_tracks: list[np.ndarray] = []
view_visibility: list[np.ndarray] = []
for view in VIEW_ORDER:
count = VIEW_POINT_COUNTS[view]
track_key = f"tracks_{view}"
vis_key = f"vis_{view}"
track = np.asarray(payload[track_key], dtype=np.float32)
vis = np.asarray(payload[vis_key], dtype=np.float32)
if track.shape != (num_frames, count, 2):
raise DatasetSchemaError(f"{path}: {track_key} has shape {track.shape}")
if vis.shape != (num_frames, count):
raise DatasetSchemaError(f"{path}: {vis_key} has shape {vis.shape}")
_validate_unit_interval(track_key, track)
_validate_unit_interval(vis_key, vis)
view_tracks.append(track)
view_visibility.append(vis)
if not np.array_equal(np.concatenate(view_tracks, axis=1), tracks):
raise DatasetSchemaError(f"{path}: combined tracks differ from per-view tracks")
if not np.array_equal(np.concatenate(view_visibility, axis=1), visibility):
raise DatasetSchemaError(f"{path}: combined visibility differs from per-view visibility")
return payload
def track_features_from_payload(payload: dict[str, np.ndarray]) -> dict[str, np.ndarray]:
features: dict[str, np.ndarray] = {
TRACK_XY_COLUMN: np.asarray(payload["tracks"], dtype=np.float32),
TRACK_VISIBILITY_COLUMN: np.asarray(payload["vis"], dtype=np.float32),
}
for view, column in TRACK_COLUMNS.items():
xy = np.asarray(payload[f"tracks_{view}"], dtype=np.float32)
vis = np.asarray(payload[f"vis_{view}"], dtype=np.float32)[..., None]
features[column] = np.concatenate([xy, vis], axis=-1).astype(np.float32)
return features
@lru_cache(maxsize=1)
def _load_lerobot_common():
path = DEFAULT_TREX_ROOT / "utils" / "lerobot_common.py"
if not path.is_file():
raise FileNotFoundError(f"T-Rex pose semantics module not found: {path}")
spec = importlib.util.spec_from_file_location("_trex_lerobot_common_schema", path)
if spec is None or spec.loader is None:
raise ImportError(f"cannot load {path}")
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
for name in ("pose_matrix_to_9d", "get_rot_mat"):
if not hasattr(module, name):
raise ImportError(f"{path} does not define {name}")
return module
def _validate_transform(matrix: np.ndarray, *, label: str) -> None:
transform = np.asarray(matrix, dtype=np.float64)
if transform.shape != (4, 4) or not np.isfinite(transform).all():
raise DatasetSchemaError(f"{label}: FK returned an invalid transform")
if not np.allclose(transform[3], [0.0, 0.0, 0.0, 1.0], atol=1e-8):
raise DatasetSchemaError(f"{label}: FK transform has an invalid homogeneous row")
rotation = transform[:3, :3]
if not np.allclose(rotation.T @ rotation, np.eye(3), atol=1e-5):
raise DatasetSchemaError(f"{label}: FK rotation is not orthonormal")
if not np.isclose(np.linalg.det(rotation), 1.0, atol=1e-5):
raise DatasetSchemaError(f"{label}: FK rotation determinant is not +1")
def joint58_to_eef62_batch(joints: np.ndarray) -> np.ndarray:
"""Convert joint-space state/action rows to absolute 62-D EEF semantics.
This function deliberately has no fallback. Missing robot assets, invalid
joints, or unreliable FK raise an exception rather than fabricating poses.
"""
source = np.asarray(joints, dtype=np.float64)
if source.ndim != 2 or source.shape[1] != 58:
raise DatasetSchemaError(f"FK expects (T,58), got {source.shape}")
if not np.isfinite(source).all():
raise DatasetSchemaError("FK input contains NaN/Inf")
try:
from trex_track.trex_fk import (
frame_pose_matrix,
get_bimanual_robot,
state_to_components,
)
robot, assemble_qpos, _ = get_bimanual_robot()
common = _load_lerobot_common()
output = np.empty((source.shape[0], 62), dtype=np.float32)
for index, row in enumerate(source):
components = state_to_components(row)
qpos = assemble_qpos(
{
"left_arm": components["left_arm"],
"right_arm": components["right_arm"],
}
)
left_pose = frame_pose_matrix(robot, qpos, "L_ee")
right_pose = frame_pose_matrix(robot, qpos, "R_ee")
_validate_transform(left_pose, label=f"row {index} left")
_validate_transform(right_pose, label=f"row {index} right")
left_9d = common.pose_matrix_to_9d(left_pose[None])[0]
right_9d = common.pose_matrix_to_9d(right_pose[None])[0]
output[index] = np.concatenate(
[
left_9d,
components["left_hand"],
right_9d,
components["right_hand"],
]
)
except DatasetSchemaError:
raise
except Exception as exc:
raise DatasetSchemaError(
"reliable T-Rex FK failed; refusing to synthesize EEF values"
) from exc
validate_eef62(output, source_joint58=source, label="FK output")
return output
def validate_eef62(
values: np.ndarray,
*,
source_joint58: np.ndarray | None = None,
label: str,
) -> None:
"""Validate shape, hand preservation, rotations, and pose/rot6d roundtrip."""
array = np.asarray(values, dtype=np.float64)
if array.ndim != 2 or array.shape[1] != 62:
raise DatasetSchemaError(f"{label}: expected (T,62), got {array.shape}")
if not np.isfinite(array).all():
raise DatasetSchemaError(f"{label}: contains NaN/Inf")
common = _load_lerobot_common()
for side, arm_slice in (("left", LEFT_EEF), ("right", RIGHT_EEF)):
arm = array[:, arm_slice]
for row_index, pose9 in enumerate(arm):
rotation = np.asarray(common.get_rot_mat(pose9[3:9]), dtype=np.float64)
if not np.allclose(rotation.T @ rotation, np.eye(3), atol=2e-5):
raise DatasetSchemaError(
f"{label}: {side} row {row_index} rot6d is not orthonormal"
)
if not np.isclose(np.linalg.det(rotation), 1.0, atol=2e-5):
raise DatasetSchemaError(
f"{label}: {side} row {row_index} rotation determinant is not +1"
)
transform = np.eye(4, dtype=np.float64)
transform[:3, :3] = rotation
transform[:3, 3] = pose9[:3]
roundtrip = common.pose_matrix_to_9d(transform[None])[0]
if not np.allclose(roundtrip, pose9, atol=2e-5, rtol=1e-5):
raise DatasetSchemaError(
f"{label}: {side} row {row_index} pose/rot6d roundtrip failed"
)
if source_joint58 is not None:
source = np.asarray(source_joint58, dtype=np.float64)
if source.shape != (array.shape[0], 58):
raise DatasetSchemaError(
f"{label}: source shape {source.shape} does not match EEF rows"
)
if not np.allclose(array[:, LEFT_HAND_EEF], source[:, 7:29], atol=1e-6):
raise DatasetSchemaError(f"{label}: left hand values were not preserved")
if not np.allclose(array[:, RIGHT_HAND_EEF], source[:, 36:58], atol=1e-6):
raise DatasetSchemaError(f"{label}: right hand values were not preserved")
def convert_eef_columns(
state58: np.ndarray,
action58: np.ndarray,
*,
converter: EefConverter | None = None,
) -> tuple[np.ndarray, np.ndarray]:
convert = converter or joint58_to_eef62_batch
state = np.asarray(state58, dtype=np.float64)
action = np.asarray(action58, dtype=np.float64)
if state.ndim != 2 or state.shape[1] != 58:
raise DatasetSchemaError(f"observation.state must be (T,58), got {state.shape}")
if action.shape != state.shape:
raise DatasetSchemaError(f"action shape {action.shape} != state shape {state.shape}")
state_eef = np.asarray(convert(state), dtype=np.float32)
action_eef = np.asarray(convert(action), dtype=np.float32)
validate_eef62(state_eef, source_joint58=state, label=STATE_EEF_COLUMN)
validate_eef62(action_eef, source_joint58=action, label=ACTION_EEF_COLUMN)
return state_eef, action_eef
def episode_parquet_path(dataset_root: Path, episode_index: int, info: dict | None = None) -> Path:
metadata = info or _load_json(dataset_root / "meta" / "info.json")
pattern = metadata.get(
"data_path",
"data/chunk-{episode_chunk:03d}/episode_{episode_index:06d}.parquet",
)
chunk_size = int(metadata.get("chunks_size", 1000))
return dataset_root / pattern.format(
episode_chunk=int(episode_index) // chunk_size,
episode_index=int(episode_index),
)
def default_track_cache(dataset_root: Path) -> Path:
return dataset_root.with_name(dataset_root.name + "_tracks")
def track_npz_path(track_cache: Path, episode_index: int) -> Path:
return track_cache / f"episode_{int(episode_index):06d}.npz"
def validate_episode_parquet(
path: Path,
*,
expected_frames: int | None = None,
verify_source_fk: bool = False,
converter: EefConverter | None = None,
) -> dict[str, object]:
"""Validate fixed-size Arrow types, values, frame count, and EEF semantics."""
_, pq = _import_pyarrow()
if not path.is_file():
raise FileNotFoundError(path)
table = pq.read_table(path)
schema_version = (table.schema.metadata or {}).get(PARQUET_SCHEMA_METADATA_KEY)
if schema_version != SCHEMA_VERSION.encode("utf-8"):
found = schema_version.decode("utf-8") if schema_version is not None else None
raise DatasetSchemaError(
f"{path}: parquet schema version {found!r} != {SCHEMA_VERSION!r}; rebuild required"
)
if expected_frames is not None and table.num_rows != int(expected_frames):
raise DatasetSchemaError(
f"{path}: {table.num_rows} rows != expected {expected_frames}"
)
for old_column in ("observation.state", "action"):
if old_column not in table.column_names:
raise DatasetSchemaError(f"{path}: original column {old_column!r} is missing")
timestamps = _column_to_numpy(table, "timestamp", dtype=np.float64)
if timestamps.shape == (table.num_rows, 1):
timestamps = timestamps[:, 0]
if timestamps.shape != (table.num_rows,):
raise DatasetSchemaError(
f"{path}: timestamp must have shape ({table.num_rows},), got {timestamps.shape}"
)
sampling_summary = summarize_timestamp_sampling(timestamps)
force_summary = validate_tactile_force(table)
combined_xy_field = (
table.schema.field(TRACK_XY_COLUMN)
if TRACK_XY_COLUMN in table.column_names
else None
)
if combined_xy_field is None or not _is_fixed_shape(
combined_xy_field.type, (NUM_COMBINED_POINTS, 2)
):
raise DatasetSchemaError(
f"{path}: {TRACK_XY_COLUMN} must be Arrow fixed-size float32 "
f"({NUM_COMBINED_POINTS}, 2)"
)
combined_visibility_field = (
table.schema.field(TRACK_VISIBILITY_COLUMN)
if TRACK_VISIBILITY_COLUMN in table.column_names
else None
)
if combined_visibility_field is None or not _is_fixed_shape(
combined_visibility_field.type, (NUM_COMBINED_POINTS,)
):
raise DatasetSchemaError(
f"{path}: {TRACK_VISIBILITY_COLUMN} must be Arrow fixed-size float32 "
f"({NUM_COMBINED_POINTS},)"
)
combined_xy = _column_to_numpy(table, TRACK_XY_COLUMN)
combined_visibility = _column_to_numpy(table, TRACK_VISIBILITY_COLUMN)
_validate_unit_interval(TRACK_XY_COLUMN, combined_xy)
_validate_unit_interval(TRACK_VISIBILITY_COLUMN, combined_visibility)
if not np.all(
(combined_visibility == 0.0) | (combined_visibility == 1.0)
):
raise DatasetSchemaError(
f"{path}: {TRACK_VISIBILITY_COLUMN} visibility is not binary"
)
view_xy: list[np.ndarray] = []
view_visibility: list[np.ndarray] = []
for view, column in TRACK_COLUMNS.items():
field = table.schema.field(column) if column in table.column_names else None
shape = (VIEW_POINT_COUNTS[view], 3)
if field is None or not _is_fixed_shape(field.type, shape):
raise DatasetSchemaError(
f"{path}: {column} must be Arrow fixed-size float32 {shape}"
)
values = _column_to_numpy(table, column)
if values.shape != (table.num_rows, *shape):
raise DatasetSchemaError(f"{path}: {column} has shape {values.shape}")
_validate_unit_interval(column, values)
visibility = values[..., 2]
if not np.all((visibility == 0.0) | (visibility == 1.0)):
raise DatasetSchemaError(f"{path}: {column} visibility is not binary")
view_xy.append(values[..., :2])
view_visibility.append(visibility)
if not np.array_equal(np.concatenate(view_xy, axis=1), combined_xy):
raise DatasetSchemaError(f"{path}: combined and per-view track XY differ")
if not np.array_equal(
np.concatenate(view_visibility, axis=1), combined_visibility
):
raise DatasetSchemaError(f"{path}: combined and per-view visibility differ")
for column in (STATE_EEF_COLUMN, ACTION_EEF_COLUMN):
field = table.schema.field(column) if column in table.column_names else None
if field is None or not _is_fixed_shape(field.type, (62,)):
raise DatasetSchemaError(
f"{path}: {column} must be Arrow fixed-size float32 (62,)"
)
source_state = _column_to_numpy(table, "observation.state")
source_action = _column_to_numpy(table, "action")
state_eef = _column_to_numpy(table, STATE_EEF_COLUMN)
action_eef = _column_to_numpy(table, ACTION_EEF_COLUMN)
validate_eef62(state_eef, source_joint58=source_state, label=f"{path}:{STATE_EEF_COLUMN}")
validate_eef62(action_eef, source_joint58=source_action, label=f"{path}:{ACTION_EEF_COLUMN}")
if verify_source_fk:
convert = converter or joint58_to_eef62_batch
expected_state = np.asarray(convert(source_state), dtype=np.float32)
expected_action = np.asarray(convert(source_action), dtype=np.float32)
if not np.allclose(state_eef, expected_state, atol=2e-5, rtol=1e-5):
raise DatasetSchemaError(f"{path}: state EEF does not roundtrip through FK")
if not np.allclose(action_eef, expected_action, atol=2e-5, rtol=1e-5):
raise DatasetSchemaError(f"{path}: action EEF does not roundtrip through FK")
return {
"path": str(path),
"num_frames": int(table.num_rows),
"schema_version": SCHEMA_VERSION,
"sampling_20hz": sampling_summary,
"force_only": force_summary,
}
def output_is_valid(path: Path) -> tuple[bool, str]:
try:
validate_episode_parquet(path, verify_source_fk=False)
except Exception as exc: # Validation intentionally collapses to a skip decision.
return False, str(exc)
return True, "valid"
def build_episode(
*,
dataset_root: Path,
episode_index: int,
track_path: Path,
converter: EefConverter | None = None,
verify_source_fk: bool = True,
) -> dict[str, object]:
"""Merge one validated track cache and reliable FK columns into a parquet."""
_, pq = _import_pyarrow()
info = _load_json(dataset_root / "meta" / "info.json")
parquet_path = episode_parquet_path(dataset_root, episode_index, info)
if not parquet_path.is_file():
raise FileNotFoundError(parquet_path)
table = pq.read_table(parquet_path)
old_names = tuple(table.column_names)
expected_frames = int(table.num_rows)
payload = load_track_payload(
track_path,
expected_frames=expected_frames,
episode_index=episode_index,
)
state58 = _column_to_numpy(table, "observation.state")
action58 = _column_to_numpy(table, "action")
state_eef, action_eef = convert_eef_columns(
state58,
action58,
converter=converter,
)
output = table
for column, values in track_features_from_payload(payload).items():
output = _set_or_append_column(output, column, values)
output = _set_or_append_column(output, STATE_EEF_COLUMN, state_eef)
output = _set_or_append_column(output, ACTION_EEF_COLUMN, action_eef)
if any(name not in output.column_names for name in old_names):
raise AssertionError("an original parquet column was dropped")
schema_metadata = dict(output.schema.metadata or {})
schema_metadata[PARQUET_SCHEMA_METADATA_KEY] = SCHEMA_VERSION.encode("utf-8")
output = output.replace_schema_metadata(schema_metadata)
_atomic_write_parquet(output, parquet_path, expected_rows=expected_frames)
validation = validate_episode_parquet(
parquet_path,
expected_frames=expected_frames,
verify_source_fk=verify_source_fk,
converter=converter,
)
return {
"episode_index": int(episode_index),
"num_frames": expected_frames,
"parquet": str(parquet_path.relative_to(dataset_root)),
"track_npz": str(track_path),
"track_sha256": _sha256(track_path),
"sampling_20hz": validation["sampling_20hz"],
"force_only": validation["force_only"],
"validated_at": _utc_now(),
}
def _eef_feature_names(prefix: str) -> list[str]:
rotation_names = [
"rot6d_col1_x",
"rot6d_col1_y",
"rot6d_col1_z",
"rot6d_col2_x",
"rot6d_col2_y",
"rot6d_col2_z",
]
names = [
f"left_{prefix}_x",
f"left_{prefix}_y",
f"left_{prefix}_z",
*(f"left_{prefix}_{name}" for name in rotation_names),
*(f"left_hand_q_{index}" for index in range(22)),
f"right_{prefix}_x",
f"right_{prefix}_y",
f"right_{prefix}_z",
*(f"right_{prefix}_{name}" for name in rotation_names),
*(f"right_hand_q_{index}" for index in range(22)),
]
if len(names) != 62:
raise AssertionError("EEF feature names must have length 62")
return names
def _new_feature_metadata() -> dict[str, dict]:
features: dict[str, dict] = {
TRACK_XY_COLUMN: {
"dtype": "float32",
"shape": [NUM_COMBINED_POINTS, 2],
"names": None,
},
TRACK_VISIBILITY_COLUMN: {
"dtype": "float32",
"shape": [NUM_COMBINED_POINTS],
"names": None,
},
}
features.update(
{
column: {
"dtype": "float32",
"shape": [VIEW_POINT_COUNTS[view], 3],
"names": None,
}
for view, column in TRACK_COLUMNS.items()
}
)
features[STATE_EEF_COLUMN] = {
"dtype": "float32",
"shape": [62],
"names": _eef_feature_names("eef"),
}
features[ACTION_EEF_COLUMN] = {
"dtype": "float32",
"shape": [62],
"names": _eef_feature_names("eef_target"),
}
return features
def _force_only_metadata() -> dict[str, object]:
return {
"source_column": FORCE_COLUMN,
"stored_shape": [FORCE_FLAT_DIM],
"reshape": [FORCE_SENSOR_COUNT, FORCE_SENSOR_DIM],
"target_rate_hz": 5.0,
"action_rate_hz": TARGET_RATE_HZ,
"action_update_stride": 4,
"action_chunk_offsets": [0, 4, 8, 12],
"history_frames": FORCE_HISTORY_FRAMES,
"history_duration_seconds": FORCE_HISTORY_FRAMES / 5.0,
"history_encoding": "online_model_encoder",
"vq_codes_on_disk": False,
"deformation_maps_used": False,
}
def _autoregressive_metadata() -> dict[str, object]:
return {
"blocks": AUTOREGRESSIVE_BLOCKS,
"action_steps_per_block": ACTION_CHUNK_STEPS,
"action_steps_per_sample": AUTOREGRESSIVE_BLOCKS * ACTION_CHUNK_STEPS,
"video_conditioning_frames": 1,
"video_frames_per_block": VIDEO_FRAMES_PER_BLOCK,
"video_frames_per_sample": TRAINING_VIDEO_FRAMES,
}
def _eef_modality_entries(original_key: str, *, action: bool) -> dict[str, dict]:
prefix = "eef62_absolute" if action else "eef62"
def entry(start: int, end: int, rotation_type: str | None = None) -> dict:
return {
"original_key": original_key,
"start": start,
"end": end,
"rotation_type": rotation_type,
"absolute": True,
"dtype": "float32",
"range": None,
}
return {
prefix: entry(0, 62),
f"left_{prefix}_position": entry(0, 3),
f"left_{prefix}_rotation_6d": entry(3, 9, "rotation_6d"),
f"left_{prefix}_hand": entry(9, 31),
f"right_{prefix}_position": entry(31, 34),
f"right_{prefix}_rotation_6d": entry(34, 40, "rotation_6d"),
f"right_{prefix}_hand": entry(40, 62),
}
def _statistics(values: np.ndarray) -> dict[str, list]:
array = np.asarray(values, dtype=np.float64)
if array.ndim < 2 or not np.isfinite(array).all():
raise DatasetSchemaError(f"cannot compute stats for shape {array.shape}")
return {
"mean": np.mean(array, axis=0).tolist(),
"std": np.std(array, axis=0).tolist(),
"min": np.min(array, axis=0).tolist(),
"max": np.max(array, axis=0).tolist(),
"q01": np.quantile(array, 0.01, axis=0).tolist(),
"q99": np.quantile(array, 0.99, axis=0).tolist(),
}
def _rotation_6d_to_matrix(rotation_6d: np.ndarray) -> np.ndarray:
values = np.asarray(rotation_6d, dtype=np.float64)
if values.shape[-1] != 6:
raise DatasetSchemaError("rotation_6d must end in six values")
first = values[..., :3]
first /= np.linalg.norm(first, axis=-1, keepdims=True).clip(min=1e-8)
second = values[..., 3:6]
second = second - np.sum(first * second, axis=-1, keepdims=True) * first
second /= np.linalg.norm(second, axis=-1, keepdims=True).clip(min=1e-8)
third = np.cross(first, second)
return np.stack((first, second, third), axis=-1)
def eef62_delta_base(
reference_state: np.ndarray, absolute_targets: np.ndarray
) -> np.ndarray:
"""T-Rex chunk-start-frame action: relative EEF pose + absolute hand joints."""
reference = np.asarray(reference_state, dtype=np.float64)
targets = np.asarray(absolute_targets, dtype=np.float64)
if reference.shape != (62,) or targets.shape[-1] != 62:
raise DatasetSchemaError("delta-base conversion expects [62] and [...,62]")
output = np.empty_like(targets, dtype=np.float64)
for pose_slice, hand_slice in (
(LEFT_EEF, LEFT_HAND_EEF),
(RIGHT_EEF, RIGHT_HAND_EEF),
):
reference_pose = reference[pose_slice]
target_pose = targets[..., pose_slice]
reference_rotation = _rotation_6d_to_matrix(reference_pose[3:9])
target_rotation = _rotation_6d_to_matrix(target_pose[..., 3:9])
delta_xyz = np.einsum(
"ji,...j->...i",
reference_rotation,
target_pose[..., :3] - reference_pose[:3],
)
delta_rotation = np.einsum(
"ji,...jk->...ik", reference_rotation, target_rotation
)
output[..., pose_slice] = np.concatenate(
(
delta_xyz,
delta_rotation[..., :, 0],
delta_rotation[..., :, 1],
),
axis=-1,
)
output[..., hand_slice] = targets[..., hand_slice]
return output.astype(np.float32)
def compute_delta_base_stats(parquet_paths: Iterable[Path]) -> dict[str, list]:
"""Pool all complete 16-step 20 Hz chunks for action normalization."""
_, pq = _import_pyarrow()
chunks: list[np.ndarray] = []
boundary_fallback_chunks: list[np.ndarray] = []
offsets = np.arange(ACTION_CHUNK_STEPS, dtype=np.int64)
for path in parquet_paths:
table = pq.read_table(
path,
columns=["timestamp", STATE_EEF_COLUMN, ACTION_EEF_COLUMN],
)
timestamps = np.asarray(table["timestamp"].to_numpy(), dtype=np.float64)
state = _column_to_numpy(table, STATE_EEF_COLUMN)
action = _column_to_numpy(table, ACTION_EEF_COLUMN)
anchors = sample_timestamps_nearest(
timestamps,
target_rate_hz=TARGET_RATE_HZ,
anchor_index=0,
)
anchor_indices = np.asarray(anchors["indices"], dtype=np.int64)
anchor_times = np.asarray(anchors["target_timestamps"], dtype=np.float64)
for anchor_index, anchor_time in zip(anchor_indices, anchor_times):
selection = sample_timestamps_nearest(
timestamps,
target_rate_hz=TARGET_RATE_HZ,
anchor_timestamp=float(anchor_time),
offsets=offsets,
)
target_indices = np.asarray(selection["indices"], dtype=np.int64)
delta_chunk = eef62_delta_base(
state[int(anchor_index)], action[target_indices]
)
if np.asarray(selection["padding_mask"], dtype=bool).any():
boundary_fallback_chunks.append(delta_chunk)
else:
chunks.append(delta_chunk)
if not chunks:
# Tiny schema fixtures and unusually short episodes cannot contain a
# complete horizon. Edge-clamped values keep metadata writable, while
# the runtime loader still excludes these anchors from training.
chunks = boundary_fallback_chunks
if not chunks:
raise DatasetSchemaError("no rows are available for delta-base stats")
return _statistics(np.concatenate(chunks, axis=0))
def compute_new_stats(parquet_paths: Iterable[Path]) -> dict[str, dict]:
_, pq = _import_pyarrow()
stat_columns = (*NEW_COLUMNS, FORCE_COLUMN)
buffers: dict[str, list[np.ndarray]] = {column: [] for column in stat_columns}
for path in parquet_paths:
table = pq.read_table(path, columns=list(stat_columns))
for column in stat_columns:
buffers[column].append(_column_to_numpy(table, column))
if not all(buffers.values()):
raise DatasetSchemaError("no valid converted episodes are available for stats")
return {
column: _statistics(np.concatenate(parts, axis=0))
for column, parts in buffers.items()
}
def _valid_converted_episodes(
dataset_root: Path,
*,
info: dict,
) -> tuple[list[int], list[Path]]:
indices: list[int] = []
paths: list[Path] = []
for episode_index in range(int(info["total_episodes"])):
path = episode_parquet_path(dataset_root, episode_index, info)
valid, _ = output_is_valid(path)
if valid:
indices.append(episode_index)
paths.append(path)
return indices, paths
def update_metadata(
dataset_root: Path,
*,
assume_all_converted: bool = False,
) -> dict[str, object]:
"""Atomically update info/modality/stats while preserving all old entries."""
meta_dir = dataset_root / "meta"
info_path = meta_dir / "info.json"
modality_path = meta_dir / "modality.json"
stats_path = meta_dir / "stats.json"
info = _load_json(info_path)
modality = _load_json(modality_path)
stats = _load_json(stats_path) if stats_path.exists() else {}
if assume_all_converted:
converted_indices = list(range(int(info["total_episodes"])))
converted_paths = [
episode_parquet_path(dataset_root, episode_index, info)
for episode_index in converted_indices
]
missing = [path for path in converted_paths if not path.is_file()]
if missing:
raise FileNotFoundError(missing[0])
else:
converted_indices, converted_paths = _valid_converted_episodes(
dataset_root,
info=info,
)
if not converted_paths:
raise DatasetSchemaError("metadata cannot be updated before one valid episode exists")
new_stats = compute_new_stats(converted_paths)
features = info.setdefault("features", {})
force_feature = features.get(FORCE_COLUMN)
if not isinstance(force_feature, dict) or force_feature.get("shape") != [
FORCE_FLAT_DIM
]:
raise DatasetSchemaError(
f"info.json must declare existing {FORCE_COLUMN} with shape [{FORCE_FLAT_DIM}]"
)
if "float" not in str(force_feature.get("dtype", "")):
raise DatasetSchemaError(f"info.json {FORCE_COLUMN} must be floating-point")
features.update(_new_feature_metadata())
info["trex_track_force"] = {
"schema_version": SCHEMA_VERSION,
"track_layout": layout_metadata(),
"sampling_20hz": {
"source_column": "timestamp",
"target_rate_hz": TARGET_RATE_HZ,
"method": "deterministic_nearest_earlier_on_tie",
"source_data_overwritten": False,
"action_chunk_steps": ACTION_CHUNK_STEPS,
"action_chunk_duration_seconds": ACTION_CHUNK_DURATION_SECONDS,
"action_chunk_timestamp_span_seconds": (
ACTION_CHUNK_TIMESTAMP_SPAN_SECONDS
),
},
"autoregressive_training": _autoregressive_metadata(),
"force_only": _force_only_metadata(),
"eef62_layout": {
"order": [
"left_eef_pose9",
"left_hand22",
"right_eef_pose9",
"right_hand22",
],
"slices": {
"left_eef_pose9": [0, 9],
"left_hand22": [9, 31],
"right_eef_pose9": [31, 40],
"right_hand22": [40, 62],
},
"pose9": "translation_xyz + rotation_matrix_column_1 + rotation_matrix_column_2",
"source": "T-Rex trex_fk + utils/lerobot_common.py",
},
"converted_episode_indices": converted_indices,
"complete": len(converted_indices) == int(info["total_episodes"]),
"updated_at": _utc_now(),
}
modality.setdefault("state", {}).update(
_eef_modality_entries(STATE_EEF_COLUMN, action=False)
)
modality.setdefault("action", {}).update(
_eef_modality_entries(ACTION_EEF_COLUMN, action=True)
)
# The dedicated loader exposes this alias after applying T-Rex delta-base
# conversion. The source parquet remains absolute and is never overwritten.
modality["action"]["eef62"] = {
"original_key": ACTION_EEF_COLUMN,
"start": 0,
"end": 62,
"rotation_type": None,
"absolute": False,
"dtype": "float32",
"range": None,
}
modality["track"] = {
"xy": {
"original_key": TRACK_XY_COLUMN,
"shape": [NUM_COMBINED_POINTS, 2],
"coordinate_space": "normalized_xy_div_wh",
},
"visibility": {
"original_key": TRACK_VISIBILITY_COLUMN,
"shape": [NUM_COMBINED_POINTS],
"range": [0.0, 1.0],
},
"views": {
view: {
"original_key": column,
"shape": [VIEW_POINT_COUNTS[view], 3],
"value_order": ["x", "y", "visibility"],
"coordinate_space": "normalized_xy_div_wh",
"slice": list(VIEW_SLICES[view]),
}
for view, column in TRACK_COLUMNS.items()
},
}
modality["force"] = {
"current": {
"original_key": FORCE_COLUMN,
"stored_shape": [FORCE_FLAT_DIM],
"reshape": [FORCE_SENSOR_COUNT, FORCE_SENSOR_DIM],
},
"history": {
"original_key": FORCE_COLUMN,
"frames": FORCE_HISTORY_FRAMES,
"target_rate_hz": 5.0,
"action_update_stride": 4,
"encoding": "online_model_encoder",
"vq_codes_on_disk": False,
},
}
stats.update(new_stats)
relative_stats_path = meta_dir / RELATIVE_ACTION_STATS_FILENAME
relative_stats = {"eef62": compute_delta_base_stats(converted_paths)}
for path in (info_path, modality_path, stats_path, relative_stats_path):
_atomic_backup(path)
_atomic_write_json(info_path, info)
_atomic_write_json(modality_path, modality)
_atomic_write_json(stats_path, stats)
_atomic_write_json(relative_stats_path, relative_stats)
return {
"converted_episode_indices": converted_indices,
"complete": info["trex_track_force"]["complete"],
"stats_episode_count": len(converted_indices),
}
def validate_metadata(dataset_root: Path) -> None:
info = _load_json(dataset_root / "meta" / "info.json")
modality = _load_json(dataset_root / "meta" / "modality.json")
stats = _load_json(dataset_root / "meta" / "stats.json")
relative_stats = _load_json(
dataset_root / "meta" / RELATIVE_ACTION_STATS_FILENAME
)
feature_specs = _new_feature_metadata()
for column, expected in feature_specs.items():
if info.get("features", {}).get(column) != expected:
raise DatasetSchemaError(f"info.json has invalid feature metadata for {column}")
if column not in stats:
raise DatasetSchemaError(f"stats.json is missing {column}")
schema_block = info.get("trex_track_force", {})
if schema_block.get("schema_version") != SCHEMA_VERSION:
raise DatasetSchemaError("info.json is missing the track-force schema version")
if schema_block.get("track_layout") != layout_metadata():
raise DatasetSchemaError("info.json has unstable track identity metadata")
if schema_block.get("sampling_20hz", {}).get("target_rate_hz") != TARGET_RATE_HZ:
raise DatasetSchemaError("info.json is missing the 20 Hz sampling contract")
if schema_block.get("force_only") != _force_only_metadata():
raise DatasetSchemaError("info.json has invalid force-only metadata")
if schema_block.get("autoregressive_training") != _autoregressive_metadata():
raise DatasetSchemaError("info.json has invalid autoregressive training metadata")
force_feature = info.get("features", {}).get(FORCE_COLUMN, {})
if force_feature.get("shape") != [FORCE_FLAT_DIM]:
raise DatasetSchemaError(f"info.json has invalid {FORCE_COLUMN} shape")
if FORCE_COLUMN not in stats or any(
len(stats[FORCE_COLUMN].get(name, [])) != FORCE_FLAT_DIM
for name in ("mean", "std", "min", "max", "q01", "q99")
):
raise DatasetSchemaError(f"stats.json is missing 60-D {FORCE_COLUMN} stats")
if "observation.force_history_vq" in info.get("features", {}):
raise DatasetSchemaError("metadata must not declare fabricated force VQ codes")
for name in _eef_modality_entries(STATE_EEF_COLUMN, action=False):
if name not in modality.get("state", {}):
raise DatasetSchemaError(f"modality.json is missing state.{name}")
for name in _eef_modality_entries(ACTION_EEF_COLUMN, action=True):
if name not in modality.get("action", {}):
raise DatasetSchemaError(f"modality.json is missing action.{name}")
if modality.get("action", {}).get("eef62", {}).get("absolute") is not False:
raise DatasetSchemaError("modality.json is missing delta-base action.eef62")
delta_stats = relative_stats.get("eef62", {})
if set(delta_stats) != {"mean", "std", "min", "max", "q01", "q99"}:
raise DatasetSchemaError("relative action stats are missing action.eef62")
if any(len(delta_stats[name]) != 62 for name in delta_stats):
raise DatasetSchemaError("relative action.eef62 stats must have 62 values")
track_meta = modality.get("track", {})
if track_meta.get("xy", {}).get("original_key") != TRACK_XY_COLUMN:
raise DatasetSchemaError("modality.json has invalid track XY mapping")
if (
track_meta.get("visibility", {}).get("original_key")
!= TRACK_VISIBILITY_COLUMN
):
raise DatasetSchemaError("modality.json has invalid track visibility mapping")
if set(track_meta.get("views", {})) != set(VIEW_ORDER):
raise DatasetSchemaError("modality.json has invalid track views")
force_meta = modality.get("force", {})
if force_meta.get("current", {}).get("original_key") != FORCE_COLUMN:
raise DatasetSchemaError("modality.json has invalid current force source")
if force_meta.get("history", {}).get("encoding") != "online_model_encoder":
raise DatasetSchemaError("modality.json must encode force history online")
if force_meta.get("history", {}).get("vq_codes_on_disk") is not False:
raise DatasetSchemaError("modality.json must not claim on-disk VQ codes")
def _new_manifest(dataset_root: Path, track_cache: Path) -> dict:
return {
"schema_version": SCHEMA_VERSION,
"track_layout_version": TRACK_LAYOUT_VERSION,
"track_layout": layout_metadata(),
"sampling_contract": {
"source_column": "timestamp",
"target_rate_hz": TARGET_RATE_HZ,
"action_chunk_steps": ACTION_CHUNK_STEPS,
"action_chunk_duration_seconds": ACTION_CHUNK_DURATION_SECONDS,
"action_chunk_timestamp_span_seconds": (
ACTION_CHUNK_TIMESTAMP_SPAN_SECONDS
),
},
"autoregressive_training": _autoregressive_metadata(),
"force_only": _force_only_metadata(),
"dataset_root": str(dataset_root),
"track_cache": str(track_cache),
"created_at": _utc_now(),
"updated_at": _utc_now(),
"episodes": {},
}
def load_manifest(path: Path, *, dataset_root: Path, track_cache: Path) -> dict:
if not path.exists():
return _new_manifest(dataset_root, track_cache)
manifest = _load_json(path)
if manifest.get("schema_version") != SCHEMA_VERSION:
fresh = _new_manifest(dataset_root, track_cache)
fresh["supersedes_schema_version"] = manifest.get("schema_version")
fresh["stale_episode_entries_discarded"] = len(manifest.get("episodes", {}))
return fresh
if manifest.get("track_layout") != layout_metadata():
raise DatasetSchemaError(f"{path}: manifest point layout is not canonical")
manifest["autoregressive_training"] = _autoregressive_metadata()
manifest.setdefault("episodes", {})
return manifest
def write_manifest(path: Path, manifest: dict) -> None:
manifest["updated_at"] = _utc_now()
_atomic_backup(path)
_atomic_write_json(path, manifest)
def select_episode_indices(
total_episodes: int,
*,
episode_index: int | None = None,
episode_range: Sequence[int] | None = None,
all_episodes: bool = False,
) -> list[int]:
modes = int(episode_index is not None) + int(episode_range is not None) + int(all_episodes)
if modes != 1:
raise ValueError("select exactly one of episode_index, episode_range, or all_episodes")
if episode_index is not None:
result = [int(episode_index)]
elif episode_range is not None:
if len(episode_range) != 2:
raise ValueError("episode_range must contain START END")
start, end = map(int, episode_range)
if end <= start:
raise ValueError("episode range is half-open and requires END > START")
result = list(range(start, end))
else:
result = list(range(int(total_episodes)))
invalid = [index for index in result if index < 0 or index >= int(total_episodes)]
if invalid:
raise ValueError(
f"episode indices out of range [0,{total_episodes}): {invalid[:5]}"
)
return result
def validate_dataset(
*,
dataset_root: Path,
episode_indices: Sequence[int],
manifest_path: Path,
verify_fk: bool,
) -> list[dict[str, object]]:
info = _load_json(dataset_root / "meta" / "info.json")
manifest = load_manifest(
manifest_path,
dataset_root=dataset_root,
track_cache=default_track_cache(dataset_root),
)
results: list[dict[str, object]] = []
for episode_index in episode_indices:
path = episode_parquet_path(dataset_root, episode_index, info)
result = validate_episode_parquet(
path,
verify_source_fk=verify_fk,
)
manifest_entry = manifest.get("episodes", {}).get(f"{episode_index:06d}")
if not manifest_entry or manifest_entry.get("status") != "complete":
raise DatasetSchemaError(
f"manifest has no complete entry for episode {episode_index}"
)
recorded_sampling = manifest_entry.get("sampling_20hz", {})
current_sampling = result["sampling_20hz"]
for name in (
"source_frame_count",
"target_rate_hz",
"target_sample_count",
"action_chunk_steps",
"action_chunk_duration_seconds",
"action_chunk_timestamp_span_seconds",
):
if recorded_sampling.get(name) != current_sampling[name]:
raise DatasetSchemaError(
f"manifest episode {episode_index} has stale sampling field {name}"
)
if not np.isclose(
float(recorded_sampling.get("source_rate_hz", np.nan)),
float(current_sampling["source_rate_hz"]),
rtol=1e-9,
atol=1e-9,
):
raise DatasetSchemaError(
f"manifest episode {episode_index} has stale source_rate_hz"
)
if manifest_entry.get("force_only") != result["force_only"]:
raise DatasetSchemaError(
f"manifest episode {episode_index} has stale force-only metadata"
)
results.append(result)
validate_metadata(dataset_root)
return results
def _ensure_track_npz(
*,
dataset_root: Path,
track_cache: Path,
episode_index: int,
expected_frames: int,
args: argparse.Namespace,
runtime_holder: dict[str, object],
) -> Path:
path = track_npz_path(track_cache, episode_index)
try:
load_track_payload(
path,
expected_frames=expected_frames,
episode_index=episode_index,
)
return path
except (FileNotFoundError, DatasetSchemaError) as exc:
if not args.extract_missing:
raise DatasetSchemaError(
f"episode {episode_index}: no valid track cache and extraction is disabled"
) from exc
print(f"episode {episode_index}: extracting tracks ({exc})")
if "runtime" not in runtime_holder:
import extract_track
runtime_holder["module"] = extract_track
runtime_holder["runtime"] = extract_track.create_tracking_runtime(
calib_path=args.calib_path,
openpi_root=args.openpi_root,
cotracker_checkpoint=args.cotracker_checkpoint,
cotracker_device=args.cotracker_device,
sam2_model=args.sam2_model,
sam2_device=args.sam2_device,
sam2_libs=args.sam2_libs,
image_height=args.image_height,
image_width=args.image_width,
)
module = runtime_holder["module"]
runtime = runtime_holder["runtime"]
viz_dir = track_cache / "viz_tracks"
masks_dir = track_cache / "sam2_masks"
module.process_episode(
dataset_root=dataset_root,
episode_index=episode_index,
output_path=track_cache,
calib=runtime.calib,
out_hw=runtime.out_hw,
cotracker_model=runtime.cotracker_model,
cotracker_device=runtime.cotracker_device,
save_viz=bool(args.save_viz),
viz_out_dir=viz_dir,
viz_fps=int(args.viz_fps),
viz_trail=int(args.viz_trail),
sam2_predictor=runtime.sam2_predictor,
sam2_seed=int(args.sam2_seed),
save_sam2_masks_flag=bool(args.save_sam2_masks),
sam2_masks_dir=masks_dir,
)
load_track_payload(
path,
expected_frames=expected_frames,
episode_index=episode_index,
)
return path
def _build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(
description=__doc__,
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
)
parser.add_argument("--dataset-root", type=Path, default=DEFAULT_DATASET_ROOT)
selection = parser.add_mutually_exclusive_group(required=True)
selection.add_argument("--episode-index", type=int)
selection.add_argument(
"--episode-range",
type=int,
nargs=2,
metavar=("START", "END"),
help="Half-open episode range [START, END)",
)
selection.add_argument("--all", dest="all_episodes", action="store_true")
parser.add_argument("--track-cache", type=Path, default=None)
parser.add_argument("--manifest-path", type=Path, default=None)
parser.add_argument("--dry-run", action="store_true")
parser.add_argument("--validate-only", action="store_true")
parser.add_argument("--force", action="store_true", help="Rebuild even valid parquets")
parser.add_argument(
"--extract-missing",
action=argparse.BooleanOptionalAction,
default=True,
)
parser.add_argument(
"--update-metadata",
action=argparse.BooleanOptionalAction,
default=True,
)
parser.add_argument(
"--verify-fk",
action=argparse.BooleanOptionalAction,
default=True,
)
# Heavy extraction dependencies are only imported if a cache is missing.
parser.add_argument(
"--calib-path",
type=Path,
default=_DREAMZERO_ROOT / "assets" / "trex_camera_calib.json",
)
parser.add_argument(
"--openpi-root",
type=Path,
default=Path("/scratch2/home/zhicao/openpi"),
)
parser.add_argument("--cotracker-checkpoint", type=str, default="")
parser.add_argument("--cotracker-device", type=str, default="")
parser.add_argument(
"--sam2-model",
type=str,
default=os.environ.get("SAM2_MODEL", "facebook/sam2-hiera-large"),
)
parser.add_argument("--sam2-device", type=str, default="")
parser.add_argument("--sam2-seed", type=int, default=0)
parser.add_argument(
"--sam2-libs",
type=Path,
default=Path(os.environ.get("SAM2_LIBS", "/scratch1/home/zhicao/physctrl/libs")),
)
parser.add_argument("--image-height", type=int, default=0)
parser.add_argument("--image-width", type=int, default=0)
parser.add_argument(
"--save-viz",
action=argparse.BooleanOptionalAction,
default=False,
)
parser.add_argument("--viz-fps", type=int, default=10)
parser.add_argument("--viz-trail", type=int, default=15)
parser.add_argument(
"--save-sam2-masks",
action=argparse.BooleanOptionalAction,
default=False,
)
return parser
def main(argv: Sequence[str] | None = None) -> int:
args = _build_parser().parse_args(argv)
if args.dry_run and args.validate_only:
raise ValueError("--dry-run and --validate-only are mutually exclusive")
dataset_root = args.dataset_root.expanduser().resolve()
info = _load_json(dataset_root / "meta" / "info.json")
episode_indices = select_episode_indices(
int(info["total_episodes"]),
episode_index=args.episode_index,
episode_range=args.episode_range,
all_episodes=bool(args.all_episodes),
)
track_cache = (
args.track_cache.expanduser().resolve()
if args.track_cache is not None
else default_track_cache(dataset_root)
)
manifest_path = (
args.manifest_path.expanduser().resolve()
if args.manifest_path is not None
else dataset_root / "meta" / "trex_track_force_manifest.json"
)
if args.validate_only:
results = validate_dataset(
dataset_root=dataset_root,
episode_indices=episode_indices,
manifest_path=manifest_path,
verify_fk=bool(args.verify_fk),
)
for result in results:
sampling = result["sampling_20hz"]
print(
f"{Path(str(result['path'])).name}: "
f"source={sampling['source_rate_hz']:.6f}Hz "
f"target={sampling['target_rate_hz']:.1f}Hz "
f"samples={sampling['target_sample_count']} "
f"chunk={sampling['action_chunk_steps']} steps/"
f"{sampling['action_chunk_duration_seconds']:.1f}s "
f"(timestamp span "
f"{sampling['action_chunk_timestamp_span_seconds']:.2f}s)"
)
print(f"validated {len(results)} episode(s)")
return 0
if args.dry_run:
_, pq = _import_pyarrow()
for episode_index in episode_indices:
parquet_path = episode_parquet_path(dataset_root, episode_index, info)
valid, reason = output_is_valid(parquet_path)
cache_path = track_npz_path(track_cache, episode_index)
cache_valid = False
cache_reason = "missing"
if cache_path.exists() and parquet_path.exists():
try:
expected_frames = int(pq.read_metadata(parquet_path).num_rows)
load_track_payload(
cache_path,
expected_frames=expected_frames,
episode_index=episode_index,
)
cache_valid = True
cache_reason = "valid"
except Exception as exc:
cache_reason = str(exc)
if valid and not args.force:
action = "skip valid output"
elif cache_valid:
action = "merge cache + FK"
elif args.extract_missing:
action = f"extract SAM2/CoTracker, merge + FK (cache: {cache_reason})"
else:
action = f"FAIL: no valid track cache ({cache_reason})"
print(f"[dry-run] episode {episode_index:06d}: {action} ({reason})")
print("[dry-run] no files were changed")
return 0
manifest = load_manifest(
manifest_path,
dataset_root=dataset_root,
track_cache=track_cache,
)
runtime_holder: dict[str, object] = {}
_, pq = _import_pyarrow()
for episode_index in episode_indices:
key = f"{episode_index:06d}"
parquet_path = episode_parquet_path(dataset_root, episode_index, info)
valid, reason = output_is_valid(parquet_path)
try:
validated_summary: dict[str, object] | None = None
if valid and not args.force:
try:
validated_summary = validate_episode_parquet(
parquet_path,
verify_source_fk=bool(args.verify_fk),
)
except DatasetSchemaError as exc:
valid = False
reason = f"deep validation failed: {exc}"
if valid and not args.force:
if validated_summary is None:
raise AssertionError("valid output was not validated")
summary = validated_summary
summary.update(
{
"episode_index": episode_index,
"status": "complete",
"skipped": True,
"validated_at": _utc_now(),
}
)
print(f"episode {episode_index:06d}: skip valid output")
else:
if not parquet_path.is_file():
raise FileNotFoundError(parquet_path)
expected_frames = int(pq.read_metadata(parquet_path).num_rows)
cache_path = _ensure_track_npz(
dataset_root=dataset_root,
track_cache=track_cache,
episode_index=episode_index,
expected_frames=expected_frames,
args=args,
runtime_holder=runtime_holder,
)
summary = build_episode(
dataset_root=dataset_root,
episode_index=episode_index,
track_path=cache_path,
verify_source_fk=bool(args.verify_fk),
)
summary["status"] = "complete"
summary["skipped"] = False
print(f"episode {episode_index:06d}: built and validated ({reason})")
manifest["episodes"][key] = summary
write_manifest(manifest_path, manifest)
except Exception as exc:
manifest["episodes"][key] = {
"episode_index": episode_index,
"status": "failed",
"error": f"{type(exc).__name__}: {exc}",
"failed_at": _utc_now(),
}
write_manifest(manifest_path, manifest)
raise
if args.update_metadata:
manifest["metadata"] = update_metadata(dataset_root)
write_manifest(manifest_path, manifest)
if args.update_metadata:
results = validate_dataset(
dataset_root=dataset_root,
episode_indices=episode_indices,
manifest_path=manifest_path,
# Every built/skipped episode was already deep-FK checked above.
verify_fk=False,
)
else:
results = [
validate_episode_parquet(
episode_parquet_path(dataset_root, episode_index, info),
verify_source_fk=False,
)
for episode_index in episode_indices
]
print(
f"completed {len(results)} episode(s); manifest={manifest_path}; "
f"metadata={'updated' if args.update_metadata else 'unchanged'}"
)
return 0
if __name__ == "__main__":
raise SystemExit(main())