import os import time import math import glob import wandb import random import torch import torchaudio import itertools import numpy as np import torch.nn as nn import torch.nn.functional as F import torch.distributed as dist from torch.utils.data import DataLoader from tqdm import tqdm from typing import List def mdct(x): N = x.shape[-1] n = torch.arange(N, device=x.device) k = torch.arange(N // 2, device=x.device) arg = ( (math.pi / (2 * N)) * ((2 * n + 1 + N // 2).view(-1, 1)) * ((2 * k + 1).view(1, -1)) ) mdct_matrix = torch.cos(arg) * (2.0 / N) ** 0.5 return torch.matmul(x, mdct_matrix) def imdct(X): half_N = X.shape[-1] N = half_N * 2 n = torch.arange(N, device=X.device) k = torch.arange(half_N, device=X.device) arg = ( (math.pi / (2 * N)) * ((2 * n + 1 + N // 2).view(-1, 1)) * ((2 * k + 1).view(1, -1)) ) imdct_matrix = torch.cos(arg) * (2.0 / N) ** 0.5 return torch.matmul(X, imdct_matrix.T) * 2.0 def audio_to_mdct_frames(audio, frame_size=1920, midside=False): batch, channels, samples = audio.shape hop_size = frame_size // 2 if midside: mid = (audio[:, 0, :] + audio[:, 1, :]) / 2.0 side = (audio[:, 0, :] - audio[:, 1, :]) / 2.0 audio = torch.stack([mid, side], dim=1) n_frames = (samples - frame_size) // hop_size + 1 frames = [] for i in range(n_frames): start = i * hop_size frame = audio[:, :, start : start + frame_size] if frame.shape[-1] == frame_size: frames.append(frame) frames = torch.stack(frames, dim=2) window = torch.sin( torch.pi / frame_size * (torch.arange(frame_size, device=audio.device) + 0.5) ) frames = frames * window.view(1, 1, 1, -1) shape = frames.shape frames_reshaped = frames.reshape(-1, frame_size) mdct_coeffs = mdct(frames_reshaped) return mdct_coeffs.reshape(shape[0], shape[1], shape[2], -1) def mdct_frames_to_audio_slow(mdct_coeffs, frame_size=1920, midside=False): batch, channels, n_frames, half_frame_size = mdct_coeffs.shape hop_size = frame_size // 2 shape = mdct_coeffs.shape coeffs_reshaped = mdct_coeffs.reshape(-1, half_frame_size) frames = imdct(coeffs_reshaped) frames = frames.reshape(shape[0], shape[1], shape[2], -1) window = torch.sin( torch.pi / frame_size * (torch.arange(frame_size, device=mdct_coeffs.device) + 0.5) ) frames = frames * window.view(1, 1, 1, -1) total_samples = (n_frames - 1) * hop_size + frame_size output = torch.zeros(batch, channels, total_samples, device=mdct_coeffs.device) for i in range(n_frames): start = i * hop_size output[:, :, start : start + frame_size] += frames[:, :, i] if midside: mid = output[:, 0, :] side = output[:, 1, :] left = mid + side right = mid - side output = torch.stack([left, right], dim=1) return output def mdct_frames_to_audio(mdct_coeffs, frame_size=1920, midside=False): batch, channels, n_frames, half_frame_size = mdct_coeffs.shape hop_size = frame_size // 2 # Reshape and apply IMDCT shape = mdct_coeffs.shape coeffs_reshaped = mdct_coeffs.reshape(-1, half_frame_size) frames = imdct(coeffs_reshaped) frames = frames.reshape(shape[0], shape[1], shape[2], -1) # Apply window window = torch.sin( torch.pi / frame_size * (torch.arange(frame_size, device=mdct_coeffs.device) + 0.5) ) frames = frames * window.view(1, 1, 1, -1) # Calculate total samples and create output shape total_samples = (n_frames - 1) * hop_size + frame_size # Reshape frames to prepare for folding frames = frames.permute(0, 1, 3, 2) # [batch, channels, frame_size, n_frames] frames = frames.reshape(batch * channels, frame_size, n_frames) # Use fold operation to overlap-add frames output = torch.nn.functional.fold( frames, output_size=(1, total_samples), kernel_size=(1, frame_size), stride=(1, hop_size), ) # Reshape output to expected dimensions output = output.view(batch, channels, total_samples) if midside: mid = output[:, 0, :] side = output[:, 1, :] left = mid + side right = mid - side output = torch.stack([left, right], dim=1) return output class WaveformMDCTVAE(nn.Module): def __init__( self, frame_size: int = 1920, latent_dim: int = 256, hidden_dims: list = None, dropout: float = 0.1, midside: bool = False, ): super().__init__() self.frame_size = frame_size self.n_coeffs = frame_size // 2 self.latent_dim = latent_dim self.midside = midside if hidden_dims is None: hidden_dims = [512, 256] # Encoder layers modules = [] input_dim = 2 * self.n_coeffs for h_dim in hidden_dims: modules.append( nn.Sequential( nn.Linear(input_dim, h_dim), # nn.LayerNorm(h_dim), nn.LeakyReLU(), nn.Dropout(dropout), nn.Linear(h_dim, h_dim), ) ) input_dim = h_dim self.encoder = nn.Sequential(*modules) self.fc_mu = nn.Linear(hidden_dims[-1], latent_dim) self.fc_var = nn.Linear(hidden_dims[-1], latent_dim) # Decoder layers modules = [] hidden_dims.reverse() self.decoder_input = nn.Sequential( nn.Linear(latent_dim, hidden_dims[0]), nn.LayerNorm(hidden_dims[0]), nn.LeakyReLU(), nn.Dropout(dropout), ) for i in range(len(hidden_dims) - 1): modules.append( nn.Sequential( nn.Linear(hidden_dims[i], hidden_dims[i + 1]), # nn.LayerNorm(hidden_dims[i + 1]), nn.LeakyReLU(), nn.Dropout(dropout), nn.Linear(hidden_dims[i + 1], hidden_dims[i + 1]), ) ) self.decoder = nn.Sequential(*modules) self.final_layer = nn.Linear(hidden_dims[-1], 2 * self.n_coeffs) def _encode(self, mdct_frames: torch.Tensor) -> list[torch.Tensor]: batch_size, _, n_frames, _ = mdct_frames.shape ch1 = mdct_frames[:, 0, :, :] ch2 = mdct_frames[:, 1, :, :] x = torch.cat((ch1, ch2), dim=-1) result = self.encoder(x) mu = self.fc_mu(result) log_var = self.fc_var(result) return [mu, log_var] def _decode(self, z: torch.Tensor) -> torch.Tensor: batch_size, n_frames, _ = z.shape result = self.decoder_input(z) result = self.decoder(result) result = self.final_layer(result) ch1 = result[..., : self.n_coeffs] ch2 = result[..., self.n_coeffs :] result = torch.stack((ch1, ch2), dim=1) return result def reparameterize(self, mu: torch.Tensor, log_var: torch.Tensor) -> torch.Tensor: if self.training: std = torch.exp(0.5 * log_var) eps = torch.randn_like(std) return eps * std + mu else: return mu def forward( self, waveform: torch.Tensor ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: """ Forward pass handling both waveform conversion and VAE operations. Args: waveform: Input audio tensor of shape (batch_size, channels, samples) Returns: tuple containing: - reconstructed waveform - reconstructed MDCT coefficients - original MDCT coefficients (for loss computation) - mu - log_var """ # Convert input waveform to MDCT frames mdct_frames = audio_to_mdct_frames( waveform, frame_size=self.frame_size, midside=self.midside ) # (batch, channels, n_frames, n_coeffs) # Encode and decode mu, log_var = self._encode(mdct_frames) z = self.reparameterize(mu, log_var) mdct_recon = self._decode(z) # Convert back to waveform waveform_recon = mdct_frames_to_audio( mdct_recon, frame_size=self.frame_size, midside=self.midside ) return waveform_recon, mdct_recon, mdct_frames, mu, log_var def vae_loss( mdct_recon: torch.Tensor, mdct_original: torch.Tensor, waveform_recon: torch.Tensor, waveform_original: torch.Tensor, mu: torch.Tensor, log_var: torch.Tensor, kld_weight: float = 0.0, ) -> dict: """ Separate loss function that can be called after forward pass. """ # Reconstruction loss in MDCT domain # print(mdct_recon[0:2, 0, 0, :10], mdct_original[0:2, 0, 0, :10]) recons_loss = F.mse_loss(mdct_recon, mdct_original) # recons_loss = F.mse_loss(waveform_recon, waveform_original) # KL divergence loss kld_loss = torch.mean(-0.5 * torch.sum(1 + log_var - mu**2 - log_var.exp(), dim=2)) kld_loss = torch.mean(kld_loss) # Total loss loss = recons_loss # + kld_weight * kld_loss return {"loss": loss, "reconstruction_loss": recons_loss, "kld_loss": kld_loss} class BufferedAudioDataset(torch.utils.data.Dataset): def __init__( self, filepaths: List[str], sample_rate: int, num_workers: int = 1, chunk_size_s: float = 10.0, buffer_size: int = 50_000, ): self.filepaths = filepaths self.sample_rate = sample_rate self.chunk_size_s = chunk_size_s self.buffer_size = buffer_size self.chunk_size_samples = int(chunk_size_s * sample_rate) self.num_workers = num_workers self.items_since_last_reload = buffer_size # force a reload self.buffer = [] def __len__(self): return self.buffer_size * self.num_workers def _reload_buffer(self): self.buffer = [] rand_idxs = torch.randperm(len(self.filepaths)) print("Reloading buffer...") # max rand_idxs repeat endlessly rand_idxs = itertools.cycle(rand_idxs) pbar = tqdm(rand_idxs, total=len(self.filepaths), desc="Loading audio buffer") for idx in pbar: if len(self.buffer) >= self.buffer_size: break try: filepath = self.filepaths[idx] audio, sr = torchaudio.load(filepath) if sr != self.sample_rate: audio = torchaudio.functional.resample(audio, sr, self.sample_rate) # Pad if needed to ensure consistent chunk size if audio.shape[-1] < self.chunk_size_samples: continue # Split into chunks chunks = audio.unfold( -1, self.chunk_size_samples, self.chunk_size_samples ) chunks = chunks.chunk(chunks.shape[1], dim=1) # Filter chunks by minimum length valid_chunks = [ chunk.squeeze(1) for chunk in chunks if chunk.shape[-1] >= self.chunk_size_samples ] # filter out chunks of silence valid_chunks = [ chunk for chunk in valid_chunks if (chunk.abs() ** 2).mean() > 0.001 ] self.buffer.extend(valid_chunks) pbar.set_postfix({"buffer_size": len(self.buffer)}) except Exception as e: print(f"Error loading {filepath}: {e}") continue self.items_since_last_reload = 0 def __getitem__(self, _): if self.items_since_last_reload >= len(self.buffer): self._reload_buffer() # get a random preset and apply it to the audio buffer_idx = np.random.randint(0, len(self.buffer)) audio = self.buffer[buffer_idx] # ensure nothing is out of range if audio.abs().max() > 1.0: audio = audio / audio.abs().max() # apply random gain reduction # if np.random.uniform() < 0.5: # gain_reduction_db = np.random.uniform(-10, 0) # audio *= 10 ** (gain_reduction_db / 20.0) # self.items_since_last_reload += 1 return audio 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, f"last_ckpt.pt") print(f"Saving checkpoint to {checkpoint_path}") torch.save(checkpoint, checkpoint_path) def validate(model, val_loader, kld_weight: float = 0.0): print("Validating...") pbar = tqdm(val_loader, total=len(val_loader)) # accumulate loss, reconstruction loss, and kld loss total_loss = 0.0 total_recons_loss = 0.0 total_kld_loss = 0.0 total_samples = 0 for audio_batch in pbar: waveform_original = audio_batch.cuda() with torch.no_grad(): waveform_recon, mdct_recon, mdct_original, mu, log_var = model( waveform_original ) loss_dict = vae_loss( mdct_recon, mdct_original, waveform_recon, waveform_original, mu, log_var, kld_weight=kld_weight, ) loss = loss_dict["loss"] loss = loss.mean() total_loss += loss.item() total_recons_loss += loss_dict["reconstruction_loss"].item() total_kld_loss += loss_dict["kld_loss"].item() total_samples += len(waveform_original) return ( total_loss / total_samples, total_recons_loss / total_samples, total_kld_loss / total_samples, ) if __name__ == "__main__": run_start_time = time.strftime("%Y-%m-%d_%H-%M-%S") checkpoint_dir = f"/app/suno/christian/checkpoints/mdct-codec/{run_start_time}_s{random.randint(0, 9999)}" os.makedirs(checkpoint_dir, exist_ok=False) torch.set_float32_matmul_precision("medium") # Initialize distributed process group local_rank = int(os.environ.get("LOCAL_RANK", 0)) print(f"Local rank: {local_rank}") # dist.init_process_group(backend="nccl") torch.cuda.set_device(local_rank) # set the seed differently for each process torch.manual_seed(local_rank) run_config = { "training": { "max_steps": 1_000_000, "run_name": "test", "project_name": "mdct-codec", "lr": 1e-4, "grad_clip_norm": 1.0, "kld_weight": 0.0001, }, "model": { "frame_size": 1920, "latent_dim": 1920, "hidden_dims": [2048, 2048, 2048, 2048], "dropout": 0.0, "midside": False, }, "dataset": { "train_audio_dir": "/app/suno/data/audio_2ch_48khz_lg/train/genius_hq", "val_audio_dir": "/app/suno/data/audio_2ch_48khz_lg/val/genius_hq", "batch_size": 32, "num_workers": 4, "chunk_size_s": 10.0, "buffer_size": 5_000, "sample_rate": 48_000, }, } # Initialize wandb (only one process should do this) if local_rank == 0: wandb.init( project=run_config["training"]["project_name"], name=run_config["training"]["run_name"], ) wandb.config.update( {"checkpoint_dir": checkpoint_dir, "run_config": run_config} ) # find all filepaths in the train and val dirs train_filepaths = glob.glob( os.path.join(run_config["dataset"]["train_audio_dir"], "*.wav"), recursive=True ) val_filepaths = glob.glob( os.path.join(run_config["dataset"]["val_audio_dir"], "*.wav"), recursive=True ) # create train dataset train_dataset = BufferedAudioDataset( filepaths=train_filepaths, sample_rate=run_config["dataset"]["sample_rate"], num_workers=run_config["dataset"]["num_workers"], chunk_size_s=run_config["dataset"]["chunk_size_s"], buffer_size=run_config["dataset"]["buffer_size"], ) train_loader = DataLoader( train_dataset, batch_size=run_config["dataset"]["batch_size"], num_workers=run_config["dataset"]["num_workers"], shuffle=True, pin_memory=True, persistent_workers=True, drop_last=True, ) # create val dataset val_dataset = BufferedAudioDataset( filepaths=val_filepaths, sample_rate=run_config["dataset"]["sample_rate"], num_workers=run_config["dataset"]["num_workers"], chunk_size_s=run_config["dataset"]["chunk_size_s"], buffer_size=run_config["dataset"]["buffer_size"], ) val_loader = DataLoader( val_dataset, batch_size=run_config["dataset"]["batch_size"], num_workers=run_config["dataset"]["num_workers"], persistent_workers=True, pin_memory=True, shuffle=False, ) # create model and wrap in DDP model = WaveformMDCTVAE( frame_size=run_config["model"]["frame_size"], latent_dim=run_config["model"]["latent_dim"], hidden_dims=run_config["model"]["hidden_dims"], dropout=run_config["model"]["dropout"], midside=run_config["model"]["midside"], ) num_params = sum(p.numel() for p in model.parameters()) print(f"Number of parameters: {num_params/1e6:0.1f}M") model.cuda() # model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[local_rank]) # model = torch.compile(model) # Add dynamo compilation optimizer = torch.optim.AdamW(model.parameters(), lr=run_config["training"]["lr"]) warmup_scheduler = torch.optim.lr_scheduler.LinearLR( optimizer, start_factor=0.001, end_factor=1.0, total_iters=1000 ) cosine_scheduler = torch.optim.lr_scheduler.CosineAnnealingLR( optimizer, run_config["training"]["max_steps"] - 1000 ) scheduler = torch.optim.lr_scheduler.ChainedScheduler( [warmup_scheduler, cosine_scheduler] ) global_step = 0 while global_step < run_config["training"]["max_steps"]: pbar = tqdm(train_loader, total=len(train_loader)) for audio_batch in pbar: optimizer.zero_grad() # move to gpu audio_batch = audio_batch.cuda() # run the compression model waveform_recon, mdct_recon, mdct_original, mu, log_var = model(audio_batch) # Compute loss loss_dict = vae_loss( mdct_recon, mdct_original, waveform_recon, audio_batch, mu, log_var, kld_weight=run_config["training"]["kld_weight"], ) loss = loss_dict["loss"] loss.backward() torch.nn.utils.clip_grad_norm_( model.parameters(), run_config["training"]["grad_clip_norm"] ) optimizer.step() scheduler.step() loss = loss.mean() grad_norm = torch.norm( torch.stack( [ torch.norm(p.grad) for p in model.parameters() if p.grad is not None ] ) ) if local_rank == 0: pbar.set_postfix({"loss": loss.item()}) wandb.log( { "train/loss": loss.item(), "train/grad_norm": grad_norm.item(), "train/reconstruction_loss": loss_dict[ "reconstruction_loss" ].item(), "train/kld_loss": loss_dict["kld_loss"].item(), "trainer/lr": optimizer.param_groups[0]["lr"], "trainer/global_step": global_step, } ) global_step += 1 if local_rank == 0: print(f"Step {global_step} loss: {loss.item():.4f}") val_loss, val_recons_loss, val_kld_loss = validate( model, val_loader, kld_weight=run_config["training"]["kld_weight"] ) wandb.log( { "val/loss": val_loss, "val/reconstruction_loss": val_recons_loss, "val/kld_loss": val_kld_loss, "trainer/global_step": global_step, } ) save_checkpoint(model, optimizer, run_config, global_step, checkpoint_dir) print("Done!")