[低延迟端点]:完成首句保留和快速结束修正,包含ACK缓冲策略、批量读帧和主说话人端点回归测试
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
@@ -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):
|
||||
|
||||
@@ -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 = [
|
||||
|
||||
Reference in New Issue
Block a user