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