Download codec/audio_processing/dune_codec.py from Respair/dune_codec: direct link, hf CLI and curl.
- Browser
- Download file 12.7 kB
-
https://huggingface.co/Respair/dune_codec/resolve/main/codec/audio_processing/dune_codec.py
- Command line
-
hf download hf://Respair/dune_codec/codec/audio_processing/dune_codec.py
-
curl -L -o dune_codec.py https://huggingface.co/Respair/dune_codec/resolve/main/codec/audio_processing/dune_codec.py
12.7 kB
| 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 | |
| def tps(self): | |
| return self.frame_rate | |
| 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 |