Download training/train.py from DancingNow/swag-train-bundle: direct link, hf CLI and curl.
- Browser
- Download file 8.94 kB
-
https://huggingface.co/DancingNow/swag-train-bundle/resolve/main/training/train.py
- Command line
-
hf download hf://DancingNow/swag-train-bundle/training/train.py
-
curl -L -o train.py https://huggingface.co/DancingNow/swag-train-bundle/resolve/main/training/train.py
8.94 kB
| 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) | |
| 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() | |