[本地唤醒模型]:完成KWS模型下载和检查,包含manifest、配置和model-check

This commit is contained in:
mkbk
2026-06-17 20:35:21 +08:00
parent 57e447b2fc
commit e565164e6e
11 changed files with 286 additions and 8 deletions
+4
View File
@@ -6,6 +6,10 @@ OWNER_AUDIO_INPUT_DEVICE=
OWNER_AUDIO_OUTPUT_DEVICE= OWNER_AUDIO_OUTPUT_DEVICE=
OWNER_ASSET_DIR=assets/pet OWNER_ASSET_DIR=assets/pet
OWNER_LOG_DIR=logs OWNER_LOG_DIR=logs
OWNER_WAKE_PROVIDER=local_kws
OWNER_WAKE_KEYWORDS_FILE=
OWNER_WAKE_KWS_THRESHOLD=0.25
OWNER_WAKE_KWS_SCORE=1.0
OWNER_SPEECH_PROVIDER=cloud OWNER_SPEECH_PROVIDER=cloud
OWNER_ASR_MODEL=mimo-v2.5-asr OWNER_ASR_MODEL=mimo-v2.5-asr
OWNER_TTS_MODEL=mimo-v2.5-tts OWNER_TTS_MODEL=mimo-v2.5-tts
@@ -10,11 +10,11 @@
## 2. 本地 KWS 模型与配置 ## 2. 本地 KWS 模型与配置
- [ ] 2.1 扩展 `speech_models.py` manifest 和路径 helper;前置条件:OpenSpec 提交完成;验收标准:包含 KWS URL、目录、tokens、encoder、decoder、joiner、keywords;测试要点:manifest/default paths 测试;优先级:P0;预计:45 分钟。 - [x] 2.1 扩展 `speech_models.py` manifest 和路径 helper;前置条件:OpenSpec 提交完成;验收标准:包含 KWS URL、目录、tokens、encoder、decoder、joiner、keywords;测试要点:manifest/default paths 测试;优先级:P0;预计:45 分钟。
- [ ] 2.2 扩展 `download_speech_models.py` 下载 KWS 模型并写入 keywords 文件;前置条件:2.1 完成;验收标准:脚本幂等,能补齐 `models/wake`;测试要点:真实下载或 skip existing;优先级:P0;预计:60 分钟。 - [x] 2.2 扩展 `download_speech_models.py` 下载 KWS 模型并写入 keywords 文件;前置条件:2.1 完成;验收标准:脚本幂等,能补齐 `models/wake`;测试要点:真实下载或 skip existing;优先级:P0;预计:60 分钟。
- [ ] 2.3 扩展 `AppConfig``.env.example` wake 配置;前置条件:字段确定;验收标准:默认 `local_kws`;测试要点:配置加载测试;优先级:P0;预计:30 分钟。 - [x] 2.3 扩展 `AppConfig``.env.example` wake 配置;前置条件:字段确定;验收标准:默认 `local_kws`;测试要点:配置加载测试;优先级:P0;预计:30 分钟。
- [ ] 2.4 扩展 `model-check` 检查 wake 模型并尝试加载;前置条件:2.1 至 2.3;验收标准:完整模型通过,缺模型结构化失败;测试要点:临时目录和真实模型;优先级:P0;预计:45 分钟。 - [x] 2.4 扩展 `model-check` 检查 wake 模型并尝试加载;前置条件:2.1 至 2.3;验收标准:完整模型通过,缺模型结构化失败;测试要点:临时目录和真实模型;优先级:P0;预计:45 分钟。
- [ ] 2.5 验证并提交“本地唤醒模型”模块;前置条件:2.1 至 2.4 完成;验收标准:compileall、相关 unittest、security-check、model-check、OpenSpec 通过后 commit;优先级:P0;预计:20 分钟。 - [x] 2.5 验证并提交“本地唤醒模型”模块;前置条件:2.1 至 2.4 完成;验收标准:compileall、相关 unittest、security-check、model-check、OpenSpec 通过后 commit;优先级:P0;预计:20 分钟。
## 3. Runtime 独立唤醒与实时转写 ## 3. Runtime 独立唤醒与实时转写
+36
View File
@@ -16,6 +16,9 @@ if str(SRC_DIR) not in sys.path:
sys.path.insert(0, str(SRC_DIR)) sys.path.insert(0, str(SRC_DIR))
from owner_voice_pet.speech_models import ( from owner_voice_pet.speech_models import (
DEFAULT_KWS_DIR,
DEFAULT_KWS_KEYWORDS,
DEFAULT_KWS_URL,
DEFAULT_STT_DIR, DEFAULT_STT_DIR,
DEFAULT_STT_URL, DEFAULT_STT_URL,
DEFAULT_VAD_URL, DEFAULT_VAD_URL,
@@ -32,6 +35,7 @@ def main(argv: list[str] | None = None) -> int:
target = Path(args.dir) target = Path(args.dir)
target.mkdir(parents=True, exist_ok=True) target.mkdir(parents=True, exist_ok=True)
download_kws(target, force=args.force)
download_vad(target, force=args.force) download_vad(target, force=args.force)
download_stt(target, force=args.force) download_stt(target, force=args.force)
manifest = write_default_manifest(target) manifest = write_default_manifest(target)
@@ -42,6 +46,38 @@ def main(argv: list[str] | None = None) -> int:
return 0 if not status.missing_files else 1 return 0 if not status.missing_files else 1
def download_kws(models_dir: Path, *, force: bool = False) -> Path:
output_dir = models_dir / "wake" / DEFAULT_KWS_DIR
if output_dir.exists() and any(output_dir.iterdir()) and not force:
print(f"skip existing {output_dir}")
else:
output_dir.parent.mkdir(parents=True, exist_ok=True)
last_error: Exception | None = None
for attempt in range(1, 4):
with tempfile.TemporaryDirectory() as tmp:
archive = Path(tmp) / f"{DEFAULT_KWS_DIR}.tar.bz2"
try:
download_file(DEFAULT_KWS_URL, archive)
extract_tar_bz2(archive, output_dir.parent)
last_error = None
break
except Exception as exc:
last_error = exc
print(f"download/extract KWS attempt {attempt} failed: {exc}")
if output_dir.exists():
shutil.rmtree(output_dir)
if last_error is not None:
raise last_error
keywords = models_dir / "wake" / "keywords.txt"
if force or not keywords.exists():
keywords.parent.mkdir(parents=True, exist_ok=True)
keywords.write_text(DEFAULT_KWS_KEYWORDS, encoding="utf-8")
print(f"wrote {keywords}")
else:
print(f"skip existing {keywords}")
return output_dir
def download_vad(models_dir: Path, *, force: bool = False) -> Path: def download_vad(models_dir: Path, *, force: bool = False) -> Path:
output = models_dir / "vad" / "silero_vad.onnx" output = models_dir / "vad" / "silero_vad.onnx"
if output.exists() and not force: if output.exists() and not force:
+2 -1
View File
@@ -16,7 +16,7 @@ from .models import (
WakeEvent, WakeEvent,
) )
from .transport import AudioRingBuffer, FileReplayTransport, MemoryAudioTransport from .transport import AudioRingBuffer, FileReplayTransport, MemoryAudioTransport
from .wakeword import KeywordWakeWordProvider from .wakeword import KeywordWakeWordProvider, SherpaOnnxKeywordWakeWordProvider
from .vad import EnergyVadProvider, VadRecorder from .vad import EnergyVadProvider, VadRecorder
from .stt import CloudAsrSttProvider, MetadataSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text from .stt import CloudAsrSttProvider, MetadataSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
from .conversation import ConversationContext from .conversation import ConversationContext
@@ -36,6 +36,7 @@ __all__ = [
"FileReplayTransport", "FileReplayTransport",
"MemoryAudioTransport", "MemoryAudioTransport",
"KeywordWakeWordProvider", "KeywordWakeWordProvider",
"SherpaOnnxKeywordWakeWordProvider",
"EnergyVadProvider", "EnergyVadProvider",
"VadRecorder", "VadRecorder",
"CloudAsrSttProvider", "CloudAsrSttProvider",
+16 -1
View File
@@ -18,7 +18,7 @@ from .stt import MetadataSttProvider, SherpaOnnxSttProvider
from .transport import MemoryAudioTransport, sounddevice_device_report from .transport import MemoryAudioTransport, sounddevice_device_report
from .tts import SineTtsProvider from .tts import SineTtsProvider
from .vad import EnergyVadProvider, SherpaOnnxVadProvider, VadRecorder from .vad import EnergyVadProvider, SherpaOnnxVadProvider, VadRecorder
from .wakeword import KeywordWakeWordProvider from .wakeword import KeywordWakeWordProvider, SherpaOnnxKeywordWakeWordProvider
def main(argv: list[str] | None = None) -> int: def main(argv: list[str] | None = None) -> int:
@@ -53,6 +53,10 @@ def main(argv: list[str] | None = None) -> int:
"llm_stream": config.llm_stream, "llm_stream": config.llm_stream,
"llm_api_key_present": bool(config.llm_api_key), "llm_api_key_present": bool(config.llm_api_key),
"asset_dir": str(config.asset_dir), "asset_dir": str(config.asset_dir),
"wake_provider": config.wake_provider,
"wake_keywords_file": str(config.wake_keywords_file) if config.wake_keywords_file else "",
"wake_kws_threshold": config.wake_kws_threshold,
"wake_kws_score": config.wake_kws_score,
"speech_provider": config.speech_provider, "speech_provider": config.speech_provider,
"asr_model": config.asr_model, "asr_model": config.asr_model,
"tts_model": config.tts_model, "tts_model": config.tts_model,
@@ -83,6 +87,13 @@ def main(argv: list[str] | None = None) -> int:
provider_load_checked = False provider_load_checked = False
if not errors: if not errors:
try: try:
SherpaOnnxKeywordWakeWordProvider(
models_dir,
keyword=config.wake_word,
keywords_file=config.wake_keywords_file,
threshold=config.wake_kws_threshold,
score=config.wake_kws_score,
).load()
SherpaOnnxVadProvider(models_dir).load() SherpaOnnxVadProvider(models_dir).load()
SherpaOnnxSttProvider(str(models_dir)).load() SherpaOnnxSttProvider(str(models_dir)).load()
provider_load_checked = True provider_load_checked = True
@@ -129,6 +140,10 @@ def main(argv: list[str] | None = None) -> int:
audio_output_device=config.audio_output_device, audio_output_device=config.audio_output_device,
asset_dir=config.asset_dir, asset_dir=config.asset_dir,
log_dir=config.log_dir, log_dir=config.log_dir,
wake_provider=config.wake_provider,
wake_keywords_file=config.wake_keywords_file,
wake_kws_threshold=config.wake_kws_threshold,
wake_kws_score=config.wake_kws_score,
speech_provider=config.speech_provider, speech_provider=config.speech_provider,
asr_model=config.asr_model, asr_model=config.asr_model,
tts_model=config.tts_model, tts_model=config.tts_model,
+38
View File
@@ -20,6 +20,10 @@ class AppConfig:
audio_output_device: str | None = None audio_output_device: str | None = None
asset_dir: Path = Path("assets/pet") asset_dir: Path = Path("assets/pet")
log_dir: Path = Path("logs") log_dir: Path = Path("logs")
wake_provider: str = "local_kws"
wake_keywords_file: Path | None = None
wake_kws_threshold: float = 0.25
wake_kws_score: float = 1.0
speech_provider: str = "cloud" speech_provider: str = "cloud"
asr_model: str = "mimo-v2.5-asr" asr_model: str = "mimo-v2.5-asr"
tts_model: str = "mimo-v2.5-tts" tts_model: str = "mimo-v2.5-tts"
@@ -49,6 +53,10 @@ class AppConfig:
audio_output_device=get("AUDIO_OUTPUT_DEVICE"), audio_output_device=get("AUDIO_OUTPUT_DEVICE"),
asset_dir=Path(get("ASSET_DIR", "assets/pet") or "assets/pet"), asset_dir=Path(get("ASSET_DIR", "assets/pet") or "assets/pet"),
log_dir=Path(get("LOG_DIR", "logs") or "logs"), log_dir=Path(get("LOG_DIR", "logs") or "logs"),
wake_provider=(get("WAKE_PROVIDER", "local_kws") or "local_kws").lower(),
wake_keywords_file=Path(value) if (value := get("WAKE_KEYWORDS_FILE")) else None,
wake_kws_threshold=float(get("WAKE_KWS_THRESHOLD", "0.25") or "0.25"),
wake_kws_score=float(get("WAKE_KWS_SCORE", "1.0") or "1.0"),
speech_provider=(get("SPEECH_PROVIDER", "cloud") or "cloud").lower(), speech_provider=(get("SPEECH_PROVIDER", "cloud") or "cloud").lower(),
asr_model=get("ASR_MODEL", "mimo-v2.5-asr") or "mimo-v2.5-asr", asr_model=get("ASR_MODEL", "mimo-v2.5-asr") or "mimo-v2.5-asr",
tts_model=get("TTS_MODEL", "mimo-v2.5-tts") or "mimo-v2.5-tts", tts_model=get("TTS_MODEL", "mimo-v2.5-tts") or "mimo-v2.5-tts",
@@ -114,6 +122,36 @@ class AppConfig:
"startup", "startup",
) )
) )
if self.wake_provider not in {"local_kws"}:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_WAKE_PROVIDER must be local_kws",
False,
"config",
"startup",
)
)
if self.wake_kws_threshold <= 0:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_WAKE_KWS_THRESHOLD must be positive",
False,
"config",
"startup",
)
)
if self.wake_kws_score <= 0:
errors.append(
ProviderError(
ErrorCode.CONFIG_MISSING_VALUE,
"OWNER_WAKE_KWS_SCORE must be positive",
False,
"config",
"startup",
)
)
if not self.llm_base_url.startswith(("http://", "https://")): if not self.llm_base_url.startswith(("http://", "https://")):
errors.append( errors.append(
ProviderError( ProviderError(
+38
View File
@@ -14,8 +14,19 @@ DEFAULT_STT_URL = (
"sherpa-onnx-streaming-zipformer-zh-14M-2023-02-23.tar.bz2" "sherpa-onnx-streaming-zipformer-zh-14M-2023-02-23.tar.bz2"
) )
DEFAULT_STT_DIR = "sherpa-onnx-streaming-zipformer-zh-14M-2023-02-23" DEFAULT_STT_DIR = "sherpa-onnx-streaming-zipformer-zh-14M-2023-02-23"
DEFAULT_KWS_URL = (
"https://github.com/k2-fsa/sherpa-onnx/releases/download/kws-models/"
"sherpa-onnx-kws-zipformer-wenetspeech-3.3M-2024-01-01-mobile.tar.bz2"
)
DEFAULT_KWS_DIR = "sherpa-onnx-kws-zipformer-wenetspeech-3.3M-2024-01-01-mobile"
DEFAULT_KWS_KEYWORDS = "x iǎo j ié x iǎo j ié @小杰小杰\n"
REQUIRED_MODEL_FILES = ( REQUIRED_MODEL_FILES = (
f"wake/{DEFAULT_KWS_DIR}/tokens.txt",
f"wake/{DEFAULT_KWS_DIR}/encoder-epoch-12-avg-2-chunk-16-left-64.int8.onnx",
f"wake/{DEFAULT_KWS_DIR}/decoder-epoch-12-avg-2-chunk-16-left-64.onnx",
f"wake/{DEFAULT_KWS_DIR}/joiner-epoch-12-avg-2-chunk-16-left-64.int8.onnx",
"wake/keywords.txt",
"vad/silero_vad.onnx", "vad/silero_vad.onnx",
f"stt/{DEFAULT_STT_DIR}/tokens.txt", f"stt/{DEFAULT_STT_DIR}/tokens.txt",
f"stt/{DEFAULT_STT_DIR}/encoder-epoch-99-avg-1.int8.onnx", f"stt/{DEFAULT_STT_DIR}/encoder-epoch-99-avg-1.int8.onnx",
@@ -54,10 +65,20 @@ def default_manifest() -> dict[str, Any]:
"schema": "owner_voice_pet.speech_models.v1", "schema": "owner_voice_pet.speech_models.v1",
"sample_rate": 16000, "sample_rate": 16000,
"sources": { "sources": {
"wake": DEFAULT_KWS_URL,
"vad": DEFAULT_VAD_URL, "vad": DEFAULT_VAD_URL,
"stt": DEFAULT_STT_URL, "stt": DEFAULT_STT_URL,
}, },
"providers": { "providers": {
"wake": {
"type": "sherpa-onnx-keyword-spotter",
"model_dir": f"wake/{DEFAULT_KWS_DIR}",
"tokens": f"wake/{DEFAULT_KWS_DIR}/tokens.txt",
"encoder": f"wake/{DEFAULT_KWS_DIR}/encoder-epoch-12-avg-2-chunk-16-left-64.int8.onnx",
"decoder": f"wake/{DEFAULT_KWS_DIR}/decoder-epoch-12-avg-2-chunk-16-left-64.onnx",
"joiner": f"wake/{DEFAULT_KWS_DIR}/joiner-epoch-12-avg-2-chunk-16-left-64.int8.onnx",
"keywords": "wake/keywords.txt",
},
"vad": { "vad": {
"type": "silero-vad", "type": "silero-vad",
"path": "vad/silero_vad.onnx", "path": "vad/silero_vad.onnx",
@@ -106,6 +127,23 @@ def vad_model_path(models_dir: str | Path) -> Path:
return root / str(path) return root / str(path)
def wake_model_paths(models_dir: str | Path) -> dict[str, Path]:
root = Path(models_dir)
manifest = load_manifest(root)
wake = manifest.get("providers", {}).get("wake", {})
return {
"model_dir": root / str(wake.get("model_dir", f"wake/{DEFAULT_KWS_DIR}")),
"tokens": root / str(wake.get("tokens", f"wake/{DEFAULT_KWS_DIR}/tokens.txt")),
"encoder": root
/ str(wake.get("encoder", f"wake/{DEFAULT_KWS_DIR}/encoder-epoch-12-avg-2-chunk-16-left-64.int8.onnx")),
"decoder": root
/ str(wake.get("decoder", f"wake/{DEFAULT_KWS_DIR}/decoder-epoch-12-avg-2-chunk-16-left-64.onnx")),
"joiner": root
/ str(wake.get("joiner", f"wake/{DEFAULT_KWS_DIR}/joiner-epoch-12-avg-2-chunk-16-left-64.int8.onnx")),
"keywords": root / str(wake.get("keywords", "wake/keywords.txt")),
}
def stt_model_paths(model_path: str | Path) -> dict[str, Path]: def stt_model_paths(model_path: str | Path) -> dict[str, Path]:
root = Path(model_path) root = Path(model_path)
if (root / "manifest.json").exists() or (root / "stt").exists(): if (root / "manifest.json").exists() or (root / "stt").exists():
+125
View File
@@ -1,6 +1,10 @@
from __future__ import annotations from __future__ import annotations
from pathlib import Path
from typing import Any
from .models import AudioFrame, ErrorCode, ProviderError, WakeEvent from .models import AudioFrame, ErrorCode, ProviderError, WakeEvent
from .speech_models import wake_model_paths
class KeywordWakeWordProvider: class KeywordWakeWordProvider:
@@ -51,3 +55,124 @@ class MissingWakeWordModelProvider:
def reset(self) -> None: def reset(self) -> None:
return None return None
class SherpaOnnxKeywordWakeWordProvider:
def __init__(
self,
models_dir: str | Path,
*,
keyword: str = "小杰小杰",
keywords_file: str | Path | None = None,
threshold: float = 0.25,
score: float = 1.0,
sherpa_module: Any | None = None,
) -> None:
self.models_dir = Path(models_dir)
self.keyword = keyword
self.keywords_file = Path(keywords_file) if keywords_file else None
self.threshold = threshold
self.score = score
self.loaded = False
self._sherpa = sherpa_module
self._spotter: Any | None = None
self._stream: Any | None = None
def load(self) -> None:
paths = wake_model_paths(self.models_dir)
if self.keywords_file is not None:
paths["keywords"] = self.keywords_file
missing = [name for name, path in paths.items() if name != "model_dir" and not path.exists()]
if missing:
raise ProviderError(
ErrorCode.WAKE_MODEL_MISSING,
"sherpa-onnx KWS model files are missing: " + ", ".join(missing),
False,
"sherpa-onnx-kws",
"wakeword",
)
if not paths["keywords"].read_text(encoding="utf-8").strip():
raise ProviderError(
ErrorCode.WAKE_MODEL_MISSING,
f"wake keyword file is empty: {paths['keywords']}",
False,
"sherpa-onnx-kws",
"wakeword",
)
sherpa_onnx = self._sherpa
if sherpa_onnx is None:
try:
import sherpa_onnx # type: ignore[import-not-found]
except Exception as exc:
raise ProviderError(
ErrorCode.WAKE_MODEL_LOAD_FAILED,
f"sherpa_onnx is not available: {exc}",
False,
"sherpa-onnx-kws",
"wakeword",
) from exc
try:
self._spotter = sherpa_onnx.KeywordSpotter(
tokens=str(paths["tokens"]),
encoder=str(paths["encoder"]),
decoder=str(paths["decoder"]),
joiner=str(paths["joiner"]),
keywords_file=str(paths["keywords"]),
num_threads=2,
sample_rate=16000,
keywords_score=self.score,
keywords_threshold=self.threshold,
provider="cpu",
)
self._stream = self._spotter.create_stream()
except Exception as exc:
raise ProviderError(
ErrorCode.WAKE_MODEL_LOAD_FAILED,
f"failed to load sherpa-onnx KWS model: {exc}",
False,
"sherpa-onnx-kws",
"wakeword",
) from exc
self.loaded = True
def detect(self, frame: AudioFrame) -> WakeEvent | None:
if not self.loaded or self._spotter is None or self._stream is None:
raise ProviderError(
ErrorCode.WAKE_MODEL_LOAD_FAILED,
"sherpa-onnx KWS provider is not loaded",
False,
"sherpa-onnx-kws",
"wakeword",
)
try:
import numpy as np
samples = np.frombuffer(frame.pcm, dtype=np.int16).astype(np.float32) / 32768.0
if frame.channels > 1 and samples.size:
samples = samples.reshape(-1, frame.channels).mean(axis=1)
self._stream.accept_waveform(frame.sample_rate, samples)
while self._spotter.is_ready(self._stream):
self._spotter.decode_stream(self._stream)
result = str(self._spotter.get_result(self._stream) or "").strip()
if result:
self._spotter.reset_stream(self._stream)
if not self.keyword:
return WakeEvent(result, 1.0, frame.timestamp_ms)
if self.keyword in result or result in self.keyword:
return WakeEvent(self.keyword, 1.0, frame.timestamp_ms)
except Exception as exc:
raise ProviderError(
ErrorCode.WAKE_MODEL_LOAD_FAILED,
f"sherpa-onnx KWS detection failed: {exc}",
True,
"sherpa-onnx-kws",
"wakeword",
) from exc
return None
def reset(self) -> None:
if self._spotter is not None and self._stream is not None:
try:
self._spotter.reset_stream(self._stream)
except Exception:
self._stream = self._spotter.create_stream()
+2
View File
@@ -59,9 +59,11 @@ class CliAcceptanceTests(unittest.TestCase):
path.write_bytes(b"placeholder") path.write_bytes(b"placeholder")
with ( with (
patch("importlib.util.find_spec", return_value=object()), 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.SherpaOnnxVadProvider") as vad_cls,
patch("owner_voice_pet.cli.SherpaOnnxSttProvider") as stt_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 vad_cls.return_value.load.return_value = None
stt_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)) code, data = self.call("model-check", "--models-dir", str(root))
+8
View File
@@ -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_base_url, "https://token-plan-cn.xiaomimimo.com/v1")
self.assertEqual(config.llm_api_key, "secret-value") self.assertEqual(config.llm_api_key, "secret-value")
self.assertEqual(config.llm_model, "test-model") 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.speech_provider, "cloud")
self.assertEqual(config.asr_model, "mimo-v2.5-asr") self.assertEqual(config.asr_model, "mimo-v2.5-asr")
self.assertEqual(config.tts_model, "mimo-v2.5-tts") self.assertEqual(config.tts_model, "mimo-v2.5-tts")
@@ -82,6 +85,11 @@ class ModelsConfigTests(unittest.TestCase):
errors = config.validate_basic() errors = config.validate_basic()
self.assertTrue(any("OWNER_SPEECH_PROVIDER" in error.message for error in errors)) 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: def test_missing_dotenv_uses_non_secret_defaults(self) -> None:
config = AppConfig.from_dotenv("/tmp/owner-voice-pet-missing.env") config = AppConfig.from_dotenv("/tmp/owner-voice-pet-missing.env")
self.assertEqual(config.llm_base_url, "https://token-plan-cn.xiaomimimo.com/v1") self.assertEqual(config.llm_base_url, "https://token-plan-cn.xiaomimimo.com/v1")
+12 -1
View File
@@ -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.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
from owner_voice_pet.vad import EnergyVadProvider, VadRecorder 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( def make_frame(
@@ -47,6 +51,13 @@ class WakeVadSttTests(unittest.TestCase):
MissingWakeWordModelProvider("/missing/model.onnx").load() MissingWakeWordModelProvider("/missing/model.onnx").load()
self.assertEqual(raised.exception.code, ErrorCode.WAKE_MODEL_MISSING) 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: def test_vad_recorder_returns_segment_after_silence(self) -> None:
provider = EnergyVadProvider() provider = EnergyVadProvider()
provider.load() provider.load()