import os import time import wandb import math import torch import random import auraloss import torchaudio import numpy as np import pyloudnorm as pyln import scipy.signal as signal from tqdm import tqdm from suno_utils.utils.text import read_jsonl from typing import Tuple import torch.distributed as dist import torch.multiprocessing as mp from torch.nn.parallel import DistributedDataParallel as DDP def biquad( gain_db: float, cutoff_freq: float, q_factor: float, sample_rate: float, filter_type: str, ) -> Tuple[np.ndarray, np.ndarray]: """Use design parameters to generate coefficients for a specific filter type.""" A = 10 ** (gain_db / 40.0) w0 = 2.0 * np.pi * (cutoff_freq / sample_rate) alpha = np.sin(w0) / (2.0 * q_factor) cos_w0 = np.cos(w0) sqrt_A = np.sqrt(A) if filter_type == "high_shelf": b0 = A * ((A + 1) + (A - 1) * cos_w0 + 2 * sqrt_A * alpha) b1 = -2 * A * ((A - 1) + (A + 1) * cos_w0) b2 = A * ((A + 1) + (A - 1) * cos_w0 - 2 * sqrt_A * alpha) a0 = (A + 1) - (A - 1) * cos_w0 + 2 * sqrt_A * alpha a1 = 2 * ((A - 1) - (A + 1) * cos_w0) a2 = (A + 1) - (A - 1) * cos_w0 - 2 * sqrt_A * alpha elif filter_type == "low_shelf": b0 = A * ((A + 1) - (A - 1) * cos_w0 + 2 * sqrt_A * alpha) b1 = 2 * A * ((A - 1) - (A + 1) * cos_w0) b2 = A * ((A + 1) - (A - 1) * cos_w0 - 2 * sqrt_A * alpha) a0 = (A + 1) + (A - 1) * cos_w0 + 2 * sqrt_A * alpha a1 = -2 * ((A - 1) + (A + 1) * cos_w0) a2 = (A + 1) + (A - 1) * cos_w0 - 2 * sqrt_A * alpha elif filter_type == "peaking": b0 = 1 + alpha * A b1 = -2 * cos_w0 b2 = 1 - alpha * A a0 = 1 + alpha / A a1 = -2 * cos_w0 a2 = 1 - alpha / A b = np.array([b0, b1, b2]) / a0 a = np.array([1.0, a1 / a0, a2 / a0]) return b, a def apply_stereo_to_mono(audio: torch.Tensor, sample_rate: float): return audio.mean(dim=0, keepdims=True).repeat(2, 1) def apply_channel_imbalance( audio: torch.Tensor, sample_rate: float, imbalance: float = 0.0 ): if not -1 <= imbalance <= 1 or audio.shape[-2] != 2: raise ValueError("Invalid input") out = audio.clone() l_gain, r_gain = (1.0 - imbalance, 1.0) if imbalance > 0 else (1.0, 1.0 + imbalance) out[0, :], out[1, :] = out[0, :] * l_gain, out[1, :] * r_gain return out def apply_highpass(audio: torch.Tensor, sample_rate: float, cutoff_hz: float = 1000.0): return torchaudio.functional.highpass_biquad(audio, sample_rate, cutoff_hz) def apply_lowpass(audio: torch.Tensor, sample_rate: float, cutoff_hz: float = 1000.0): return torchaudio.functional.lowpass_biquad(audio, sample_rate, cutoff_hz) def apply_noise( audio: torch.Tensor, sample_rate: float, gain_db: float = 0.0, noise_type: str = "white", ): gain_lin = 10 ** (gain_db / 20.0) noise = torch.randn_like(audio) if noise_type == "white": return audio + gain_lin * noise elif noise_type == "pink": b = torch.tensor([0.049922035, -0.095993537, 0.050612699, -0.004408786]) a = torch.tensor([1, -2.494956002, 2.017265875, -0.522189400]) noise = torchaudio.functional.filtfilt(noise, a, b) noise /= noise.abs().max() return audio + gain_lin * noise else: raise ValueError(f"Invalid noise type: {noise_type}") def apply_shelving_filter( audio: torch.Tensor, sample_rate: float, gain_db: float, cutoff_freq: float, q_factor: float, filter_type: str, ): # convert x to numpy audio = audio.numpy() b, a = biquad( gain_db, cutoff_freq, q_factor, sample_rate, filter_type, ) x = signal.lfilter(b, a, audio).astype(np.float32) return torch.from_numpy(x) # randomized corrputions def apply_random_noise(audio: torch.Tensor, sample_rate: float): noise_type = random.choice(["white", "pink"]) if noise_type == "white": noise_gain = random.uniform(-96, -48) else: noise_gain = random.uniform(-48, -12) return apply_noise(audio, sample_rate, noise_gain, noise_type) def apply_random_stereo_to_mono(audio: torch.Tensor, sample_rate: float): return apply_stereo_to_mono(audio, sample_rate) def apply_random_channel_imbalance(audio: torch.Tensor, sample_rate: float): imbalance = random.uniform(-1.0, 1.0) return apply_channel_imbalance(audio, sample_rate, imbalance) def apply_random_filter(audio: torch.Tensor, sample_rate: float): filter_type = random.choice(["highpass", "lowpass", "high_shelf", "low_shelf"]) if filter_type == "highpass": cutoff_freq = random.uniform(20, 4000) return apply_highpass(audio, sample_rate, cutoff_freq) elif filter_type == "lowpass": cutoff_freq = random.uniform(1000, 16000) return apply_lowpass(audio, sample_rate, cutoff_freq) else: gain_db = random.uniform(-12, 12) if filter_type == "high_shelf": cutoff_freq = random.uniform(6000, 20000) else: cutoff_freq = random.uniform(20, 2000) q_factor = random.uniform(0.1, 10.0) return apply_shelving_filter( audio, sample_rate, gain_db, cutoff_freq, q_factor, filter_type ) def corrupt(waveform_tensor, sample_rate): """ Apply a random number of corruptions (at least one) to the input waveform. """ # List of available random corruption functions corruption_fns = [ apply_random_noise, apply_random_stereo_to_mono, apply_random_channel_imbalance, apply_random_filter, ] n_corr = random.randint(1, len(corruption_fns)) # at least one selected = random.sample(corruption_fns, n_corr) print(selected) out = waveform_tensor.clone() for fn in selected: out = fn(out, sample_rate) return torch.tanh(out) # audio pair dataset # takes in a metas files # metas has the following structure: # { # "id": "", # "input_local_path": "", # "target_local_path": "", # } class AudioPairDataset(torch.utils.data.Dataset): def __init__(self, metas_filepath, chunk_size_samples=262144, buffer_size=1000): self.metas_filepath = metas_filepath self.chunk_size_samples = chunk_size_samples self.buffer_size = buffer_size metas = read_jsonl(metas_filepath) print(f"Loaded {len(metas)} metas from {metas_filepath}") self.metas = metas self.buffer = [] self.items_since_last_reload = self.buffer_size def __len__(self): return len(self.metas) def _reload_buffer(self): self.buffer = [] rand_indices = np.random.permutation(len(self.metas)) finished = False if int(os.environ["LOCAL_RANK"]) == 0: print(f"Reloading buffer with {len(rand_indices)} items") for idx in rand_indices: meta = self.metas[idx] # input_audio, sr = torchaudio.load(meta["input_filepath"]) target_audio, sr = torchaudio.load(meta["target_filepath"]) # make inut audio lowpass filtered # input_audio = torchaudio.functional.lowpass_filter(input_audio, sr, 4000) # compute number of chunks in the input audio num_chunks = target_audio.shape[1] // self.chunk_size_samples for i in range(num_chunks): start_sample = i * self.chunk_size_samples end_sample = start_sample + self.chunk_size_samples target_chunk = target_audio[:, start_sample:end_sample] # apply random corruptions corrupted_chunk = corrupt(target_chunk, sr) self.buffer.append((corrupted_chunk, target_chunk)) if len(self.buffer) >= self.buffer_size: finished = True break if finished: break pbar.set_postfix(buffer_size=len(self.buffer)) def __getitem__(self, idx): if self.items_since_last_reload == self.buffer_size: self._reload_buffer() self.items_since_last_reload = 0 input_audio_chunk, target_audio_chunk = self.buffer[ self.items_since_last_reload ] self.items_since_last_reload += 1 return input_audio_chunk, target_audio_chunk class AudioPatcher: def __init__(self, patch_size): self.patch_size = patch_size def to_patches(self, audio): # audio: (bs, 2, seq_len) bs, channels, seq_len = audio.shape assert ( seq_len % self.patch_size == 0 ), "Sequence length must be divisible by patch size" num_patches = seq_len // self.patch_size patches = audio.view( bs, channels, num_patches, self.patch_size ) # (bs, 2, num_patches, patch_size) return patches def flatten_patches(self, patches): # patches: (bs, 2, num_patches, patch_size) bs, channels, num_patches, patch_size = patches.shape flattened = patches.permute(0, 2, 1, 3).reshape( bs, num_patches * channels, patch_size ) # (bs, num_patches * 2, patch_size) return flattened def unflatten_patches(self, flattened, original_channels=2): # flattened: (bs, num_patches * channels, patch_size) bs, total_patches, patch_size = flattened.shape num_patches = total_patches // original_channels patches = flattened.view( bs, num_patches, original_channels, patch_size ).permute(0, 2, 1, 3) # (bs, 2, num_patches, patch_size) return patches def reconstruct_audio(self, patches): # patches: (bs, 2, num_patches, patch_size) bs, channels, num_patches, patch_size = patches.shape audio = patches.reshape(bs, channels, num_patches * patch_size) return audio class SinusoidalPositionalEncoding(torch.nn.Module): def __init__(self, hidden_dim): super().__init__() position = torch.arange(10000).unsqueeze(1) div_term = torch.exp( torch.arange(0, hidden_dim, 2) * -(math.log(10000.0) / hidden_dim) ) pe = torch.zeros(10000, hidden_dim) pe[:, 0::2] = torch.sin(position * div_term) pe[:, 1::2] = torch.cos(position * div_term) self.register_buffer("pe", pe) def forward(self, x): # x: (batch, num_patches, hidden_dim) return x + self.pe[: x.size(1)] class RefinerTransformer(torch.nn.Module): def __init__( self, hidden_dim=1024, num_heads=8, num_transformer_layers=12, dropout=0.1, patch_size=1024, ): super().__init__() self.patch_size = patch_size self.patcher = AudioPatcher(patch_size) self.pos_embed = SinusoidalPositionalEncoding(hidden_dim) self.input_layer = torch.nn.Linear(patch_size, hidden_dim) # Transformer encoder encoder_layer = torch.nn.TransformerEncoderLayer( d_model=hidden_dim, nhead=num_heads, dim_feedforward=hidden_dim * 4, dropout=dropout, ) self.transformer = torch.nn.TransformerEncoder( encoder_layer, num_transformer_layers ) self.output_layer = torch.nn.Linear(hidden_dim, patch_size) def forward(self, input_audio): # input_audio shape: (batch_size, 2, seq_len) # first split into patches, which becomes (batch_size, 2, num_patches, patch_size) patches = self.patcher.to_patches(input_audio) flattened = self.patcher.flatten_patches(patches) x = self.input_layer(flattened) # (batch_size, num_patches, hidden_dim) x = self.pos_embed(x) # (batch_size, num_patches, hidden_dim) x = self.transformer(x) # (batch_size, num_patches, hidden_dim) x = self.output_layer(x) # (batch_size, num_patches, 1) x = self.patcher.unflatten_patches(x) # (batch_size, 2, seq_len) # fold the patches back together output_audio = self.patcher.reconstruct_audio(x) return input_audio + output_audio class HDemucsRefiner(torch.nn.Module): def __init__(self, **kwargs): super().__init__() self.hdemucs = torchaudio.models.HDemucs(sources=["output"], **kwargs) def forward(self, input_audio): return self.hdemucs.forward(input_audio).sum(dim=1) class DummyRefiner(torch.nn.Module): def __init__(self): super().__init__() self.linear = torch.nn.Conv1d(2, 2, kernel_size=3, padding=1) def forward(self, x): return self.linear(x) def save_checkpoint( model, optimizer, run_config, global_step, checkpoint_dir, ): checkpoint = { "model": model.state_dict(), "optimizer": optimizer.state_dict(), "run_config": run_config, "global_step": global_step, } checkpoint_path = os.path.join(checkpoint_dir, "last_ckpt.pt") torch.save(checkpoint, checkpoint_path) print(f"Saved checkpoint to {checkpoint_path}") def set_seed(base_seed=42): rank = int(os.environ.get("LOCAL_RANK", 0)) seed = base_seed + rank random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) def setup(): dist.init_process_group(backend="nccl") torch.cuda.set_device(int(os.environ["LOCAL_RANK"])) def reduce_tensor(tensor, average=True): """ Reduces a tensor from all processes so rank 0 can log the average value. """ rt = tensor.clone() dist.all_reduce(rt, op=dist.ReduceOp.SUM) if average: rt /= dist.get_world_size() return rt def train_step( batch, model, optimizer, time_loss_fn, freq_loss_fn, run_config, ): input_audio, target_audio = batch input_audio = input_audio.cuda() target_audio = target_audio.cuda() output_audio = model(input_audio) time_loss = time_loss_fn(output_audio, target_audio) freq_loss = freq_loss_fn(output_audio, target_audio) loss = ( run_config["training"]["time_loss_weight"] * time_loss + run_config["training"]["freq_loss_weight"] * freq_loss ) return time_loss, freq_loss, loss if __name__ == "__main__": run_start_time = time.strftime("%Y-%m-%d_%H-%M-%S") checkpoint_dir = f"/app/suno/christian/checkpoints/refiner-v1/{run_start_time}_s{random.randint(0, 9999)}" os.makedirs(checkpoint_dir, exist_ok=False) torch.set_float32_matmul_precision("medium") setup() set_seed() local_rank = int(os.environ["LOCAL_RANK"]) world_size = int(os.environ["WORLD_SIZE"]) device = torch.device("cuda", local_rank) torch.cuda.set_device(local_rank) run_config = { "training": { "max_steps": 100_000, "run_name": "refiner-10s", "project_name": "refiner", "lr": 1e-4, "grad_clip_norm": 1.0, "preload_ckpt": None, "preload_optimizer": False, "warmup_steps": 500, "ckpt_every": 1000, "val_every": 1000, "time_loss_weight": 2.0, "freq_loss_weight": 1.0, }, "model": { "hidden_dim": 2048, "num_heads": 16, "num_transformer_layers": 12, "dropout": 0.1, }, "dataset": { "train_metas_filepath": "/mnt/localdisk/tmp_cjs/v2-infill-data-v1/metas_tr.jsonl", "val_metas_filepath": "/mnt/localdisk/tmp_cjs/v2-infill-data-v1/metas_val.jsonl", "batch_size": 16, "num_workers": 4, "chunk_size_samples": 524288, "buffer_size": 1000, }, } # setup the data train_dataset = AudioPairDataset( run_config["dataset"]["train_metas_filepath"], run_config["dataset"]["chunk_size_samples"], run_config["dataset"]["buffer_size"], ) val_dataset = AudioPairDataset( run_config["dataset"]["val_metas_filepath"], run_config["dataset"]["chunk_size_samples"], run_config["dataset"]["buffer_size"], ) train_sampler = torch.utils.data.distributed.DistributedSampler( train_dataset, num_replicas=world_size, rank=local_rank, shuffle=True ) train_loader = torch.utils.data.DataLoader( train_dataset, batch_size=run_config["dataset"]["batch_size"], sampler=train_sampler, num_workers=run_config["dataset"]["num_workers"], pin_memory=True, ) val_loader = torch.utils.data.DataLoader( val_dataset, batch_size=run_config["dataset"]["batch_size"], num_workers=run_config["dataset"]["num_workers"], shuffle=False, pin_memory=True, ) # setup the model # model = Refiner(**run_config["model"]) model = HDemucsRefiner() model = model.cuda() print(f"Model has {sum(p.numel() for p in model.parameters()):,} parameters") # wrap the model in DDP model = DDP(model, device_ids=[local_rank], find_unused_parameters=True) # setup the optimizer optimizer = torch.optim.AdamW(model.parameters(), lr=run_config["training"]["lr"]) # setup the scheduler warmup_scheduler = torch.optim.lr_scheduler.LinearLR( optimizer, start_factor=0.001, end_factor=1.0, total_iters=run_config["training"]["warmup_steps"], ) cosine_scheduler = torch.optim.lr_scheduler.CosineAnnealingLR( optimizer, run_config["training"]["max_steps"] - run_config["training"]["warmup_steps"], ) scheduler = torch.optim.lr_scheduler.ChainedScheduler( [warmup_scheduler, cosine_scheduler] ) # setup losses time_loss_fn = torch.nn.L1Loss() # time_loss_fn = auraloss.time.SI_SDR() freq_loss_fn = auraloss.freq.SumAndDifferenceSTFTLoss( fft_sizes=[1024, 2048, 4096, 8192], hop_sizes=[128, 256, 512, 1024], win_lengths=[1024, 2048, 4096, 8192], ) global_step = 0 # setup wandb if local_rank == 0: wandb.init( project="refiner", name=run_config["training"]["run_name"], config=run_config, ) wandb.config.update( {"checkpoint_dir": checkpoint_dir, "run_config": run_config} ) while global_step < run_config["training"]["max_steps"]: print(f"[{local_rank}] Starting training loop") torch.distributed.barrier() train_sampler.set_epoch(global_step) pbar = tqdm(train_loader, total=len(train_loader)) # Track time for iterations per second calculation start_time = time.time() iter_times = [] for batch in pbar: optimizer.zero_grad() loss = torch.tensor(0.0, device=device) time_loss, freq_loss, loss = train_step( batch, model, optimizer, time_loss_fn, freq_loss_fn, run_config, ) loss.backward() torch.nn.utils.clip_grad_norm_( model.parameters(), run_config["training"]["grad_clip_norm"] ) grad_norm = torch.norm( torch.stack( [ torch.norm(p.grad) for p in model.parameters() if p.grad is not None ] ) ) optimizer.step() scheduler.step() global_step += 1 if local_rank == 0: # Calculate iterations per second iter_time = time.time() - start_time iter_times.append(iter_time) if len(iter_times) > 100: # Keep a moving window iter_times.pop(0) ips = 1.0 / (sum(iter_times) / len(iter_times)) start_time = time.time() pbar.set_postfix( loss=loss.item(), grad_norm=grad_norm.item(), ips=f"{ips:.2f}" ) if ( global_step % run_config["training"]["ckpt_every"] == 0 and local_rank == 0 ): save_checkpoint( model, optimizer, run_config, global_step, checkpoint_dir ) if local_rank == 0: wandb.log( { "train/loss": loss.item(), "train/grad_norm": grad_norm.item(), "train/time_loss": time_loss.item(), "train/freq_loss": freq_loss.item(), "train/global_step": global_step, "train/lr": optimizer.param_groups[0]["lr"], "train/iterations_per_second": ips, }, )