| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552 |
- """
- Server integration module for the ML Trainer Framework.
- This module provides integration capabilities for server script inclusion
- and programmatic access to the training framework from other components
- of the Trixy voice assistant system.
- """
- import os
- import json
- import logging
- from pathlib import Path
- from typing import Dict, Any, Optional, List, Callable, Union
- import threading
- import queue
- from dataclasses import dataclass
- from enum import Enum
- from .base import TrainerConfig
- from .wakeword.trainer import WakewordTrainer, create_wakeword_trainer_config
- from .voice_recognition.trainer import VoiceRecognitionTrainer, create_voice_recognition_trainer_config
- from .utils import setup_logging
- class TrainingStatus(Enum):
- """Training status enumeration."""
- PENDING = "pending"
- RUNNING = "running"
- COMPLETED = "completed"
- FAILED = "failed"
- CANCELLED = "cancelled"
- @dataclass
- class TrainingJob:
- """Training job definition."""
- job_id: str
- trainer_type: str # 'wakeword' or 'voice_recognition'
- config: TrainerConfig
- status: TrainingStatus = TrainingStatus.PENDING
- progress: float = 0.0
- current_epoch: int = 0
- total_epochs: int = 0
- error_message: Optional[str] = None
- result: Optional[Dict[str, Any]] = None
- log_messages: List[str] = None
-
- def __post_init__(self):
- if self.log_messages is None:
- self.log_messages = []
- class TrainingJobManager:
- """
- Manages training jobs with support for background execution,
- progress tracking, and result retrieval.
- """
-
- def __init__(self, max_concurrent_jobs: int = 1):
- """
- Initialize training job manager.
-
- Args:
- max_concurrent_jobs: Maximum number of concurrent training jobs
- """
- self.max_concurrent_jobs = max_concurrent_jobs
- self.jobs: Dict[str, TrainingJob] = {}
- self.job_queue = queue.Queue()
- self.active_jobs = set()
- self.lock = threading.Lock()
- self.worker_threads = []
- self.shutdown_flag = threading.Event()
-
- # Setup logging
- self.logger = setup_logging("training_job_manager", "INFO")
-
- # Start worker threads
- for i in range(max_concurrent_jobs):
- worker = threading.Thread(target=self._worker_loop, daemon=True)
- worker.start()
- self.worker_threads.append(worker)
-
- def create_wakeword_job(self, job_id: str, config_dict: Dict[str, Any]) -> str:
- """
- Create a wakeword training job.
-
- Args:
- job_id: Unique job identifier
- config_dict: Training configuration dictionary
-
- Returns:
- Job ID
- """
- try:
- config = create_wakeword_trainer_config(**config_dict)
- job = TrainingJob(
- job_id=job_id,
- trainer_type='wakeword',
- config=config,
- total_epochs=config.num_epochs
- )
-
- with self.lock:
- self.jobs[job_id] = job
- self.job_queue.put(job_id)
-
- self.logger.info(f"Created wakeword training job: {job_id}")
- return job_id
-
- except Exception as e:
- self.logger.error(f"Failed to create wakeword job {job_id}: {str(e)}")
- raise
-
- def create_voice_recognition_job(self, job_id: str, config_dict: Dict[str, Any]) -> str:
- """
- Create a voice recognition training job.
-
- Args:
- job_id: Unique job identifier
- config_dict: Training configuration dictionary
-
- Returns:
- Job ID
- """
- try:
- config = create_voice_recognition_trainer_config(**config_dict)
- job = TrainingJob(
- job_id=job_id,
- trainer_type='voice_recognition',
- config=config,
- total_epochs=config.num_epochs
- )
-
- with self.lock:
- self.jobs[job_id] = job
- self.job_queue.put(job_id)
-
- self.logger.info(f"Created voice recognition training job: {job_id}")
- return job_id
-
- except Exception as e:
- self.logger.error(f"Failed to create voice recognition job {job_id}: {str(e)}")
- raise
-
- def get_job_status(self, job_id: str) -> Optional[Dict[str, Any]]:
- """
- Get training job status.
-
- Args:
- job_id: Job identifier
-
- Returns:
- Job status dictionary or None if job not found
- """
- with self.lock:
- job = self.jobs.get(job_id)
- if job is None:
- return None
-
- return {
- 'job_id': job.job_id,
- 'trainer_type': job.trainer_type,
- 'status': job.status.value,
- 'progress': job.progress,
- 'current_epoch': job.current_epoch,
- 'total_epochs': job.total_epochs,
- 'error_message': job.error_message,
- 'has_result': job.result is not None
- }
-
- def get_job_result(self, job_id: str) -> Optional[Dict[str, Any]]:
- """
- Get training job result.
-
- Args:
- job_id: Job identifier
-
- Returns:
- Job result or None
- """
- with self.lock:
- job = self.jobs.get(job_id)
- if job is None or job.result is None:
- return None
- return job.result.copy()
-
- def get_job_logs(self, job_id: str) -> List[str]:
- """
- Get training job logs.
-
- Args:
- job_id: Job identifier
-
- Returns:
- List of log messages
- """
- with self.lock:
- job = self.jobs.get(job_id)
- if job is None:
- return []
- return job.log_messages.copy()
-
- def cancel_job(self, job_id: str) -> bool:
- """
- Cancel a training job.
-
- Args:
- job_id: Job identifier
-
- Returns:
- True if job was cancelled successfully
- """
- with self.lock:
- job = self.jobs.get(job_id)
- if job is None:
- return False
-
- if job.status in [TrainingStatus.COMPLETED, TrainingStatus.FAILED, TrainingStatus.CANCELLED]:
- return False
-
- job.status = TrainingStatus.CANCELLED
- job.error_message = "Job cancelled by user"
-
- self.logger.info(f"Cancelled training job: {job_id}")
- return True
-
- def list_jobs(self) -> List[Dict[str, Any]]:
- """
- List all training jobs.
-
- Returns:
- List of job status dictionaries
- """
- with self.lock:
- return [self.get_job_status(job_id) for job_id in self.jobs.keys()]
-
- def cleanup_completed_jobs(self, max_age_hours: int = 24):
- """
- Clean up completed jobs older than specified age.
-
- Args:
- max_age_hours: Maximum age in hours for completed jobs
- """
- import time
- current_time = time.time()
- cutoff_time = current_time - (max_age_hours * 3600)
-
- with self.lock:
- jobs_to_remove = []
- for job_id, job in self.jobs.items():
- if job.status in [TrainingStatus.COMPLETED, TrainingStatus.FAILED, TrainingStatus.CANCELLED]:
- # In a real implementation, you'd track job completion time
- # For now, just keep a reasonable number of jobs
- if len(self.jobs) > 100: # Keep last 100 jobs
- jobs_to_remove.append(job_id)
-
- for job_id in jobs_to_remove[:len(jobs_to_remove)//2]: # Remove half
- del self.jobs[job_id]
- self.logger.info(f"Cleaned up old job: {job_id}")
-
- def _worker_loop(self):
- """Worker thread loop for processing training jobs."""
- while not self.shutdown_flag.is_set():
- try:
- # Get next job from queue (with timeout)
- job_id = self.job_queue.get(timeout=1.0)
-
- with self.lock:
- if job_id not in self.jobs:
- continue
-
- job = self.jobs[job_id]
- if job.status != TrainingStatus.PENDING:
- continue
-
- job.status = TrainingStatus.RUNNING
- self.active_jobs.add(job_id)
-
- # Execute training job
- self._execute_job(job)
-
- with self.lock:
- self.active_jobs.discard(job_id)
-
- except queue.Empty:
- continue
- except Exception as e:
- self.logger.error(f"Error in worker loop: {str(e)}")
-
- def _execute_job(self, job: TrainingJob):
- """
- Execute a training job.
-
- Args:
- job: Training job to execute
- """
- try:
- self.logger.info(f"Starting training job: {job.job_id}")
-
- # Create progress callback
- def progress_callback(epoch: int, total_epochs: int, metrics: Dict[str, Any]):
- job.current_epoch = epoch
- job.progress = epoch / total_epochs
- job.log_messages.append(f"Epoch {epoch}/{total_epochs}: {metrics}")
-
- # Create trainer
- if job.trainer_type == 'wakeword':
- trainer = WakewordTrainer(job.config)
- elif job.trainer_type == 'voice_recognition':
- trainer = VoiceRecognitionTrainer(job.config)
- else:
- raise ValueError(f"Unknown trainer type: {job.trainer_type}")
-
- # Add progress callback to trainer (if supported)
- if hasattr(trainer, 'add_progress_callback'):
- trainer.add_progress_callback(progress_callback)
-
- # Execute training
- result = trainer.train()
-
- # Store result
- job.result = result
- job.status = TrainingStatus.COMPLETED
- job.progress = 1.0
-
- self.logger.info(f"Completed training job: {job.job_id}")
-
- except Exception as e:
- job.status = TrainingStatus.FAILED
- job.error_message = str(e)
- job.log_messages.append(f"ERROR: {str(e)}")
- self.logger.error(f"Training job failed {job.job_id}: {str(e)}")
-
- def shutdown(self):
- """Shutdown the job manager."""
- self.logger.info("Shutting down training job manager...")
- self.shutdown_flag.set()
-
- for worker in self.worker_threads:
- worker.join(timeout=5.0)
-
- self.logger.info("Training job manager shut down")
- class ServerInterface:
- """
- Main server interface for ML training integration.
-
- This class provides a clean API for server scripts to interact
- with the ML training framework.
- """
-
- def __init__(self, config_dir: str = "./config", models_dir: str = "./models"):
- """
- Initialize server interface.
-
- Args:
- config_dir: Directory containing configuration files
- models_dir: Directory for storing trained models
- """
- self.config_dir = Path(config_dir)
- self.models_dir = Path(models_dir)
- self.job_manager = TrainingJobManager()
-
- # Ensure directories exist
- self.config_dir.mkdir(parents=True, exist_ok=True)
- self.models_dir.mkdir(parents=True, exist_ok=True)
-
- self.logger = setup_logging("server_interface", "INFO")
- self.logger.info("Server interface initialized")
-
- def train_wakeword_model(self, config_name: str = "wakeword_default",
- config_overrides: Optional[Dict[str, Any]] = None,
- job_id: Optional[str] = None) -> str:
- """
- Start wakeword model training.
-
- Args:
- config_name: Name of configuration file (without .json extension)
- config_overrides: Dictionary of configuration overrides
- job_id: Optional job ID (auto-generated if not provided)
-
- Returns:
- Training job ID
- """
- if job_id is None:
- import uuid
- job_id = f"wakeword_{uuid.uuid4().hex[:8]}"
-
- # Load base configuration
- config_file = self.config_dir / f"{config_name}.json"
- if config_file.exists():
- with open(config_file) as f:
- config_dict = json.load(f)
- else:
- # Use default configuration
- config_dict = {
- 'model_name': f'wakeword_{job_id}',
- 'data_dir': './trainer/data/wakeword',
- 'output_dir': str(self.models_dir / 'wakeword')
- }
-
- # Apply overrides
- if config_overrides:
- config_dict.update(config_overrides)
-
- return self.job_manager.create_wakeword_job(job_id, config_dict)
-
- def train_voice_recognition_model(self, config_name: str = "voice_recognition_default",
- config_overrides: Optional[Dict[str, Any]] = None,
- job_id: Optional[str] = None) -> str:
- """
- Start voice recognition model training.
-
- Args:
- config_name: Name of configuration file (without .json extension)
- config_overrides: Dictionary of configuration overrides
- job_id: Optional job ID (auto-generated if not provided)
-
- Returns:
- Training job ID
- """
- if job_id is None:
- import uuid
- job_id = f"voice_rec_{uuid.uuid4().hex[:8]}"
-
- # Load base configuration
- config_file = self.config_dir / f"{config_name}.json"
- if config_file.exists():
- with open(config_file) as f:
- config_dict = json.load(f)
- else:
- # Use default configuration
- config_dict = {
- 'model_name': f'voice_recognition_{job_id}',
- 'data_dir': './trainer/data/voice_recognition',
- 'output_dir': str(self.models_dir / 'voice_recognition')
- }
-
- # Apply overrides
- if config_overrides:
- config_dict.update(config_overrides)
-
- return self.job_manager.create_voice_recognition_job(job_id, config_dict)
-
- def get_training_status(self, job_id: str) -> Optional[Dict[str, Any]]:
- """Get training job status."""
- return self.job_manager.get_job_status(job_id)
-
- def get_training_result(self, job_id: str) -> Optional[Dict[str, Any]]:
- """Get training job result."""
- return self.job_manager.get_job_result(job_id)
-
- def get_training_logs(self, job_id: str) -> List[str]:
- """Get training job logs."""
- return self.job_manager.get_job_logs(job_id)
-
- def cancel_training(self, job_id: str) -> bool:
- """Cancel training job."""
- return self.job_manager.cancel_job(job_id)
-
- def list_training_jobs(self) -> List[Dict[str, Any]]:
- """List all training jobs."""
- return self.job_manager.list_jobs()
-
- def list_available_models(self) -> Dict[str, List[str]]:
- """
- List available trained models.
-
- Returns:
- Dictionary with model types as keys and lists of model names as values
- """
- models = {
- 'wakeword': [],
- 'voice_recognition': []
- }
-
- # Scan for wakeword models
- wakeword_dir = self.models_dir / 'wakeword'
- if wakeword_dir.exists():
- for model_file in wakeword_dir.rglob('*.pth'):
- models['wakeword'].append(str(model_file.relative_to(wakeword_dir)))
-
- # Scan for voice recognition models
- voice_dir = self.models_dir / 'voice_recognition'
- if voice_dir.exists():
- for model_file in voice_dir.rglob('*.pth'):
- models['voice_recognition'].append(str(model_file.relative_to(voice_dir)))
-
- return models
-
- def get_model_metadata(self, model_path: str) -> Optional[Dict[str, Any]]:
- """
- Get metadata for a trained model.
-
- Args:
- model_path: Path to model file (relative to models directory)
-
- Returns:
- Model metadata dictionary or None
- """
- try:
- full_path = self.models_dir / model_path
- metadata_path = full_path.parent / "metadata.json"
-
- if metadata_path.exists():
- with open(metadata_path) as f:
- return json.load(f)
- return None
-
- except Exception as e:
- self.logger.error(f"Failed to load model metadata: {str(e)}")
- return None
-
- def cleanup_old_jobs(self, max_age_hours: int = 24):
- """Clean up old training jobs."""
- self.job_manager.cleanup_completed_jobs(max_age_hours)
-
- def shutdown(self):
- """Shutdown the server interface."""
- self.logger.info("Shutting down server interface...")
- self.job_manager.shutdown()
- # Global server interface instance for easy import
- _server_interface = None
- def get_server_interface(config_dir: str = "./config",
- models_dir: str = "./models") -> ServerInterface:
- """
- Get the global server interface instance.
-
- Args:
- config_dir: Directory containing configuration files
- models_dir: Directory for storing trained models
-
- Returns:
- ServerInterface instance
- """
- global _server_interface
- if _server_interface is None:
- _server_interface = ServerInterface(config_dir, models_dir)
- return _server_interface
- def shutdown_server_interface():
- """Shutdown the global server interface."""
- global _server_interface
- if _server_interface is not None:
- _server_interface.shutdown()
- _server_interface = None
|