detector.py 26 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661
  1. """
  2. Main wakeword detector class with full system integration.
  3. This module provides the complete WakewordDetector class that integrates
  4. all components of the wakeword detection system including audio processing,
  5. model inference, event handling, and performance monitoring.
  6. """
  7. import threading
  8. import time
  9. import traceback
  10. import warnings
  11. from typing import Optional, Callable, Dict, Any, List
  12. from pathlib import Path
  13. import logging
  14. import torch
  15. import numpy as np
  16. from .config import WakewordConfig, create_default_config
  17. from .model_loader import SecureModelLoader, ModelMetadata, ModelLoadError
  18. from .audio_features import AudioFeatureExtractor, create_feature_extractor
  19. from .audio_buffer import CircularAudioBuffer, AudioChunk, create_audio_buffer
  20. from .detection_engine import (
  21. WakewordDetectionEngine, DetectionResult, WakewordType,
  22. create_detection_engine
  23. )
  24. from ..events import WakewordEventData as WakewordReceivedEventData
  25. from ...events.event_data import SpeakerInfo, SatelliteInfo
  26. from ...events.decorators import register_event_handlers, unregister_event_handlers
  27. class WakewordDetectorError(Exception):
  28. """Exception raised by WakewordDetector."""
  29. pass
  30. class WakewordDetector:
  31. """
  32. Complete wakeword detection system for Trixy voice assistant.
  33. This class provides a high-level interface for real-time wakeword detection
  34. with automatic model loading, audio processing, and event integration.
  35. Features:
  36. - Real-time audio streaming and buffering
  37. - Dual wakeword detection (custom + system_command)
  38. - Password-protected model loading
  39. - Event system integration
  40. - Performance monitoring and optimization
  41. - Error handling and recovery
  42. - Configurable confidence thresholds
  43. Usage:
  44. # Basic usage
  45. config = WakewordConfig()
  46. config.model_config.model_path = "path/to/model.pth"
  47. config.model_config.model_password = "password"
  48. detector = WakewordDetector(config, event_handler)
  49. detector.start()
  50. # Process audio data
  51. detector.process_audio(audio_data)
  52. # Stop detection
  53. detector.stop()
  54. """
  55. def __init__(self, config: Optional[WakewordConfig] = None,
  56. event_handler: Optional[Any] = None,
  57. logger: Optional[logging.Logger] = None):
  58. """
  59. Initialize the wakeword detector.
  60. Args:
  61. config: Wakeword detection configuration
  62. event_handler: Event handler for triggering events
  63. logger: Logger instance for debugging
  64. """
  65. self.config = config or create_default_config()
  66. self.event_handler = event_handler
  67. self.logger = logger or logging.getLogger(__name__)
  68. # Validate configuration (skip overall validation due to AudioConfig issue, rely on ML Manager validation)
  69. # Note: ML Manager already validates sub-configs individually before creating the detector
  70. # if not self.config.validate():
  71. # raise WakewordDetectorError("Invalid configuration provided")
  72. # Initialize state
  73. self._is_running = False
  74. self._is_initialized = False
  75. self._lock = threading.RLock()
  76. self._audio_thread = None
  77. self._processing_thread = None
  78. # Components (initialized in _initialize_components)
  79. self.device = None
  80. self.model_loader = None
  81. self.model = None
  82. self.model_metadata = None
  83. self.feature_extractor = None
  84. self.audio_buffer = None
  85. self.detection_engine = None
  86. # Performance monitoring
  87. self._stats = {
  88. 'start_time': None,
  89. 'total_audio_processed_seconds': 0.0,
  90. 'total_chunks_processed': 0,
  91. 'total_detections': 0,
  92. 'detection_counts': {'custom': 0, 'system_command': 0, 'negative': 0},
  93. 'avg_processing_time_ms': 0.0,
  94. 'errors_count': 0,
  95. 'last_detection_time': None
  96. }
  97. # Event tracking
  98. self._last_wakeword_event = None
  99. self._event_callbacks: List[Callable] = []
  100. # Initialize components
  101. try:
  102. self._initialize_components()
  103. self._is_initialized = True
  104. self.logger.info("WakewordDetector initialized successfully")
  105. except Exception as e:
  106. self.logger.error(f"Failed to initialize WakewordDetector: {e}")
  107. raise WakewordDetectorError(f"Initialization failed: {e}")
  108. def _initialize_components(self):
  109. """Initialize all detector components."""
  110. # Set up device
  111. self.device = self._setup_device()
  112. self.logger.info(f"Using device: {self.device}")
  113. # Initialize model loader
  114. self.model_loader = SecureModelLoader(self.device)
  115. # Load model
  116. self._load_model()
  117. # Initialize feature extractor
  118. self.feature_extractor = create_feature_extractor(
  119. config=self.config.audio_config.spectrogram_config,
  120. device=self.device
  121. )
  122. # Initialize audio buffer
  123. self.audio_buffer = self._create_audio_buffer()
  124. # Initialize detection engine
  125. self.detection_engine = create_detection_engine(
  126. model=self.model,
  127. metadata=self.model_metadata,
  128. feature_extractor=self.feature_extractor,
  129. config=self.config.detection_config,
  130. device=self.device
  131. )
  132. # Warm up the model for consistent performance
  133. if self.config.performance_config.enable_performance_monitoring:
  134. self.detection_engine.warm_up()
  135. self.logger.info("All components initialized successfully")
  136. def _setup_device(self) -> torch.device:
  137. """Set up PyTorch device based on configuration."""
  138. device_config = self.config.model_config.device.lower()
  139. if device_config == "auto":
  140. # Auto-detect best available device
  141. if torch.cuda.is_available():
  142. device = torch.device("cuda")
  143. elif hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
  144. device = torch.device("mps")
  145. else:
  146. device = torch.device("cpu")
  147. else:
  148. device = torch.device(device_config)
  149. # Validate device availability
  150. if device.type == "cuda" and not torch.cuda.is_available():
  151. self.logger.warning("CUDA requested but not available, falling back to CPU")
  152. device = torch.device("cpu")
  153. elif device.type == "mps" and not (hasattr(torch.backends, "mps") and torch.backends.mps.is_available()):
  154. self.logger.warning("MPS requested but not available, falling back to CPU")
  155. device = torch.device("cpu")
  156. return device
  157. def _load_model(self):
  158. """Load the wakeword detection model."""
  159. model_path = self.config.model_config.model_path
  160. model_password = self.config.model_config.model_password
  161. if not model_path or not Path(model_path).exists():
  162. raise WakewordDetectorError(f"Model file not found: {model_path}")
  163. try:
  164. self.model, self.model_metadata = self.model_loader.load_model(
  165. model_path=model_path,
  166. password=model_password,
  167. verify_hash=self.config.model_config.verify_model_hash
  168. )
  169. # Validate model metadata for wakeword detection
  170. if not self._validate_model_metadata():
  171. raise WakewordDetectorError("Model metadata validation failed")
  172. self.logger.info(f"Loaded model: {self.model_metadata.model_name} v{self.model_metadata.model_version}")
  173. except ModelLoadError as e:
  174. # Try fallback model if configured
  175. if self.config.model_config.fallback_on_error and self.config.model_config.fallback_model_path:
  176. self.logger.warning(f"Primary model load failed: {e}, trying fallback")
  177. try:
  178. self.model, self.model_metadata = self.model_loader.load_model(
  179. model_path=self.config.model_config.fallback_model_path,
  180. password=model_password,
  181. verify_hash=False # Less strict for fallback
  182. )
  183. self.logger.info("Fallback model loaded successfully")
  184. except Exception as fallback_error:
  185. raise WakewordDetectorError(f"Both primary and fallback model loading failed: {e}, {fallback_error}")
  186. else:
  187. raise WakewordDetectorError(f"Model loading failed: {e}")
  188. def _validate_model_metadata(self) -> bool:
  189. """Validate that the loaded model is suitable for wakeword detection."""
  190. if not self.model_metadata:
  191. return False
  192. # Check model type
  193. if self.model_metadata.model_type != "wakeword":
  194. self.logger.error(f"Invalid model type: {self.model_metadata.model_type}, expected 'wakeword'")
  195. return False
  196. # Check class labels
  197. expected_labels = {"custom", "system_command", "negative"}
  198. model_labels = set(self.model_metadata.class_labels)
  199. if not expected_labels.issubset(model_labels):
  200. self.logger.error(f"Model missing required labels. Expected: {expected_labels}, Got: {model_labels}")
  201. return False
  202. # Check input shape compatibility
  203. audio_config = self.config.audio_config.spectrogram_config
  204. expected_shape = [1, audio_config.n_mels, audio_config.time_frames]
  205. if self.model_metadata.input_shape != expected_shape:
  206. self.logger.warning(f"Input shape mismatch. Model: {self.model_metadata.input_shape}, Config: {expected_shape}")
  207. # Update config to match model if reasonable
  208. if len(self.model_metadata.input_shape) == 3:
  209. audio_config.n_mels = self.model_metadata.input_shape[1]
  210. audio_config.time_frames = self.model_metadata.input_shape[2]
  211. self.logger.info("Updated audio config to match model input shape")
  212. return True
  213. def _create_audio_buffer(self) -> CircularAudioBuffer:
  214. """Create audio buffer based on configuration."""
  215. audio_config = self.config.audio_config
  216. return CircularAudioBuffer(
  217. buffer_duration=audio_config.buffer_duration,
  218. chunk_duration=audio_config.chunk_duration,
  219. sample_rate=audio_config.sample_rate,
  220. overlap_ratio=audio_config.overlap_ratio,
  221. chunk_callback=self._process_audio_chunk if self.config.performance_config.enable_threading else None
  222. )
  223. def start(self):
  224. """Start the wakeword detection system."""
  225. with self._lock:
  226. if self._is_running:
  227. self.logger.warning("WakewordDetector is already running")
  228. return
  229. if not self._is_initialized:
  230. raise WakewordDetectorError("WakewordDetector not properly initialized")
  231. self._is_running = True
  232. self._stats['start_time'] = time.time()
  233. # Register event handlers if event system is enabled
  234. if self.config.enable_event_system and self.event_handler:
  235. register_event_handlers(self)
  236. # Start processing threads if threading is enabled
  237. if self.config.performance_config.enable_threading:
  238. self._start_threads()
  239. self.logger.info("WakewordDetector started")
  240. def stop(self):
  241. """Stop the wakeword detection system."""
  242. with self._lock:
  243. if not self._is_running:
  244. return
  245. self._is_running = False
  246. # Stop threads
  247. self._stop_threads()
  248. # Unregister event handlers
  249. if self.config.enable_event_system and self.event_handler:
  250. unregister_event_handlers(self)
  251. # Clear audio buffer
  252. if self.audio_buffer:
  253. self.audio_buffer.clear()
  254. self.logger.info("WakewordDetector stopped")
  255. def _start_threads(self):
  256. """Start background processing threads."""
  257. if not self.config.performance_config.enable_threading:
  258. return
  259. # Audio processing thread is handled by buffer callback
  260. # Create additional processing thread for batched inference if needed
  261. if self.config.detection_config.batch_inference:
  262. self._processing_thread = threading.Thread(
  263. target=self._batch_processing_loop,
  264. name="WakewordBatchProcessor",
  265. daemon=True
  266. )
  267. self._processing_thread.start()
  268. def _stop_threads(self):
  269. """Stop background processing threads."""
  270. # Threads will stop automatically when _is_running becomes False
  271. if self._processing_thread and self._processing_thread.is_alive():
  272. self._processing_thread.join(timeout=2.0)
  273. def process_audio(self, audio_data: np.ndarray) -> List[DetectionResult]:
  274. """
  275. Process incoming audio data for wakeword detection.
  276. Args:
  277. audio_data: Raw audio samples (numpy array)
  278. Returns:
  279. List of detection results (empty if no processing occurred)
  280. """
  281. if not self._is_running:
  282. return []
  283. try:
  284. # Convert to float32 if needed
  285. if audio_data.dtype != np.float32:
  286. audio_data = audio_data.astype(np.float32)
  287. # Normalize audio range
  288. if audio_data.max() > 1.0 or audio_data.min() < -1.0:
  289. audio_data = audio_data / max(abs(audio_data.max()), abs(audio_data.min()))
  290. # Add to audio buffer
  291. self.audio_buffer.write_audio(audio_data)
  292. # Update statistics
  293. duration_seconds = len(audio_data) / self.config.audio_config.sample_rate
  294. self._stats['total_audio_processed_seconds'] += duration_seconds
  295. # Process immediately if threading is disabled
  296. if not self.config.performance_config.enable_threading:
  297. chunk = self.audio_buffer.get_latest_chunk()
  298. if chunk:
  299. return [self._process_audio_chunk(chunk)]
  300. return []
  301. except Exception as e:
  302. self.logger.error(f"Error processing audio: {e}")
  303. self._stats['errors_count'] += 1
  304. return []
  305. def _process_audio_chunk(self, audio_chunk: AudioChunk) -> Optional[DetectionResult]:
  306. """
  307. Process a single audio chunk for wakeword detection.
  308. Args:
  309. audio_chunk: Audio chunk to process
  310. Returns:
  311. Detection result or None
  312. """
  313. try:
  314. # Run detection
  315. result = self.detection_engine.detect_wakeword(audio_chunk)
  316. # Update statistics
  317. self._update_detection_stats(result)
  318. # Trigger event if wakeword detected
  319. if result.is_wakeword_detected:
  320. self.logger.info(f"[CLIENT] Wakeword detected: {result.wakeword_type.value} "
  321. f"(confidence: {result.confidence:.3f}, "
  322. f"processing: {result.processing_time_ms:.1f}ms)")
  323. self._trigger_wakeword_event(result, audio_chunk)
  324. # Log debug information
  325. if self.config.debug_mode:
  326. self.logger.debug(f"Detection result: {result.wakeword_type.value} "
  327. f"(confidence: {result.confidence:.3f}, "
  328. f"processing: {result.processing_time_ms:.1f}ms)")
  329. return result
  330. except Exception as e:
  331. self.logger.error(f"Error processing audio chunk: {e}")
  332. if self.config.debug_mode:
  333. self.logger.debug(traceback.format_exc())
  334. self._stats['errors_count'] += 1
  335. return None
  336. def _batch_processing_loop(self):
  337. """Background loop for batch processing (if enabled)."""
  338. # This is a placeholder for batch processing implementation
  339. # In practice, you would collect chunks and process them in batches
  340. while self._is_running:
  341. try:
  342. time.sleep(0.1) # Avoid busy waiting
  343. # Implement batch collection and processing here
  344. except Exception as e:
  345. self.logger.error(f"Error in batch processing loop: {e}")
  346. def _update_detection_stats(self, result: DetectionResult):
  347. """Update detection statistics."""
  348. self._stats['total_chunks_processed'] += 1
  349. self._stats['detection_counts'][result.wakeword_type.value] += 1
  350. if result.is_wakeword_detected:
  351. self._stats['total_detections'] += 1
  352. self._stats['last_detection_time'] = time.time()
  353. # Update average processing time
  354. total_time = self._stats['avg_processing_time_ms'] * (self._stats['total_chunks_processed'] - 1)
  355. self._stats['avg_processing_time_ms'] = (total_time + result.processing_time_ms) / self._stats['total_chunks_processed']
  356. def _trigger_wakeword_event(self, result: DetectionResult, audio_chunk: AudioChunk):
  357. """Trigger wakeword received event."""
  358. if not self.config.enable_event_system or not self.event_handler:
  359. return
  360. try:
  361. # Create event data
  362. event_data = WakewordReceivedEventData(
  363. wakeword_id=result.wakeword_type.value,
  364. wakeword_type=result.wakeword_type.value,
  365. confidence=result.confidence,
  366. raw_scores=result.raw_scores,
  367. processing_time_ms=result.processing_time_ms,
  368. audio_buffer_length=audio_chunk.duration,
  369. chunk_id=result.chunk_id,
  370. features_shape=list(result.features_shape),
  371. model_name=self.model_metadata.model_name if self.model_metadata else "",
  372. temporal_filtered=self.config.detection_config.use_temporal_smoothing,
  373. satellite_info=self._create_satellite_info(),
  374. volume=self._estimate_audio_volume(audio_chunk)
  375. )
  376. # Trigger event
  377. self.event_handler.trigger("wakeword_received", event_data)
  378. # Cache last event
  379. self._last_wakeword_event = event_data
  380. self.logger.info(f"Wakeword detected: {result.wakeword_type.value} "
  381. f"(confidence: {result.confidence:.3f})")
  382. except Exception as e:
  383. self.logger.error(f"Error triggering wakeword event: {e}")
  384. def _create_satellite_info(self) -> Optional[SatelliteInfo]:
  385. """Create satellite info for events."""
  386. if not self.config.satellite_id:
  387. return None
  388. return SatelliteInfo(
  389. satellite_id=self.config.satellite_id,
  390. mac_address="", # Would be filled by satellite system
  391. room_id="", # Would be filled by satellite system
  392. alias="", # Would be filled by satellite system
  393. version="1.0.0", # Trixy version
  394. capabilities=["wakeword_detection"]
  395. )
  396. def _estimate_audio_volume(self, audio_chunk: AudioChunk) -> float:
  397. """Estimate volume/amplitude of audio chunk."""
  398. try:
  399. rms = np.sqrt(np.mean(audio_chunk.data ** 2))
  400. # Convert to dB-like scale
  401. db = 20 * np.log10(max(rms, 1e-8))
  402. # Normalize to 0-1 range (assuming typical range -60dB to 0dB)
  403. volume = max(0.0, min(1.0, (db + 60) / 60))
  404. return volume
  405. except Exception:
  406. return 0.0
  407. def update_thresholds(self, custom_threshold: Optional[float] = None,
  408. system_command_threshold: Optional[float] = None):
  409. """
  410. Update confidence thresholds for wakeword detection.
  411. Args:
  412. custom_threshold: New threshold for custom wakeword
  413. system_command_threshold: New threshold for system command wakeword
  414. """
  415. if self.detection_engine:
  416. self.detection_engine.update_thresholds(
  417. custom_threshold=custom_threshold,
  418. system_command_threshold=system_command_threshold
  419. )
  420. # Update config as well
  421. if custom_threshold is not None:
  422. self.config.detection_config.custom_threshold = custom_threshold
  423. if system_command_threshold is not None:
  424. self.config.detection_config.system_command_threshold = system_command_threshold
  425. self.logger.info(f"Updated thresholds - custom: {self.config.detection_config.custom_threshold}, "
  426. f"system: {self.config.detection_config.system_command_threshold}")
  427. def get_performance_stats(self) -> Dict[str, Any]:
  428. """Get comprehensive performance statistics."""
  429. stats = self._stats.copy()
  430. # Add runtime information
  431. if stats['start_time']:
  432. stats['uptime_seconds'] = time.time() - stats['start_time']
  433. # Add component statistics
  434. if self.detection_engine:
  435. stats['detection_engine'] = self.detection_engine.get_performance_stats()
  436. if self.feature_extractor:
  437. stats['feature_extractor'] = self.feature_extractor.get_performance_stats()
  438. if self.audio_buffer:
  439. stats['audio_buffer'] = self.audio_buffer.get_stats()
  440. # Add configuration summary
  441. stats['config_summary'] = {
  442. 'model_path': self.config.model_config.model_path,
  443. 'device': str(self.device),
  444. 'sample_rate': self.config.audio_config.sample_rate,
  445. 'chunk_duration': self.config.audio_config.chunk_duration,
  446. 'custom_threshold': self.config.detection_config.custom_threshold,
  447. 'system_threshold': self.config.detection_config.system_command_threshold,
  448. }
  449. return stats
  450. def reset_stats(self):
  451. """Reset performance statistics."""
  452. self._stats = {
  453. 'start_time': time.time() if self._is_running else None,
  454. 'total_audio_processed_seconds': 0.0,
  455. 'total_chunks_processed': 0,
  456. 'total_detections': 0,
  457. 'detection_counts': {'custom': 0, 'system_command': 0, 'negative': 0},
  458. 'avg_processing_time_ms': 0.0,
  459. 'errors_count': 0,
  460. 'last_detection_time': None
  461. }
  462. if self.detection_engine:
  463. self.detection_engine.reset_stats()
  464. if self.feature_extractor:
  465. self.feature_extractor.reset_stats()
  466. def is_running(self) -> bool:
  467. """Check if detector is currently running."""
  468. return self._is_running
  469. def is_initialized(self) -> bool:
  470. """Check if detector is properly initialized."""
  471. return self._is_initialized
  472. def get_last_wakeword_event(self) -> Optional[WakewordReceivedEventData]:
  473. """Get the last wakeword event that was triggered."""
  474. return self._last_wakeword_event
  475. def add_detection_callback(self, callback: Callable[[DetectionResult], None]):
  476. """
  477. Add a callback function that will be called for each detection result.
  478. Args:
  479. callback: Function to call with detection results
  480. """
  481. self._event_callbacks.append(callback)
  482. def remove_detection_callback(self, callback: Callable[[DetectionResult], None]):
  483. """
  484. Remove a detection callback.
  485. Args:
  486. callback: Callback function to remove
  487. """
  488. if callback in self._event_callbacks:
  489. self._event_callbacks.remove(callback)
  490. def __enter__(self):
  491. """Context manager entry."""
  492. self.start()
  493. return self
  494. def __exit__(self, exc_type, exc_val, exc_tb):
  495. """Context manager exit."""
  496. self.stop()
  497. def __del__(self):
  498. """Cleanup on deletion."""
  499. try:
  500. self.stop()
  501. except Exception:
  502. pass # Ignore errors during cleanup
  503. def create_wakeword_detector(config_path: Optional[str] = None,
  504. model_path: Optional[str] = None,
  505. model_password: Optional[str] = None,
  506. event_handler: Optional[Any] = None,
  507. **kwargs) -> WakewordDetector:
  508. """
  509. Factory function to create a wakeword detector with simplified setup.
  510. Args:
  511. config_path: Path to configuration file (optional)
  512. model_path: Path to model file (overrides config)
  513. model_password: Model password (overrides config)
  514. event_handler: Event handler instance
  515. **kwargs: Additional configuration parameters
  516. Returns:
  517. Configured WakewordDetector instance
  518. """
  519. # Load configuration
  520. if config_path and Path(config_path).exists():
  521. config = WakewordConfig.load_from_file(config_path)
  522. else:
  523. config = create_default_config()
  524. # Override with provided parameters
  525. if model_path:
  526. config.model_config.model_path = model_path
  527. if model_password:
  528. config.model_config.model_password = model_password
  529. # Apply additional configuration parameters
  530. for key, value in kwargs.items():
  531. if hasattr(config, key):
  532. setattr(config, key, value)
  533. elif hasattr(config.detection_config, key):
  534. setattr(config.detection_config, key, value)
  535. elif hasattr(config.audio_config, key):
  536. setattr(config.audio_config, key, value)
  537. return WakewordDetector(config, event_handler)