import shutil from contextlib import nullcontext import datetime import functools import logging import math import os import random import time import gc import copy import numpy as np from tqdm import tqdm from einops import rearrange import torch from torch.nn.parallel import DistributedDataParallel as DDP from torch.distributed.fsdp import ( FullyShardedDataParallel as FSDP, ShardingStrategy, ) from torch.distributed.fsdp.wrap import transformer_auto_wrap_policy from torch.distributed import destroy_process_group, init_process_group from torch.utils.data import DataLoader from data_utils import PreprocessDataset from audioloader import AudioLoaderDataset, AudioConfig from modules.base import ( apply_fsdp_checkpointing, configure_optimizers as base_configure_optimizers, estimate_mfu_no_model, LayerNorm, CausalSelfAttention, MLP, Block, ) from utils.fsdp_policies import bfSixteen from utils.helpers import ( dist_barrier, hash_string_to_number, load_checkpoint, load_old_state_dict, load_old_optimizer_state_dict, print_with_time, print_with_time_master, save_checkpoint, save_old_checkpoint, save_dual_model_checkpoint, suppress_logging, verify_preload_model_args, ) from utils.logging import build_gpu_memory_monitor, Color, NoColor from utils.profiling import maybe_enable_memory_snapshot, maybe_enable_profiling from models.model_selector import CodecConfig, DiscriminatorConfig, get_model from modules.losses import LSGANLoss, MelSpectrogramLoss color = Color if True else NoColor # turn down some annoying fsdp logging logging.getLogger("torch.distributed.fsdp._debug_utils").setLevel(logging.ERROR) logging.getLogger("torch.distributed.fsdp._optim_utils").setLevel(logging.ERROR) logging.getLogger("torch.distributed.checkpoint._dedup_tensors").setLevel(logging.ERROR) os.umask(0o003) # set umask to 0o003 to allow group write for created directories data_dir = None out_dir = None enable_profiling = False dump_folder = "/app/suno/gpt_profiling/" save_traces_folder = "traces" profile_freq = 50 enable_memory_snapshot = False save_memory_snapshot_folder = "memory_snapshots" master_addr = "localhost" master_port = 12355 train_metas_filename = "metas_tr.jsonl" val_metas_filename = "metas_val.jsonl" debug_val_only = False preload_checkpoint = None preload_optimizer = False local_cache_dir = None # checkpoint will get copied here, 1 per node to allow for faster loading preload_strict = True # enforce keys in dict on load suppress_compile_warnings = True grad_checkpointing = False checkpoint_save_old_format = True # model params codec_type = "dac_vae" encoder_dim = 128 encoder_rates = [2, 3, 5, 8, 8] latent_dim = 128 decoder_dim = 2048 decoder_rates = [8, 8, 5, 3, 2] vae_dim = 128 is_frozen_encoder = False discriminator_type = "dac_discriminator" discriminator_rates = [] discriminator_periods = [2, 3, 5, 7, 11] discriminator_fft_sizes = [2048, 1024, 512] discriminator_bands = [(0.0, 0.1), (0.1, 0.25), (0.25, 0.5), (0.5, 0.75), (0.75, 1.0)] # data n_channels = 2 sample_rate = 48000 duration_s = 1.0 is_vae = False is_mert = False is_musicfm = False gradient_accumulation_steps = 1 batch_size = 2 # if gradient_accumulation_steps > 1, this is the micro-batch size batch_store_size = 4 # train params weight_mel_loss = 15.0 weight_kl_loss = 0.0001 weight_feat_loss = 2.0 weight_adv_loss = 1.0 weight_disc_loss = 1.0 # eval items custom_seed_offset = 0 eval_interval = 2000 log_interval = 25 eval_iters = 50 eval_only = False # if True, script exits right after the first eval debug_gradients = False # wandb logging wandb_log = False wandb_project = "suno-test" wandb_run_name = "test" wandb_dir = None # model n_transformer_layers = 0 # adamw optimizer learning_rate_codec = 1.5e-4 # codec/generator learning rate (matches impl 1) learning_rate_disc = 3e-4 # discriminator learning rate (matches impl 1) max_iters = 100_000 # total number of training iterations step_save_iters = 20_000 # at this checkpoint we save the model weight_decay = 0.0 # removed weight decay to match impl 3 beta1 = 0.8 beta2 = 0.99 grad_clip_gen = 10.0 # clip gradients at this value, or disable if == 0.0 grad_clip_disc = 1000.0 # clip gradients at this value, or disable if == 0.0 # LR scheduler selection: "inverse", "cosine", "exponential" lr_scheduler_type = "inverse" # "inverse", "cosine", or "exponential" # InverseLR scheduler params inverse_lr_gamma = 200000 inverse_lr_power = 0.5 inverse_lr_warmup = 0.999 # ExponentialLR scheduler params exponential_lr_gamma = 0.999996 # system device = "cuda" dtype = "bfloat16" # "float32", "bfloat16" compile = False # use PyTorch 2.0 to compile the model to be faster fsdp = False # fully sharded data parallel sharding_strategy = "full_shard" # ----------------------------------------------------------------------------- config_keys = [ k for k, v in globals().items() if not k.startswith("_") and isinstance(v, (int, float, bool, str, type(None))) ] exec(open("configurator.py").read()) # overrides from command line or config file config = {k: globals()[k] for k in config_keys} # will be useful for logging # ----------------------------------------------------------------------------- # Validate scheduler type if lr_scheduler_type not in ["inverse", "cosine", "exponential"]: raise ValueError(f"lr_scheduler_type must be 'inverse', 'cosine', or 'exponential', got: {lr_scheduler_type}") use_raw_audio = True assert dtype in ("bfloat16", "float32", "float16") if debug_val_only or eval_only: train_metas_filename = val_metas_filename if not eval_only: wandb_log = False # set up distributed variables os.environ["MASTER_ADDR"] = str(master_addr) os.environ["MASTER_PORT"] = str(master_port) if "SLURM_PROCID" in os.environ: # Running on SLURM if int(os.environ["SLURM_NTASKS_PER_NODE"]) != torch.cuda.device_count(): raise ValueError( f"SLURM_NTASKS_PER_NODE ({os.environ['SLURM_NTASKS_PER_NODE']}) does not match" f" the number of CUDA devices ({torch.cuda.device_count()}) on node {os.environ['HOSTNAME']}" ) ddp_rank = int(os.environ["SLURM_PROCID"]) ddp_local_rank = int(os.environ["SLURM_LOCALID"]) world_size = int(os.environ["SLURM_JOB_NUM_NODES"]) * int(os.environ["SLURM_NTASKS_PER_NODE"]) elif "RANK" in os.environ and "LOCAL_RANK" in os.environ and "WORLD_SIZE" in os.environ: # Running with torchrun ddp_rank = int(os.environ["RANK"]) ddp_local_rank = int(os.environ["LOCAL_RANK"]) world_size = int(os.environ["WORLD_SIZE"]) print(f"Detected torchrun: rank={ddp_rank}, local_rank={ddp_local_rank}, world_size={world_size}") else: # Running locally single GPU ddp_rank = 0 ddp_local_rank = 0 world_size = 1 os.environ["RANK"] = str(ddp_rank) os.environ["LOCAL_RANK"] = str(ddp_local_rank) # various inits, derived attributes, I/O setup ddp = int(os.environ.get("RANK", -1)) != -1 # is this a ddp run? if fsdp: assert ddp, "found fsdp = True but ddp is False" if ddp: print( f"Attempting DDP initialization: rank={ddp_rank}, world_size={world_size}, local_rank={ddp_local_rank}" ) try: print(f"Calling init_process_group with backend=nccl...") init_process_group( backend="nccl", timeout=datetime.timedelta(seconds=2 * 60 * 60), rank=ddp_rank, world_size=world_size, device_id=torch.device(f"cuda:{ddp_local_rank}"), ) print(f"init_process_group completed successfully for rank {ddp_rank}") except Exception as e: print(f"Distributed error on rank {ddp_rank} with host {os.environ.get('HOSTNAME', 'Unknown')}") print(f"Exception details: {e}") raise e device = f"cuda:{ddp_local_rank}" torch.cuda.set_device(device) master_process = ddp_rank == 0 # this process will do logging, checkpointing etc. local_master_process = ddp_local_rank == 0 seed_offset = ddp_rank + 1 # each process gets a different seed print_with_time(f"ddp init, rank {ddp_rank}, local_rank {ddp_local_rank}") else: # if not ddp, we are running on a single gpu, and one process master_process = True local_master_process = True seed_offset = 1 n_gpus_per_node = torch.cuda.device_count() dist_barrier() print_with_time_master(f"ddp init: world size {world_size} ddp_rank {ddp_rank}.") # make sure we offset seeds in a clever way seed_offset += ( custom_seed_offset * world_size + 0 if preload_checkpoint is None else hash_string_to_number(preload_checkpoint) ) torch.manual_seed(6006 + seed_offset) random.seed(6006 + seed_offset) np.random.seed(6006 + seed_offset) torch.backends.cuda.matmul.allow_tf32 = True # allow tf32 on matmul torch.backends.cudnn.allow_tf32 = True # allow tf32 on cudnn device_type = "cuda" if "cuda" in device else "cpu" # for later use in torch.autocast ptdtype = {"float32": torch.float32, "bfloat16": torch.bfloat16, "float16": torch.float16}[dtype] # Use float32 for input audio when frozen encoder is enabled (for better quality) input_dtype = torch.float32 if is_frozen_encoder else ptdtype if is_frozen_encoder and dtype in ("bfloat16", "float16"): print_with_time_master( f"Mixed precision training enabled: Input+Encoder+Quantizer (float32), Decoder ({dtype})" ) ctx = ( nullcontext() if device_type == "cpu" or fsdp else torch.amp.autocast(device_type=device_type, dtype=ptdtype) ) loss_discount_map = {} loss_discount_map["mel_loss"] = weight_mel_loss loss_discount_map["kl_loss"] = weight_kl_loss loss_discount_map["feat_loss"] = weight_feat_loss loss_discount_map["adv_loss"] = weight_adv_loss loss_discount_map["disc_loss"] = weight_disc_loss for k, v in loss_discount_map.items(): print_with_time_master(f"loss weight for {k}: {v:.3f}") date_time_str = datetime.datetime.now().strftime("%Y-%m-%d_%H-%M-%S") out_dir = os.path.join(out_dir, date_time_str) if master_process and not debug_val_only: os.makedirs(out_dir, exist_ok=True) print_with_time_master(f"logging checkpoint here: {out_dir}") # logging if wandb_log and master_process: import wandb import importlib.metadata import sys wandb_log_cfg = {} wandb_log_cfg["run_config"] = {k: v for k, v in config.items()} wandb_log_cfg["world_size"] = world_size wandb_log_cfg["slurm_id"] = os.environ.get("SLURM_JOB_ID") wandb_log_cfg["slurm_name"] = os.environ.get("SLURM_JOB_NAME") wandb_log_cfg["slurm_script_path"] = os.environ.get("SLURM_SCRIPT_PATH") wandb_log_cfg["checkpoint_dir"] = out_dir wandb_log_cfg["pip_freeze"] = { dist.metadata["Name"]: dist.version for dist in importlib.metadata.distributions() } wandb_log_cfg["python_path"] = sys.executable wandb.init(project=wandb_project, name=wandb_run_name, config=wandb_log_cfg, dir=wandb_dir) wandb.run.log_code(".") print_with_time_master(f"Total world size {world_size}") if not use_raw_audio: raise NotImplementedError("Raw audio not implemented") dist_barrier() # model init codec_args = dict( model_type=codec_type, encoder_dim=encoder_dim, encoder_rates=encoder_rates, latent_dim=latent_dim, decoder_dim=decoder_dim, decoder_rates=decoder_rates, vae_dim=vae_dim, sample_rate=sample_rate, is_frozen_encoder=is_frozen_encoder, n_transformer_layers=n_transformer_layers, ) discriminator_args = dict( model_type=discriminator_type, rates=discriminator_rates, periods=discriminator_periods, fft_sizes=discriminator_fft_sizes, sample_rate=sample_rate, bands=discriminator_bands, ) if preload_checkpoint is not None and not preload_checkpoint.endswith(".pt"): # verification for old style we do later on checkpoint load verify_preload_model_args(codec_args, preload_checkpoint, preload_strict=preload_strict) gpu_memory_monitor = build_gpu_memory_monitor() # init a new model from scratch print_with_time_master("Initializing a new model from scratch") codec_config = CodecConfig(**codec_args) codec_model = get_model(codec_config) discriminator_config = DiscriminatorConfig(**discriminator_args) discriminator_model = get_model(discriminator_config) if not fsdp: codec_model.to(device) discriminator_model.to(device) # Convert models to target dtype (bfloat16, float16, or float32) if dtype in ("bfloat16", "float16"): print_with_time_master(f"Converting models to {dtype} for DDP...") codec_model = codec_model.to(dtype=ptdtype) discriminator_model = discriminator_model.to(dtype=ptdtype) # If frozen encoder mode, keep encoder/quantizer in float32 for better quality if is_frozen_encoder: print_with_time_master("Keeping frozen encoder/quantizer in float32 for better quality...") codec_model.encoder = codec_model.encoder.float() codec_model.quantizer = codec_model.quantizer.float() print(codec_args) print(discriminator_args) print(f"codec model params: {codec_model.get_num_params()}") print(f"discriminator model params: {discriminator_model.get_num_params()}") # this is needed to calculate MFU later # it will get messed up by FSDP, so calculate now raw_model_n_params = codec_model.get_num_params() # compile the model if compile: import torch._dynamo torch._dynamo.config.cache_size_limit = 512 # 64 print_with_time_master("compiling the model... (takes a ~minute)") compile_ctx = suppress_logging if suppress_compile_warnings else nullcontext with compile_ctx(): # codec_model = torch.compile(codec_model, fullgraph=True, mode="max-autotune") codec_model = torch.compile(codec_model) discriminator_model = torch.compile(discriminator_model) dist_barrier() else: print_with_time_master("not compiling model.") iter_num = 0 total_hours_processed = 0 rel_hours_processed = 0 best_val_loss = 1e9 # Tracks best validation mel loss (not overall loss) if local_cache_dir is not None and local_master_process: shutil.rmtree(local_cache_dir, ignore_errors=True) os.makedirs(local_cache_dir, exist_ok=True) os.chmod(local_cache_dir, 0o774) # load old single-file checkpoint if preload_checkpoint is not None and preload_checkpoint.endswith(".pt"): # Load checkpoint file if local_cache_dir is not None: local_ckpt_fp = os.path.join(local_cache_dir, "ckpt.pt") if master_process: print_with_time_master("copying checkpoint file to local cachedir...") os.makedirs(local_cache_dir, exist_ok=True) shutil.copy2(preload_checkpoint, local_ckpt_fp) dist_barrier() else: local_ckpt_fp = preload_checkpoint # Load both models from checkpoint print_with_time_master("loading codec and discriminator state_dicts...") checkpoint = torch.load(local_ckpt_fp, map_location=device) # Load codec model if "codec_model" in checkpoint: codec_model.load_state_dict(checkpoint["codec_model"], strict=preload_strict) print_with_time_master("loaded codec model state_dict") elif "model" in checkpoint: # Fallback to old format where only codec was saved codec_model.load_state_dict(checkpoint["model"], strict=preload_strict) print_with_time_master("loaded codec model state_dict (legacy format)") # Load discriminator model if available if "discriminator_model" in checkpoint: discriminator_model.load_state_dict(checkpoint["discriminator_model"], strict=preload_strict) print_with_time_master("loaded discriminator model state_dict") else: print_with_time_master("discriminator state_dict not found in checkpoint - starting fresh") del checkpoint dist_barrier() # order matters: # FSDP: load model ckpt, wrap model, make optim (sharded), shard ckpt into optimizer # DDP: load model ckpt, make optimizer, load model and optimizer checkpoints if fsdp: print_with_time_master("wrapping models in FSDP...") codec_model = FSDP( codec_model, mixed_precision=bfSixteen, sharding_strategy=getattr(ShardingStrategy, sharding_strategy.upper()), device_id=torch.cuda.current_device(), sync_module_states=True, use_orig_params=True, # cpu_offload=torch.distributed.fsdp.CPUOffload(offload_params=True), ) discriminator_model = FSDP( discriminator_model, mixed_precision=bfSixteen, sharding_strategy=getattr(ShardingStrategy, sharding_strategy.upper()), device_id=torch.cuda.current_device(), sync_module_states=True, use_orig_params=True, # cpu_offload=torch.distributed.fsdp.CPUOffload(offload_params=True), ) gpu_mem_stats = gpu_memory_monitor.get_peak_stats() print_with_time_master( f"GPU memory usage for model: " f"{gpu_mem_stats.max_reserved_gib:.2f}GiB" f"({gpu_mem_stats.max_reserved_pct:.2f}%)" ) if grad_checkpointing: apply_fsdp_checkpointing(codec_model) apply_fsdp_checkpointing(discriminator_model) generator_optimizer = base_configure_optimizers( codec_model, weight_decay, learning_rate_codec, # Use separate LR for codec (beta1, beta2), device_type, use_fused=False, is_fsdp=True, ) discriminator_optimizer = base_configure_optimizers( discriminator_model, weight_decay, learning_rate_disc, # Use separate LR for discriminator (beta1, beta2), device_type, use_fused=False, is_fsdp=True, ) else: # both DDP and single-worker # Create separate optimizers for generator (codec) and discriminator generator_optimizer = base_configure_optimizers( codec_model, weight_decay, learning_rate_codec, # Use separate LR for codec (beta1, beta2), device_type, use_fused=False, is_fsdp=False, ) discriminator_optimizer = base_configure_optimizers( discriminator_model, weight_decay, learning_rate_disc, # Use separate LR for discriminator (beta1, beta2), device_type, use_fused=False, is_fsdp=False, ) if ddp: print_with_time_master("wrapping models in DDP") codec_model = DDP(codec_model, device_ids=[ddp_local_rank], find_unused_parameters=True) discriminator_model = DDP(discriminator_model, device_ids=[ddp_local_rank], find_unused_parameters=True) # After DDP wrapping, reconvert frozen encoder/quantizer to float32 for mixed precision if is_frozen_encoder and dtype in ("bfloat16", "float16"): print_with_time_master("Reconverting frozen encoder/quantizer to float32 for mixed precision...") # Access the actual model (unwrap DDP if needed) actual_codec_model = codec_model.module if ddp else codec_model actual_codec_model.encoder = actual_codec_model.encoder.float() actual_codec_model.quantizer = actual_codec_model.quantizer.float() # Verify dtype conversion if master_process: enc_dtype = next(actual_codec_model.encoder.parameters()).dtype quant_params = list(actual_codec_model.quantizer.parameters()) quant_dtype = next(actual_codec_model.quantizer.parameters()).dtype if quant_params else "no params" dec_dtype = next(actual_codec_model.decoder.parameters()).dtype print_with_time_master(f"Dtype check - Encoder: {enc_dtype}, Quantizer: {quant_dtype}, Decoder: {dec_dtype}") torch.cuda.empty_cache() dist_barrier() # load old single-file checkpoint optimizers if preload_checkpoint is not None and preload_checkpoint.endswith(".pt") and preload_optimizer: local_ckpt_fp = ( preload_checkpoint if local_cache_dir is None else os.path.join(local_cache_dir, "ckpt.pt") ) print_with_time_master("loading optimizer state_dicts...") checkpoint = torch.load(local_ckpt_fp, map_location=device) # Load optimizer states if "generator_optimizer" in checkpoint: generator_optimizer.load_state_dict(checkpoint["generator_optimizer"]) print_with_time_master("loaded generator optimizer state_dict") if "discriminator_optimizer" in checkpoint: discriminator_optimizer.load_state_dict(checkpoint["discriminator_optimizer"]) print_with_time_master("loaded discriminator optimizer state_dict") # Load training state if "iter_num" in checkpoint: iter_num = checkpoint["iter_num"] if "best_val_loss" in checkpoint: best_val_loss = checkpoint["best_val_loss"] if "n_hours" in checkpoint: total_hours_processed = checkpoint["n_hours"] del checkpoint print_with_time_master(f"resumed from iteration {iter_num}, best_val_loss: {best_val_loss:.4f}") # load new distributed checkpoint if preload_checkpoint is not None and not preload_checkpoint.endswith(".pt"): # For now, distributed checkpoint loading for dual models not implemented # iter_num, total_tokens_processed, best_val_loss = load_checkpoint( # preload_checkpoint, # preload_optimizer, # codec_model, # generator_optimizer, # ) print_with_time_master("distributed checkpoint loading for dual models not implemented yet") dist_barrier() print_with_time_master("model setup done") # set up loss functions gan_loss = LSGANLoss(discriminator_model) mel_loss = MelSpectrogramLoss().to(device) # InverseLR scheduler def get_inverse_lr(it, base_lr, inv_gamma=200000, power=0.5, warmup=0.999, final_lr=0.0): """Inverse decay learning rate schedule with exponential warmup. Args: it: Current iteration base_lr: Base learning rate inv_gamma: Inverse multiplicative factor of learning rate decay power: Exponential factor of learning rate decay warmup: Exponential warmup factor (0 <= warmup < 1, 0 to disable) final_lr: The final learning rate """ warmup_factor = 1 - warmup ** (it + 1) lr_mult = (1 + it / inv_gamma) ** -power return warmup_factor * max(final_lr, base_lr * lr_mult) # ExponentialLR scheduler def get_exponential_lr(it, base_lr, gamma=0.999996): """Exponential decay learning rate schedule. Args: it: Current iteration base_lr: Base learning rate gamma: Multiplicative factor of learning rate decay per step Returns: Current learning rate """ return base_lr * (gamma ** it) def flatten_stereo(batch): return rearrange(batch, "b c t -> (b c) 1 t") # audio config audio_cfg = AudioConfig( sample_rate=sample_rate, n_channels=n_channels, duration_s=duration_s, is_vae=is_vae, is_mert=is_mert, is_musicfm=is_musicfm, ) def make_data_iter(split, is_eval=False): metas_filename = train_metas_filename if split == "train" else val_metas_filename audio_dataset = AudioLoaderDataset( audio_cfg, os.path.join(data_dir, metas_filename), split=split, ) audio_dataloader = DataLoader( audio_dataset, shuffle=False, num_workers=3 if is_eval else 6, prefetch_factor=batch_store_size * (2 if is_eval else 4), batch_size=None, worker_init_fn=lambda x: random.seed(x + seed_offset * 1000), ) audio_dataloader_iter = iter(audio_dataloader) dataset = PreprocessDataset( audio_dataloader_iter, batch_size, audio_cfg, ) dataloader = DataLoader( dataset, shuffle=False, num_workers=0, # no multiprocessing so its in sync across workers batch_size=None, worker_init_fn=lambda x: random.seed(x + seed_offset * 1000), ) def cached_dataloader_iter_fn(dataloader): # preload batch_store_size batches in memory. # #this way we make new batches every 10 steps, smoothing out the load dataloader_iter = iter(dataloader) batch_store = [] while True: if not batch_store: for _ in tqdm(range(batch_store_size), desc="preloading batches", disable=True): try: batch = next(dataloader_iter) except StopIteration: print("dataloader_iter exhausted, resetting") dataloader_iter = iter(dataloader) batch = next(dataloader_iter) batch_store.append(batch) batch = batch_store.pop(0) yield batch return cached_dataloader_iter_fn(dataloader) tr_dataloader_iter = make_data_iter("train", is_eval=False) val_dataloaders = {} for k in ["train", "val"]: val_dataloaders[k] = {} val_dataloaders[k][0] = make_data_iter(k, is_eval=True) train_dataset_names = ["main"] val_dataset_names = ["main"] @torch.no_grad() def estimate_loss(): n_loss_modalities = len(train_dataset_names) + len(val_dataset_names) effective_eval_iters = int(round(eval_iters / n_loss_modalities)) if fsdp: modules_for_eval = ( torch.nn.Linear, torch.nn.Dropout, torch.nn.Embedding, torch.nn.SiLU, LayerNorm, MLP, CausalSelfAttention, ) for name, module in codec_model.named_modules(): module.train(False) for name, module in discriminator_model.named_modules(): module.train(False) else: codec_model.eval() discriminator_model.eval() loss_prefix = "loss" n_loss_entries = n_loss_modalities * len(loss_discount_map) loss_tensor = torch.zeros(n_loss_entries, device=device) loss_tensor_keys = [] n_loss_entry = 0 for split in ["train", "val"]: n_datasets = len(train_dataset_names) if split == "train" else len(val_dataset_names) dataset_names = train_dataset_names if split == "train" else val_dataset_names for dataset_idx in range(n_datasets): losses = [] for _ in range(effective_eval_iters): data_list = next(val_dataloaders[split][dataset_idx]) with ctx: # For evaluation, just compute codec reconstruction loss raw_audio_batch = [line.data_wav for line in data_list] if is_mert: input_batch = [line.data_mert for line in data_list] elif is_musicfm: input_batch = [line.data_musicfm for line in data_list] elif is_vae: input_batch = [line.data_vae for line in data_list] else: input_batch = raw_audio_batch input_batch = torch.from_numpy(np.stack(input_batch)).to(device, dtype=input_dtype) raw_audio_batch = torch.from_numpy(np.stack(raw_audio_batch)).to( device, dtype=input_dtype ) outp = codec_model(input_batch) gen_audio = outp["audio"] # Compute all losses for evaluation disc_loss_val = gan_loss.discriminator_loss(gen_audio, raw_audio_batch) mel_loss_val = mel_loss(flatten_stereo(gen_audio), flatten_stereo(raw_audio_batch)) mel_loss_val += mel_loss(gen_audio.mean(dim=1), raw_audio_batch.mean(dim=1)) mel_loss_val /= 2 kl_loss_val = outp.get("kl", torch.tensor(0.0, device=device, dtype=ptdtype)) adv_loss_val, feat_loss_val = gan_loss.generator_loss(gen_audio, raw_audio_batch) # Create loss dict matching loss_discount_map keys loss_dict = { "mel_loss": mel_loss_val, "kl_loss": kl_loss_val, "feat_loss": feat_loss_val, "adv_loss": adv_loss_val, "disc_loss": disc_loss_val, } losses.append( [ loss_dict[k].item() if isinstance(loss_dict[k], torch.Tensor) else loss_dict[k] for k in loss_discount_map.keys() ] ) for n, loss_name in enumerate(loss_discount_map.keys()): loss_tensor[n_loss_entry] = float(np.mean([e[n] for e in losses])) loss_tensor_keys.append( f"{split}/{loss_prefix}_{dataset_names[dataset_idx]}_{loss_name}" ) n_loss_entry += 1 if ddp: torch.distributed.all_reduce(loss_tensor, op=torch.distributed.ReduceOp.AVG) tmp_out = {k: loss_tensor[n].item() for n, k in enumerate(loss_tensor_keys)} # add extra loss items out = {k: v for k, v in tmp_out.items()} out[f"train/{loss_prefix}"] = float( np.mean([v for k, v in tmp_out.items() if k.startswith("train/")]) ) out[f"val/{loss_prefix}"] = float(np.mean([v for k, v in tmp_out.items() if k.startswith("val/")])) codec_model.train() discriminator_model.train() for name, module in codec_model.named_modules(): assert module.training # make sure we can undo everything for name, module in discriminator_model.named_modules(): assert module.training # make sure we can undo everything return out gc.disable() # manually gc to avoid slowdowns https://imbue.com/research/70b-infrastructure/ # training loop print_with_time_master("training...") t0 = time.time() t00 = time.time() t_start = time.time() # absolute time since starting to train t_data = 0 # time spent loading data t_wait = 0 # time spent waiting in loop t_model = 0 # time spent in master node model fw+bw local_iter_num = 0 # number of iterations in the lifetime of this process # raw_model = model.module if ddp else model # unwrap DDP container if needed mfu = 0 batch_per_s = 0 effective_batch_per_s_per_node = 0 running_gen_loss = [] running_disc_loss = [] running_mel_loss = [] running_kl_loss = [] running_feat_loss = [] running_adv_loss = [] # with torch.autograd.set_detect_anomaly(True): gpu_memory_monitor.reset_peak_stats() with ( maybe_enable_profiling( enable_profiling, dump_folder, save_traces_folder, profile_freq, global_step=iter_num ) as torch_profiler, maybe_enable_memory_snapshot( enable_memory_snapshot, dump_folder, save_memory_snapshot_folder, profile_freq, global_step=iter_num, ) as memory_profiler, ): while True: # determine and set the learning rate for this iteration if lr_scheduler_type == "inverse": lr_gen = get_inverse_lr(iter_num, learning_rate_codec, inv_gamma=inverse_lr_gamma, power=inverse_lr_power, warmup=inverse_lr_warmup) lr_disc = get_inverse_lr(iter_num, learning_rate_disc, inv_gamma=inverse_lr_gamma, power=inverse_lr_power, warmup=inverse_lr_warmup) elif lr_scheduler_type == "exponential": lr_gen = get_exponential_lr(iter_num, learning_rate_codec, gamma=exponential_lr_gamma) lr_disc = get_exponential_lr(iter_num, learning_rate_disc, gamma=exponential_lr_gamma) else: raise ValueError(f"Unknown lr_scheduler_type: {lr_scheduler_type}") for param_group in generator_optimizer.param_groups: param_group["lr"] = lr_gen for param_group in discriminator_optimizer.param_groups: param_group["lr"] = lr_disc # evaluate the loss on train/val sets and write checkpoints if iter_num % eval_interval == 0 or iter_num == max_iters - 1: dist_barrier() time_since_last_loss = time.time() - t00 t00 = time.time() losses = estimate_loss() estimation_time = time.time() - t00 eval_time_pct = np.clip(estimation_time / time_since_last_loss * 100, 0, 100) # Only master process runs evaluation to avoid conflicts if master_process: print_with_time_master( f"loss estimation took {estimation_time:.1f} seconds. ({eval_time_pct:.1f}% of loop)" ) print_with_time_master(f"step {iter_num}: val loss {losses['val/loss']:.4f}, val mel loss {losses['val/loss_main_mel_loss']:.4f}") if wandb_log: log_dict = { "iter": iter_num, "n_hours": total_hours_processed, } for k, v in losses.items(): log_dict[k] = v wandb.log(log_dict) else: # Other processes create dummy loss dict for checkpoint logic losses = {"val/loss": best_val_loss, "val/loss_main_mel_loss": best_val_loss} dist_barrier() if iter_num > 0: # Save checkpoints using proper FSDP/DDP handling if checkpoint_save_old_format: # Use the dual model checkpoint saving function (best determined by mel loss) best_val_loss = save_dual_model_checkpoint( out_dir=out_dir, codec_model=codec_model, discriminator_model=discriminator_model, generator_optimizer=generator_optimizer, discriminator_optimizer=discriminator_optimizer, best_val_loss=best_val_loss, current_val_loss=losses["val/loss_main_mel_loss"], step_save_iters=step_save_iters, time_since_last_loss=time_since_last_loss, codec_args=codec_args, discriminator_args=discriminator_args, iter_num=iter_num, n_hours=total_hours_processed, debug_val_only=False, save_best_ckpt=True, save_periodic_ckpt=True, ) else: # Use distributed checkpoint format (implement later if needed) print_with_time_master("distributed checkpoint saving not implemented yet") pass dist_barrier() # end if eval test only if eval_only: print_with_time_master("eval test done.") break # GAN training - alternating discriminator and generator updates latest_loss_dict = {} # Store latest loss values for logging grad_norm_disc = None grad_norm_gen = None grad_norm_gen_raw = None accum_gen_loss = 0.0 accum_disc_loss = 0.0 accum_mel_loss = 0.0 accum_kl_loss = 0.0 accum_feat_loss = 0.0 accum_adv_loss = 0.0 # =============================== # Load batches for this training iteration # =============================== batches = [] for micro_step in range(gradient_accumulation_steps): t0_tmp = time.time() data_list = next(tr_dataloader_iter) t_data += time.time() - t0_tmp # Prepare audio data raw_audio_batch = [line.data_wav for line in data_list] if is_mert: input_batch = [line.data_mert for line in data_list] elif is_musicfm: input_batch = [line.data_musicfm for line in data_list] elif is_vae: input_batch = [line.data_vae for line in data_list] else: input_batch = raw_audio_batch input_batch = torch.from_numpy(np.stack(input_batch)).to(device, dtype=input_dtype) raw_audio_batch = torch.from_numpy(np.stack(raw_audio_batch)).to(device, dtype=input_dtype) batches.append((input_batch, raw_audio_batch)) # =============================== # Phase 1: Accumulate and Update Discriminator # =============================== discriminator_optimizer.zero_grad(set_to_none=True) for micro_step in range(gradient_accumulation_steps): if ddp and micro_step < gradient_accumulation_steps - 1: disc_grad_sync_context = discriminator_model.no_sync else: disc_grad_sync_context = nullcontext t0_tmp = time.time() input_batch, raw_audio_batch = batches[micro_step] # Generate fake audio for discriminator training with torch.no_grad(): codec_output = codec_model(input_batch) fake_audio = codec_output["audio"] # Train Discriminator with disc_grad_sync_context(): with ctx: # Compute discriminator loss # Note: discriminator_loss internally detaches fake_audio disc_loss = gan_loss.discriminator_loss(fake_audio, raw_audio_batch) disc_loss = disc_loss / gradient_accumulation_steps # Backward pass for discriminator disc_loss.backward() accum_disc_loss += float(disc_loss.item()) t_model += time.time() - t0_tmp # Clip discriminator gradients after accumulation if grad_clip_disc != 0.0: if fsdp: grad_norm_disc = discriminator_model.clip_grad_norm_(grad_clip_disc) if isinstance(grad_norm_disc, torch.Tensor): if torch.isnan(grad_norm_disc).any(): raise RuntimeError("Found NaN in discriminator grad") elif math.isnan(float(grad_norm_disc)): raise RuntimeError("Found NaN in discriminator grad") else: grad_norm_disc = torch.nn.utils.clip_grad_norm_( discriminator_model.parameters(), grad_clip_disc, error_if_nonfinite=True ) grad_norm_disc = ( grad_norm_disc.item() if isinstance(grad_norm_disc, torch.Tensor) else grad_norm_disc ) else: if isinstance(grad_norm_disc, torch.Tensor): grad_norm_disc = grad_norm_disc.item() # Step discriminator discriminator_optimizer.step() discriminator_optimizer.zero_grad(set_to_none=True) # =============================== # Phase 2: Accumulate and Update Generator (with updated discriminator) # =============================== generator_optimizer.zero_grad(set_to_none=True) for micro_step in range(gradient_accumulation_steps): if ddp and micro_step < gradient_accumulation_steps - 1: codec_grad_sync_context = codec_model.no_sync else: codec_grad_sync_context = nullcontext t0_tmp = time.time() # Use the SAME batches as discriminator training input_batch, raw_audio_batch = batches[micro_step] # Generate fake audio for generator training (with gradients) codec_output = codec_model(input_batch) fake_audio = codec_output["audio"] # Train Generator with codec_grad_sync_context(): with ctx: # Compute all generator losses (discriminator now uses updated weights) mel_loss_val = mel_loss(flatten_stereo(fake_audio), flatten_stereo(raw_audio_batch)) mel_loss_val += mel_loss(fake_audio.mean(dim=1), raw_audio_batch.mean(dim=1)) mel_loss_val /= 2 kl_loss_val = codec_output.get("kl", torch.tensor(0.0, device=device, dtype=ptdtype)) adv_loss_val, feat_loss_val = gan_loss.generator_loss(fake_audio, raw_audio_batch) # Combine losses with weights total_gen_loss = ( weight_mel_loss * mel_loss_val + weight_kl_loss * kl_loss_val + weight_feat_loss * feat_loss_val + weight_adv_loss * adv_loss_val ) total_gen_loss = total_gen_loss / gradient_accumulation_steps # Store loss values for logging loss_dict = { "mel_loss": mel_loss_val, "kl_loss": kl_loss_val, "feat_loss": feat_loss_val, "adv_loss": adv_loss_val, "disc_loss": accum_disc_loss / gradient_accumulation_steps, # Use accumulated disc loss } latest_loss_dict = loss_dict # Store for later logging # Backward pass for generator total_gen_loss.backward() accum_gen_loss += float(total_gen_loss.item()) # Component losses need to be scaled for correct logging accum_mel_loss += float((mel_loss_val / gradient_accumulation_steps).item() if isinstance(mel_loss_val, torch.Tensor) else mel_loss_val / gradient_accumulation_steps) accum_kl_loss += float((kl_loss_val / gradient_accumulation_steps).item() if isinstance(kl_loss_val, torch.Tensor) else kl_loss_val / gradient_accumulation_steps) accum_feat_loss += float((feat_loss_val / gradient_accumulation_steps).item() if isinstance(feat_loss_val, torch.Tensor) else feat_loss_val / gradient_accumulation_steps) accum_adv_loss += float((adv_loss_val / gradient_accumulation_steps).item() if isinstance(adv_loss_val, torch.Tensor) else adv_loss_val / gradient_accumulation_steps) # Update processed hours (assuming each sample is duration_s seconds) batch_hours = (batch_size * duration_s / 3600.0) * world_size total_hours_processed += batch_hours rel_hours_processed += batch_hours if debug_gradients and wandb_log and master_process: d = { "iter": iter_num, "n_hours": total_hours_processed, "total_gen_loss": total_gen_loss.item(), "disc_loss": accum_disc_loss / gradient_accumulation_steps, } # Log individual loss components for k, v in loss_dict.items(): d[f"loss/{k}"] = v.item() if isinstance(v, torch.Tensor) else v wandb.log(d) t_model += time.time() - t0_tmp # Clip generator gradients after accumulation if grad_clip_gen != 0.0: if fsdp: grad_norm_gen = codec_model.clip_grad_norm_(grad_clip_gen) if isinstance(grad_norm_gen, torch.Tensor): if torch.isnan(grad_norm_gen).any(): raise RuntimeError("Found NaN in generator grad") elif math.isnan(float(grad_norm_gen)): raise RuntimeError("Found NaN in generator grad") else: grad_norm_gen_raw = torch.nn.utils.clip_grad_norm_( codec_model.parameters(), float("inf") ) grad_norm_gen_raw = ( grad_norm_gen_raw.item() if isinstance(grad_norm_gen_raw, torch.Tensor) else float(grad_norm_gen_raw) ) if grad_norm_gen_raw > 100.0: print(f"⚠️ WARNING: Large generator gradient: {grad_norm_gen_raw:.2f}") grad_norm_gen = torch.nn.utils.clip_grad_norm_( codec_model.parameters(), grad_clip_gen, error_if_nonfinite=True ) grad_norm_gen = ( grad_norm_gen.item() if isinstance(grad_norm_gen, torch.Tensor) else grad_norm_gen ) if grad_norm_gen_raw is not None and isinstance(grad_norm_gen_raw, torch.Tensor): grad_norm_gen_raw = grad_norm_gen_raw.item() else: if isinstance(grad_norm_gen, torch.Tensor): grad_norm_gen = grad_norm_gen.item() # Step generator generator_optimizer.step() generator_optimizer.zero_grad(set_to_none=True) running_gen_loss.append(accum_gen_loss) running_disc_loss.append(accum_disc_loss) running_mel_loss.append(accum_mel_loss) running_kl_loss.append(accum_kl_loss) running_feat_loss.append(accum_feat_loss) running_adv_loss.append(accum_adv_loss) if master_process and wandb_log: wandb.log( { "iter": iter_num, "n_hours": total_hours_processed, "misc/grad_norm_gen": grad_norm_gen if grad_norm_gen is not None else 0.0, "misc/grad_norm_disc": grad_norm_disc if grad_norm_disc is not None else 0.0, } ) if iter_num % 300 == 1: # manually collect garbage in sync to avoid slowdowns over time # this slows down the iteration by ~30%, dont gc too often # do it on mod 1 to avoid it showing up in the logs assert not gc.isenabled() gc.collect() # timing and logging t1 = time.time() dt = t1 - t0 t0 = t1 if iter_num % log_interval == 0 or iter_num == max_iters - 1: if master_process: if local_iter_num >= 5: # let the training loop settle a bit # Note: we process gradient_accumulation_steps batches (same batches for both D and G) batch_per_s = world_size * batch_size * gradient_accumulation_steps / dt effective_batch_per_s_per_node = ( rel_hours_processed * 3600 / duration_s # convert hours back to batches / (time.time() - t_start) / max(1, int(round(world_size / n_gpus_per_node))) ) pct_wait = t_wait / (time.time() - t_start) * 100 pct_data = t_data / (time.time() - t_start) * 100 pct_model = t_model / (time.time() - t_start) * 100 pct_overhead = 100 - pct_wait - pct_data - pct_model avg_gen_loss = np.mean(running_gen_loss) if running_gen_loss else 0.0 avg_disc_loss = np.mean(running_disc_loss) if running_disc_loss else 0.0 avg_mel_loss = np.mean(running_mel_loss) if running_mel_loss else 0.0 avg_kl_loss = np.mean(running_kl_loss) if running_kl_loss else 0.0 avg_feat_loss = np.mean(running_feat_loss) if running_feat_loss else 0.0 avg_adv_loss = np.mean(running_adv_loss) if running_adv_loss else 0.0 running_gen_loss = [] running_disc_loss = [] running_mel_loss = [] running_kl_loss = [] running_feat_loss = [] running_adv_loss = [] gpu_mem_stats = gpu_memory_monitor.get_peak_stats() print_with_time_master( f"{color.cyan}iter {iter_num}:" f"{color.green} gen_loss {avg_gen_loss:.3f}," f"{color.red} disc_loss {avg_disc_loss:.3f}," f"{color.magenta} mel_loss {avg_mel_loss:.3f}," f"{color.cyan} kl_loss {avg_kl_loss:.6f}," f"{color.yellow} feat_loss {avg_feat_loss:.3f}," f"{color.blue} adv_loss {avg_adv_loss:.3f}," f"{color.cyan} step_time {dt * 1000:.1f}ms," f"{color.green} throughput {batch_per_s:.1f} batch/s," f"{color.yellow} memory {gpu_mem_stats.max_reserved_gib:5.2f}GiB" f"({gpu_mem_stats.max_reserved_pct:.2f}%)" f"{color.reset}" ) if wandb_log: log_dict = { "iter": iter_num, "n_hours": total_hours_processed, "avg_gen_loss": avg_gen_loss, "avg_disc_loss": avg_disc_loss, "avg_mel_loss": avg_mel_loss, "avg_kl_loss": avg_kl_loss, "avg_feat_loss": avg_feat_loss, "avg_adv_loss": avg_adv_loss, "lr_gen": lr_gen, "lr_disc": lr_disc, "batch/s": batch_per_s, "eff_batch/s/node": effective_batch_per_s_per_node, "perf/pct_wait": pct_wait, "perf/pct_data": pct_data, "perf/pct_model": pct_model, "perf/pct_overhead": pct_overhead, "perf/t_data": t_data, "perf/t_model": t_model, "perf/t_wait": t_wait, "memory/max_active(GiB)": gpu_mem_stats.max_active_gib, "memory/max_active(%)": gpu_mem_stats.max_active_pct, "memory/max_reserved(GiB)": gpu_mem_stats.max_reserved_gib, "memory/max_reserved(%)": gpu_mem_stats.max_reserved_pct, "memory/num_alloc_retries": gpu_mem_stats.num_alloc_retries, "memory/num_ooms": gpu_mem_stats.num_ooms, } wandb.log(log_dict) gpu_memory_monitor.reset_peak_stats() iter_num += 1 local_iter_num += 1 # signals the profiler that the next profiling step has started if torch_profiler: torch_profiler.step() if memory_profiler: memory_profiler.step() # termination conditions if iter_num >= max_iters: print_with_time_master("done.") break dist_barrier() print_with_time_master("removing cache dir.") if local_cache_dir is not None and local_master_process: shutil.rmtree(local_cache_dir, ignore_errors=True) dist_barrier() print_with_time_master("done.") dist_barrier() if ddp: destroy_process_group()