Download src/turn_detection/runtime/features.py from suvradeepp/tiny-hinglish-turn-detector: direct link, hf CLI and curl.
- Browser
- Download file 8.25 kB
-
https://huggingface.co/suvradeepp/tiny-hinglish-turn-detector/resolve/main/src/turn_detection/runtime/features.py
- Command line
-
hf download hf://suvradeepp/tiny-hinglish-turn-detector/src/turn_detection/runtime/features.py
-
curl -L -o features.py https://huggingface.co/suvradeepp/tiny-hinglish-turn-detector/resolve/main/src/turn_detection/runtime/features.py
8.25 kB
| """Dependency-light audio normalization and log-mel feature extraction. | |
| The frontend configuration is serialized next to every exported model. Keeping | |
| one implementation for export validation and inference prevents a very common | |
| failure mode: a correct ONNX graph fed subtly different features in production. | |
| """ | |
| from __future__ import annotations | |
| from dataclasses import asdict, dataclass | |
| from functools import lru_cache | |
| from typing import Any | |
| def _numpy() -> Any: | |
| try: | |
| import numpy as np | |
| except ImportError as exc: # pragma: no cover - exercised in minimal installs | |
| raise RuntimeError("Audio inference requires numpy; install the base package") from exc | |
| return np | |
| class FrontendConfig: | |
| sample_rate: int = 16_000 | |
| max_seconds: float = 8.0 | |
| n_fft: int = 400 | |
| win_length: int = 400 | |
| hop_length: int = 160 | |
| n_mels: int = 80 | |
| f_min: float = 0.0 | |
| f_max: float = 8_000.0 | |
| normalization: str = "whisper" | |
| pad_side: str = "left" | |
| def __post_init__(self) -> None: | |
| if self.sample_rate <= 0 or self.max_seconds <= 0: | |
| raise ValueError("sample_rate and max_seconds must be positive") | |
| if self.n_fft <= 0 or self.win_length <= 0 or self.hop_length <= 0: | |
| raise ValueError("FFT and window sizes must be positive") | |
| if self.win_length > self.n_fft: | |
| raise ValueError("win_length cannot exceed n_fft") | |
| if self.n_mels <= 0: | |
| raise ValueError("n_mels must be positive") | |
| if not 0.0 <= self.f_min < self.f_max <= self.sample_rate / 2: | |
| raise ValueError("mel frequency bounds must lie inside Nyquist") | |
| if self.normalization not in {"whisper", "log10", "none"}: | |
| raise ValueError("unsupported normalization") | |
| if self.pad_side not in {"left", "right"}: | |
| raise ValueError("pad_side must be 'left' or 'right'") | |
| def max_samples(self) -> int: | |
| return round(self.sample_rate * self.max_seconds) | |
| def target_frames(self) -> int: | |
| return self.max_samples // self.hop_length | |
| def to_dict(self) -> dict[str, Any]: | |
| return asdict(self) | |
| def normalize_waveform(audio: Any) -> Any: | |
| """Convert mono/stereo integer/float audio to finite mono float32 in [-1, 1].""" | |
| np = _numpy() | |
| samples = np.asarray(audio) | |
| if samples.size == 0: | |
| raise ValueError("audio cannot be empty") | |
| original_dtype = samples.dtype | |
| if samples.ndim == 2: | |
| # Gradio commonly returns [samples, channels]; accept [channels, samples] too. | |
| channel_axis = 1 if samples.shape[1] <= 8 else 0 | |
| samples = samples.astype(np.float32).mean(axis=channel_axis) | |
| elif samples.ndim != 1: | |
| raise ValueError(f"expected mono/stereo audio, got shape {samples.shape}") | |
| if np.issubdtype(original_dtype, np.integer): | |
| info = np.iinfo(original_dtype) | |
| scale = float(max(abs(info.min), info.max)) | |
| samples = samples.astype(np.float32) / scale | |
| else: | |
| samples = samples.astype(np.float32, copy=False) | |
| samples = np.nan_to_num(samples, nan=0.0, posinf=1.0, neginf=-1.0) | |
| peak = float(np.max(np.abs(samples))) | |
| if peak > 1.0: | |
| samples = samples / peak | |
| return np.clip(samples, -1.0, 1.0) | |
| def resample_waveform(audio: Any, source_rate: int, target_rate: int) -> Any: | |
| """Resample with the same deterministic linear rule used during training. | |
| Optional packages must not change model inputs, so this deliberately avoids | |
| a SciPy-dependent branch. Higher-quality telephony resampling can happen | |
| upstream, but training and serving always agree at this boundary. | |
| """ | |
| np = _numpy() | |
| if source_rate <= 0 or target_rate <= 0: | |
| raise ValueError("sample rates must be positive") | |
| samples = normalize_waveform(audio) | |
| if source_rate == target_rate: | |
| return samples | |
| output_length = max(1, round(len(samples) * target_rate / source_rate)) | |
| old_x = np.linspace(0.0, 1.0, len(samples), endpoint=False) | |
| new_x = np.linspace(0.0, 1.0, output_length, endpoint=False) | |
| return np.interp(new_x, old_x, samples).astype(np.float32) | |
| def pad_or_trim(audio: Any, config: FrontendConfig) -> tuple[Any, int]: | |
| """Return fixed-length audio and number of genuine (non-padding) samples.""" | |
| np = _numpy() | |
| samples = normalize_waveform(audio) | |
| if len(samples) >= config.max_samples: | |
| return samples[-config.max_samples :].copy(), config.max_samples | |
| pad = config.max_samples - len(samples) | |
| widths = (pad, 0) if config.pad_side == "left" else (0, pad) | |
| return np.pad(samples, widths).astype(np.float32), len(samples) | |
| def _hz_to_mel(value: Any) -> Any: | |
| np = _numpy() | |
| return 2595.0 * np.log10(1.0 + np.asarray(value) / 700.0) | |
| def _mel_to_hz(value: Any) -> Any: | |
| np = _numpy() | |
| return 700.0 * (10.0 ** (np.asarray(value) / 2595.0) - 1.0) | |
| def mel_filterbank(config: FrontendConfig) -> Any: | |
| """Create a deterministic triangular mel filter bank.""" | |
| np = _numpy() | |
| mel_points = np.linspace(_hz_to_mel(config.f_min), _hz_to_mel(config.f_max), config.n_mels + 2) | |
| hz_points = _mel_to_hz(mel_points) | |
| fft_hz = np.linspace(0.0, config.sample_rate / 2, config.n_fft // 2 + 1) | |
| filters = np.zeros((config.n_mels, len(fft_hz)), dtype=np.float32) | |
| for index in range(config.n_mels): | |
| left, center, right = hz_points[index : index + 3] | |
| filters[index] = np.maximum( | |
| 0.0, | |
| np.minimum( | |
| (fft_hz - left) / max(center - left, 1e-12), | |
| (right - fft_hz) / max(right - center, 1e-12), | |
| ), | |
| ) | |
| # Area normalization reduces frequency-dependent scale drift. | |
| enorm = 2.0 / np.maximum(hz_points[2 : config.n_mels + 2] - hz_points[: config.n_mels], 1e-12) | |
| result = filters * enorm[:, None] | |
| result.flags.writeable = False | |
| return result | |
| def _hann_window(length: int) -> Any: | |
| np = _numpy() | |
| window = np.hanning(length).astype(np.float32) | |
| window.flags.writeable = False | |
| return window | |
| def log_mel_spectrogram( | |
| audio: Any, | |
| source_rate: int, | |
| config: FrontendConfig | None = None, | |
| ) -> tuple[Any, Any]: | |
| """Return ``[n_mels, frames]`` features and a valid-frame mask. | |
| The implementation follows Whisper's centered-STFT and dynamic-range | |
| normalization convention closely, while remaining free of torch/librosa at | |
| inference time. Export parity tests compare it against the training path. | |
| """ | |
| np = _numpy() | |
| cfg = config or FrontendConfig() | |
| if source_rate <= 0: | |
| raise ValueError("source sample rate must be positive") | |
| # Bound work before resampling: an uploaded meeting can be hours long, while | |
| # endpoint intent uses only the configured suffix. Training uses the same rule. | |
| normalized = normalize_waveform(audio) | |
| source_suffix_samples = max(1, round(cfg.max_seconds * source_rate)) | |
| normalized = normalized[-source_suffix_samples:] | |
| resampled = resample_waveform(normalized, source_rate, cfg.sample_rate) | |
| fixed, valid_samples = pad_or_trim(resampled, cfg) | |
| pad = cfg.n_fft // 2 | |
| padded = np.pad(fixed, (pad, pad), mode="reflect") | |
| frames = np.lib.stride_tricks.sliding_window_view(padded, cfg.win_length)[:: cfg.hop_length] | |
| frames = frames[: cfg.target_frames] | |
| window = _hann_window(cfg.win_length) | |
| spectrum = np.fft.rfft(frames * window[None, :], n=cfg.n_fft, axis=1) | |
| power = (spectrum.real**2 + spectrum.imag**2).astype(np.float32) | |
| mel = np.maximum(mel_filterbank(cfg) @ power.T, 1e-10) | |
| features = np.log10(mel) | |
| if cfg.normalization == "whisper": | |
| features = np.maximum(features, float(features.max()) - 8.0) | |
| features = (features + 4.0) / 4.0 | |
| elif cfg.normalization == "none": | |
| features = mel | |
| frame_mask = np.zeros(cfg.target_frames, dtype=np.float32) | |
| valid_frames = min( | |
| cfg.target_frames, max(1, (valid_samples + cfg.hop_length - 1) // cfg.hop_length) | |
| ) | |
| if cfg.pad_side == "left": | |
| frame_mask[-valid_frames:] = 1.0 | |
| else: | |
| frame_mask[:valid_frames] = 1.0 | |
| return features.astype(np.float32), frame_mask | |