[唤醒应答与快速端点]:完成唤醒后我在播报和本地VAD端点优化,包含缓冲清理、配置项和测试覆盖

This commit is contained in:
mkbk
2026-06-17 20:51:14 +08:00
parent 860c2909fc
commit 44f708a3dc
12 changed files with 224 additions and 7 deletions
+16
View File
@@ -57,6 +57,14 @@ def main(argv: list[str] | None = None) -> int:
"wake_keywords_file": str(config.wake_keywords_file) if config.wake_keywords_file else "",
"wake_kws_threshold": config.wake_kws_threshold,
"wake_kws_score": config.wake_kws_score,
"wake_ack_text": config.wake_ack_text,
"post_playback_drain_ms": config.post_playback_drain_ms,
"vad_provider": config.vad_provider,
"vad_threshold": config.vad_threshold,
"vad_min_duration_ms": config.vad_min_duration_ms,
"vad_end_silence_ms": config.vad_end_silence_ms,
"vad_no_speech_timeout_ms": config.vad_no_speech_timeout_ms,
"vad_max_recording_ms": config.vad_max_recording_ms,
"speech_provider": config.speech_provider,
"asr_model": config.asr_model,
"tts_model": config.tts_model,
@@ -144,6 +152,14 @@ def main(argv: list[str] | None = None) -> int:
wake_keywords_file=config.wake_keywords_file,
wake_kws_threshold=config.wake_kws_threshold,
wake_kws_score=config.wake_kws_score,
wake_ack_text=config.wake_ack_text,
post_playback_drain_ms=config.post_playback_drain_ms,
vad_provider=config.vad_provider,
vad_threshold=config.vad_threshold,
vad_min_duration_ms=config.vad_min_duration_ms,
vad_end_silence_ms=config.vad_end_silence_ms,
vad_no_speech_timeout_ms=config.vad_no_speech_timeout_ms,
vad_max_recording_ms=config.vad_max_recording_ms,
speech_provider=config.speech_provider,
asr_model=config.asr_model,
tts_model=config.tts_model,
+62
View File
@@ -24,6 +24,14 @@ class AppConfig:
wake_keywords_file: Path | None = None
wake_kws_threshold: float = 0.25
wake_kws_score: float = 1.0
wake_ack_text: str = "我在"
post_playback_drain_ms: int = 250
vad_provider: str = "local"
vad_threshold: float = 0.5
vad_min_duration_ms: int = 250
vad_end_silence_ms: int = 350
vad_no_speech_timeout_ms: int = 5000
vad_max_recording_ms: int = 12000
speech_provider: str = "cloud"
asr_model: str = "mimo-v2.5-asr"
tts_model: str = "mimo-v2.5-tts"
@@ -57,6 +65,14 @@ class AppConfig:
wake_keywords_file=Path(value) if (value := get("WAKE_KEYWORDS_FILE")) else None,
wake_kws_threshold=float(get("WAKE_KWS_THRESHOLD", "0.25") or "0.25"),
wake_kws_score=float(get("WAKE_KWS_SCORE", "1.0") or "1.0"),
wake_ack_text=get("WAKE_ACK_TEXT", "我在") or "我在",
post_playback_drain_ms=int(get("POST_PLAYBACK_DRAIN_MS", "250") or "250"),
vad_provider=(get("VAD_PROVIDER", "local") or "local").lower(),
vad_threshold=float(get("VAD_THRESHOLD", "0.5") or "0.5"),
vad_min_duration_ms=int(get("VAD_MIN_DURATION_MS", "250") or "250"),
vad_end_silence_ms=int(get("VAD_END_SILENCE_MS", "350") or "350"),
vad_no_speech_timeout_ms=int(get("VAD_NO_SPEECH_TIMEOUT_MS", "5000") or "5000"),
vad_max_recording_ms=int(get("VAD_MAX_RECORDING_MS", "12000") or "12000"),
speech_provider=(get("SPEECH_PROVIDER", "cloud") or "cloud").lower(),
asr_model=get("ASR_MODEL", "mimo-v2.5-asr") or "mimo-v2.5-asr",
tts_model=get("TTS_MODEL", "mimo-v2.5-tts") or "mimo-v2.5-tts",
@@ -152,6 +168,52 @@ class AppConfig:
"startup",
)
)
if self.post_playback_drain_ms < 0:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_POST_PLAYBACK_DRAIN_MS must be non-negative",
False,
"config",
"startup",
)
)
if self.vad_provider not in {"local", "energy"}:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_VAD_PROVIDER must be local or energy",
False,
"config",
"startup",
)
)
if self.vad_threshold <= 0:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_VAD_THRESHOLD must be positive",
False,
"config",
"startup",
)
)
for name, value in {
"OWNER_VAD_MIN_DURATION_MS": self.vad_min_duration_ms,
"OWNER_VAD_END_SILENCE_MS": self.vad_end_silence_ms,
"OWNER_VAD_NO_SPEECH_TIMEOUT_MS": self.vad_no_speech_timeout_ms,
"OWNER_VAD_MAX_RECORDING_MS": self.vad_max_recording_ms,
}.items():
if value <= 0:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
f"{name} must be positive",
False,
"config",
"startup",
)
)
if not self.llm_base_url.startswith(("http://", "https://")):
errors.append(
ProviderError(
+3
View File
@@ -28,6 +28,9 @@ class AudioTransport(Protocol):
def play_pcm(self, segment: AudioSegment, interrupt: bool = False) -> PlaybackResult:
...
def flush_input(self) -> int:
...
def stop(self) -> None:
...
+48 -2
View File
@@ -12,7 +12,7 @@ from .protocols import AudioTransport, LlmProvider, SttProvider, TtsProvider, Wa
from .stt import CloudAsrSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
from .transport import SoundDeviceAudioTransport
from .tts import CloudTtsProvider, MacSayTtsProvider, SentenceBuffer
from .vad import EnergyVadProvider, VadRecorder
from .vad import EnergyVadProvider, SherpaOnnxVadProvider, VadRecorder
from .wakeword import SherpaOnnxKeywordWakeWordProvider
@@ -71,6 +71,7 @@ class LiveVoiceRuntime:
llm: LlmProvider,
tts: TtsProvider,
context: ConversationContext,
ack_tts: TtsProvider | None = None,
reporter: RuntimeReporter | None = None,
sentence_buffer: SentenceBuffer | None = None,
) -> None:
@@ -81,6 +82,7 @@ class LiveVoiceRuntime:
self.stt = stt
self.llm = llm
self.tts = tts
self.ack_tts = ack_tts or tts
self.context = context
self.reporter = reporter or TerminalRuntimeReporter()
self.sentence_buffer = sentence_buffer or SentenceBuffer()
@@ -91,6 +93,8 @@ class LiveVoiceRuntime:
self.vad_recorder.provider.load()
self.stt.load()
self.tts.load()
if self.ack_tts is not self.tts:
self.ack_tts.load()
def run(self, *, once: bool = False, max_turns: int | None = None) -> RuntimeSummary:
completed = 0
@@ -142,6 +146,9 @@ class LiveVoiceRuntime:
if wake_error is not None:
return wake_error
self._state(PipelineState.SPEECH_DETECTING, "唤醒命中:请说出问题", turn_id=turn_id)
ack_error = self._acknowledge_wake(turn_id)
if ack_error is not None:
return ack_error
user_segment = self._capture_segment(turn_id, state_message="录音中:正在听取问题")
if isinstance(user_segment, ProviderError):
return user_segment
@@ -186,6 +193,33 @@ class LiveVoiceRuntime:
if isinstance(result, AudioSegment):
return result
def _acknowledge_wake(self, turn_id: int) -> ProviderError | None:
text = self.config.wake_ack_text.strip()
if not text:
self._drain_input_after_playback()
return None
try:
self._state(PipelineState.SPEAKING, f"应答中:{text}", turn_id=turn_id)
segment = self.ack_tts.synthesize(text)
playback = self.transport.play_pcm(segment)
if playback.error:
return playback.error
self._drain_input_after_playback()
return None
except ProviderError as exc:
return exc
def _drain_input_after_playback(self) -> None:
if self.config.post_playback_drain_ms <= 0:
return
self.transport.flush_input()
remaining_ms = self.config.post_playback_drain_ms
while remaining_ms > 0:
timeout_ms = min(50, remaining_ms)
self.transport.read_frames(timeout_ms=timeout_ms)
remaining_ms -= timeout_ms
self.transport.flush_input()
def _reply_to_user(self, user_text: str, turn_id: int) -> TurnResult:
self.context.append_user(user_text)
self._state(PipelineState.THINKING, "思考中:正在生成回复", turn_id=turn_id)
@@ -220,6 +254,7 @@ class LiveVoiceRuntime:
playback = self.transport.play_pcm(segment)
if playback.error:
raise playback.error
self._drain_input_after_playback()
def _recover(self, error: ProviderError, turn_id: int) -> TurnResult:
self.reporter.error(error.stage, error.code.value, error.message, turn_id=turn_id)
@@ -242,6 +277,10 @@ def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = Non
else:
stt = SherpaOnnxSttProvider(str(config.speech_models_dir))
tts = MacSayTtsProvider()
if config.vad_provider == "local":
vad_provider = SherpaOnnxVadProvider(config.speech_models_dir, threshold=config.vad_threshold)
else:
vad_provider = EnergyVadProvider()
return LiveVoiceRuntime(
config=config,
transport=SoundDeviceAudioTransport(output_device=config.audio_output_device),
@@ -252,10 +291,17 @@ def build_live_runtime(config: AppConfig, reporter: RuntimeReporter | None = Non
threshold=config.wake_kws_threshold,
score=config.wake_kws_score,
),
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=300, end_silence_ms=500, no_speech_timeout_ms=8000),
vad_recorder=VadRecorder(
vad_provider,
min_duration_ms=config.vad_min_duration_ms,
end_silence_ms=config.vad_end_silence_ms,
no_speech_timeout_ms=config.vad_no_speech_timeout_ms,
max_recording_ms=config.vad_max_recording_ms,
),
stt=stt,
llm=OpenAICompatibleLlmProvider(config),
tts=tts,
ack_tts=MacSayTtsProvider(),
context=ConversationContext(
max_messages=config.context_max_messages,
max_chars=config.context_max_chars,
+15
View File
@@ -103,6 +103,11 @@ class MemoryAudioTransport:
self.played_segments.append(segment)
return PlaybackResult(True, segment.duration_ms)
def flush_input(self) -> int:
count = len(self._frames)
self._frames.clear()
return count
def stop(self) -> None:
self.started = False
@@ -290,6 +295,16 @@ class SoundDeviceAudioTransport:
break
return None
def flush_input(self) -> int:
count = 0
while not self._queue.empty():
try:
self._queue.get_nowait()
count += 1
except queue.Empty:
break
return count
def health(self) -> TransportHealth:
if self._sd is None:
return TransportHealth(False, False, "sounddevice unavailable")