[唤醒应答加速]:完成ACK音频预热缓存,包含唤醒后复用播放和回归测试

This commit is contained in:
mkbk
2026-06-18 13:44:44 +08:00
parent 3408a30e25
commit 86a429f018
3 changed files with 82 additions and 2 deletions
+19 -1
View File
@@ -122,6 +122,8 @@ class TurnController:
self.sentence_buffer = sentence_buffer or SentenceBuffer()
self._states: list[PipelineState] = []
self._pending_capture_frames: list[AudioFrame] = []
self._cached_ack_text: str | None = None
self._cached_ack_segment: AudioSegment | None = None
def run_turn(self, turn_id: int) -> TurnResult:
self._states = []
@@ -139,6 +141,15 @@ class TurnController:
except ProviderError as exc:
return self._recover(exc, turn_id)
def prepare_ack_audio(self) -> None:
text = self.config.wake_ack_text.strip()
if not text:
self._cached_ack_text = None
self._cached_ack_segment = None
return
self._cached_ack_text = text
self._cached_ack_segment = self.ack_tts.synthesize(text)
def _wait_for_wake_and_user_text(self, turn_id: int) -> str | ProviderError:
wake_error = self._wait_for_local_wake()
if wake_error is not None:
@@ -300,7 +311,7 @@ class TurnController:
return None
try:
self._event(ACK_STARTED, PipelineState.SPEAKING, f"应答中:{text}", turn_id=turn_id)
segment = self.ack_tts.synthesize(text)
segment = self._ack_segment(text)
playback = self.transport.play_pcm(segment)
if playback.error:
return playback.error
@@ -309,6 +320,12 @@ class TurnController:
except ProviderError as exc:
return exc
def _ack_segment(self, text: str) -> AudioSegment:
if self._cached_ack_text != text or self._cached_ack_segment is None:
self._cached_ack_text = text
self._cached_ack_segment = self.ack_tts.synthesize(text)
return self._cached_ack_segment
def _reply_to_user(self, user_text: str, turn_id: int) -> TurnResult:
completed_turns = 0
current_user_text = user_text
@@ -664,6 +681,7 @@ class VoiceAssistantPipeline:
self.tts.load()
if self.ack_tts is not self.tts:
self.ack_tts.load()
self.controller.prepare_ack_audio()
def run(self, *, once: bool = False, max_turns: int | None = None) -> RuntimeSummary:
completed = 0
+19 -1
View File
@@ -115,6 +115,8 @@ class LiveVoiceRuntime:
self.event_bus.subscribe(self._report_event)
self.sentence_buffer = sentence_buffer or SentenceBuffer()
self._states: list[PipelineState] = []
self._cached_ack_text: str | None = None
self._cached_ack_segment: AudioSegment | None = None
def load(self) -> None:
self.wakeword.load()
@@ -123,6 +125,16 @@ class LiveVoiceRuntime:
self.tts.load()
if self.ack_tts is not self.tts:
self.ack_tts.load()
self.prepare_ack_audio()
def prepare_ack_audio(self) -> None:
text = self.config.wake_ack_text.strip()
if not text:
self._cached_ack_text = None
self._cached_ack_segment = None
return
self._cached_ack_text = text
self._cached_ack_segment = self.ack_tts.synthesize(text)
def run(self, *, once: bool = False, max_turns: int | None = None) -> RuntimeSummary:
completed = 0
@@ -243,7 +255,7 @@ class LiveVoiceRuntime:
return None
try:
self._event(ACK_STARTED, PipelineState.SPEAKING, f"应答中:{text}", turn_id=turn_id)
segment = self.ack_tts.synthesize(text)
segment = self._ack_segment(text)
playback = self.transport.play_pcm(segment)
if playback.error:
return playback.error
@@ -252,6 +264,12 @@ class LiveVoiceRuntime:
except ProviderError as exc:
return exc
def _ack_segment(self, text: str) -> AudioSegment:
if self._cached_ack_text != text or self._cached_ack_segment is None:
self._cached_ack_text = text
self._cached_ack_segment = self.ack_tts.synthesize(text)
return self._cached_ack_segment
def _drain_input_after_playback(self) -> None:
self.transport.flush_input()
if self.config.post_playback_drain_ms <= 0:
+44
View File
@@ -113,6 +113,19 @@ class QueueLlmProvider:
yield ReplyDelta("", finish_reason="stop")
class CountingTtsProvider:
def __init__(self) -> None:
self.delegate = SineTtsProvider()
self.synthesized_texts: list[str] = []
def load(self) -> None:
self.delegate.load()
def synthesize(self, text: str) -> AudioSegment:
self.synthesized_texts.append(text)
return self.delegate.synthesize(text)
class RecordingReporter:
def __init__(self) -> None:
self.statuses: list[str] = []
@@ -282,6 +295,37 @@ class LiveRuntimeTests(unittest.TestCase):
positions = [event_types.index(item) for item in expected_order]
self.assertEqual(positions, sorted(positions))
def test_wake_ack_audio_is_prepared_once_and_reused(self) -> None:
frames = []
for idx, _text in enumerate(["第一问", "第二问"]):
base_id = idx * 5
base_ms = idx * 120
frames.append(wake_frame(base_id, base_ms))
frames.extend(segment_frames(base_id + 1, base_ms + 20))
transport = MemoryAudioTransport(frames, flush_clears_input=False)
ack_tts = CountingTtsProvider()
runtime = VoiceAssistantPipeline(
config=AppConfig(llm_api_key="secret", speech_provider="cloud", wake_ack_text="我在"),
transport=transport,
wakeword=KeywordWakeWordProvider(),
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40),
stt=QueueSttProvider(["第一问", "第二问"]),
realtime_stt=None,
llm=MockLlmProvider(["这是答复。"]),
tts=SineTtsProvider(),
ack_tts=ack_tts,
context=ConversationContext(),
reporter=RecordingReporter(),
event_bus=PipelineEventBus(),
)
summary = runtime.run(max_turns=2)
self.assertEqual(summary.completed_turns, 2)
self.assertEqual(ack_tts.synthesized_texts, ["我在"])
self.assertEqual(transport.played_segments[0].metadata["text"], "我在")
self.assertEqual(transport.played_segments[2].metadata["text"], "我在")
def test_zero_post_playback_drain_flushes_without_dropping_prefilled_question(self) -> None:
runtime, stt, _, transport, reporter = make_runtime(["第一问"])
self.assertEqual(runtime.config.post_playback_drain_ms, 0)