Download neucodec/model.py from puterijessica/WideCodec: direct link, hf CLI and curl.
- Browser
- Download file 8.04 kB
-
https://huggingface.co/puterijessica/WideCodec/resolve/main/neucodec/model.py
- Command line
-
hf download hf://puterijessica/WideCodec/neucodec/model.py
-
curl -L -o model.py https://huggingface.co/puterijessica/WideCodec/resolve/main/neucodec/model.py
8.04 kB
| from typing import Optional, Dict | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| import torchaudio | |
| from torchaudio import transforms as T | |
| from huggingface_hub import PyTorchModelHubMixin, ModelHubMixin, hf_hub_download | |
| from transformers import AutoFeatureExtractor, HubertModel, Wav2Vec2BertModel | |
| from .codec_encoder import CodecEncoder | |
| from .codec_encoder_distill import DistillCodecEncoder | |
| from .codec_decoder_vocos import CodecDecoderVocos | |
| from .module import SemanticEncoder | |
| class NeuCodec( | |
| nn.Module, | |
| PyTorchModelHubMixin, | |
| repo_url="https://github.com/neuphonic/neucodec", | |
| license="apache-2.0", | |
| ): | |
| def __init__(self, sample_rate: int, hop_length: int, decoder_depth: int = 12): | |
| super().__init__() | |
| self.sample_rate = sample_rate | |
| self.hop_length = hop_length | |
| self.semantic_model = Wav2Vec2BertModel.from_pretrained( | |
| "facebook/w2v-bert-2.0", output_hidden_states=True | |
| ) | |
| self.feature_extractor = AutoFeatureExtractor.from_pretrained( | |
| "facebook/w2v-bert-2.0" | |
| ) | |
| self.SemanticEncoder_module = SemanticEncoder(1024, 1024, 1024) | |
| self.CodecEnc = CodecEncoder() | |
| self.generator = CodecDecoderVocos(hop_length=hop_length, depth=decoder_depth) | |
| self.fc_prior = nn.Linear(2048, 2048) | |
| self.fc_post_a = nn.Linear(2048, 1024) | |
| def device(self): | |
| return next(self.parameters()).device | |
| def _from_pretrained( | |
| cls, | |
| *, | |
| model_id: str = None, | |
| revision: Optional[str] = None, | |
| cache_dir: Optional[str] = None, | |
| force_download: bool = False, | |
| proxies: Optional[Dict] = None, | |
| resume_download: bool = False, | |
| local_files_only: bool = False, | |
| token: Optional[str] = None, | |
| map_location: str = "cpu", | |
| strict: bool = False, | |
| local_ckpt_path: str = None, | |
| **model_kwargs, | |
| ): | |
| if model_id == "neuphonic/neucodec": | |
| ignore_keys = ["fc_post_s", "SemanticDecoder"] | |
| elif model_id == "neuphonic/distill-neucodec": | |
| ignore_keys = [] | |
| else: | |
| ignore_keys = [] | |
| if model_id is not None: | |
| ckpt_path = hf_hub_download( | |
| repo_id=model_id, | |
| filename="pytorch_model.bin", | |
| revision=revision, | |
| cache_dir=cache_dir, | |
| force_download=force_download, | |
| proxies=proxies, | |
| resume_download=resume_download, | |
| local_files_only=local_files_only, | |
| token=token, | |
| ) | |
| else: | |
| # incase we interpolate the weight to become 960 instead train from scratch | |
| ckpt_path = local_ckpt_path | |
| # initialize model | |
| decoder_depth = model_kwargs.pop('decoder_depth', 12) | |
| model = cls(44_100, 882, decoder_depth=decoder_depth) | |
| # load weights | |
| state_dict = torch.load(ckpt_path, map_location) | |
| contains_list = lambda s, l: any(i in s for i in l) | |
| state_dict = { | |
| k:v for k, v in state_dict.items() | |
| if not contains_list(k, ignore_keys) | |
| } | |
| # Filter out keys with shape mismatches (e.g. 48k model vs 24k checkpoint) | |
| model_state = model.state_dict() | |
| state_dict = { | |
| k: v for k, v in state_dict.items() | |
| if k in model_state and v.shape == model_state[k].shape | |
| } | |
| model.load_state_dict(state_dict, strict=False) | |
| return model | |
| def _prepare_audio(self, audio_or_path: torch.Tensor | Path | str): | |
| # load from file | |
| if isinstance(audio_or_path, (Path, str)): | |
| y, sr = torchaudio.load(audio_or_path) | |
| if sr != 16_000: | |
| y, sr = (T.Resample(sr, 16_000)(y), 16_000) | |
| y = y[None, :] # [1, T] -> [B, 1, T] | |
| # ensure input tensor is of correct shape | |
| elif isinstance(audio_or_path, torch.Tensor): | |
| y = audio_or_path | |
| if len(y.shape) == 3: | |
| y = audio_or_path | |
| else: | |
| raise ValueError( | |
| f"NeuCodec expects tensor audio input to be of shape [B, 1, T] -- received shape: {y.shape}" | |
| ) | |
| # pad audio | |
| pad_for_wav = 320 - (y.shape[-1] % 320) | |
| y = torch.nn.functional.pad(y, (0, pad_for_wav)) | |
| return y | |
| def encode_code(self, audio_or_path: torch.Tensor | Path | str) -> torch.Tensor: | |
| """ | |
| Args: | |
| audio_or_path: torch.Tensor [B, 1, T] | Path | str, input audio | |
| Returns: | |
| fsq_codes: torch.Tensor [B, 1, F], 50hz FSQ codes | |
| """ | |
| # prepare inputs | |
| y = self._prepare_audio(audio_or_path) | |
| semantic_features = self.feature_extractor( | |
| [w for w in y.squeeze(1).cpu()], sampling_rate=16_000, return_tensors="pt" | |
| ).input_features.to(self.device) | |
| # acoustic encoding | |
| acoustic_emb = self.CodecEnc(y.to(self.device)) | |
| acoustic_emb = acoustic_emb.transpose(1, 2) | |
| # semantic encoding | |
| semantic_output = ( | |
| self.semantic_model(semantic_features).hidden_states[16].transpose(1, 2) | |
| ) | |
| semantic_encoded = self.SemanticEncoder_module(semantic_output) | |
| # concatenate embeddings | |
| if acoustic_emb.shape[-1] != semantic_encoded.shape[-1]: | |
| min_len = min(acoustic_emb.shape[-1], semantic_encoded.shape[-1]) | |
| acoustic_emb = acoustic_emb[:, :, :min_len] | |
| semantic_encoded = semantic_encoded[:, :, :min_len] | |
| concat_emb = torch.cat([semantic_encoded, acoustic_emb], dim=1) | |
| concat_emb = self.fc_prior(concat_emb.transpose(1, 2)).transpose(1, 2) | |
| # quantize | |
| _, fsq_codes, _ = self.generator(concat_emb, vq=True) | |
| return fsq_codes | |
| def encode_code_from_features(self, audio: torch.Tensor, semantic_features: torch.Tensor) -> torch.Tensor: | |
| """Encode using pre-computed semantic features, avoiding CPU feature extraction. | |
| Args: | |
| audio: torch.Tensor [B, 1, T], 16kHz input audio | |
| semantic_features: torch.Tensor [B, seq_len, feat_dim], pre-computed features | |
| Returns: | |
| fsq_codes: torch.Tensor [B, 1, F], 50hz FSQ codes | |
| """ | |
| y = self._prepare_audio(audio) | |
| semantic_features = semantic_features.to(self.device) | |
| # acoustic encoding | |
| acoustic_emb = self.CodecEnc(y.to(self.device)) | |
| acoustic_emb = acoustic_emb.transpose(1, 2) | |
| # semantic encoding | |
| semantic_output = ( | |
| self.semantic_model(semantic_features).hidden_states[16].transpose(1, 2) | |
| ) | |
| semantic_encoded = self.SemanticEncoder_module(semantic_output) | |
| # concatenate embeddings | |
| if acoustic_emb.shape[-1] != semantic_encoded.shape[-1]: | |
| min_len = min(acoustic_emb.shape[-1], semantic_encoded.shape[-1]) | |
| acoustic_emb = acoustic_emb[:, :, :min_len] | |
| semantic_encoded = semantic_encoded[:, :, :min_len] | |
| concat_emb = torch.cat([semantic_encoded, acoustic_emb], dim=1) | |
| concat_emb = self.fc_prior(concat_emb.transpose(1, 2)).transpose(1, 2) | |
| # quantize | |
| _, fsq_codes, _ = self.generator(concat_emb, vq=True) | |
| return fsq_codes | |
| def decode_code(self, fsq_codes: torch.Tensor) -> torch.Tensor: | |
| """ | |
| Args: | |
| fsq_codes: torch.Tensor [B, 1, F], 50hz FSQ codes | |
| Returns: | |
| recon: torch.Tensor [B, 1, T], reconstructed 48kHz audio | |
| """ | |
| fsq_post_emb = self.generator.quantizer.get_output_from_indices(fsq_codes.transpose(1, 2)) | |
| fsq_post_emb = fsq_post_emb.transpose(1, 2) | |
| fsq_post_emb = self.fc_post_a(fsq_post_emb.transpose(1, 2)).transpose(1, 2) | |
| recon = self.generator(fsq_post_emb.transpose(1, 2), vq=False)[0] | |
| return recon | |