[唤醒应答加速]:完成ACK音频预热缓存,包含唤醒后复用播放和回归测试
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user