config.py 17 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462
  1. """
  2. Configuration management for wakeword detection system.
  3. This module provides comprehensive configuration management for all aspects
  4. of the wakeword detection pipeline including model loading, audio processing,
  5. detection parameters, and performance tuning.
  6. """
  7. import os
  8. import json
  9. from pathlib import Path
  10. from typing import Optional, Dict, Any, List, Union
  11. from dataclasses import dataclass, asdict, field
  12. from .audio_features import SpectrogramConfig
  13. from .detection_engine import DetectionConfig
  14. @dataclass
  15. class ModelConfig:
  16. """Configuration for model loading and management."""
  17. # Model file paths
  18. model_path: str = ""
  19. model_password: str = ""
  20. backup_model_paths: List[str] = field(default_factory=list)
  21. # Model selection
  22. model_type: str = "auto" # "auto", "standard", "improved", "lightweight"
  23. use_model_optimization: bool = True
  24. verify_model_hash: bool = True
  25. # Device configuration
  26. device: str = "auto" # "auto", "cpu", "cuda", "mps"
  27. use_half_precision: bool = False
  28. enable_model_compilation: bool = True
  29. # Fallback behavior
  30. fallback_on_error: bool = True
  31. fallback_model_path: str = ""
  32. def validate(self) -> bool:
  33. """Validate model configuration."""
  34. if not self.model_path:
  35. return False
  36. # Check if model file exists
  37. if not Path(self.model_path).exists():
  38. return False
  39. # Validate device setting
  40. valid_devices = ["auto", "cpu", "cuda", "mps"]
  41. if self.device not in valid_devices:
  42. return False
  43. return True
  44. @dataclass
  45. class AudioConfig:
  46. """Configuration for audio processing."""
  47. # Audio input parameters
  48. sample_rate: int = 16000
  49. channels: int = 1 # Mono audio
  50. bit_depth: int = 16
  51. # Buffer configuration
  52. buffer_duration: float = 2.0 # seconds
  53. chunk_duration: float = 1.0 # seconds for processing
  54. overlap_ratio: float = 0.5
  55. # Feature extraction
  56. spectrogram_config: SpectrogramConfig = field(default_factory=SpectrogramConfig)
  57. # Audio preprocessing
  58. enable_noise_reduction: bool = False
  59. enable_gain_control: bool = True
  60. enable_voice_activity_detection: bool = False
  61. # Streaming parameters
  62. max_chunks_in_memory: int = 10
  63. use_circular_buffer: bool = True
  64. def validate(self) -> bool:
  65. """Validate audio configuration."""
  66. if self.sample_rate <= 0:
  67. return False
  68. if self.channels not in [1, 2]:
  69. return False
  70. if self.bit_depth not in [16, 24, 32]:
  71. return False
  72. if self.buffer_duration <= 0 or self.chunk_duration <= 0:
  73. return False
  74. if not (0.0 <= self.overlap_ratio <= 1.0):
  75. return False
  76. return self.spectrogram_config and hasattr(self.spectrogram_config, 'validate')
  77. @dataclass
  78. class PerformanceConfig:
  79. """Configuration for performance optimization."""
  80. # Threading and concurrency
  81. enable_threading: bool = True
  82. max_worker_threads: int = 2
  83. thread_priority: str = "normal" # "low", "normal", "high"
  84. # Memory management
  85. max_memory_usage_mb: float = 512.0
  86. enable_memory_pool: bool = True
  87. gc_frequency: int = 100 # Garbage collection every N chunks
  88. # Inference optimization
  89. enable_batch_inference: bool = False
  90. max_batch_size: int = 8
  91. inference_timeout_ms: float = 100.0
  92. # Monitoring and logging
  93. enable_performance_monitoring: bool = True
  94. log_performance_stats: bool = False
  95. stats_update_interval: int = 10 # seconds
  96. def validate(self) -> bool:
  97. """Validate performance configuration."""
  98. if self.max_worker_threads < 1:
  99. return False
  100. if self.max_memory_usage_mb <= 0:
  101. return False
  102. if self.max_batch_size < 1:
  103. return False
  104. if self.inference_timeout_ms <= 0:
  105. return False
  106. return True
  107. @dataclass
  108. class WakewordConfig:
  109. """
  110. Complete configuration for wakeword detection system.
  111. This is the main configuration class that encompasses all aspects
  112. of the wakeword detection pipeline.
  113. """
  114. # Sub-configurations
  115. model_config: ModelConfig = field(default_factory=ModelConfig)
  116. audio_config: AudioConfig = field(default_factory=AudioConfig)
  117. detection_config: DetectionConfig = field(default_factory=DetectionConfig)
  118. performance_config: PerformanceConfig = field(default_factory=PerformanceConfig)
  119. # System integration
  120. enable_event_system: bool = True
  121. event_handler: Optional[Any] = None
  122. satellite_id: str = ""
  123. # Debugging and development
  124. debug_mode: bool = False
  125. save_debug_audio: bool = False
  126. debug_output_path: str = "/tmp/trixy_wakeword_debug"
  127. # Feature flags
  128. enable_voice_recognition_integration: bool = False
  129. enable_conversation_context: bool = True
  130. enable_adaptive_thresholds: bool = False
  131. def validate(self) -> bool:
  132. """Validate complete configuration."""
  133. if not self.model_config.validate():
  134. return False
  135. if not self.audio_config.validate():
  136. return False
  137. if not self.detection_config.validate():
  138. return False
  139. if not self.performance_config.validate():
  140. return False
  141. return True
  142. def to_dict(self) -> Dict[str, Any]:
  143. """Convert configuration to dictionary."""
  144. return {
  145. 'model_config': asdict(self.model_config),
  146. 'audio_config': asdict(self.audio_config),
  147. 'detection_config': asdict(self.detection_config),
  148. 'performance_config': asdict(self.performance_config),
  149. 'enable_event_system': self.enable_event_system,
  150. 'satellite_id': self.satellite_id,
  151. 'debug_mode': self.debug_mode,
  152. 'save_debug_audio': self.save_debug_audio,
  153. 'debug_output_path': self.debug_output_path,
  154. 'enable_voice_recognition_integration': self.enable_voice_recognition_integration,
  155. 'enable_conversation_context': self.enable_conversation_context,
  156. 'enable_adaptive_thresholds': self.enable_adaptive_thresholds
  157. }
  158. @classmethod
  159. def from_dict(cls, config_dict: Dict[str, Any]) -> 'WakewordConfig':
  160. """Create configuration from dictionary."""
  161. # Extract sub-configurations
  162. model_config = ModelConfig(**config_dict.get('model_config', {}))
  163. audio_config = AudioConfig(**config_dict.get('audio_config', {}))
  164. detection_config = DetectionConfig(**config_dict.get('detection_config', {}))
  165. performance_config = PerformanceConfig(**config_dict.get('performance_config', {}))
  166. # Create main config
  167. return cls(
  168. model_config=model_config,
  169. audio_config=audio_config,
  170. detection_config=detection_config,
  171. performance_config=performance_config,
  172. enable_event_system=config_dict.get('enable_event_system', True),
  173. satellite_id=config_dict.get('satellite_id', ''),
  174. debug_mode=config_dict.get('debug_mode', False),
  175. save_debug_audio=config_dict.get('save_debug_audio', False),
  176. debug_output_path=config_dict.get('debug_output_path', '/tmp/trixy_wakeword_debug'),
  177. enable_voice_recognition_integration=config_dict.get('enable_voice_recognition_integration', False),
  178. enable_conversation_context=config_dict.get('enable_conversation_context', True),
  179. enable_adaptive_thresholds=config_dict.get('enable_adaptive_thresholds', False)
  180. )
  181. def save_to_file(self, file_path: Union[str, Path]):
  182. """Save configuration to JSON file."""
  183. file_path = Path(file_path)
  184. file_path.parent.mkdir(parents=True, exist_ok=True)
  185. with open(file_path, 'w') as f:
  186. json.dump(self.to_dict(), f, indent=2)
  187. @classmethod
  188. def load_from_file(cls, file_path: Union[str, Path]) -> 'WakewordConfig':
  189. """Load configuration from JSON file."""
  190. file_path = Path(file_path)
  191. if not file_path.exists():
  192. raise FileNotFoundError(f"Configuration file not found: {file_path}")
  193. with open(file_path, 'r') as f:
  194. config_dict = json.load(f)
  195. return cls.from_dict(config_dict)
  196. def update_from_env(self, prefix: str = "TRIXY_WAKEWORD_"):
  197. """Update configuration from environment variables."""
  198. env_mappings = {
  199. f"{prefix}MODEL_PATH": ("model_config", "model_path"),
  200. f"{prefix}MODEL_PASSWORD": ("model_config", "model_password"),
  201. f"{prefix}DEVICE": ("model_config", "device"),
  202. f"{prefix}SAMPLE_RATE": ("audio_config", "sample_rate"),
  203. f"{prefix}CUSTOM_THRESHOLD": ("detection_config", "custom_threshold"),
  204. f"{prefix}SYSTEM_THRESHOLD": ("detection_config", "system_command_threshold"),
  205. f"{prefix}DEBUG_MODE": ("debug_mode", None),
  206. f"{prefix}SATELLITE_ID": ("satellite_id", None),
  207. }
  208. for env_var, (section, key) in env_mappings.items():
  209. value = os.getenv(env_var)
  210. if value is not None:
  211. if section == "debug_mode" or section == "satellite_id":
  212. # Direct attribute
  213. if section == "debug_mode":
  214. setattr(self, section, value.lower() in ['true', '1', 'yes'])
  215. else:
  216. setattr(self, section, value)
  217. else:
  218. # Sub-configuration attribute
  219. config_obj = getattr(self, section)
  220. if hasattr(config_obj, key):
  221. current_value = getattr(config_obj, key)
  222. # Type conversion based on current value type
  223. if isinstance(current_value, bool):
  224. setattr(config_obj, key, value.lower() in ['true', '1', 'yes'])
  225. elif isinstance(current_value, int):
  226. setattr(config_obj, key, int(value))
  227. elif isinstance(current_value, float):
  228. setattr(config_obj, key, float(value))
  229. else:
  230. setattr(config_obj, key, value)
  231. def create_default_config() -> WakewordConfig:
  232. """Create a default wakeword configuration."""
  233. return WakewordConfig(
  234. model_config=ModelConfig(
  235. model_path="models/wakeword/default/model.pth",
  236. model_password="",
  237. device="auto",
  238. use_model_optimization=True
  239. ),
  240. audio_config=AudioConfig(
  241. sample_rate=16000,
  242. buffer_duration=2.0,
  243. chunk_duration=1.0,
  244. overlap_ratio=0.5
  245. ),
  246. detection_config=DetectionConfig(
  247. custom_threshold=0.7,
  248. system_command_threshold=0.8,
  249. use_temporal_smoothing=True,
  250. smoothing_window_size=5
  251. ),
  252. performance_config=PerformanceConfig(
  253. enable_threading=True,
  254. max_worker_threads=2,
  255. enable_performance_monitoring=True
  256. )
  257. )
  258. def create_lightweight_config() -> WakewordConfig:
  259. """Create a lightweight configuration for resource-constrained devices."""
  260. config = create_default_config()
  261. # Optimize for low resource usage
  262. config.model_config.use_half_precision = True
  263. config.model_config.enable_model_compilation = False
  264. config.audio_config.buffer_duration = 1.5
  265. config.audio_config.max_chunks_in_memory = 5
  266. config.detection_config.use_temporal_smoothing = False
  267. config.detection_config.batch_inference = False
  268. config.performance_config.max_worker_threads = 1
  269. config.performance_config.max_memory_usage_mb = 128.0
  270. config.performance_config.enable_memory_pool = True
  271. return config
  272. def create_high_performance_config() -> WakewordConfig:
  273. """Create a high-performance configuration for powerful devices."""
  274. config = create_default_config()
  275. # Optimize for high performance
  276. config.model_config.use_model_optimization = True
  277. config.model_config.enable_model_compilation = True
  278. config.audio_config.buffer_duration = 3.0
  279. config.audio_config.max_chunks_in_memory = 20
  280. config.detection_config.use_temporal_smoothing = True
  281. config.detection_config.batch_inference = True
  282. config.detection_config.enable_model_optimization = True
  283. config.performance_config.max_worker_threads = 4
  284. config.performance_config.max_memory_usage_mb = 1024.0
  285. config.performance_config.enable_batch_inference = True
  286. config.performance_config.max_batch_size = 16
  287. return config
  288. def validate_config_compatibility(config: WakewordConfig) -> List[str]:
  289. """
  290. Validate configuration for compatibility and potential issues.
  291. Args:
  292. config: Configuration to validate
  293. Returns:
  294. List of warning/error messages
  295. """
  296. warnings = []
  297. # Check model and audio compatibility
  298. if config.model_config.use_half_precision and config.model_config.device == "cpu":
  299. warnings.append("Half precision not supported on CPU, will be disabled")
  300. # Check buffer and chunk size compatibility
  301. if config.audio_config.chunk_duration >= config.audio_config.buffer_duration:
  302. warnings.append("Chunk duration should be smaller than buffer duration")
  303. # Check memory constraints
  304. estimated_memory = estimate_memory_usage(config)
  305. if estimated_memory > config.performance_config.max_memory_usage_mb:
  306. warnings.append(f"Estimated memory usage ({estimated_memory:.1f}MB) exceeds limit")
  307. # Check performance settings
  308. if config.detection_config.batch_inference and not config.performance_config.enable_batch_inference:
  309. warnings.append("Batch inference enabled in detection but disabled in performance config")
  310. # Check temporal smoothing compatibility
  311. if config.detection_config.use_temporal_smoothing and config.performance_config.max_worker_threads < 2:
  312. warnings.append("Temporal smoothing may need multiple threads for optimal performance")
  313. return warnings
  314. def estimate_memory_usage(config: WakewordConfig) -> float:
  315. """
  316. Estimate memory usage in MB for the given configuration.
  317. Args:
  318. config: Configuration to analyze
  319. Returns:
  320. Estimated memory usage in MB
  321. """
  322. # Audio buffer memory
  323. samples_per_second = config.audio_config.sample_rate
  324. buffer_samples = int(config.audio_config.buffer_duration * samples_per_second)
  325. chunk_samples = int(config.audio_config.chunk_duration * samples_per_second)
  326. # Assuming 32-bit float samples
  327. buffer_memory = buffer_samples * 4 / (1024 * 1024) # MB
  328. chunk_memory = chunk_samples * config.audio_config.max_chunks_in_memory * 4 / (1024 * 1024)
  329. # Feature extraction memory (spectrograms)
  330. spec_config = config.audio_config.spectrogram_config
  331. spec_memory = spec_config.n_mels * spec_config.time_frames * 4 / (1024 * 1024)
  332. # Model memory (rough estimate)
  333. model_memory = 50.0 # MB - rough estimate for RepCNN models
  334. # Additional overhead
  335. overhead = 20.0 # MB
  336. total_memory = buffer_memory + chunk_memory + spec_memory + model_memory + overhead
  337. return total_memory
  338. def auto_tune_config(config: WakewordConfig, target_latency_ms: float = 100.0,
  339. available_memory_mb: float = 512.0) -> WakewordConfig:
  340. """
  341. Automatically tune configuration parameters for target performance.
  342. Args:
  343. config: Base configuration to tune
  344. target_latency_ms: Target inference latency in milliseconds
  345. available_memory_mb: Available memory in MB
  346. Returns:
  347. Tuned configuration
  348. """
  349. tuned_config = WakewordConfig.from_dict(config.to_dict())
  350. # Adjust for latency target
  351. if target_latency_ms < 50.0:
  352. # Very low latency requirements
  353. tuned_config.audio_config.chunk_duration = 0.5
  354. tuned_config.detection_config.use_temporal_smoothing = False
  355. tuned_config.model_config.use_half_precision = True
  356. tuned_config.performance_config.max_worker_threads = 1
  357. elif target_latency_ms > 200.0:
  358. # Can afford higher latency for better accuracy
  359. tuned_config.audio_config.chunk_duration = 1.5
  360. tuned_config.detection_config.use_temporal_smoothing = True
  361. tuned_config.detection_config.smoothing_window_size = 7
  362. # Adjust for memory constraints
  363. estimated_memory = estimate_memory_usage(tuned_config)
  364. if estimated_memory > available_memory_mb:
  365. # Reduce memory usage
  366. tuned_config.audio_config.buffer_duration = max(1.0, tuned_config.audio_config.buffer_duration * 0.7)
  367. tuned_config.audio_config.max_chunks_in_memory = max(3, tuned_config.audio_config.max_chunks_in_memory // 2)
  368. tuned_config.performance_config.max_memory_usage_mb = available_memory_mb * 0.8
  369. return tuned_config