| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623 |
- """
- Command-line interface for the ML Trainer Framework.
- This module provides a comprehensive CLI for training wakeword detection,
- voice recognition, and other models with support for configuration files
- and command-line argument overrides.
- """
- import os
- import sys
- import json
- import argparse
- import logging
- from pathlib import Path
- from typing import Dict, Any, Optional, List
- from .config import ConfigManager, TrainerConfig, ModelFormat
- from .wakeword.trainer import WakewordTrainer, create_wakeword_trainer_config
- from .voice_recognition.trainer import VoiceRecognitionTrainer, create_voice_recognition_trainer_config
- from .base import TrainerConfig as BaseTrainerConfig
- from dataclasses import fields
- # Note: setup_logging function was moved to TrainerLogger class in utils
- def setup_cli_logging(level: str = "INFO", log_file: Optional[str] = None) -> logging.Logger:
- """Setup logging for CLI operations."""
- logger = logging.getLogger("trainer_cli")
- logger.setLevel(getattr(logging, level.upper()))
- # Prevent duplicate logging by disabling propagation to root logger
- logger.propagate = False
- # CRITICAL: Always clear existing handlers to prevent duplicates
- logger.handlers.clear()
- # Console handler
- console_handler = logging.StreamHandler()
- console_formatter = logging.Formatter('%(asctime)s - %(name)s - %(levelname)s - %(message)s')
- console_handler.setFormatter(console_formatter)
- logger.addHandler(console_handler)
- # File handler if requested
- if log_file:
- file_handler = logging.FileHandler(log_file)
- file_handler.setFormatter(console_formatter)
- logger.addHandler(file_handler)
- return logger
- def load_config_file(config_path: str) -> Dict[str, Any]:
- """Load configuration from JSON file."""
- try:
- with open(config_path, 'r') as f:
- return json.load(f)
- except FileNotFoundError:
- raise FileNotFoundError(f"Configuration file not found: {config_path}")
- except json.JSONDecodeError as e:
- raise ValueError(f"Invalid JSON in configuration file: {e}")
- def create_parser() -> argparse.ArgumentParser:
- """Create the main argument parser."""
- parser = argparse.ArgumentParser(
- description="Trixy ML Trainer Framework",
- formatter_class=argparse.RawDescriptionHelpFormatter,
- epilog="""
- Examples:
- # Train wakeword detection model
- python -m trainer.cli train wakeword --config ./config/wakeword.json
-
- # Train voice recognition model with overrides
- python -m trainer.cli train voice-recognition --model-type ecapa_tdnn --batch-size 64
-
- # List available models
- python -m trainer.cli list-models
-
- # Validate model
- python -m trainer.cli validate ./models/my_model.pth --test-data ./data/test
- """
- )
-
- # Global arguments
- parser.add_argument(
- "--config", "-c", type=str,
- help="Path to configuration file (JSON)"
- )
- parser.add_argument(
- "--log-level", choices=["DEBUG", "INFO", "WARNING", "ERROR"],
- default="INFO", help="Logging level"
- )
- parser.add_argument(
- "--log-file", type=str,
- help="Path to log file (default: console only)"
- )
- parser.add_argument(
- "--device", choices=["auto", "cpu", "cuda"],
- default="auto", help="Device to use for training/inference"
- )
-
- # Subcommands
- subparsers = parser.add_subparsers(dest="command", help="Available commands")
-
- # Train command
- train_parser = subparsers.add_parser("train", help="Train a model")
- train_subparsers = train_parser.add_subparsers(dest="train_type", help="Type of model to train")
-
- # Wakeword training
- wakeword_parser = train_subparsers.add_parser("wakeword", help="Train wakeword detection model")
- add_wakeword_arguments(wakeword_parser)
-
- # Voice recognition training
- voice_parser = train_subparsers.add_parser("voice-recognition", help="Train voice recognition model")
- add_voice_recognition_arguments(voice_parser)
-
- # Validate command
- validate_parser = subparsers.add_parser("validate", help="Validate a trained model")
- add_validation_arguments(validate_parser)
-
- # List models command
- list_parser = subparsers.add_parser("list-models", help="List available trained models")
- add_list_arguments(list_parser)
-
- # Convert model command
- convert_parser = subparsers.add_parser("convert", help="Convert model between formats")
- add_convert_arguments(convert_parser)
-
- return parser
- def add_wakeword_arguments(parser: argparse.ArgumentParser):
- """Add wakeword-specific arguments."""
- # Model configuration
- parser.add_argument("--model-name", type=str, default="wakeword_model",
- help="Name of the model")
- parser.add_argument("--model-type", choices=["standard", "improved", "lightweight"],
- default="standard", help="Type of RepCNN model")
- parser.add_argument("--model-format", choices=[fmt.value for fmt in ModelFormat],
- default=ModelFormat.PTH.value, help="Model output format")
-
- # Data configuration
- parser.add_argument("--data-dir", type=str, required=True,
- help="Directory containing training data")
- parser.add_argument("--output-dir", type=str, default="./models/wakeword",
- help="Output directory for trained model")
-
- # Training parameters
- parser.add_argument("--batch-size", type=int, default=64,
- help="Training batch size")
- parser.add_argument("--learning-rate", type=float, default=0.001,
- help="Learning rate")
- parser.add_argument("--num-epochs", type=int, default=100,
- help="Number of training epochs")
- parser.add_argument("--min-epochs", type=int, default=20,
- help="Minimum number of training epochs")
- parser.add_argument("--early-stopping-patience", type=int, default=15,
- help="Early stopping patience")
-
- # Audio parameters
- parser.add_argument("--sample-rate", type=int, default=16000,
- help="Audio sample rate")
- parser.add_argument("--audio-length", type=float, default=1.5,
- help="Target audio length in seconds")
- parser.add_argument("--n-mels", type=int, default=40,
- help="Number of mel filterbank features")
-
- # Augmentation
- parser.add_argument("--use-augmentation", action="store_true",
- help="Enable data augmentation")
- parser.add_argument("--noise-factor", type=float, default=0.1,
- help="Noise augmentation factor")
-
- # Model specific
- parser.add_argument("--use-focal-loss", action="store_true",
- help="Use focal loss for imbalanced data")
- parser.add_argument("--detection-threshold", type=float, default=0.5,
- help="Detection threshold for wakeword")
-
- # Security
- parser.add_argument("--use-password-protection", action="store_true",
- help="Enable password protection for model")
- parser.add_argument("--password", type=str,
- help="Password for model protection")
- def add_voice_recognition_arguments(parser: argparse.ArgumentParser):
- """Add voice recognition specific arguments."""
- # Model configuration
- parser.add_argument("--model-name", type=str, default="voice_recognition_model",
- help="Name of the model")
- parser.add_argument("--model-type", choices=["ecapa_tdnn", "titanet_s", "speakernet_m"],
- default="ecapa_tdnn", help="Type of voice recognition model")
- parser.add_argument("--model-format", choices=[fmt.value for fmt in ModelFormat],
- default=ModelFormat.PTH.value, help="Model output format")
-
- # Data configuration
- parser.add_argument("--data-dir", type=str, required=True,
- help="Directory containing speaker data")
- parser.add_argument("--output-dir", type=str, default="./models/voice_recognition",
- help="Output directory for trained model")
-
- # Training parameters
- parser.add_argument("--batch-size", type=int, default=32,
- help="Training batch size")
- parser.add_argument("--learning-rate", type=float, default=0.001,
- help="Learning rate")
- parser.add_argument("--num-epochs", type=int, default=200,
- help="Number of training epochs")
- parser.add_argument("--min-epochs", type=int, default=30,
- help="Minimum number of training epochs")
- parser.add_argument("--early-stopping-patience", type=int, default=20,
- help="Early stopping patience")
-
- # Audio parameters
- parser.add_argument("--sample-rate", type=int, default=16000,
- help="Audio sample rate")
- parser.add_argument("--audio-length", type=float, default=3.0,
- help="Target audio length in seconds")
- parser.add_argument("--n-mels", type=int, default=40,
- help="Number of mel filterbank features")
-
- # Loss function
- parser.add_argument("--loss-type",
- choices=["angular_margin", "ge2e", "contrastive", "triplet", "cross_entropy"],
- default="angular_margin", help="Loss function type")
- parser.add_argument("--embedding-dim", type=int, default=192,
- help="Embedding dimension")
-
- # Augmentation
- parser.add_argument("--use-augmentation", action="store_true",
- help="Enable data augmentation")
-
- # Security
- parser.add_argument("--use-password-protection", action="store_true",
- help="Enable password protection for model")
- parser.add_argument("--password", type=str,
- help="Password for model protection")
- def add_validation_arguments(parser: argparse.ArgumentParser):
- """Add validation arguments."""
- parser.add_argument("model_path", type=str,
- help="Path to model file to validate")
- parser.add_argument("--test-data", type=str, required=True,
- help="Path to test data directory")
- parser.add_argument("--batch-size", type=int, default=32,
- help="Validation batch size")
- parser.add_argument("--output-dir", type=str,
- help="Output directory for validation results")
- def add_list_arguments(parser: argparse.ArgumentParser):
- """Add list models arguments."""
- parser.add_argument("--models-dir", type=str, default="./models",
- help="Directory to search for models")
- parser.add_argument("--model-type", choices=["all", "wakeword", "voice-recognition"],
- default="all", help="Type of models to list")
- def add_convert_arguments(parser: argparse.ArgumentParser):
- """Add model conversion arguments."""
- parser.add_argument("input_model", type=str,
- help="Path to input model file")
- parser.add_argument("output_model", type=str,
- help="Path for output model file")
- parser.add_argument("--target-format", choices=[fmt.value for fmt in ModelFormat],
- required=True, help="Target model format")
- parser.add_argument("--password", type=str,
- help="Password for password-protected models")
- def merge_config_and_args(config_dict: Dict[str, Any], args: argparse.Namespace) -> Dict[str, Any]:
- """Merge configuration file and command-line arguments."""
- # Command-line arguments override config file
- merged = config_dict.copy()
-
- # Convert args to dict and filter None values and control arguments
- # These arguments are for CLI control and should not be passed to TrainerConfig
- cli_control_args = {'command', 'train_type', 'config', 'log_level', 'log_file'}
- args_dict = {k: v for k, v in vars(args).items()
- if v is not None and k not in cli_control_args}
-
- # Merge, giving priority to command-line arguments
- merged.update(args_dict)
-
- return merged
- def organize_voice_recognition_params(merged_config: Dict[str, Any]) -> Dict[str, Any]:
- """
- Organize parameters for voice recognition trainer.
-
- Separates base TrainerConfig parameters from model-specific parameters,
- putting model-specific parameters into custom_params.
- """
- # Get valid TrainerConfig field names
- trainer_config_fields = {f.name for f in fields(BaseTrainerConfig)}
-
- # Parameters that should go to create_voice_recognition_trainer_config directly
- direct_params = {
- 'model_name', 'model_type', 'loss_type', 'data_dir'
- }
-
- # Model-specific parameters that go into custom_params
- model_specific_params = {
- 'embedding_dim', 'model_format', 'detection_threshold',
- 'use_focal_loss', 'use_password_protection', 'password'
- }
-
- # Organize parameters
- organized = {}
- custom_params = {}
-
- for key, value in merged_config.items():
- # Convert hyphens to underscores
- clean_key = key.replace('-', '_')
-
- if clean_key in direct_params:
- # Parameters that go directly to the creation function
- organized[clean_key] = value
- elif clean_key in model_specific_params:
- # Model-specific parameters go into custom_params
- custom_params[clean_key] = value
- elif clean_key in trainer_config_fields:
- # Valid TrainerConfig parameters
- organized[clean_key] = value
- else:
- # Unknown parameters go into custom_params for safety
- custom_params[clean_key] = value
-
- # Add custom_params if any
- if custom_params:
- organized['custom_params'] = custom_params
-
- return organized
- def organize_wakeword_params(merged_config: Dict[str, Any]) -> Dict[str, Any]:
- """
- Organize parameters for wakeword trainer.
-
- Separates base TrainerConfig parameters from model-specific parameters,
- putting model-specific parameters into custom_params.
- """
- # Get valid TrainerConfig field names
- trainer_config_fields = {f.name for f in fields(BaseTrainerConfig)}
-
- # Parameters that should go to create_wakeword_trainer_config directly
- direct_params = {
- 'model_name', 'model_type', 'data_dir'
- }
-
- # Model-specific parameters that go into custom_params
- model_specific_params = {
- 'model_format', 'detection_threshold', 'use_focal_loss',
- 'use_password_protection', 'password', 'wakeword_classes'
- }
-
- # Organize parameters
- organized = {}
- custom_params = {}
-
- for key, value in merged_config.items():
- # Convert hyphens to underscores
- clean_key = key.replace('-', '_')
-
- if clean_key in direct_params:
- # Parameters that go directly to the creation function
- organized[clean_key] = value
- elif clean_key in model_specific_params:
- # Model-specific parameters go into custom_params
- custom_params[clean_key] = value
- elif clean_key in trainer_config_fields:
- # Valid TrainerConfig parameters
- organized[clean_key] = value
- else:
- # Unknown parameters go into custom_params for safety
- custom_params[clean_key] = value
-
- # Add custom_params if any
- if custom_params:
- organized['custom_params'] = custom_params
-
- return organized
- def train_wakeword(args: argparse.Namespace, logger: logging.Logger):
- """Train wakeword detection model."""
- logger.info("Starting wakeword detection training...")
-
- # Load configuration if provided
- config_dict = {}
- if args.config:
- config_dict = load_config_file(args.config)
-
- # Merge with command-line arguments
- merged_config = merge_config_and_args(config_dict, args)
-
- # Separate TrainerConfig parameters from model-specific parameters
- trainer_config_params = organize_wakeword_params(merged_config)
-
- # Create trainer configuration
- try:
- trainer_config = create_wakeword_trainer_config(**trainer_config_params)
- except Exception as e:
- logger.error(f"Failed to create trainer configuration: {e}")
- logger.error(f"Parameters passed: {list(trainer_config_params.keys())}")
- return False
-
- # Create and run trainer
- try:
- trainer = WakewordTrainer(trainer_config)
- results = trainer.train()
-
- logger.info("Wakeword training completed successfully!")
- logger.info(f"Model saved to: {trainer_config.output_dir}/{trainer_config.model_name}")
-
- return True
-
- except Exception as e:
- logger.error(f"Training failed: {e}")
- return False
- def train_voice_recognition(args: argparse.Namespace, logger: logging.Logger):
- """Train voice recognition model."""
- logger.info("Starting voice recognition training...")
-
- # Load configuration if provided
- config_dict = {}
- if args.config:
- config_dict = load_config_file(args.config)
-
- # Merge with command-line arguments
- merged_config = merge_config_and_args(config_dict, args)
-
- # Separate TrainerConfig parameters from model-specific parameters
- trainer_config_params = organize_voice_recognition_params(merged_config)
-
- # Create trainer configuration
- try:
- trainer_config = create_voice_recognition_trainer_config(**trainer_config_params)
- except Exception as e:
- logger.error(f"Failed to create trainer configuration: {e}")
- logger.error(f"Parameters passed: {list(trainer_config_params.keys())}")
- return False
-
- # Create and run trainer
- try:
- trainer = VoiceRecognitionTrainer(trainer_config)
- results = trainer.train()
-
- logger.info("Voice recognition training completed successfully!")
- logger.info(f"Model saved to: {trainer_config.output_dir}/{trainer_config.model_name}")
-
- return True
-
- except Exception as e:
- logger.error(f"Training failed: {e}")
- return False
- def validate_model(args: argparse.Namespace, logger: logging.Logger):
- """Validate a trained model."""
- logger.info(f"Validating model: {args.model_path}")
-
- try:
- from .model_formats import ModelFormatManager
- from .validation import ModelValidator, ValidationConfig
-
- # Load model
- format_manager = ModelFormatManager(logger)
- model, metadata = format_manager.load_model(args.model_path, args.password)
-
- if model is None:
- logger.error("Failed to load model")
- return False
-
- # Create validator
- validator = ModelValidator(model, args.device, logger)
-
- # Load test data (simplified - would need proper data loading)
- logger.warning("Test data loading not fully implemented in CLI")
-
- logger.info("Model validation completed")
- return True
-
- except Exception as e:
- logger.error(f"Validation failed: {e}")
- return False
- def list_models(args: argparse.Namespace, logger: logging.Logger):
- """List available trained models."""
- models_dir = Path(args.models_dir)
-
- if not models_dir.exists():
- logger.warning(f"Models directory not found: {models_dir}")
- return True
-
- logger.info(f"Searching for models in: {models_dir}")
-
- # Find model files
- model_extensions = ['.pth', '.pt', '.onnx']
- models_found = []
-
- for ext in model_extensions:
- models_found.extend(list(models_dir.rglob(f"*{ext}")))
-
- if not models_found:
- logger.info("No models found")
- return True
-
- # Display models
- print(f"\nFound {len(models_found)} models:")
- print("-" * 60)
-
- for model_path in sorted(models_found):
- rel_path = model_path.relative_to(models_dir)
- model_type = "Unknown"
-
- # Try to determine model type from path
- if "wakeword" in str(rel_path).lower():
- model_type = "Wakeword"
- elif "voice" in str(rel_path).lower() or "speaker" in str(rel_path).lower():
- model_type = "Voice Recognition"
-
- size_mb = model_path.stat().st_size / (1024 * 1024)
-
- print(f"📁 {rel_path}")
- print(f" Type: {model_type}")
- print(f" Size: {size_mb:.1f} MB")
- print(f" Format: {model_path.suffix}")
- print()
-
- return True
- def convert_model(args: argparse.Namespace, logger: logging.Logger):
- """Convert model between formats."""
- logger.info(f"Converting model from {args.input_model} to {args.output_model}")
-
- try:
- from .model_formats import ModelFormatManager
-
- format_manager = ModelFormatManager(logger)
-
- # Load model
- model, metadata = format_manager.load_model(args.input_model, args.password)
-
- if model is None:
- logger.error("Failed to load input model")
- return False
-
- # Save in target format
- success = format_manager.save_model(
- model, args.output_model, metadata,
- password=args.password if args.password else None
- )
-
- if success:
- logger.info(f"Model converted successfully to: {args.output_model}")
- return True
- else:
- logger.error("Model conversion failed")
- return False
-
- except Exception as e:
- logger.error(f"Conversion failed: {e}")
- return False
- def main():
- """Main CLI entry point."""
- parser = create_parser()
- args = parser.parse_args()
-
- # Setup logging
- logger = setup_cli_logging(args.log_level, args.log_file)
-
- # Handle device selection
- if args.device == "auto":
- import torch
- device = "cuda" if torch.cuda.is_available() else "cpu"
- args.device = device
- logger.info(f"Auto-selected device: {device}")
-
- success = False
-
- try:
- if args.command == "train":
- if args.train_type == "wakeword":
- success = train_wakeword(args, logger)
- elif args.train_type == "voice-recognition":
- success = train_voice_recognition(args, logger)
- else:
- logger.error("Please specify training type: wakeword or voice-recognition")
-
- elif args.command == "validate":
- success = validate_model(args, logger)
-
- elif args.command == "list-models":
- success = list_models(args, logger)
-
- elif args.command == "convert":
- success = convert_model(args, logger)
-
- else:
- parser.print_help()
- success = True
-
- except KeyboardInterrupt:
- logger.info("Operation cancelled by user")
- success = True
- except Exception as e:
- logger.error(f"Unexpected error: {e}")
- success = False
-
- sys.exit(0 if success else 1)
- if __name__ == "__main__":
- main()
|