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