Download src/turn_detection/runtime/controller.py from suvradeepp/tiny-hinglish-turn-detector: direct link, hf CLI and curl.
- Browser
- Download file 6.36 kB
-
https://huggingface.co/suvradeepp/tiny-hinglish-turn-detector/resolve/main/src/turn_detection/runtime/controller.py
- Command line
-
hf download hf://suvradeepp/tiny-hinglish-turn-detector/src/turn_detection/runtime/controller.py
-
curl -L -o controller.py https://huggingface.co/suvradeepp/tiny-hinglish-turn-detector/resolve/main/src/turn_detection/runtime/controller.py
6.36 kB
| """A deterministic controller around an audio endpoint classifier. | |
| The neural model is deliberately not allowed to own the whole product policy. | |
| It is queried only at VAD pause checkpoints. The controller enforces minimum | |
| and maximum silence bounds, optional debounce, and a gradual long-pause | |
| threshold relaxation. This keeps latency/false-interruption trade-offs explicit | |
| and testable. | |
| """ | |
| from __future__ import annotations | |
| from dataclasses import dataclass | |
| from .types import Prediction, TurnDecision, TurnState | |
| class ControllerConfig: | |
| endpoint_threshold: float = 0.60 | |
| long_pause_threshold: float = 0.42 | |
| min_silence_ms: float = 200.0 | |
| relax_after_ms: float = 800.0 | |
| max_silence_ms: float = 1800.0 | |
| required_confirmations: int = 1 | |
| def __post_init__(self) -> None: | |
| for name, value in ( | |
| ("endpoint_threshold", self.endpoint_threshold), | |
| ("long_pause_threshold", self.long_pause_threshold), | |
| ): | |
| if not 0.0 <= value <= 1.0: | |
| raise ValueError(f"{name} must be in [0, 1]") | |
| if self.long_pause_threshold > self.endpoint_threshold: | |
| raise ValueError("long_pause_threshold cannot exceed endpoint_threshold") | |
| if self.min_silence_ms < 0.0: | |
| raise ValueError("min_silence_ms cannot be negative") | |
| if not self.min_silence_ms <= self.relax_after_ms <= self.max_silence_ms: | |
| raise ValueError( | |
| "silence bounds must satisfy min_silence_ms <= relax_after_ms <= max_silence_ms" | |
| ) | |
| if self.required_confirmations < 1: | |
| raise ValueError("required_confirmations must be at least one") | |
| class TurnController: | |
| """Stateful SPEAKING/HOLD/END controller. | |
| Call :meth:`observe_speech` whenever VAD sees speech. Call | |
| :meth:`evaluate_pause` at silence checkpoints. A new speech event after an | |
| END decision starts a fresh turn automatically. | |
| """ | |
| def __init__(self, config: ControllerConfig | None = None) -> None: | |
| self.config = config or ControllerConfig() | |
| self._state = TurnState.SPEAKING | |
| self._confirmations = 0 | |
| self._last_timestamp_ms: float | None = None | |
| def state(self) -> TurnState: | |
| return self._state | |
| def reset(self) -> None: | |
| self._state = TurnState.SPEAKING | |
| self._confirmations = 0 | |
| self._last_timestamp_ms = None | |
| def observe_speech(self, timestamp_ms: float | None = None) -> TurnDecision: | |
| self._validate_timestamp(timestamp_ms) | |
| self._state = TurnState.SPEAKING | |
| self._confirmations = 0 | |
| return TurnDecision( | |
| state=self._state, | |
| endpoint_probability=None, | |
| threshold=None, | |
| silence_ms=0.0, | |
| reason="speech_observed", | |
| timestamp_ms=timestamp_ms, | |
| ) | |
| def threshold_for_silence(self, silence_ms: float) -> float: | |
| """Return the decision threshold at the current silence duration.""" | |
| if silence_ms <= self.config.relax_after_ms: | |
| return self.config.endpoint_threshold | |
| span = self.config.max_silence_ms - self.config.relax_after_ms | |
| if span <= 0.0: | |
| return self.config.long_pause_threshold | |
| progress = min(1.0, (silence_ms - self.config.relax_after_ms) / span) | |
| delta = self.config.endpoint_threshold - self.config.long_pause_threshold | |
| return self.config.endpoint_threshold - progress * delta | |
| def evaluate_pause( | |
| self, | |
| prediction: Prediction | float, | |
| silence_ms: float, | |
| timestamp_ms: float | None = None, | |
| ) -> TurnDecision: | |
| """Combine a model score and silence duration into a product decision.""" | |
| self._validate_timestamp(timestamp_ms) | |
| if silence_ms < 0.0: | |
| raise ValueError("silence_ms cannot be negative") | |
| if isinstance(prediction, Prediction): | |
| probability = prediction.endpoint_probability | |
| inference_ms = prediction.inference_ms | |
| metadata = { | |
| "model_name": prediction.model_name, | |
| "auxiliary": prediction.auxiliary, | |
| } | |
| else: | |
| probability = float(prediction) | |
| if not 0.0 <= probability <= 1.0: | |
| raise ValueError("endpoint probability must be in [0, 1]") | |
| inference_ms = 0.0 | |
| metadata = {} | |
| threshold = self.threshold_for_silence(silence_ms) | |
| emit_response = False | |
| if self._state is TurnState.END: | |
| reason = "endpoint_latched" | |
| elif silence_ms < self.config.min_silence_ms: | |
| self._state = TurnState.HOLD | |
| self._confirmations = 0 | |
| reason = "minimum_silence_not_reached" | |
| elif silence_ms >= self.config.max_silence_ms: | |
| self._state = TurnState.END | |
| self._confirmations = self.config.required_confirmations | |
| reason = "maximum_timeout" | |
| emit_response = True | |
| elif probability >= threshold: | |
| self._confirmations += 1 | |
| if self._confirmations >= self.config.required_confirmations: | |
| self._state = TurnState.END | |
| reason = "model_endpoint" | |
| emit_response = True | |
| else: | |
| self._state = TurnState.HOLD | |
| reason = "awaiting_confirmation" | |
| else: | |
| self._state = TurnState.HOLD | |
| self._confirmations = 0 | |
| reason = "model_hold" | |
| return TurnDecision( | |
| state=self._state, | |
| endpoint_probability=probability, | |
| threshold=threshold, | |
| silence_ms=silence_ms, | |
| reason=reason, | |
| timestamp_ms=timestamp_ms, | |
| inference_ms=inference_ms, | |
| confirmations=self._confirmations, | |
| emit_response=emit_response, | |
| metadata=metadata, | |
| ) | |
| def _validate_timestamp(self, timestamp_ms: float | None) -> None: | |
| if timestamp_ms is None: | |
| return | |
| if timestamp_ms < 0.0: | |
| raise ValueError("timestamp_ms cannot be negative") | |
| if self._last_timestamp_ms is not None and timestamp_ms < self._last_timestamp_ms: | |
| raise ValueError("timestamps must be monotonic") | |
| self._last_timestamp_ms = timestamp_ms | |