"""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