base.py 35 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891892893894895896
  1. """
  2. Base trainer classes and interfaces for the Trixy ML training framework.
  3. This module provides the fundamental building blocks for all ML trainers,
  4. including abstract base classes, configuration management, and training state tracking.
  5. """
  6. import os
  7. import time
  8. import logging
  9. import hashlib
  10. from abc import ABC, abstractmethod
  11. from typing import Any, Dict, List, Optional, Union, Callable, Tuple
  12. from dataclasses import dataclass, field
  13. from enum import Enum
  14. import torch
  15. import torch.nn as nn
  16. import torch.optim as optim
  17. from torch.utils.data import DataLoader
  18. import numpy as np
  19. from .validation_manager import ValidationManager, ValidationMode, create_validation_manager
  20. class TrainingState(Enum):
  21. """Training state enumeration for tracking trainer status."""
  22. IDLE = "idle"
  23. PREPARING = "preparing"
  24. TRAINING = "training"
  25. VALIDATING = "validating"
  26. TESTING = "testing"
  27. PAUSED = "paused"
  28. COMPLETED = "completed"
  29. FAILED = "failed"
  30. STOPPING = "stopping"
  31. class ModelFormat(Enum):
  32. """Supported model formats."""
  33. PTH = ".pth"
  34. PT = ".pt"
  35. ONNX = ".onnx"
  36. @dataclass
  37. class TrainerConfig:
  38. """Configuration class for ML trainers."""
  39. # General settings
  40. trainer_name: str = "BaseTrainer"
  41. model_name: str = "model"
  42. output_dir: str = "./models"
  43. data_dir: str = "./trainer/data"
  44. config_dir: str = "./config"
  45. # Training parameters
  46. batch_size: int = 32
  47. learning_rate: float = 0.0001 # Reduced from 0.001 for better stability with metric learning
  48. num_epochs: int = 100
  49. min_epochs: int = 10
  50. max_epochs: int = 1000
  51. early_stopping_patience: int = 10
  52. # Model parameters
  53. model_format: ModelFormat = ModelFormat.PTH
  54. use_mixed_precision: bool = True
  55. gradient_clip_norm: float = 5.0 # Increased from 1.0 for better gradient stability
  56. weight_decay: float = 1e-4
  57. # Data parameters
  58. validation_split: float = 0.2
  59. test_split: float = 0.1
  60. shuffle_data: bool = True
  61. num_workers: int = 4
  62. pin_memory: bool = True
  63. # Audio processing parameters
  64. sample_rate: int = 16000
  65. bit_depth: int = 16
  66. audio_length: float = 1.5 # seconds
  67. n_mels: int = 40
  68. n_fft: int = 512
  69. hop_length: int = 160 # 10ms at 16kHz
  70. win_length: int = 400 # 25ms at 16kHz
  71. # Augmentation parameters
  72. use_augmentation: bool = True
  73. noise_factor: float = 0.1
  74. speed_factor: float = 0.1
  75. pitch_factor: float = 0.1
  76. volume_factor: float = 0.2
  77. # Security settings
  78. use_password_protection: bool = True
  79. password: Optional[str] = None
  80. encryption_algorithm: str = "AES256"
  81. # Logging and monitoring
  82. log_level: str = "INFO"
  83. save_checkpoints: bool = True
  84. checkpoint_interval: int = 10 # epochs
  85. validate_interval: int = 1 # epochs
  86. log_interval: int = 100 # batches
  87. # Enhanced validation settings
  88. use_lightweight_validation: bool = True
  89. lightweight_sample_ratio: float = 0.3
  90. lightweight_max_samples: int = 1000
  91. comprehensive_validation_epochs: Optional[List[int]] = None
  92. # Hardware settings
  93. device: str = "auto" # auto, cpu, cuda, mps
  94. use_distributed: bool = False
  95. num_gpus: int = 1
  96. # Advanced settings
  97. resume_from_checkpoint: bool = False
  98. checkpoint_path: Optional[str] = None
  99. freeze_backbone: bool = False
  100. use_scheduler: bool = True
  101. scheduler_type: str = "cosine" # cosine, step, exponential
  102. # Custom parameters (extensible)
  103. custom_params: Dict[str, Any] = field(default_factory=dict)
  104. def __post_init__(self):
  105. """Validate and process configuration after initialization."""
  106. self._validate_config()
  107. self._process_device()
  108. def _validate_config(self):
  109. """Validate configuration parameters."""
  110. if self.batch_size <= 0:
  111. raise ValueError("batch_size must be positive")
  112. if self.learning_rate <= 0:
  113. raise ValueError("learning_rate must be positive")
  114. if self.num_epochs <= 0:
  115. raise ValueError("num_epochs must be positive")
  116. if self.min_epochs < 0:
  117. raise ValueError("min_epochs must be non-negative")
  118. if self.min_epochs > self.num_epochs:
  119. raise ValueError("min_epochs cannot be greater than num_epochs")
  120. if self.max_epochs < self.num_epochs:
  121. raise ValueError("max_epochs cannot be less than num_epochs")
  122. if self.early_stopping_patience <= 0:
  123. raise ValueError("early_stopping_patience must be positive")
  124. if not 0 < self.validation_split < 1:
  125. raise ValueError("validation_split must be between 0 and 1")
  126. if not 0 <= self.test_split < 1:
  127. raise ValueError("test_split must be between 0 and 1")
  128. if self.validation_split + self.test_split >= 1:
  129. raise ValueError("validation_split + test_split must be less than 1")
  130. def _process_device(self):
  131. """Automatically determine the best available device."""
  132. if self.device == "auto":
  133. if torch.cuda.is_available():
  134. self.device = "cuda"
  135. elif hasattr(torch.backends, 'mps') and torch.backends.mps.is_available():
  136. self.device = "mps"
  137. else:
  138. self.device = "cpu"
  139. def to_dict(self) -> Dict[str, Any]:
  140. """Convert configuration to dictionary."""
  141. from .utils import make_json_serializable
  142. config_dict = {}
  143. for key, value in self.__dict__.items():
  144. if isinstance(value, Enum):
  145. config_dict[key] = value.value
  146. else:
  147. config_dict[key] = make_json_serializable(value)
  148. return config_dict
  149. @classmethod
  150. def from_dict(cls, config_dict: Dict[str, Any]) -> 'TrainerConfig':
  151. """Create configuration from dictionary."""
  152. # Handle enum conversion
  153. if 'model_format' in config_dict:
  154. config_dict['model_format'] = ModelFormat(config_dict['model_format'])
  155. return cls(**config_dict)
  156. class TrainingMetrics:
  157. """Class for tracking and managing training metrics."""
  158. def __init__(self):
  159. self.metrics = {
  160. 'train_loss': [],
  161. 'train_accuracy': [],
  162. 'val_loss': [],
  163. 'val_accuracy': [],
  164. 'learning_rates': [],
  165. 'epoch_times': [],
  166. 'total_time': 0.0
  167. }
  168. self.current_epoch = 0
  169. self.best_val_loss = float('inf')
  170. self.best_val_accuracy = 0.0
  171. self.best_epoch = 0
  172. def update(self, epoch: int, train_loss: float, train_acc: float = None,
  173. val_loss: float = None, val_acc: float = None, lr: float = None,
  174. epoch_time: float = None):
  175. """Update metrics for the current epoch."""
  176. self.current_epoch = epoch
  177. self.metrics['train_loss'].append(train_loss)
  178. if train_acc is not None:
  179. self.metrics['train_accuracy'].append(train_acc)
  180. if val_loss is not None:
  181. self.metrics['val_loss'].append(val_loss)
  182. if val_loss < self.best_val_loss:
  183. self.best_val_loss = val_loss
  184. self.best_epoch = epoch
  185. if val_acc is not None:
  186. self.metrics['val_accuracy'].append(val_acc)
  187. if val_acc > self.best_val_accuracy:
  188. self.best_val_accuracy = val_acc
  189. if lr is not None:
  190. self.metrics['learning_rates'].append(lr)
  191. if epoch_time is not None:
  192. self.metrics['epoch_times'].append(epoch_time)
  193. self.metrics['total_time'] += epoch_time
  194. def get_summary(self) -> Dict[str, Any]:
  195. """Get a summary of training metrics."""
  196. return {
  197. 'current_epoch': self.current_epoch,
  198. 'best_epoch': self.best_epoch,
  199. 'best_val_loss': self.best_val_loss,
  200. 'best_val_accuracy': self.best_val_accuracy,
  201. 'total_training_time': self.metrics['total_time'],
  202. 'avg_epoch_time': np.mean(self.metrics['epoch_times']) if self.metrics['epoch_times'] else 0,
  203. 'final_train_loss': self.metrics['train_loss'][-1] if self.metrics['train_loss'] else None,
  204. 'final_val_loss': self.metrics['val_loss'][-1] if self.metrics['val_loss'] else None
  205. }
  206. class BaseTrainer(ABC):
  207. """
  208. Abstract base class for all ML trainers in the Trixy framework.
  209. This class provides the fundamental structure and common functionality
  210. for all training implementations, including state management, logging,
  211. checkpointing, and validation.
  212. """
  213. def __init__(self, config: TrainerConfig):
  214. """
  215. Initialize the base trainer.
  216. Args:
  217. config: TrainerConfig object containing training parameters
  218. """
  219. self.config = config
  220. self.state = TrainingState.IDLE
  221. self.metrics = TrainingMetrics()
  222. self.logger = self._setup_logger()
  223. # Training components (to be initialized by subclasses)
  224. self.model: Optional[nn.Module] = None
  225. self.optimizer: Optional[optim.Optimizer] = None
  226. self.scheduler: Optional[optim.lr_scheduler._LRScheduler] = None
  227. self.criterion: Optional[nn.Module] = None
  228. self.scaler: Optional[torch.cuda.amp.GradScaler] = None
  229. # Data loaders (to be initialized by subclasses)
  230. self.train_loader: Optional[DataLoader] = None
  231. self.val_loader: Optional[DataLoader] = None
  232. self.test_loader: Optional[DataLoader] = None
  233. # Training state
  234. self.current_epoch = 0
  235. self.global_step = 0
  236. self.best_model_state = None
  237. self.early_stopping_counter = 0
  238. # Validation manager (initialized after model setup)
  239. self.validation_manager: Optional[ValidationManager] = None
  240. # Callbacks
  241. self.callbacks: List[Callable] = []
  242. # Setup mixed precision if requested
  243. if self.config.use_mixed_precision and self.config.device != "cpu":
  244. if self.config.device.startswith("cuda"):
  245. self.scaler = torch.cuda.amp.GradScaler()
  246. else:
  247. # For other devices like MPS, use CPU scaler or disable mixed precision
  248. self.logger.warning(f"Mixed precision not fully supported on {self.config.device}, disabling scaler")
  249. self.scaler = None
  250. self.logger.info(f"Initialized {self.__class__.__name__} with config: {self.config.trainer_name}")
  251. def _setup_logger(self) -> logging.Logger:
  252. """Setup logger for the trainer."""
  253. logger = logging.getLogger(f"trixy.trainer.{self.config.trainer_name}")
  254. logger.setLevel(getattr(logging, self.config.log_level.upper()))
  255. # Prevent duplicate logging by disabling propagation to parent loggers
  256. logger.propagate = False
  257. # CRITICAL: Always clear existing handlers to prevent duplicates
  258. # This handles cases where the logger was created elsewhere or reused
  259. logger.handlers.clear()
  260. # Create console handler
  261. console_handler = logging.StreamHandler()
  262. console_handler.setLevel(logging.INFO)
  263. # Create file handler
  264. os.makedirs(f"{self.config.output_dir}/logs", exist_ok=True)
  265. file_handler = logging.FileHandler(
  266. f"{self.config.output_dir}/logs/{self.config.trainer_name}.log"
  267. )
  268. file_handler.setLevel(logging.DEBUG)
  269. # Create formatter
  270. formatter = logging.Formatter(
  271. '%(asctime)s - %(name)s - %(levelname)s - %(message)s'
  272. )
  273. console_handler.setFormatter(formatter)
  274. file_handler.setFormatter(formatter)
  275. logger.addHandler(console_handler)
  276. logger.addHandler(file_handler)
  277. return logger
  278. @abstractmethod
  279. def prepare_data(self) -> Tuple[DataLoader, DataLoader, DataLoader]:
  280. """
  281. Prepare training, validation, and test data loaders.
  282. Returns:
  283. Tuple of (train_loader, val_loader, test_loader)
  284. """
  285. pass
  286. @abstractmethod
  287. def build_model(self) -> nn.Module:
  288. """
  289. Build and return the model architecture.
  290. Returns:
  291. PyTorch model
  292. """
  293. pass
  294. @abstractmethod
  295. def create_criterion(self) -> nn.Module:
  296. """
  297. Create and return the loss criterion.
  298. Returns:
  299. PyTorch loss function
  300. """
  301. pass
  302. def create_optimizer(self) -> optim.Optimizer:
  303. """Create and return the optimizer."""
  304. return optim.AdamW(
  305. self.model.parameters(),
  306. lr=self.config.learning_rate,
  307. weight_decay=self.config.weight_decay
  308. )
  309. def create_scheduler(self) -> Optional[optim.lr_scheduler._LRScheduler]:
  310. """Create and return the learning rate scheduler."""
  311. if not self.config.use_scheduler:
  312. return None
  313. if self.config.scheduler_type == "cosine":
  314. return optim.lr_scheduler.CosineAnnealingLR(
  315. self.optimizer, T_max=self.config.num_epochs
  316. )
  317. elif self.config.scheduler_type == "step":
  318. return optim.lr_scheduler.StepLR(
  319. self.optimizer, step_size=30, gamma=0.1
  320. )
  321. elif self.config.scheduler_type == "exponential":
  322. return optim.lr_scheduler.ExponentialLR(
  323. self.optimizer, gamma=0.95
  324. )
  325. else:
  326. self.logger.warning(f"Unknown scheduler type: {self.config.scheduler_type}")
  327. return None
  328. def setup_training(self):
  329. """Setup all training components."""
  330. self.state = TrainingState.PREPARING
  331. self.logger.info("Setting up training components...")
  332. # Prepare data
  333. self.train_loader, self.val_loader, self.test_loader = self.prepare_data()
  334. # Build model
  335. self.model = self.build_model()
  336. self.model.to(self.config.device)
  337. # Create training components
  338. self.criterion = self.create_criterion()
  339. # Move criterion to device if it has parameters (custom loss functions)
  340. if hasattr(self.criterion, 'parameters'):
  341. try:
  342. # Check if the criterion has any parameters
  343. params = list(self.criterion.parameters())
  344. if params:
  345. self.criterion.to(self.config.device)
  346. except Exception:
  347. # Fallback: just try to move it if it has parameters method
  348. self.criterion.to(self.config.device)
  349. self.optimizer = self.create_optimizer()
  350. self.scheduler = self.create_scheduler()
  351. # Initialize validation manager
  352. if self.config.use_lightweight_validation:
  353. self.validation_manager = create_validation_manager(
  354. model=self.model,
  355. device=self.config.device,
  356. logger=self.logger,
  357. lightweight_sample_ratio=self.config.lightweight_sample_ratio,
  358. lightweight_max_samples=self.config.lightweight_max_samples
  359. )
  360. self.logger.info("Initialized enhanced validation manager")
  361. # Load checkpoint if requested
  362. if self.config.resume_from_checkpoint and self.config.checkpoint_path:
  363. self.load_checkpoint(self.config.checkpoint_path)
  364. self.logger.info("Training setup completed successfully")
  365. def _compute_loss(self, outputs: torch.Tensor, targets: torch.Tensor, data: torch.Tensor) -> torch.Tensor:
  366. """
  367. Compute loss. Can be overridden by subclasses for custom loss computation.
  368. Args:
  369. outputs: Model outputs
  370. targets: Target labels
  371. data: Input data (for custom loss functions that need access to inputs)
  372. Returns:
  373. Loss value
  374. """
  375. return self.criterion(outputs, targets)
  376. def train_epoch(self) -> Tuple[float, float]:
  377. """
  378. Train for one epoch.
  379. Returns:
  380. Tuple of (average_loss, average_accuracy)
  381. """
  382. self.model.train()
  383. total_loss = 0.0
  384. total_correct = 0
  385. total_samples = 0
  386. valid_batches = 0
  387. skipped_batches = 0
  388. for batch_idx, (data, targets) in enumerate(self.train_loader):
  389. data = data.to(self.config.device)
  390. targets = targets.to(self.config.device)
  391. # Zero gradients
  392. self.optimizer.zero_grad()
  393. batch_skipped = False
  394. # Forward pass with mixed precision
  395. if self.scaler is not None:
  396. with torch.cuda.amp.autocast():
  397. outputs = self.model(data)
  398. loss = self._compute_loss(outputs, targets, data)
  399. # Check if batch should be skipped (loss < 0 indicates skip signal)
  400. # Valid losses are always >= 0, so negative loss is a clear skip indicator
  401. if loss.item() < 0 or torch.isnan(loss) or torch.isinf(loss):
  402. batch_skipped = True
  403. skipped_batches += 1
  404. self.logger.debug(f"Skipping batch {batch_idx} due to invalid loss: {loss.item()}")
  405. continue
  406. # Backward pass
  407. self.scaler.scale(loss).backward()
  408. # Check for NaN/Inf gradients before proceeding
  409. has_nan_grads = False
  410. for param in self.model.parameters():
  411. if param.grad is not None:
  412. if torch.isnan(param.grad).any() or torch.isinf(param.grad).any():
  413. has_nan_grads = True
  414. break
  415. if has_nan_grads:
  416. batch_skipped = True
  417. skipped_batches += 1
  418. self.logger.debug(f"Skipping batch {batch_idx} due to NaN/Inf gradients")
  419. self.optimizer.zero_grad() # Clear the bad gradients
  420. continue
  421. # Gradient clipping with logging
  422. if self.config.gradient_clip_norm > 0:
  423. self.scaler.unscale_(self.optimizer)
  424. grad_norm = torch.nn.utils.clip_grad_norm_(
  425. self.model.parameters(), self.config.gradient_clip_norm
  426. )
  427. # Check if gradient norm is reasonable
  428. if torch.isnan(grad_norm) or torch.isinf(grad_norm):
  429. batch_skipped = True
  430. skipped_batches += 1
  431. self.logger.debug(f"Skipping batch {batch_idx} due to invalid gradient norm: {grad_norm}")
  432. self.optimizer.zero_grad()
  433. continue
  434. # Only step if we have valid gradients
  435. self.scaler.step(self.optimizer)
  436. self.scaler.update()
  437. else:
  438. outputs = self.model(data)
  439. loss = self._compute_loss(outputs, targets, data)
  440. # Check if batch should be skipped (loss < 0 indicates skip signal)
  441. # Valid losses are always >= 0, so negative loss is a clear skip indicator
  442. if loss.item() < 0 or torch.isnan(loss) or torch.isinf(loss):
  443. batch_skipped = True
  444. skipped_batches += 1
  445. self.logger.debug(f"Skipping batch {batch_idx} due to invalid loss: {loss.item()}")
  446. continue
  447. loss.backward()
  448. # Check for NaN/Inf gradients
  449. has_nan_grads = False
  450. for param in self.model.parameters():
  451. if param.grad is not None:
  452. if torch.isnan(param.grad).any() or torch.isinf(param.grad).any():
  453. has_nan_grads = True
  454. break
  455. if has_nan_grads:
  456. batch_skipped = True
  457. skipped_batches += 1
  458. self.logger.debug(f"Skipping batch {batch_idx} due to NaN/Inf gradients")
  459. self.optimizer.zero_grad()
  460. continue
  461. # Gradient clipping with logging
  462. if self.config.gradient_clip_norm > 0:
  463. grad_norm = torch.nn.utils.clip_grad_norm_(
  464. self.model.parameters(), self.config.gradient_clip_norm
  465. )
  466. if torch.isnan(grad_norm) or torch.isinf(grad_norm):
  467. batch_skipped = True
  468. skipped_batches += 1
  469. self.logger.debug(f"Skipping batch {batch_idx} due to invalid gradient norm: {grad_norm}")
  470. self.optimizer.zero_grad()
  471. continue
  472. self.optimizer.step()
  473. # Update metrics only for valid batches
  474. if not batch_skipped:
  475. total_loss += loss.item()
  476. if len(outputs.shape) > 1 and outputs.shape[1] > 1: # Classification
  477. _, predicted = torch.max(outputs.data, 1)
  478. total_correct += (predicted == targets).sum().item()
  479. total_samples += targets.size(0)
  480. valid_batches += 1
  481. self.global_step += 1
  482. # Log progress
  483. if batch_idx % self.config.log_interval == 0:
  484. self.logger.debug(
  485. f"Epoch {self.current_epoch}, Batch {batch_idx}/{len(self.train_loader)}, "
  486. f"Loss: {loss.item():.6f}, Valid batches: {valid_batches}, Skipped: {skipped_batches}"
  487. )
  488. # Log skipped batches summary
  489. if skipped_batches > 0:
  490. self.logger.info(f"Epoch {self.current_epoch}: Skipped {skipped_batches}/{len(self.train_loader)} batches due to NaN/Inf issues")
  491. avg_loss = total_loss / max(valid_batches, 1)
  492. avg_accuracy = total_correct / max(total_samples, 1)
  493. return avg_loss, avg_accuracy
  494. def validate_epoch(self) -> Tuple[float, float]:
  495. """
  496. Validate for one epoch.
  497. Returns:
  498. Tuple of (average_loss, average_accuracy)
  499. """
  500. self.model.eval()
  501. total_loss = 0.0
  502. total_correct = 0
  503. total_samples = 0
  504. with torch.no_grad():
  505. for data, targets in self.val_loader:
  506. data = data.to(self.config.device)
  507. targets = targets.to(self.config.device)
  508. outputs = self.model(data)
  509. loss = self._compute_loss(outputs, targets, data)
  510. total_loss += loss.item()
  511. if len(outputs.shape) > 1 and outputs.shape[1] > 1: # Classification
  512. _, predicted = torch.max(outputs.data, 1)
  513. total_correct += (predicted == targets).sum().item()
  514. total_samples += targets.size(0)
  515. avg_loss = total_loss / len(self.val_loader)
  516. avg_accuracy = total_correct / total_samples if total_samples > 0 else 0.0
  517. return avg_loss, avg_accuracy
  518. def train(self) -> TrainingMetrics:
  519. """
  520. Main training loop.
  521. Returns:
  522. TrainingMetrics object containing training history
  523. """
  524. try:
  525. self.setup_training()
  526. self.state = TrainingState.TRAINING
  527. self.logger.info(f"Starting training for {self.config.num_epochs} epochs")
  528. start_time = time.time()
  529. for epoch in range(self.current_epoch, self.config.num_epochs):
  530. epoch_start_time = time.time()
  531. self.current_epoch = epoch
  532. # Training phase
  533. train_loss, train_acc = self.train_epoch()
  534. # Validation phase
  535. if epoch % self.config.validate_interval == 0:
  536. self.state = TrainingState.VALIDATING
  537. if self.validation_manager is not None:
  538. # Determine validation mode
  539. comprehensive_epochs = self.config.comprehensive_validation_epochs
  540. if comprehensive_epochs is None:
  541. # Default: comprehensive at start, middle, and end
  542. comprehensive_epochs = [0, self.config.num_epochs // 2, self.config.num_epochs - 1]
  543. use_lightweight = self.validation_manager.should_use_lightweight(
  544. epoch=epoch,
  545. total_epochs=self.config.num_epochs,
  546. lightweight_interval=self.config.validate_interval,
  547. comprehensive_epochs=comprehensive_epochs
  548. )
  549. validation_mode = ValidationMode.LIGHTWEIGHT if use_lightweight else ValidationMode.COMPREHENSIVE
  550. validation_results = self.validation_manager.validate(
  551. val_loader=self.val_loader,
  552. criterion=self.criterion,
  553. compute_loss_fn=self._compute_loss,
  554. mode=validation_mode
  555. )
  556. val_loss = validation_results['loss']
  557. val_acc = validation_results['accuracy']
  558. # Log validation mode
  559. self.logger.debug(f"Used {validation_mode.value} validation "
  560. f"(samples: {validation_results.get('samples_used', 'all')})")
  561. else:
  562. # Fallback to original validation
  563. val_loss, val_acc = self.validate_epoch()
  564. self.state = TrainingState.TRAINING
  565. else:
  566. val_loss, val_acc = None, None
  567. # Update learning rate
  568. if self.scheduler is not None:
  569. self.scheduler.step()
  570. # Calculate epoch time
  571. epoch_time = time.time() - epoch_start_time
  572. # Update metrics
  573. current_lr = self.optimizer.param_groups[0]['lr']
  574. self.metrics.update(
  575. epoch=epoch,
  576. train_loss=train_loss,
  577. train_acc=train_acc,
  578. val_loss=val_loss,
  579. val_acc=val_acc,
  580. lr=current_lr,
  581. epoch_time=epoch_time
  582. )
  583. # Log progress
  584. log_msg = f"Epoch {epoch}/{self.config.num_epochs} - "
  585. log_msg += f"Train Loss: {train_loss:.6f}, Train Acc: {train_acc:.4f}"
  586. if val_loss is not None:
  587. log_msg += f", Val Loss: {val_loss:.6f}, Val Acc: {val_acc:.4f}"
  588. log_msg += f", LR: {current_lr:.2e}, Time: {epoch_time:.2f}s"
  589. self.logger.info(log_msg)
  590. # Save checkpoint
  591. if self.config.save_checkpoints and epoch % self.config.checkpoint_interval == 0:
  592. self.save_checkpoint(epoch)
  593. # Early stopping check (only after minimum epochs)
  594. if val_loss is not None:
  595. if val_loss < self.metrics.best_val_loss:
  596. self.best_model_state = self.model.state_dict().copy()
  597. self.early_stopping_counter = 0
  598. else:
  599. # Only increment early stopping counter after min_epochs
  600. if epoch >= self.config.min_epochs:
  601. self.early_stopping_counter += 1
  602. if self.early_stopping_counter >= self.config.early_stopping_patience:
  603. self.logger.info(f"Early stopping triggered after {epoch + 1} epochs "
  604. f"(minimum {self.config.min_epochs} epochs completed)")
  605. break
  606. else:
  607. # Reset counter if we're still below min_epochs
  608. self.early_stopping_counter = 0
  609. # Execute callbacks
  610. for callback in self.callbacks:
  611. callback(self, epoch)
  612. total_time = time.time() - start_time
  613. self.logger.info(f"Training completed in {total_time:.2f} seconds")
  614. # Load best model if available
  615. if self.best_model_state is not None:
  616. self.model.load_state_dict(self.best_model_state)
  617. self.state = TrainingState.COMPLETED
  618. return self.metrics
  619. except Exception as e:
  620. self.state = TrainingState.FAILED
  621. self.logger.error(f"Training failed: {str(e)}")
  622. raise
  623. def test(self) -> Dict[str, float]:
  624. """
  625. Test the trained model.
  626. Returns:
  627. Dictionary containing test metrics
  628. """
  629. self.state = TrainingState.TESTING
  630. self.logger.info("Starting model testing...")
  631. if self.test_loader is None:
  632. raise ValueError("Test loader not available")
  633. self.model.eval()
  634. total_loss = 0.0
  635. total_correct = 0
  636. total_samples = 0
  637. with torch.no_grad():
  638. for data, targets in self.test_loader:
  639. data = data.to(self.config.device)
  640. targets = targets.to(self.config.device)
  641. outputs = self.model(data)
  642. loss = self.criterion(outputs, targets)
  643. total_loss += loss.item()
  644. if len(outputs.shape) > 1 and outputs.shape[1] > 1: # Classification
  645. _, predicted = torch.max(outputs.data, 1)
  646. total_correct += (predicted == targets).sum().item()
  647. total_samples += targets.size(0)
  648. test_loss = total_loss / len(self.test_loader)
  649. test_accuracy = total_correct / total_samples if total_samples > 0 else 0.0
  650. test_results = {
  651. 'test_loss': test_loss,
  652. 'test_accuracy': test_accuracy,
  653. 'total_samples': total_samples
  654. }
  655. self.logger.info(f"Test Results - Loss: {test_loss:.6f}, Accuracy: {test_accuracy:.4f}")
  656. return test_results
  657. def save_checkpoint(self, epoch: int, filepath: Optional[str] = None):
  658. """Save training checkpoint."""
  659. if filepath is None:
  660. os.makedirs(f"{self.config.output_dir}/checkpoints", exist_ok=True)
  661. filepath = f"{self.config.output_dir}/checkpoints/{self.config.model_name}_epoch_{epoch}.pth"
  662. checkpoint = {
  663. 'epoch': epoch,
  664. 'model_state_dict': self.model.state_dict(),
  665. 'optimizer_state_dict': self.optimizer.state_dict(),
  666. 'scheduler_state_dict': self.scheduler.state_dict() if self.scheduler else None,
  667. 'scaler_state_dict': self.scaler.state_dict() if self.scaler else None,
  668. 'metrics': self.metrics,
  669. 'config': self.config.to_dict(),
  670. 'best_model_state': self.best_model_state
  671. }
  672. torch.save(checkpoint, filepath)
  673. self.logger.info(f"Checkpoint saved: {filepath}")
  674. def load_checkpoint(self, filepath: str):
  675. """Load training checkpoint."""
  676. if not os.path.exists(filepath):
  677. raise FileNotFoundError(f"Checkpoint not found: {filepath}")
  678. checkpoint = torch.load(filepath, map_location=self.config.device)
  679. # Load model state
  680. if self.model is not None:
  681. self.model.load_state_dict(checkpoint['model_state_dict'])
  682. # Load optimizer state
  683. if self.optimizer is not None:
  684. self.optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
  685. # Load scheduler state
  686. if self.scheduler is not None and checkpoint.get('scheduler_state_dict'):
  687. self.scheduler.load_state_dict(checkpoint['scheduler_state_dict'])
  688. # Load scaler state
  689. if self.scaler is not None and checkpoint.get('scaler_state_dict'):
  690. self.scaler.load_state_dict(checkpoint['scaler_state_dict'])
  691. # Load training state
  692. self.current_epoch = checkpoint['epoch'] + 1
  693. self.metrics = checkpoint.get('metrics', TrainingMetrics())
  694. self.best_model_state = checkpoint.get('best_model_state')
  695. self.logger.info(f"Checkpoint loaded: {filepath}")
  696. def add_callback(self, callback: Callable):
  697. """Add a callback function to be executed during training."""
  698. self.callbacks.append(callback)
  699. def get_model_summary(self) -> str:
  700. """Get a summary of the model architecture."""
  701. if self.model is None:
  702. return "Model not initialized"
  703. total_params = sum(p.numel() for p in self.model.parameters())
  704. trainable_params = sum(p.numel() for p in self.model.parameters() if p.requires_grad)
  705. summary = f"Model: {self.model.__class__.__name__}\n"
  706. summary += f"Total parameters: {total_params:,}\n"
  707. summary += f"Trainable parameters: {trainable_params:,}\n"
  708. summary += f"Device: {next(self.model.parameters()).device}\n"
  709. return summary
  710. def stop_training(self):
  711. """Stop training gracefully."""
  712. self.state = TrainingState.STOPPING
  713. self.logger.info("Training stop requested")
  714. def pause_training(self):
  715. """Pause training."""
  716. self.state = TrainingState.PAUSED
  717. self.logger.info("Training paused")
  718. def resume_training(self):
  719. """Resume training from paused state."""
  720. if self.state == TrainingState.PAUSED:
  721. self.state = TrainingState.TRAINING
  722. self.logger.info("Training resumed")
  723. def get_training_status(self) -> Dict[str, Any]:
  724. """Get current training status and progress."""
  725. return {
  726. 'state': self.state.value,
  727. 'current_epoch': self.current_epoch,
  728. 'total_epochs': self.config.num_epochs,
  729. 'progress_percentage': (self.current_epoch / self.config.num_epochs) * 100,
  730. 'global_step': self.global_step,
  731. 'metrics_summary': self.metrics.get_summary(),
  732. 'best_epoch': self.metrics.best_epoch,
  733. 'early_stopping_counter': self.early_stopping_counter
  734. }