#!/usr/bin/env python """Training loop for shard-streamed data, tuned for a single H100. Differences from scripts/train.py (all of them are speed or correctness, none change the objective): data ShardSource (sequential parquet + JPEG) instead of LeRobotDataset random access into h264. This is the change that matters. attention expert routed through SDPA (see patches/modules_expert.py.diff) vision cam0 goes through the Qwen tower ONCE per step, not twice (patches/modeling_tinyvla.py.diff) tokens task strings pre-tokenized once into a lookup, not per step optimizer fused AdamW, foreach off, set_to_none precision bf16 autocast + TF32 matmuls, grad_accum 1 at a large batch compile torch.compile on the expert (many small ops; biggest compile win) resume shape-aware load: tensors whose shape changed (state/action projections when the action dim changes) are re-initialised and REPORTED, instead of raising or silently loading garbage Usage: # 1 GPU python train_fast.py --config configs/physical_ai_ft.yaml # N GPU (DDP): batch_size в конфиге — ПЕР-GPU; глобальный батч = batch_size*N torchrun --standalone --nproc_per_node 8 train_fast.py --config configs/physical_ai_ft.yaml """ from __future__ import annotations import argparse import json import math import time from pathlib import Path import torch import yaml # --------------------------------------------------------------------- policy def make_policy(cfg: dict, state_dim: int, action_dim: int): from lerobot.configs import FeatureType, PolicyFeature from tinyvla.configuration_tinyvla import TinyVLAConfig from tinyvla.modeling_tinyvla import TinyVLAPolicy pcfg = TinyVLAConfig(**cfg.get("policy", {})) if action_dim > pcfg.max_action_dim or state_dim > pcfg.max_state_dim: raise SystemExit( f"dataset has state {state_dim}d / action {action_dim}d but the config caps them at " f"{pcfg.max_state_dim} / {pcfg.max_action_dim}. Raise max_state_dim/max_action_dim — " f"CanonicalSource._pad would TRUNCATE silently." ) s = pcfg.image_size pcfg.input_features = { "observation.images.cam0": PolicyFeature(type=FeatureType.VISUAL, shape=(3, s, s)), "observation.images.cam1": PolicyFeature(type=FeatureType.VISUAL, shape=(3, s, s)), "observation.state": PolicyFeature(type=FeatureType.STATE, shape=(pcfg.max_state_dim,)), } pcfg.output_features = {"action": PolicyFeature(type=FeatureType.ACTION, shape=(action_dim,))} pcfg.validate_features() return TinyVLAPolicy(pcfg), pcfg def load_compatible(model, path: Path): """Load a checkpoint, keeping only tensors whose shape still matches. Continuing C-scaled on a 58-DoF bimanual robot changes the action/state projections. torch's strict=False does NOT tolerate a shape change (it raises), so filter explicitly and say out loud what was dropped — a silently re-initialised action head is the difference between "fine-tuning" and "training a new head on a frozen trunk". """ from safetensors.torch import load_file sd = load_file(path / "model.safetensors") own = model.state_dict() keep, reshaped, expanded, unexpected = {}, [], [], [] for k, v in sd.items(): if k not in own: unexpected.append(k) elif own[k].shape != v.shape: # Rows grew, trailing dims intact (num_embodiments 16 -> 32): keep # the trained rows, new rows stay at init. A plain re-init here # would silently discard every trained embodiment embedding. if (own[k].ndim == v.ndim and own[k].shape[0] > v.shape[0] and own[k].shape[1:] == v.shape[1:]): merged = own[k].clone() merged[: v.shape[0]] = v keep[k] = merged expanded.append((k, tuple(v.shape), tuple(own[k].shape))) else: reshaped.append((k, tuple(v.shape), tuple(own[k].shape))) else: keep[k] = v missing = [k for k in own if k not in keep] model.load_state_dict(keep, strict=False) print(f"resume {path}: loaded {len(keep)}/{len(own)} tensors") for k, a, b in expanded: print(f" EXPANDED (rows grew, old rows kept) {k}: {a} -> {b}") for k, a, b in reshaped: print(f" RE-INIT (shape changed) {k}: {a} -> {b}") if unexpected: print(f" ignored {len(unexpected)} unexpected keys, e.g. {unexpected[:3]}") left = [k for k in missing if all(k != r[0] for r in reshaped)] if left: print(f" {len(left)} tensors kept at init, e.g. {left[:3]}") return len(keep) class Collate: """Picklable collate: tokenizes task strings with a per-process cache. A module-level class (not a closure) so DataLoader workers can be started with the 'spawn' context. spawn matters under DDP: forking workers from a process that already initialized CUDA/NCCL (with ffmpeg + rust-tokenizer threads alive) is exactly the fork-after-CUDA hazard that hung rank dataloaders in testing; spawned workers start clean. """ def __init__(self, model_name: str, max_length: int, morph_text_max_len: int = 32): self.model_name = model_name self.max_length = max_length self.morph_text_max_len = morph_text_max_len self._tokenizer = None self._cache: dict = {} self._mcache: dict = {} def _tok(self, tasks): if self._tokenizer is None: import os os.environ.setdefault("TOKENIZERS_PARALLELISM", "false") from transformers import AutoTokenizer self._tokenizer = AutoTokenizer.from_pretrained(self.model_name) new = [t for t in set(tasks) if t not in self._cache] if new: e = self._tokenizer(new, padding="max_length", truncation=True, max_length=self.max_length, return_tensors="pt") for i, t in enumerate(new): self._cache[t] = (e["input_ids"][i], e["attention_mask"][i].bool()) return (torch.stack([self._cache[t][0] for t in tasks]), torch.stack([self._cache[t][1] for t in tasks])) def _mtok(self, texts): if self._tokenizer is None: self._tok([""]) # init tokenizer new = [t for t in set(texts) if t not in self._mcache] if new: e = self._tokenizer(new, padding="max_length", truncation=True, max_length=self.morph_text_max_len, return_tensors="pt") for i, t in enumerate(new): self._mcache[t] = (e["input_ids"][i], e["attention_mask"][i].bool()) return (torch.stack([self._mcache[t][0] for t in texts]), torch.stack([self._mcache[t][1] for t in texts])) def __getstate__(self): return {"model_name": self.model_name, "max_length": self.max_length, "morph_text_max_len": self.morph_text_max_len} def __setstate__(self, st): self.__init__(st["model_name"], st["max_length"], st.get("morph_text_max_len", 32)) def __call__(self, items): out = {} for k in items[0]: if k.startswith("__"): continue # webdataset re-injects __key__/__url__ after every stage if k == "task": ids, mask = self._tok([it["task"] for it in items]) out["observation.language.tokens"] = ids out["observation.language.attention_mask"] = mask elif k == "morph_text": ids, mask = self._mtok([it["morph_text"] for it in items]) out["morph_text_ids"] = ids out["morph_text_mask"] = mask else: out[k] = torch.stack([it[k] for it in items]) return out # ----------------------------------------------------------------------- main def main(): ap = argparse.ArgumentParser() ap.add_argument("--config", type=Path, required=True) ap.add_argument("--profile-steps", type=int, default=0, help="run N steps, print throughput, exit (use to replace the estimate with a measurement)") args = ap.parse_args() cfg = yaml.safe_load(args.config.read_text()) torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = True torch.backends.cudnn.benchmark = True # ---- DDP: активируется сам под torchrun, иначе одиночный GPU ---- import os as _osenv world = int(_osenv.environ.get("WORLD_SIZE", "1")) rank = int(_osenv.environ.get("RANK", "0")) local_rank = int(_osenv.environ.get("LOCAL_RANK", "0")) ddp = world > 1 if ddp: torch.distributed.init_process_group("nccl") torch.cuda.set_device(local_rank) is_main = rank == 0 def log(*a, **k): if is_main: print(*a, **k) dev = torch.device("cuda", local_rank if ddp else 0) # V100 (sm70) has no bf16: fall back to fp16 + GradScaler so debug runs on # older cards work; H100/B200 keep bf16 with the scaler disabled (no-op). # including_emulation=False: plain is_bf16_supported() returns True on # Volta (software emulation) and silently makes training ~10x slower. _bf16 = torch.cuda.is_bf16_supported(including_emulation=False) amp_dtype = torch.bfloat16 if _bf16 else torch.float16 scaler = torch.amp.GradScaler("cuda", enabled=amp_dtype is torch.float16) src_cfg = cfg["source"] morph = None if (src_cfg.get("type") != "hub" and cfg.get("morphology_descriptors") and cfg["policy"].get("conditioning") == "morph"): from tinyvla.modules.embodiment import MORPH_FIELDS raw = yaml.safe_load(Path(cfg["morphology_descriptors"]).read_text()) sc = {"arm_dof": 0.1, "reach_m": 2, "gripper_width_m": 10, "num_cameras": 1 / 3, "control_hz": 1 / 30, "joint_lo_mean": 1 / 3.1416, "joint_hi_mean": 1 / 3.1416, "workspace_x": 2, "workspace_y": 2, "workspace_z": 2, "payload_kg": 0.2} d = raw[src_cfg["morph_key"]] morph = torch.tensor([d.get(f, 0) * sc.get(f, 1) for f in MORPH_FIELDS], dtype=torch.float32) chunk = cfg["policy"]["chunk_size"] if src_cfg.get("type") == "hub": # stream episodes straight from HF hub (LeRobot v2.x), no local data at all import os as _os from tinyvla.data.streaming_hub import HubEpisodeStream specs = yaml.safe_load(Path(src_cfg["specs"]).read_text())["datasets"] tok_path = _os.path.expanduser("~/.cache/huggingface/token") pol = cfg["policy"] n_support = pol.get("n_support", 0) if (pol.get("use_demo_conditioning") or pol.get("vlm_native")) else 0 morph_map, prompt_map = {}, {} if cfg.get("morphology_descriptors"): from tinyvla.modules.embodiment import MORPH_FIELDS raw = yaml.safe_load(Path(cfg["morphology_descriptors"]).read_text()) sc = {"arm_dof": 0.1, "reach_m": 2, "gripper_width_m": 10, "num_cameras": 1 / 3, "control_hz": 1 / 30, "joint_lo_mean": 1 / 3.1416, "joint_hi_mean": 1 / 3.1416, "workspace_x": 2, "workspace_y": 2, "workspace_z": 2, "payload_kg": 0.2} for krobot, d in raw.items(): morph_map[krobot] = torch.tensor([d.get(f, 0) * sc.get(f, 1) for f in MORPH_FIELDS], dtype=torch.float32) if cfg.get("robot_prompts"): prompt_map = yaml.safe_load(Path(cfg["robot_prompts"]).read_text()) source = HubEpisodeStream( specs, token=open(tok_path).read().strip() if _os.path.exists(tok_path) else None, chunk=chunk, image_size=pol.get("image_size", 256), max_state_dim=pol["max_state_dim"], max_action_dim=pol["max_action_dim"], shuffle_buffer=cfg.get("shuffle_buffer", 4096), seed=cfg.get("seed", 42), rank=rank, world_size=world, n_support=n_support, morph_descriptors=morph_map, robot_prompts=prompt_map, ) log(f"source: HUB STREAM, {len(specs)} datasets (infinite mixture), world={world}") elif src_cfg.get("type") == "wds": # local self-contained webdataset tars (pack_wds.py / *_to_wds.py output), # mixed per-sample by weight; instructions resolved from tasks.json from tinyvla.data.wds_mix import MixtureSource specs = yaml.safe_load(Path(src_cfg["specs"]).read_text())["datasets"] root = Path(src_cfg.get("root", ".")) for s in specs: if "dir" not in s: # hf_repo: fetched beforehand by scripts/fetch_wds.py s["dir"] = str(root / s["hf_repo"].split("/")[-1] / s.get("subdir", "")) source = MixtureSource( specs, batch_size=cfg["batch_size"], num_workers=cfg.get("num_workers", 12), image_size=cfg["policy"].get("image_size", 256), max_state_dim=cfg["policy"]["max_state_dim"], max_action_dim=cfg["policy"]["max_action_dim"], shuffle_buffer=cfg.get("shuffle_buffer", 4096), steps_per_epoch=cfg.get("steps_per_epoch", 1000), seed=cfg.get("seed", 42) + rank, # decorrelate ranks: resampled streams ) log(f"source: WDS MIX, {len(specs)} datasets (infinite mixture), world={world}") elif src_cfg.get("type") == "wds_pack": # мелкие self-contained датасеты формата pack_wds (файнтюн на утёнке и т.п.) import json as _json from tinyvla.data.wds_pack import WdsPackSource pol = cfg["policy"] morph = None if cfg.get("morphology_descriptors") and src_cfg.get("morph_key"): from tinyvla.modules.embodiment import MORPH_FIELDS raw = yaml.safe_load(Path(cfg["morphology_descriptors"]).read_text()) sc = {"arm_dof": 0.1, "reach_m": 2, "gripper_width_m": 10, "num_cameras": 1 / 3, "control_hz": 1 / 30, "joint_lo_mean": 1 / 3.1416, "joint_hi_mean": 1 / 3.1416, "workspace_x": 2, "workspace_y": 2, "workspace_z": 2, "payload_kg": 0.2} d = raw[src_cfg["morph_key"]] morph = torch.tensor([d.get(f, 0) * sc.get(f, 1) for f in MORPH_FIELDS], dtype=torch.float32) tn = {} if src_cfg.get("task_names"): tn = {int(k): v for k, v in _json.loads(Path(src_cfg["task_names"]).read_text()).items()} source = WdsPackSource( root=src_cfg["root"], split=src_cfg.get("split", "train"), image_size=pol.get("image_size", 256), max_state_dim=pol["max_state_dim"], max_action_dim=pol["max_action_dim"], embodiment_id=src_cfg.get("embodiment_id", 0), n_support=pol.get("n_support", 0) if (pol.get("use_demo_conditioning") or pol.get("vlm_native")) else 0, task_names=tn, task_group_size=src_cfg.get("task_group_size", 0), morphology=morph, morph_text=src_cfg.get("robot_text", ""), seed=cfg.get("seed", 42), ) print(f"source: WDS_PACK {src_cfg['root']} [{src_cfg.get('split','train')}]: {len(source)} samples") else: from tinyvla.data.shards import ShardSource source = ShardSource( root=src_cfg["root"], embodiment_id=src_cfg.get("embodiment_id", 0), chunk=chunk, image_size=cfg["policy"].get("image_size", 256), max_state_dim=cfg["policy"]["max_state_dim"], max_action_dim=cfg["policy"]["max_action_dim"], morphology=morph, robot_prompt=src_cfg.get("robot_prompt"), shuffle_buffer=cfg.get("shuffle_buffer", 8192), seed=cfg.get("seed", 42), rank=rank, world_size=world, ) log(f"source: {source.num_frames:,} frames @ {source.fps} Hz, " f"state {source.state_dim}d action {source.action_dim}d, {len(source.shards)} shards") if src_cfg.get("type") in ("hub", "wds", "wds_pack"): # per-dataset dims vary; the source pads everything to the config caps policy, pcfg = make_policy(cfg, cfg["policy"]["max_state_dim"], cfg["policy"]["max_action_dim"]) else: policy, pcfg = make_policy(cfg, source.state_dim, source.action_dim) policy = policy.to(dev) if cfg.get("resume_from"): if is_main: load_compatible(policy, Path(cfg["resume_from"])) else: import contextlib, io as _io with contextlib.redirect_stdout(_io.StringIO()): load_compatible(policy, Path(cfg["resume_from"])) raw_policy = policy if ddp: policy = torch.nn.parallel.DistributedDataParallel( policy, device_ids=[local_rank], gradient_as_bucket_view=True ) # ---- pre-tokenize every distinct instruction once ----------------------- collate = Collate(pcfg.lm_model_name, pcfg.tokenizer_max_length, getattr(pcfg, "morph_text_max_len", 32)) nw = cfg.get("num_workers", 12) if ddp: # 16 воркеров/ранг x N рангов душат CPU и heartbeat torchrun-агента nw = cfg.get("num_workers_per_rank", max(4, nw // world)) log(f"DDP: {nw} dataloader workers per rank") hub = src_cfg.get("type") == "hub" map_style = not isinstance(source, torch.utils.data.IterableDataset) loader = torch.utils.data.DataLoader( source, batch_size=cfg["batch_size"], shuffle=map_style, num_workers=nw, pin_memory=True, persistent_workers=(nw > 0) and not hub, # hub: воркеры умирают каждую "эпоху" — сброс утечек prefetch_factor=cfg.get("prefetch_factor", 6) if nw > 0 else None, drop_last=True, collate_fn=collate, multiprocessing_context="spawn" if nw > 0 else None, ) # ---- optimizer: backbone at a lower lr, exactly as train.py does -------- backbone = [p for n, p in policy.named_parameters() if p.requires_grad and "semantic.vlm" in n] head = [p for n, p in policy.named_parameters() if p.requires_grad and "semantic.vlm" not in n] groups = [{"params": head, "lr": cfg["lr"]}] if backbone: groups.append({"params": backbone, "lr": cfg["lr"] * cfg.get("backbone_lr_mult", 0.1)}) opt = torch.optim.AdamW(groups, betas=(0.9, 0.95), weight_decay=1e-10, fused=True) log(f"trainable: head {sum(p.numel() for p in head)/1e6:.1f}M, " f"backbone {sum(p.numel() for p in backbone)/1e6:.1f}M at {cfg.get('backbone_lr_mult',0.1)}x lr") steps, warmup = cfg["steps"], cfg.get("warmup_steps", 1000) def lr_lambda(s): if s < warmup: return s / max(1, warmup) p = (s - warmup) / max(1, steps - warmup) return 0.025 + 0.975 * 0.5 * (1 + math.cos(math.pi * p)) sched = torch.optim.lr_scheduler.LambdaLR(opt, lr_lambda) orig_expert = raw_policy.expert # компилированный wrapper пишет state_dict с if cfg.get("compile", True): # префиксом expert._orig_mod.* — сейвим оригинал raw_policy.expert = torch.compile(raw_policy.expert, dynamic=False) log("torch.compile: expert") def save_clean(path): compiled = raw_policy.expert raw_policy.expert = orig_expert # параметры общие — это только про имена ключей raw_policy.save_pretrained(path) raw_policy.expert = compiled out_dir = Path(cfg["output_dir"]) if is_main: out_dir.mkdir(parents=True, exist_ok=True) if hasattr(source, "manifest"): (out_dir / "source_manifest.json").write_text(json.dumps(source.manifest, indent=2)) if cfg.get("wandb") and is_main: import wandb wandb.init(project=cfg["wandb"], config=cfg) grad_accum = cfg.get("grad_accum", 1) eff_batch = cfg["batch_size"] * grad_accum * world # глобальный батч log_freq = cfg.get("log_freq", 50) target = args.profile_steps or steps step, seen, t0 = 0, 0, time.time() data_iter = iter(loader) while step < target: opt.zero_grad(set_to_none=True) for _ in range(grad_accum): try: batch = next(data_iter) except StopIteration: if hasattr(source, "set_epoch"): source.set_epoch(getattr(source, "epoch", 0) + 1) data_iter = iter(loader) batch = next(data_iter) batch = {k: (v.to(dev, non_blocking=True) if torch.is_tensor(v) else v) for k, v in batch.items()} with torch.autocast("cuda", dtype=amp_dtype): loss, info = policy(batch) scaler.scale(loss / grad_accum).backward() seen += cfg["batch_size"] scaler.unscale_(opt) torch.nn.utils.clip_grad_norm_(raw_policy.parameters(), cfg.get("grad_clip", 10.0)) scaler.step(opt) scaler.update() sched.step() step += 1 if step % log_freq == 0: torch.cuda.synchronize() el = time.time() - t0 sps = seen * world / el # глобально по всем рангам # float(): yaml 1.1 парсит "2250.0e12" без знака экспоненты как строку gflops = float(cfg.get("gflops_per_sample", 360.0)) mfu = sps * gflops * 1e9 / float(cfg.get("peak_flops", 989e12)) * 100 eta = (steps - step) * eff_batch / sps / 3600 log(f"step {step}/{steps} loss {info['loss']:.4f} | {sps:.0f} samples/s " f"| {step/el:.2f} it/s | MFU~{mfu:.0f}% | eta {eta:.2f} h " f"| mem {torch.cuda.max_memory_allocated()/1e9:.1f} GB", flush=True) if cfg.get("wandb") and is_main: wandb.log({"loss": info["loss"], "lr": sched.get_last_lr()[0], "samples_per_s": sps}, step=step) seen, t0 = 0, time.time() if args.profile_steps == 0 and step % cfg.get("save_freq", 10000) == 0: if is_main: save_clean(out_dir / f"step_{step}") if ddp: torch.distributed.barrier() if args.profile_steps: log(f"\nprofile done: умножь samples/s выше на {steps * eff_batch:,} " f"глобальных семплов, чтобы получить реальный wall clock.") elif is_main: save_clean(out_dir / "final") if ddp: torch.distributed.barrier() torch.distributed.destroy_process_group() if __name__ == "__main__": main()