Feature Extraction
Transformers
Safetensors
resemblyzer
audio
speaker-embedding
speaker-verification
custom_code
Instructions to use prj-beatrice/resemblyzer-torch-native with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use prj-beatrice/resemblyzer-torch-native with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("feature-extraction", model="prj-beatrice/resemblyzer-torch-native", trust_remote_code=True)# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("prj-beatrice/resemblyzer-torch-native", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
| """Self-contained Resemblyzer inference with its WebRTC VAD preprocessing. | |
| The Apache-2.0 Resemblyzer and BSD-licensed WebRTC implementations were | |
| adapted for tensor-only batched execution; see LICENSE. | |
| """ | |
| from __future__ import annotations | |
| import math | |
| import operator | |
| from collections.abc import Sequence | |
| from dataclasses import dataclass | |
| import torch | |
| import torch.nn.functional as F | |
| import torchaudio.functional as AF | |
| from torch import nn | |
| from transformers import PretrainedConfig, PreTrainedModel | |
| from transformers.utils import ModelOutput | |
| _TARGET_SAMPLING_RATE = 16_000 | |
| _PARTIAL_RATE = 1.3 | |
| _MIN_COVERAGE = 0.75 | |
| def _validate_sampling_rate(sampling_rate: int) -> int: | |
| if isinstance(sampling_rate, bool): | |
| raise TypeError("sampling_rate must be an integer") | |
| try: | |
| sampling_rate = operator.index(sampling_rate) | |
| except TypeError as error: | |
| raise TypeError("sampling_rate must be an integer") from error | |
| if sampling_rate <= 0: | |
| raise ValueError("sampling_rate must be positive") | |
| return sampling_rate | |
| def _as_waveform_sequence(input_values, input_lengths=None): | |
| if input_lengths is not None: | |
| if not isinstance(input_values, torch.Tensor) or input_values.ndim != 2: | |
| raise ValueError( | |
| "input_lengths requires dense input_values with shape [B, T]" | |
| ) | |
| if not isinstance(input_lengths, torch.Tensor): | |
| raise TypeError("input_lengths must be an integer tensor with shape [B]") | |
| if ( | |
| input_lengths.dtype == torch.bool | |
| or input_lengths.is_floating_point() | |
| or input_lengths.is_complex() | |
| ): | |
| raise TypeError("input_lengths must use an integer dtype") | |
| if input_lengths.shape != (input_values.shape[0],): | |
| raise ValueError("input_lengths must have shape [B]") | |
| lengths = input_lengths.tolist() | |
| if any(length <= 0 or length > input_values.shape[1] for length in lengths): | |
| raise ValueError("input_lengths values are outside the dense input") | |
| waveforms = tuple( | |
| input_values[index, :length] for index, length in enumerate(lengths) | |
| ) | |
| elif isinstance(input_values, torch.Tensor): | |
| if input_values.ndim == 1: | |
| waveforms = (input_values,) | |
| elif input_values.ndim == 2: | |
| waveforms = tuple(input_values.unbind(0)) | |
| else: | |
| raise ValueError("input_values tensor must have shape [T] or [B, T]") | |
| elif isinstance(input_values, Sequence): | |
| waveforms = tuple(input_values) | |
| else: | |
| raise TypeError("input_values must be a tensor or a sequence of tensors") | |
| if not waveforms: | |
| raise ValueError("input_values must contain at least one waveform") | |
| device = None | |
| for waveform in waveforms: | |
| if not isinstance(waveform, torch.Tensor) or waveform.ndim != 1: | |
| raise ValueError("every waveform must be a one-dimensional torch.Tensor") | |
| if waveform.numel() == 0: | |
| raise ValueError("empty waveforms are not supported") | |
| if not waveform.is_floating_point(): | |
| raise TypeError("waveforms must use a floating-point dtype") | |
| if device is None: | |
| device = waveform.device | |
| elif waveform.device != device: | |
| raise ValueError("all waveforms must be on the same device") | |
| return waveforms | |
| def _resample_waveform(waveform, original_rate, target_rate): | |
| if original_rate == target_rate: | |
| return waveform | |
| return AF.resample(waveform, original_rate, target_rate) | |
| def _segment_mean(values, counts): | |
| lengths = torch.tensor(counts, device=values.device) | |
| indices = torch.repeat_interleave( | |
| torch.arange(len(counts), device=values.device), | |
| lengths, | |
| output_size=values.shape[0], | |
| ) | |
| result = values.new_zeros((len(counts), *values.shape[1:])) | |
| result.index_add_(0, indices, values) | |
| return result.div_(lengths.reshape((-1,) + (1,) * (values.ndim - 1))) | |
| def _slaney_filterbank(n_freqs: int, n_mels: int, max_frequency: float): | |
| log_step = math.log(6.4) / 27.0 | |
| max_mel = 15.0 + math.log(8.0) / log_step | |
| mel_points = torch.linspace(0.0, max_mel, n_mels + 2) | |
| frequency_points = torch.where( | |
| mel_points >= 15.0, | |
| 1000.0 * torch.exp(log_step * (mel_points - 15.0)), | |
| (200.0 / 3.0) * mel_points, | |
| ) | |
| frequencies = torch.linspace(0.0, max_frequency, n_freqs) | |
| slopes = frequency_points.unsqueeze(0) - frequencies.unsqueeze(1) | |
| differences = frequency_points[1:] - frequency_points[:-1] | |
| filters = torch.minimum( | |
| -slopes[:, :-2] / differences[:-1], | |
| slopes[:, 2:] / differences[1:], | |
| ).clamp_min(0) | |
| normalization = 2.0 / (frequency_points[2:] - frequency_points[:-2]) | |
| return (filters * normalization.unsqueeze(0)).transpose(0, 1) | |
| def _trunc_div(numerator, denominator): | |
| return torch.div(numerator, denominator, rounding_mode="trunc") | |
| class VoiceActivityDetector(nn.Module): | |
| """Torch port of the 16 kHz, 30 ms, mode-3 WebRTC VAD path.""" | |
| def __init__(self): | |
| super().__init__() | |
| constants = { | |
| "noise_weights": [[34, 62, 72, 66, 53, 25], [94, 66, 56, 62, 75, 103]], | |
| "speech_weights": [[48, 82, 45, 87, 50, 47], [80, 46, 83, 41, 78, 81]], | |
| "noise_means": [ | |
| [6738, 4892, 7065, 6715, 6771, 3369], | |
| [7646, 3863, 7820, 7266, 5020, 4362], | |
| ], | |
| "speech_means": [ | |
| [8306, 10085, 10078, 11823, 11843, 6309], | |
| [9473, 9571, 10879, 7581, 8180, 7483], | |
| ], | |
| "noise_stds": [ | |
| [378, 1064, 493, 582, 688, 593], | |
| [474, 697, 475, 688, 421, 455], | |
| ], | |
| "speech_stds": [ | |
| [555, 505, 567, 524, 585, 1231], | |
| [509, 828, 492, 1540, 1079, 850], | |
| ], | |
| } | |
| for name, value in constants.items(): | |
| self.register_buffer(name, torch.tensor(value)) | |
| self.register_buffer("spectrum_weights", torch.tensor([6, 8, 10, 12, 14, 16])) | |
| self.register_buffer( | |
| "minimum_difference", torch.tensor([544, 544, 576, 576, 576, 576]) | |
| ) | |
| self.register_buffer( | |
| "maximum_speech", torch.tensor([11392, 11392, 11520, 11520, 11520, 11520]) | |
| ) | |
| self.register_buffer( | |
| "maximum_noise", torch.tensor([9216, 9088, 8960, 8832, 8704, 8576]) | |
| ) | |
| def _all_pass(waveform: torch.Tensor, coefficient: float): | |
| numerator = torch.tensor( | |
| [coefficient / 2, 0.5], device=waveform.device, dtype=waveform.dtype | |
| ) | |
| denominator = torch.tensor( | |
| [1.0, coefficient], device=waveform.device, dtype=waveform.dtype | |
| ) | |
| return torch.floor(AF.lfilter(waveform, denominator, numerator, clamp=False)) | |
| def _split(self, waveform): | |
| upper = self._all_pass(waveform[..., ::2], 20972 / 32768) | |
| lower = self._all_pass(waveform[..., 1::2], 5571 / 32768) | |
| return upper - lower, upper + lower | |
| def _features(self, waveform): | |
| batch_size = waveform.shape[0] | |
| quantized = torch.round(waveform * 32767).to(torch.int64) | |
| quantized = (quantized + 32768).remainder(65536) - 32768 | |
| quantized = quantized.to(waveform.dtype) | |
| upper = self._all_pass(quantized[..., ::2], 5243 / 8192) | |
| lower = self._all_pass(quantized[..., 1::2], 1392 / 8192) | |
| downsampled = upper + lower | |
| high_4, low_4 = self._split(downsampled) | |
| high_5, low_5 = self._split(high_4) | |
| high_3, low_3 = self._split(low_4) | |
| high_2, low_2 = self._split(low_3) | |
| high_1, low_1 = self._split(low_2) | |
| numerator = ( | |
| torch.tensor( | |
| [6631, -13262, 6631], device=waveform.device, dtype=waveform.dtype | |
| ) | |
| / 16384 | |
| ) | |
| denominator = ( | |
| torch.tensor( | |
| [16384, -7756, 5620], device=waveform.device, dtype=waveform.dtype | |
| ) | |
| / 16384 | |
| ) | |
| high_0 = AF.lfilter(low_1, denominator, numerator, clamp=False) | |
| bands = (high_0, high_1, high_2, high_3, low_5, high_5) | |
| lengths = (15, 15, 30, 60, 60, 60) | |
| offsets = (368, 368, 272, 176, 176, 176) | |
| features, power = [], [] | |
| for band, length, offset in zip(bands, lengths, offsets): | |
| frames = band.reshape(batch_size, -1, length) | |
| integer_frames = torch.floor(frames).to(torch.int64) | |
| maximum = integer_frames.abs().amax(-1) | |
| maximum_square = maximum.square() | |
| maximum_bit = torch.floor( | |
| torch.log2(maximum_square.clamp_min(1).double()) | |
| ).to(torch.int64) | |
| scaling = (maximum_bit + (length.bit_length() - 30)).clamp_min(0) | |
| energy = (integer_frames.square() >> scaling[:, :, None]).sum(-1) | |
| energy_bit = torch.floor(torch.log2(energy.clamp_min(1).double())).to( | |
| torch.int64 | |
| ) | |
| normalization = energy_bit - 14 | |
| normalized = torch.where( | |
| normalization < 0, | |
| energy << (-normalization).clamp_min(0), | |
| energy >> normalization.clamp_min(0), | |
| ) | |
| log2_energy = 14336 + ((normalized & 0x3FFF) >> 4) | |
| total_shifts = scaling + normalization | |
| value = ((24660 * log2_energy) >> 19) + ((total_shifts * 24660) >> 9) | |
| value = value.clamp_min(0) + offset | |
| features.append( | |
| torch.where(energy > 0, value, torch.full_like(value, offset)) | |
| ) | |
| contribution = torch.where( | |
| total_shifts >= 0, | |
| 11, | |
| normalized >> (-total_shifts).clamp_min(0), | |
| ) | |
| power.append(torch.where(energy > 0, contribution, 0)) | |
| return torch.stack(features, -1).to(torch.int64), torch.stack(power).sum(0) > 10 | |
| def _probability(features, means, stds): | |
| inverse_std = (131072 + (stds >> 1)) // stds | |
| inverse_variance = ((inverse_std >> 2).square()) >> 2 | |
| difference = (features[:, None] << 3) - means | |
| delta = (inverse_variance * difference) >> 10 | |
| exponent = (delta * difference) >> 9 | |
| logarithm = (5909 * exponent) >> 12 | |
| negative = (-logarithm).to(torch.int16).to(torch.int64) | |
| exponential = torch.bitwise_or( | |
| torch.full_like(negative, 1024), torch.bitwise_and(negative, 1023) | |
| ) | |
| complemented = torch.bitwise_xor(negative, torch.full_like(negative, 65535)) | |
| shift = (complemented.to(torch.int16).to(torch.int64) >> 10) + 1 | |
| exponential = torch.where(exponent < 22005, exponential >> shift, 0) | |
| return inverse_std * exponential, delta | |
| def _classify(self, features, active): | |
| batch_size = features.shape[0] | |
| noise_means = self.noise_means.expand(batch_size, -1, -1).clone() | |
| speech_means = self.speech_means.expand(batch_size, -1, -1).clone() | |
| noise_stds = self.noise_stds.expand(batch_size, -1, -1).clone() | |
| speech_stds = self.speech_stds.expand(batch_size, -1, -1).clone() | |
| low_values = torch.full( | |
| (batch_size, 6, 16), 10000, dtype=torch.int64, device=features.device | |
| ) | |
| ages = torch.zeros_like(low_values) | |
| mean_values = torch.full( | |
| (batch_size, 6), 1600, dtype=torch.int64, device=features.device | |
| ) | |
| overhang = torch.zeros((batch_size,), dtype=torch.int64, device=features.device) | |
| speech_count = torch.zeros_like(overhang) | |
| model_frame_count = torch.zeros_like(overhang) | |
| minimum_means = torch.tensor( | |
| [640, 768], dtype=torch.int64, device=features.device | |
| )[:, None] | |
| maximum_means = ( | |
| torch.tensor( | |
| [[72, 71, 70, 69, 68, 67], [73, 72, 71, 70, 69, 68]], | |
| dtype=torch.int64, | |
| device=features.device, | |
| ) | |
| << 7 | |
| ) | |
| speech_limits = ( | |
| torch.cat( | |
| ( | |
| torch.tensor([12800], dtype=torch.int64, device=features.device), | |
| self.maximum_speech[:-1], | |
| ) | |
| )[:, None].transpose(0, 1) | |
| + 640 | |
| ) | |
| all_active = bool(active.all()) | |
| decisions = [] | |
| for frame_index in range(features.shape[1]): | |
| frame = features[:, frame_index] | |
| has_energy = active[:, frame_index] | |
| old_noise_means = noise_means | |
| old_speech_means = speech_means | |
| old_noise_stds = noise_stds | |
| old_speech_stds = speech_stds | |
| if not all_active: | |
| old_low_values = low_values | |
| old_ages = ages | |
| old_mean_values = mean_values | |
| noise_probability, noise_delta = self._probability( | |
| frame, noise_means, noise_stds | |
| ) | |
| speech_probability, speech_delta = self._probability( | |
| frame, speech_means, speech_stds | |
| ) | |
| noise_probability = noise_probability * self.noise_weights | |
| speech_probability = speech_probability * self.speech_weights | |
| noise_total = noise_probability.sum(1) | |
| speech_total = speech_probability.sum(1) | |
| noise_log = torch.where( | |
| noise_total > 0, | |
| torch.floor(torch.log2(noise_total.double())).to(torch.int64), | |
| -1, | |
| ) | |
| speech_log = torch.where( | |
| speech_total > 0, | |
| torch.floor(torch.log2(speech_total.double())).to(torch.int64), | |
| -1, | |
| ) | |
| ratio = speech_log - noise_log | |
| decision = has_energy & ( | |
| (ratio * 4 > 94).any(1) | |
| | ((ratio * self.spectrum_weights).sum(1) >= 1100) | |
| ) | |
| noise_total_q15 = (noise_total >> 12).to(torch.int16).to(torch.int64) | |
| speech_total_q15 = (speech_total >> 12).to(torch.int16).to(torch.int64) | |
| noise_first = torch.where( | |
| noise_total_q15 > 0, | |
| _trunc_div( | |
| torch.bitwise_and(noise_probability[:, 0], 0xFFFFF000) << 2, | |
| noise_total_q15.clamp_min(1), | |
| ), | |
| 16384, | |
| ) | |
| speech_first = torch.where( | |
| speech_total_q15 > 0, | |
| _trunc_div( | |
| torch.bitwise_and(speech_probability[:, 0], 0xFFFFF000) << 2, | |
| speech_total_q15.clamp_min(1), | |
| ), | |
| 0, | |
| ) | |
| noise_condition = torch.stack((noise_first, 16384 - noise_first), dim=1) | |
| speech_condition = torch.stack( | |
| ( | |
| speech_first, | |
| torch.where(speech_total_q15 > 0, 16384 - speech_first, 0), | |
| ), | |
| dim=1, | |
| ) | |
| valid = ages < 100 | |
| candidates = torch.cat( | |
| (torch.where(valid, low_values, 10000), frame[:, :, None]), 2 | |
| ) | |
| candidate_ages = torch.cat( | |
| ( | |
| torch.where(valid, ages + 1, 101), | |
| torch.ones( | |
| (batch_size, 6, 1), | |
| dtype=torch.int64, | |
| device=features.device, | |
| ), | |
| ), | |
| 2, | |
| ) | |
| order = torch.argsort(candidates, dim=2, stable=True)[:, :, :16] | |
| low_values = torch.gather(candidates, 2, order) | |
| ages = torch.gather(candidate_ages, 2, order) | |
| if all_active: | |
| if frame_index > 2: | |
| current_median = low_values[:, :, 2] | |
| elif frame_index > 0: | |
| current_median = low_values[:, :, 0] | |
| else: | |
| current_median = torch.full_like(mean_values, 1600) | |
| if frame_index > 0: | |
| alpha = torch.where(current_median < mean_values, 6553, 32439) | |
| else: | |
| alpha = torch.zeros_like(mean_values) | |
| else: | |
| current_median = torch.where( | |
| model_frame_count[:, None] > 2, | |
| low_values[:, :, 2], | |
| torch.where( | |
| model_frame_count[:, None] > 0, low_values[:, :, 0], 1600 | |
| ), | |
| ) | |
| alpha = torch.where( | |
| model_frame_count[:, None] > 0, | |
| torch.where(current_median < mean_values, 6553, 32439), | |
| 0, | |
| ) | |
| mean_values = ( | |
| (alpha + 1) * mean_values + (32767 - alpha) * current_median + 16384 | |
| ) >> 15 | |
| noise_global = (noise_means * self.noise_weights).sum(1) | |
| noise_update = (noise_condition * noise_delta) >> 11 | |
| adapted_noise = old_noise_means + ((noise_update * 655) >> 22) | |
| adapted_noise = torch.where( | |
| decision[:, None, None], old_noise_means, adapted_noise | |
| ) | |
| correction = (mean_values << 4) - (noise_global >> 6) | |
| adapted_noise = adapted_noise + ((correction[:, None] * 154) >> 9) | |
| noise_means = torch.maximum( | |
| torch.minimum(adapted_noise, maximum_means), minimum_means | |
| ) | |
| speech_update = (speech_condition * speech_delta) >> 11 | |
| speech_step = (speech_update * 6554) >> 21 | |
| adapted_speech = old_speech_means + ((speech_step + 1) >> 1) | |
| adapted_speech = torch.maximum( | |
| torch.minimum(adapted_speech, speech_limits), minimum_means | |
| ) | |
| speech_means = torch.where( | |
| decision[:, None, None], adapted_speech, old_speech_means | |
| ) | |
| speech_error = frame[:, None] - ((old_speech_means + 4) >> 3) | |
| speech_variance = (speech_delta * speech_error) >> 3 | |
| speech_variance = (speech_condition >> 2) * (speech_variance - 4096) | |
| speech_variance = speech_variance >> 4 | |
| speech_std_step = ( | |
| _trunc_div(speech_variance.abs(), old_speech_stds * 10) | |
| * speech_variance.sign() | |
| ) | |
| adapted_speech_std = old_speech_stds + ((speech_std_step + 128) >> 8) | |
| speech_stds = torch.where( | |
| decision[:, None, None], | |
| adapted_speech_std.clamp_min(384), | |
| old_speech_stds, | |
| ) | |
| noise_error = frame[:, None] - (old_noise_means >> 3) | |
| noise_variance = (noise_delta * noise_error) >> 3 | |
| noise_variance = ((noise_condition + 2) >> 2) * (noise_variance - 4096) | |
| noise_variance = noise_variance >> 14 | |
| noise_std_step = ( | |
| _trunc_div(noise_variance.abs(), old_noise_stds) * noise_variance.sign() | |
| ) | |
| adapted_noise_std = old_noise_stds + ((noise_std_step + 32) >> 6) | |
| noise_stds = torch.where( | |
| decision[:, None, None], | |
| old_noise_stds, | |
| adapted_noise_std.clamp_min(384), | |
| ) | |
| noise_global = (noise_means * self.noise_weights).sum(1) | |
| speech_global = (speech_means * self.speech_weights).sum(1) | |
| separation = (speech_global >> 9) - (noise_global >> 9) | |
| missing = (self.minimum_difference - separation).clamp_min(0) | |
| speech_means = speech_means + (((13 * missing) >> 2)[:, None]) | |
| noise_means = noise_means - (((3 * missing) >> 2)[:, None]) | |
| speech_global = (speech_means * self.speech_weights).sum(1) | |
| noise_global = (noise_means * self.noise_weights).sum(1) | |
| speech_means = ( | |
| speech_means | |
| - ((speech_global >> 7) - self.maximum_speech).clamp_min(0)[:, None] | |
| ) | |
| noise_means = ( | |
| noise_means | |
| - ((noise_global >> 7) - self.maximum_noise).clamp_min(0)[:, None] | |
| ) | |
| if not all_active: | |
| update = has_energy[:, None, None] | |
| noise_means = torch.where(update, noise_means, old_noise_means) | |
| speech_means = torch.where(update, speech_means, old_speech_means) | |
| noise_stds = torch.where(update, noise_stds, old_noise_stds) | |
| speech_stds = torch.where(update, speech_stds, old_speech_stds) | |
| low_values = torch.where(update, low_values, old_low_values) | |
| ages = torch.where(update, ages, old_ages) | |
| mean_values = torch.where( | |
| has_energy[:, None], mean_values, old_mean_values | |
| ) | |
| model_frame_count = model_frame_count + has_energy.to(torch.int64) | |
| emitted = decision | (overhang > 0) | |
| next_count = speech_count + 1 | |
| speech_count = torch.where( | |
| decision, next_count.clamp_max(6), torch.zeros_like(speech_count) | |
| ) | |
| speech_overhang = torch.where( | |
| next_count > 6, | |
| torch.full_like(overhang, 3), | |
| torch.full_like(overhang, 2), | |
| ) | |
| overhang = torch.where( | |
| decision, speech_overhang, (overhang - 1).clamp_min(0) | |
| ) | |
| decisions.append(emitted) | |
| return torch.stack(decisions, dim=1) | |
| def forward(self, waveform, input_lengths=None): | |
| if waveform.ndim != 2: | |
| raise ValueError("VAD input must have shape [B, T]") | |
| if input_lengths is None: | |
| input_lengths = (waveform.shape[1],) * waveform.shape[0] | |
| usable_lengths = tuple(length // 480 * 480 for length in input_lengths) | |
| maximum = max(usable_lengths) | |
| if maximum == 0: | |
| return tuple(item[:0] for item in waveform) | |
| waveform = waveform[:, :maximum] | |
| features, active = self._features(waveform.float()) | |
| frame_lengths = torch.tensor(usable_lengths, device=waveform.device) // 480 | |
| valid_frames = ( | |
| torch.arange(features.shape[1], device=waveform.device)[None] | |
| < frame_lengths[:, None] | |
| ) | |
| flags = self._classify(features, active & valid_frames) & valid_frames | |
| flags = flags.to(waveform.dtype) | |
| averaged = F.avg_pool1d(F.pad(flags[:, None], (3, 4)), 8, stride=1)[:, 0] | |
| smoothed = torch.round(averaged).to(torch.bool) | |
| dilated = F.max_pool1d( | |
| smoothed.to(waveform.dtype)[:, None], 7, stride=1, padding=3 | |
| )[:, 0].to(torch.bool) | |
| masks = dilated.repeat_interleave(480, dim=1) | |
| return tuple( | |
| item[:length][mask[:length]] | |
| for item, mask, length in zip(waveform, masks, usable_lengths) | |
| ) | |
| class ResemblyzerConfig(PretrainedConfig): | |
| model_type = "resemblyzer" | |
| class ResemblyzerOutput(ModelOutput): | |
| embeddings: torch.Tensor | None = None | |
| class MelFrontend(nn.Module): | |
| def __init__(self): | |
| super().__init__() | |
| self.register_buffer("window", torch.hann_window(400)) | |
| self.register_buffer("mel_filters", _slaney_filterbank(201, 40, 8_000.0)) | |
| def forward(self, waveform): | |
| spectrum = torch.stft( | |
| waveform, | |
| 400, | |
| 160, | |
| 400, | |
| self.window, | |
| center=True, | |
| pad_mode="constant", | |
| return_complex=True, | |
| ) | |
| return torch.matmul(self.mel_filters, spectrum.abs().square()).transpose(0, 1) | |
| class ResemblyzerModel(PreTrainedModel): | |
| config_class = ResemblyzerConfig | |
| main_input_name = "input_values" | |
| def __init__(self, config): | |
| super().__init__(config) | |
| self.lstm = nn.LSTM(40, 256, 3, batch_first=True, dropout=0.0) | |
| self.linear = nn.Linear(256, 256) | |
| self.mel_frontend = MelFrontend() | |
| self.voice_activity_detector = VoiceActivityDetector() | |
| self.post_init() | |
| def train(self, mode: bool = True): | |
| super().train(mode) | |
| # Dropout is zero, so this does not change the output. Keeping only the | |
| # LSTM in training mode makes cuDNN retain the state required by backward. | |
| self.lstm.train() | |
| return self | |
| def _normalize_volume(waveform): | |
| rms = waveform.square().mean().sqrt() | |
| target = waveform.new_tensor(10 ** (-30 / 20)) | |
| gain = target / rms.clamp_min(torch.finfo(waveform.dtype).tiny) | |
| normalized = waveform * torch.maximum(gain, waveform.new_tensor(1.0)) | |
| return torch.where(rms > 0, normalized, waveform) | |
| def _partial_mels(self, waveform): | |
| samples_per_frame = 160 | |
| n_samples = waveform.shape[-1] | |
| n_frames = math.ceil((n_samples + 1) / samples_per_frame) | |
| frame_step = round((_TARGET_SAMPLING_RATE / _PARTIAL_RATE) / samples_per_frame) | |
| steps = max(1, n_frames - 160 + frame_step + 1) | |
| starts = list(range(0, steps, frame_step)) | |
| coverage = (n_samples - starts[-1] * samples_per_frame) / ( | |
| 160 * samples_per_frame | |
| ) | |
| if coverage < _MIN_COVERAGE and len(starts) > 1: | |
| starts.pop() | |
| maximum = (starts[-1] + 160) * samples_per_frame | |
| waveform = F.pad(waveform, (0, max(0, maximum - n_samples))) | |
| mel = self.mel_frontend(waveform) | |
| return mel.unfold(0, 160, frame_step)[: len(starts)].transpose(1, 2) | |
| def _trim_silences(self, waveforms): | |
| usable_lengths = tuple(waveform.shape[0] // 480 * 480 for waveform in waveforms) | |
| maximum = max(usable_lengths) | |
| batch = torch.stack( | |
| [ | |
| F.pad(waveform[:length], (0, maximum - length)) | |
| for waveform, length in zip(waveforms, usable_lengths) | |
| ] | |
| ) | |
| return self.voice_activity_detector(batch, usable_lengths) | |
| def _encode(self, mels): | |
| _, (hidden, _) = self.lstm(mels) | |
| raw = F.relu(self.linear(hidden[-1])) | |
| return F.normalize(raw, dim=1) | |
| def forward( | |
| self, | |
| input_values: torch.Tensor | Sequence[torch.Tensor], | |
| sampling_rate: int, | |
| input_lengths: torch.Tensor | None = None, | |
| *, | |
| return_dict: bool | None = None, | |
| ) -> ResemblyzerOutput | tuple[torch.Tensor]: | |
| sampling_rate = _validate_sampling_rate(sampling_rate) | |
| return_dict = self.config.return_dict if return_dict is None else return_dict | |
| waveforms = _as_waveform_sequence(input_values, input_lengths) | |
| dtype = self.linear.weight.dtype | |
| waveforms = tuple( | |
| self._normalize_volume( | |
| _resample_waveform( | |
| waveform.float(), | |
| sampling_rate, | |
| _TARGET_SAMPLING_RATE, | |
| ) | |
| ) | |
| for waveform in waveforms | |
| ) | |
| waveforms = self._trim_silences(waveforms) | |
| if any(waveform.numel() == 0 for waveform in waveforms): | |
| raise ValueError("voice activity detection removed the entire waveform") | |
| waveforms = tuple(waveform.to(dtype=dtype) for waveform in waveforms) | |
| partials, counts = [], [] | |
| for waveform in waveforms: | |
| item = self._partial_mels(waveform) | |
| partials.append(item) | |
| counts.append(item.shape[0]) | |
| partials = torch.cat(partials) | |
| encoded = self._encode(partials) | |
| result = F.normalize( | |
| _segment_mean(encoded, counts), | |
| dim=1, | |
| ) | |
| if not return_dict: | |
| return (result,) | |
| return ResemblyzerOutput(embeddings=result) | |