[本地唤醒模型]:完成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
+2 -1
View File
@@ -16,7 +16,7 @@ from .models import (
WakeEvent,
)
from .transport import AudioRingBuffer, FileReplayTransport, MemoryAudioTransport
from .wakeword import KeywordWakeWordProvider
from .wakeword import KeywordWakeWordProvider, SherpaOnnxKeywordWakeWordProvider
from .vad import EnergyVadProvider, VadRecorder
from .stt import CloudAsrSttProvider, MetadataSttProvider, SherpaOnnxSttProvider, is_valid_transcript_text
from .conversation import ConversationContext
@@ -36,6 +36,7 @@ __all__ = [
"FileReplayTransport",
"MemoryAudioTransport",
"KeywordWakeWordProvider",
"SherpaOnnxKeywordWakeWordProvider",
"EnergyVadProvider",
"VadRecorder",
"CloudAsrSttProvider",
+16 -1
View File
@@ -18,7 +18,7 @@ from .stt import MetadataSttProvider, SherpaOnnxSttProvider
from .transport import MemoryAudioTransport, sounddevice_device_report
from .tts import SineTtsProvider
from .vad import EnergyVadProvider, SherpaOnnxVadProvider, VadRecorder
from .wakeword import KeywordWakeWordProvider
from .wakeword import KeywordWakeWordProvider, SherpaOnnxKeywordWakeWordProvider
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_api_key_present": bool(config.llm_api_key),
"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,
"asr_model": config.asr_model,
"tts_model": config.tts_model,
@@ -83,6 +87,13 @@ def main(argv: list[str] | None = None) -> int:
provider_load_checked = False
if not errors:
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()
SherpaOnnxSttProvider(str(models_dir)).load()
provider_load_checked = True
@@ -129,6 +140,10 @@ def main(argv: list[str] | None = None) -> int:
audio_output_device=config.audio_output_device,
asset_dir=config.asset_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,
asr_model=config.asr_model,
tts_model=config.tts_model,
+38
View File
@@ -20,6 +20,10 @@ class AppConfig:
audio_output_device: str | None = None
asset_dir: Path = Path("assets/pet")
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"
asr_model: str = "mimo-v2.5-asr"
tts_model: str = "mimo-v2.5-tts"
@@ -49,6 +53,10 @@ class AppConfig:
audio_output_device=get("AUDIO_OUTPUT_DEVICE"),
asset_dir=Path(get("ASSET_DIR", "assets/pet") or "assets/pet"),
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(),
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",
@@ -114,6 +122,36 @@ class AppConfig:
"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://")):
errors.append(
ProviderError(
+38
View File
@@ -14,8 +14,19 @@ DEFAULT_STT_URL = (
"sherpa-onnx-streaming-zipformer-zh-14M-2023-02-23.tar.bz2"
)
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 = (
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",
f"stt/{DEFAULT_STT_DIR}/tokens.txt",
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",
"sample_rate": 16000,
"sources": {
"wake": DEFAULT_KWS_URL,
"vad": DEFAULT_VAD_URL,
"stt": DEFAULT_STT_URL,
},
"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": {
"type": "silero-vad",
"path": "vad/silero_vad.onnx",
@@ -106,6 +127,23 @@ def vad_model_path(models_dir: str | Path) -> 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]:
root = Path(model_path)
if (root / "manifest.json").exists() or (root / "stt").exists():
+125
View File
@@ -1,6 +1,10 @@
from __future__ import annotations
from pathlib import Path
from typing import Any
from .models import AudioFrame, ErrorCode, ProviderError, WakeEvent
from .speech_models import wake_model_paths
class KeywordWakeWordProvider:
@@ -51,3 +55,124 @@ class MissingWakeWordModelProvider:
def reset(self) -> 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()