[本地语音降噪]:完成本地语音链路和噪音过滤,包含降噪模型、本地ASR和实时字幕稳定策略

This commit is contained in:
mkbk
2026-06-17 22:55:33 +08:00
parent a77a172412
commit 8c75fc5baf
22 changed files with 803 additions and 64 deletions
+97
View File
@@ -0,0 +1,97 @@
from __future__ import annotations
import tempfile
import unittest
from pathlib import Path
from owner_voice_pet.audio_preprocess import SherpaOnnxDenoiserPreprocessor
from owner_voice_pet.models import AudioFrame, ErrorCode, ProviderError
from owner_voice_pet.speech_models import DEFAULT_DENOISER_PATH
class FakeDenoisedAudio:
def __init__(self, samples, sample_rate: int) -> None:
self.samples = samples
self.sample_rate = sample_rate
class FakeSherpa:
class OfflineSpeechDenoiserGtcrnModelConfig:
def __init__(self, model: str) -> None:
self.model = model
class OfflineSpeechDenoiserModelConfig:
def __init__(self, gtcrn, num_threads: int, provider: str) -> None:
self.gtcrn = gtcrn
self.num_threads = num_threads
self.provider = provider
class OnlineSpeechDenoiserConfig:
def __init__(self, model) -> None:
self.model = model
class OnlineSpeechDenoiser:
def __init__(self, config) -> None:
self.config = config
self.reset_called = False
def run(self, samples, sample_rate: int):
import numpy as np
return FakeDenoisedAudio(np.zeros_like(samples, dtype=np.float32), sample_rate)
def reset(self) -> None:
self.reset_called = True
class FailingSherpa(FakeSherpa):
class OnlineSpeechDenoiser(FakeSherpa.OnlineSpeechDenoiser):
def run(self, samples, sample_rate: int):
raise RuntimeError("denoise failed")
class AudioPreprocessTests(unittest.TestCase):
def test_denoiser_missing_model_is_structured(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
provider = SherpaOnnxDenoiserPreprocessor(tmp, sherpa_module=FakeSherpa)
with self.assertRaises(ProviderError) as raised:
provider.load()
self.assertEqual(raised.exception.code, ErrorCode.NOISE_FILTER_MODEL_MISSING)
def test_denoiser_processes_frame_and_marks_metadata(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
model = Path(tmp) / DEFAULT_DENOISER_PATH
model.parent.mkdir(parents=True)
model.write_bytes(b"fake")
provider = SherpaOnnxDenoiserPreprocessor(tmp, sherpa_module=FakeSherpa)
provider.load()
frame = AudioFrame(
b"\xff\x7f",
16000,
1,
100,
7,
{"duration_ms": 20, "speech": True},
)
processed = provider.process_frame(frame)
self.assertNotEqual(processed.pcm, frame.pcm)
self.assertTrue(processed.metadata["denoised"])
self.assertEqual(processed.metadata["noise_filter_provider"], "sherpa_onnx_gtcrn")
self.assertEqual(processed.timestamp_ms, 100)
self.assertEqual(processed.frame_id, 7)
def test_denoiser_runtime_failure_is_structured(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
model = Path(tmp) / DEFAULT_DENOISER_PATH
model.parent.mkdir(parents=True)
model.write_bytes(b"fake")
provider = SherpaOnnxDenoiserPreprocessor(tmp, sherpa_module=FailingSherpa)
provider.load()
with self.assertRaises(ProviderError) as raised:
provider.process_frame(AudioFrame(b"\xff\x7f", 16000, 1, 0, 0))
self.assertEqual(raised.exception.code, ErrorCode.NOISE_FILTER_FAILED)
if __name__ == "__main__":
unittest.main()
+2
View File
@@ -65,10 +65,12 @@ class CliAcceptanceTests(unittest.TestCase):
patch("owner_voice_pet.cli.SherpaOnnxKeywordWakeWordProvider") as wake_cls,
patch("owner_voice_pet.cli.SherpaOnnxVadProvider") as vad_cls,
patch("owner_voice_pet.cli.SherpaOnnxSttProvider") as stt_cls,
patch("owner_voice_pet.cli.SherpaOnnxDenoiserPreprocessor") as denoiser_cls,
):
wake_cls.return_value.load.return_value = None
vad_cls.return_value.load.return_value = None
stt_cls.return_value.load.return_value = None
denoiser_cls.return_value.load.return_value = None
code, data = self.call("model-check", "--models-dir", str(root))
self.assertEqual(code, 0)
self.assertTrue(data["ok"])
+54 -2
View File
@@ -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)
+41 -1
View File
@@ -1,7 +1,9 @@
from __future__ import annotations
import json
import tempfile
import unittest
from pathlib import Path
from owner_voice_pet.config import AppConfig
from owner_voice_pet.models import (
@@ -12,6 +14,12 @@ from owner_voice_pet.models import (
PipelineState,
ProviderError,
)
from owner_voice_pet.speech_models import (
DEFAULT_DENOISER_PATH,
default_manifest,
denoiser_model_path,
stt_model_paths,
)
class ModelsConfigTests(unittest.TestCase):
@@ -77,6 +85,9 @@ class ModelsConfigTests(unittest.TestCase):
self.assertEqual(config.post_playback_drain_ms, 0)
self.assertEqual(config.pipeline_mode, "live_turn_based")
self.assertEqual(config.endpoint_mode, "primary_speaker")
self.assertTrue(config.noise_filter_enabled)
self.assertEqual(config.noise_filter_provider, "sherpa_onnx_gtcrn")
self.assertFalse(config.wake_denoise_enabled)
self.assertEqual(config.speaker_profile_ms, 600)
self.assertEqual(config.speaker_profile_min_ms, 120)
self.assertEqual(config.speaker_absent_ms, 300)
@@ -88,7 +99,7 @@ class ModelsConfigTests(unittest.TestCase):
self.assertEqual(config.vad_end_silence_ms, 350)
self.assertEqual(config.vad_no_speech_timeout_ms, 5000)
self.assertEqual(config.vad_max_recording_ms, 12000)
self.assertEqual(config.speech_provider, "cloud")
self.assertEqual(config.speech_provider, "local")
self.assertEqual(config.asr_model, "mimo-v2.5-asr")
self.assertEqual(config.tts_model, "mimo-v2.5-tts")
self.assertEqual(config.tts_voice, "mimo_default")
@@ -117,6 +128,11 @@ class ModelsConfigTests(unittest.TestCase):
errors = config.validate_basic()
self.assertTrue(any("OWNER_ENDPOINT_MODE" in error.message for error in errors))
def test_noise_filter_provider_must_be_gtcrn(self) -> None:
config = AppConfig(noise_filter_provider="invalid")
errors = config.validate_basic()
self.assertTrue(any("OWNER_NOISE_FILTER_PROVIDER" in error.message for error in errors))
def test_speaker_similarity_threshold_range_is_validated(self) -> None:
config = AppConfig(speaker_similarity_threshold=1.5)
errors = config.validate_basic()
@@ -153,6 +169,30 @@ class ModelsConfigTests(unittest.TestCase):
message = Message(role="user", content="你好", created_at=1.0)
self.assertEqual(message.role, "user")
def test_default_manifest_uses_ctc_stt_and_denoiser(self) -> None:
manifest = default_manifest()
stt = manifest["providers"]["stt"]
denoiser = manifest["providers"]["denoiser"]
self.assertEqual(stt["type"], "sherpa-onnx-streaming-zipformer2-ctc")
self.assertTrue(stt["model"].endswith("model.int8.onnx"))
self.assertEqual(denoiser["path"], DEFAULT_DENOISER_PATH)
self.assertIn(DEFAULT_DENOISER_PATH, manifest["required_files"])
def test_model_path_helpers_support_ctc_manifest_and_denoiser(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
root = Path(tmp)
manifest_path = root / "manifest.json"
manifest_path.write_text(
json.dumps(default_manifest(), ensure_ascii=False),
encoding="utf-8",
)
paths = stt_model_paths(root)
self.assertEqual(paths["type"], "sherpa-onnx-streaming-zipformer2-ctc")
self.assertTrue(str(paths["model"]).endswith("model.int8.onnx"))
self.assertTrue(str(denoiser_model_path(root)).endswith(DEFAULT_DENOISER_PATH))
if __name__ == "__main__":
unittest.main()
+68 -7
View File
@@ -7,7 +7,13 @@ from pathlib import Path
from owner_voice_pet.models import AudioFrame, AudioSegment, ErrorCode, ProviderError
from owner_voice_pet.config import AppConfig
from owner_voice_pet.stt import CloudAsrSttProvider, MetadataSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
from owner_voice_pet.stt import (
CloudAsrSttProvider,
MetadataSttProvider,
SherpaOnnxSttProvider,
is_valid_transcript_text,
should_emit_partial_transcript,
)
from owner_voice_pet.speech_models import stt_model_paths
from owner_voice_pet.vad import EnergyVadProvider, HybridVadProvider, PrimarySpeakerVadRecorder, VadRecorder
from owner_voice_pet.wakeword import (
@@ -241,6 +247,13 @@ class WakeVadSttTests(unittest.TestCase):
self.assertTrue(is_valid_transcript_text("hello"))
self.assertFalse(is_valid_transcript_text("?! 。"))
def test_partial_transcript_filter_rejects_short_noise(self) -> None:
self.assertFalse(should_emit_partial_transcript("", ""))
self.assertFalse(should_emit_partial_transcript("家确", ""))
self.assertTrue(should_emit_partial_transcript("你是谁", ""))
self.assertFalse(should_emit_partial_transcript("", "你是谁"))
self.assertTrue(should_emit_partial_transcript("你是谁呀", "你是谁"))
def test_sherpa_stt_missing_model_is_structured(self) -> None:
with tempfile.TemporaryDirectory() as tmp:
provider = SherpaOnnxSttProvider(f"{tmp}/missing")
@@ -260,7 +273,7 @@ class WakeVadSttTests(unittest.TestCase):
def accept_waveform(self, sample_rate, samples) -> None:
self.ready = True
self.text = "" if not self.text else "你好"
self.text = "" if not self.text else "你好"
class FakeRecognizer:
def create_stream(self):
@@ -286,7 +299,7 @@ class WakeVadSttTests(unittest.TestCase):
with tempfile.TemporaryDirectory() as tmp:
paths = stt_model_paths(Path(tmp))
for name, path in paths.items():
if name == "model_dir":
if name in {"type", "model_dir"}:
continue
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text("fake", encoding="utf-8")
@@ -296,11 +309,59 @@ class WakeVadSttTests(unittest.TestCase):
first = session.accept_frame(make_frame(1, 0, speech=True))
second = session.accept_frame(make_frame(2, 20, speech=True))
self.assertIsNotNone(first)
self.assertIsNone(first)
self.assertIsNotNone(second)
assert first is not None and second is not None
self.assertEqual(first.text, "")
self.assertEqual(second.text, "你好")
assert second is not None
self.assertEqual(second.text, "好吗")
def test_sherpa_stt_loads_ctc_manifest(self) -> None:
class FakeRecognizer:
def create_stream(self):
return object()
class FakeOnlineRecognizer:
ctc_kwargs = None
@staticmethod
def from_zipformer2_ctc(**kwargs):
FakeOnlineRecognizer.ctc_kwargs = kwargs
return FakeRecognizer()
class FakeSherpa:
OnlineRecognizer = FakeOnlineRecognizer
with tempfile.TemporaryDirectory() as tmp:
root = Path(tmp)
stt_dir = root / "stt" / "ctc"
stt_dir.mkdir(parents=True)
(stt_dir / "tokens.txt").write_text("你 1\n", encoding="utf-8")
(stt_dir / "model.int8.onnx").write_bytes(b"fake")
(root / "manifest.json").write_text(
json.dumps(
{
"providers": {
"stt": {
"type": "sherpa-onnx-streaming-zipformer2-ctc",
"model_dir": "stt/ctc",
"tokens": "stt/ctc/tokens.txt",
"model": "stt/ctc/model.int8.onnx",
}
},
"required_files": [
"stt/ctc/tokens.txt",
"stt/ctc/model.int8.onnx",
],
},
ensure_ascii=False,
),
encoding="utf-8",
)
provider = SherpaOnnxSttProvider(tmp, sherpa_module=FakeSherpa)
provider.load()
self.assertIsNotNone(FakeOnlineRecognizer.ctc_kwargs)
assert FakeOnlineRecognizer.ctc_kwargs is not None
self.assertTrue(FakeOnlineRecognizer.ctc_kwargs["model"].endswith("model.int8.onnx"))
def test_cloud_asr_posts_audio_transcription_request(self) -> None:
class FakeResponse: