"""Dataset utilities for cached A2C2 BEHAVIOR/OpenPI parquet exports.""" from __future__ import annotations import math from pathlib import Path import random from typing import Iterator, NamedTuple import numpy as np import pyarrow as pa import pyarrow.parquet as pq import torch from torch import Tensor from torch.utils.data import IterableDataset, get_worker_info class EpisodePair(NamedTuple): data_path: Path latent_path: Path def resolve_dataset_root(path: Path) -> Path: """Resolve either an A2C2 root or its parent directory.""" path = path.expanduser().resolve() if (path / "data").is_dir() and (path / "latent" / "data").is_dir(): return path candidates = [p for p in path.iterdir() if (p / "data").is_dir() and (p / "latent" / "data").is_dir()] if len(candidates) == 1: return candidates[0].resolve() if not candidates: raise FileNotFoundError(f"No A2C2 dataset root found under {path}") names = ", ".join(str(p) for p in candidates) raise ValueError(f"Multiple dataset roots found under {path}; pass one explicitly: {names}") def discover_episode_pairs(dataset_root: Path, task_dir: str | None = None) -> list[EpisodePair]: """Find matching data/latent parquet pairs.""" dataset_root = resolve_dataset_root(dataset_root) pattern = f"{task_dir}/episode_*.parquet" if task_dir else "task-*/episode_*.parquet" data_paths = sorted((dataset_root / "data").glob(pattern)) pairs: list[EpisodePair] = [] for data_path in data_paths: rel = data_path.relative_to(dataset_root / "data") latent_path = dataset_root / "latent" / "data" / rel if not latent_path.is_file(): raise FileNotFoundError(f"Missing latent parquet for {data_path}: {latent_path}") pairs.append(EpisodePair(data_path=data_path, latent_path=latent_path)) if not pairs: raise FileNotFoundError(f"No episode parquet files found in {dataset_root / 'data'}") return pairs def split_episode_pairs( pairs: list[EpisodePair], val_ratio: float, seed: int, max_episodes: int | None = None, ) -> tuple[list[EpisodePair], list[EpisodePair]]: """Shuffle episode pairs and split into train/validation subsets.""" pairs = list(pairs) rng = random.Random(seed) rng.shuffle(pairs) if max_episodes is not None: pairs = pairs[:max_episodes] val_count = int(round(len(pairs) * val_ratio)) if val_ratio > 0 and val_count == 0 and len(pairs) > 1: val_count = 1 val_pairs = pairs[:val_count] train_pairs = pairs[val_count:] if not train_pairs: raise ValueError("No training episodes left after split.") return train_pairs, val_pairs def fixed_or_variable_list_to_numpy(column: pa.ChunkedArray, dtype: np.dtype) -> np.ndarray: """Convert Arrow list/fixed-size-list columns to dense numpy arrays.""" array = column.combine_chunks() if pa.types.is_fixed_size_list(array.type): outer_size = array.type.list_size inner = array.values if pa.types.is_fixed_size_list(inner.type): inner_size = inner.type.list_size flat = inner.values.to_numpy(zero_copy_only=False) return np.asarray(flat, dtype=dtype).reshape(len(array), outer_size, inner_size) flat = inner.to_numpy(zero_copy_only=False) return np.asarray(flat, dtype=dtype).reshape(len(array), outer_size) return np.asarray(array.to_pylist(), dtype=dtype) def load_episode(pair: EpisodePair) -> dict[str, np.ndarray]: """Load one episode's state/action/chunk rows plus aligned base-policy latents.""" data = pq.read_table( pair.data_path, columns=[ "observation.state", "action", "a2c2.base_action_chunk", "a2c2.valid_action_mask", ], ) latent = pq.read_table(pair.latent_path, columns=["a2c2.base_policy_z"]) if data.num_rows != latent.num_rows: raise ValueError(f"Row mismatch: {pair.data_path} has {data.num_rows}, {pair.latent_path} has {latent.num_rows}") return { "states": fixed_or_variable_list_to_numpy(data.column("observation.state"), np.float32), "actions": fixed_or_variable_list_to_numpy(data.column("action"), np.float32), "chunks": fixed_or_variable_list_to_numpy(data.column("a2c2.base_action_chunk"), np.float32), "masks": fixed_or_variable_list_to_numpy(data.column("a2c2.valid_action_mask"), np.bool_), "zs": fixed_or_variable_list_to_numpy(latent.column("a2c2.base_policy_z"), np.float32), } class A2C2RandomSampleDataset(IterableDataset): """Randomly sample valid (source frame t, chunk offset k) training examples.""" def __init__( self, episode_pairs: list[EpisodePair], action_horizon: int, samples_per_episode: int, seed: int, total_samples: int | None = None, ) -> None: super().__init__() self.episode_pairs = list(episode_pairs) self.action_horizon = action_horizon self.samples_per_episode = samples_per_episode self.seed = seed self.total_samples = total_samples def __iter__(self) -> Iterator[dict[str, np.ndarray]]: worker = get_worker_info() worker_id = worker.id if worker else 0 num_workers = worker.num_workers if worker else 1 pairs = self.episode_pairs[worker_id::num_workers] if not pairs: return rng = np.random.default_rng(self.seed + worker_id) yielded = 0 while self.total_samples is None or yielded < self.total_samples: order = rng.permutation(len(pairs)) for episode_idx in order: episode = load_episode(pairs[int(episode_idx)]) rows = episode["actions"].shape[0] for _ in range(self.samples_per_episode): if self.total_samples is not None and yielded >= self.total_samples: return source_idx = int(rng.integers(0, rows)) valid_offsets = np.flatnonzero(episode["masks"][source_idx]) if valid_offsets.size == 0: continue k = int(rng.choice(valid_offsets)) target_idx = source_idx + k if target_idx >= rows: continue base_action = episode["chunks"][source_idx, k] expert_action = episode["actions"][target_idx] denom = max(self.action_horizon - 1, 1) phase = 2.0 * math.pi * float(k) / denom yield { "observation_state": episode["states"][target_idx], "base_action_chunk": episode["chunks"][source_idx], "base_policy_z": episode["zs"][source_idx], "time_feature": np.asarray([math.sin(phase), math.cos(phase)], dtype=np.float32), "valid_action_mask": episode["masks"][source_idx], "base_action": base_action, "target_delta": expert_action - base_action, "expert_action": expert_action, } yielded += 1 def move_batch_to_device(batch: dict[str, Tensor], device: torch.device) -> dict[str, Tensor]: return {key: value.to(device, non_blocking=True) if torch.is_tensor(value) else value for key, value in batch.items()} def pick_device(raw: str) -> torch.device: if raw != "auto": return torch.device(raw) if torch.cuda.is_available(): return torch.device("cuda") if torch.backends.mps.is_available(): return torch.device("mps") return torch.device("cpu")