import datetime import inspect import logging import os import random import time import math from contextlib import contextmanager, nullcontext import numpy as np import torch import torch.nn.functional as F from torch.nn.parallel import DistributedDataParallel as DDP from torch.distributed import init_process_group, destroy_process_group import torchaudio.functional as aF from torchaudio.transforms import MelSpectrogram from discriminator import MultiScaleSTFTDiscriminator from loss import stft from model import CirceNet @contextmanager def suppress_logging(highest_level=logging.CRITICAL): previous_level = logging.root.manager.disable logging.disable(highest_level) try: yield finally: logging.disable(previous_level) # data params data_dir = None out_dir = None train_filename = None val_filename = None # model params dimension = 512 n_filters = 64 sample_rate = 48_000 ratios = (8, 5, 4, 4) # (8, 5, 4, 3, 2) n_codebooks = 8 causal = False # misc n_steps_adv_start = 5000 cycle_sample_rate = True disc_only = False gen_only = False skip_quantization = False custom_seed_offset = 123 match_val_weights = True randomize_stft = True preload_checkpoint = None preload_optimizer = False preload_strict = True preload_checkpoint_adv = None preload_optimizer_adv = False preload_strict_adv = True suppress_compile_warnings = True eval_interval = 2000 log_interval = 25 eval_iters = 500 eval_only = False # if True, script exits right after the first eval always_save_checkpoint = True # if True, always save a checkpoint after each eval # wandb logging wandb_log = False # disabled by default wandb_project = "suno" wandb_run_name = "base" # data gradient_accumulation_steps = 1 # used to simulate larger batch sizes batch_size = 16 # if gradient_accumulation_steps > 1, this is the micro-batch size # adamw optimizer learning_rate = 3e-4 max_iters = 250000 # total number of training iterations beta1 = 0.5 beta2 = 0.9 grad_clip = 1.0 # clip gradients at this value, or disable if == 0.0 # learning rate decay settings decay_lr = True # whether to decay the learning rate warmup_iters = 1000 # how many steps to warm up for lr_decay_iters = None # should be ~= max_iters per Chinchilla min_lr = 0 # minimum learning rate, should be ~= learning_rate/10 per Chinchilla # DDP settings backend = "nccl" # "nccl", "gloo", etc. # system device = "cuda" # examples: "cpu", "cuda", "cuda:0", "cuda:1" etc., or try "mps" on macbooks dtype = "float32" # "float32", "bfloat16", or "float16" (implements a GradScaler) compile = False # use PyTorch 2.0 to compile the model to be faster # TODO: technically missing the 0.1 l1_loss in time domain loss_factor_map = { "rec": 25.0, # 100 "disc": 1.0, "gen_h": 4.0, "gen_f": 4.0, "comm": 1.0, # 100 } # ----------------------------------------------------------------------------- config_keys = [ k for k,v in globals().items() if not k.startswith("_") and isinstance(v, (int, float, bool, str)) ] exec(open("custom_configurator.py").read()) # overrides from command line or config file config = {k: globals()[k] for k in config_keys} # will be useful for logging # ----------------------------------------------------------------------------- assert(dtype == "float32") assert(not (gen_only and disc_only)) assert(gradient_accumulation_steps == 1) eval_iters = int(eval_iters * gradient_accumulation_steps / 5) if lr_decay_iters is None: lr_decay_iters = max_iters # various inits, derived attributes, I/O setup ddp = int(os.environ.get("RANK", -1)) != -1 # is this a ddp run? if ddp: init_process_group(backend=backend) ddp_rank = int(os.environ["RANK"]) ddp_local_rank = int(os.environ["LOCAL_RANK"]) world_size = torch.distributed.get_world_size() device = f"cuda:{ddp_local_rank}" torch.cuda.set_device(device) master_process = ddp_rank == 0 # this process will do logging, checkpointing etc. seed_offset = ddp_rank # each process gets a different seed else: # if not ddp, we are running on a single gpu, and one process master_process = True seed_offset = 0 world_size = 1 seed_offset += 1 seed_offset *= custom_seed_offset + 1 # multiply to not just shift 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 # note: float16 data type will automatically use a GradScaler ptdtype = {"float32": torch.float32, "bfloat16": torch.bfloat16, "float16": torch.float16}[dtype] ctx = ( nullcontext() if device_type == "cpu" else torch.amp.autocast(device_type=device_type, dtype=ptdtype) ) # logging if wandb_log and master_process: import wandb wandb.init(project=wandb_project, name=wandb_run_name, config=config) 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: os.makedirs(out_dir, exist_ok=True) print(f"logging checkpoint here: {out_dir}") # load data train_data = np.memmap(os.path.join(data_dir, train_filename), dtype=np.int16, mode="r") val_data = np.memmap(os.path.join(data_dir, val_filename), dtype=np.int16, mode="r") block_size = sample_rate COMMON_SAMPLE_RATES = [8000, 16000, 24000, 32000, 44100, 48000] def _cycle_sample_rate(waveform, from_sample_rate=48_000, to_sample_rate=8_000): assert(isinstance(waveform, torch.Tensor)) assert(len(waveform.shape) == 2) resampled_waveform = aF.resample(waveform, from_sample_rate, to_sample_rate) cycled_waveform = aF.resample(resampled_waveform, to_sample_rate, from_sample_rate) assert(waveform.shape == cycled_waveform.shape) return cycled_waveform def get_sample(split, is_training=False): data = train_data if split == "train" else val_data idx = random.randint(0, len(data)-block_size) arr = np.array(data[idx:idx+block_size]) arr = torch.from_numpy(arr.astype(np.float32) / np.iinfo(np.int16).max)[None] if cycle_sample_rate and is_training and random.random() >= 0.75: arr = _cycle_sample_rate( arr, from_sample_rate=sample_rate, to_sample_rate=random.choice(COMMON_SAMPLE_RATES) ) x = arr y = x.clone() return x, y def get_batch(split, is_training=False): x_list = [] y_list = [] for _ in range(batch_size): x, y = get_sample(split, is_training=is_training) x_list.append(x) y_list.append(y) x = torch.stack(x_list) y = torch.stack(y_list) if device_type == "cuda": # pin arrays x,y, which allows us to move them to GPU asynchronously (non_blocking=True) x = x.pin_memory().to(device, non_blocking=True) y = y.pin_memory().to(device, non_blocking=True) else: x, y = x.to(device), y.to(device) del x_list, y_list return x, y iter_num = 0 best_val_loss = 1e9 # init a new model from scratch print("Initializing a new model from scratch") model = CirceNet( dimension=dimension, n_filters=n_filters, ratios=ratios, causal=causal, skip_quantization=skip_quantization, n_codebooks=n_codebooks, ) model.to(device) model_adv = MultiScaleSTFTDiscriminator(32) model_adv.to(device) # optimizer use_fused = (device_type == "cuda") and ("fused" in inspect.signature(torch.optim.AdamW).parameters) print(f"using fused AdamW: {use_fused}") extra_args = dict(fused=True) if use_fused else dict() optimizer = torch.optim.AdamW( [{"params": model.parameters(), "lr": learning_rate}], betas=(beta1, beta2), **extra_args ) optimizer_adv = torch.optim.AdamW( [{"params": model_adv.parameters(), "lr": learning_rate}], betas=(beta1, beta2), **extra_args ) # load checkpoint if preload_checkpoint is not None: print("preloading checkpoint") checkpoint = torch.load(preload_checkpoint, map_location="cpu") if "model" in checkpoint: state_dict = checkpoint["model"] # fix the keys of the state dictionary :( # honestly no idea how checkpoints sometimes get this prefix, have to debug more unwanted_prefix = "_orig_mod." for k, v in list(state_dict.items()): if k.startswith(unwanted_prefix): state_dict[k[len(unwanted_prefix):]] = state_dict.pop(k) else: state_dict = checkpoint if not preload_strict: # clean up orig_state_dict = model.state_dict() n_dropped = 0 n_groups = len(state_dict) for k in list(state_dict.keys()): if k not in orig_state_dict or state_dict[k].shape != orig_state_dict[k].shape: state_dict.pop(k) n_dropped += 1 print(f"dropped {n_dropped}/{n_groups} state dict groups") del orig_state_dict model.load_state_dict(state_dict, strict=preload_strict) if preload_optimizer: print("preloading optimizer") optimizer.load_state_dict(checkpoint["optimizer"]) del state_dict, checkpoint if preload_checkpoint_adv is not None: print("preloading checkpoint adv") checkpoint = torch.load(preload_checkpoint_adv, map_location="cpu") # hack if "raw_model_adv" in checkpoint: checkpoint["model_adv"] = checkpoint["raw_model_adv"] if "model_adv" in checkpoint: state_dict = checkpoint["model_adv"] # fix the keys of the state dictionary :( # honestly no idea how checkpoints sometimes get this prefix, have to debug more unwanted_prefix = "_orig_mod." for k, v in list(state_dict.items()): if k.startswith(unwanted_prefix): state_dict[k[len(unwanted_prefix):]] = state_dict.pop(k) else: state_dict = checkpoint if not preload_strict_adv: # clean up orig_state_dict = model_adv.state_dict() n_dropped = 0 n_groups = len(state_dict) for k in list(state_dict.keys()): if k not in orig_state_dict or state_dict[k].shape != orig_state_dict[k].shape: state_dict.pop(k) n_dropped += 1 print(f"dropped {n_dropped}/{n_groups} state dict groups") del orig_state_dict model_adv.load_state_dict(state_dict, strict=preload_strict_adv) if preload_optimizer: print("preloading optimizer") optimizer_adv.load_state_dict(checkpoint["optimizer_adv"]) del state_dict, checkpoint torch.cuda.empty_cache() torch.cuda.synchronize() # compile the model if compile: print("compiling the model... (takes a ~minute)") compile_ctx = suppress_logging if suppress_compile_warnings else nullcontext with compile_ctx(): model = torch.compile(model) # requires PyTorch 2.0 model_adv = torch.compile(model_adv) # wrap model into DDP container if ddp: model = DDP(model, device_ids=[ddp_local_rank]) model_adv = DDP(model_adv, device_ids=[ddp_local_rank]) def reconstruction_loss(x, G_x, eps=1e-7): L = 100 * F.mse_loss(x, G_x) # wav L1 loss for i in range(6, 11): s = 2**i melspec = MelSpectrogram( sample_rate=sample_rate, n_fft=s, hop_length=s // 4, n_mels=64, wkwargs={"device": x.device}).to(x.device) S_x = melspec(x) S_G_x = melspec(G_x) loss = ((S_x - S_G_x).abs().mean() + ( ((torch.log(S_x.abs() + eps) - torch.log(S_G_x.abs() + eps))**2 ).mean(dim=-2)**0.5).mean()) / (i) L += loss return L def gen_loss(fmap_real, fmap_fake, logits_real, logits_fake): # hinge loss loss_h = torch.tensor(0.0, device=logits_fake[0].device, requires_grad=True) n_logits = len(logits_fake) for lf in logits_fake: loss_h = loss_h + F.relu(1 - lf).mean() / n_logits # f1 for features loss_f = torch.tensor(0.0, device=logits_fake[0].device, requires_grad=True) for er, ef in zip(fmap_real, fmap_fake): for eer, eef in zip(er, ef): # loss_f = loss_f + F.l1_loss(eer, eef) loss_f = loss_f + ((eer - eef).abs() / (eer.abs().mean())).mean() # missing factor of 100 loss_f = loss_f / (len(logits_fake) * len(logits_fake[0])) # sim loss loss_s = torch.tensor(0.0, device=logits_fake[0].device, requires_grad=True) for lr, lf in zip(logits_real, logits_fake): loss_s = loss_s + F.mse_loss(lr, lf) / len(logits_fake) loss_s = loss_s / len(logits_fake) return loss_h, loss_f + loss_s def disc_loss(logits_real, logits_fake): loss = torch.tensor(0.0, device=logits_fake[0].device, requires_grad=True) n_logits = len(logits_fake) for lr, lf in zip(logits_real, logits_fake): loss = loss + (F.relu(1-lr) + F.relu(1+lf)).mean() / n_logits return loss # helps estimate an arbitrarily accurate loss over either split using many batches @torch.no_grad() def estimate_loss(): loss_types = ["rec_loss", "disc_loss", "gen_h_loss", "gen_f_loss", "comm_loss"] n_loss_modalities = 1 + 1 # train + val effective_eval_iters = int(round(eval_iters / n_loss_modalities)) model.eval() n_loss_entries = n_loss_modalities * len(loss_types) loss_tensor = torch.zeros(n_loss_entries, device=device) tensor_keys = [] n_loss_entry = 0 for split in ["train", "val"]: losses = [[] for _ in loss_types] for k in range(effective_eval_iters): X, Y = get_batch(split) with ctx: y_pred, loss_dict = model(X, Y) y_adv_real, fmap_real = model_adv(Y) y_adv_fake, fmap_fake = model_adv(y_pred) loss_adv = disc_loss(y_adv_real, y_adv_fake) loss_gen_h, loss_gen_f = gen_loss(fmap_real, fmap_fake, y_adv_real, y_adv_fake) # losses[0].append(reconstruction_loss(Y, y_pred).item()) losses[0].append((loss_dict["sc_loss"] + loss_dict["mag_loss"]).item()) losses[1].append(loss_adv.item()) losses[2].append(loss_gen_h.item()) losses[3].append(loss_gen_f.item()) losses[4].append(0 if loss_dict["comm_loss"] is None else loss_dict["comm_loss"].item()) for loss_type_idx, loss_type in enumerate(loss_types): loss_tensor[n_loss_entry] = np.mean(losses[loss_type_idx]) tensor_keys.append(f"{split}/{loss_type}") n_loss_entry += 1 if ddp: torch.distributed.all_reduce(loss_tensor, op=torch.distributed.ReduceOp.AVG) out = {k: loss_tensor[n].item() for n, k in enumerate(tensor_keys)} # add extra loss items for split in ["train", "val"]: out[f"{split}/loss"] = ( out[f"{split}/rec_loss"] + out[f"{split}/gen_h_loss"] + out[f"{split}/gen_f_loss"] ) model.train() for name, module in model.named_modules(): assert module.training # make sure we can undo everything return out # save spec images def _get_im(x, fft_size=2048, win_length=1024, hop_size=64*2): assert(isinstance(x, torch.Tensor)) assert(len(x.shape) == 2) assert(x.shape[0] == 1) x = x.detach().cpu() window = torch.hann_window(win_length) x_mag = stft(x, fft_size, hop_size, win_length, window).numpy()[0] im_data = np.fliplr(np.log(x_mag)).T del window return im_data @torch.no_grad() def _get_images(n_samples=10): model.eval() im_data_real = [] im_data_fake = [] for n in range(n_samples): offs = int(len(val_data) * n / n_samples) x = torch.from_numpy( np.array(val_data[offs:offs+block_size])[None].astype(np.float32) / np.iinfo(np.int16).max ) x_cycle, _ = model(x[None]) x_cycle = x_cycle[0] im_data_real.append(_get_im(x)) im_data_fake.append(_get_im(x_cycle)) model.train() for name, module in model.named_modules(): assert module.training # make sure we can undo everything return im_data_real, im_data_fake # learning rate decay scheduler (cosine with warmup) def get_lr(it): # 1) linear warmup for warmup_iters steps if it < warmup_iters: return learning_rate * it / warmup_iters # 2) if it > lr_decay_iters, return min learning rate if it > lr_decay_iters: return min_lr # 3) in between, use cosine decay down to min learning rate decay_ratio = (it - warmup_iters) / (lr_decay_iters - warmup_iters) assert 0 <= decay_ratio <= 1 coeff = 0.5 * (1.0 + math.cos(math.pi * decay_ratio)) # coeff ranges 0..1 return min_lr + coeff * (learning_rate - min_lr) # training loop X, Y = get_batch("train", is_training=True) # fetch the very first batch t0 = time.time() t00 = time.time() 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 raw_model_adv = model_adv.module if ddp else model_adv mfu = 0 tokens_per_s = 0 running_loss = [] running_loss_adv = [] # with torch.autograd.set_detect_anomaly(True): while True: # determine and set the learning rate for this iteration lr = get_lr(iter_num) if decay_lr else learning_rate for param_group in optimizer.param_groups: param_group["lr"] = lr for param_group in optimizer_adv.param_groups: param_group["lr"] = lr # evaluate the loss on train/val sets and write checkpoints if iter_num % eval_interval == 0: 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) if master_process: print( f"loss estimation took {estimation_time:.1f} seconds." f" ({eval_time_pct:.1f}% of loop)" ) print( f"step {iter_num}: train loss {losses['train/loss']:.4f}," f" val loss {losses['val/loss']:.4f}" ) if wandb_log: log_dict = { "iter": iter_num, "lr": lr, "mfu": mfu, # convert to percentage "tok/s": tokens_per_s, } for k, v in losses.items(): log_dict[k] = v im_data_real, im_data_fake = _get_images() for n, im_data in enumerate(im_data_real): log_dict[f"image/real_{n}"] = wandb.Image(im_data) for n, im_data in enumerate(im_data_fake): log_dict[f"image/fake_{n}"] = wandb.Image(im_data) wandb.log(log_dict) if losses["val/loss"] < best_val_loss or always_save_checkpoint: if iter_num > 0: checkpoint = { "model": raw_model.state_dict(), "model_adv": raw_model_adv.state_dict(), "optimizer": optimizer.state_dict(), "optimizer_adv": optimizer_adv.state_dict(), "iter_num": iter_num, "best_val_loss": losses["val/loss"], "config": config, } print(f"saving checkpoint to {out_dir}") if losses["val/loss"] < best_val_loss: torch.save(checkpoint, os.path.join(out_dir, "best_ckpt.pt")) if always_save_checkpoint: torch.save(checkpoint, os.path.join(out_dir, "last_ckpt.pt")) reduced_checkpoint = { k: v for k, v in checkpoint.items() if k in ["model", "config", "best_val_loss"] } torch.save( reduced_checkpoint, os.path.join(out_dir, "last_ckpt_infer.pt") ) if losses["val/loss"] < best_val_loss: best_val_loss = losses["val/loss"] if iter_num == 0 and eval_only: break # forward backward update, with optional gradient accumulation to simulate larger batch size # and using the GradScaler if data type is float16 if ddp: # in DDP training we only need to sync gradients at the last micro step. # the official way to do this is with model.no_sync() context manager, but # I really dislike that this bloats the code and forces us to repeat code # looking at the source of that context manager, it just toggles this variable model.require_backward_grad_sync = True model_adv.require_backward_grad_sync = True with ctx: for update_type in ["generator", "discriminator"]: model.zero_grad() model_adv.zero_grad() y_pred, loss_dict = model(X, Y, randomize_stft=randomize_stft) if update_type == "generator": # update generator # reconstruction loss # loss_rec = reconstruction_loss(Y, y_pred) loss_rec = loss_dict["sc_loss"] + loss_dict["mag_loss"] if skip_quantization: loss_comm = loss_dict["comm_loss"] assert(loss_comm is None) else: loss_comm = loss_dict["comm_loss"] # generative loss y_adv_real, fmap_real = model_adv(Y) y_adv_fake, fmap_fake = model_adv(y_pred) loss_gen_h, loss_gen_f = gen_loss(fmap_real, fmap_fake, y_adv_real, y_adv_fake) if not disc_only: loss = loss_rec * loss_factor_map["rec"] if loss_comm is not None: loss = loss + loss_comm * loss_factor_map["comm"] if not gen_only and iter_num >= n_steps_adv_start: loss = ( loss + loss_gen_h * loss_factor_map["gen_h"] + loss_gen_f * loss_factor_map["gen_f"] ) loss.backward() for _, param in model.named_parameters(): if torch.isnan(param.grad).any() or torch.isinf(param.grad).any(): print("nan found, setting to 0") param.grad = torch.nan_to_num(param.grad, nan=0.0, posinf=0.0, neginf=0.0) if grad_clip != 0.0: grad_norm = torch.nn.utils.clip_grad_norm_( model.parameters(), grad_clip, error_if_nonfinite=True ) else: grad_norm = 0 optimizer.step() else: # update discriminator if iter_num < n_steps_adv_start or iter_num % 2 == 0: # only update every other step loss_adv = None grad_norm_adv = None continue y_adv_real, _ = model_adv(Y) y_adv_fake, _ = model_adv(y_pred.detach()) loss_adv = disc_loss(y_adv_real, y_adv_fake) if not gen_only: (loss_adv * loss_factor_map["disc"]).backward() for _, param in model_adv.named_parameters(): if torch.isnan(param.grad).any() or torch.isinf(param.grad).any(): print("nan found, setting to 0") param.grad = torch.nan_to_num( param.grad, nan=0.0, posinf=0.0, neginf=0.0 ) if grad_clip != 0.0: grad_norm_adv = torch.nn.utils.clip_grad_norm_( model_adv.parameters(), grad_clip, error_if_nonfinite=True ) else: grad_norm_adv = 0 optimizer_adv.step() optimizer.zero_grad(set_to_none=True) optimizer_adv.zero_grad(set_to_none=True) running_loss.append(loss_rec.item()) if loss_adv is not None: running_loss_adv.append(loss_adv.item()) # immediately async prefetch next batch while model is doing the forward pass on the GPU X, Y = get_batch("train", is_training=True) if master_process and wandb_log: log_dict = { "misc/loss_recon": loss_rec.item(), "misc/loss_gen_h": loss_gen_h.item(), "misc/loss_gen_f": loss_gen_f.item(), "misc/grad_norm": grad_norm.item(), } if grad_norm_adv is not None: log_dict["misc/grad_norm_adv"] = grad_norm_adv.item() if loss_adv is not None: log_dict["misc/loss_adv"] = loss_adv.item() if loss_comm is not None: log_dict["misc/loss_comm"] = loss_comm.item() wandb.log(log_dict) # timing and logging t1 = time.time() dt = t1 - t0 t0 = t1 if iter_num % log_interval == 0 and master_process: if local_iter_num >= 5: # let the training loop settle a bit mfu = raw_model.estimate_mfu(batch_size * gradient_accumulation_steps, dt) tokens_per_s = ( world_size * batch_size * block_size * gradient_accumulation_steps / dt ) avg_loss = np.mean(running_loss) avg_loss_adv = np.mean(running_loss_adv) running_loss = [] running_loss_adv = [] print( f"iter {iter_num}: loss_recon {avg_loss:.3f}, loss_adv {avg_loss_adv:.3f}," f" step_time {dt*1000:.1f}ms, mfu {mfu:.1f}," f" throughput {tokens_per_s/1e3:,.0f}k tok/s" ) iter_num += 1 local_iter_num += 1 # termination conditions if iter_num > max_iters: break if ddp: destroy_process_group()