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