Instructions to use SleepMastger/fruit-picking-lingbot with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use SleepMastger/fruit-picking-lingbot with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("SleepMastger/fruit-picking-lingbot", dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
Download training_code/train.py from SleepMastger/fruit-picking-lingbot: direct link, hf CLI and curl.
- Browser
- Download file 39.5 kB
-
https://huggingface.co/SleepMastger/fruit-picking-lingbot/resolve/main/training_code/train.py
- Command line
-
hf download hf://SleepMastger/fruit-picking-lingbot/training_code/train.py
-
curl -L -o train.py https://huggingface.co/SleepMastger/fruit-picking-lingbot/resolve/main/training_code/train.py
39.5 kB
| # Copyright 2024-2025 The Robbyant Team Authors. All rights reserved. | |
| import argparse | |
| import os | |
| import sys | |
| from pathlib import Path | |
| import wandb | |
| import torch | |
| import torch.distributed as dist | |
| import torch.nn.functional as F | |
| from torch.utils.data import DataLoader, DistributedSampler, Subset | |
| from tqdm import tqdm | |
| from torch.distributed.checkpoint.state_dict import ( | |
| get_model_state_dict, | |
| get_optimizer_state_dict, | |
| set_optimizer_state_dict, | |
| StateDictOptions, | |
| ) | |
| from safetensors.torch import save_file, load_file | |
| import json | |
| sys.path.append(os.path.dirname(os.path.abspath(__file__))) | |
| from configs import VA_CONFIGS | |
| from distributed.fsdp import shard_model, apply_ac | |
| from distributed.util import ( | |
| _configure_model, | |
| init_distributed, | |
| dist_mean, | |
| dist_max, | |
| dist_sum | |
| ) | |
| from einops import rearrange | |
| from modules.utils import ( | |
| load_transformer, | |
| ) | |
| from utils import ( | |
| init_logger, | |
| logger, | |
| get_mesh_id, | |
| sample_timestep_id, | |
| data_seq_to_patch, | |
| warmup_constant_lambda, | |
| FlowMatchScheduler | |
| ) | |
| from dataset import MultiLatentLeRobotDataset | |
| import gc | |
| class Trainer: | |
| def __init__(self, config): | |
| if config.enable_wandb and config.rank == 0: | |
| wandb.login(host=os.environ['WANDB_BASE_URL'], key=os.environ['WANDB_API_KEY']) | |
| self.wandb = wandb | |
| # Was hardcoded to the literal string 'test_lln' -- every run collided under the | |
| # same name in the dashboard. run_name (set by run() from --config-name) plus a | |
| # timestamp keeps runs distinguishable and re-runs of the same config from colliding. | |
| import time as _time | |
| run_name = getattr(config, 'run_name', None) or config.get('__name__', 'run') | |
| run_name = f"{run_name}_{_time.strftime('%Y%m%d_%H%M%S')}" | |
| self.wandb.init( | |
| entity=os.environ["WANDB_TEAM_NAME"], | |
| project=os.getenv("WANDB_PROJECT", "va_robotwin"), | |
| # dir=log_dir, | |
| config=config, | |
| mode="online", | |
| name=run_name, | |
| ) | |
| logger.info(f"WandB logging enabled, run name: {run_name}") | |
| self.step = 0 | |
| self.config = config | |
| self.device = torch.device(f"cuda:{config.local_rank}") | |
| self.dtype = config.param_dtype | |
| self.patch_size = config.patch_size | |
| # Load models | |
| logger.info("Loading models...") | |
| # Load and shard transformer with FSDP | |
| logger.info("Loading transformer...") | |
| if hasattr(config, 'resume_from') and config.resume_from: | |
| transformer_path = os.path.join(config.resume_from, 'transformer') | |
| if config.rank == 0: | |
| logger.info(f"Resuming from checkpoint: {transformer_path}") | |
| else: | |
| transformer_path = os.path.join(config.wan22_pretrained_model_name_or_path, 'transformer') | |
| self.transformer = load_transformer( | |
| transformer_path, | |
| torch_dtype=torch.float32, | |
| torch_device='cpu', | |
| attn_mode="flex" | |
| ) | |
| logger.info("Setting up activation checkpointing ...") | |
| apply_ac(self.transformer) | |
| logger.info("Setting up FSDP...") | |
| shard_fn = shard_model | |
| self.transformer = _configure_model( | |
| model=self.transformer, | |
| shard_fn=shard_fn, | |
| param_dtype=self.dtype, | |
| device=self.device, | |
| eval_mode=False, | |
| ) | |
| self.transformer.train() | |
| self.transformer.requires_grad_(True) | |
| # Optimizer | |
| self.optimizer = torch.optim.AdamW( | |
| [p for p in self.transformer.parameters() if p.requires_grad], | |
| lr=config.learning_rate, | |
| betas=(config.beta1, config.beta2), | |
| eps=1e-8, | |
| weight_decay=config.weight_decay, | |
| fused=True, | |
| foreach=False, | |
| ) | |
| self.lr_scheduler = torch.optim.lr_scheduler.LambdaLR(self.optimizer, | |
| lr_lambda=lambda step: warmup_constant_lambda(step, warmup_steps=config.warmup_steps)) | |
| # Setup dataloaders | |
| logger.info("Setting up datasets...") | |
| full_dataset = MultiLatentLeRobotDataset(config=config) | |
| # Held-out validation split. Two mutually exclusive modes, both absent for every config | |
| # outside lift_new_*/place_cube_* -- so this stays a no-op for libero/robotwin/franka/demo. | |
| # | |
| # config.num_train_episodes -- ordered split: episodes [0, N) train, everything from N | |
| # on validates. place_cube_bowl_new needs this because its dataset is a chronological | |
| # slice of George's recordings (first 100 train, next 20 val) and the val tail comes | |
| # from later teleop sessions on purpose. A random split would put clips from the same | |
| # session on both sides and report an optimistically low val loss. | |
| # config.val_split -- seeded random fraction (lift_new_*). | |
| num_train_episodes = getattr(config, 'num_train_episodes', 0) | |
| val_split = getattr(config, 'val_split', 0.0) | |
| self.val_loader = None | |
| if num_train_episodes > 0: | |
| # episode_index is only unique within one LatentLeRobotDataset, so an ordered split | |
| # is ambiguous once several are concatenated; these configs load exactly one. | |
| if len(full_dataset._datasets) != 1: | |
| raise ValueError( | |
| f"num_train_episodes requires a single underlying dataset, got " | |
| f"{len(full_dataset._datasets)} -- use val_split instead" | |
| ) | |
| # new_metas[i] is the sample at dataset index i (LatentLeRobotDataset.__getitem__), | |
| # so split on its episode_index rather than on i: that stays correct if an episode | |
| # ever contributes more than one clip. | |
| ep_of = [m["episode_index"] for m in full_dataset._datasets[0].new_metas] | |
| train_indices = [i for i, e in enumerate(ep_of) if e < num_train_episodes] | |
| val_indices = [i for i, e in enumerate(ep_of) if e >= num_train_episodes] | |
| if not train_indices or not val_indices: | |
| raise ValueError( | |
| f"num_train_episodes={num_train_episodes} gives " | |
| f"{len(train_indices)} train / {len(val_indices)} val samples over " | |
| f"episodes {min(ep_of)}..{max(ep_of)}" | |
| ) | |
| train_dataset = Subset(full_dataset, train_indices) | |
| val_dataset = Subset(full_dataset, val_indices) | |
| if config.rank == 0: | |
| logger.info(f"train/val split: {len(train_dataset)} train / {len(val_dataset)} " | |
| f"val (ordered, episodes <{num_train_episodes} train)") | |
| elif val_split > 0: | |
| n = len(full_dataset) | |
| n_val = max(1, round(n * val_split)) | |
| g = torch.Generator().manual_seed(getattr(config, 'val_seed', 42)) | |
| perm = torch.randperm(n, generator=g).tolist() | |
| val_indices, train_indices = perm[:n_val], perm[n_val:] | |
| train_dataset = Subset(full_dataset, train_indices) | |
| val_dataset = Subset(full_dataset, val_indices) | |
| if config.rank == 0: | |
| logger.info(f"train/val split: {len(train_dataset)} train / {len(val_dataset)} " | |
| f"val (val_split={val_split}, seed={getattr(config, 'val_seed', 42)})") | |
| else: | |
| train_dataset = full_dataset | |
| val_dataset = None | |
| train_sampler = DistributedSampler( | |
| train_dataset, | |
| num_replicas=config.world_size, | |
| rank=config.rank, | |
| shuffle=True, | |
| seed=42 | |
| ) if config.world_size > 1 else None | |
| self.train_loader = DataLoader( | |
| train_dataset, | |
| batch_size=config.batch_size, | |
| shuffle=(train_sampler is None), | |
| num_workers=config.load_worker, | |
| sampler=train_sampler, | |
| ) | |
| if val_dataset is not None: | |
| # shuffle=False: deterministic order is preferable for validation. Every rank must | |
| # still see the same NUMBER of batches -- FSDP forward is a collective op, so a rank | |
| # that runs out of batches early would leave the others hanging on a collective that | |
| # never comes. DistributedSampler guarantees equal counts (via padding) same as it | |
| # does for train_sampler above. | |
| val_sampler = DistributedSampler( | |
| val_dataset, | |
| num_replicas=config.world_size, | |
| rank=config.rank, | |
| shuffle=False, | |
| seed=42, | |
| ) if config.world_size > 1 else None | |
| self.val_loader = DataLoader( | |
| val_dataset, | |
| batch_size=config.batch_size, | |
| shuffle=False, | |
| num_workers=min(config.load_worker, 4), | |
| sampler=val_sampler, | |
| ) | |
| self.val_interval = getattr(config, 'val_interval', config.save_interval) | |
| self.train_scheduler_latent = FlowMatchScheduler(shift=self.config.snr_shift, sigma_min=0.0, extra_one_step=True) | |
| self.train_scheduler_latent.set_timesteps(1000, training=True) | |
| self.train_scheduler_action = FlowMatchScheduler(shift=self.config.action_snr_shift, sigma_min=0.0, extra_one_step=True) | |
| self.train_scheduler_action.set_timesteps(1000, training=True) | |
| # Per-run subdirectory so runs NEVER overwrite each other's checkpoints. Previously | |
| # every run wrote to save_root/checkpoints/checkpoint_step_N, keyed only by step number, | |
| # so a second run (e.g. place_cube_bowl_base) silently clobbered the first | |
| # (lift_new_base) at each shared step (200/400/...). The tag must be identical across all | |
| # ranks AND unique per launch: SLURM_JOB_ID (this repo's Slurm path) and | |
| # TORCHELASTIC_RUN_ID (bare torchrun) are both shared across a launch's ranks; the | |
| # timestamp fallback is only reached single-process (world_size 1), where ranks can't | |
| # disagree. run_name is the config name, so dirs read e.g. place_cube_bowl_base_job542/. | |
| run_name = getattr(config, 'run_name', None) or config.get('__name__', 'run') | |
| _slurm_id = os.environ.get('SLURM_JOB_ID') | |
| if _slurm_id: | |
| run_subdir = f"{run_name}_job{_slurm_id}" | |
| elif os.environ.get('TORCHELASTIC_RUN_ID'): | |
| run_subdir = f"{run_name}_{os.environ['TORCHELASTIC_RUN_ID']}" | |
| else: | |
| import time as _time | |
| run_subdir = f"{run_name}_{_time.strftime('%Y%m%d_%H%M%S')}" | |
| self.save_dir = Path(config.save_root) / run_subdir / "checkpoints" | |
| self.save_dir.mkdir(parents=True, exist_ok=True) | |
| if config.rank == 0: | |
| logger.info(f"Checkpoints for this run -> {self.save_dir}") | |
| self.gradient_accumulation_steps = getattr(config, 'gradient_accumulation_steps', 1) | |
| self.train_loader_iter = None | |
| # if hasattr(config, 'resume_from') and config.resume_from: | |
| # self._load_training_state(config.resume_from) | |
| def _get_next_batch(self): | |
| """Get next batch from iterator, reset if epoch is finished.""" | |
| if self.train_loader_iter is None: | |
| self.train_loader_iter = iter(self.train_loader) | |
| try: | |
| batch = next(self.train_loader_iter) | |
| except StopIteration: | |
| # Reset sampler and iterator when epoch finishes | |
| if hasattr(self.train_loader.sampler, 'set_epoch'): | |
| self.train_loader.sampler.set_epoch(self.train_loader.sampler.epoch + 1) | |
| self.train_loader_iter = iter(self.train_loader) | |
| batch = next(self.train_loader_iter) | |
| return batch | |
| def _add_noise(self, latent, train_scheduler, action_mask=False, action_mode=False, noisy_cond_prob=0.): | |
| B, C, F, H, W = latent.shape | |
| timestep_ids = sample_timestep_id(batch_size=F, num_train_timesteps=train_scheduler.num_train_timesteps) | |
| noise = torch.zeros_like(latent).normal_() | |
| timesteps = train_scheduler.timesteps[timestep_ids].to(device=self.device) | |
| noisy_latents =train_scheduler.add_noise(latent, noise, timesteps, t_dim=2) | |
| targets =train_scheduler.training_target(latent, noise, timesteps) | |
| patch_f, patch_h, patch_w = self.patch_size | |
| if action_mode: | |
| patch_f = patch_h = patch_w = 1 | |
| latent_grid_id = get_mesh_id( | |
| latent.shape[-3] // patch_f, # F | |
| latent.shape[-2] // patch_h, # H | |
| latent.shape[-1] // patch_w, # W | |
| t=1 if action_mode else 0, # 1 for action mode (0 for latent), not used | |
| f_w=1, | |
| f_shift=0, | |
| action=action_mode | |
| ).to(self.device) # shape: [4, seq_len] | |
| latent_grid_id = latent_grid_id[None].repeat(B, 1, 1) | |
| if torch.rand(1).item() < noisy_cond_prob: | |
| cond_timestep_ids = sample_timestep_id( | |
| batch_size=F, | |
| min_timestep_bd=0.5, | |
| max_timestep_bd=1.0, | |
| num_train_timesteps=train_scheduler.num_train_timesteps, | |
| ) | |
| noise = torch.zeros_like(latent).normal_() | |
| cond_timesteps = train_scheduler.timesteps[cond_timestep_ids].to(device=self.device) | |
| latent = train_scheduler.add_noise(latent, noise, cond_timesteps, t_dim=2) | |
| else: | |
| cond_timesteps = torch.zeros_like(timesteps) | |
| if action_mask is not None: | |
| noisy_latents *= action_mask.float() | |
| targets *= action_mask.float() | |
| latent *= action_mask.float() | |
| return dict( | |
| timesteps=timesteps[None].repeat(B, 1), | |
| noisy_latents=noisy_latents, | |
| targets=targets, | |
| latent=latent, | |
| cond_timesteps=cond_timesteps[None].repeat(B, 1), | |
| grid_id=latent_grid_id, | |
| ) | |
| def _prepare_input_dict(self, batch_dict): | |
| """Prepare input dict following infer code pattern from wan_va_server.py.""" | |
| # Generate grid_id following infer code (no batch dimension yet) | |
| # For action mode: get_mesh_id(shape[-3], shape[-2], shape[-1], t=1, f_w=1, f_shift, action=True) | |
| latent_dict = self._add_noise( | |
| latent=batch_dict['latents'], | |
| train_scheduler=self.train_scheduler_latent, | |
| action_mask=None, | |
| action_mode=False, | |
| noisy_cond_prob=0.5) | |
| action_dict = self._add_noise( | |
| latent=batch_dict['actions'], | |
| train_scheduler=self.train_scheduler_action, | |
| action_mask=batch_dict['actions_mask'], | |
| action_mode=True, | |
| noisy_cond_prob=0.0) | |
| latent_dict['text_emb'] = batch_dict['text_emb'] | |
| action_dict['text_emb'] = batch_dict['text_emb'] | |
| action_dict['actions_mask'] = batch_dict['actions_mask'] | |
| input_dict = { | |
| 'latent_dict': latent_dict, | |
| 'action_dict': action_dict, | |
| 'chunk_size': torch.randint(1, 5, (1,)).item(), | |
| 'window_size': torch.randint(4, 65, (1,)).item(), | |
| } | |
| return input_dict | |
| def convert_input_format(self, input_dict): | |
| """Convert input dict to match transformer input format if needed.""" | |
| for key, value in input_dict.items(): | |
| input_dict[key] = value.to(self.device)#.to(self.dtype) | |
| return input_dict | |
| def compute_loss(self, | |
| input_dict, | |
| pred | |
| ): | |
| latent_pred, action_pred = pred | |
| action_pred = rearrange(action_pred, 'b (f n) c -> b c f n 1', f=input_dict['action_dict']['targets'].shape[-3]) | |
| latent_pred = data_seq_to_patch( | |
| self.patch_size, latent_pred, | |
| input_dict['latent_dict']['targets'].shape[-3], input_dict['latent_dict']['targets'].shape[-2], | |
| input_dict['latent_dict']['targets'].shape[-1], batch_size=latent_pred.shape[0]) | |
| Bn, Fn = input_dict['latent_dict']['timesteps'].shape | |
| latent_loss_weight = self.train_scheduler_latent.training_weight(input_dict['latent_dict']['timesteps'].flatten()).reshape(Bn, Fn) | |
| action_loss_weight = self.train_scheduler_action.training_weight(input_dict['action_dict']['timesteps'].flatten()).reshape(Bn, Fn) | |
| # Frame-wise video loss calculation | |
| latent_loss = F.mse_loss(latent_pred.float(), input_dict['latent_dict']['targets'].float().detach(), reduction='none') | |
| latent_loss = latent_loss * latent_loss_weight[:, None, :, None, None] | |
| # Permute to (B, F, H, W, C) and flatten to (B*F, H*W*C) | |
| latent_loss = latent_loss.permute(0, 2, 3, 4, 1) # (B, C, F, H, W) -> (B, F, H, W, C) | |
| latent_loss = latent_loss.flatten(0, 1).flatten(1) # (B, F, H, W, C) -> (B*F, H*W*C) | |
| # Sum per frame and compute mask per frame | |
| latent_loss_per_frame = latent_loss.sum(dim=1) # (B*F,) | |
| latent_mask_per_frame = torch.ones_like(latent_loss).sum(dim=1) # (B*F,) | |
| latent_loss = (latent_loss_per_frame / (latent_mask_per_frame + 1e-6)).mean() | |
| # Frame-wise action loss calculation | |
| action_loss = F.mse_loss(action_pred.float(), input_dict['action_dict']['targets'].float().detach(), reduction='none') | |
| action_loss = action_loss * action_loss_weight[:, None, :, None, None] | |
| action_loss = action_loss * input_dict['action_dict']['actions_mask'].float() | |
| # Permute to (B, F, H, W, C) and flatten to (B*F, H*W*C) | |
| action_loss = action_loss.permute(0, 2, 3, 4, 1) # (B, C, F, H, W) -> (B, F, H, W, C) | |
| action_mask = input_dict['action_dict']['actions_mask'].float().permute(0, 2, 3, 4, 1) # (B, C, F, H, W) -> (B, F, H, W, C) | |
| action_loss = action_loss.flatten(0, 1).flatten(1) # (B, F, H, W, C) -> (B*F, H*W*C) | |
| action_mask = action_mask.flatten(0, 1).flatten(1) # (B, F, H, W, C) -> (B*F, H*W*C) | |
| # Sum per frame and normalize by mask per frame | |
| action_loss_per_frame = action_loss.sum(dim=1) # (B*F,) | |
| action_mask_per_frame = action_mask.sum(dim=1) # (B*F,) | |
| action_loss = (action_loss_per_frame / (action_mask_per_frame + 1e-6)).mean() | |
| return latent_loss / self.gradient_accumulation_steps, action_loss / self.gradient_accumulation_steps | |
| def _action_l2_metrics(self, pred, input_dict): | |
| """De-normalized L2 error between a one-step flow-match reconstruction of the clean | |
| action and ground truth, restricted to the real (non-padded) action channels. | |
| One-step reconstruction: FlowMatchScheduler defines noisy = (1-sigma)*clean + sigma*noise | |
| and target = noise - clean (see wan_va/utils/scheduler.py), so a model that predicts the | |
| true target exactly satisfies clean = noisy - sigma*target -- same identity | |
| FlowMatchScheduler.step(..., to_final=True) uses, just applied here with per-frame sigmas | |
| (scheduler.step() itself only handles a single scalar timestep, not our per-frame ones, | |
| so this reimplements the same broadcast pattern _add_noise/scheduler.add_noise use). | |
| This is an approximation (real inference denoises over many steps, not one) -- cheap | |
| enough to log every training step, not a substitute for real rollout eval. | |
| Fully generic / safe for every OTHER config in this repo: returns per-real-channel sums | |
| only; gripper-accuracy and named/grouped breakdowns are opt-in via config fields | |
| (gripper_channel_index, action_channel_groups / action_channel_names) that only | |
| lift_new_configs/va_lift_new_cfg.py sets. Absent those fields, this still returns valid | |
| per-channel-index sums, just without the semantic labels. | |
| """ | |
| _, action_pred = pred | |
| action_dict = input_dict['action_dict'] | |
| F_dim = action_dict['targets'].shape[-3] | |
| action_pred = rearrange(action_pred, 'b (f n) c -> b c f n 1', f=F_dim) | |
| noisy = action_dict['noisy_latents'] | |
| timesteps = action_dict['timesteps'][0] # [F] -- identical across the batch dim (see _add_noise) | |
| timestep_id = torch.argmin( | |
| (self.train_scheduler_action.timesteps[:, None].to(timesteps.device) - timesteps[None]).abs(), | |
| dim=0, | |
| ) | |
| sigma = self.train_scheduler_action.sigmas.to(timesteps.device)[timestep_id] # [F] | |
| shape = [1] * noisy.ndim | |
| shape[2] = sigma.shape[0] # t_dim=2, matches _add_noise's convention | |
| sigma = sigma.view(shape) | |
| x0_pred = noisy - sigma * action_pred.float() | |
| x0_true = action_dict['latent'] | |
| mask = action_dict['actions_mask'].float() | |
| n_real = len(self.config.used_action_channel_ids) | |
| q01 = torch.tensor(self.config.norm_stat['q01'][:n_real], device=noisy.device, dtype=torch.float32) | |
| q99 = torch.tensor(self.config.norm_stat['q99'][:n_real], device=noisy.device, dtype=torch.float32) | |
| denorm_shape = [1, n_real] + [1] * (noisy.ndim - 2) | |
| def denorm(x): | |
| return (x[:, :n_real].float() + 1) / 2 * (q99 - q01 + 1e-6).view(denorm_shape) + q01.view(denorm_shape) | |
| pred_real = denorm(x0_pred) | |
| true_real = denorm(x0_true) | |
| m = mask[:, :n_real] | |
| sq_err = (pred_real - true_real) ** 2 * m | |
| sum_dims = (0, 2, 3, 4) | |
| out = { | |
| 'sq_err_sum': sq_err.sum(dim=sum_dims).detach(), # [n_real] | |
| 'mask_sum': (m.sum(dim=sum_dims) + 1e-6).detach(), # [n_real] | |
| } | |
| gripper_idx = getattr(self.config, 'gripper_channel_index', None) | |
| if gripper_idx is not None: | |
| gripper_pred = (pred_real[:, gripper_idx:gripper_idx + 1] > 0.5).float() | |
| gripper_true = (true_real[:, gripper_idx:gripper_idx + 1] > 0.5).float() | |
| gripper_mask = m[:, gripper_idx:gripper_idx + 1] | |
| out['gripper_match_sum'] = ((gripper_pred == gripper_true).float() * gripper_mask).sum().detach() | |
| out['gripper_mask_sum'] = (gripper_mask.sum() + 1e-6).detach() | |
| return out | |
| def _rmse_report(self, sq_err_sum, mask_sum, gripper_match_sum=None, gripper_mask_sum=None): | |
| """Turn accumulated (already cross-rank-summed) sq_err_sum/mask_sum into final scalar | |
| metrics. Per-group/per-name breakdowns are opt-in via config.action_channel_groups / | |
| config.action_channel_names -- absent those, only 'overall' (and gripper_accuracy, if | |
| gripper_match_sum is given) is reported.""" | |
| sq_err_sum = sq_err_sum.detach().cpu() | |
| mask_sum = mask_sum.detach().cpu() | |
| report = {'overall': (sq_err_sum.sum() / (mask_sum.sum() + 1e-6)).sqrt().item()} | |
| groups = getattr(self.config, 'action_channel_groups', None) | |
| names = getattr(self.config, 'action_channel_names', None) | |
| if groups: | |
| for name, idxs in groups.items(): | |
| idxs_t = torch.tensor(idxs) | |
| report[f'rmse_{name}'] = (sq_err_sum[idxs_t].sum() / (mask_sum[idxs_t].sum() + 1e-6)).sqrt().item() | |
| elif names: | |
| for i, name in enumerate(names): | |
| report[f'rmse_{name}'] = (sq_err_sum[i] / (mask_sum[i] + 1e-6)).sqrt().item() | |
| if gripper_match_sum is not None: | |
| report['gripper_accuracy'] = ( | |
| gripper_match_sum.detach().cpu() / (gripper_mask_sum.detach().cpu() + 1e-6) | |
| ).item() | |
| return report | |
| def validate(self): | |
| """Run the held-out validation split (config.val_split). All ranks must call this | |
| together -- FSDP forward is a collective op, so skipping it on some ranks (e.g. "only | |
| rank 0 validates") would hang the others. Only rank 0 needs the returned report for | |
| logging, but every rank has to process its shard of val_loader for that report to be | |
| correct (and to avoid a hang).""" | |
| if self.val_loader is None: | |
| return {} | |
| self.transformer.eval() | |
| n_real = len(self.config.used_action_channel_ids) | |
| sq_err_sum = torch.zeros(n_real, device=self.device) | |
| mask_sum = torch.zeros(n_real, device=self.device) | |
| gripper_match_sum = torch.zeros((), device=self.device) | |
| gripper_mask_sum = torch.zeros((), device=self.device) | |
| has_gripper = getattr(self.config, 'gripper_channel_index', None) is not None | |
| latent_losses, action_losses = [], [] | |
| for batch in self.val_loader: | |
| batch = self.convert_input_format(batch) | |
| input_dict = self._prepare_input_dict(batch) | |
| output = self.transformer(input_dict, train_mode=True) | |
| latent_loss, action_loss = self.compute_loss(input_dict, output) | |
| # compute_loss divides by gradient_accumulation_steps for training's backward | |
| # scaling; undo that here so validation loss is directly comparable to a fresh | |
| # per-batch loss, not scaled by an accumulation window that doesn't apply here. | |
| latent_losses.append(latent_loss.detach() * self.gradient_accumulation_steps) | |
| action_losses.append(action_loss.detach() * self.gradient_accumulation_steps) | |
| l2 = self._action_l2_metrics(output, input_dict) | |
| sq_err_sum += l2['sq_err_sum'] | |
| mask_sum += l2['mask_sum'] | |
| if has_gripper: | |
| gripper_match_sum += l2['gripper_match_sum'] | |
| gripper_mask_sum += l2['gripper_mask_sum'] | |
| self.transformer.train() | |
| n_batches = max(1, len(self.val_loader)) | |
| latent_loss_show = dist_mean(torch.stack(latent_losses).sum() / n_batches) | |
| action_loss_show = dist_mean(torch.stack(action_losses).sum() / n_batches) | |
| sq_err_sum = dist_sum(sq_err_sum) | |
| mask_sum = dist_sum(mask_sum) | |
| if has_gripper: | |
| gripper_match_sum = dist_sum(gripper_match_sum) | |
| gripper_mask_sum = dist_sum(gripper_mask_sum) | |
| report = self._rmse_report( | |
| sq_err_sum, mask_sum, | |
| gripper_match_sum if has_gripper else None, | |
| gripper_mask_sum if has_gripper else None, | |
| ) | |
| report['latent_loss'] = latent_loss_show.item() | |
| report['action_loss'] = action_loss_show.item() | |
| return report | |
| def _train_step(self, batch, batch_idx): | |
| """Train a single batch, returns losses for logging.""" | |
| batch = self.convert_input_format(batch) | |
| input_dict = self._prepare_input_dict(batch) | |
| should_sync = (batch_idx + 1) % self.gradient_accumulation_steps == 0 | |
| if not should_sync: | |
| self.transformer.set_requires_gradient_sync(False) | |
| else: | |
| self.transformer.set_requires_gradient_sync(True) | |
| output = self.transformer(input_dict, train_mode=True) | |
| latent_loss, action_loss = self.compute_loss(input_dict, output) | |
| loss = latent_loss + action_loss | |
| loss.backward() | |
| # Reuses the already-computed forward output -- no extra forward pass needed to also | |
| # report a physically-interpretable action error alongside the flow-matching loss. | |
| l2 = self._action_l2_metrics(output, input_dict) | |
| losses = { | |
| 'latent_loss': latent_loss.detach(), | |
| 'action_loss': action_loss.detach(), | |
| 'sq_err_sum': l2['sq_err_sum'], | |
| 'mask_sum': l2['mask_sum'], | |
| } | |
| if 'gripper_match_sum' in l2: | |
| losses['gripper_match_sum'] = l2['gripper_match_sum'] | |
| losses['gripper_mask_sum'] = l2['gripper_mask_sum'] | |
| # Only update weights after accumulating gradients | |
| if should_sync: | |
| total_norm = torch.nn.utils.clip_grad_norm_(self.transformer.parameters(), 2.0) | |
| self.optimizer.step() | |
| self.lr_scheduler.step() | |
| self.optimizer.zero_grad() | |
| losses['total_norm'] = total_norm | |
| losses['should_log'] = True | |
| else: | |
| losses['should_log'] = False | |
| return losses | |
| def save_checkpoint(self,): | |
| """Save model checkpoint in the same format as pretrained model.""" | |
| try: | |
| state_dict = get_model_state_dict( | |
| self.transformer, | |
| options=StateDictOptions(full_state_dict=True, cpu_offload=True), | |
| ) | |
| state_dict_bf16 = {k: v.to(torch.bfloat16) for k, v in state_dict.items()} | |
| # optim_state = get_optimizer_state_dict( | |
| # self.transformer, self.optimizer, | |
| # options=StateDictOptions(full_state_dict=True, cpu_offload=True), | |
| # ) | |
| # Only rank 0 saves the checkpoint | |
| if self.config.rank == 0: | |
| checkpoint_dir = self.save_dir / f"checkpoint_step_{self.step}" | |
| checkpoint_dir.mkdir(parents=True, exist_ok=True) | |
| # Save transformer in the same format as pretrained model | |
| transformer_dir = checkpoint_dir / "transformer" | |
| transformer_dir.mkdir(parents=True, exist_ok=True) | |
| logger.info(f"Saving transformer to {transformer_dir}") | |
| # Manually save in diffusers format (outside FSDP context to avoid deadlock) | |
| # Save model weights | |
| model_file = transformer_dir / "diffusion_pytorch_model.safetensors" | |
| save_file(state_dict_bf16, model_file) | |
| # Save config (copy from original transformer config and update _name_or_path) | |
| config_file = transformer_dir / "config.json" | |
| config_dict = dict(self.transformer.config) | |
| config_dict.pop('_name_or_path', None) | |
| with open(config_file, 'w') as f: | |
| json.dump(config_dict, f, indent=2) | |
| # # Save optimizer state and training metadata in PyTorch format | |
| # training_state_path = checkpoint_dir / "training_state.pt" | |
| # logger.info(f"Saving training state to {training_state_path}") | |
| # torch.save({ | |
| # 'step': self.step, | |
| # 'optimizer_state_dict': optim_state, | |
| # 'config': vars(self.config), | |
| # }, training_state_path) | |
| logger.info(f"Checkpoint saved successfully at step {self.step}") | |
| # Synchronize all processes after saving | |
| if dist.is_initialized(): | |
| dist.barrier() | |
| except Exception as e: | |
| if self.config.rank == 0: | |
| logger.error(f"Failed to save checkpoint: {e}") | |
| import traceback | |
| logger.error(traceback.format_exc()) | |
| # Ensure all processes stay synchronized even on error | |
| if dist.is_initialized(): | |
| dist.barrier() | |
| def _load_training_state(self, checkpoint_path): | |
| """Load training state (optimizer + step) after FSDP and optimizer creation.""" | |
| checkpoint_dir = Path(checkpoint_path) | |
| training_state_path = checkpoint_dir / "training_state.pt" | |
| if not training_state_path.exists(): | |
| if self.config.rank == 0: | |
| logger.warning(f"Training state not found: {training_state_path}, starting from step 0") | |
| return | |
| if self.config.rank == 0: | |
| logger.info(f"Loading training state from {training_state_path}") | |
| # All ranks load the training state directly | |
| training_state = torch.load(training_state_path, map_location='cpu', weights_only=False) | |
| # All ranks load optimizer state (required for FSDP) | |
| set_optimizer_state_dict( | |
| self.transformer, self.optimizer, | |
| optim_state_dict=training_state['optimizer_state_dict'], | |
| options=StateDictOptions(full_state_dict=True, strict=False) | |
| ) | |
| self.step = training_state.get('step', 0) | |
| if self.config.rank == 0: | |
| logger.info(f"Training state loaded, resuming from step {self.step}") | |
| # Synchronize all ranks | |
| if dist.is_initialized(): | |
| dist.barrier() | |
| def train(self): | |
| """Main training loop - train by steps instead of epochs.""" | |
| logger.info(f"Starting training for {self.config.num_steps} steps...") | |
| self.transformer.train() | |
| progress_bar = tqdm( | |
| total=self.config.num_steps, | |
| desc="Training", | |
| disable=(self.config.rank != 0), | |
| leave=True, | |
| dynamic_ncols=True, | |
| initial=self.step | |
| ) | |
| self.optimizer.zero_grad() | |
| accumulated_latent_losses = [] | |
| accumulated_action_losses = [] | |
| accumulated_sq_err_sum = None | |
| accumulated_mask_sum = None | |
| accumulated_gripper_match_sum = None | |
| accumulated_gripper_mask_sum = None | |
| step_in_accumulation = 0 | |
| while self.step < self.config.num_steps: | |
| # Get next batch (handles epoch reset automatically) | |
| batch = self._get_next_batch() | |
| losses = self._train_step(batch, step_in_accumulation) | |
| # Accumulate losses for logging | |
| accumulated_latent_losses.append(losses['latent_loss']) | |
| accumulated_action_losses.append(losses['action_loss']) | |
| accumulated_sq_err_sum = (losses['sq_err_sum'] if accumulated_sq_err_sum is None | |
| else accumulated_sq_err_sum + losses['sq_err_sum']) | |
| accumulated_mask_sum = (losses['mask_sum'] if accumulated_mask_sum is None | |
| else accumulated_mask_sum + losses['mask_sum']) | |
| if 'gripper_match_sum' in losses: | |
| accumulated_gripper_match_sum = (losses['gripper_match_sum'] if accumulated_gripper_match_sum is None | |
| else accumulated_gripper_match_sum + losses['gripper_match_sum']) | |
| accumulated_gripper_mask_sum = (losses['gripper_mask_sum'] if accumulated_gripper_mask_sum is None | |
| else accumulated_gripper_mask_sum + losses['gripper_mask_sum']) | |
| step_in_accumulation += 1 | |
| # Log and checkpoint when optimizer steps | |
| if losses['should_log']: | |
| lr = self.lr_scheduler.get_last_lr()[0] | |
| # Average accumulated losses | |
| latent_loss_show = dist_mean(torch.stack(accumulated_latent_losses).sum()).detach().cpu().item() | |
| action_loss_show = dist_mean(torch.stack(accumulated_action_losses).sum()).detach().cpu().item() | |
| max_latent_loss_show = dist_max(torch.stack(accumulated_latent_losses).sum()).detach().cpu().item() | |
| max_action_loss_show = dist_max(torch.stack(accumulated_action_losses).sum()).detach().cpu().item() | |
| train_action_report = self._rmse_report( | |
| dist_sum(accumulated_sq_err_sum), | |
| dist_sum(accumulated_mask_sum), | |
| dist_sum(accumulated_gripper_match_sum) if accumulated_gripper_match_sum is not None else None, | |
| dist_sum(accumulated_gripper_mask_sum) if accumulated_gripper_mask_sum is not None else None, | |
| ) | |
| # Clear accumulated losses | |
| accumulated_latent_losses = [] | |
| accumulated_action_losses = [] | |
| accumulated_sq_err_sum = None | |
| accumulated_mask_sum = None | |
| accumulated_gripper_match_sum = None | |
| accumulated_gripper_mask_sum = None | |
| step_in_accumulation = 0 | |
| torch.cuda.synchronize() | |
| if self.step % self.config.gc_interval == 0: | |
| torch.cuda.empty_cache() | |
| gc.collect() | |
| if self.config.rank == 0: | |
| total_norm = losses['total_norm'] | |
| progress_bar.n += 1 | |
| progress_bar.set_postfix({ | |
| 'latent_loss': f'{latent_loss_show:.4f}', | |
| 'action_loss': f'{action_loss_show:.4f}', | |
| 'act_rmse': f"{train_action_report['overall']:.4f}", | |
| 'step': self.step, | |
| 'grad_norm': f'{total_norm.item():.2f}', | |
| 'lr': f'{lr:.2e}' | |
| }) | |
| logger.info(f"step {self.step} train_action_report: {train_action_report}") | |
| if self.config.enable_wandb: | |
| wandb_log = { | |
| 'loss_metrics/global_avg_video_loss': latent_loss_show, | |
| 'loss_metrics/global_avg_action_loss': action_loss_show, | |
| 'loss_metrics/global_max_video_loss': max_latent_loss_show, | |
| 'loss_metrics/global_max_action_loss': max_action_loss_show, | |
| 'grad_norm': total_norm.item(), | |
| 'lr': lr, | |
| } | |
| wandb_log.update({f'train_action/{k}': v for k, v in train_action_report.items()}) | |
| self.wandb.log(wandb_log, step=self.step) | |
| self.step += 1 | |
| if self.step % self.config.save_interval == 0: | |
| if self.config.rank == 0: | |
| logger.info(f"Starting save model at step {self.step}") | |
| self.save_checkpoint() | |
| if self.val_loader is not None and self.step % self.val_interval == 0: | |
| if self.config.rank == 0: | |
| logger.info(f"Running validation at step {self.step}") | |
| val_report = self.validate() | |
| if self.config.rank == 0: | |
| logger.info(f"step {self.step} validation: {val_report}") | |
| if self.config.enable_wandb: | |
| self.wandb.log({f'val/{k}': v for k, v in val_report.items()}, step=self.step) | |
| if dist.is_initialized(): | |
| dist.barrier() | |
| progress_bar.close() | |
| logger.info("Training completed!") | |
| def run(args): | |
| """Main entry point.""" | |
| config = VA_CONFIGS[args.config_name] | |
| config.run_name = args.config_name | |
| rank = int(os.getenv("RANK", 0)) | |
| local_rank = int(os.environ.get('LOCAL_RANK', 0)) | |
| world_size = int(os.environ.get("WORLD_SIZE", 1)) | |
| init_distributed(world_size, local_rank, rank) | |
| config.rank = rank | |
| config.local_rank = local_rank | |
| config.world_size = world_size | |
| if args.save_root is not None: | |
| config.save_root = args.save_root | |
| if rank == 0: | |
| logger.info(f"Using config: {args.config_name}") | |
| logger.info(f"World size: {world_size}, Local rank: {local_rank}") | |
| trainer = Trainer(config) | |
| trainer.train() | |
| def main(): | |
| """Parse arguments and run training.""" | |
| parser = argparse.ArgumentParser(description="Train WAN model for robotics") | |
| parser.add_argument( | |
| "--config-name", | |
| type=str, | |
| default='robotwin_train', | |
| help="Config name", | |
| ) | |
| parser.add_argument( | |
| "--save-root", | |
| type=str, | |
| default=None, | |
| help="Root directory for saving checkpoints", | |
| ) | |
| args = parser.parse_args() | |
| run(args) | |
| if __name__ == "__main__": | |
| init_logger() | |
| main() |