"""Whole-episode, strictly causal Bridge data for standalone CTE Stage 1.""" from __future__ import annotations import json import math import re from collections import defaultdict from pathlib import Path import av import numpy as np import pyarrow.parquet as pq import torch from torch.utils.data import Dataset, Sampler from cosmos_framework.data.generator.action.datasets.widowx_bridge_v3_dataset import normalize_actions def canonical_instruction(text: str) -> str: # Conservative equality only: do NOT merge directions, objects, or synonyms. return re.sub(r"\s+", " ", text.strip().lower()).rstrip(" .") def load_metadata(root: str | Path, seed: int = 42, val_fraction: float = 0.02): root = Path(root) info = json.loads((root / "meta/info.json").read_text()) episodes = [json.loads(line) for line in (root / "meta/episodes.jsonl").read_text().splitlines()] tasks = [json.loads(line) for line in (root / "meta/tasks.jsonl").read_text().splitlines()] if info["robot_type"] != "widowx" or info["fps"] != 5: raise ValueError("This contract requires WidowX Bridge at 5 Hz") if info["features"]["action"]["shape"] != [7]: raise ValueError("Expected seven-dimensional Bridge actions") if len({e["episode_index"] for e in episodes}) != len(episodes): raise ValueError("Duplicate episode IDs") if any(e["length"] < 1 for e in episodes): raise ValueError("Empty episode") known = sorted({canonical_instruction(t["task"]) for t in tasks} - {""}) labels = {text: index for index, text in enumerate(known)} for ep in episodes: texts = {canonical_instruction(t) for t in ep["tasks"]} # Ambiguous multi-instruction episodes do not receive a semantic label. ep["semantic_id"] = labels.get(next(iter(texts)), -1) if len(texts) == 1 else -1 order = np.random.default_rng(seed).permutation(len(episodes)) nval = max(1, round(len(episodes) * val_fraction)) if not 0 < nval < len(episodes): raise ValueError("Need nonempty training and validation splits") val_indices = set(order[:nval].tolist()) train = [e for i, e in enumerate(episodes) if i not in val_indices] val = [e for i, e in enumerate(episodes) if i in val_indices] return info, train, val class BridgeCTEDataset(Dataset): def __init__(self, root, info, episodes): self.root, self.info, self.episodes = Path(root), info, episodes tasks = [json.loads(line) for line in (self.root / "meta/tasks.jsonl").read_text().splitlines()] self.task_text = {int(t["task_index"]): canonical_instruction(t["task"]) for t in tasks} def __len__(self): return len(self.episodes) def __getitem__(self, index): ep = self.episodes[index] eid, length = int(ep["episode_index"]), int(ep["length"]) fmt = dict(episode_index=eid, episode_chunk=eid // self.info.get("chunks_size", 1000), video_key="observation.images.image_0") table = pq.read_table(self.root / self.info["data_path"].format(**fmt), columns=["action", "frame_index", "episode_index", "task_index"]) actions = np.asarray(table["action"].to_pylist(), dtype=np.float32) if actions.shape != (length, 7) or not np.isfinite(actions).all(): raise ValueError(f"Invalid actions in episode {eid}") if not np.array_equal(table["frame_index"].to_numpy(), np.arange(length)): raise ValueError(f"Nonconsecutive frame indices in episode {eid}") if not (table["episode_index"].to_numpy() == eid).all(): raise ValueError(f"Episode metadata mismatch: {eid}") row_tasks = {self.task_text[int(t)] for t in np.unique(table["task_index"].to_numpy())} if row_tasks != {canonical_instruction(t) for t in ep["tasks"]}: raise ValueError(f"Instruction metadata mismatch: {eid}") if ((actions[:, 6] < 0) | (actions[:, 6] > 1)).any(): raise ValueError(f"Gripper is not in the expected [0,1] convention: {eid}") frames = [] count = 0 with av.open(str(self.root / self.info["video_path"].format(**fmt))) as container: container.streams.video[0].thread_count = 1 for i, frame in enumerate(container.decode(video=0)): count += 1 if i % 4 == 0: image = frame.to_ndarray(format="rgb24") if image.shape != (256, 256, 3): raise ValueError(f"Unexpected image shape: {eid} {image.shape}") frames.append(image) if count != length: raise ValueError(f"Video/action alignment mismatch: {eid} {count} != {length}") # Boundary b+1 is observed only after actions [4b:4b+4]. No terminal # action or incomplete final group is an observed visual transition. ntransitions = len(frames) - 1 actions = normalize_actions(actions[:4 * ntransitions]).reshape(ntransitions, 4, 7) return dict(frames=torch.from_numpy(np.stack(frames)).permute(0, 3, 1, 2).contiguous(), actions=torch.from_numpy(actions), episode_id=eid, semantic_id=ep["semantic_id"]) def collate_episodes(samples): # Padding to >=5 keeps all effect heads in the graph, including short # episodes. Masks, not padding values, decide which losses are supervised. steps = max(5, max(len(s["frames"]) for s in samples)) batch = len(samples) valid = torch.zeros(batch, steps, dtype=torch.bool) actions = torch.zeros(batch, steps - 1, 4, 7) transition_valid = torch.zeros(batch, steps - 1, 4, dtype=torch.bool) for i, sample in enumerate(samples): count = len(sample["frames"]) valid[i, :count] = True actions[i, :count-1] = sample["actions"] transition_valid[i, :count-1] = True return dict(frames=torch.cat([s["frames"] for s in samples]), valid=valid, actions=actions, transition_valid=transition_valid, episode_ids=torch.tensor([s["episode_id"] for s in samples]), semantic_ids=torch.tensor([s["semantic_id"] for s in samples])) class PairedEpisodeSampler(Sampler): """One exhaustive shuffled anchor pass plus a same-task partner per anchor. No episode is dropped or length-truncated. Unknown/singleton anchors get a random extra episode, NOT a fake same-task positive. The loss also excludes duplicate episode IDs. Epoch/cursor reconstruction is independent of loader prefetch, so restarting at a saved consumed-batch cursor is exact. """ def __init__(self, episodes, batch_size, world_size=1, rank=0, seed=42, epoch=0, start_batch=0): if batch_size < 2 or batch_size % 2: raise ValueError("Per-rank batch size must be positive and even") self.episodes = episodes self.batch_size, self.world_size, self.rank = batch_size, world_size, rank self.seed, self.epoch, self.start_batch = seed, epoch, start_batch self.anchors_per_batch = batch_size * world_size // 2 self.total_batches = math.ceil(len(episodes) / self.anchors_per_batch) self.groups = defaultdict(list) for i, ep in enumerate(episodes): if ep["semantic_id"] >= 0: self.groups[ep["semantic_id"]].append(i) def __len__(self): return self.total_batches - self.start_batch def __iter__(self): rng = np.random.default_rng(self.seed + self.epoch) anchors = rng.permutation(len(self.episodes)) anchors = np.resize(anchors, self.total_batches * self.anchors_per_batch) pairs = [] for anchor in anchors: candidates = self.groups.get(self.episodes[anchor]["semantic_id"], []) if len(candidates) > 1: partner = anchor while partner == anchor: partner = candidates[int(rng.integers(len(candidates)))] else: partner = int(rng.integers(len(self.episodes))) pairs.extend((int(anchor), int(partner))) array = np.asarray(pairs).reshape(self.total_batches, self.world_size, self.batch_size) for b in range(self.start_batch, self.total_batches): yield array[b, self.rank].tolist() class ValidationSampler(Sampler): """Deterministic paired validation; includes all held-out anchors. Partners are also exclusively held out. The small last-batch wrap is explicit in metadata; validation metrics are sampled-batch diagnostics, not claimed to be a uniformly weighted dataset score. """ def __init__(self, episodes, batch_size, world_size, rank, seed): self.inner = PairedEpisodeSampler(episodes, batch_size, world_size, rank, seed) def __len__(self): return len(self.inner) def __iter__(self): return iter(self.inner)