DynaTTT / scripts /smoke_train_bridge_data.py
DanTim05's picture
Upload Zeva source and project assets
68eed61 verified
Raw History Blame Contribute Delete
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())