File size: 7,843 Bytes
0adab2f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 | """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")
|