suvradeepp's picture
Publish Tiny Hinglish Turn Detector development preview
35d483e verified
Raw
History Blame Contribute Delete
8.22 kB
"""Torch-native waveform to log-mel feature extraction.
The frontend intentionally avoids torchaudio so the deployable student has one
fewer binary dependency. Exported production models normally accept log-mel
features; keeping this implementation in the repository gives training, demo,
and parity tests one canonical preprocessing contract.
"""
from __future__ import annotations
from collections.abc import Mapping
from dataclasses import dataclass
import torch
import torch.nn.functional as F
from torch import Tensor, nn
@dataclass(frozen=True)
class LogMelConfig:
sample_rate: int = 16_000
n_fft: int = 400
hop_length: int = 160
win_length: int = 400
n_mels: int = 80
f_min: float = 0.0
f_max: float | None = 8_000.0
log_floor: float = 1e-10
normalize: bool = True
mel_scale: str = "htk"
log_scale: str = "standard"
center: bool = False
drop_last_frame: bool = False
pad_side: str = "right"
@classmethod
def from_mapping(cls, values: Mapping[str, object]) -> LogMelConfig:
known = {field.name for field in cls.__dataclass_fields__.values()}
return cls(**{key: value for key, value in values.items() if key in known}) # type: ignore[arg-type]
def _hz_to_mel(freq: Tensor) -> Tensor:
# HTK mel convention. It is stable, simple, and matches common speech
# frontends closely enough for a model trained with the same extractor.
return 2595.0 * torch.log10(1.0 + freq / 700.0)
def _mel_to_hz(mels: Tensor) -> Tensor:
return 700.0 * (torch.pow(10.0, mels / 2595.0) - 1.0)
def _hz_to_slaney_mel(freq: Tensor) -> Tensor:
linear_spacing = 200.0 / 3.0
min_log_hz = 1_000.0
min_log_mel = min_log_hz / linear_spacing
log_step = torch.log(torch.tensor(6.4, dtype=freq.dtype, device=freq.device)) / 27.0
linear = freq / linear_spacing
logarithmic = min_log_mel + torch.log((freq / min_log_hz).clamp_min(1e-12)) / log_step
return torch.where(freq >= min_log_hz, logarithmic, linear)
def _slaney_mel_to_hz(mels: Tensor) -> Tensor:
linear_spacing = 200.0 / 3.0
min_log_hz = 1_000.0
min_log_mel = min_log_hz / linear_spacing
log_step = torch.log(torch.tensor(6.4, dtype=mels.dtype, device=mels.device)) / 27.0
linear = mels * linear_spacing
logarithmic = min_log_hz * torch.exp(log_step * (mels - min_log_mel))
return torch.where(mels >= min_log_mel, logarithmic, linear)
def make_mel_filterbank(config: LogMelConfig) -> Tensor:
"""Create a ``[n_mels, n_fft // 2 + 1]`` triangular filterbank."""
max_hz = float(config.sample_rate / 2 if config.f_max is None else config.f_max)
if not 0.0 <= config.f_min < max_hz <= config.sample_rate / 2:
raise ValueError("expected 0 <= f_min < f_max <= sample_rate / 2")
fft_freqs = torch.linspace(0.0, config.sample_rate / 2, config.n_fft // 2 + 1)
if config.mel_scale == "htk":
to_mel, to_hz = _hz_to_mel, _mel_to_hz
elif config.mel_scale == "slaney":
to_mel, to_hz = _hz_to_slaney_mel, _slaney_mel_to_hz
else:
raise ValueError("mel_scale must be 'htk' or 'slaney'")
mel_edges = torch.linspace(
to_mel(torch.tensor(float(config.f_min))),
to_mel(torch.tensor(max_hz)),
config.n_mels + 2,
)
hz_edges = to_hz(mel_edges)
lower = hz_edges[:-2, None]
center = hz_edges[1:-1, None]
upper = hz_edges[2:, None]
rising = (fft_freqs[None, :] - lower) / (center - lower).clamp_min(1e-12)
falling = (upper - fft_freqs[None, :]) / (upper - center).clamp_min(1e-12)
filters = torch.minimum(rising, falling).clamp_min(0.0)
# Area normalization reduces sensitivity to mel-band width.
enorm = 2.0 / (upper - lower).clamp_min(1e-12)
return filters * enorm
class LogMelFrontend(nn.Module):
"""Convert padded mono waveforms to normalized log-mel features.
Parameters
----------
waveforms:
Float tensor shaped ``[batch, samples]`` (or ``[samples]``).
lengths:
Optional valid sample counts. The returned mask is ``[batch, frames]``.
"""
def __init__(self, config: LogMelConfig | None = None) -> None:
super().__init__()
config = config or LogMelConfig()
self.config = config
if config.pad_side not in {"left", "right"}:
raise ValueError("pad_side must be 'left' or 'right'")
# numpy.hanning in the dependency-light runtime uses a symmetric window.
self.register_buffer(
"window", torch.hann_window(config.win_length, periodic=False), persistent=False
)
self.register_buffer("mel_filters", make_mel_filterbank(config), persistent=True)
def forward(self, waveforms: Tensor, lengths: Tensor | None = None) -> tuple[Tensor, Tensor]:
if waveforms.ndim == 1:
waveforms = waveforms.unsqueeze(0)
if waveforms.ndim != 2:
raise ValueError("waveforms must have shape [batch, samples]")
batch, original_samples = waveforms.shape
if lengths is None:
lengths = torch.full(
(batch,), original_samples, dtype=torch.long, device=waveforms.device
)
else:
lengths = lengths.to(device=waveforms.device, dtype=torch.long).clamp(
min=0, max=original_samples
)
if original_samples < self.config.n_fft:
waveforms = F.pad(waveforms, (0, self.config.n_fft - original_samples))
spectrum = torch.stft(
waveforms,
n_fft=self.config.n_fft,
hop_length=self.config.hop_length,
win_length=self.config.win_length,
window=self.window.to(dtype=waveforms.dtype),
center=self.config.center,
return_complex=True,
)
if self.config.drop_last_frame:
spectrum = spectrum[..., :-1]
power = spectrum.abs().square()
mel = torch.matmul(self.mel_filters.to(dtype=power.dtype), power)
log_mel = torch.log10(mel.clamp_min(self.config.log_floor))
if self.config.log_scale == "whisper":
dynamic_floor = log_mel.amax(dim=(-2, -1), keepdim=True) - 8.0
log_mel = torch.maximum(log_mel, dynamic_floor)
log_mel = (log_mel + 4.0) / 4.0
elif self.config.log_scale != "standard":
raise ValueError("log_scale must be 'standard' or 'whisper'")
if self.config.center:
if self.config.drop_last_frame:
frame_lengths = torch.div(
lengths + self.config.hop_length - 1,
self.config.hop_length,
rounding_mode="floor",
)
else:
frame_lengths = 1 + torch.div(
lengths, self.config.hop_length, rounding_mode="floor"
)
else:
padded_lengths = lengths.clamp_min(self.config.n_fft)
frame_lengths = 1 + torch.div(
padded_lengths - self.config.n_fft,
self.config.hop_length,
rounding_mode="floor",
)
if self.config.drop_last_frame and not self.config.center:
frame_lengths = (frame_lengths - 1).clamp_min(0)
frame_lengths = torch.where(lengths > 0, frame_lengths, torch.zeros_like(frame_lengths))
frame_lengths = frame_lengths.clamp(max=log_mel.shape[-1])
positions = torch.arange(log_mel.shape[-1], device=waveforms.device)
if self.config.pad_side == "left":
mask = positions.unsqueeze(0) >= (log_mel.shape[-1] - frame_lengths).unsqueeze(1)
else:
mask = positions.unsqueeze(0) < frame_lengths.unsqueeze(1)
if self.config.normalize:
valid = mask.unsqueeze(1).to(log_mel.dtype)
denominator = valid.sum(dim=-1, keepdim=True).clamp_min(1.0)
mean = (log_mel * valid).sum(dim=-1, keepdim=True) / denominator
variance = ((log_mel - mean).square() * valid).sum(dim=-1, keepdim=True)
variance = variance / denominator
log_mel = (log_mel - mean) * torch.rsqrt(variance + 1e-5)
log_mel = log_mel * valid
return log_mel, mask