[低延迟端点]:完成首句保留和快速结束修正,包含ACK缓冲策略、批量读帧和主说话人端点回归测试

This commit is contained in:
mkbk
2026-06-17 22:04:25 +08:00
parent e74ec28e7d
commit a64bb86da4
15 changed files with 136 additions and 15 deletions
+2 -1
View File
@@ -97,7 +97,7 @@ def make_runtime(texts: list[str], context: ConversationContext | None = None) -
reporter = RecordingReporter()
event_bus = PipelineEventBus()
runtime = VoiceAssistantPipeline(
config=AppConfig(llm_api_key="secret", speech_provider="cloud", post_playback_drain_ms=0),
config=AppConfig(llm_api_key="secret", speech_provider="cloud"),
transport=transport,
wakeword=KeywordWakeWordProvider(),
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40),
@@ -119,6 +119,7 @@ class LiveRuntimeTests(unittest.TestCase):
summary = runtime.run(max_turns=2)
self.assertEqual(summary.completed_turns, 2)
self.assertEqual(runtime.config.post_playback_drain_ms, 0)
self.assertEqual(len(stt.calls), 2)
self.assertEqual(len(llm.calls), 2)
self.assertEqual(len(transport.played_segments), 4)
+2 -1
View File
@@ -73,10 +73,11 @@ class ModelsConfigTests(unittest.TestCase):
self.assertEqual(config.wake_kws_threshold, 0.15)
self.assertEqual(config.wake_kws_score, 1.0)
self.assertEqual(config.wake_ack_text, "我在")
self.assertEqual(config.post_playback_drain_ms, 50)
self.assertEqual(config.post_playback_drain_ms, 0)
self.assertEqual(config.pipeline_mode, "live_turn_based")
self.assertEqual(config.endpoint_mode, "primary_speaker")
self.assertEqual(config.speaker_profile_ms, 600)
self.assertEqual(config.speaker_profile_min_ms, 120)
self.assertEqual(config.speaker_absent_ms, 300)
self.assertEqual(config.speaker_similarity_threshold, 0.70)
self.assertEqual(config.speaker_min_rms, 0.012)
+16 -2
View File
@@ -82,6 +82,15 @@ class TransportTests(unittest.TestCase):
self.assertEqual(frames[0].pcm, b"\x01\x00\x02\x00")
self.assertEqual(frames[0].sample_rate, 16000)
def test_sounddevice_transport_returns_queued_frame_batch(self) -> None:
fake = FakeSoundDevice(callback_payloads=[b"\x01\x00", b"\x02\x00", b"\x03\x00"])
transport = SoundDeviceAudioTransport(sounddevice_module=fake)
transport.start_input(sample_rate=16000, channels=1)
frames = transport.read_frames(10)
transport.stop()
self.assertEqual([item.frame_id for item in frames], [0, 1, 2])
def test_sounddevice_transport_flushes_queued_input(self) -> None:
fake = FakeSoundDevice()
transport = SoundDeviceAudioTransport(sounddevice_module=fake)
@@ -115,7 +124,9 @@ class FakeInputStream:
def start(self) -> None:
self.started = True
self.callback(b"\x01\x00\x02\x00", 2, None, "")
payloads = self.kwargs.get("callback_payloads") or [b"\x01\x00\x02\x00"]
for payload in payloads:
self.callback(payload, max(1, len(payload) // 2), None, "")
def stop(self) -> None:
self.started = False
@@ -140,10 +151,13 @@ class FakeOutputStream:
class FakeSoundDevice:
def __init__(self) -> None:
def __init__(self, callback_payloads: list[bytes] | None = None) -> None:
self.callback_payloads = callback_payloads
self.output_writes: list[bytes] = []
def RawInputStream(self, **kwargs):
if self.callback_payloads is not None:
kwargs["callback_payloads"] = self.callback_payloads
return FakeInputStream(**kwargs)
def RawOutputStream(self, **kwargs):
+32 -1
View File
@@ -116,9 +116,10 @@ class WakeVadSttTests(unittest.TestCase):
provider.load()
recorder = PrimarySpeakerVadRecorder(
provider,
min_duration_ms=40,
min_duration_ms=250,
end_silence_ms=1000,
speaker_profile_ms=40,
speaker_profile_min_ms=40,
speaker_absent_ms=40,
)
frames = [
@@ -142,6 +143,35 @@ class WakeVadSttTests(unittest.TestCase):
self.assertEqual(consumed, 4)
self.assertEqual(segment.metadata["transcript"], "你是谁")
def test_primary_speaker_endpoint_does_not_wait_for_vad_min_duration(self) -> None:
provider = EnergyVadProvider()
provider.load()
recorder = PrimarySpeakerVadRecorder(
provider,
min_duration_ms=1000,
end_silence_ms=1000,
speaker_profile_ms=120,
speaker_profile_min_ms=40,
speaker_absent_ms=40,
)
frames = [
make_frame(1, 0, speech=True, metadata={"speaker_id": "owner", "transcript": "你在做什么"}),
make_frame(2, 20, speech=True, metadata={"speaker_id": "owner"}),
make_frame(3, 40, speech=True, metadata={"speaker_id": "background"}),
make_frame(4, 60, speech=True, metadata={"speaker_id": "background"}),
make_frame(5, 80, speech=True, metadata={"speaker_id": "owner", "transcript": "第二次重复"}),
]
segment = None
for item in frames:
result = recorder.feed(item)
if isinstance(result, AudioSegment):
segment = result
break
self.assertIsNotNone(segment)
assert segment is not None
self.assertEqual(segment.metadata["end_reason"], "primary_speaker_absent")
self.assertEqual(segment.metadata["transcript"], "你在做什么")
def test_primary_speaker_endpoint_allows_short_pause(self) -> None:
provider = EnergyVadProvider()
provider.load()
@@ -150,6 +180,7 @@ class WakeVadSttTests(unittest.TestCase):
min_duration_ms=40,
end_silence_ms=1000,
speaker_profile_ms=40,
speaker_profile_min_ms=40,
speaker_absent_ms=60,
)
frames = [