[本地唤醒模型]:完成KWS模型下载和检查,包含manifest、配置和model-check
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user