PyTorch
Russian
English
proxymos
audio
mos
ProxyMos / inference_model.py
assskelad's picture
Upload 3 files
47a117e verified
Raw
History Blame Contribute Delete
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),
)
@torch.inference_mode()
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]
@torch.inference_mode()
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}")