suvradeepp's picture
Deploy Tiny Hinglish Turn Detector development preview
2d70679 verified
Raw
History Blame Contribute Delete
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
@dataclass(frozen=True, slots=True)
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
@property
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