Download scripts/smoke_train_bridge_data.py from DanTim05/DynaTTT: direct link, hf CLI and curl.
- Browser
- Download file 5.09 kB
-
https://huggingface.co/spaces/DanTim05/DynaTTT/resolve/main/scripts/smoke_train_bridge_data.py
- Command line
-
hf download hf://spaces/DanTim05/DynaTTT/scripts/smoke_train_bridge_data.py
-
curl -L -o smoke_train_bridge_data.py https://huggingface.co/spaces/DanTim05/DynaTTT/resolve/main/scripts/smoke_train_bridge_data.py
5.09 kB
| #!/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()) | |