DancingNow's picture
Add files using upload-large-folder tool
923fdff verified
Raw History Blame Contribute Delete
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)
@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()