Download scripts/bridge_cte_data.py from DanTim05/DynaTTT: direct link, hf CLI and curl.
- Browser
- Download file 8.91 kB
-
https://huggingface.co/spaces/DanTim05/DynaTTT/resolve/main/scripts/bridge_cte_data.py
- Command line
-
hf download hf://spaces/DanTim05/DynaTTT/scripts/bridge_cte_data.py
-
curl -L -o bridge_cte_data.py https://huggingface.co/spaces/DanTim05/DynaTTT/resolve/main/scripts/bridge_cte_data.py
8.91 kB
| """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) | |