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