| 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) |
|
|
| @property |
| def device(self): |
| return next(self.parameters()).device |
|
|
| @classmethod |
| 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: |
| |
| ckpt_path = local_ckpt_path |
|
|
| |
| decoder_depth = model_kwargs.pop('decoder_depth', 12) |
| model = cls(44_100, 882, decoder_depth=decoder_depth) |
|
|
| |
| 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) |
| } |
|
|
| |
| 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): |
| |
| |
| 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, :] |
|
|
| |
| 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_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 |
| """ |
| |
| |
| 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_emb = self.CodecEnc(y.to(self.device)) |
| acoustic_emb = acoustic_emb.transpose(1, 2) |
|
|
| |
| semantic_output = ( |
| self.semantic_model(semantic_features).hidden_states[16].transpose(1, 2) |
| ) |
| semantic_encoded = self.SemanticEncoder_module(semantic_output) |
|
|
| |
| 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) |
|
|
| |
| _, 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_emb = self.CodecEnc(y.to(self.device)) |
| acoustic_emb = acoustic_emb.transpose(1, 2) |
|
|
| |
| semantic_output = ( |
| self.semantic_model(semantic_features).hidden_states[16].transpose(1, 2) |
| ) |
| semantic_encoded = self.SemanticEncoder_module(semantic_output) |
|
|
| |
| 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) |
|
|
| |
| _, 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 |
|
|