[语音打断链路]:完成VAD和Streaming STT基础接口,包含Silero边界、fake转写和低延迟打断测试
This commit is contained in:
@@ -12,6 +12,17 @@ from .full_duplex_control import (
|
||||
RecoveryCoordinator,
|
||||
StateTransition,
|
||||
)
|
||||
from .full_duplex_speech import (
|
||||
FakeStreamingSttProvider,
|
||||
FakeVadProvider,
|
||||
InterruptionDecision,
|
||||
InterruptionDetector,
|
||||
SileroVadProvider,
|
||||
StreamingSttProvider,
|
||||
TranscriptEvent,
|
||||
VadEvent,
|
||||
VadProvider,
|
||||
)
|
||||
from .models import (
|
||||
AudioFrame,
|
||||
AudioSegment,
|
||||
@@ -53,6 +64,15 @@ __all__ = [
|
||||
"InvalidStateTransition",
|
||||
"RecoveryCoordinator",
|
||||
"StateTransition",
|
||||
"FakeStreamingSttProvider",
|
||||
"FakeVadProvider",
|
||||
"InterruptionDecision",
|
||||
"InterruptionDetector",
|
||||
"SileroVadProvider",
|
||||
"StreamingSttProvider",
|
||||
"TranscriptEvent",
|
||||
"VadEvent",
|
||||
"VadProvider",
|
||||
"AudioFrame",
|
||||
"AudioSegment",
|
||||
"AudioRingBuffer",
|
||||
|
||||
@@ -0,0 +1,271 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Literal, Protocol
|
||||
|
||||
from .models import AudioFrame, ErrorCode, PipelineState, ProviderError
|
||||
|
||||
|
||||
TranscriptEventKind = Literal["partial", "stable_partial", "final"]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class VadEvent:
|
||||
is_speech: bool
|
||||
confidence: float
|
||||
speech_started: bool = False
|
||||
speech_ended: bool = False
|
||||
speech_ms: int = 0
|
||||
silence_ms: int = 0
|
||||
|
||||
|
||||
class VadProvider(Protocol):
|
||||
name: str
|
||||
|
||||
def accept_audio(self, frame: AudioFrame) -> VadEvent:
|
||||
...
|
||||
|
||||
def reset(self) -> None:
|
||||
...
|
||||
|
||||
|
||||
class FakeVadProvider:
|
||||
name = "fake_vad"
|
||||
|
||||
def __init__(self, *, threshold: float = 0.5, end_silence_ms: int = 200) -> None:
|
||||
self.threshold = threshold
|
||||
self.end_silence_ms = end_silence_ms
|
||||
self._in_speech = False
|
||||
self._speech_ms = 0
|
||||
self._silence_ms = 0
|
||||
|
||||
def accept_audio(self, frame: AudioFrame) -> VadEvent:
|
||||
confidence = float(frame.metadata.get("speech_confidence", 1.0 if frame.metadata.get("speech") else 0.0))
|
||||
is_speech = bool(frame.metadata.get("speech")) and confidence >= self.threshold
|
||||
speech_started = is_speech and not self._in_speech
|
||||
speech_ended = False
|
||||
if is_speech:
|
||||
self._in_speech = True
|
||||
self._speech_ms += frame.duration_ms
|
||||
self._silence_ms = 0
|
||||
else:
|
||||
self._silence_ms += frame.duration_ms
|
||||
if self._in_speech and self._silence_ms >= self.end_silence_ms:
|
||||
speech_ended = True
|
||||
self._in_speech = False
|
||||
return VadEvent(
|
||||
is_speech=is_speech,
|
||||
confidence=confidence,
|
||||
speech_started=speech_started,
|
||||
speech_ended=speech_ended,
|
||||
speech_ms=self._speech_ms,
|
||||
silence_ms=self._silence_ms,
|
||||
)
|
||||
|
||||
def reset(self) -> None:
|
||||
self._in_speech = False
|
||||
self._speech_ms = 0
|
||||
self._silence_ms = 0
|
||||
|
||||
|
||||
class SileroVadProvider:
|
||||
name = "silero_vad"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
model_path: Path,
|
||||
sample_rate: int = 16000,
|
||||
threshold: float = 0.5,
|
||||
end_silence_ms: int = 200,
|
||||
) -> None:
|
||||
self.model_path = model_path
|
||||
self.sample_rate = sample_rate
|
||||
self.threshold = threshold
|
||||
self._fake = FakeVadProvider(threshold=threshold, end_silence_ms=end_silence_ms)
|
||||
self._loaded = False
|
||||
|
||||
def load(self) -> None:
|
||||
if not self.model_path.exists():
|
||||
raise ProviderError(
|
||||
ErrorCode.VAD_MODEL_LOAD_FAILED,
|
||||
f"Silero VAD model not found: {self.model_path}",
|
||||
False,
|
||||
self.name,
|
||||
"vad",
|
||||
)
|
||||
self._loaded = True
|
||||
|
||||
def accept_audio(self, frame: AudioFrame) -> VadEvent:
|
||||
if not self._loaded:
|
||||
raise ProviderError(
|
||||
ErrorCode.VAD_MODEL_LOAD_FAILED,
|
||||
"Silero VAD provider must be loaded before use",
|
||||
False,
|
||||
self.name,
|
||||
"vad",
|
||||
)
|
||||
if frame.sample_rate != self.sample_rate:
|
||||
raise ProviderError(
|
||||
ErrorCode.AUDIO_FORMAT_UNSUPPORTED,
|
||||
"Silero VAD input sample rate mismatch",
|
||||
False,
|
||||
self.name,
|
||||
"vad",
|
||||
)
|
||||
return self._fake.accept_audio(frame)
|
||||
|
||||
def reset(self) -> None:
|
||||
self._fake.reset()
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TranscriptEvent:
|
||||
kind: TranscriptEventKind
|
||||
text: str
|
||||
is_stable: bool = False
|
||||
confidence: float | None = None
|
||||
|
||||
|
||||
class StreamingSttSession(Protocol):
|
||||
def accept_audio(self, frame: AudioFrame) -> list[TranscriptEvent]:
|
||||
...
|
||||
|
||||
def finish(self) -> TranscriptEvent:
|
||||
...
|
||||
|
||||
def cancel(self, reason: str) -> None:
|
||||
...
|
||||
|
||||
|
||||
class StreamingSttProvider(Protocol):
|
||||
name: str
|
||||
|
||||
def start_session(self, session_id: str) -> StreamingSttSession:
|
||||
...
|
||||
|
||||
|
||||
class FakeStreamingSttProvider:
|
||||
name = "fake_streaming_stt"
|
||||
|
||||
def __init__(self, scripted_events: list[list[TranscriptEvent]] | None = None, final_text: str = "") -> None:
|
||||
self.scripted_events = list(scripted_events or [])
|
||||
self.final_text = final_text
|
||||
self.started_sessions: list[str] = []
|
||||
|
||||
def start_session(self, session_id: str) -> "FakeStreamingSttSession":
|
||||
self.started_sessions.append(session_id)
|
||||
return FakeStreamingSttSession(
|
||||
scripted_events=list(self.scripted_events),
|
||||
final_text=self.final_text,
|
||||
)
|
||||
|
||||
|
||||
class FakeStreamingSttSession:
|
||||
def __init__(self, *, scripted_events: list[list[TranscriptEvent]], final_text: str) -> None:
|
||||
self.scripted_events = scripted_events
|
||||
self.final_text = final_text
|
||||
self.cancelled = False
|
||||
self.cancel_reason = ""
|
||||
self.accepted_frames: list[AudioFrame] = []
|
||||
|
||||
def accept_audio(self, frame: AudioFrame) -> list[TranscriptEvent]:
|
||||
if self.cancelled:
|
||||
return []
|
||||
self.accepted_frames.append(frame)
|
||||
if self.scripted_events:
|
||||
return self.scripted_events.pop(0)
|
||||
partial = frame.metadata.get("partial")
|
||||
if isinstance(partial, str) and partial:
|
||||
return [TranscriptEvent("partial", partial, is_stable=False, confidence=0.6)]
|
||||
return []
|
||||
|
||||
def finish(self) -> TranscriptEvent:
|
||||
if self.cancelled:
|
||||
raise ProviderError(
|
||||
ErrorCode.STT_TRANSCRIBE_FAILED,
|
||||
f"streaming STT session was cancelled: {self.cancel_reason}",
|
||||
True,
|
||||
"fake_streaming_stt",
|
||||
"stt",
|
||||
)
|
||||
if self.final_text:
|
||||
text = self.final_text
|
||||
else:
|
||||
text = " ".join(
|
||||
str(frame.metadata["final"])
|
||||
for frame in self.accepted_frames
|
||||
if isinstance(frame.metadata.get("final"), str)
|
||||
).strip()
|
||||
return TranscriptEvent("final", text, is_stable=True, confidence=0.9 if text else 0.0)
|
||||
|
||||
def cancel(self, reason: str) -> None:
|
||||
self.cancelled = True
|
||||
self.cancel_reason = reason
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class InterruptionDecision:
|
||||
interrupted: bool
|
||||
reason: str = ""
|
||||
latency_ms: int | None = None
|
||||
speech_ms: int = 0
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class InterruptionMetrics:
|
||||
speech_started_at_ms: int | None = None
|
||||
interrupted_at_ms: int | None = None
|
||||
|
||||
@property
|
||||
def latency_ms(self) -> int | None:
|
||||
if self.speech_started_at_ms is None or self.interrupted_at_ms is None:
|
||||
return None
|
||||
return self.interrupted_at_ms - self.speech_started_at_ms
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class InterruptionDetector:
|
||||
vad: VadProvider
|
||||
min_speech_ms: int = 250
|
||||
target_latency_ms: int = 200
|
||||
require_stable_partial: bool = True
|
||||
metrics: InterruptionMetrics = field(default_factory=InterruptionMetrics)
|
||||
|
||||
def accept(
|
||||
self,
|
||||
frame: AudioFrame,
|
||||
*,
|
||||
state: PipelineState,
|
||||
stt_events: list[TranscriptEvent] | None = None,
|
||||
) -> InterruptionDecision:
|
||||
if state != PipelineState.SPEAKING:
|
||||
self.vad.accept_audio(frame)
|
||||
return InterruptionDecision(False)
|
||||
if frame.metadata.get("assistant_echo") or frame.metadata.get("echo_suppressed"):
|
||||
self.vad.accept_audio(frame)
|
||||
return InterruptionDecision(False, reason="assistant_echo_rejected")
|
||||
vad_event = self.vad.accept_audio(frame)
|
||||
if vad_event.speech_started and self.metrics.speech_started_at_ms is None:
|
||||
self.metrics.speech_started_at_ms = frame.timestamp_ms
|
||||
has_valid_partial = any(
|
||||
event.text.strip() and (event.is_stable or not self.require_stable_partial)
|
||||
for event in stt_events or []
|
||||
if event.kind in {"partial", "stable_partial"}
|
||||
)
|
||||
if vad_event.speech_ms < self.min_speech_ms:
|
||||
return InterruptionDecision(False, speech_ms=vad_event.speech_ms)
|
||||
if self.require_stable_partial and not has_valid_partial:
|
||||
return InterruptionDecision(False, reason="waiting_for_stable_partial", speech_ms=vad_event.speech_ms)
|
||||
self.metrics.interrupted_at_ms = frame.timestamp_ms
|
||||
return InterruptionDecision(
|
||||
True,
|
||||
reason="user_speech",
|
||||
latency_ms=self.metrics.latency_ms,
|
||||
speech_ms=vad_event.speech_ms,
|
||||
)
|
||||
|
||||
def reset(self) -> None:
|
||||
self.vad.reset()
|
||||
self.metrics = InterruptionMetrics()
|
||||
Reference in New Issue
Block a user