The NVIDIA Cosmos3-DROID dataset sits at 707 GB, but this tutorial shows how to train a robot policy without downloading the full repository. The approach uses HTTP byte-range access to pull specific Parquet row groups and AV1 video windows only when needed.
Repository structure
The script connects to the Hugging Face repository via HfFileSystem. It lists the available files and counts them by prefix. The data is split into shards for success and failure cases, alongside metadata files.
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())
The environment setup installs PyArrow, Pandas, and Matplotlib. The code then loads the info.json file to extract the total episode count, frame count, and task list. It identifies state variables like joint positions and action variables like joint velocity. The script also loads the tasks parquet file to map indices to human-readable task names.
Streaming data access
Instead of reading the whole file, the code defines a function to open a Parquet file via the Hugging Face filesystem. It calculates the byte offsets for specific row groups. This allows the script to request only the data needed for a specific episode range.
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}")
The code identifies the start and end row positions for every episode in the shard. It creates a dictionary mapping



