#!/usr/bin/env python3 """Small real-data training smoke for bridge_orig_lerobot_smoke. This deliberately tests decoding, batching, CUDA forward/backward, optimizer updates, and checkpoint writing without claiming to train the full Cosmos/Zeva model. The Zeva framework's Bridge adapter currently expects older LeRobot metadata, while this smoke set is v3. """ from __future__ import annotations import argparse import json from pathlib import Path import numpy as np import pyarrow.parquet as pq import torch from torch import nn import av class TinyActionNet(nn.Module): def __init__(self) -> None: super().__init__() self.encoder = nn.Sequential( nn.Conv2d(3, 16, 5, stride=2, padding=2), nn.GELU(), nn.Conv2d(16, 32, 5, stride=2, padding=2), nn.GELU(), nn.Conv2d(32, 64, 5, stride=2, padding=2), nn.GELU(), nn.AdaptiveAvgPool2d(1), ) self.head = nn.Sequential(nn.Linear(64 + 8, 64), nn.GELU(), nn.Linear(64, 7)) def forward(self, image: torch.Tensor, state: torch.Tensor) -> torch.Tensor: feat = self.encoder(image).flatten(1) return self.head(torch.cat([feat, state], dim=1)) def load_subset(root: Path, episodes: int, size: int) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: rows_i, images, states, actions = [], [], [], [] for episode in range(episodes): table = pq.read_table(root / "data" / "chunk-000" / f"episode_{episode:06d}.parquet").to_pylist() video = root / "videos" / "chunk-000" / "observation.images.image_0" / f"episode_{episode:06d}.mp4" container = av.open(str(video)) decoded = [frame.to_ndarray(format="rgb24") for frame in container.decode(video=0)] container.close() if len(decoded) < len(table): raise RuntimeError(f"decoded {len(decoded)} frames but parquet has {len(table)}: {video}") for row, frame in zip(table, decoded): # Avoid an OpenCV dependency; torch interpolation handles the resize. frame = torch.from_numpy(frame).permute(2, 0, 1).float().unsqueeze(0) frame = torch.nn.functional.interpolate(frame, size=(size, size), mode="bilinear", align_corners=False) frame = frame.squeeze(0).byte().permute(1, 2, 0).numpy() images.append(frame) states.append(row["observation.state"]) actions.append(row["action"]) if not images: raise RuntimeError("smoke subset is empty") return ( torch.from_numpy(np.asarray(images)).permute(0, 3, 1, 2).float().div_(255.0), torch.tensor(np.asarray(states), dtype=torch.float32), torch.tensor(np.asarray(actions), dtype=torch.float32), ) def main() -> int: p = argparse.ArgumentParser() p.add_argument("--root", type=Path, required=True) p.add_argument("--steps", type=int, default=2000) p.add_argument("--episodes", type=int, default=3) p.add_argument("--output", type=Path, required=True) p.add_argument("--device", default="cuda") args = p.parse_args() if args.steps < 1: raise ValueError("steps must be positive") device = torch.device(args.device if args.device != "cuda" or torch.cuda.is_available() else "cpu") images, states, actions = load_subset(args.root, args.episodes, 64) images, states, actions = images.to(device), states.to(device), actions.to(device) model = TinyActionNet().to(device) opt = torch.optim.AdamW(model.parameters(), lr=3e-4) loss_fn = nn.MSELoss() losses: list[float] = [] generator = torch.Generator(device=device).manual_seed(42) model.train() for step in range(args.steps): idx = torch.randint(images.shape[0], (min(8, images.shape[0]),), generator=generator, device=device) loss = loss_fn(model(images.index_select(0, idx), states.index_select(0, idx)), actions.index_select(0, idx)) opt.zero_grad(set_to_none=True) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) opt.step() losses.append(float(loss.detach().cpu())) if (step + 1) % max(1, args.steps // 10) == 0: print(f"step={step + 1} loss={losses[-1]:.8f}", flush=True) args.output.parent.mkdir(parents=True, exist_ok=True) torch.save({"model": model.state_dict(), "steps": args.steps}, args.output.with_suffix(".pt")) report = { "status": "PASS_DATA_GPU_TRAINING_SMOKE", "device": str(device), "episodes": args.episodes, "frames": int(images.shape[0]), "image_shape": list(images.shape[1:]), "state_shape": list(states.shape[1:]), "action_shape": list(actions.shape[1:]), "steps": args.steps, "loss_first": losses[0], "loss_last": losses[-1], "loss_min": min(losses), "checkpoint": str(args.output.with_suffix(".pt").resolve()), } args.output.write_text(json.dumps(report, indent=2), encoding="utf-8") print(json.dumps(report, indent=2)) return 0 if __name__ == "__main__": raise SystemExit(main())