DynaTTT / scripts /bridge_cte_data.py
DanTim05's picture
Upload Zeva source and project assets
68eed61 verified
Raw History Blame Contribute Delete
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)