from __future__ import annotations import argparse import copy import json import os import random from pathlib import Path import numpy as np import torch import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel from torch.utils.data import DataLoader from torch.utils.data.distributed import DistributedSampler from tqdm import tqdm from config import load_config, resolve_path from diffusion import create_diffusion from hdf5_dataset import WaveformDataset from models import EmptyConditionSWaG, SWaG PROJECT_ROOT = Path(__file__).resolve().parents[2] def distributed_context() -> tuple[int, int, int]: world_size = int(os.environ.get("WORLD_SIZE", "1")) rank = int(os.environ.get("RANK", "0")) local_rank = int(os.environ.get("LOCAL_RANK", "0")) if world_size > 1: dist.init_process_group("nccl") return rank, local_rank, world_size def set_seed(seed: int, rank: int) -> None: value = seed + rank random.seed(value) np.random.seed(value) torch.manual_seed(value) torch.cuda.manual_seed_all(value) @torch.no_grad() def update_ema(ema: torch.nn.Module, model: torch.nn.Module, decay: float) -> None: source = model.module if isinstance(model, DistributedDataParallel) else model for target_parameter, source_parameter in zip(ema.parameters(), source.parameters()): target_parameter.mul_(decay).add_(source_parameter, alpha=1.0 - decay) for target_buffer, source_buffer in zip(ema.buffers(), source.buffers()): target_buffer.copy_(source_buffer) def save_checkpoint( path: Path, model: torch.nn.Module, ema: torch.nn.Module, optimizer: torch.optim.Optimizer, epoch: int, step: int, config: dict, ) -> None: source = model.module if isinstance(model, DistributedDataParallel) else model payload = { "model": source.state_dict(), "ema": ema.state_dict(), "optimizer": optimizer.state_dict(), "epoch": epoch, "step": step, "config": config, } temporary = path.with_suffix(".tmp") torch.save(payload, temporary) temporary.replace(path) def main() -> None: parser = argparse.ArgumentParser(description="Train SWaG") parser.add_argument("--config", type=Path, required=True) args = parser.parse_args() if not torch.cuda.is_available(): raise RuntimeError("CUDA is required for training") config = load_config(args.config.resolve()) rank, local_rank, world_size = distributed_context() device = torch.device("cuda", local_rank) torch.cuda.set_device(device) training = config["training"] set_seed(int(training["seed"]), rank) data_path = resolve_path(PROJECT_ROOT, config["data"]["train_h5"]) output_dir = resolve_path(PROJECT_ROOT, config["output"]["directory"]) checkpoint_dir = output_dir / "checkpoints" if rank == 0: checkpoint_dir.mkdir(parents=True, exist_ok=True) (output_dir / "config.yaml").write_text(args.config.read_text(encoding="utf-8"), encoding="utf-8") if world_size > 1: dist.barrier() dataset = WaveformDataset( data_path, dataset_key=str(config["data"].get("dataset_key", "data")), max_samples=int(config["data"].get("max_samples", 0)), ) model_config = dict(config["model"]) if dataset.channels != int(model_config["in_channels"]) or dataset.waveform_length != int(model_config["length"]): raise ValueError( f"Data shape [C={dataset.channels}, L={dataset.waveform_length}] does not match model " f"[C={model_config['in_channels']}, L={model_config['length']}]" ) global_batch = int(training["global_batch_size"]) accumulation = int(training.get("gradient_accumulation_steps", 1)) divisor = world_size * accumulation if global_batch % divisor: raise ValueError(f"global_batch_size must be divisible by world_size * accumulation ({divisor})") local_batch = global_batch // divisor sampler = DistributedSampler(dataset, shuffle=True, seed=int(training["seed"])) if world_size > 1 else None loader = DataLoader( dataset, batch_size=local_batch, shuffle=sampler is None, sampler=sampler, num_workers=int(training.get("num_workers", 8)), pin_memory=True, persistent_workers=int(training.get("num_workers", 8)) > 0, drop_last=True, ) model_class = EmptyConditionSWaG if int(model_config.get("condition_slot_count", 0)) == 8 else SWaG model = model_class(**model_config).to(device) ema = copy.deepcopy(model).to(device).eval() for parameter in ema.parameters(): parameter.requires_grad_(False) if world_size > 1: model = DistributedDataParallel(model, device_ids=[local_rank]) optimizer = torch.optim.AdamW( model.parameters(), lr=float(training["learning_rate"]), weight_decay=float(training.get("weight_decay", 0.0)), ) diffusion = create_diffusion(timestep_respacing="", **config["diffusion"]) start_epoch = 0 step = 0 resume = str(training.get("resume_checkpoint", "")).strip() if resume: checkpoint = torch.load(resolve_path(PROJECT_ROOT, resume), map_location="cpu", weights_only=False) source = model.module if isinstance(model, DistributedDataParallel) else model source.load_state_dict(checkpoint["model"], strict=True) ema.load_state_dict(checkpoint["ema"], strict=True) optimizer.load_state_dict(checkpoint["optimizer"]) start_epoch = int(checkpoint["epoch"]) step = int(checkpoint["step"]) use_amp = bool(training.get("use_amp", True)) # PyTorch changed GradScaler from torch.cuda.amp to torch.amp; support # both APIs so the public training script works across DiT environments. if hasattr(torch, "amp") and hasattr(torch.amp, "GradScaler"): scaler = torch.amp.GradScaler("cuda", enabled=use_amp) else: scaler = torch.cuda.amp.GradScaler(enabled=use_amp) epochs = int(training["epochs"]) ema_decay = float(training.get("ema_decay", 0.9999)) checkpoint_interval = int(config["output"].get("checkpoint_every_epochs", 5)) optimizer.zero_grad(set_to_none=True) for epoch in range(start_epoch, epochs): if sampler is not None: sampler.set_epoch(epoch) progress = tqdm(loader, disable=rank != 0, desc=f"epoch {epoch + 1}/{epochs}") running_loss = 0.0 for batch_index, waveforms in enumerate(progress): waveforms = waveforms.to(device, non_blocking=True) timesteps = torch.randint(0, diffusion.num_timesteps, (waveforms.shape[0],), device=device) model_kwargs = None if int(model_config.get("condition_slot_count", 0)) == 8: model_kwargs = { "conditions": torch.zeros( waveforms.shape[0], 8, device=device, dtype=torch.float32 ) } with torch.autocast("cuda", enabled=use_amp, dtype=torch.float16): loss = diffusion.training_losses( model, waveforms, timesteps, model_kwargs=model_kwargs )["loss"].mean() / accumulation scaler.scale(loss).backward() running_loss += float(loss.detach()) * accumulation should_step = (batch_index + 1) % accumulation == 0 if should_step: max_norm = training.get("grad_clip_norm") if max_norm is not None: scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), float(max_norm)) scaler.step(optimizer) scaler.update() optimizer.zero_grad(set_to_none=True) update_ema(ema, model, ema_decay) step += 1 if rank == 0: progress.set_postfix(loss=f"{running_loss / (batch_index + 1):.5f}") mean_loss = torch.tensor(running_loss / max(len(loader), 1), device=device) if world_size > 1: dist.all_reduce(mean_loss, op=dist.ReduceOp.SUM) mean_loss /= world_size if rank == 0: metrics = {"epoch": epoch + 1, "step": step, "loss": mean_loss.item()} with (output_dir / "metrics.jsonl").open("a", encoding="utf-8") as handle: handle.write(json.dumps(metrics) + "\n") if (epoch + 1) % checkpoint_interval == 0 or epoch + 1 == epochs: save_checkpoint( checkpoint_dir / f"checkpoint_epoch_{epoch + 1:05d}.pt", model, ema, optimizer, epoch + 1, step, config, ) if world_size > 1: dist.barrier() dataset.close() if world_size > 1: dist.destroy_process_group() if __name__ == "__main__": main()