cli.py 23 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623
  1. """
  2. Command-line interface for the ML Trainer Framework.
  3. This module provides a comprehensive CLI for training wakeword detection,
  4. voice recognition, and other models with support for configuration files
  5. and command-line argument overrides.
  6. """
  7. import os
  8. import sys
  9. import json
  10. import argparse
  11. import logging
  12. from pathlib import Path
  13. from typing import Dict, Any, Optional, List
  14. from .config import ConfigManager, TrainerConfig, ModelFormat
  15. from .wakeword.trainer import WakewordTrainer, create_wakeword_trainer_config
  16. from .voice_recognition.trainer import VoiceRecognitionTrainer, create_voice_recognition_trainer_config
  17. from .base import TrainerConfig as BaseTrainerConfig
  18. from dataclasses import fields
  19. # Note: setup_logging function was moved to TrainerLogger class in utils
  20. def setup_cli_logging(level: str = "INFO", log_file: Optional[str] = None) -> logging.Logger:
  21. """Setup logging for CLI operations."""
  22. logger = logging.getLogger("trainer_cli")
  23. logger.setLevel(getattr(logging, level.upper()))
  24. # Prevent duplicate logging by disabling propagation to root logger
  25. logger.propagate = False
  26. # CRITICAL: Always clear existing handlers to prevent duplicates
  27. logger.handlers.clear()
  28. # Console handler
  29. console_handler = logging.StreamHandler()
  30. console_formatter = logging.Formatter('%(asctime)s - %(name)s - %(levelname)s - %(message)s')
  31. console_handler.setFormatter(console_formatter)
  32. logger.addHandler(console_handler)
  33. # File handler if requested
  34. if log_file:
  35. file_handler = logging.FileHandler(log_file)
  36. file_handler.setFormatter(console_formatter)
  37. logger.addHandler(file_handler)
  38. return logger
  39. def load_config_file(config_path: str) -> Dict[str, Any]:
  40. """Load configuration from JSON file."""
  41. try:
  42. with open(config_path, 'r') as f:
  43. return json.load(f)
  44. except FileNotFoundError:
  45. raise FileNotFoundError(f"Configuration file not found: {config_path}")
  46. except json.JSONDecodeError as e:
  47. raise ValueError(f"Invalid JSON in configuration file: {e}")
  48. def create_parser() -> argparse.ArgumentParser:
  49. """Create the main argument parser."""
  50. parser = argparse.ArgumentParser(
  51. description="Trixy ML Trainer Framework",
  52. formatter_class=argparse.RawDescriptionHelpFormatter,
  53. epilog="""
  54. Examples:
  55. # Train wakeword detection model
  56. python -m trainer.cli train wakeword --config ./config/wakeword.json
  57. # Train voice recognition model with overrides
  58. python -m trainer.cli train voice-recognition --model-type ecapa_tdnn --batch-size 64
  59. # List available models
  60. python -m trainer.cli list-models
  61. # Validate model
  62. python -m trainer.cli validate ./models/my_model.pth --test-data ./data/test
  63. """
  64. )
  65. # Global arguments
  66. parser.add_argument(
  67. "--config", "-c", type=str,
  68. help="Path to configuration file (JSON)"
  69. )
  70. parser.add_argument(
  71. "--log-level", choices=["DEBUG", "INFO", "WARNING", "ERROR"],
  72. default="INFO", help="Logging level"
  73. )
  74. parser.add_argument(
  75. "--log-file", type=str,
  76. help="Path to log file (default: console only)"
  77. )
  78. parser.add_argument(
  79. "--device", choices=["auto", "cpu", "cuda"],
  80. default="auto", help="Device to use for training/inference"
  81. )
  82. # Subcommands
  83. subparsers = parser.add_subparsers(dest="command", help="Available commands")
  84. # Train command
  85. train_parser = subparsers.add_parser("train", help="Train a model")
  86. train_subparsers = train_parser.add_subparsers(dest="train_type", help="Type of model to train")
  87. # Wakeword training
  88. wakeword_parser = train_subparsers.add_parser("wakeword", help="Train wakeword detection model")
  89. add_wakeword_arguments(wakeword_parser)
  90. # Voice recognition training
  91. voice_parser = train_subparsers.add_parser("voice-recognition", help="Train voice recognition model")
  92. add_voice_recognition_arguments(voice_parser)
  93. # Validate command
  94. validate_parser = subparsers.add_parser("validate", help="Validate a trained model")
  95. add_validation_arguments(validate_parser)
  96. # List models command
  97. list_parser = subparsers.add_parser("list-models", help="List available trained models")
  98. add_list_arguments(list_parser)
  99. # Convert model command
  100. convert_parser = subparsers.add_parser("convert", help="Convert model between formats")
  101. add_convert_arguments(convert_parser)
  102. return parser
  103. def add_wakeword_arguments(parser: argparse.ArgumentParser):
  104. """Add wakeword-specific arguments."""
  105. # Model configuration
  106. parser.add_argument("--model-name", type=str, default="wakeword_model",
  107. help="Name of the model")
  108. parser.add_argument("--model-type", choices=["standard", "improved", "lightweight"],
  109. default="standard", help="Type of RepCNN model")
  110. parser.add_argument("--model-format", choices=[fmt.value for fmt in ModelFormat],
  111. default=ModelFormat.PTH.value, help="Model output format")
  112. # Data configuration
  113. parser.add_argument("--data-dir", type=str, required=True,
  114. help="Directory containing training data")
  115. parser.add_argument("--output-dir", type=str, default="./models/wakeword",
  116. help="Output directory for trained model")
  117. # Training parameters
  118. parser.add_argument("--batch-size", type=int, default=64,
  119. help="Training batch size")
  120. parser.add_argument("--learning-rate", type=float, default=0.001,
  121. help="Learning rate")
  122. parser.add_argument("--num-epochs", type=int, default=100,
  123. help="Number of training epochs")
  124. parser.add_argument("--min-epochs", type=int, default=20,
  125. help="Minimum number of training epochs")
  126. parser.add_argument("--early-stopping-patience", type=int, default=15,
  127. help="Early stopping patience")
  128. # Audio parameters
  129. parser.add_argument("--sample-rate", type=int, default=16000,
  130. help="Audio sample rate")
  131. parser.add_argument("--audio-length", type=float, default=1.5,
  132. help="Target audio length in seconds")
  133. parser.add_argument("--n-mels", type=int, default=40,
  134. help="Number of mel filterbank features")
  135. # Augmentation
  136. parser.add_argument("--use-augmentation", action="store_true",
  137. help="Enable data augmentation")
  138. parser.add_argument("--noise-factor", type=float, default=0.1,
  139. help="Noise augmentation factor")
  140. # Model specific
  141. parser.add_argument("--use-focal-loss", action="store_true",
  142. help="Use focal loss for imbalanced data")
  143. parser.add_argument("--detection-threshold", type=float, default=0.5,
  144. help="Detection threshold for wakeword")
  145. # Security
  146. parser.add_argument("--use-password-protection", action="store_true",
  147. help="Enable password protection for model")
  148. parser.add_argument("--password", type=str,
  149. help="Password for model protection")
  150. def add_voice_recognition_arguments(parser: argparse.ArgumentParser):
  151. """Add voice recognition specific arguments."""
  152. # Model configuration
  153. parser.add_argument("--model-name", type=str, default="voice_recognition_model",
  154. help="Name of the model")
  155. parser.add_argument("--model-type", choices=["ecapa_tdnn", "titanet_s", "speakernet_m"],
  156. default="ecapa_tdnn", help="Type of voice recognition model")
  157. parser.add_argument("--model-format", choices=[fmt.value for fmt in ModelFormat],
  158. default=ModelFormat.PTH.value, help="Model output format")
  159. # Data configuration
  160. parser.add_argument("--data-dir", type=str, required=True,
  161. help="Directory containing speaker data")
  162. parser.add_argument("--output-dir", type=str, default="./models/voice_recognition",
  163. help="Output directory for trained model")
  164. # Training parameters
  165. parser.add_argument("--batch-size", type=int, default=32,
  166. help="Training batch size")
  167. parser.add_argument("--learning-rate", type=float, default=0.001,
  168. help="Learning rate")
  169. parser.add_argument("--num-epochs", type=int, default=200,
  170. help="Number of training epochs")
  171. parser.add_argument("--min-epochs", type=int, default=30,
  172. help="Minimum number of training epochs")
  173. parser.add_argument("--early-stopping-patience", type=int, default=20,
  174. help="Early stopping patience")
  175. # Audio parameters
  176. parser.add_argument("--sample-rate", type=int, default=16000,
  177. help="Audio sample rate")
  178. parser.add_argument("--audio-length", type=float, default=3.0,
  179. help="Target audio length in seconds")
  180. parser.add_argument("--n-mels", type=int, default=40,
  181. help="Number of mel filterbank features")
  182. # Loss function
  183. parser.add_argument("--loss-type",
  184. choices=["angular_margin", "ge2e", "contrastive", "triplet", "cross_entropy"],
  185. default="angular_margin", help="Loss function type")
  186. parser.add_argument("--embedding-dim", type=int, default=192,
  187. help="Embedding dimension")
  188. # Augmentation
  189. parser.add_argument("--use-augmentation", action="store_true",
  190. help="Enable data augmentation")
  191. # Security
  192. parser.add_argument("--use-password-protection", action="store_true",
  193. help="Enable password protection for model")
  194. parser.add_argument("--password", type=str,
  195. help="Password for model protection")
  196. def add_validation_arguments(parser: argparse.ArgumentParser):
  197. """Add validation arguments."""
  198. parser.add_argument("model_path", type=str,
  199. help="Path to model file to validate")
  200. parser.add_argument("--test-data", type=str, required=True,
  201. help="Path to test data directory")
  202. parser.add_argument("--batch-size", type=int, default=32,
  203. help="Validation batch size")
  204. parser.add_argument("--output-dir", type=str,
  205. help="Output directory for validation results")
  206. def add_list_arguments(parser: argparse.ArgumentParser):
  207. """Add list models arguments."""
  208. parser.add_argument("--models-dir", type=str, default="./models",
  209. help="Directory to search for models")
  210. parser.add_argument("--model-type", choices=["all", "wakeword", "voice-recognition"],
  211. default="all", help="Type of models to list")
  212. def add_convert_arguments(parser: argparse.ArgumentParser):
  213. """Add model conversion arguments."""
  214. parser.add_argument("input_model", type=str,
  215. help="Path to input model file")
  216. parser.add_argument("output_model", type=str,
  217. help="Path for output model file")
  218. parser.add_argument("--target-format", choices=[fmt.value for fmt in ModelFormat],
  219. required=True, help="Target model format")
  220. parser.add_argument("--password", type=str,
  221. help="Password for password-protected models")
  222. def merge_config_and_args(config_dict: Dict[str, Any], args: argparse.Namespace) -> Dict[str, Any]:
  223. """Merge configuration file and command-line arguments."""
  224. # Command-line arguments override config file
  225. merged = config_dict.copy()
  226. # Convert args to dict and filter None values and control arguments
  227. # These arguments are for CLI control and should not be passed to TrainerConfig
  228. cli_control_args = {'command', 'train_type', 'config', 'log_level', 'log_file'}
  229. args_dict = {k: v for k, v in vars(args).items()
  230. if v is not None and k not in cli_control_args}
  231. # Merge, giving priority to command-line arguments
  232. merged.update(args_dict)
  233. return merged
  234. def organize_voice_recognition_params(merged_config: Dict[str, Any]) -> Dict[str, Any]:
  235. """
  236. Organize parameters for voice recognition trainer.
  237. Separates base TrainerConfig parameters from model-specific parameters,
  238. putting model-specific parameters into custom_params.
  239. """
  240. # Get valid TrainerConfig field names
  241. trainer_config_fields = {f.name for f in fields(BaseTrainerConfig)}
  242. # Parameters that should go to create_voice_recognition_trainer_config directly
  243. direct_params = {
  244. 'model_name', 'model_type', 'loss_type', 'data_dir'
  245. }
  246. # Model-specific parameters that go into custom_params
  247. model_specific_params = {
  248. 'embedding_dim', 'model_format', 'detection_threshold',
  249. 'use_focal_loss', 'use_password_protection', 'password'
  250. }
  251. # Organize parameters
  252. organized = {}
  253. custom_params = {}
  254. for key, value in merged_config.items():
  255. # Convert hyphens to underscores
  256. clean_key = key.replace('-', '_')
  257. if clean_key in direct_params:
  258. # Parameters that go directly to the creation function
  259. organized[clean_key] = value
  260. elif clean_key in model_specific_params:
  261. # Model-specific parameters go into custom_params
  262. custom_params[clean_key] = value
  263. elif clean_key in trainer_config_fields:
  264. # Valid TrainerConfig parameters
  265. organized[clean_key] = value
  266. else:
  267. # Unknown parameters go into custom_params for safety
  268. custom_params[clean_key] = value
  269. # Add custom_params if any
  270. if custom_params:
  271. organized['custom_params'] = custom_params
  272. return organized
  273. def organize_wakeword_params(merged_config: Dict[str, Any]) -> Dict[str, Any]:
  274. """
  275. Organize parameters for wakeword trainer.
  276. Separates base TrainerConfig parameters from model-specific parameters,
  277. putting model-specific parameters into custom_params.
  278. """
  279. # Get valid TrainerConfig field names
  280. trainer_config_fields = {f.name for f in fields(BaseTrainerConfig)}
  281. # Parameters that should go to create_wakeword_trainer_config directly
  282. direct_params = {
  283. 'model_name', 'model_type', 'data_dir'
  284. }
  285. # Model-specific parameters that go into custom_params
  286. model_specific_params = {
  287. 'model_format', 'detection_threshold', 'use_focal_loss',
  288. 'use_password_protection', 'password', 'wakeword_classes'
  289. }
  290. # Organize parameters
  291. organized = {}
  292. custom_params = {}
  293. for key, value in merged_config.items():
  294. # Convert hyphens to underscores
  295. clean_key = key.replace('-', '_')
  296. if clean_key in direct_params:
  297. # Parameters that go directly to the creation function
  298. organized[clean_key] = value
  299. elif clean_key in model_specific_params:
  300. # Model-specific parameters go into custom_params
  301. custom_params[clean_key] = value
  302. elif clean_key in trainer_config_fields:
  303. # Valid TrainerConfig parameters
  304. organized[clean_key] = value
  305. else:
  306. # Unknown parameters go into custom_params for safety
  307. custom_params[clean_key] = value
  308. # Add custom_params if any
  309. if custom_params:
  310. organized['custom_params'] = custom_params
  311. return organized
  312. def train_wakeword(args: argparse.Namespace, logger: logging.Logger):
  313. """Train wakeword detection model."""
  314. logger.info("Starting wakeword detection training...")
  315. # Load configuration if provided
  316. config_dict = {}
  317. if args.config:
  318. config_dict = load_config_file(args.config)
  319. # Merge with command-line arguments
  320. merged_config = merge_config_and_args(config_dict, args)
  321. # Separate TrainerConfig parameters from model-specific parameters
  322. trainer_config_params = organize_wakeword_params(merged_config)
  323. # Create trainer configuration
  324. try:
  325. trainer_config = create_wakeword_trainer_config(**trainer_config_params)
  326. except Exception as e:
  327. logger.error(f"Failed to create trainer configuration: {e}")
  328. logger.error(f"Parameters passed: {list(trainer_config_params.keys())}")
  329. return False
  330. # Create and run trainer
  331. try:
  332. trainer = WakewordTrainer(trainer_config)
  333. results = trainer.train()
  334. logger.info("Wakeword training completed successfully!")
  335. logger.info(f"Model saved to: {trainer_config.output_dir}/{trainer_config.model_name}")
  336. return True
  337. except Exception as e:
  338. logger.error(f"Training failed: {e}")
  339. return False
  340. def train_voice_recognition(args: argparse.Namespace, logger: logging.Logger):
  341. """Train voice recognition model."""
  342. logger.info("Starting voice recognition training...")
  343. # Load configuration if provided
  344. config_dict = {}
  345. if args.config:
  346. config_dict = load_config_file(args.config)
  347. # Merge with command-line arguments
  348. merged_config = merge_config_and_args(config_dict, args)
  349. # Separate TrainerConfig parameters from model-specific parameters
  350. trainer_config_params = organize_voice_recognition_params(merged_config)
  351. # Create trainer configuration
  352. try:
  353. trainer_config = create_voice_recognition_trainer_config(**trainer_config_params)
  354. except Exception as e:
  355. logger.error(f"Failed to create trainer configuration: {e}")
  356. logger.error(f"Parameters passed: {list(trainer_config_params.keys())}")
  357. return False
  358. # Create and run trainer
  359. try:
  360. trainer = VoiceRecognitionTrainer(trainer_config)
  361. results = trainer.train()
  362. logger.info("Voice recognition training completed successfully!")
  363. logger.info(f"Model saved to: {trainer_config.output_dir}/{trainer_config.model_name}")
  364. return True
  365. except Exception as e:
  366. logger.error(f"Training failed: {e}")
  367. return False
  368. def validate_model(args: argparse.Namespace, logger: logging.Logger):
  369. """Validate a trained model."""
  370. logger.info(f"Validating model: {args.model_path}")
  371. try:
  372. from .model_formats import ModelFormatManager
  373. from .validation import ModelValidator, ValidationConfig
  374. # Load model
  375. format_manager = ModelFormatManager(logger)
  376. model, metadata = format_manager.load_model(args.model_path, args.password)
  377. if model is None:
  378. logger.error("Failed to load model")
  379. return False
  380. # Create validator
  381. validator = ModelValidator(model, args.device, logger)
  382. # Load test data (simplified - would need proper data loading)
  383. logger.warning("Test data loading not fully implemented in CLI")
  384. logger.info("Model validation completed")
  385. return True
  386. except Exception as e:
  387. logger.error(f"Validation failed: {e}")
  388. return False
  389. def list_models(args: argparse.Namespace, logger: logging.Logger):
  390. """List available trained models."""
  391. models_dir = Path(args.models_dir)
  392. if not models_dir.exists():
  393. logger.warning(f"Models directory not found: {models_dir}")
  394. return True
  395. logger.info(f"Searching for models in: {models_dir}")
  396. # Find model files
  397. model_extensions = ['.pth', '.pt', '.onnx']
  398. models_found = []
  399. for ext in model_extensions:
  400. models_found.extend(list(models_dir.rglob(f"*{ext}")))
  401. if not models_found:
  402. logger.info("No models found")
  403. return True
  404. # Display models
  405. print(f"\nFound {len(models_found)} models:")
  406. print("-" * 60)
  407. for model_path in sorted(models_found):
  408. rel_path = model_path.relative_to(models_dir)
  409. model_type = "Unknown"
  410. # Try to determine model type from path
  411. if "wakeword" in str(rel_path).lower():
  412. model_type = "Wakeword"
  413. elif "voice" in str(rel_path).lower() or "speaker" in str(rel_path).lower():
  414. model_type = "Voice Recognition"
  415. size_mb = model_path.stat().st_size / (1024 * 1024)
  416. print(f"📁 {rel_path}")
  417. print(f" Type: {model_type}")
  418. print(f" Size: {size_mb:.1f} MB")
  419. print(f" Format: {model_path.suffix}")
  420. print()
  421. return True
  422. def convert_model(args: argparse.Namespace, logger: logging.Logger):
  423. """Convert model between formats."""
  424. logger.info(f"Converting model from {args.input_model} to {args.output_model}")
  425. try:
  426. from .model_formats import ModelFormatManager
  427. format_manager = ModelFormatManager(logger)
  428. # Load model
  429. model, metadata = format_manager.load_model(args.input_model, args.password)
  430. if model is None:
  431. logger.error("Failed to load input model")
  432. return False
  433. # Save in target format
  434. success = format_manager.save_model(
  435. model, args.output_model, metadata,
  436. password=args.password if args.password else None
  437. )
  438. if success:
  439. logger.info(f"Model converted successfully to: {args.output_model}")
  440. return True
  441. else:
  442. logger.error("Model conversion failed")
  443. return False
  444. except Exception as e:
  445. logger.error(f"Conversion failed: {e}")
  446. return False
  447. def main():
  448. """Main CLI entry point."""
  449. parser = create_parser()
  450. args = parser.parse_args()
  451. # Setup logging
  452. logger = setup_cli_logging(args.log_level, args.log_file)
  453. # Handle device selection
  454. if args.device == "auto":
  455. import torch
  456. device = "cuda" if torch.cuda.is_available() else "cpu"
  457. args.device = device
  458. logger.info(f"Auto-selected device: {device}")
  459. success = False
  460. try:
  461. if args.command == "train":
  462. if args.train_type == "wakeword":
  463. success = train_wakeword(args, logger)
  464. elif args.train_type == "voice-recognition":
  465. success = train_voice_recognition(args, logger)
  466. else:
  467. logger.error("Please specify training type: wakeword or voice-recognition")
  468. elif args.command == "validate":
  469. success = validate_model(args, logger)
  470. elif args.command == "list-models":
  471. success = list_models(args, logger)
  472. elif args.command == "convert":
  473. success = convert_model(args, logger)
  474. else:
  475. parser.print_help()
  476. success = True
  477. except KeyboardInterrupt:
  478. logger.info("Operation cancelled by user")
  479. success = True
  480. except Exception as e:
  481. logger.error(f"Unexpected error: {e}")
  482. success = False
  483. sys.exit(0 if success else 1)
  484. if __name__ == "__main__":
  485. main()