[本地语音降噪]:完成本地语音链路和噪音过滤,包含降噪模型、本地ASR和实时字幕稳定策略
This commit is contained in:
@@ -24,7 +24,7 @@ from owner_voice_pet.llm import MockLlmProvider
|
||||
from owner_voice_pet.models import AudioFrame, AudioSegment, Transcript
|
||||
from owner_voice_pet.assistant_pipeline import VoiceAssistantPipeline
|
||||
from owner_voice_pet.runtime import build_live_runtime
|
||||
from owner_voice_pet.stt import MetadataSttProvider
|
||||
from owner_voice_pet.stt import MetadataSttProvider, SherpaOnnxSttProvider
|
||||
from owner_voice_pet.transport import MemoryAudioTransport
|
||||
from owner_voice_pet.tts import SineTtsProvider
|
||||
from owner_voice_pet.vad import EnergyVadProvider, PrimarySpeakerVadRecorder, VadRecorder
|
||||
@@ -97,10 +97,43 @@ class RecordingReporter:
|
||||
self.errors.append(f"{stage}:{code}:{message}")
|
||||
|
||||
|
||||
class MarkerAudioPreprocessor:
|
||||
def __init__(self, partial_text: str = "降噪后问题") -> None:
|
||||
self.partial_text = partial_text
|
||||
self.loaded = False
|
||||
self.reset_calls = 0
|
||||
self.frames: list[AudioFrame] = []
|
||||
|
||||
def load(self) -> None:
|
||||
self.loaded = True
|
||||
|
||||
def reset(self) -> None:
|
||||
self.reset_calls += 1
|
||||
|
||||
def process_frame(self, frame: AudioFrame) -> AudioFrame:
|
||||
metadata = dict(frame.metadata)
|
||||
metadata["denoised"] = True
|
||||
metadata["partial_transcript"] = self.partial_text
|
||||
processed = AudioFrame(
|
||||
b"\x01\x00",
|
||||
frame.sample_rate,
|
||||
frame.channels,
|
||||
frame.timestamp_ms,
|
||||
frame.frame_id,
|
||||
metadata,
|
||||
)
|
||||
self.frames.append(processed)
|
||||
return processed
|
||||
|
||||
def flush(self) -> list[AudioFrame]:
|
||||
return []
|
||||
|
||||
|
||||
def make_runtime(
|
||||
texts: list[str],
|
||||
context: ConversationContext | None = None,
|
||||
partial_texts: list[list[str]] | None = None,
|
||||
audio_preprocessor: MarkerAudioPreprocessor | None = None,
|
||||
) -> tuple[VoiceAssistantPipeline, QueueSttProvider, MockLlmProvider, MemoryAudioTransport, RecordingReporter]:
|
||||
frames = []
|
||||
for idx, _text in enumerate(texts):
|
||||
@@ -121,6 +154,7 @@ def make_runtime(
|
||||
wakeword=KeywordWakeWordProvider(),
|
||||
vad_recorder=VadRecorder(EnergyVadProvider(), min_duration_ms=40, end_silence_ms=40),
|
||||
stt=stt,
|
||||
audio_preprocessor=audio_preprocessor,
|
||||
realtime_stt=MetadataSttProvider() if partial_texts is not None else None,
|
||||
llm=llm,
|
||||
tts=tts,
|
||||
@@ -200,7 +234,7 @@ class LiveRuntimeTests(unittest.TestCase):
|
||||
runtime, _, llm, _, reporter = make_runtime(["第一问"], partial_texts=[["第一", "第一问"]])
|
||||
runtime.run(max_turns=1)
|
||||
|
||||
self.assertEqual(reporter.partials, ["第一", "第一问"])
|
||||
self.assertEqual(reporter.partials, ["第一问"])
|
||||
self.assertEqual(reporter.transcripts, ["第一问"])
|
||||
self.assertEqual(llm.calls[0][-1].content, "第一问")
|
||||
event_types = [event.type for event in runtime.event_bus.events]
|
||||
@@ -208,6 +242,22 @@ class LiveRuntimeTests(unittest.TestCase):
|
||||
self.assertLess(event_types.index(TRANSCRIPT_PARTIAL), event_types.index(SPEECH_ENDED))
|
||||
self.assertLess(event_types.index(TRANSCRIPT_PARTIAL), event_types.index(TRANSCRIPT_FINAL))
|
||||
|
||||
def test_capture_uses_denoised_frames_for_partial_and_final_stt(self) -> None:
|
||||
preprocessor = MarkerAudioPreprocessor(partial_text="降噪后问题")
|
||||
runtime, stt, _, _, reporter = make_runtime(
|
||||
["第一问"],
|
||||
partial_texts=[["原始噪声", "原始噪声"]],
|
||||
audio_preprocessor=preprocessor,
|
||||
)
|
||||
runtime.run(max_turns=1)
|
||||
|
||||
self.assertTrue(preprocessor.loaded)
|
||||
self.assertGreaterEqual(preprocessor.reset_calls, 1)
|
||||
self.assertEqual(reporter.partials, ["降噪后问题"])
|
||||
self.assertEqual(len(stt.calls), 1)
|
||||
self.assertTrue(stt.calls[0].metadata["denoised"])
|
||||
self.assertIn(b"\x01\x00", stt.calls[0].pcm)
|
||||
|
||||
def test_wake_keyword_does_not_pollute_llm_user_message(self) -> None:
|
||||
runtime, _, llm, _, _ = make_runtime(["第一问"])
|
||||
runtime.run(max_turns=1)
|
||||
@@ -218,6 +268,8 @@ class LiveRuntimeTests(unittest.TestCase):
|
||||
def test_live_runtime_uses_primary_speaker_endpoint_by_default(self) -> None:
|
||||
runtime = build_live_runtime(AppConfig(llm_api_key="secret"))
|
||||
self.assertIsInstance(runtime.vad_recorder, PrimarySpeakerVadRecorder)
|
||||
self.assertEqual(runtime.config.speech_provider, "local")
|
||||
self.assertIsInstance(runtime.stt, SherpaOnnxSttProvider)
|
||||
self.assertIsNotNone(runtime.realtime_stt)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user