| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661 |
- """
- Main wakeword detector class with full system integration.
- This module provides the complete WakewordDetector class that integrates
- all components of the wakeword detection system including audio processing,
- model inference, event handling, and performance monitoring.
- """
- import threading
- import time
- import traceback
- import warnings
- from typing import Optional, Callable, Dict, Any, List
- from pathlib import Path
- import logging
- import torch
- import numpy as np
- from .config import WakewordConfig, create_default_config
- from .model_loader import SecureModelLoader, ModelMetadata, ModelLoadError
- from .audio_features import AudioFeatureExtractor, create_feature_extractor
- from .audio_buffer import CircularAudioBuffer, AudioChunk, create_audio_buffer
- from .detection_engine import (
- WakewordDetectionEngine, DetectionResult, WakewordType,
- create_detection_engine
- )
- from ..events import WakewordEventData as WakewordReceivedEventData
- from ...events.event_data import SpeakerInfo, SatelliteInfo
- from ...events.decorators import register_event_handlers, unregister_event_handlers
- class WakewordDetectorError(Exception):
- """Exception raised by WakewordDetector."""
- pass
- class WakewordDetector:
- """
- Complete wakeword detection system for Trixy voice assistant.
-
- This class provides a high-level interface for real-time wakeword detection
- with automatic model loading, audio processing, and event integration.
-
- Features:
- - Real-time audio streaming and buffering
- - Dual wakeword detection (custom + system_command)
- - Password-protected model loading
- - Event system integration
- - Performance monitoring and optimization
- - Error handling and recovery
- - Configurable confidence thresholds
-
- Usage:
- # Basic usage
- config = WakewordConfig()
- config.model_config.model_path = "path/to/model.pth"
- config.model_config.model_password = "password"
-
- detector = WakewordDetector(config, event_handler)
- detector.start()
-
- # Process audio data
- detector.process_audio(audio_data)
-
- # Stop detection
- detector.stop()
- """
-
- def __init__(self, config: Optional[WakewordConfig] = None,
- event_handler: Optional[Any] = None,
- logger: Optional[logging.Logger] = None):
- """
- Initialize the wakeword detector.
-
- Args:
- config: Wakeword detection configuration
- event_handler: Event handler for triggering events
- logger: Logger instance for debugging
- """
- self.config = config or create_default_config()
- self.event_handler = event_handler
- self.logger = logger or logging.getLogger(__name__)
- # Validate configuration (skip overall validation due to AudioConfig issue, rely on ML Manager validation)
- # Note: ML Manager already validates sub-configs individually before creating the detector
- # if not self.config.validate():
- # raise WakewordDetectorError("Invalid configuration provided")
-
- # Initialize state
- self._is_running = False
- self._is_initialized = False
- self._lock = threading.RLock()
- self._audio_thread = None
- self._processing_thread = None
-
- # Components (initialized in _initialize_components)
- self.device = None
- self.model_loader = None
- self.model = None
- self.model_metadata = None
- self.feature_extractor = None
- self.audio_buffer = None
- self.detection_engine = None
-
- # Performance monitoring
- self._stats = {
- 'start_time': None,
- 'total_audio_processed_seconds': 0.0,
- 'total_chunks_processed': 0,
- 'total_detections': 0,
- 'detection_counts': {'custom': 0, 'system_command': 0, 'negative': 0},
- 'avg_processing_time_ms': 0.0,
- 'errors_count': 0,
- 'last_detection_time': None
- }
-
- # Event tracking
- self._last_wakeword_event = None
- self._event_callbacks: List[Callable] = []
-
- # Initialize components
- try:
- self._initialize_components()
- self._is_initialized = True
- self.logger.info("WakewordDetector initialized successfully")
- except Exception as e:
- self.logger.error(f"Failed to initialize WakewordDetector: {e}")
- raise WakewordDetectorError(f"Initialization failed: {e}")
-
- def _initialize_components(self):
- """Initialize all detector components."""
- # Set up device
- self.device = self._setup_device()
- self.logger.info(f"Using device: {self.device}")
-
- # Initialize model loader
- self.model_loader = SecureModelLoader(self.device)
-
- # Load model
- self._load_model()
-
- # Initialize feature extractor
- self.feature_extractor = create_feature_extractor(
- config=self.config.audio_config.spectrogram_config,
- device=self.device
- )
-
- # Initialize audio buffer
- self.audio_buffer = self._create_audio_buffer()
-
- # Initialize detection engine
- self.detection_engine = create_detection_engine(
- model=self.model,
- metadata=self.model_metadata,
- feature_extractor=self.feature_extractor,
- config=self.config.detection_config,
- device=self.device
- )
-
- # Warm up the model for consistent performance
- if self.config.performance_config.enable_performance_monitoring:
- self.detection_engine.warm_up()
-
- self.logger.info("All components initialized successfully")
-
- def _setup_device(self) -> torch.device:
- """Set up PyTorch device based on configuration."""
- device_config = self.config.model_config.device.lower()
-
- if device_config == "auto":
- # Auto-detect best available device
- if torch.cuda.is_available():
- device = torch.device("cuda")
- elif hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
- device = torch.device("mps")
- else:
- device = torch.device("cpu")
- else:
- device = torch.device(device_config)
-
- # Validate device availability
- if device.type == "cuda" and not torch.cuda.is_available():
- self.logger.warning("CUDA requested but not available, falling back to CPU")
- device = torch.device("cpu")
- elif device.type == "mps" and not (hasattr(torch.backends, "mps") and torch.backends.mps.is_available()):
- self.logger.warning("MPS requested but not available, falling back to CPU")
- device = torch.device("cpu")
-
- return device
-
- def _load_model(self):
- """Load the wakeword detection model."""
- model_path = self.config.model_config.model_path
- model_password = self.config.model_config.model_password
-
- if not model_path or not Path(model_path).exists():
- raise WakewordDetectorError(f"Model file not found: {model_path}")
-
- try:
- self.model, self.model_metadata = self.model_loader.load_model(
- model_path=model_path,
- password=model_password,
- verify_hash=self.config.model_config.verify_model_hash
- )
-
- # Validate model metadata for wakeword detection
- if not self._validate_model_metadata():
- raise WakewordDetectorError("Model metadata validation failed")
-
- self.logger.info(f"Loaded model: {self.model_metadata.model_name} v{self.model_metadata.model_version}")
-
- except ModelLoadError as e:
- # Try fallback model if configured
- if self.config.model_config.fallback_on_error and self.config.model_config.fallback_model_path:
- self.logger.warning(f"Primary model load failed: {e}, trying fallback")
- try:
- self.model, self.model_metadata = self.model_loader.load_model(
- model_path=self.config.model_config.fallback_model_path,
- password=model_password,
- verify_hash=False # Less strict for fallback
- )
- self.logger.info("Fallback model loaded successfully")
- except Exception as fallback_error:
- raise WakewordDetectorError(f"Both primary and fallback model loading failed: {e}, {fallback_error}")
- else:
- raise WakewordDetectorError(f"Model loading failed: {e}")
-
- def _validate_model_metadata(self) -> bool:
- """Validate that the loaded model is suitable for wakeword detection."""
- if not self.model_metadata:
- return False
-
- # Check model type
- if self.model_metadata.model_type != "wakeword":
- self.logger.error(f"Invalid model type: {self.model_metadata.model_type}, expected 'wakeword'")
- return False
-
- # Check class labels
- expected_labels = {"custom", "system_command", "negative"}
- model_labels = set(self.model_metadata.class_labels)
- if not expected_labels.issubset(model_labels):
- self.logger.error(f"Model missing required labels. Expected: {expected_labels}, Got: {model_labels}")
- return False
-
- # Check input shape compatibility
- audio_config = self.config.audio_config.spectrogram_config
- expected_shape = [1, audio_config.n_mels, audio_config.time_frames]
- if self.model_metadata.input_shape != expected_shape:
- self.logger.warning(f"Input shape mismatch. Model: {self.model_metadata.input_shape}, Config: {expected_shape}")
- # Update config to match model if reasonable
- if len(self.model_metadata.input_shape) == 3:
- audio_config.n_mels = self.model_metadata.input_shape[1]
- audio_config.time_frames = self.model_metadata.input_shape[2]
- self.logger.info("Updated audio config to match model input shape")
-
- return True
-
- def _create_audio_buffer(self) -> CircularAudioBuffer:
- """Create audio buffer based on configuration."""
- audio_config = self.config.audio_config
-
- return CircularAudioBuffer(
- buffer_duration=audio_config.buffer_duration,
- chunk_duration=audio_config.chunk_duration,
- sample_rate=audio_config.sample_rate,
- overlap_ratio=audio_config.overlap_ratio,
- chunk_callback=self._process_audio_chunk if self.config.performance_config.enable_threading else None
- )
-
- def start(self):
- """Start the wakeword detection system."""
- with self._lock:
- if self._is_running:
- self.logger.warning("WakewordDetector is already running")
- return
-
- if not self._is_initialized:
- raise WakewordDetectorError("WakewordDetector not properly initialized")
-
- self._is_running = True
- self._stats['start_time'] = time.time()
-
- # Register event handlers if event system is enabled
- if self.config.enable_event_system and self.event_handler:
- register_event_handlers(self)
-
- # Start processing threads if threading is enabled
- if self.config.performance_config.enable_threading:
- self._start_threads()
-
- self.logger.info("WakewordDetector started")
-
- def stop(self):
- """Stop the wakeword detection system."""
- with self._lock:
- if not self._is_running:
- return
-
- self._is_running = False
-
- # Stop threads
- self._stop_threads()
-
- # Unregister event handlers
- if self.config.enable_event_system and self.event_handler:
- unregister_event_handlers(self)
-
- # Clear audio buffer
- if self.audio_buffer:
- self.audio_buffer.clear()
-
- self.logger.info("WakewordDetector stopped")
-
- def _start_threads(self):
- """Start background processing threads."""
- if not self.config.performance_config.enable_threading:
- return
-
- # Audio processing thread is handled by buffer callback
- # Create additional processing thread for batched inference if needed
- if self.config.detection_config.batch_inference:
- self._processing_thread = threading.Thread(
- target=self._batch_processing_loop,
- name="WakewordBatchProcessor",
- daemon=True
- )
- self._processing_thread.start()
-
- def _stop_threads(self):
- """Stop background processing threads."""
- # Threads will stop automatically when _is_running becomes False
- if self._processing_thread and self._processing_thread.is_alive():
- self._processing_thread.join(timeout=2.0)
-
- def process_audio(self, audio_data: np.ndarray) -> List[DetectionResult]:
- """
- Process incoming audio data for wakeword detection.
-
- Args:
- audio_data: Raw audio samples (numpy array)
-
- Returns:
- List of detection results (empty if no processing occurred)
- """
- if not self._is_running:
- return []
-
- try:
- # Convert to float32 if needed
- if audio_data.dtype != np.float32:
- audio_data = audio_data.astype(np.float32)
-
- # Normalize audio range
- if audio_data.max() > 1.0 or audio_data.min() < -1.0:
- audio_data = audio_data / max(abs(audio_data.max()), abs(audio_data.min()))
-
- # Add to audio buffer
- self.audio_buffer.write_audio(audio_data)
-
- # Update statistics
- duration_seconds = len(audio_data) / self.config.audio_config.sample_rate
- self._stats['total_audio_processed_seconds'] += duration_seconds
-
- # Process immediately if threading is disabled
- if not self.config.performance_config.enable_threading:
- chunk = self.audio_buffer.get_latest_chunk()
- if chunk:
- return [self._process_audio_chunk(chunk)]
-
- return []
-
- except Exception as e:
- self.logger.error(f"Error processing audio: {e}")
- self._stats['errors_count'] += 1
- return []
-
- def _process_audio_chunk(self, audio_chunk: AudioChunk) -> Optional[DetectionResult]:
- """
- Process a single audio chunk for wakeword detection.
-
- Args:
- audio_chunk: Audio chunk to process
-
- Returns:
- Detection result or None
- """
- try:
- # Run detection
- result = self.detection_engine.detect_wakeword(audio_chunk)
-
- # Update statistics
- self._update_detection_stats(result)
-
- # Trigger event if wakeword detected
- if result.is_wakeword_detected:
- self.logger.info(f"[CLIENT] Wakeword detected: {result.wakeword_type.value} "
- f"(confidence: {result.confidence:.3f}, "
- f"processing: {result.processing_time_ms:.1f}ms)")
- self._trigger_wakeword_event(result, audio_chunk)
- # Log debug information
- if self.config.debug_mode:
- self.logger.debug(f"Detection result: {result.wakeword_type.value} "
- f"(confidence: {result.confidence:.3f}, "
- f"processing: {result.processing_time_ms:.1f}ms)")
-
- return result
-
- except Exception as e:
- self.logger.error(f"Error processing audio chunk: {e}")
- if self.config.debug_mode:
- self.logger.debug(traceback.format_exc())
- self._stats['errors_count'] += 1
- return None
-
- def _batch_processing_loop(self):
- """Background loop for batch processing (if enabled)."""
- # This is a placeholder for batch processing implementation
- # In practice, you would collect chunks and process them in batches
- while self._is_running:
- try:
- time.sleep(0.1) # Avoid busy waiting
- # Implement batch collection and processing here
- except Exception as e:
- self.logger.error(f"Error in batch processing loop: {e}")
-
- def _update_detection_stats(self, result: DetectionResult):
- """Update detection statistics."""
- self._stats['total_chunks_processed'] += 1
- self._stats['detection_counts'][result.wakeword_type.value] += 1
-
- if result.is_wakeword_detected:
- self._stats['total_detections'] += 1
- self._stats['last_detection_time'] = time.time()
-
- # Update average processing time
- total_time = self._stats['avg_processing_time_ms'] * (self._stats['total_chunks_processed'] - 1)
- self._stats['avg_processing_time_ms'] = (total_time + result.processing_time_ms) / self._stats['total_chunks_processed']
-
- def _trigger_wakeword_event(self, result: DetectionResult, audio_chunk: AudioChunk):
- """Trigger wakeword received event."""
- if not self.config.enable_event_system or not self.event_handler:
- return
-
- try:
- # Create event data
- event_data = WakewordReceivedEventData(
- wakeword_id=result.wakeword_type.value,
- wakeword_type=result.wakeword_type.value,
- confidence=result.confidence,
- raw_scores=result.raw_scores,
- processing_time_ms=result.processing_time_ms,
- audio_buffer_length=audio_chunk.duration,
- chunk_id=result.chunk_id,
- features_shape=list(result.features_shape),
- model_name=self.model_metadata.model_name if self.model_metadata else "",
- temporal_filtered=self.config.detection_config.use_temporal_smoothing,
- satellite_info=self._create_satellite_info(),
- volume=self._estimate_audio_volume(audio_chunk)
- )
-
- # Trigger event
- self.event_handler.trigger("wakeword_received", event_data)
-
- # Cache last event
- self._last_wakeword_event = event_data
-
- self.logger.info(f"Wakeword detected: {result.wakeword_type.value} "
- f"(confidence: {result.confidence:.3f})")
-
- except Exception as e:
- self.logger.error(f"Error triggering wakeword event: {e}")
-
- def _create_satellite_info(self) -> Optional[SatelliteInfo]:
- """Create satellite info for events."""
- if not self.config.satellite_id:
- return None
-
- return SatelliteInfo(
- satellite_id=self.config.satellite_id,
- mac_address="", # Would be filled by satellite system
- room_id="", # Would be filled by satellite system
- alias="", # Would be filled by satellite system
- version="1.0.0", # Trixy version
- capabilities=["wakeword_detection"]
- )
-
- def _estimate_audio_volume(self, audio_chunk: AudioChunk) -> float:
- """Estimate volume/amplitude of audio chunk."""
- try:
- rms = np.sqrt(np.mean(audio_chunk.data ** 2))
- # Convert to dB-like scale
- db = 20 * np.log10(max(rms, 1e-8))
- # Normalize to 0-1 range (assuming typical range -60dB to 0dB)
- volume = max(0.0, min(1.0, (db + 60) / 60))
- return volume
- except Exception:
- return 0.0
-
- def update_thresholds(self, custom_threshold: Optional[float] = None,
- system_command_threshold: Optional[float] = None):
- """
- Update confidence thresholds for wakeword detection.
-
- Args:
- custom_threshold: New threshold for custom wakeword
- system_command_threshold: New threshold for system command wakeword
- """
- if self.detection_engine:
- self.detection_engine.update_thresholds(
- custom_threshold=custom_threshold,
- system_command_threshold=system_command_threshold
- )
-
- # Update config as well
- if custom_threshold is not None:
- self.config.detection_config.custom_threshold = custom_threshold
- if system_command_threshold is not None:
- self.config.detection_config.system_command_threshold = system_command_threshold
-
- self.logger.info(f"Updated thresholds - custom: {self.config.detection_config.custom_threshold}, "
- f"system: {self.config.detection_config.system_command_threshold}")
-
- def get_performance_stats(self) -> Dict[str, Any]:
- """Get comprehensive performance statistics."""
- stats = self._stats.copy()
-
- # Add runtime information
- if stats['start_time']:
- stats['uptime_seconds'] = time.time() - stats['start_time']
-
- # Add component statistics
- if self.detection_engine:
- stats['detection_engine'] = self.detection_engine.get_performance_stats()
-
- if self.feature_extractor:
- stats['feature_extractor'] = self.feature_extractor.get_performance_stats()
-
- if self.audio_buffer:
- stats['audio_buffer'] = self.audio_buffer.get_stats()
-
- # Add configuration summary
- stats['config_summary'] = {
- 'model_path': self.config.model_config.model_path,
- 'device': str(self.device),
- 'sample_rate': self.config.audio_config.sample_rate,
- 'chunk_duration': self.config.audio_config.chunk_duration,
- 'custom_threshold': self.config.detection_config.custom_threshold,
- 'system_threshold': self.config.detection_config.system_command_threshold,
- }
-
- return stats
-
- def reset_stats(self):
- """Reset performance statistics."""
- self._stats = {
- 'start_time': time.time() if self._is_running else None,
- 'total_audio_processed_seconds': 0.0,
- 'total_chunks_processed': 0,
- 'total_detections': 0,
- 'detection_counts': {'custom': 0, 'system_command': 0, 'negative': 0},
- 'avg_processing_time_ms': 0.0,
- 'errors_count': 0,
- 'last_detection_time': None
- }
-
- if self.detection_engine:
- self.detection_engine.reset_stats()
-
- if self.feature_extractor:
- self.feature_extractor.reset_stats()
-
- def is_running(self) -> bool:
- """Check if detector is currently running."""
- return self._is_running
-
- def is_initialized(self) -> bool:
- """Check if detector is properly initialized."""
- return self._is_initialized
-
- def get_last_wakeword_event(self) -> Optional[WakewordReceivedEventData]:
- """Get the last wakeword event that was triggered."""
- return self._last_wakeword_event
-
- def add_detection_callback(self, callback: Callable[[DetectionResult], None]):
- """
- Add a callback function that will be called for each detection result.
-
- Args:
- callback: Function to call with detection results
- """
- self._event_callbacks.append(callback)
-
- def remove_detection_callback(self, callback: Callable[[DetectionResult], None]):
- """
- Remove a detection callback.
-
- Args:
- callback: Callback function to remove
- """
- if callback in self._event_callbacks:
- self._event_callbacks.remove(callback)
-
- def __enter__(self):
- """Context manager entry."""
- self.start()
- return self
-
- def __exit__(self, exc_type, exc_val, exc_tb):
- """Context manager exit."""
- self.stop()
-
- def __del__(self):
- """Cleanup on deletion."""
- try:
- self.stop()
- except Exception:
- pass # Ignore errors during cleanup
- def create_wakeword_detector(config_path: Optional[str] = None,
- model_path: Optional[str] = None,
- model_password: Optional[str] = None,
- event_handler: Optional[Any] = None,
- **kwargs) -> WakewordDetector:
- """
- Factory function to create a wakeword detector with simplified setup.
-
- Args:
- config_path: Path to configuration file (optional)
- model_path: Path to model file (overrides config)
- model_password: Model password (overrides config)
- event_handler: Event handler instance
- **kwargs: Additional configuration parameters
-
- Returns:
- Configured WakewordDetector instance
- """
- # Load configuration
- if config_path and Path(config_path).exists():
- config = WakewordConfig.load_from_file(config_path)
- else:
- config = create_default_config()
-
- # Override with provided parameters
- if model_path:
- config.model_config.model_path = model_path
- if model_password:
- config.model_config.model_password = model_password
-
- # Apply additional configuration parameters
- for key, value in kwargs.items():
- if hasattr(config, key):
- setattr(config, key, value)
- elif hasattr(config.detection_config, key):
- setattr(config.detection_config, key, value)
- elif hasattr(config.audio_config, key):
- setattr(config.audio_config, key, value)
-
- return WakewordDetector(config, event_handler)
|