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