a2c2 / src /dataset.py
dennis96's picture
Upload folder using huggingface_hub
0adab2f verified
Raw
History Blame Contribute Delete
7.84 kB
"""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")