import contextlib import inspect import json import logging import math import os import librosa import numpy as np import torch import torch.nn as nn import torch.nn.functional as F import torchaudio from huggingface_hub import snapshot_download from nemo.collections.tts.models import AudioCodecModel import pyloudnorm as pyln logger = logging.getLogger(__name__) def WNConv1d(*args, **kwargs): return nn.utils.weight_norm(nn.Conv1d(*args, **kwargs)) def WNConvTranspose1d(*args, **kwargs): return nn.utils.weight_norm(nn.ConvTranspose1d(*args, **kwargs)) class Snake1d(nn.Module): def __init__(self, channels): super().__init__() self.alpha = nn.Parameter(torch.ones(1, channels, 1)) def forward(self, x): return x + (1.0 / (self.alpha + 1e-9)) * torch.sin(self.alpha * x).pow(2) class ResidualUnit(nn.Module): def __init__(self, dim=16, dilation=1): super().__init__() pad = ((7 - 1) * dilation) // 2 self.block = nn.Sequential( Snake1d(dim), WNConv1d(dim, dim, kernel_size=7, dilation=dilation, padding=pad), Snake1d(dim), WNConv1d(dim, dim, kernel_size=1), ) def forward(self, x): y = self.block(x) pad = (x.shape[-1] - y.shape[-1]) // 2 if pad > 0: x = x[..., pad:-pad] return x + y class DACDecoderBlock(nn.Module): def __init__(self, input_dim=16, output_dim=8, stride=1): super().__init__() self.block = nn.Sequential( Snake1d(input_dim), WNConvTranspose1d( input_dim, output_dim, kernel_size=2 * stride, stride=stride, padding=math.ceil(stride / 2), output_padding=stride % 2, ), ResidualUnit(output_dim, dilation=1), ResidualUnit(output_dim, dilation=3), ResidualUnit(output_dim, dilation=9), ) def forward(self, x): return self.block(x) class DACStyleDecoder(nn.Module): def __init__(self, input_channels, decoder_dim, upsample_rates, d_out=1): super().__init__() layers = [WNConv1d(input_channels, decoder_dim, kernel_size=7, padding=3)] for i, stride in enumerate(upsample_rates): layers.append( DACDecoderBlock(decoder_dim // (2 ** i), decoder_dim // (2 ** (i + 1)), stride) ) final_dim = decoder_dim // (2 ** len(upsample_rates)) layers += [ Snake1d(final_dim), WNConv1d(final_dim, d_out, kernel_size=7, padding=3), nn.Tanh(), ] self.model = nn.Sequential(*layers) def forward(self, x): return self.model(x) class DuneAudioTokenizer(nn.Module): def __init__( self, nemo_model="nvidia/nemo-nano-codec-22khz-1.78kbps-12.5fps", # i only borrow its encoder as training a codec encoder (even FSQ) from scratch is a pain in the 🍑 sample_rate=44100, encoder_sample_rate=None, output_sample_rate=None, latent_dim=52, upsample_ratio=None, decoder_dim=1024, device="cuda", **kwargs, ): super().__init__() self.device = device self.nemo_model = nemo_model self.codec = AudioCodecModel.from_pretrained(nemo_model) self.codec.to(device) self.codec.eval() self.encoder_sample_rate = int(getattr(self.codec, "sample_rate", None) or encoder_sample_rate) self.samples_per_frame_in = int( getattr(self.codec, "samples_per_frame", None) or self._infer_samples_per_frame_in() ) self.frame_rate = self.encoder_sample_rate / self.samples_per_frame_in self.output_sample_rate = int(output_sample_rate or sample_rate) self.samples_per_frame_out = self._compute_samples_per_frame_out() self.latent_dim = int(latent_dim) self._backbone_frozen = False if upsample_ratio: self.upsample_ratio = list(upsample_ratio) elif self.output_sample_rate == self.encoder_sample_rate: self.upsample_ratio = [] else: sr_ratio = self.output_sample_rate // self.encoder_sample_rate self.upsample_ratio = list(self._infer_codec_upsample_rates()) + [sr_ratio] self.is_upsampling_model = bool(self.upsample_ratio) if not self.is_upsampling_model: self.dac_decoder = None else: self._validate_upsample_ratio() self.dac_decoder = DACStyleDecoder( input_channels=self.latent_dim, decoder_dim=decoder_dim, upsample_rates=self.upsample_ratio, d_out=1, ).to(device) def _infer_samples_per_frame_in(self): return int(np.prod([int(r) for r in self.codec.audio_encoder.down_sample_rates])) def _infer_codec_upsample_rates(self): return [int(r) for r in self.codec.audio_decoder.up_sample_rates] def _compute_samples_per_frame_out(self): num = self.output_sample_rate * self.samples_per_frame_in if num % self.encoder_sample_rate != 0: raise ValueError( f"{self.output_sample_rate}Hz output is not reachable from " f"{self.encoder_sample_rate}Hz at {self.samples_per_frame_in} samples/frame" ) return int(num // self.encoder_sample_rate) def _validate_upsample_ratio(self): total = int(np.prod(self.upsample_ratio)) if self.upsample_ratio else 1 if total != self.samples_per_frame_out: raise ValueError( f"upsample_ratio product {total} != samples_per_frame_out " f"{self.samples_per_frame_out}" ) def _set_frozen_eval(self): self.codec.audio_encoder.eval() self.codec.vector_quantizer.eval() def freeze_for_upsampling_finetune(self): prefixes = ("dac_decoder",) if self.dac_decoder is not None else ("codec.audio_decoder",) for name, param in self.named_parameters(): param.requires_grad = name.startswith(prefixes) self._backbone_frozen = True self._set_frozen_eval() total = sum(p.numel() for p in self.parameters()) trainable = sum(p.numel() for p in self.parameters() if p.requires_grad) logger.info(f"trainable {trainable / 1e6:.2f}M / {total / 1e6:.2f}M params") def train(self, mode=True): super().train(mode) if self._backbone_frozen: self._set_frozen_eval() return self @property def tps(self): return self.frame_rate @property def sampling_rate(self): return self.output_sample_rate def _maybe_no_grad(self): return torch.no_grad() if self._backbone_frozen else contextlib.nullcontext() def _dequantize(self, tokens, tokens_len): return self.codec.dequantize(tokens=tokens, tokens_len=tokens_len) def forward(self, x, bw=None): target_length = x.shape[-1] x_mono = x[:, 0, :] if x.dim() == 3 else x if self.output_sample_rate != self.encoder_sample_rate: x_enc = torchaudio.functional.resample( x_mono, self.output_sample_rate, self.encoder_sample_rate ) else: x_enc = x_mono audio_len = torch.full( (x_enc.shape[0],), x_enc.shape[1], device=x_enc.device, dtype=torch.long ) with self._maybe_no_grad(): tokens, tokens_len = self.codec.encode(audio=x_enc, audio_len=audio_len) if self.dac_decoder is not None: with self._maybe_no_grad(): dequant = self._dequantize(tokens, tokens_len) o = self.dac_decoder(dequant) else: o, _ = self.codec.decode(tokens=tokens, tokens_len=tokens_len) if o.dim() == 2: o = o.unsqueeze(1) if o.shape[-1] > target_length: o = o[..., :target_length] elif o.shape[-1] < target_length: o = F.pad(o, (0, target_length - o.shape[-1])) zero = torch.zeros((), device=x.device) return o, zero, zero, None def encode(self, audio_path_or_wv, sr=None, loudness_normalize=False, loudness_threshold=-23.0): if isinstance(audio_path_or_wv, str): wv, sr = librosa.load(audio_path_or_wv, mono=True, sr=None) else: wv = audio_path_or_wv if sr is None: raise ValueError("sr is required when passing a waveform") if loudness_normalize: meter = pyln.Meter(sr) wv = pyln.normalize.loudness(wv, meter.integrated_loudness(wv), loudness_threshold) if sr != self.encoder_sample_rate: wv = librosa.resample(wv, orig_sr=sr, target_sr=self.encoder_sample_rate) audio = torch.from_numpy(wv).float().unsqueeze(0).to(self.device) audio_len = torch.tensor([audio.shape[-1]], device=self.device, dtype=torch.long) with torch.no_grad(): tokens, _ = self.codec.encode(audio=audio, audio_len=audio_len) return tokens[0] def _post_filter(self, audio): """Spectral post-filter over the reconstructed waveform. Applied per item at the output rate. A failure here must not cost the caller their audio, so it degrades to the unfiltered signal. """ try: from ._postfilter import get_post_filter pf = get_post_filter(device="cpu") except Exception: return audio out = np.array(audio, dtype=np.float32, copy=True) flat = out.reshape(-1, out.shape[-1]) if out.ndim > 1 else out[None] for i in range(flat.shape[0]): try: filtered = pf(flat[i], self.output_sample_rate) except Exception: continue n = min(filtered.size, flat.shape[1]) flat[i, :n] = filtered[:n] return flat.reshape(out.shape) if out.ndim > 1 else flat[0] def decode(self, vq_code): tokens = vq_code if vq_code.dim() == 3 else vq_code.unsqueeze(0) tokens = tokens.to(self.device) tokens_len = torch.full( (tokens.shape[0],), tokens.shape[-1], device=self.device, dtype=torch.long ) with torch.no_grad(): if self.dac_decoder is not None: audio = self.dac_decoder(self._dequantize(tokens, tokens_len)) if audio.dim() == 3: audio = audio[:, 0, :] else: audio, _ = self.codec.decode(tokens=tokens, tokens_len=tokens_len) return self._post_filter(audio.cpu().numpy()) def _state_dict_from(ckpt): state_dict = ckpt.get("model_state_dict") or ckpt.get("state_dict") or ckpt out = {} for key, value in state_dict.items(): for prefix in ("module.", "_orig_mod."): if key.startswith(prefix): key = key[len(prefix):] out[key] = value return out def _model_kwargs(cfg): cfg = dict(cfg) if "nemo_model" not in cfg and "nemo_model_name" in cfg: cfg["nemo_model"] = cfg.pop("nemo_model_name") accepted = set(inspect.signature(DuneAudioTokenizer.__init__).parameters) return {k: v for k, v in cfg.items() if k in accepted - {"self", "device", "kwargs"}} def prepare(checkpoint_path, config_path=None, device="cuda", compile_after_load=False): ckpt = torch.load(checkpoint_path, map_location="cpu", weights_only=False) cfg = ckpt.get("config") if not isinstance(cfg, dict): with open(config_path, "r") as f: cfg = json.load(f) model = DuneAudioTokenizer(**_model_kwargs(cfg), device=device).to(device) missing, unexpected = model.load_state_dict(_state_dict_from(ckpt), strict=False) logger.info(f"loaded {checkpoint_path} | missing={len(missing)} unexpected={len(unexpected)}") model.eval() if compile_after_load: model = torch.compile(model, mode="default").eval() return model def load_dune_audio_tokenizer(tokenizer_name_or_path, device="cuda"): is_local = os.path.exists(tokenizer_name_or_path) if not is_local: tokenizer_path = snapshot_download(tokenizer_name_or_path) else: tokenizer_path = tokenizer_name_or_path config_path = os.path.join(tokenizer_path, "config.json") checkpoint_path = os.path.join(tokenizer_path, "model_209k.pth") config = json.load(open(config_path)) model = prepare(checkpoint_path, config_path, device) model.eval() return model