import os import sys import warnings from dataclasses import dataclass from pathlib import Path import random import argbind import auraloss import torch from audiotools import AudioSignal from audiotools import ml from audiotools.core import util from audiotools.data import transforms from audiotools.ml.decorators import timer from audiotools.ml.decorators import Tracker from audiotools.ml.decorators import when import numpy as np from torch.utils.tensorboard import SummaryWriter import torchaudio.functional as aF from dac.model.dac2 import DAC as DAC_import from dac.model.discriminator2 import Discriminator as Discriminator_import from dac.nn import loss as loss_import warnings.filterwarnings("ignore", category=UserWarning) # Enable cudnn autotuner to speed up training # (can be altered by the funcs.seed function) torch.backends.cudnn.benchmark = bool(int(os.getenv("CUDNN_BENCHMARK", 1))) # Uncomment to trade memory for speed. # Optimizers AdamW = argbind.bind(torch.optim.AdamW, "generator", "discriminator") Accelerator = argbind.bind(ml.Accelerator, without_prefix=True) @argbind.bind("generator", "discriminator") def ExponentialLR(optimizer, gamma: float = 1.0): return torch.optim.lr_scheduler.ExponentialLR(optimizer, gamma) # Models DAC = argbind.bind(DAC_import) Discriminator = argbind.bind(Discriminator_import) # Transforms filter_fn = lambda fn: hasattr(fn, "transform") and fn.__qualname__ not in [ "BaseTransform", "Compose", "Choose", ] tfm = argbind.bind_module(transforms, "train", "val", filter_fn=filter_fn) # Loss filter_fn = lambda fn: hasattr(fn, "forward") and "Loss" in fn.__name__ losses = argbind.bind_module(loss_import, filter_fn=filter_fn) @dataclass class State: generator: DAC optimizer_g: AdamW scheduler_g: ExponentialLR discriminator: Discriminator optimizer_d: AdamW scheduler_d: ExponentialLR mel_loss: auraloss.freq.SumAndDifferenceSTFTLoss gan_loss: losses.GANLoss tracker: Tracker sample_rate: int n_samples_val: int @argbind.bind(without_prefix=True) def load( args, accel: ml.Accelerator, tracker: Tracker, save_path: str, resume: bool = False, tag: str = "latest", load_weights: bool = False, ): generator, g_extra = None, {} discriminator, d_extra = None, {} if resume: kwargs = { "folder": f"{save_path}/{tag}", "map_location": "cpu", "package": not load_weights, } tracker.print(f"Resuming from {str(Path('.').absolute())}/{kwargs['folder']}") if (Path(kwargs["folder"]) / "dac").exists(): generator, g_extra = DAC.load_from_folder(**kwargs) if (Path(kwargs["folder"]) / "discriminator").exists(): discriminator, d_extra = Discriminator.load_from_folder(**kwargs) generator = DAC() if generator is None else generator discriminator = Discriminator() if discriminator is None else discriminator # tracker.print(generator) # tracker.print(discriminator) generator = accel.prepare_model(generator) discriminator = accel.prepare_model(discriminator) with argbind.scope(args, "generator"): optimizer_g = AdamW(generator.parameters(), use_zero=accel.use_ddp) scheduler_g = ExponentialLR(optimizer_g) with argbind.scope(args, "discriminator"): optimizer_d = AdamW(discriminator.parameters(), use_zero=accel.use_ddp) scheduler_d = ExponentialLR(optimizer_d) if "optimizer.pth" in g_extra: optimizer_g.load_state_dict(g_extra["optimizer.pth"]) if "scheduler.pth" in g_extra: scheduler_g.load_state_dict(g_extra["scheduler.pth"]) if "tracker.pth" in g_extra: tracker.load_state_dict(g_extra["tracker.pth"]) if "optimizer.pth" in d_extra: optimizer_d.load_state_dict(d_extra["optimizer.pth"]) if "scheduler.pth" in d_extra: scheduler_d.load_state_dict(d_extra["scheduler.pth"]) mel_loss = auraloss.freq.SumAndDifferenceSTFTLoss( sample_rate=args["DAC.sample_rate"], fft_sizes=[2048, 1024, 512, 256, 128, 64, 32], hop_sizes=[512, 256, 128, 64, 32, 16, 8], win_lengths=[2048, 1024, 512, 256, 128, 64, 32], perceptual_weighting=True, ) gan_loss = losses.GANLoss(discriminator) return State( generator=generator, optimizer_g=optimizer_g, scheduler_g=scheduler_g, discriminator=discriminator, optimizer_d=optimizer_d, scheduler_d=scheduler_d, mel_loss=mel_loss, gan_loss=gan_loss, tracker=tracker, sample_rate=args["DAC.sample_rate"], n_samples_val=args["val/n_examples"], ) @timer() @torch.no_grad() def val_loop(batch, state, accel): state.generator.eval() batch = batch.to(accel.device) recons = state.generator(batch, state.sample_rate)["audio"] return { "loss": state.mel_loss(recons, batch), "mel/loss": state.mel_loss(recons, batch), } @timer() def train_loop(state, batch, accel, lambdas): state.generator.train() state.discriminator.train() output = {} batch = batch.to(accel.device) signal = AudioSignal(batch, state.sample_rate) global_batch_size = signal.batch_size * accel.world_size with accel.autocast(): out = state.generator(batch, state.sample_rate) recons = AudioSignal(out["audio"], state.sample_rate) commitment_loss = out["vq/commitment_loss"] codebook_loss = out["vq/codebook_loss"] with accel.autocast(): output["adv/disc_loss"] = state.gan_loss.discriminator_loss(recons, signal) state.optimizer_d.zero_grad() accel.backward(output["adv/disc_loss"]) accel.scaler.unscale_(state.optimizer_d) output["other/grad_norm_d"] = torch.nn.utils.clip_grad_norm_( state.discriminator.parameters(), 10.0 ) accel.step(state.optimizer_d) state.scheduler_d.step() b, _, t = signal.audio_data.shape with accel.autocast(): output["mel/loss"] = state.mel_loss(recons.audio_data, signal.audio_data) ( output["adv/gen_loss"], output["adv/feat_loss"], ) = state.gan_loss.generator_loss(recons, signal) output["vq/commitment_loss"] = commitment_loss output["vq/codebook_loss"] = codebook_loss output["loss"] = sum([v * output[k] for k, v in lambdas.items() if k in output]) state.optimizer_g.zero_grad() accel.backward(output["loss"]) accel.scaler.unscale_(state.optimizer_g) output["other/grad_norm"] = torch.nn.utils.clip_grad_norm_( state.generator.parameters(), 1e3 ) accel.step(state.optimizer_g) state.scheduler_g.step() accel.update() output["other/learning_rate"] = state.optimizer_g.param_groups[0]["lr"] output["other/batch_size"] = global_batch_size return {k: v for k, v in sorted(output.items())} def checkpoint(state, save_iters, save_path): metadata = {"logs": state.tracker.history} tags = ["latest"] state.tracker.print(f"Saving to {str(Path('.').absolute())}") if state.tracker.is_best("val", "mel/loss"): state.tracker.print(f"Best generator so far") tags.append("best") if state.tracker.step in save_iters: tags.append(f"{state.tracker.step // 1000}k") for tag in tags: generator_extra = { "optimizer.pth": state.optimizer_g.state_dict(), "scheduler.pth": state.scheduler_g.state_dict(), "tracker.pth": state.tracker.state_dict(), "metadata.pth": metadata, } accel.unwrap(state.generator).metadata = metadata accel.unwrap(state.generator).save_to_folder( f"{save_path}/{tag}", generator_extra, package=False ) discriminator_extra = { "optimizer.pth": state.optimizer_d.state_dict(), "scheduler.pth": state.scheduler_d.state_dict(), } accel.unwrap(state.discriminator).save_to_folder( f"{save_path}/{tag}", discriminator_extra, package=False ) @torch.no_grad() def save_samples(state, val_dataloader, val_idx, writer): state.tracker.print("Saving audio samples to TensorBoard") state.generator.eval() samples = [val_dataloader[idx] for idx in val_idx] batch = torch.stack(samples).to(accel.device) signal = AudioSignal(batch, state.sample_rate) out = state.generator(signal.audio_data, signal.sample_rate) recons = AudioSignal(out["audio"], signal.sample_rate) audio_dict = {"recons": recons} if state.tracker.step == 0: audio_dict["signal"] = signal for k, v in audio_dict.items(): for nb in range(v.batch_size): v[nb].cpu().write_audio_to_tb( f"{k}/sample_{nb}.wav", writer, state.tracker.step ) @torch.no_grad() def save_golden_samples(state, golden_dir, output_dir): state.tracker.print("Saving golden audio samples") state.generator.eval() golden_dir = Path(golden_dir) output_dir = Path(output_dir) golden_files = list(golden_dir.glob("*.wav")) golden_files.sort() for golden_file in golden_files: signal = AudioSignal(golden_file, state.sample_rate, device=accel.device) recons = state.generator(signal.audio_data, signal.sample_rate)["audio"] recons = AudioSignal(recons, signal.sample_rate) recons.cpu().write(output_dir / f"{golden_file.stem}_{state.tracker.step}.wav") def validate(state, val_dataloader, accel): for n, batch in enumerate(val_dataloader): output = val_loop(batch, state, accel) if n == state.n_samples_val - 1: break # Consolidate state dicts if using ZeroRedundancyOptimizer if hasattr(state.optimizer_g, "consolidate_state_dict"): state.optimizer_g.consolidate_state_dict() state.optimizer_d.consolidate_state_dict() return output 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 class CustomDataLoader: def __init__(self, mem_map_path, duration_s=5, batch_size=8, sample_rate=48000, is_val=False): self.data = np.memmap(mem_map_path, dtype=np.int16, mode="r") self.duration_s = duration_s self.batch_size = batch_size self.sample_rate = sample_rate self.n_samples = int(round(self.duration_s*self.sample_rate))*2 self.is_val = is_val # def __len__(self): # return self.length def __getitem__(self, idx): arr = np.array(self.data[idx:idx+self.n_samples].reshape(-1, 2).T) arr = torch.from_numpy(arr.astype(np.float32) / np.iinfo(np.int16).max) return arr def __iter__(self): return self def __next__(self): out = [] for _ in range(self.batch_size): idx = int(random.randint(0, len(self.data) - self.n_samples - 1) / 2) * 2 arr = np.array(self.data[idx:idx+self.n_samples].reshape(-1, 2).T) arr = torch.from_numpy(arr.astype(np.float32) / np.iinfo(np.int16).max) # if not self.is_val and random.random() >= 0.9: # arr = _cycle_sample_rate( # arr, # from_sample_rate=self.sample_rate, # to_sample_rate=random.choice(COMMON_SAMPLE_RATES), # ) out.append(arr) return torch.stack(out) @argbind.bind(without_prefix=True) def train( args, accel: ml.Accelerator, seed: int = 0, save_path: str = "ckpt", num_iters: int = 250000, save_iters: list = [10000, 50000, 100000, 200000], sample_freq: int = 10000, valid_freq: int = 1000, lambdas: dict = { "mel/loss": 100.0, "adv/feat_loss": 2.0, "adv/gen_loss": 1.0, "vq/commitment_loss": 0.25, "vq/codebook_loss": 1.0, }, wandb_log: bool = False, wandb_project: str = "dac_test", wandb_run_name: str = "test", log_interval: int = 100, golden_dir: str = None, output_dir: str = None, ): util.seed(seed) Path(save_path).mkdir(exist_ok=True, parents=True) writer = ( SummaryWriter(log_dir=f"{save_path}/logs") if accel.local_rank == 0 else None ) tracker = Tracker( writer=writer, log_file=f"{save_path}/log.txt", rank=accel.local_rank ) state = load(args, accel, tracker, save_path) # Wrap the functions so that they neatly track in TensorBoard + progress bars # and only run when specific conditions are met. global train_loop, val_loop, validate, save_samples, checkpoint train_loop = tracker.log("train", "value", history=False)( tracker.track("train", num_iters, completed=state.tracker.step)(train_loop) ) val_loop = tracker.track("val", args["val/n_examples"])(val_loop) validate = tracker.log("val", "mean")(validate) # These functions run only on the 0-rank process save_samples = when(lambda: accel.local_rank == 0)(save_samples) checkpoint = when(lambda: accel.local_rank == 0)(checkpoint) ddp = int(os.environ.get("RANK", -1)) != -1 if ddp: seed_offset = torch.distributed.get_rank() else: seed_offset = 0 random.seed(6006 + seed_offset) train_dataloader = CustomDataLoader( args["train/memmap_path"], duration_s=args["train/duration"], batch_size=args["batch_size"], sample_rate=args["DAC.sample_rate"], is_val=False, ) val_dataloader = CustomDataLoader( args["val/memmap_path"], duration_s=args["val/duration"], batch_size=args["batch_size"], sample_rate=args["DAC.sample_rate"], is_val=True, ) master_process = accel.local_rank == 0 if master_process and wandb_log: import wandb wandb.init( project=wandb_project, name=wandb_run_name, config=args, ) with tracker.live: for tracker.step, batch in enumerate(train_dataloader, start=tracker.step): # batch = { # 'idx': tensor([0, 0]), # 'signal': , # 'source_idx': tensor([0, 0]), # 'item_idx': tensor([2, 2]), # 'source': ['../gpt/samples/mini', '../gpt/samples/mini'], # 'path': ['../gpt/samples/mini/obama_short.mp3', '../gpt/samples/mini/obama_short.mp3'], # } train_out = train_loop(state, batch, accel, lambdas) if master_process and wandb_log and tracker.step % log_interval == 0: wandb.log(train_out, step=tracker.step) last_iter = ( tracker.step == num_iters - 1 if num_iters is not None else False ) if tracker.step % sample_freq == 0 or last_iter: # save_samples(state, val_dataloader, args["val/samples_offsets"], writer) if golden_dir is not None and output_dir is not None: save_golden_samples(state, golden_dir, output_dir) if tracker.step % valid_freq == 0 or last_iter: val_out = validate(state, val_dataloader, accel) if master_process and wandb_log: # add val_ to keys to avoid confusion with train metrics val_out = {f"val_{k}": v for k, v in val_out.items()} wandb.log(val_out, step=tracker.step) checkpoint(state, save_iters, save_path) # Reset validation progress bar, print summary since last validation. tracker.done("val", f"Iteration {tracker.step}") if last_iter: break if __name__ == "__main__": args = argbind.parse_args() args["args.debug"] = int(os.getenv("LOCAL_RANK", 0)) == 0 with argbind.scope(args): with Accelerator() as accel: if accel.local_rank != 0: sys.tracebacklimit = 0 train(args, accel)