Download inference_model.py from lab260/ProxyMos: direct link, hf CLI and curl.
- Browser
- Download file 4.74 kB
-
https://huggingface.co/lab260/ProxyMos/resolve/main/inference_model.py
- Command line
-
hf download hf://lab260/ProxyMos/inference_model.py
-
curl -L -o inference_model.py https://huggingface.co/lab260/ProxyMos/resolve/main/inference_model.py
4.74 kB
| import torch | |
| import torch.nn as nn | |
| from torch.utils.data import DataLoader | |
| from pathlib import Path | |
| from tqdm import tqdm | |
| import numpy as np | |
| from accelerate import Accelerator | |
| from torch.optim.lr_scheduler import CosineAnnealingLR | |
| from scipy.stats import spearmanr, pearsonr | |
| from sklearn.metrics import mean_squared_error, mean_absolute_error | |
| from fairseq2.nn import BatchLayout | |
| import torchaudio | |
| import torchaudio.transforms as T | |
| TARGET_SR = 16_000 | |
| class AttentiveStatsPooling(nn.Module): | |
| def __init__(self, dim: int): | |
| super().__init__() | |
| self.att = nn.Sequential( | |
| nn.Linear(dim, dim), | |
| nn.Tanh(), | |
| nn.Linear(dim, 1) | |
| ) | |
| def forward( | |
| self, | |
| x: torch.Tensor, # [B, T, D] | |
| padding_mask: torch.Tensor | None = None # [B, T], True = pad | |
| ) -> torch.Tensor: | |
| """Returns: [B, 2D]""" | |
| scores = self.att(x).squeeze(-1) # [B, T] | |
| if padding_mask is not None: | |
| scores = scores.masked_fill(padding_mask, -1e9) | |
| weights = torch.softmax(scores, dim=1).unsqueeze(-1) | |
| mean = torch.sum(weights * x, dim=1) | |
| var = torch.sum(weights * (x - mean.unsqueeze(1)) ** 2, dim=1) | |
| std = torch.sqrt(var + 1e-6) | |
| return torch.cat([mean, std], dim=-1) | |
| class OmniMOS(nn.Module): | |
| """ | |
| MOS prediction model built on top of a Wav2Vec2-style encoder. | |
| Args: | |
| encoder (nn.Module): Feature extraction encoder (e.g. Wav2Vec2). | |
| hidden_dim (int): Hidden dimensionality. Default: 1024. | |
| attentive_pooling (bool): Use attentive stats pooling instead of mean pooling. | |
| """ | |
| def __init__( | |
| self, | |
| encoder: nn.Module, | |
| hidden_dim: int = 1024, | |
| attentive_pooling: bool = True, | |
| ): | |
| super().__init__() | |
| self.encoder = encoder | |
| dim = hidden_dim | |
| if attentive_pooling: | |
| self.pool = AttentiveStatsPooling(dim) | |
| pooled_dim = dim * 2 | |
| else: | |
| self.pool = None | |
| pooled_dim = dim | |
| self.head = nn.Sequential( | |
| nn.Linear(pooled_dim, hidden_dim), | |
| nn.GELU(), | |
| nn.Linear(hidden_dim, 1), | |
| ) | |
| def inference(self, wave: torch.Tensor) -> torch.Tensor: | |
| self.eval() | |
| return self.forward(wave) | |
| def forward(self, wave: torch.Tensor) -> torch.Tensor: | |
| """ | |
| Args: | |
| wave (torch.Tensor): Waveform tensor of shape [B, T] or [B, 1, T]. | |
| Returns: | |
| torch.Tensor: MOS scores of shape [B]. | |
| """ | |
| wave = wave.float() | |
| if wave.dim() == 3 and wave.shape[1] == 1: | |
| wave = wave.squeeze(1) | |
| if wave.dim() == 3: | |
| wave = wave.mean(dim=1) | |
| B, T = wave.shape | |
| seqs_layout = BatchLayout( | |
| shape=(B, T), | |
| seq_lens=[T] * B, | |
| packed=False, | |
| device=wave.device, | |
| ) | |
| features = self.encoder.extract_features(wave, seqs_layout) | |
| if hasattr(features, "seqs"): | |
| feats = features.seqs | |
| elif hasattr(features, "encoder_output"): | |
| feats = features.encoder_output | |
| elif isinstance(features, tuple): | |
| feats = features[0] | |
| else: | |
| feats = features | |
| if self.pool is not None: | |
| pooled = self.pool(feats, None) # [B, 2D] | |
| else: | |
| pooled = feats.mean(dim=1) # [B, D] | |
| return self.head(pooled).squeeze(-1) | |
| def load_audio(path: str) -> torch.Tensor: | |
| wave, sr = torchaudio.load(path) | |
| if wave.shape[0] > 1: | |
| wave = wave.mean(dim=0, keepdim=True) | |
| if sr != TARGET_SR: | |
| wave = T.Resample(sr, TARGET_SR)(wave) | |
| return wave # [1, T] | |
| def predict_mos(model: OmniMOS, path: str, device: torch.device) -> float: | |
| wave = load_audio(path).unsqueeze(0).to(device) # [1, 1, T] | |
| return model(wave).item() | |
| def load_model(checkpoint_path: str, device: torch.device) -> OmniMOS: | |
| from fairseq2.models.wav2vec2 import get_wav2vec2_model_hub | |
| hub = get_wav2vec2_model_hub() | |
| fs2_config = hub.get_model_config('omniASR_W2V_300M') | |
| encoder = hub.create_new_model(fs2_config, device=torch.device("cpu")) | |
| model = OmniMOS(encoder=encoder) | |
| model.load_state_dict(torch.load(checkpoint_path, map_location="cpu")) | |
| model.to(device).eval() | |
| return model | |
| if __name__ == "__main__": | |
| import sys | |
| audio_path = sys.argv[1] | |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| model = load_model("best_model_full.pt", device) | |
| score = predict_mos(model, audio_path, device) | |
| print(f"MOS: {score:.4f}") | |