server_integration.py 18 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552
  1. """
  2. Server integration module for the ML Trainer Framework.
  3. This module provides integration capabilities for server script inclusion
  4. and programmatic access to the training framework from other components
  5. of the Trixy voice assistant system.
  6. """
  7. import os
  8. import json
  9. import logging
  10. from pathlib import Path
  11. from typing import Dict, Any, Optional, List, Callable, Union
  12. import threading
  13. import queue
  14. from dataclasses import dataclass
  15. from enum import Enum
  16. from .base import TrainerConfig
  17. from .wakeword.trainer import WakewordTrainer, create_wakeword_trainer_config
  18. from .voice_recognition.trainer import VoiceRecognitionTrainer, create_voice_recognition_trainer_config
  19. from .utils import setup_logging
  20. class TrainingStatus(Enum):
  21. """Training status enumeration."""
  22. PENDING = "pending"
  23. RUNNING = "running"
  24. COMPLETED = "completed"
  25. FAILED = "failed"
  26. CANCELLED = "cancelled"
  27. @dataclass
  28. class TrainingJob:
  29. """Training job definition."""
  30. job_id: str
  31. trainer_type: str # 'wakeword' or 'voice_recognition'
  32. config: TrainerConfig
  33. status: TrainingStatus = TrainingStatus.PENDING
  34. progress: float = 0.0
  35. current_epoch: int = 0
  36. total_epochs: int = 0
  37. error_message: Optional[str] = None
  38. result: Optional[Dict[str, Any]] = None
  39. log_messages: List[str] = None
  40. def __post_init__(self):
  41. if self.log_messages is None:
  42. self.log_messages = []
  43. class TrainingJobManager:
  44. """
  45. Manages training jobs with support for background execution,
  46. progress tracking, and result retrieval.
  47. """
  48. def __init__(self, max_concurrent_jobs: int = 1):
  49. """
  50. Initialize training job manager.
  51. Args:
  52. max_concurrent_jobs: Maximum number of concurrent training jobs
  53. """
  54. self.max_concurrent_jobs = max_concurrent_jobs
  55. self.jobs: Dict[str, TrainingJob] = {}
  56. self.job_queue = queue.Queue()
  57. self.active_jobs = set()
  58. self.lock = threading.Lock()
  59. self.worker_threads = []
  60. self.shutdown_flag = threading.Event()
  61. # Setup logging
  62. self.logger = setup_logging("training_job_manager", "INFO")
  63. # Start worker threads
  64. for i in range(max_concurrent_jobs):
  65. worker = threading.Thread(target=self._worker_loop, daemon=True)
  66. worker.start()
  67. self.worker_threads.append(worker)
  68. def create_wakeword_job(self, job_id: str, config_dict: Dict[str, Any]) -> str:
  69. """
  70. Create a wakeword training job.
  71. Args:
  72. job_id: Unique job identifier
  73. config_dict: Training configuration dictionary
  74. Returns:
  75. Job ID
  76. """
  77. try:
  78. config = create_wakeword_trainer_config(**config_dict)
  79. job = TrainingJob(
  80. job_id=job_id,
  81. trainer_type='wakeword',
  82. config=config,
  83. total_epochs=config.num_epochs
  84. )
  85. with self.lock:
  86. self.jobs[job_id] = job
  87. self.job_queue.put(job_id)
  88. self.logger.info(f"Created wakeword training job: {job_id}")
  89. return job_id
  90. except Exception as e:
  91. self.logger.error(f"Failed to create wakeword job {job_id}: {str(e)}")
  92. raise
  93. def create_voice_recognition_job(self, job_id: str, config_dict: Dict[str, Any]) -> str:
  94. """
  95. Create a voice recognition training job.
  96. Args:
  97. job_id: Unique job identifier
  98. config_dict: Training configuration dictionary
  99. Returns:
  100. Job ID
  101. """
  102. try:
  103. config = create_voice_recognition_trainer_config(**config_dict)
  104. job = TrainingJob(
  105. job_id=job_id,
  106. trainer_type='voice_recognition',
  107. config=config,
  108. total_epochs=config.num_epochs
  109. )
  110. with self.lock:
  111. self.jobs[job_id] = job
  112. self.job_queue.put(job_id)
  113. self.logger.info(f"Created voice recognition training job: {job_id}")
  114. return job_id
  115. except Exception as e:
  116. self.logger.error(f"Failed to create voice recognition job {job_id}: {str(e)}")
  117. raise
  118. def get_job_status(self, job_id: str) -> Optional[Dict[str, Any]]:
  119. """
  120. Get training job status.
  121. Args:
  122. job_id: Job identifier
  123. Returns:
  124. Job status dictionary or None if job not found
  125. """
  126. with self.lock:
  127. job = self.jobs.get(job_id)
  128. if job is None:
  129. return None
  130. return {
  131. 'job_id': job.job_id,
  132. 'trainer_type': job.trainer_type,
  133. 'status': job.status.value,
  134. 'progress': job.progress,
  135. 'current_epoch': job.current_epoch,
  136. 'total_epochs': job.total_epochs,
  137. 'error_message': job.error_message,
  138. 'has_result': job.result is not None
  139. }
  140. def get_job_result(self, job_id: str) -> Optional[Dict[str, Any]]:
  141. """
  142. Get training job result.
  143. Args:
  144. job_id: Job identifier
  145. Returns:
  146. Job result or None
  147. """
  148. with self.lock:
  149. job = self.jobs.get(job_id)
  150. if job is None or job.result is None:
  151. return None
  152. return job.result.copy()
  153. def get_job_logs(self, job_id: str) -> List[str]:
  154. """
  155. Get training job logs.
  156. Args:
  157. job_id: Job identifier
  158. Returns:
  159. List of log messages
  160. """
  161. with self.lock:
  162. job = self.jobs.get(job_id)
  163. if job is None:
  164. return []
  165. return job.log_messages.copy()
  166. def cancel_job(self, job_id: str) -> bool:
  167. """
  168. Cancel a training job.
  169. Args:
  170. job_id: Job identifier
  171. Returns:
  172. True if job was cancelled successfully
  173. """
  174. with self.lock:
  175. job = self.jobs.get(job_id)
  176. if job is None:
  177. return False
  178. if job.status in [TrainingStatus.COMPLETED, TrainingStatus.FAILED, TrainingStatus.CANCELLED]:
  179. return False
  180. job.status = TrainingStatus.CANCELLED
  181. job.error_message = "Job cancelled by user"
  182. self.logger.info(f"Cancelled training job: {job_id}")
  183. return True
  184. def list_jobs(self) -> List[Dict[str, Any]]:
  185. """
  186. List all training jobs.
  187. Returns:
  188. List of job status dictionaries
  189. """
  190. with self.lock:
  191. return [self.get_job_status(job_id) for job_id in self.jobs.keys()]
  192. def cleanup_completed_jobs(self, max_age_hours: int = 24):
  193. """
  194. Clean up completed jobs older than specified age.
  195. Args:
  196. max_age_hours: Maximum age in hours for completed jobs
  197. """
  198. import time
  199. current_time = time.time()
  200. cutoff_time = current_time - (max_age_hours * 3600)
  201. with self.lock:
  202. jobs_to_remove = []
  203. for job_id, job in self.jobs.items():
  204. if job.status in [TrainingStatus.COMPLETED, TrainingStatus.FAILED, TrainingStatus.CANCELLED]:
  205. # In a real implementation, you'd track job completion time
  206. # For now, just keep a reasonable number of jobs
  207. if len(self.jobs) > 100: # Keep last 100 jobs
  208. jobs_to_remove.append(job_id)
  209. for job_id in jobs_to_remove[:len(jobs_to_remove)//2]: # Remove half
  210. del self.jobs[job_id]
  211. self.logger.info(f"Cleaned up old job: {job_id}")
  212. def _worker_loop(self):
  213. """Worker thread loop for processing training jobs."""
  214. while not self.shutdown_flag.is_set():
  215. try:
  216. # Get next job from queue (with timeout)
  217. job_id = self.job_queue.get(timeout=1.0)
  218. with self.lock:
  219. if job_id not in self.jobs:
  220. continue
  221. job = self.jobs[job_id]
  222. if job.status != TrainingStatus.PENDING:
  223. continue
  224. job.status = TrainingStatus.RUNNING
  225. self.active_jobs.add(job_id)
  226. # Execute training job
  227. self._execute_job(job)
  228. with self.lock:
  229. self.active_jobs.discard(job_id)
  230. except queue.Empty:
  231. continue
  232. except Exception as e:
  233. self.logger.error(f"Error in worker loop: {str(e)}")
  234. def _execute_job(self, job: TrainingJob):
  235. """
  236. Execute a training job.
  237. Args:
  238. job: Training job to execute
  239. """
  240. try:
  241. self.logger.info(f"Starting training job: {job.job_id}")
  242. # Create progress callback
  243. def progress_callback(epoch: int, total_epochs: int, metrics: Dict[str, Any]):
  244. job.current_epoch = epoch
  245. job.progress = epoch / total_epochs
  246. job.log_messages.append(f"Epoch {epoch}/{total_epochs}: {metrics}")
  247. # Create trainer
  248. if job.trainer_type == 'wakeword':
  249. trainer = WakewordTrainer(job.config)
  250. elif job.trainer_type == 'voice_recognition':
  251. trainer = VoiceRecognitionTrainer(job.config)
  252. else:
  253. raise ValueError(f"Unknown trainer type: {job.trainer_type}")
  254. # Add progress callback to trainer (if supported)
  255. if hasattr(trainer, 'add_progress_callback'):
  256. trainer.add_progress_callback(progress_callback)
  257. # Execute training
  258. result = trainer.train()
  259. # Store result
  260. job.result = result
  261. job.status = TrainingStatus.COMPLETED
  262. job.progress = 1.0
  263. self.logger.info(f"Completed training job: {job.job_id}")
  264. except Exception as e:
  265. job.status = TrainingStatus.FAILED
  266. job.error_message = str(e)
  267. job.log_messages.append(f"ERROR: {str(e)}")
  268. self.logger.error(f"Training job failed {job.job_id}: {str(e)}")
  269. def shutdown(self):
  270. """Shutdown the job manager."""
  271. self.logger.info("Shutting down training job manager...")
  272. self.shutdown_flag.set()
  273. for worker in self.worker_threads:
  274. worker.join(timeout=5.0)
  275. self.logger.info("Training job manager shut down")
  276. class ServerInterface:
  277. """
  278. Main server interface for ML training integration.
  279. This class provides a clean API for server scripts to interact
  280. with the ML training framework.
  281. """
  282. def __init__(self, config_dir: str = "./config", models_dir: str = "./models"):
  283. """
  284. Initialize server interface.
  285. Args:
  286. config_dir: Directory containing configuration files
  287. models_dir: Directory for storing trained models
  288. """
  289. self.config_dir = Path(config_dir)
  290. self.models_dir = Path(models_dir)
  291. self.job_manager = TrainingJobManager()
  292. # Ensure directories exist
  293. self.config_dir.mkdir(parents=True, exist_ok=True)
  294. self.models_dir.mkdir(parents=True, exist_ok=True)
  295. self.logger = setup_logging("server_interface", "INFO")
  296. self.logger.info("Server interface initialized")
  297. def train_wakeword_model(self, config_name: str = "wakeword_default",
  298. config_overrides: Optional[Dict[str, Any]] = None,
  299. job_id: Optional[str] = None) -> str:
  300. """
  301. Start wakeword model training.
  302. Args:
  303. config_name: Name of configuration file (without .json extension)
  304. config_overrides: Dictionary of configuration overrides
  305. job_id: Optional job ID (auto-generated if not provided)
  306. Returns:
  307. Training job ID
  308. """
  309. if job_id is None:
  310. import uuid
  311. job_id = f"wakeword_{uuid.uuid4().hex[:8]}"
  312. # Load base configuration
  313. config_file = self.config_dir / f"{config_name}.json"
  314. if config_file.exists():
  315. with open(config_file) as f:
  316. config_dict = json.load(f)
  317. else:
  318. # Use default configuration
  319. config_dict = {
  320. 'model_name': f'wakeword_{job_id}',
  321. 'data_dir': './trainer/data/wakeword',
  322. 'output_dir': str(self.models_dir / 'wakeword')
  323. }
  324. # Apply overrides
  325. if config_overrides:
  326. config_dict.update(config_overrides)
  327. return self.job_manager.create_wakeword_job(job_id, config_dict)
  328. def train_voice_recognition_model(self, config_name: str = "voice_recognition_default",
  329. config_overrides: Optional[Dict[str, Any]] = None,
  330. job_id: Optional[str] = None) -> str:
  331. """
  332. Start voice recognition model training.
  333. Args:
  334. config_name: Name of configuration file (without .json extension)
  335. config_overrides: Dictionary of configuration overrides
  336. job_id: Optional job ID (auto-generated if not provided)
  337. Returns:
  338. Training job ID
  339. """
  340. if job_id is None:
  341. import uuid
  342. job_id = f"voice_rec_{uuid.uuid4().hex[:8]}"
  343. # Load base configuration
  344. config_file = self.config_dir / f"{config_name}.json"
  345. if config_file.exists():
  346. with open(config_file) as f:
  347. config_dict = json.load(f)
  348. else:
  349. # Use default configuration
  350. config_dict = {
  351. 'model_name': f'voice_recognition_{job_id}',
  352. 'data_dir': './trainer/data/voice_recognition',
  353. 'output_dir': str(self.models_dir / 'voice_recognition')
  354. }
  355. # Apply overrides
  356. if config_overrides:
  357. config_dict.update(config_overrides)
  358. return self.job_manager.create_voice_recognition_job(job_id, config_dict)
  359. def get_training_status(self, job_id: str) -> Optional[Dict[str, Any]]:
  360. """Get training job status."""
  361. return self.job_manager.get_job_status(job_id)
  362. def get_training_result(self, job_id: str) -> Optional[Dict[str, Any]]:
  363. """Get training job result."""
  364. return self.job_manager.get_job_result(job_id)
  365. def get_training_logs(self, job_id: str) -> List[str]:
  366. """Get training job logs."""
  367. return self.job_manager.get_job_logs(job_id)
  368. def cancel_training(self, job_id: str) -> bool:
  369. """Cancel training job."""
  370. return self.job_manager.cancel_job(job_id)
  371. def list_training_jobs(self) -> List[Dict[str, Any]]:
  372. """List all training jobs."""
  373. return self.job_manager.list_jobs()
  374. def list_available_models(self) -> Dict[str, List[str]]:
  375. """
  376. List available trained models.
  377. Returns:
  378. Dictionary with model types as keys and lists of model names as values
  379. """
  380. models = {
  381. 'wakeword': [],
  382. 'voice_recognition': []
  383. }
  384. # Scan for wakeword models
  385. wakeword_dir = self.models_dir / 'wakeword'
  386. if wakeword_dir.exists():
  387. for model_file in wakeword_dir.rglob('*.pth'):
  388. models['wakeword'].append(str(model_file.relative_to(wakeword_dir)))
  389. # Scan for voice recognition models
  390. voice_dir = self.models_dir / 'voice_recognition'
  391. if voice_dir.exists():
  392. for model_file in voice_dir.rglob('*.pth'):
  393. models['voice_recognition'].append(str(model_file.relative_to(voice_dir)))
  394. return models
  395. def get_model_metadata(self, model_path: str) -> Optional[Dict[str, Any]]:
  396. """
  397. Get metadata for a trained model.
  398. Args:
  399. model_path: Path to model file (relative to models directory)
  400. Returns:
  401. Model metadata dictionary or None
  402. """
  403. try:
  404. full_path = self.models_dir / model_path
  405. metadata_path = full_path.parent / "metadata.json"
  406. if metadata_path.exists():
  407. with open(metadata_path) as f:
  408. return json.load(f)
  409. return None
  410. except Exception as e:
  411. self.logger.error(f"Failed to load model metadata: {str(e)}")
  412. return None
  413. def cleanup_old_jobs(self, max_age_hours: int = 24):
  414. """Clean up old training jobs."""
  415. self.job_manager.cleanup_completed_jobs(max_age_hours)
  416. def shutdown(self):
  417. """Shutdown the server interface."""
  418. self.logger.info("Shutting down server interface...")
  419. self.job_manager.shutdown()
  420. # Global server interface instance for easy import
  421. _server_interface = None
  422. def get_server_interface(config_dir: str = "./config",
  423. models_dir: str = "./models") -> ServerInterface:
  424. """
  425. Get the global server interface instance.
  426. Args:
  427. config_dir: Directory containing configuration files
  428. models_dir: Directory for storing trained models
  429. Returns:
  430. ServerInterface instance
  431. """
  432. global _server_interface
  433. if _server_interface is None:
  434. _server_interface = ServerInterface(config_dir, models_dir)
  435. return _server_interface
  436. def shutdown_server_interface():
  437. """Shutdown the global server interface."""
  438. global _server_interface
  439. if _server_interface is not None:
  440. _server_interface.shutdown()
  441. _server_interface = None