In this tutorial, we design an end-to-end streaming robotics learning pipeline around the NVIDIA Cosmos3-DROID dataset without downloading its 707 GB repository locally. We first introspect the LeRobotDataset v3.0 structure and construct a metadata graph from info.json, task metadata, episode tables, and dataset statistics, then use HTTP byte-range access with PyArrow to selectively read Parquet row groups and columns. We convert individual episodes into state-action trajectories and analyze joint motion, gripper events, Cartesian end-effector paths, and action-frequency spectra before decoding only the required AV1 video windows through seek-based PyAV/FFmpeg access. We then normalize observations and actions using dataset statistics, construct an ACT-style chunked PyTorch dataset with optional visual conditioning, and train a multimodal behavior-cloning policy. Finally, we evaluate the learned policy through open-loop rollout with temporally ensembled action chunks, report per-joint MSE and R^2 against a mean-action baseline, visualize predicted versus ground-truth actions, and save the complete policy checkpoint for downstream use.
import subprocess, sys, os, json, math, time, warnings, random, tempfile
warnings.filterwarnings("ignore")
subprocess.run([sys.executable, "-m", "pip", "install", "-q",
"huggingface_hub>=0.34.0", "pyarrow>=15.0", "av>=12.0",
"pandas", "matplotlib", "tqdm"], check=False)
import numpy as np, pandas as pd, pyarrow as pa, pyarrow.parquet as pq
import matplotlib.pyplot as plt
from huggingface_hub import HfApi, HfFileSystem, hf_hub_download, hf_hub_url
import torch, torch.nn as nn, torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader
REPO_ID = "nvidia/Cosmos3-DROID"
ROOT = "success"
VIDEO_KEY = "observation.image.wrist_image_left"
FPS = 15
N_EPISODES = 48
HORIZON = 8
OBS_HISTORY = 2
USE_VISION = True
N_VIS_EPS = 6
VIS_SIZE = 96
EPOCHS = 12
BATCH = 256
SEED = 0
random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
DEV = "cuda" if torch.cuda.is_available() else "cpu"
print(f"[env] torch={torch.__version__} device={DEV}")
if os.environ.get("HF_TOKEN"):
from huggingface_hub import login; login(os.environ["HF_TOKEN"])
api = HfApi()
fs = HfFileSystem()
HFS = lambda rel: f"datasets/{REPO_ID}/{rel}"
URL = lambda rel: hf_hub_url(REPO_ID, rel, repo_type="dataset")
print("\n" + "="*78 + "\n1. REPO INTROSPECTION\n" + "="*78)
all_files = api.list_repo_files(REPO_ID, repo_type="dataset")
print(f"total files in repo : {len(all_files):,}")
for prefix in ("success/data", "success/videos", "success/meta",
"failure/data", "failure/videos", "failure/meta"):
print(f" {prefix:<18} {sum(f.startswith(prefix) for f in all_files):>6,} files")
data_shards = sorted(f for f in all_files if f.startswith(f"{ROOT}/data/") and f.endswith(".parquet"))
vid_shards = sorted(f for f in all_files if f.startswith(f"{ROOT}/videos/{VIDEO_KEY}/"))
meta_files = sorted(f for f in all_files if f.startswith(f"{ROOT}/meta/"))
print(f"\n[{ROOT}] data shards={len(data_shards)} video shards({VIDEO_KEY})={len(vid_shards)}")
print("first data shard :", data_shards[0])
print("first video shard:", vid_shards[0])
print("\n" + "="*78 + "\n2. METADATA\n" + "="*78)
info = json.load(open(hf_hub_download(REPO_ID, f"{ROOT}/meta/info.json", repo_type="dataset")))
print(f"episodes={info.get('total_episodes'):,} frames={info.get('total_frames'):,} "
f"tasks={info.get('total_tasks'):,} fps={info.get('fps')}")
print("data_path template :", info.get("data_path"))
print("video_path template:", info.get("video_path"))
FEATURES = info["features"]
state_keys = sorted(k for k in FEATURES if k.startswith("observation.state"))
action_keys = sorted(k for k in FEATURES if k.startswith("action."))
video_keys = sorted(k for k in FEATURES if FEATURES[k]["dtype"] == "video")
print("\nstate :", [f"{k.split('.')[-1]}{tuple(FEATURES[k]['shape'])}" for k in state_keys])
print("action :", [f"{k.split('.')[-1]}{tuple(FEATURES[k]['shape'])}" for k in action_keys])
print("video :", video_keys)
tdf = pd.read_parquet(hf_hub_download(REPO_ID, f"{ROOT}/meta/tasks.parquet", repo_type="dataset"))
tdf = tdf.reset_index()
tcol = "task" if "task" in tdf.columns else tdf.columns[0]
TASKS = dict(zip(tdf["task_index"].astype(int), tdf[tcol].astype(str))) if "task_index" in tdf \
else {i: str(v) for i, v in enumerate(tdf[tcol])}
print(f"\n{len(TASKS):,} task strings. Random sample:")
for t in random.sample(list(TASKS.values()), min(8, len(TASKS))): print(" ·", t[:90])
ep_files = [f for f in meta_files if "/episodes/" in f and f.endswith(".parquet")]
eps = pd.concat([pd.read_parquet(hf_hub_download(REPO_ID, f, repo_type="dataset"))
for f in ep_files[:4]], ignore_index=True)
print(f"\nepisodes table: {len(eps):,} rows")
print("columns:", [c for c in eps.columns if not c.startswith("stats")][:14], "...")
print(eps[[c for c in ("episode_index", "length", "data/chunk_index", "data/file_index")
if c in eps.columns]].head())
We initialize the Colab environment, install the required libraries, and configure the Cosmos3-DROID dataset, episode, video, and training parameters. We inspect the repository structure and identify the available data, video, and metadata shards without downloading the complete dataset. We then load the core metadata and task descriptions to understand the dataset schema, available state/action features, and episode organization.
print("\n" + "="*78 + "\n3. BYTE-RANGE PARQUET READER\n" + "="*78)
def open_pf(rel_path):
return pq.ParquetFile(fs.open(HFS(rel_path), "rb"))
def rowgroup_span(pf):
md, starts, c = pf.metadata, [], 0
for i in range(md.num_row_groups):
starts.append(c); c += md.row_group(i).num_rows
return np.array(starts), c
def read_rows(pf, lo, hi, columns):
starts, total = rowgroup_span(pf)
ends = np.append(starts[1:], total)
rgs = [i for i in range(len(starts)) if starts[i] < hi and ends[i] > lo]
tbl = pf.read_row_groups(rgs, columns=columns)
return tbl.slice(lo - starts[rgs[0]], hi - lo)
def col2np(tbl, name):
ca = tbl.column(name).combine_chunks()
if pa.types.is_list(ca.type) or pa.types.is_large_list(ca.type) or pa.types.is_fixed_size_list(ca.type):
flat = np.asarray(ca.flatten().to_numpy(zero_copy_only=False))
return flat.reshape(len(ca), -1).astype(np.float32)
return np.asarray(ca.to_numpy(zero_copy_only=False)).reshape(-1, 1).astype(np.float32)
SHARD = data_shards[0]
pf = open_pf(SHARD)
md = pf.metadata
print(f"shard : {SHARD}")
print(f"rows : {md.num_rows:,} row_groups: {md.num_row_groups} "
f"compressed: {md.serialized_size/1e6:.1f} MB footer")
print(f"columns : {len(pf.schema_arrow.names)}")
t0 = time.time()
ep_idx_all = pf.read(columns=["episode_index"]).column("episode_index").to_numpy()
print(f"pulled episode_index column ({len(ep_idx_all):,} rows) in {time.time()-t0:.1f}s")
uniq, first_pos = np.unique(ep_idx_all, return_index=True)
order = np.argsort(first_pos)
uniq = uniq[order]; first_pos = first_pos[order]
last_pos = np.append(first_pos[1:], len(ep_idx_all))
EP_BOUNDS = {int(e): (int(a), int(b)) for e, a, b in zip(uniq, first_pos, last_pos)}
print(f"{len(EP_BOUNDS)} episodes live in this shard "
f"(ids {uniq.min()}..{uniq.max()}, mean len {np.mean(last_pos-first_pos):.0f} frames)")
STATE_USE = ["observation.state.joint_positions", "observation.state.gripper_position",
"observation.state.cartesian_position"]
ACTION_USE = ["action.joint_velocity", "action.gripper_position"]
READ_COLS = STATE_USE + ACTION_USE + ["timestamp", "frame_index", "task_index", "episode_index"]
def load_episode(ep):
lo, hi = EP_BOUNDS[ep]
tbl = read_rows(pf, lo, hi, READ_COLS)
out = {k: col2np(tbl, k) for k in STATE_USE + ACTION_USE}
out["timestamp"] = col2np(tbl, "timestamp").ravel()
out["task_index"] = int(col2np(tbl, "task_index").ravel()[0])
out["task"] = TASKS.get(out["task_index"], "<unknown>")
out["state"] = np.concatenate([out[k] for k in STATE_USE], axis=1)
out["action"] = np.concatenate([out[k] for k in ACTION_USE], axis=1)
return out
EP0 = int(uniq[0]); traj = load_episode(EP0)
print(f"\nepisode {EP0}: T={len(traj['state'])} state_dim={traj['state'].shape[1]} "
f"action_dim={traj['action'].shape[1]}")
print(f"task: {traj['task']!r}")
print("\n" + "="*78 + "\n5. TRAJECTORY ANALYTICS\n" + "="*78)
q = traj["observation.state.joint_positions"]
grip = traj["observation.state.gripper_position"].ravel()
cart = traj["observation.state.cartesian_position"]
dq = traj["action.joint_velocity"]
t = traj["timestamp"]
fig = plt.figure(figsize=(15, 9))
ax = fig.add_subplot(2, 3, 1)
for j in range(q.shape[1]): ax.plot(t, q[:, j], lw=1.1, label=f"j{j+1}")
ax.set_title("joint positions [rad]"); ax.set_xlabel("s"); ax.legend(fontsize=6, ncol=2)
ax = fig.add_subplot(2, 3, 2)
ax.plot(t, grip, color="crimson", lw=1.4)
opens = np.where(np.abs(np.diff(grip)) > 0.05)[0]
for k in opens[:40]: ax.axvline(t[k], color="k", alpha=.15, lw=.8)
ax.set_title(f"gripper (|Δ|>0.05 events: {len(opens)})"); ax.set_xlabel("s")
ax = fig.add_subplot(2, 3, 3, projection="3d")
ax.plot(cart[:, 0], cart[:, 1], cart[:, 2], lw=1.2)
ax.scatter(*cart[0, :3], c="g", s=45, label="start"); ax.scatter(*cart[-1, :3], c="r", s=45, label="end")
ax.set_title("EE cartesian path [m]"); ax.legend(fontsize=7)
ax = fig.add_subplot(2, 3, 4)
im = ax.imshow(dq.T, aspect="auto", cmap="RdBu_r", vmin=-np.abs(dq).max(), vmax=np.abs(dq).max())
ax.set_title("action.joint_velocity (7 x T)"); ax.set_ylabel("joint"); plt.colorbar(im, ax=ax)
ax = fig.add_subplot(2, 3, 5)
freqs = np.fft.rfftfreq(len(dq), d=1/FPS)
for j in range(dq.shape[1]):
ax.semilogy(freqs, np.abs(np.fft.rfft(dq[:, j] - dq[:, j].mean())) + 1e-9, lw=.9)
ax.set_title("action spectra (Nyquist=7.5 Hz)"); ax.set_xlabel("Hz")
ax = fig.add_subplot(2, 3, 6)
lens = [EP_BOUNDS[e][1] - EP_BOUNDS[e][0] for e in list(EP_BOUNDS)[:2000]]
ax.hist(np.array(lens)/FPS, bins=40, color="steelblue")
ax.set_title(f"episode duration [s] (n={len(lens)})"); ax.set_xlabel("s")
plt.suptitle(f"{REPO_ID} · {ROOT} · ep {EP0} · {traj['task'][:70]}", y=1.0)
plt.tight_layout(); plt.show()
We implement a byte-range Parquet reader that accesses only the required row groups and columns directly through the Hugging Face filesystem. We identify episode boundaries within a data shard and convert selected state and action fields into NumPy trajectories. We then visualize joint positions, gripper activity, Cartesian motion, action distributions, frequency spectra, and episode-duration statistics.
print("\n" + "="*78 + "\n6. VIDEO: SEEK-BASED AV1 DECODE (no full download)\n" + "="*78)
def video_window(ep):
row = eps.loc[eps["episode_index"] == ep]
if len(row) == 0: return None
row = row.iloc[0]
ci = int(row.get(f"videos/{VIDEO_KEY}/chunk_index", row.get("data/chunk_index", 0)))
fi = int(row.get(f"videos/{VIDEO_KEY}/file_index", row.get("data/file_index", 0)))
f0 = float(row.get(f"videos/{VIDEO_KEY}/from_timestamp", 0.0))
f1 = float(row.get(f"videos/{VIDEO_KEY}/to_timestamp",
f0 + int(row.get("length", 100))/FPS))
return f"{ROOT}/videos/{VIDEO_KEY}/chunk-{ci:03d}/file-{fi:03d}.mp4", f0, f1
def decode_pyav(url, t0, t1, max_frames, stride, size):
import av
c = av.open(url, options={"rw_timeout": "30000000"})
s = c.streams.video[0]; s.thread_type = "AUTO"
if t0 > 0: c.seek(int(t0 / s.time_base), stream=s)
out, k = [], 0
for fr in c.decode(s):
ts = float(fr.pts * s.time_base)
if ts < t0 - 1e-3: continue
if ts > t1 + 1e-3 or len(out) >= max_frames: break
if k % stride == 0:
out.append(fr.reformat(width=size, height=size, format="rgb24").to_ndarray())
k += 1
c.close()
return np.stack(out) if out else None
def decode_ffmpeg(url, t0, t1, max_frames, stride, size):
cmd = ["ffmpeg", "-v", "error", "-ss", f"{t0:.3f}", "-i", url,
"-t", f"{max(t1-t0, 0.5):.3f}",
"-vf", f"select=not(mod(n\\,{stride})),scale={size}:{size}",
"-vsync", "0", "-frames:v", str(max_frames),
"-f", "rawvideo", "-pix_fmt", "rgb24", "-"]
buf = subprocess.run(cmd, capture_output=True).stdout
n = len(buf) // (size*size*3)
return np.frombuffer(buf[:n*size*size*3], np.uint8).reshape(n, size, size, 3) if n else None
def get_frames(ep, max_frames=64, stride=2, size=VIS_SIZE):
w = video_window(ep)
if w is None: return None
rel, t0, t1 = w; url = URL(rel)
for fn in (decode_pyav, decode_ffmpeg):
try:
f = fn(url, t0, t1, max_frames, stride, size)
if f is not None and len(f): return f
except Exception as e:
print(f" {fn.__name__} failed: {type(e).__name__}: {str(e)[:80]}")
return None
frames = get_frames(EP0, max_frames=12, stride=max(1, len(q)//12), size=160)
if frames is not None:
print(f"decoded {frames.shape} from {video_window(EP0)[0]}")
fig, axs = plt.subplots(2, 6, figsize=(15, 5.2))
for i, ax in enumerate(axs.ravel()):
ax.axis("off")
if i < len(frames):
ax.imshow(frames[i]); ax.set_title(f"t≈{i*(len(q)//12)/FPS:.1f}s", fontsize=8)
plt.suptitle(f"{VIDEO_KEY} · ep {EP0} · {traj['task'][:60]}"); plt.tight_layout(); plt.show()
else:
print("video decode unavailable (AV1 codec missing) — continuing state-only.")
USE_VISION = False
print("\n" + "="*78 + "\n7. NORMALIZATION\n" + "="*78)
try:
stats = json.load(open(hf_hub_download(REPO_ID, f"{ROOT}/meta/stats.json", repo_type="dataset")))
def cat_stat(keys, field):
return np.concatenate([np.atleast_1d(np.asarray(stats[k][field], dtype=np.float32).ravel())
for k in keys])
S_MEAN, S_STD = cat_stat(STATE_USE, "mean"), cat_stat(STATE_USE, "std")
A_MEAN, A_STD = cat_stat(ACTION_USE, "mean"), cat_stat(ACTION_USE, "std")
print("using dataset-level stats from meta/stats.json")
except Exception as e:
print("stats.json unusable, will compute empirically:", type(e).__name__)
S_MEAN = S_STD = A_MEAN = A_STD = None
We build a seek-based video pipeline that retrieves only the required temporal window from an episode, rather than downloading an entire video shard. We support both PyAV and FFmpeg decoding paths to handle AV1 video efficiently and resize selected frames for lightweight processing. We also load dataset-level normalization statistics from stats.json, with an empirical fallback when those statistics are unavailable.
print("\n" + "="*78 + "\n8. BUILDING TRAINING SET\n" + "="*78)
ep_ids = [e for e in list(EP_BOUNDS) if EP_BOUNDS[e][1]-EP_BOUNDS[e][0] > HORIZON+OBS_HISTORY+4][:N_EPISODES]
EPISODES = {}
for i, e in enumerate(ep_ids):
EPISODES[e] = load_episode(e)
if (i+1) % 8 == 0: print(f" loaded {i+1}/{len(ep_ids)} episodes")
print(f"loaded {len(EPISODES)} episodes, {sum(len(v['state']) for v in EPISODES.values()):,} frames")
VIS_CACHE = {}
if USE_VISION:
for e in ep_ids[:N_VIS_EPS]:
T = len(EPISODES[e]["state"])
f = get_frames(e, max_frames=min(T, 200), stride=1, size=VIS_SIZE)
if f is not None:
VIS_CACHE[e] = f
print(f" video ep {e}: {f.shape}")
USE_VISION = len(VIS_CACHE) >= 2
print(f"vision enabled: {USE_VISION} ({len(VIS_CACHE)} episodes cached)")
if S_MEAN is None:
allS = np.concatenate([v["state"] for v in EPISODES.values()])
allA = np.concatenate([v["action"] for v in EPISODES.values()])
S_MEAN, S_STD = allS.mean(0), allS.std(0) + 1e-6
A_MEAN, A_STD = allA.mean(0), allA.std(0) + 1e-6
S_STD = np.maximum(S_STD, 1e-4); A_STD = np.maximum(A_STD, 1e-4)
class DroidChunks(Dataset):
def __init__(self, episodes, ep_list, vision):
self.eps, self.vision, self.items = episodes, vision, []
for e in ep_list:
if vision and e not in VIS_CACHE: continue
T = len(episodes[e]["state"])
if vision: T = min(T, len(VIS_CACHE[e]))
for i in range(OBS_HISTORY-1, T-HORIZON): self.items.append((e, i))
def __len__(self): return len(self.items)
def __getitem__(self, k):
e, i = self.items[k]; d = self.eps[e]
s = (d["state"][i-OBS_HISTORY+1:i+1] - S_MEAN) / S_STD
a = (d["action"][i:i+HORIZON] - A_MEAN) / A_STD
out = [torch.from_numpy(s.ravel().astype(np.float32)),
torch.from_numpy(a.astype(np.float32))]
if self.vision:
img = VIS_CACHE[e][i].astype(np.float32) / 255.0
out.insert(1, torch.from_numpy(img.transpose(2, 0, 1)))
return tuple(out)
pool = list(VIS_CACHE) if USE_VISION else ep_ids
tr_eps, te_eps = pool[:-2], pool[-2:]
tr, te =