"""Standalone ONNX inference for the Kaggle model bundle. This file deliberately has no dependency on the source repository. It mirrors the canonical NumPy frontend used during export and reads every serving choice from the adjacent ``model_metadata.json`` file. """ from __future__ import annotations import json import math import time from dataclasses import dataclass from functools import lru_cache from pathlib import Path from typing import Any import numpy as np @dataclass(frozen=True, slots=True) class FrontendConfig: sample_rate: int = 16_000 max_seconds: float = 4.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") @property def max_samples(self) -> int: return round(self.sample_rate * self.max_seconds) @property def target_frames(self) -> int: return self.max_samples // self.hop_length @dataclass(frozen=True, slots=True) class ControllerConfig: endpoint_threshold: float long_pause_threshold: float min_silence_ms: float relax_after_ms: float max_silence_ms: float required_confirmations: int def __post_init__(self) -> None: if not 0.0 <= self.long_pause_threshold <= self.endpoint_threshold <= 1.0: raise ValueError("controller thresholds are invalid") if not self.min_silence_ms <= self.relax_after_ms <= self.max_silence_ms: raise ValueError("controller silence bounds are invalid") if self.required_confirmations < 1: raise ValueError("required_confirmations must be positive") def normalize_waveform(audio: Any) -> np.ndarray: """Convert mono/stereo integer/float audio to finite mono float32.""" samples = np.asarray(audio) if samples.size == 0: raise ValueError("audio cannot be empty") original_dtype = samples.dtype if samples.ndim == 2: 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) samples = samples.astype(np.float32) / float(max(abs(info.min), info.max)) 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) -> np.ndarray: """Apply the deterministic linear resampler used during training.""" 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[np.ndarray, int]: samples = normalize_waveform(audio) if len(samples) >= config.max_samples: return samples[-config.max_samples :].copy(), config.max_samples padding = config.max_samples - len(samples) widths = (padding, 0) if config.pad_side == "left" else (0, padding) return np.pad(samples, widths).astype(np.float32), len(samples) def _hz_to_mel(value: Any) -> np.ndarray: return 2595.0 * np.log10(1.0 + np.asarray(value) / 700.0) def _mel_to_hz(value: Any) -> np.ndarray: return 700.0 * (10.0 ** (np.asarray(value) / 2595.0) - 1.0) @lru_cache(maxsize=16) def mel_filterbank(config: FrontendConfig) -> np.ndarray: 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), ), ) normalization = 2.0 / np.maximum( hz_points[2 : config.n_mels + 2] - hz_points[: config.n_mels], 1e-12, ) result = filters * normalization[:, None] result.flags.writeable = False return result @lru_cache(maxsize=16) def _hann_window(length: int) -> np.ndarray: window = np.hanning(length).astype(np.float32) window.flags.writeable = False return window def log_mel_spectrogram( audio: Any, source_rate: int, config: FrontendConfig, ) -> tuple[np.ndarray, np.ndarray]: """Return canonical ``[80, frames]`` features and valid-frame mask.""" if source_rate <= 0: raise ValueError("source sample rate must be positive") normalized = normalize_waveform(audio) source_suffix_samples = max(1, round(config.max_seconds * source_rate)) normalized = normalized[-source_suffix_samples:] resampled = resample_waveform(normalized, source_rate, config.sample_rate) fixed, valid_samples = pad_or_trim(resampled, config) padding = config.n_fft // 2 padded = np.pad(fixed, (padding, padding), mode="reflect") frames = np.lib.stride_tricks.sliding_window_view(padded, config.win_length)[ :: config.hop_length ] frames = frames[: config.target_frames] spectrum = np.fft.rfft( frames * _hann_window(config.win_length)[None, :], n=config.n_fft, axis=1, ) power = (spectrum.real**2 + spectrum.imag**2).astype(np.float32) mel = np.maximum(mel_filterbank(config) @ power.T, 1e-10) features = np.log10(mel) if config.normalization == "whisper": features = np.maximum(features, float(features.max()) - 8.0) features = (features + 4.0) / 4.0 elif config.normalization == "none": features = mel frame_mask = np.zeros(config.target_frames, dtype=np.float32) valid_frames = min( config.target_frames, max(1, (valid_samples + config.hop_length - 1) // config.hop_length), ) if config.pad_side == "left": frame_mask[-valid_frames:] = 1.0 else: frame_mask[:valid_frames] = 1.0 return features.astype(np.float32), frame_mask class TurnDetector: """Load the adjacent ONNX artifact and score VAD pause checkpoints.""" def __init__(self, model_directory: str | Path) -> None: try: import onnxruntime as ort except ImportError as exc: raise RuntimeError("Install requirements.txt before inference") from exc self.model_directory = Path(model_directory).expanduser().resolve() model_path = self.model_directory / "model.onnx" metadata_path = self.model_directory / "model_metadata.json" if not model_path.is_file() or not metadata_path.is_file(): raise FileNotFoundError("model.onnx and model_metadata.json must be adjacent") self.metadata = json.loads(metadata_path.read_text(encoding="utf-8")) self.frontend = FrontendConfig(**self.metadata["frontend"]) self.controller = ControllerConfig(**self.metadata["controller"]) options = ort.SessionOptions() options.intra_op_num_threads = 1 options.inter_op_num_threads = 1 options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL self.session = ort.InferenceSession( str(model_path), sess_options=options, providers=["CPUExecutionProvider"], ) expected_inputs = { self.metadata["input_features_name"], self.metadata["frame_mask_name"], } actual_inputs = {item.name for item in self.session.get_inputs()} if actual_inputs != expected_inputs: raise ValueError(f"unexpected ONNX inputs: {sorted(actual_inputs)}") expected_output = self.metadata["endpoint_output_name"] actual_outputs = [item.name for item in self.session.get_outputs()] if actual_outputs != [expected_output]: raise ValueError(f"unexpected ONNX outputs: {actual_outputs}") def threshold_for_silence(self, silence_ms: float) -> float: if silence_ms <= self.controller.relax_after_ms: return self.controller.endpoint_threshold span = self.controller.max_silence_ms - self.controller.relax_after_ms if span <= 0.0: return self.controller.long_pause_threshold progress = min(1.0, (silence_ms - self.controller.relax_after_ms) / span) delta = self.controller.endpoint_threshold - self.controller.long_pause_threshold return self.controller.endpoint_threshold - progress * delta def predict(self, audio: Any, sample_rate: int, *, silence_ms: float = 300.0) -> dict[str, Any]: """Return a stateless HOLD/END decision for one pause checkpoint.""" if silence_ms < 0.0: raise ValueError("silence_ms cannot be negative") started = time.perf_counter_ns() features, frame_mask = log_mel_spectrogram(audio, sample_rate, self.frontend) raw = self.session.run( [self.metadata["endpoint_output_name"]], { self.metadata["input_features_name"]: features[None, :, :], self.metadata["frame_mask_name"]: frame_mask[None, :], }, )[0] value = float(np.asarray(raw).reshape(-1)[0]) probability = ( 1.0 / (1.0 + math.exp(-value)) if self.metadata["output_type"] == "logits" else value ) probability = min(1.0, max(0.0, probability)) threshold = self.threshold_for_silence(float(silence_ms)) if silence_ms < self.controller.min_silence_ms: state, reason = "HOLD", "minimum_silence_not_reached" elif silence_ms >= self.controller.max_silence_ms: state, reason = "END", "maximum_timeout" elif probability >= threshold: state, reason = "END", "model_endpoint" else: state, reason = "HOLD", "model_hold" return { "state": state, "reason": reason, "p_end": probability, "threshold": threshold, "silence_ms": float(silence_ms), "inference_ms": (time.perf_counter_ns() - started) / 1_000_000, "model_name": self.metadata["model_name"], "development_only": bool(self.metadata.get("development_only", False)), } def predict_file(self, audio_path: str | Path, *, silence_ms: float = 300.0) -> dict[str, Any]: try: import soundfile as sf except ImportError as exc: raise RuntimeError("Install soundfile to load audio paths") from exc samples, sample_rate = sf.read(str(audio_path), dtype="float32", always_2d=False) return self.predict(samples, int(sample_rate), silence_ms=silence_ms)