[本地语音降噪]:完成本地语音链路和噪音过滤,包含降噪模型、本地ASR和实时字幕稳定策略
This commit is contained in:
@@ -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()
|
||||
@@ -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"])
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user