| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891892893894895896 |
- """
- Base trainer classes and interfaces for the Trixy ML training framework.
- This module provides the fundamental building blocks for all ML trainers,
- including abstract base classes, configuration management, and training state tracking.
- """
- import os
- import time
- import logging
- import hashlib
- from abc import ABC, abstractmethod
- from typing import Any, Dict, List, Optional, Union, Callable, Tuple
- from dataclasses import dataclass, field
- from enum import Enum
- import torch
- import torch.nn as nn
- import torch.optim as optim
- from torch.utils.data import DataLoader
- import numpy as np
- from .validation_manager import ValidationManager, ValidationMode, create_validation_manager
- class TrainingState(Enum):
- """Training state enumeration for tracking trainer status."""
- IDLE = "idle"
- PREPARING = "preparing"
- TRAINING = "training"
- VALIDATING = "validating"
- TESTING = "testing"
- PAUSED = "paused"
- COMPLETED = "completed"
- FAILED = "failed"
- STOPPING = "stopping"
- class ModelFormat(Enum):
- """Supported model formats."""
- PTH = ".pth"
- PT = ".pt"
- ONNX = ".onnx"
- @dataclass
- class TrainerConfig:
- """Configuration class for ML trainers."""
-
- # General settings
- trainer_name: str = "BaseTrainer"
- model_name: str = "model"
- output_dir: str = "./models"
- data_dir: str = "./trainer/data"
- config_dir: str = "./config"
-
- # Training parameters
- batch_size: int = 32
- learning_rate: float = 0.0001 # Reduced from 0.001 for better stability with metric learning
- num_epochs: int = 100
- min_epochs: int = 10
- max_epochs: int = 1000
- early_stopping_patience: int = 10
- # Model parameters
- model_format: ModelFormat = ModelFormat.PTH
- use_mixed_precision: bool = True
- gradient_clip_norm: float = 5.0 # Increased from 1.0 for better gradient stability
- weight_decay: float = 1e-4
-
- # Data parameters
- validation_split: float = 0.2
- test_split: float = 0.1
- shuffle_data: bool = True
- num_workers: int = 4
- pin_memory: bool = True
-
- # Audio processing parameters
- sample_rate: int = 16000
- bit_depth: int = 16
- audio_length: float = 1.5 # seconds
- n_mels: int = 40
- n_fft: int = 512
- hop_length: int = 160 # 10ms at 16kHz
- win_length: int = 400 # 25ms at 16kHz
-
- # Augmentation parameters
- use_augmentation: bool = True
- noise_factor: float = 0.1
- speed_factor: float = 0.1
- pitch_factor: float = 0.1
- volume_factor: float = 0.2
-
- # Security settings
- use_password_protection: bool = True
- password: Optional[str] = None
- encryption_algorithm: str = "AES256"
-
- # Logging and monitoring
- log_level: str = "INFO"
- save_checkpoints: bool = True
- checkpoint_interval: int = 10 # epochs
- validate_interval: int = 1 # epochs
- log_interval: int = 100 # batches
-
- # Enhanced validation settings
- use_lightweight_validation: bool = True
- lightweight_sample_ratio: float = 0.3
- lightweight_max_samples: int = 1000
- comprehensive_validation_epochs: Optional[List[int]] = None
-
- # Hardware settings
- device: str = "auto" # auto, cpu, cuda, mps
- use_distributed: bool = False
- num_gpus: int = 1
-
- # Advanced settings
- resume_from_checkpoint: bool = False
- checkpoint_path: Optional[str] = None
- freeze_backbone: bool = False
- use_scheduler: bool = True
- scheduler_type: str = "cosine" # cosine, step, exponential
-
- # Custom parameters (extensible)
- custom_params: Dict[str, Any] = field(default_factory=dict)
-
- def __post_init__(self):
- """Validate and process configuration after initialization."""
- self._validate_config()
- self._process_device()
-
- def _validate_config(self):
- """Validate configuration parameters."""
- if self.batch_size <= 0:
- raise ValueError("batch_size must be positive")
- if self.learning_rate <= 0:
- raise ValueError("learning_rate must be positive")
- if self.num_epochs <= 0:
- raise ValueError("num_epochs must be positive")
- if self.min_epochs < 0:
- raise ValueError("min_epochs must be non-negative")
- if self.min_epochs > self.num_epochs:
- raise ValueError("min_epochs cannot be greater than num_epochs")
- if self.max_epochs < self.num_epochs:
- raise ValueError("max_epochs cannot be less than num_epochs")
- if self.early_stopping_patience <= 0:
- raise ValueError("early_stopping_patience must be positive")
- if not 0 < self.validation_split < 1:
- raise ValueError("validation_split must be between 0 and 1")
- if not 0 <= self.test_split < 1:
- raise ValueError("test_split must be between 0 and 1")
- if self.validation_split + self.test_split >= 1:
- raise ValueError("validation_split + test_split must be less than 1")
-
- def _process_device(self):
- """Automatically determine the best available device."""
- if self.device == "auto":
- if torch.cuda.is_available():
- self.device = "cuda"
- elif hasattr(torch.backends, 'mps') and torch.backends.mps.is_available():
- self.device = "mps"
- else:
- self.device = "cpu"
-
- def to_dict(self) -> Dict[str, Any]:
- """Convert configuration to dictionary."""
- from .utils import make_json_serializable
-
- config_dict = {}
- for key, value in self.__dict__.items():
- if isinstance(value, Enum):
- config_dict[key] = value.value
- else:
- config_dict[key] = make_json_serializable(value)
- return config_dict
-
- @classmethod
- def from_dict(cls, config_dict: Dict[str, Any]) -> 'TrainerConfig':
- """Create configuration from dictionary."""
- # Handle enum conversion
- if 'model_format' in config_dict:
- config_dict['model_format'] = ModelFormat(config_dict['model_format'])
-
- return cls(**config_dict)
- class TrainingMetrics:
- """Class for tracking and managing training metrics."""
-
- def __init__(self):
- self.metrics = {
- 'train_loss': [],
- 'train_accuracy': [],
- 'val_loss': [],
- 'val_accuracy': [],
- 'learning_rates': [],
- 'epoch_times': [],
- 'total_time': 0.0
- }
- self.current_epoch = 0
- self.best_val_loss = float('inf')
- self.best_val_accuracy = 0.0
- self.best_epoch = 0
-
- def update(self, epoch: int, train_loss: float, train_acc: float = None,
- val_loss: float = None, val_acc: float = None, lr: float = None,
- epoch_time: float = None):
- """Update metrics for the current epoch."""
- self.current_epoch = epoch
- self.metrics['train_loss'].append(train_loss)
-
- if train_acc is not None:
- self.metrics['train_accuracy'].append(train_acc)
- if val_loss is not None:
- self.metrics['val_loss'].append(val_loss)
- if val_loss < self.best_val_loss:
- self.best_val_loss = val_loss
- self.best_epoch = epoch
- if val_acc is not None:
- self.metrics['val_accuracy'].append(val_acc)
- if val_acc > self.best_val_accuracy:
- self.best_val_accuracy = val_acc
- if lr is not None:
- self.metrics['learning_rates'].append(lr)
- if epoch_time is not None:
- self.metrics['epoch_times'].append(epoch_time)
- self.metrics['total_time'] += epoch_time
-
- def get_summary(self) -> Dict[str, Any]:
- """Get a summary of training metrics."""
- return {
- 'current_epoch': self.current_epoch,
- 'best_epoch': self.best_epoch,
- 'best_val_loss': self.best_val_loss,
- 'best_val_accuracy': self.best_val_accuracy,
- 'total_training_time': self.metrics['total_time'],
- 'avg_epoch_time': np.mean(self.metrics['epoch_times']) if self.metrics['epoch_times'] else 0,
- 'final_train_loss': self.metrics['train_loss'][-1] if self.metrics['train_loss'] else None,
- 'final_val_loss': self.metrics['val_loss'][-1] if self.metrics['val_loss'] else None
- }
- class BaseTrainer(ABC):
- """
- Abstract base class for all ML trainers in the Trixy framework.
-
- This class provides the fundamental structure and common functionality
- for all training implementations, including state management, logging,
- checkpointing, and validation.
- """
-
- def __init__(self, config: TrainerConfig):
- """
- Initialize the base trainer.
-
- Args:
- config: TrainerConfig object containing training parameters
- """
- self.config = config
- self.state = TrainingState.IDLE
- self.metrics = TrainingMetrics()
- self.logger = self._setup_logger()
-
- # Training components (to be initialized by subclasses)
- self.model: Optional[nn.Module] = None
- self.optimizer: Optional[optim.Optimizer] = None
- self.scheduler: Optional[optim.lr_scheduler._LRScheduler] = None
- self.criterion: Optional[nn.Module] = None
- self.scaler: Optional[torch.cuda.amp.GradScaler] = None
-
- # Data loaders (to be initialized by subclasses)
- self.train_loader: Optional[DataLoader] = None
- self.val_loader: Optional[DataLoader] = None
- self.test_loader: Optional[DataLoader] = None
-
- # Training state
- self.current_epoch = 0
- self.global_step = 0
- self.best_model_state = None
- self.early_stopping_counter = 0
-
- # Validation manager (initialized after model setup)
- self.validation_manager: Optional[ValidationManager] = None
-
- # Callbacks
- self.callbacks: List[Callable] = []
-
- # Setup mixed precision if requested
- if self.config.use_mixed_precision and self.config.device != "cpu":
- if self.config.device.startswith("cuda"):
- self.scaler = torch.cuda.amp.GradScaler()
- else:
- # For other devices like MPS, use CPU scaler or disable mixed precision
- self.logger.warning(f"Mixed precision not fully supported on {self.config.device}, disabling scaler")
- self.scaler = None
-
- self.logger.info(f"Initialized {self.__class__.__name__} with config: {self.config.trainer_name}")
-
- def _setup_logger(self) -> logging.Logger:
- """Setup logger for the trainer."""
- logger = logging.getLogger(f"trixy.trainer.{self.config.trainer_name}")
- logger.setLevel(getattr(logging, self.config.log_level.upper()))
- # Prevent duplicate logging by disabling propagation to parent loggers
- logger.propagate = False
- # CRITICAL: Always clear existing handlers to prevent duplicates
- # This handles cases where the logger was created elsewhere or reused
- logger.handlers.clear()
- # Create console handler
- console_handler = logging.StreamHandler()
- console_handler.setLevel(logging.INFO)
- # Create file handler
- os.makedirs(f"{self.config.output_dir}/logs", exist_ok=True)
- file_handler = logging.FileHandler(
- f"{self.config.output_dir}/logs/{self.config.trainer_name}.log"
- )
- file_handler.setLevel(logging.DEBUG)
- # Create formatter
- formatter = logging.Formatter(
- '%(asctime)s - %(name)s - %(levelname)s - %(message)s'
- )
- console_handler.setFormatter(formatter)
- file_handler.setFormatter(formatter)
- logger.addHandler(console_handler)
- logger.addHandler(file_handler)
- return logger
-
- @abstractmethod
- def prepare_data(self) -> Tuple[DataLoader, DataLoader, DataLoader]:
- """
- Prepare training, validation, and test data loaders.
-
- Returns:
- Tuple of (train_loader, val_loader, test_loader)
- """
- pass
-
- @abstractmethod
- def build_model(self) -> nn.Module:
- """
- Build and return the model architecture.
-
- Returns:
- PyTorch model
- """
- pass
-
- @abstractmethod
- def create_criterion(self) -> nn.Module:
- """
- Create and return the loss criterion.
-
- Returns:
- PyTorch loss function
- """
- pass
-
- def create_optimizer(self) -> optim.Optimizer:
- """Create and return the optimizer."""
- return optim.AdamW(
- self.model.parameters(),
- lr=self.config.learning_rate,
- weight_decay=self.config.weight_decay
- )
-
- def create_scheduler(self) -> Optional[optim.lr_scheduler._LRScheduler]:
- """Create and return the learning rate scheduler."""
- if not self.config.use_scheduler:
- return None
-
- if self.config.scheduler_type == "cosine":
- return optim.lr_scheduler.CosineAnnealingLR(
- self.optimizer, T_max=self.config.num_epochs
- )
- elif self.config.scheduler_type == "step":
- return optim.lr_scheduler.StepLR(
- self.optimizer, step_size=30, gamma=0.1
- )
- elif self.config.scheduler_type == "exponential":
- return optim.lr_scheduler.ExponentialLR(
- self.optimizer, gamma=0.95
- )
- else:
- self.logger.warning(f"Unknown scheduler type: {self.config.scheduler_type}")
- return None
-
- def setup_training(self):
- """Setup all training components."""
- self.state = TrainingState.PREPARING
- self.logger.info("Setting up training components...")
-
- # Prepare data
- self.train_loader, self.val_loader, self.test_loader = self.prepare_data()
-
- # Build model
- self.model = self.build_model()
- self.model.to(self.config.device)
-
- # Create training components
- self.criterion = self.create_criterion()
- # Move criterion to device if it has parameters (custom loss functions)
- if hasattr(self.criterion, 'parameters'):
- try:
- # Check if the criterion has any parameters
- params = list(self.criterion.parameters())
- if params:
- self.criterion.to(self.config.device)
- except Exception:
- # Fallback: just try to move it if it has parameters method
- self.criterion.to(self.config.device)
- self.optimizer = self.create_optimizer()
- self.scheduler = self.create_scheduler()
-
- # Initialize validation manager
- if self.config.use_lightweight_validation:
- self.validation_manager = create_validation_manager(
- model=self.model,
- device=self.config.device,
- logger=self.logger,
- lightweight_sample_ratio=self.config.lightweight_sample_ratio,
- lightweight_max_samples=self.config.lightweight_max_samples
- )
- self.logger.info("Initialized enhanced validation manager")
-
- # Load checkpoint if requested
- if self.config.resume_from_checkpoint and self.config.checkpoint_path:
- self.load_checkpoint(self.config.checkpoint_path)
-
- self.logger.info("Training setup completed successfully")
-
- def _compute_loss(self, outputs: torch.Tensor, targets: torch.Tensor, data: torch.Tensor) -> torch.Tensor:
- """
- Compute loss. Can be overridden by subclasses for custom loss computation.
-
- Args:
- outputs: Model outputs
- targets: Target labels
- data: Input data (for custom loss functions that need access to inputs)
-
- Returns:
- Loss value
- """
- return self.criterion(outputs, targets)
-
- def train_epoch(self) -> Tuple[float, float]:
- """
- Train for one epoch.
-
- Returns:
- Tuple of (average_loss, average_accuracy)
- """
- self.model.train()
- total_loss = 0.0
- total_correct = 0
- total_samples = 0
- valid_batches = 0
- skipped_batches = 0
-
- for batch_idx, (data, targets) in enumerate(self.train_loader):
- data = data.to(self.config.device)
- targets = targets.to(self.config.device)
-
- # Zero gradients
- self.optimizer.zero_grad()
-
- batch_skipped = False
-
- # Forward pass with mixed precision
- if self.scaler is not None:
- with torch.cuda.amp.autocast():
- outputs = self.model(data)
- loss = self._compute_loss(outputs, targets, data)
- # Check if batch should be skipped (loss < 0 indicates skip signal)
- # Valid losses are always >= 0, so negative loss is a clear skip indicator
- if loss.item() < 0 or torch.isnan(loss) or torch.isinf(loss):
- batch_skipped = True
- skipped_batches += 1
- self.logger.debug(f"Skipping batch {batch_idx} due to invalid loss: {loss.item()}")
- continue
- # Backward pass
- self.scaler.scale(loss).backward()
- # Check for NaN/Inf gradients before proceeding
- has_nan_grads = False
- for param in self.model.parameters():
- if param.grad is not None:
- if torch.isnan(param.grad).any() or torch.isinf(param.grad).any():
- has_nan_grads = True
- break
- if has_nan_grads:
- batch_skipped = True
- skipped_batches += 1
- self.logger.debug(f"Skipping batch {batch_idx} due to NaN/Inf gradients")
- self.optimizer.zero_grad() # Clear the bad gradients
- continue
- # Gradient clipping with logging
- if self.config.gradient_clip_norm > 0:
- self.scaler.unscale_(self.optimizer)
- grad_norm = torch.nn.utils.clip_grad_norm_(
- self.model.parameters(), self.config.gradient_clip_norm
- )
- # Check if gradient norm is reasonable
- if torch.isnan(grad_norm) or torch.isinf(grad_norm):
- batch_skipped = True
- skipped_batches += 1
- self.logger.debug(f"Skipping batch {batch_idx} due to invalid gradient norm: {grad_norm}")
- self.optimizer.zero_grad()
- continue
- # Only step if we have valid gradients
- self.scaler.step(self.optimizer)
- self.scaler.update()
-
- else:
- outputs = self.model(data)
- loss = self._compute_loss(outputs, targets, data)
- # Check if batch should be skipped (loss < 0 indicates skip signal)
- # Valid losses are always >= 0, so negative loss is a clear skip indicator
- if loss.item() < 0 or torch.isnan(loss) or torch.isinf(loss):
- batch_skipped = True
- skipped_batches += 1
- self.logger.debug(f"Skipping batch {batch_idx} due to invalid loss: {loss.item()}")
- continue
- loss.backward()
- # Check for NaN/Inf gradients
- has_nan_grads = False
- for param in self.model.parameters():
- if param.grad is not None:
- if torch.isnan(param.grad).any() or torch.isinf(param.grad).any():
- has_nan_grads = True
- break
- if has_nan_grads:
- batch_skipped = True
- skipped_batches += 1
- self.logger.debug(f"Skipping batch {batch_idx} due to NaN/Inf gradients")
- self.optimizer.zero_grad()
- continue
- # Gradient clipping with logging
- if self.config.gradient_clip_norm > 0:
- grad_norm = torch.nn.utils.clip_grad_norm_(
- self.model.parameters(), self.config.gradient_clip_norm
- )
- if torch.isnan(grad_norm) or torch.isinf(grad_norm):
- batch_skipped = True
- skipped_batches += 1
- self.logger.debug(f"Skipping batch {batch_idx} due to invalid gradient norm: {grad_norm}")
- self.optimizer.zero_grad()
- continue
- self.optimizer.step()
-
- # Update metrics only for valid batches
- if not batch_skipped:
- total_loss += loss.item()
- if len(outputs.shape) > 1 and outputs.shape[1] > 1: # Classification
- _, predicted = torch.max(outputs.data, 1)
- total_correct += (predicted == targets).sum().item()
- total_samples += targets.size(0)
- valid_batches += 1
-
- self.global_step += 1
-
- # Log progress
- if batch_idx % self.config.log_interval == 0:
- self.logger.debug(
- f"Epoch {self.current_epoch}, Batch {batch_idx}/{len(self.train_loader)}, "
- f"Loss: {loss.item():.6f}, Valid batches: {valid_batches}, Skipped: {skipped_batches}"
- )
-
- # Log skipped batches summary
- if skipped_batches > 0:
- self.logger.info(f"Epoch {self.current_epoch}: Skipped {skipped_batches}/{len(self.train_loader)} batches due to NaN/Inf issues")
-
- avg_loss = total_loss / max(valid_batches, 1)
- avg_accuracy = total_correct / max(total_samples, 1)
-
- return avg_loss, avg_accuracy
-
- def validate_epoch(self) -> Tuple[float, float]:
- """
- Validate for one epoch.
-
- Returns:
- Tuple of (average_loss, average_accuracy)
- """
- self.model.eval()
- total_loss = 0.0
- total_correct = 0
- total_samples = 0
-
- with torch.no_grad():
- for data, targets in self.val_loader:
- data = data.to(self.config.device)
- targets = targets.to(self.config.device)
-
- outputs = self.model(data)
- loss = self._compute_loss(outputs, targets, data)
-
- total_loss += loss.item()
- if len(outputs.shape) > 1 and outputs.shape[1] > 1: # Classification
- _, predicted = torch.max(outputs.data, 1)
- total_correct += (predicted == targets).sum().item()
- total_samples += targets.size(0)
-
- avg_loss = total_loss / len(self.val_loader)
- avg_accuracy = total_correct / total_samples if total_samples > 0 else 0.0
-
- return avg_loss, avg_accuracy
-
- def train(self) -> TrainingMetrics:
- """
- Main training loop.
-
- Returns:
- TrainingMetrics object containing training history
- """
- try:
- self.setup_training()
- self.state = TrainingState.TRAINING
-
- self.logger.info(f"Starting training for {self.config.num_epochs} epochs")
- start_time = time.time()
-
- for epoch in range(self.current_epoch, self.config.num_epochs):
- epoch_start_time = time.time()
- self.current_epoch = epoch
-
- # Training phase
- train_loss, train_acc = self.train_epoch()
-
- # Validation phase
- if epoch % self.config.validate_interval == 0:
- self.state = TrainingState.VALIDATING
-
- if self.validation_manager is not None:
- # Determine validation mode
- comprehensive_epochs = self.config.comprehensive_validation_epochs
- if comprehensive_epochs is None:
- # Default: comprehensive at start, middle, and end
- comprehensive_epochs = [0, self.config.num_epochs // 2, self.config.num_epochs - 1]
-
- use_lightweight = self.validation_manager.should_use_lightweight(
- epoch=epoch,
- total_epochs=self.config.num_epochs,
- lightweight_interval=self.config.validate_interval,
- comprehensive_epochs=comprehensive_epochs
- )
-
- validation_mode = ValidationMode.LIGHTWEIGHT if use_lightweight else ValidationMode.COMPREHENSIVE
-
- validation_results = self.validation_manager.validate(
- val_loader=self.val_loader,
- criterion=self.criterion,
- compute_loss_fn=self._compute_loss,
- mode=validation_mode
- )
-
- val_loss = validation_results['loss']
- val_acc = validation_results['accuracy']
-
- # Log validation mode
- self.logger.debug(f"Used {validation_mode.value} validation "
- f"(samples: {validation_results.get('samples_used', 'all')})")
- else:
- # Fallback to original validation
- val_loss, val_acc = self.validate_epoch()
-
- self.state = TrainingState.TRAINING
- else:
- val_loss, val_acc = None, None
-
- # Update learning rate
- if self.scheduler is not None:
- self.scheduler.step()
-
- # Calculate epoch time
- epoch_time = time.time() - epoch_start_time
-
- # Update metrics
- current_lr = self.optimizer.param_groups[0]['lr']
- self.metrics.update(
- epoch=epoch,
- train_loss=train_loss,
- train_acc=train_acc,
- val_loss=val_loss,
- val_acc=val_acc,
- lr=current_lr,
- epoch_time=epoch_time
- )
-
- # Log progress
- log_msg = f"Epoch {epoch}/{self.config.num_epochs} - "
- log_msg += f"Train Loss: {train_loss:.6f}, Train Acc: {train_acc:.4f}"
- if val_loss is not None:
- log_msg += f", Val Loss: {val_loss:.6f}, Val Acc: {val_acc:.4f}"
- log_msg += f", LR: {current_lr:.2e}, Time: {epoch_time:.2f}s"
- self.logger.info(log_msg)
-
- # Save checkpoint
- if self.config.save_checkpoints and epoch % self.config.checkpoint_interval == 0:
- self.save_checkpoint(epoch)
-
- # Early stopping check (only after minimum epochs)
- if val_loss is not None:
- if val_loss < self.metrics.best_val_loss:
- self.best_model_state = self.model.state_dict().copy()
- self.early_stopping_counter = 0
- else:
- # Only increment early stopping counter after min_epochs
- if epoch >= self.config.min_epochs:
- self.early_stopping_counter += 1
-
- if self.early_stopping_counter >= self.config.early_stopping_patience:
- self.logger.info(f"Early stopping triggered after {epoch + 1} epochs "
- f"(minimum {self.config.min_epochs} epochs completed)")
- break
- else:
- # Reset counter if we're still below min_epochs
- self.early_stopping_counter = 0
-
- # Execute callbacks
- for callback in self.callbacks:
- callback(self, epoch)
-
- total_time = time.time() - start_time
- self.logger.info(f"Training completed in {total_time:.2f} seconds")
-
- # Load best model if available
- if self.best_model_state is not None:
- self.model.load_state_dict(self.best_model_state)
-
- self.state = TrainingState.COMPLETED
- return self.metrics
-
- except Exception as e:
- self.state = TrainingState.FAILED
- self.logger.error(f"Training failed: {str(e)}")
- raise
-
- def test(self) -> Dict[str, float]:
- """
- Test the trained model.
-
- Returns:
- Dictionary containing test metrics
- """
- self.state = TrainingState.TESTING
- self.logger.info("Starting model testing...")
-
- if self.test_loader is None:
- raise ValueError("Test loader not available")
-
- self.model.eval()
- total_loss = 0.0
- total_correct = 0
- total_samples = 0
-
- with torch.no_grad():
- for data, targets in self.test_loader:
- data = data.to(self.config.device)
- targets = targets.to(self.config.device)
-
- outputs = self.model(data)
- loss = self.criterion(outputs, targets)
-
- total_loss += loss.item()
- if len(outputs.shape) > 1 and outputs.shape[1] > 1: # Classification
- _, predicted = torch.max(outputs.data, 1)
- total_correct += (predicted == targets).sum().item()
- total_samples += targets.size(0)
-
- test_loss = total_loss / len(self.test_loader)
- test_accuracy = total_correct / total_samples if total_samples > 0 else 0.0
-
- test_results = {
- 'test_loss': test_loss,
- 'test_accuracy': test_accuracy,
- 'total_samples': total_samples
- }
-
- self.logger.info(f"Test Results - Loss: {test_loss:.6f}, Accuracy: {test_accuracy:.4f}")
- return test_results
-
- def save_checkpoint(self, epoch: int, filepath: Optional[str] = None):
- """Save training checkpoint."""
- if filepath is None:
- os.makedirs(f"{self.config.output_dir}/checkpoints", exist_ok=True)
- filepath = f"{self.config.output_dir}/checkpoints/{self.config.model_name}_epoch_{epoch}.pth"
-
- checkpoint = {
- 'epoch': epoch,
- 'model_state_dict': self.model.state_dict(),
- 'optimizer_state_dict': self.optimizer.state_dict(),
- 'scheduler_state_dict': self.scheduler.state_dict() if self.scheduler else None,
- 'scaler_state_dict': self.scaler.state_dict() if self.scaler else None,
- 'metrics': self.metrics,
- 'config': self.config.to_dict(),
- 'best_model_state': self.best_model_state
- }
-
- torch.save(checkpoint, filepath)
- self.logger.info(f"Checkpoint saved: {filepath}")
-
- def load_checkpoint(self, filepath: str):
- """Load training checkpoint."""
- if not os.path.exists(filepath):
- raise FileNotFoundError(f"Checkpoint not found: {filepath}")
-
- checkpoint = torch.load(filepath, map_location=self.config.device)
-
- # Load model state
- if self.model is not None:
- self.model.load_state_dict(checkpoint['model_state_dict'])
-
- # Load optimizer state
- if self.optimizer is not None:
- self.optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
-
- # Load scheduler state
- if self.scheduler is not None and checkpoint.get('scheduler_state_dict'):
- self.scheduler.load_state_dict(checkpoint['scheduler_state_dict'])
-
- # Load scaler state
- if self.scaler is not None and checkpoint.get('scaler_state_dict'):
- self.scaler.load_state_dict(checkpoint['scaler_state_dict'])
-
- # Load training state
- self.current_epoch = checkpoint['epoch'] + 1
- self.metrics = checkpoint.get('metrics', TrainingMetrics())
- self.best_model_state = checkpoint.get('best_model_state')
-
- self.logger.info(f"Checkpoint loaded: {filepath}")
-
- def add_callback(self, callback: Callable):
- """Add a callback function to be executed during training."""
- self.callbacks.append(callback)
-
- def get_model_summary(self) -> str:
- """Get a summary of the model architecture."""
- if self.model is None:
- return "Model not initialized"
-
- total_params = sum(p.numel() for p in self.model.parameters())
- trainable_params = sum(p.numel() for p in self.model.parameters() if p.requires_grad)
-
- summary = f"Model: {self.model.__class__.__name__}\n"
- summary += f"Total parameters: {total_params:,}\n"
- summary += f"Trainable parameters: {trainable_params:,}\n"
- summary += f"Device: {next(self.model.parameters()).device}\n"
-
- return summary
-
- def stop_training(self):
- """Stop training gracefully."""
- self.state = TrainingState.STOPPING
- self.logger.info("Training stop requested")
-
- def pause_training(self):
- """Pause training."""
- self.state = TrainingState.PAUSED
- self.logger.info("Training paused")
-
- def resume_training(self):
- """Resume training from paused state."""
- if self.state == TrainingState.PAUSED:
- self.state = TrainingState.TRAINING
- self.logger.info("Training resumed")
-
- def get_training_status(self) -> Dict[str, Any]:
- """Get current training status and progress."""
- return {
- 'state': self.state.value,
- 'current_epoch': self.current_epoch,
- 'total_epochs': self.config.num_epochs,
- 'progress_percentage': (self.current_epoch / self.config.num_epochs) * 100,
- 'global_step': self.global_step,
- 'metrics_summary': self.metrics.get_summary(),
- 'best_epoch': self.metrics.best_epoch,
- 'early_stopping_counter': self.early_stopping_counter
- }
|