| """AmberNet language identification — standalone PyTorch, no NeMo dependency. |
| |
| Faithful re-implementation of NeMo's EncDecSpeakerLabelModel as configured for |
| AmberNet (ContextNet-style separable conv encoder + squeeze-excite, x-vector |
| stats pooling head). Parameter names match the original .nemo checkpoint so the |
| weights load verbatim. |
| """ |
|
|
| import json |
| import math |
| import os |
|
|
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
|
|
| CONSTANT = 1e-5 |
|
|
|
|
| class MaskedConv1d(nn.Module): |
| """Conv1d that zeroes padded timesteps before convolving.""" |
|
|
| def __init__(self, in_ch, out_ch, kernel_size, padding=0, groups=1): |
| super().__init__() |
| self.conv = nn.Conv1d(in_ch, out_ch, kernel_size, padding=padding, groups=groups, bias=False) |
|
|
| def forward(self, x, lens): |
| mask = torch.arange(x.shape[-1], device=x.device)[None, :] < lens[:, None] |
| return self.conv(x * mask.unsqueeze(1)), lens |
|
|
|
|
| class SqueezeExcite(nn.Module): |
| def __init__(self, channels, reduction_ratio=8): |
| super().__init__() |
| self.fc = nn.Sequential( |
| nn.Linear(channels, channels // reduction_ratio, bias=False), |
| nn.ReLU(inplace=True), |
| nn.Linear(channels // reduction_ratio, channels, bias=False), |
| ) |
|
|
| def forward(self, x, lens): |
| |
| mask = (torch.arange(x.shape[-1], device=x.device)[None, :] < lens[:, None]).unsqueeze(1) |
| x = x * mask |
| y = x.sum(dim=-1, keepdim=True) / mask.sum(dim=-1, keepdim=True).to(x.dtype) |
| y = self.fc(y.transpose(1, -1)).transpose(1, -1) |
| return x * torch.sigmoid(y), lens |
|
|
|
|
| def _conv_bn(in_ch, out_ch, kernel_size, separable): |
| padding = (kernel_size - 1) // 2 |
| if separable: |
| layers = [ |
| MaskedConv1d(in_ch, in_ch, kernel_size, padding=padding, groups=in_ch), |
| MaskedConv1d(in_ch, out_ch, 1), |
| ] |
| else: |
| layers = [MaskedConv1d(in_ch, out_ch, kernel_size, padding=padding)] |
| return layers + [nn.BatchNorm1d(out_ch, eps=1e-3, momentum=0.1)] |
|
|
|
|
| class JasperBlock(nn.Module): |
| def __init__(self, inplanes, planes, repeat, kernel_size, dropout, residual, separable=True, se=True): |
| super().__init__() |
| layers, inp = [], inplanes |
| for _ in range(repeat - 1): |
| layers += _conv_bn(inp, planes, kernel_size, separable) |
| layers += [nn.ReLU(inplace=True), nn.Dropout(dropout)] |
| inp = planes |
| layers += _conv_bn(inp, planes, kernel_size, separable) |
| if se: |
| layers.append(SqueezeExcite(planes)) |
| self.mconv = nn.ModuleList(layers) |
| self.res = nn.ModuleList([nn.ModuleList(_conv_bn(inplanes, planes, 1, separable=False))]) if residual else None |
| self.mout = nn.Sequential(nn.ReLU(inplace=True), nn.Dropout(dropout)) |
|
|
| def forward(self, x, lens): |
| out = x |
| for layer in self.mconv: |
| out, lens = layer(out, lens) if isinstance(layer, (MaskedConv1d, SqueezeExcite)) else (layer(out), lens) |
| if self.res is not None: |
| res = x |
| for layer in self.res[0]: |
| res, _ = layer(res, lens) if isinstance(layer, MaskedConv1d) else (layer(res), lens) |
| out = out + res |
| return self.mout(out), lens |
|
|
|
|
| class ConvASREncoder(nn.Module): |
| def __init__(self, feat_in, jasper): |
| super().__init__() |
| blocks = [] |
| for cfg in jasper: |
| blocks.append( |
| JasperBlock( |
| feat_in, |
| cfg["filters"], |
| cfg["repeat"], |
| cfg["kernel"][0], |
| cfg["dropout"], |
| cfg["residual"], |
| cfg.get("separable", True), |
| cfg.get("se", True), |
| ) |
| ) |
| feat_in = cfg["filters"] |
| self.encoder = nn.ModuleList(blocks) |
|
|
| def forward(self, x, lens): |
| for block in self.encoder: |
| x, lens = block(x, lens) |
| return x, lens |
|
|
|
|
| class SpeakerDecoder(nn.Module): |
| """x-vector stats pooling (mean+std) → embedding → classifier.""" |
|
|
| def __init__(self, feat_in, num_classes, emb_size): |
| super().__init__() |
| self.emb_layers = nn.ModuleList( |
| [ |
| nn.Sequential( |
| nn.Linear(feat_in * 2, emb_size), |
| nn.BatchNorm1d(emb_size, affine=False, track_running_stats=True), |
| nn.ReLU(inplace=True), |
| ) |
| ] |
| ) |
| self.final = nn.Linear(emb_size, num_classes) |
|
|
| def forward(self, x, lens): |
| mask = (torch.arange(x.shape[-1], device=x.device)[None, :] < lens[:, None]).unsqueeze(1) |
| x = x * mask |
| mean = x.sum(dim=-1) / lens.unsqueeze(-1).to(x.dtype) |
| std = ( |
| ((x - mean.unsqueeze(-1)) * mask).pow(2).sum(-1).div(lens.view(-1, 1) - 1).clamp(min=1e-10).sqrt() |
| ) |
| pool = torch.cat([mean, std], dim=-1) |
| layer = self.emb_layers[0] |
| emb = layer[:2](pool) |
| return self.final(layer(pool)), emb |
|
|
|
|
| class MelSpectrogram(nn.Module): |
| """NeMo AudioToMelSpectrogramPreprocessor, inference path (no dither/augment).""" |
|
|
| def __init__(self, sample_rate=16000, n_fft=512, win_length=400, hop_length=160, n_mels=80, preemph=0.97): |
| super().__init__() |
| self.n_fft, self.win_length, self.hop_length, self.preemph = n_fft, win_length, hop_length, preemph |
| self.register_buffer("window", torch.hann_window(win_length, periodic=False)) |
| self.register_buffer("fb", torch.zeros(1, n_mels, n_fft // 2 + 1)) |
| self.register_buffer("stft_basis", torch.zeros(n_fft + 2, 1, n_fft), persistent=False) |
| self.build_stft_basis() |
|
|
| def build_stft_basis(self): |
| """STFT as a strided conv: portable to every ONNX runtime, unlike the STFT op. |
| |
| Must be re-run after loading weights, since it is derived from `window`. |
| """ |
| n_bins = self.n_fft // 2 + 1 |
| pad = (self.n_fft - self.win_length) // 2 |
| window = F.pad(self.window, (pad, self.n_fft - self.win_length - pad)) |
| angle = torch.outer( |
| torch.arange(n_bins, dtype=torch.float64), torch.arange(self.n_fft, dtype=torch.float64) |
| ) * (-2 * math.pi / self.n_fft) |
| basis = torch.cat([torch.cos(angle), torch.sin(angle)]).float() * window |
| self.stft_basis = basis.unsqueeze(1).to(self.stft_basis.device) |
|
|
| def get_seq_len(self, seq_len): |
| return torch.div(seq_len + self.n_fft // 2 * 2 - self.n_fft, self.hop_length, rounding_mode="floor").long() |
|
|
| def forward(self, x, seq_len): |
| out_len = self.get_seq_len(seq_len) |
| time_mask = torch.arange(x.shape[1], device=x.device)[None, :] < seq_len[:, None] |
| x = torch.cat((x[:, :1], x[:, 1:] - self.preemph * x[:, :-1]), dim=1) * time_mask |
|
|
| x = F.pad(x, (self.n_fft // 2, self.n_fft // 2)) |
| spec = F.conv1d(x.unsqueeze(1), self.stft_basis, stride=self.hop_length) |
| real, imag = spec.chunk(2, dim=1) |
| power = real.pow(2) + imag.pow(2) |
| mel = torch.matmul(self.fb, power) |
| mel = torch.log(mel + 2**-24) |
|
|
| |
| valid = torch.arange(mel.shape[-1], device=mel.device)[None, :] < out_len[:, None] |
| n = valid.sum(dim=1)[:, None] |
| mean = torch.where(valid.unsqueeze(1), mel, torch.zeros_like(mel)).sum(-1) / n |
| std = torch.sqrt( |
| torch.where(valid.unsqueeze(1), mel - mean.unsqueeze(-1), torch.zeros_like(mel)).pow(2).sum(-1) / (n - 1.0) |
| ) |
| mel = (mel - mean.unsqueeze(-1)) / (std + CONSTANT).unsqueeze(-1) |
| return mel * valid.unsqueeze(1), out_len |
|
|
|
|
| class AmberNet(nn.Module): |
| """Spoken language identification over 107 languages. |
| |
| forward(audio [B, N] float32 16 kHz, audio_len [B]) -> (logits [B, 107], embedding [B, 512]) |
| """ |
|
|
| def __init__(self, config): |
| super().__init__() |
| self.config = config |
| self.labels = config["labels"] |
| self.preprocessor = nn.Module() |
| self.preprocessor.featurizer = MelSpectrogram(**config["preprocessor"]) |
| self.encoder = ConvASREncoder(config["encoder"]["feat_in"], config["encoder"]["jasper"]) |
| self.decoder = SpeakerDecoder(**config["decoder"]) |
|
|
| def load_state_dict(self, *args, **kwargs): |
| result = super().load_state_dict(*args, **kwargs) |
| self.preprocessor.featurizer.build_stft_basis() |
| return result |
|
|
| def forward(self, audio, audio_len): |
| feats, feat_len = self.preprocessor.featurizer(audio, audio_len) |
| enc, enc_len = self.encoder(feats, feat_len) |
| return self.decoder(enc, enc_len) |
|
|
| @torch.inference_mode() |
| def classify(self, audio, audio_len=None, top_k=5): |
| """Returns a list (per batch item) of (language, probability), most likely first.""" |
| if audio.ndim == 1: |
| audio = audio.unsqueeze(0) |
| if audio_len is None: |
| audio_len = torch.full((audio.shape[0],), audio.shape[1], dtype=torch.long, device=audio.device) |
| probs = self(audio, audio_len)[0].softmax(-1) |
| top = probs.topk(min(top_k, len(self.labels)), dim=-1) |
| return [ |
| [(self.labels[i], float(p)) for p, i in zip(row_p, row_i)] |
| for row_p, row_i in zip(top.values, top.indices) |
| ] |
|
|
| @classmethod |
| def from_pretrained(cls, path): |
| """Load from a directory holding config.json + model.safetensors (or pytorch_model.bin).""" |
| with open(os.path.join(path, "config.json")) as f: |
| config = json.load(f) |
| model = cls(config) |
| safetensors_path = os.path.join(path, "model.safetensors") |
| if os.path.exists(safetensors_path): |
| from safetensors.torch import load_file |
|
|
| state = load_file(safetensors_path) |
| else: |
| state = torch.load(os.path.join(path, "pytorch_model.bin"), map_location="cpu", weights_only=True) |
| model.load_state_dict(state) |
| return model.eval() |
|
|