SleepMastger's picture
add model card, conditioning, and training-time processing
0613ebe verified
Raw History Blame Contribute Delete
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
@torch.no_grad()
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,
)
@torch.no_grad()
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
@torch.no_grad()
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
@torch.no_grad()
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()