[唤醒应答与快速端点]:完成唤醒后我在播报和本地VAD端点优化,包含缓冲清理、配置项和测试覆盖
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user