สร้าง Pipeline ฝึกหุ่นยนต์แบบ Streaming ด้วย NVIDIA Cosmos3-DROID

บทช่วยสอนนี้จะแนะนำการออกแบบ pipeline การเรียนรู้ของหุ่นยนต์แบบ streaming ตั้งแต่ต้นจนจบโดยใช้ ชุดข้อมูล NVIDIA Cosmos3-DROID โดยมีจุดเด่นคือไม่จำเป็นต้องดาวน์โหลด repository ขนาด 707 GB ลงเครื่อง
ขั้นตอนแรกเริ่มจากการตรวจสอบโครงสร้าง LeRobotDataset v3.0 เพื่อสร้าง metadata graph จาก info.json, task metadata, ตาราง episode และสถิติต่างๆ ของชุดข้อมูล จากนั้นใช้การเข้าถึง HTTP byte-range ผ่าน PyArrow เพื่อเลือกอ่านเฉพาะ row groups และ columns ของ Parquet ที่ต้องการ
เราจะทำการแปลงแต่ละ episode เป็นวิถีของสถานะและการกระทำ (state-action trajectories) พร้อมวิเคราะห์ข้อมูลเชิงลึก ทั้งการเคลื่อนไหวของข้อต่อ, เหตุการณ์ของ gripper, เส้นทาง Cartesian end-effector และสเปกตรัมความถี่ ก่อนจะถอดรหัสเฉพาะช่วงวิดีโอ AV1 ที่ต้องการผ่านการเข้าถึง PyAV/FFmpeg แบบ seek-based
ถัดมาคือการทำ normalization ผลการสังเกตและการกระทำโดยใช้สถิติของชุดข้อมูล เพื่อสร้างชุดข้อมูล PyTorch แบบ chunked สไตล์ ACT พร้อมเงื่อนไขทางภาพเสริมสำหรับฝึกฝนโมเดล (policy) behavior-cloning แบบหลายรูปแบบ (multimodal)
ในขั้นตอนสุดท้าย เราจะประเมินโมเดลที่เรียนรู้ผ่านการ rollout แบบ open-loop ด้วยเทคนิค temporal ensembling ของ action chunks โดยรายงานค่า MSE และ R^2 ต่อข้อต่อเทียบกับเกณฑ์มาตรฐาน (mean-action baseline) พร้อมแสดงภาพเปรียบเทียบระหว่างการกระทำที่ทำนายกับค่าจริง (ground-truth) และบันทึก checkpoint ของโมเดลเพื่อใช้งานต่อ
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 รวมถึงพารามิเตอร์สำหรับวิดีโอและการฝึกอบรม เราจะตรวจสอบโครงสร้างของ repository เพื่อระบุตำแหน่งวิดีโอและ metadata shards โดยไม่ต้องโหลดข้อมูลทั้งหมด จากนั้นจึงโหลด metadata และรายละเอียดงานเพื่อวิเคราะห์โครงสร้างข้อมูลและฟีเจอร์ต่างๆ
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()เรานำระบบอ่าน Parquet แบบ byte-range มาใช้เพื่อเข้าถึงเฉพาะ row groups และ columns ที่จำเป็นผ่าน Hugging Face โดยตรง พร้อมกำหนดขอบเขตของ episode ภายใน shard และแสดงผลการวิเคราะห์ข้อมูลหุ่นยนต์ ทั้งวิถีการเคลื่อนที่แบบ Cartesian, การทำงานของ gripper และสถิติระยะเวลาของแต่ละ episode
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เราได้สร้างวิดีโอ pipeline แบบ seek-based เพื่อดึงเฉพาะช่วงที่ต้องการมาประมวลผลแทนการโหลดไฟล์วิดีโอทั้งหมด โดยรองรับทั้ง PyAV และ FFmpeg เพื่อจัดการกับ codec AV1 อย่างมีประสิทธิภาพ พร้อมทั้งโหลดค่าสถิติสำหรับ normalization จากไฟล์ stats.json เพื่อความแม่นยำในการฝึกโมเดล
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)เราโหลด episode และแคชข้อมูลภาพที่ซิงโครไนซ์กันไว้เพื่อลดภาระการประมวลผล จากนั้นสร้างชุดข้อมูล PyTorch แบบ ACT ที่รวมทั้งประวัติสถานะและภาพถ่ายเพื่อทำนายลำดับการกระทำในอนาคต โดยใช้โครงสร้างโมเดลที่ผสาน MLP state encoder เข้ากับ CNN vision encoder
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 optimizer ร่วมกับ OneCycle learning rate scheduler และการประมวลผลแบบ mixed-precision เพื่อประสิทธิภาพสูงสุด โดยเลือกใช้ Smooth L1 loss เพื่อช่วยให้ระบบ behavior cloning ทนทานต่อสัญญาณรบกวนจากการบังคับหุ่นยนต์ด้วยมือ (teleoperation)
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}""")สุดท้ายเราประเมินผลผ่านการ rollout แบบ open-loop โดยใช้ temporal ensembling เพื่อผสานการทำนายให้ราบรื่นยิ่งขึ้น พร้อมแสดงกราฟเปรียบเทียบกับข้อมูลจริงและบันทึกโมเดลพร้อม metadata เพื่อนำไปใช้ทดลองในอนาคต โดยสรุป เราได้สาธิตการเปลี่ยนชุดข้อมูลหุ่นยนต์ขนาดมหึมาให้เป็น pipeline การเรียนรู้ที่ประหยัดพื้นที่จัดเก็บและแบนด์วิดท์อย่างมาก ผ่านการค้นหา episode ด้วย metadata และการถอดรหัสวิดีโอแบบเฉพาะจุด เวิร์กโฟลว์นี้เป็นรากฐานที่มั่นคงสำหรับขยายผลไปสู่การฝึกหุ่นยนต์ที่ซับซ้อนยิ่งขึ้นในอนาคต
ความคิดเห็น (0)
เข้าสู่ระบบเพื่อร่วมแสดงความเห็น
สมัครสมาชิกมาเป็นคนแรกที่แสดงความเห็นกันเลยโบร
