mirror of
https://github.com/xinnan-tech/xiaozhi-esp32-server.git
synced 2026-07-28 01:53:53 +08:00
update:test迁移重命名为digital-human
This commit is contained in:
@@ -0,0 +1,3 @@
|
||||
from .detector_assets import DetectorAssets, DetectorAssetsBuilder
|
||||
from .detector import WakewordDetector
|
||||
from .microphone import MicrophoneListener
|
||||
@@ -0,0 +1,191 @@
|
||||
import queue
|
||||
import threading
|
||||
import logging
|
||||
import time
|
||||
from typing import Any, Callable
|
||||
|
||||
import numpy as np
|
||||
|
||||
from ..config import RuntimeConfig
|
||||
from .detector_assets import DetectorAssetsBuilder
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class WakewordDetector:
|
||||
def __init__(self, config: RuntimeConfig) -> None:
|
||||
self.config = config
|
||||
self.enabled = config.wakeword_enabled
|
||||
self.assets_builder = DetectorAssetsBuilder(config)
|
||||
self.audio_source = None
|
||||
self.keyword_spotter: Any | None = None
|
||||
self.stream: Any | None = None
|
||||
self.is_running_flag = False
|
||||
self.paused = False
|
||||
self.on_detected_callback: Callable[[str, str], None] | None = None
|
||||
self.on_error: Callable[[Exception], None] | None = None
|
||||
self._audio_queue: queue.Queue[np.ndarray] = queue.Queue(maxsize=100)
|
||||
self._worker_thread: threading.Thread | None = None
|
||||
self.last_detection_time = 0.0
|
||||
self.detection_cooldown = self.config.detector.cooldown_seconds
|
||||
|
||||
def initialize(self) -> None:
|
||||
if not self.enabled:
|
||||
raise RuntimeError("wakeword detector is disabled")
|
||||
|
||||
try:
|
||||
import sherpa_onnx
|
||||
except ImportError as exc:
|
||||
raise RuntimeError(
|
||||
"Missing dependency: sherpa-onnx. Install runtime dependencies before initializing detector."
|
||||
) from exc
|
||||
|
||||
assets = self.assets_builder.prepare()
|
||||
detector_cfg = self.config.detector
|
||||
|
||||
self.keyword_spotter = sherpa_onnx.KeywordSpotter(
|
||||
tokens=str(assets.tokens_file),
|
||||
encoder=str(assets.encoder_file),
|
||||
decoder=str(assets.decoder_file),
|
||||
joiner=str(assets.joiner_file),
|
||||
keywords_file=str(assets.keywords_file),
|
||||
num_threads=detector_cfg.num_threads,
|
||||
sample_rate=self.config.audio.sample_rate,
|
||||
feature_dim=80,
|
||||
max_active_paths=detector_cfg.max_active_paths,
|
||||
keywords_score=detector_cfg.keywords_score,
|
||||
keywords_threshold=detector_cfg.keywords_threshold,
|
||||
num_trailing_blanks=detector_cfg.num_trailing_blanks,
|
||||
provider=detector_cfg.provider,
|
||||
)
|
||||
self.stream = self.keyword_spotter.create_stream()
|
||||
logger.info("detector initialized")
|
||||
logger.info("detector model root: %s", assets.model_root)
|
||||
logger.info("detector keywords file: %s", assets.keywords_file)
|
||||
|
||||
def on_detected(self, callback: Callable[[str, str], None]) -> None:
|
||||
self.on_detected_callback = callback
|
||||
|
||||
def on_audio_data(self, audio_data: np.ndarray) -> None:
|
||||
if not self.enabled or not self.is_running_flag or self.paused:
|
||||
return
|
||||
|
||||
try:
|
||||
self._audio_queue.put_nowait(audio_data.copy())
|
||||
except queue.Full:
|
||||
try:
|
||||
self._audio_queue.get_nowait()
|
||||
self._audio_queue.put_nowait(audio_data.copy())
|
||||
except queue.Empty:
|
||||
pass
|
||||
except Exception as exc:
|
||||
logger.debug("audio data enqueue failed: %s", exc)
|
||||
|
||||
def start(self, audio_source) -> None:
|
||||
if not self.enabled:
|
||||
logger.info("wakeword detector disabled")
|
||||
return
|
||||
|
||||
if self.keyword_spotter is None or self.stream is None:
|
||||
self.initialize()
|
||||
|
||||
if self.is_running_flag:
|
||||
return
|
||||
|
||||
self.audio_source = audio_source
|
||||
self.audio_source.add_audio_listener(self)
|
||||
self.is_running_flag = True
|
||||
self.paused = False
|
||||
self._worker_thread = threading.Thread(
|
||||
target=self._detection_loop,
|
||||
name="wakeword-detector",
|
||||
daemon=True,
|
||||
)
|
||||
self._worker_thread.start()
|
||||
logger.info("wakeword detector started")
|
||||
|
||||
def stop(self) -> None:
|
||||
if not self.is_running_flag and self.audio_source is None and self._worker_thread is None:
|
||||
return
|
||||
|
||||
self.is_running_flag = False
|
||||
|
||||
if self.audio_source is not None:
|
||||
self.audio_source.remove_audio_listener(self)
|
||||
self.audio_source = None
|
||||
|
||||
if self._worker_thread is not None:
|
||||
self._worker_thread.join(timeout=1.0)
|
||||
self._worker_thread = None
|
||||
|
||||
while not self._audio_queue.empty():
|
||||
try:
|
||||
self._audio_queue.get_nowait()
|
||||
except queue.Empty:
|
||||
break
|
||||
|
||||
logger.info("wakeword detector stopped")
|
||||
|
||||
def _detection_loop(self) -> None:
|
||||
error_count = 0
|
||||
max_errors = 5
|
||||
|
||||
while self.is_running_flag:
|
||||
try:
|
||||
if self.paused:
|
||||
time.sleep(0.1)
|
||||
continue
|
||||
|
||||
audio_data = self._audio_queue.get(timeout=0.1)
|
||||
result = self.process_audio_chunk(audio_data)
|
||||
if result and self.on_detected_callback is not None:
|
||||
self.on_detected_callback(result, result)
|
||||
error_count = 0
|
||||
except queue.Empty:
|
||||
continue
|
||||
except Exception as exc:
|
||||
error_count += 1
|
||||
logger.error("wakeword detection loop error(%s/%s): %s", error_count, max_errors, exc)
|
||||
if self.on_error is not None:
|
||||
try:
|
||||
self.on_error(exc)
|
||||
except Exception:
|
||||
logger.exception("wakeword error callback failed")
|
||||
if error_count >= max_errors:
|
||||
logger.critical("too many wakeword detection errors, stopping detector")
|
||||
break
|
||||
time.sleep(1)
|
||||
|
||||
self.is_running_flag = False
|
||||
|
||||
def process_audio_chunk(self, audio_data: np.ndarray) -> str | None:
|
||||
if self.keyword_spotter is None or self.stream is None:
|
||||
raise RuntimeError("detector is not initialized")
|
||||
|
||||
if audio_data is None or len(audio_data) == 0:
|
||||
return None
|
||||
|
||||
if audio_data.dtype == np.int16:
|
||||
samples = audio_data.astype(np.float32) / 32768.0
|
||||
else:
|
||||
samples = audio_data.astype(np.float32)
|
||||
|
||||
sample_rate = self.config.audio.sample_rate
|
||||
self.stream.accept_waveform(sample_rate=sample_rate, waveform=samples)
|
||||
|
||||
if not self.keyword_spotter.is_ready(self.stream):
|
||||
return None
|
||||
|
||||
self.keyword_spotter.decode_stream(self.stream)
|
||||
result = self.keyword_spotter.get_result(self.stream)
|
||||
if not result:
|
||||
return None
|
||||
|
||||
self.keyword_spotter.reset_stream(self.stream)
|
||||
|
||||
current_time = time.time()
|
||||
if current_time - self.last_detection_time < self.detection_cooldown:
|
||||
return None
|
||||
|
||||
self.last_detection_time = current_time
|
||||
return str(result)
|
||||
@@ -0,0 +1,94 @@
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
from ..config import RuntimeConfig
|
||||
|
||||
REQUIRED_MODEL_FILES = (
|
||||
"encoder.onnx",
|
||||
"decoder.onnx",
|
||||
"joiner.onnx",
|
||||
"tokens.txt",
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DetectorAssets:
|
||||
model_root: Path
|
||||
tokens_file: Path
|
||||
encoder_file: Path
|
||||
decoder_file: Path
|
||||
joiner_file: Path
|
||||
keywords_file: Path
|
||||
|
||||
|
||||
class DetectorAssetsBuilder:
|
||||
def __init__(self, config: RuntimeConfig) -> None:
|
||||
self.config = config
|
||||
|
||||
def prepare(self) -> DetectorAssets:
|
||||
model_root = self._resolve_model_root()
|
||||
keywords_file = self._write_keywords_file(model_root)
|
||||
return DetectorAssets(
|
||||
model_root=model_root,
|
||||
tokens_file=model_root / "tokens.txt",
|
||||
encoder_file=model_root / "encoder.onnx",
|
||||
decoder_file=model_root / "decoder.onnx",
|
||||
joiner_file=model_root / "joiner.onnx",
|
||||
keywords_file=keywords_file,
|
||||
)
|
||||
|
||||
def _resolve_model_root(self) -> Path:
|
||||
preferred = self.config.model_dir
|
||||
if self._has_required_files(preferred):
|
||||
return preferred
|
||||
|
||||
raise FileNotFoundError(
|
||||
"No valid model directory found. Expected configured model files to exist."
|
||||
)
|
||||
|
||||
def _write_keywords_file(self, model_root: Path) -> Path:
|
||||
try:
|
||||
from pypinyin import Style, pinyin
|
||||
except ImportError as exc:
|
||||
raise RuntimeError(
|
||||
"Missing dependency: pypinyin. Install runtime dependencies before generating keywords."
|
||||
) from exc
|
||||
|
||||
wake_words = self.config.wake_words
|
||||
if not wake_words:
|
||||
raise ValueError("keywords.txt cannot be empty")
|
||||
|
||||
keywords_path = model_root / "keywords.txt"
|
||||
lines: list[str] = []
|
||||
for keyword_text in wake_words:
|
||||
initials = pinyin(keyword_text, style=Style.INITIALS, strict=False)
|
||||
finals = pinyin(
|
||||
keyword_text,
|
||||
style=Style.FINALS_TONE,
|
||||
strict=False,
|
||||
neutral_tone_with_five=True,
|
||||
)
|
||||
|
||||
tokens: list[str] = []
|
||||
for initial_parts, final_parts in zip(initials, finals):
|
||||
initial = initial_parts[0].strip()
|
||||
final = final_parts[0].strip()
|
||||
if initial:
|
||||
tokens.append(initial)
|
||||
if final:
|
||||
tokens.append(final)
|
||||
|
||||
if not tokens:
|
||||
raise ValueError(
|
||||
f"failed to generate pinyin tokens for wake word: {keyword_text}"
|
||||
)
|
||||
|
||||
lines.append(f"{' '.join(tokens)} @{keyword_text}")
|
||||
|
||||
keywords_path.write_text("\n".join(lines) + "\n", encoding="utf-8")
|
||||
return keywords_path
|
||||
|
||||
def _has_required_files(self, directory: Path) -> bool:
|
||||
if not directory.exists():
|
||||
return False
|
||||
return all((directory / file_name).exists() for file_name in REQUIRED_MODEL_FILES)
|
||||
@@ -0,0 +1,91 @@
|
||||
import logging
|
||||
from typing import Protocol
|
||||
|
||||
import numpy as np
|
||||
|
||||
from ..config import RuntimeConfig
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class AudioListener(Protocol):
|
||||
def on_audio_data(self, audio_data: np.ndarray) -> None:
|
||||
...
|
||||
|
||||
|
||||
class MicrophoneListener:
|
||||
|
||||
def __init__(self, config: RuntimeConfig) -> None:
|
||||
self.config = config
|
||||
self._stream = None
|
||||
self._running = False
|
||||
self._listeners: list[AudioListener] = []
|
||||
self._sample_rate = self.config.audio.sample_rate
|
||||
self._channels = self.config.audio.channels
|
||||
self._device = self.config.audio.input_device
|
||||
self._block_duration_ms = 30
|
||||
self._block_size = int(self._sample_rate * (self._block_duration_ms / 1000))
|
||||
|
||||
def add_audio_listener(self, listener: AudioListener) -> None:
|
||||
if listener not in self._listeners:
|
||||
self._listeners.append(listener)
|
||||
|
||||
def remove_audio_listener(self, listener: AudioListener) -> None:
|
||||
if listener in self._listeners:
|
||||
self._listeners.remove(listener)
|
||||
|
||||
def start(self) -> None:
|
||||
try:
|
||||
import sounddevice as sd
|
||||
except ImportError as exc:
|
||||
raise RuntimeError(
|
||||
"Missing dependency: sounddevice. Install runtime dependencies before starting microphone listener."
|
||||
) from exc
|
||||
|
||||
self._stream = sd.InputStream(
|
||||
device=self._device,
|
||||
samplerate=self._sample_rate,
|
||||
channels=self._channels,
|
||||
dtype="int16",
|
||||
blocksize=self._block_size,
|
||||
callback=self._input_callback,
|
||||
latency="low",
|
||||
)
|
||||
self._stream.start()
|
||||
self._running = True
|
||||
|
||||
logger.info("microphone listener started")
|
||||
logger.info("microphone sample rate: %s", self._sample_rate)
|
||||
logger.info("microphone channels: %s", self._channels)
|
||||
logger.info("microphone device: %s", self._device if self._device is not None else "default")
|
||||
logger.info("microphone block size: %s", self._block_size)
|
||||
|
||||
def stop(self) -> None:
|
||||
if not self._running and self._stream is None:
|
||||
return
|
||||
|
||||
self._running = False
|
||||
if self._stream is not None:
|
||||
try:
|
||||
self._stream.stop()
|
||||
finally:
|
||||
self._stream.close()
|
||||
self._stream = None
|
||||
logger.info("microphone listener stopped")
|
||||
|
||||
def _input_callback(self, indata, frames, time_info, status) -> None:
|
||||
_ = frames, time_info
|
||||
if status and "overflow" not in str(status).lower():
|
||||
logger.warning("microphone status: %s", status)
|
||||
|
||||
audio = np.copy(indata)
|
||||
if audio.ndim > 1:
|
||||
audio = audio[:, 0]
|
||||
else:
|
||||
audio = audio.reshape(-1)
|
||||
|
||||
for listener in list(self._listeners):
|
||||
try:
|
||||
listener.on_audio_data(audio)
|
||||
except Exception:
|
||||
logger.exception("audio listener callback failed")
|
||||
Reference in New Issue
Block a user