MarkTechPost수정일

NVIDIA Cosmos3-DROID를 활용한 스트리밍 로봇 학습 파이프라인 구축

로컬 다운로드 없이 NVIDIA Cosmos3-DROID 데이터셋을 활용하여 바이트 레인지 Parquet 읽기, 행동 클로닝(behavior cloning), 시간적 앙상블링(temporal ensembling)을 통해 엔드투엔드 스트리밍 로봇 학습 파이프라인을 구축하는 방법을 알아보세요. 이 글인 NVIDIA Cosmos3-DROID를 활용한 스트리밍 로봇 학습 파이프라인 구축은…

이 튜토리얼에서는 다음을 중심으로 엔드투엔드 스트리밍 로봇 학습 파이프라인을 설계합니다: NVIDIA Cosmos3-DROID 데이터셋 용량 707 GB의 저장소를 로컬에 다운로드하지 않고 진행합니다. 먼저 LeRobotDataset v3.0 구조를 조사하고 info.json, 태스크 메타데이터, 에피소드 테이블, 데이터셋 통계로부터 메타데이터 그래프를 구성한 뒤, HTTP 바이트 레인지 접근과 PyArrow를 사용하여 Parquet 행 그룹과 열을 선택적으로 읽습니다. 개별 에피소드를 상태-행동 궤적으로 변환하고, 필요한 AV1 비디오 윈도우만 PyAV/FFmpeg의 시크 기반 접근을 통해 디코딩하기 전에 관절 움직임, 그리퍼 이벤트, 데카르트 엔드이펙터 경로, 행동 주파수 스펙트럼을 분석합니다. 그런 다음 데이터셋 통계를 사용하여 관측값과 행동을 정규화하고, 선택적 시각 조건화(visual conditioning)가 가능한 ACT 스타일의 청크화된 PyTorch 데이터셋을 구성하며, 멀티모달 행동 클로닝 정책을 학습시킵니다. 마지막으로 시간적으로 앙상블된 행동 청크를 사용한 오픈루프 롤아웃을 통해 학습된 정책을 평가하고, 평균 행동 베이스라인 대비 관절별 MSE와 R^2를 보고하며, 예측 행동과 정답(ground-truth) 행동을 시각화하고, 이후 활용을 위해 완전한 정책 체크포인트를 저장합니다.

코드 복사
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())

Colab 환경을 초기화하고, 필요한 라이브러리를 설치하며, Cosmos3-DROID 데이터셋, 에피소드, 비디오 및 학습 파라미터를 구성합니다. 저장소 구조를 조사하여 전체 데이터셋을 다운로드하지 않고도 사용 가능한 데이터, 비디오, 메타데이터 샤드를 식별합니다. 그런 다음 핵심 메타데이터와 태스크 설명을 로드하여 데이터셋 스키마, 사용 가능한 상태/행동 특징, 에피소드 구성을 파악합니다.

코드 복사
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()

Hugging Face 파일시스템을 통해 필요한 행 그룹과 열에만 직접 접근하는 바이트 레인지 Parquet 리더를 구현합니다. 데이터 샤드 내에서 에피소드 경계를 식별하고 선택된 상태 및 행동 필드를 NumPy 궤적으로 변환합니다. 이어서 관절 위치, 그리퍼 활동, 데카르트 움직임, 행동 분포, 주파수 스펙트럼, 에피소드 길이 통계를 시각화합니다.

코드 복사
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

에피소드 전체 비디오 샤드를 다운로드하는 대신 필요한 시간적 윈도우만 검색하는 시크 기반 비디오 파이프라인을 구축합니다. PyAV와 FFmpeg 디코딩 경로를 모두 지원하여 AV1 비디오를 효율적으로 처리하고, 경량화된 처리를 위해 선택된 프레임의 크기를 조정합니다. 또한 stats.json에서 데이터셋 수준의 정규화 통계를 로드하며, 해당 통계를 사용할 수 없을 경우를 대비한 경험적 폴백(fallback)도 마련합니다.

코드 복사
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   = DroidChunks(EPISODES, tr_eps, USE_VISION), DroidChunks(EPISODES, te_eps, USE_VISION)
tl = DataLoader(tr, batch_size=BATCH, shuffle=True,  num_workers=2, drop_last=True)
vl = DataLoader(te, batch_size=BATCH, shuffle=False, num_workers=2)
print(f"train 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, h=512):
       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, img=None):
       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 스타일의 PyTorch 데이터셋을 구성합니다. 그런 다음 MLP 상태 인코더와 선택적 CNN 비전 인코더를 결합하여 미래 행동 시퀀스를 예측하는 청크화된 정책 아키텍처를 정의합니다.

코드 복사
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(), lr=3e-4, weight_decay=1e-4)
sched = torch.optim.lr_scheduler.OneCycleLR(opt, 3e-4, total_steps=EPOCHS*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_blocking=True) 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, beta=0.1)
       if train:
           opt.zero_grad(set_to_none=True); 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(f"epoch {ep+1:>2}/{EPOCHS}  train={trl:.5f}  val={vll:.5f}  ({time.time()-t0:.1f}s)")

청크화된 정책을 초기화하고 AdamW, OneCycle 학습률 스케줄링, 혼합 정밀도(mixed-precision) 실행, 그레이디언트 스케일링, 그레이디언트 클리핑으로 최적화합니다. Smooth L1 손실을 사용하여 행동 클로닝이 노이즈가 있거나 변동이 큰 텔레오퍼레이션 행동에 더 강인하도록 만듭니다. 구성된 에포크 수만큼 정책을 학습시키면서 학습 손실과 검증 손실을 모두 추적하여 학습 동향을 모니터링합니다.

코드 복사
print("\n" + "="*78 + "\n10. OPEN-LOOP ROLLOUT (temporal ensembling)\n" + "="*78)
@torch.no_grad()
def rollout(ep, m=0.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_HISTORY+1:i+1]-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 i+k < T: acc[i+k, k] = p[k]; cnt[i+k, 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 = [f"jvel_{i+1}" for i in range(7)] + ["gripper"]
print(f"episode {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), sharex=True)
for i, ax in enumerate(axs.ravel()):
   if i >= A_DIM: ax.axis("off"); continue
   ax.plot(gt[:, i], "k", lw=1.3, label="ground truth")
   ax.plot(np.where(valid, pred[:, i], np.nan), "r", lw=1.1, alpha=.85, label="policy")
   ax.set_title(names[i], fontsize=9)
   if i == 0: ax.legend(fontsize=7)
axs.ravel()[-1].axis("off")
inset = fig.add_axes([0.71, 0.08, 0.24, 0.2])
inset.plot(hist["train"], label="train"); inset.plot(hist["val"], label="val")
inset.set_yscale("log"); inset.set_title("loss", fontsize=8); inset.legend(fontsize=6)
plt.suptitle(f"Open-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_keys=STATE_USE, action_keys=ACTION_USE, horizon=HORIZON,
               obs_history=OBS_HISTORY, vision=USE_VISION, root=ROOT)},
          "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
 · ROOT="failure"                -> 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}""")

학습된 정책을 오픈루프 롤아웃을 통해 평가하고, 지수 가중 시간적 앙상블링을 사용하여 겹치는 행동 예측을 결합합니다. 관절별 MSE, 베이스라인 오차, R^2를 계산하고, 예측 행동을 정답 궤적과 함께 플롯하며, 학습/검증 손실 곡선도 함께 시각화합니다. 마지막으로 학습된 모델을 정규화 통계 및 구성 메타데이터와 함께 저장하여 이후 실험에서 정책을 재사용할 수 있도록 합니다.

결론적으로, 우리는 방대한 실제 로봇 데이터셋을 저장 및 데이터 전송 요구 사항을 극도로 낮게 유지하면서 학습 파이프라인으로 전환하는 방법을 보여주었습니다. 우리는 메타데이터 기반 에피소드 탐색, 컬럼 및 로우 그룹 수준의 Parquet 프로젝션, 그리고 탐색(seek) 기반 비디오 디코딩을 사용하여 전체 데이터셋을 실체화하는 대신 분석과 학습에 필요한 정보에만 접근했습니다. 우리는 고유감각 상태 이력과 선택적인 시각 관측을 결합하여 청킹된 행동 복제(chunked behavior-cloning) 정책을 학습시키고, 시간적 앙상블(temporal ensembling)을 사용하여 오픈루프 평가 중 더 매끄러운 행동 예측을 얻었습니다. 그 결과로 얻어진 워크플로는 컴팩트하면서도 확장 가능한 기반을 제공하며, 이를 통해 더 정교한 로보틱스 및 비전-언어-행동 실험을 위해 추가 샤드, 실패 시연, 카메라 뷰, 언어 지시문 또는 대체 행동 표현으로 확장할 수 있습니다.


다음을 확인해 보세요 전체 코드는 여기. 이 프로젝트의 연구자에게 모든 공이 돌아갑니다. 또한 언제든지 저희를 팔로우해 주세요. 트위터 그리고 저희의 다음 항목에 참여하는 것을 잊지 마세요 150k+ ML 서브레딧 및 구독하기 우리 뉴스레터잠깐! 텔레그램 쓰세요? 이제 텔레그램에서도 저희와 함께하실 수 있습니다.

GitHub 저장소, Hugging Face 페이지, 제품 출시, 웨비나 등을 홍보하기 위해 저희와 파트너가 되어야 하시나요? 저희와 소통하세요

해당 게시물 NVIDIA Cosmos3-DROID를 활용한 스트리밍 로봇 학습 파이프라인 구축 최초 게재: MarkTechPost.

원문 출처

MarkTechPost

내용 안내

원문 발행 및 권리는 출처에 있습니다.

기계 번역 · 원문을 참고하세요