| 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", |
| 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 |