[本地唤醒模型]:完成KWS模型下载和检查,包含manifest、配置和model-check
This commit is contained in:
@@ -59,9 +59,11 @@ class CliAcceptanceTests(unittest.TestCase):
|
||||
path.write_bytes(b"placeholder")
|
||||
with (
|
||||
patch("importlib.util.find_spec", return_value=object()),
|
||||
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,
|
||||
):
|
||||
wake_cls.return_value.load.return_value = None
|
||||
vad_cls.return_value.load.return_value = None
|
||||
stt_cls.return_value.load.return_value = None
|
||||
code, data = self.call("model-check", "--models-dir", str(root))
|
||||
|
||||
@@ -69,6 +69,9 @@ class ModelsConfigTests(unittest.TestCase):
|
||||
self.assertEqual(config.llm_base_url, "https://token-plan-cn.xiaomimimo.com/v1")
|
||||
self.assertEqual(config.llm_api_key, "secret-value")
|
||||
self.assertEqual(config.llm_model, "test-model")
|
||||
self.assertEqual(config.wake_provider, "local_kws")
|
||||
self.assertEqual(config.wake_kws_threshold, 0.25)
|
||||
self.assertEqual(config.wake_kws_score, 1.0)
|
||||
self.assertEqual(config.speech_provider, "cloud")
|
||||
self.assertEqual(config.asr_model, "mimo-v2.5-asr")
|
||||
self.assertEqual(config.tts_model, "mimo-v2.5-tts")
|
||||
@@ -82,6 +85,11 @@ class ModelsConfigTests(unittest.TestCase):
|
||||
errors = config.validate_basic()
|
||||
self.assertTrue(any("OWNER_SPEECH_PROVIDER" in error.message for error in errors))
|
||||
|
||||
def test_wake_provider_must_be_local_kws(self) -> None:
|
||||
config = AppConfig(wake_provider="cloud_asr")
|
||||
errors = config.validate_basic()
|
||||
self.assertTrue(any("OWNER_WAKE_PROVIDER" in error.message for error in errors))
|
||||
|
||||
def test_missing_dotenv_uses_non_secret_defaults(self) -> None:
|
||||
config = AppConfig.from_dotenv("/tmp/owner-voice-pet-missing.env")
|
||||
self.assertEqual(config.llm_base_url, "https://token-plan-cn.xiaomimimo.com/v1")
|
||||
|
||||
@@ -8,7 +8,11 @@ from owner_voice_pet.models import AudioFrame, AudioSegment, ErrorCode, Provider
|
||||
from owner_voice_pet.config import AppConfig
|
||||
from owner_voice_pet.stt import CloudAsrSttProvider, MetadataSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
|
||||
from owner_voice_pet.vad import EnergyVadProvider, VadRecorder
|
||||
from owner_voice_pet.wakeword import KeywordWakeWordProvider, MissingWakeWordModelProvider
|
||||
from owner_voice_pet.wakeword import (
|
||||
KeywordWakeWordProvider,
|
||||
MissingWakeWordModelProvider,
|
||||
SherpaOnnxKeywordWakeWordProvider,
|
||||
)
|
||||
|
||||
|
||||
def make_frame(
|
||||
@@ -47,6 +51,13 @@ class WakeVadSttTests(unittest.TestCase):
|
||||
MissingWakeWordModelProvider("/missing/model.onnx").load()
|
||||
self.assertEqual(raised.exception.code, ErrorCode.WAKE_MODEL_MISSING)
|
||||
|
||||
def test_sherpa_kws_missing_model_is_structured(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
provider = SherpaOnnxKeywordWakeWordProvider(tmp)
|
||||
with self.assertRaises(ProviderError) as raised:
|
||||
provider.load()
|
||||
self.assertEqual(raised.exception.code, ErrorCode.WAKE_MODEL_MISSING)
|
||||
|
||||
def test_vad_recorder_returns_segment_after_silence(self) -> None:
|
||||
provider = EnergyVadProvider()
|
||||
provider.load()
|
||||
|
||||
Reference in New Issue
Block a user