
基于NVIDIA Cosmos3-DROID的端到端教程存储成本直降99%小伙伴们有没有遇到过想训机器人策略一看数据集707GB硬盘直接爆红、下载三天三夜的情况今天这篇教程直接给你解决这个痛点——围绕NVIDIA Cosmos3-DROID数据集搭端到端流式机器人学习管线本地连完整数据集都不用下。整个管线核心思路就是「按需取用」只拉分析和训练需要的数据绝不把全量材料塞到本地。第一步先做环境初始化与元数据探查import subprocess, sys, os, json, math, time, warnings, random, tempfile warnings.filterwarnings(ignore) subprocess.run([sys.executable, -m, pip, install, -q, huggingface_hub0.34.0, pyarrow15.0, av12.0, pandas, matplotlib, tqdm], checkFalse) 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: fdatasets/{REPO_ID}/{rel} URL lambda rel: hf_hub_url(REPO_ID, rel, repo_typedataset) print(\n *78 \n1. REPO INTROSPECTION\n *78) all_files api.list_repo_files(REPO_ID, repo_typedataset) print(ftotal 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_typedataset))) print(fepisodes{info.get(total_episodes):,} frames{info.get(total_frames):,} ftasks{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_typedataset)) 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_typedataset)) for f in ep_files[:4]], ignore_indexTrue) 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())我们在Colab里装好依赖配置好数据集、剧集、视频、训练参数之后先不用急着拉数据直接检查远程仓库的结构定位可用的数据分片、视频分片、元数据分片再加载info.json、任务元数据、剧集表这些核心元数据搞清楚整个数据集的模式有哪些可用的状态、动作特征剧集是怎么按任务组织的相当于先摸清楚整个数据集的地图后续所有操作都基于这个元数据图谱来不会出现找半天找不到对应剧集的尴尬。 [[IMAGE_1]]第二步是选择性读取Parquet数据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, columnscolumns) 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_onlyFalse)) return flat.reshape(len(ca), -1).astype(np.float32) return np.asarray(ca.to_numpy(zero_copy_onlyFalse)).reshape(-1, 1).astype(np.float32) SHARD data_shards[0] pf open_pf(SHARD) md pf.metadata print(fshard : {SHARD}) print(frows : {md.num_rows:,} row_groups: {md.num_row_groups} fcompressed: {md.serialized_size/1e6:.1f} MB footer) print(fcolumns : {len(pf.schema_arrow.names)}) t0 time.time() ep_idx_all pf.read(columns[episode_index]).column(episode_index).to_numpy() print(fpulled episode_index column ({len(ep_idx_all):,} rows) in {time.time()-t0:.1f}s) uniq, first_pos np.unique(ep_idx_all, return_indexTrue) 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], axis1) out[action] np.concatenate([out[k] for k in ACTION_USE], axis1) return out EP0 int(uniq[0]); traj load_episode(EP0) print(f\nepisode {EP0}: T{len(traj[state])} state_dim{traj[state].shape[1]} faction_dim{traj[action].shape[1]}) print(ftask: {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], lw1.1, labelfj{j1}) ax.set_title(joint positions [rad]); ax.set_xlabel(s); ax.legend(fontsize6, ncol2) ax fig.add_subplot(2, 3, 2) ax.plot(t, grip, colorcrimson, lw1.4) opens np.where(np.abs(np.diff(grip)) 0.05)[0] for k in opens[:40]: ax.axvline(t[k], colork, alpha.15, lw.8) ax.set_title(fgripper (|Δ|0.05 events: {len(opens)})); ax.set_xlabel(s) ax fig.add_subplot(2, 3, 3, projection3d) ax.plot(cart[:, 0], cart[:, 1], cart[:, 2], lw1.2) ax.scatter(*cart[0, :3], cg, s45, labelstart); ax.scatter(*cart[-1, :3], cr, s45, labelend) ax.set_title(EE cartesian path [m]); ax.legend(fontsize7) ax fig.add_subplot(2, 3, 4) im ax.imshow(dq.T, aspectauto, cmapRdBu_r, vmin-np.abs(dq).max(), vmaxnp.abs(dq).max()) ax.set_title(action.joint_velocity (7 x T)); ax.set_ylabel(joint); plt.colorbar(im, axax) ax fig.add_subplot(2, 3, 5) freqs np.fft.rfftfreq(len(dq), d1/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 (Nyquist7.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, bins40, colorsteelblue) ax.set_title(fepisode duration [s] (n{len(lens)})); ax.set_xlabel(s) plt.suptitle(f{REPO_ID} · {ROOT} · ep {EP0} · {traj[task][:70]}, y1.0) plt.tight_layout(); plt.show()机器人的状态、动作这些表格数据都存在Parquet文件里我们不需要把整个Parquet文件拖到本地只需要通过HTTP字节范围请求用PyArrow精准读取需要的行组和列——相当于你去图书馆借书不用把整个图书馆的书都搬回家只需要拿你要的那几本。读取之后我们先识别数据分片里的剧集边界把选中的状态、动作字段转成NumPy轨迹顺手还能做一堆可视化关节运动曲线、夹爪触发事件、笛卡尔末端执行器的运动路径、动作分布、频率谱、剧集时长统计先对数据质量有个底。 [[IMAGE_2]]第三步是按需解码视频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(fvideos/{VIDEO_KEY}/chunk_index, row.get(data/chunk_index, 0))) fi int(row.get(fvideos/{VIDEO_KEY}/file_index, row.get(data/file_index, 0))) f0 float(row.get(fvideos/{VIDEO_KEY}/from_timestamp, 0.0)) f1 float(row.get(fvideos/{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), streams) 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(widthsize, heightsize, formatrgb24).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, fselectnot(mod(n\\,{stride})),scale{size}:{size}, -vsync, 0, -frames:v, str(max_frames), -f, rawvideo, -pix_fmt, rgb24, -] buf subprocess.run(cmd, capture_outputTrue).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_frames64, stride2, sizeVIS_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_frames12, stridemax(1, len(q)//12), size160) if frames is not None: print(fdecoded {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(ft≈{i*(len(q)//12)/FPS:.1f}s, fontsize8) 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_typedataset))) def cat_stat(keys, field): return np.concatenate([np.atleast_1d(np.asarray(stats[k][field], dtypenp.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视频是AV1格式的单个分片动辄几十GB我们根本不需要全下。基于seek的PyAV/FFmpeg管线可以直接定位到当前剧集需要的时间窗口只拉对应的小段视频流解码解码之后调整帧大小做轻量处理就够了。同时我们会加载stats.json里的数据集级归一化统计信息要是这个文件缺失就自动用经验值回退保证后续训练的数值稳定性。 [[IMAGE_3]]接下来是数据集构建与策略训练。先看数据集部分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] HORIZONOBS_HISTORY4][:N_EPISODES] EPISODES {} for i, e in enumerate(ep_ids): EPISODES[e] load_episode(e) if (i1) % 8 0: print(f loaded {i1}/{len(ep_ids)} episodes) print(floaded {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_framesmin(T, 200), stride1, sizeVIS_SIZE) if f is not None: VIS_CACHE[e] f print(f video ep {e}: {f.shape}) USE_VISION len(VIS_CACHE) 2 print(fvision 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_HISTORY1:i1] - S_MEAN) / S_STD a (d[action][i:iHORIZON] - 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 DroidChunks(EPISODES, tr_eps, USE_VISION), DroidChunks(EPISODES, te_eps, USE_VISION) tl DataLoader(tr, batch_sizeBATCH, shuffleTrue, num_workers2, drop_lastTrue) vl DataLoader(te, batch_sizeBATCH, shuffleFalse, num_workers2) print(ftrain windows{len(tr):,} ({len(tr_eps)} eps) val windows{len(te):,} ({len(te_eps)} eps)) S_DIM, A_DIM EPISODES[ep_ids[0]][state].shape[1], EPISODES[ep_ids[0]][action].shape[1] class ChunkPolicy(nn.Module): def __init__(self, s_dim, a_dim, horizon, vision, h512): super().__init__() self.vision, self.horizon, self.a_dim vision, horizon, a_dim feat h if vision: self.cnn nn.Sequential( nn.Conv2d(3, 32, 5, 2, 2), nn.GroupNorm(8, 32), nn.SiLU(), nn.Conv2d(32, 64, 3, 2, 1), nn.GroupNorm(8, 64), nn.SiLU(), nn.Conv2d(64,128, 3, 2, 1), nn.GroupNorm(8, 128), nn.SiLU(), nn.Conv2d(128,256,3, 2, 1), nn.GroupNorm(8, 256), nn.SiLU(), nn.AdaptiveAvgPool2d(1), nn.Flatten(), nn.Linear(256, 256)) feat 256 self.smlp nn.Sequential(nn.Linear(s_dim*OBS_HISTORY, h), nn.SiLU(), nn.Linear(h, h)) self.trunk nn.Sequential(nn.Linear(feat, h), nn.SiLU(), nn.LayerNorm(h), nn.Linear(h, h), nn.SiLU(), nn.LayerNorm(h)) self.head nn.Linear(h, horizon*a_dim) def forward(self, s, imgNone): z self.smlp(s) if self.vision: z torch.cat([z, self.cnn(img)], -1) return self.head(self.trunk(z)).view(-1, self.horizon, self.a_dim)我们加载配置好的剧集集合为了不让显存爆炸只会缓存少量同步的视觉观测剩下的按需读取。数据集是ACT风格的分块结构把历史本体感觉状态、可选的视觉观测和归一化后的未来动作块拼在一起喂给模型的时候直接出符合输入要求的样本。策略架构也很灵活用MLP编码状态特征有视觉条件的话再加个CNN视觉编码器最终输出未来一段时间的动作序列适合做机器人长程规划。 [[IMAGE_4]]然后是训练优化model ChunkPolicy(S_DIM, A_DIM, HORIZON, USE_VISION).to(DEV) print(f\nmodel params: {sum(p.numel() for p in model.parameters())/1e6:.2f} M (vision{USE_VISION})) opt torch.optim.AdamW(model.parameters(), lr3e-4, weight_decay1e-4) sched torch.optim.lr_scheduler.OneCycleLR(opt, 3e-4, total_stepsEPOCHS*max(len(tl), 1), pct_start.15) scaler torch.amp.GradScaler(DEV, enabled(DEV cuda)) hist {train: [], val: []} def run(loader, train): model.train(train); tot n 0 for batch in loader: batch [b.to(DEV, non_blockingTrue) for b in batch] s, img, a (batch[0], batch[1], batch[2]) if USE_VISION else (batch[0], None, batch[1]) with torch.set_grad_enabled(train), torch.amp.autocast(DEV, enabled(DEV cuda)): loss F.smooth_l1_loss(model(s, img), a, beta0.1) if train: opt.zero_grad(set_to_noneTrue); scaler.scale(loss).backward() scaler.unscale_(opt); nn.utils.clip_grad_norm_(model.parameters(), 1.0) scaler.step(opt); scaler.update(); sched.step() tot loss.item()*len(s); n len(s) return tot/max(n, 1) print(\n *78 \n9. TRAINING\n *78) for ep in range(EPOCHS): t0 time.time(); trl run(tl, True); vll run(vl, False) hist[train].append(trl); hist[val].append(vll) print(fepoch {ep1:2}/{EPOCHS} train{trl:.5f} val{vll:.5f} ({time.time()-t0:.1f}s))我们用AdamW做优化器搭配OneCycle学习率调度开混合精度执行、梯度缩放和梯度裁剪训练稳定性直接拉满。损失函数没用普通的MSE选了Smooth L1对遥操作数据里常见的噪声、动作抖动鲁棒性更强不会因为某几个异常动作就让梯度爆炸训出来的策略动作更平滑。训练过程中会同步跟踪训练和验证损失一眼就能看出模型是过拟合还是没训够。 [[IMAGE_5]]最后是评估与落地print(\n *78 \n10. OPEN-LOOP ROLLOUT (temporal ensembling)\n *78) torch.no_grad() def rollout(ep, m0.1): d EPISODES[ep]; T len(d[state]) if USE_VISION: T min(T, len(VIS_CACHE[ep])) acc np.zeros((T, HORIZON, A_DIM), np.float32); cnt np.zeros((T, HORIZON), np.float32) model.eval() for i in range(OBS_HISTORY-1, T-HORIZON): s torch.from_numpy(((d[state][i-OBS_HISTORY1:i1]-S_MEAN)/S_STD).ravel() .astype(np.float32))[None].to(DEV) img None if USE_VISION: img torch.from_numpy((VIS_CACHE[ep][i].astype(np.float32)/255.) .transpose(2, 0, 1))[None].to(DEV) p model(s, img)[0].float().cpu().numpy()*A_STD A_MEAN for k in range(HORIZON): if ik T: acc[ik, k] p[k]; cnt[ik, k] math.exp(-m*k) w cnt[..., None]; pred (acc*w).sum(1) / np.maximum(w.sum(1), 1e-8) valid cnt.sum(1) 0 return pred, d[action][:T], valid ep_eval te_eps[0] pred, gt, valid rollout(ep_eval) mse ((pred[valid]-gt[valid])**2).mean(0) base ((gt[valid].mean(0)-gt[valid])**2).mean(0) names [fjvel_{i1} for i in range(7)] [gripper] print(fepisode {ep_eval} · task: {EPISODES[ep_eval][task][:70]}) print(f{dim:10}{MSE:12}{mean-baseline:16}{R²:10}) for i, nm in enumerate(names[:A_DIM]): print(f{nm:10}{mse[i]:12.5f}{base[i]:16.5f}{1-mse[i]/max(base[i],1e-9):10.3f}) print(f{OVERALL:10}{mse.mean():12.5f}{base.mean():16.5f}{1-mse.mean()/base.mean():10.3f}) fig, axs plt.subplots(3, 3, figsize(15, 8), sharexTrue) for i, ax in enumerate(axs.ravel()): if i A_DIM: ax.axis(off); continue ax.plot(gt[:, i], k, lw1.3, labelground truth) ax.plot(np.where(valid, pred[:, i], np.nan), r, lw1.1, alpha.85, labelpolicy) ax.set_title(names[i], fontsize9) if i 0: ax.legend(fontsize7) axs.ravel()[-1].axis(off) inset fig.add_axes([0.71, 0.08, 0.24, 0.2]) inset.plot(hist[train], labeltrain); inset.plot(hist[val], labelval) inset.set_yscale(log); inset.set_title(loss, fontsize8); inset.legend(fontsize6) plt.suptitle(fOpen-loop chunked BC · {ROOT} ep {ep_eval} · vision{USE_VISION}) plt.tight_layout(); plt.show() torch.save({model: model.state_dict(), s_mean: S_MEAN, s_std: S_STD, a_mean: A_MEAN, a_std: A_STD, cfg: dict( state_keysSTATE_USE, action_keysACTION_USE, horizonHORIZON, obs_historyOBS_HISTORY, visionUSE_VISION, rootROOT)}, droid_chunk_policy.pt) print(\nsaved - droid_chunk_policy.pt) print(f {*78} DONE. Everything above streamed from a 707 GB repo; peak disk use ≈ a few hundred MB. Scale-up levers · N_EPISODES / more shards - data_shards[1:], rebuild EP_BOUNDS per shard · ROOTfailure - 14,268 negative episodes for success classifiers · VIDEO_KEY - exterior_image_1_left / exterior_image_2_left (3 synced views: multi-view or view-randomization) · language - 53,086 task strings; add a text encoder for VLA-style conditioning instead of the state-only trunk · targets - swap ACTION_USE to action.cartesian_velocity for end-effector control, or predict deltas · real training - pip install lerobot; LeRobotDataset(local/{ROOT}) once you have local disk (v3.0 native loader) {*78})训完之后我们做开环rollout测试用指数加权时间集成把重叠的动作预测平滑掉避免动作跳变。评估指标算得很细每个关节的MSE、和平均动作基线的误差、R²都能直观看到模型在每个自由度上的表现。最后我们会把训练好的策略、归一化统计信息、配置元数据一起打包存好后续做别的实验直接就能复用不用重新训。整个管线跑下来存储和数据传输的需求比传统全量下载的方案低了不止一个量级而且扩展性很强要加更多数据分片、失败演示样本、多相机视角、语言指令甚至换别的动作表示直接在现有框架上改就行够支撑视觉-语言-动作这类更复杂的机器人实验。现在GitHub上已经放出了完整的Colab Notebook环境都是现成的点开就能跑。你有没有遇到过大数据集训练的存储瓶颈欢迎在评论区聊聊你的解法~