solarevat/multilabel-news-classifier
0
1"""PyTorch Lightning module with enhanced experiment tracking."""2 3from typing import Dict, Any, Optional4import torch5import torch.nn as nn6import pytorch_lightning as pl7from pytorch_lightning.callbacks import Callback8from pytorch_lightning.loggers import WandbLogger, MLFlowLogger9 10from utils.experiment_tracking import WandBTracker, MLflowTracker, ExperimentTracker11import logging12 13logger = logging.getLogger(__name__)14 15 16class WandBCallback(Callback):17 """Enhanced WandB callback for PyTorch Lightning."""18 19 def __init__(self, log_model: bool = True, log_artifacts: bool = True):20 """21 Initialize WandB callback.22 23 Args:24 log_model: Whether to log model checkpoints25 log_artifacts: Whether to log artifacts26 """27 super().__init__()28 self.log_model = log_model29 self.log_artifacts = log_artifacts30 31 def on_train_epoch_end(self, trainer: pl.Trainer, pl_module: pl.LightningModule) -> None:32 """Log metrics at end of training epoch."""33 metrics = {f"train/{k}": v for k, v in trainer.callback_metrics.items()}34 if hasattr(trainer, 'logger') and isinstance(trainer.logger, WandbLogger):35 trainer.logger.experiment.log(metrics)36 37 def on_validation_epoch_end(self, trainer: pl.Trainer, pl_module: pl.LightningModule) -> None:38 """Log metrics at end of validation epoch."""39 metrics = {f"val/{k}": v for k, v in trainer.callback_metrics.items()}40 if hasattr(trainer, 'logger') and isinstance(trainer.logger, WandbLogger):41 trainer.logger.experiment.log(metrics)42 43 def on_train_end(self, trainer: pl.Trainer, pl_module: pl.LightningModule) -> None:44 """Log artifacts at end of training."""45 if self.log_artifacts and hasattr(trainer, 'logger'):46 if isinstance(trainer.logger, WandbLogger):47 # Log best model48 if trainer.checkpoint_callback and trainer.checkpoint_callback.best_model_path:49 trainer.logger.experiment.log_artifact(50 trainer.checkpoint_callback.best_model_path,51 name="best_model"52 )53 54 55class MLflowCallback(Callback):56 """MLflow callback for PyTorch Lightning."""57 58 def __init__(self, log_model: bool = True):59 """60 Initialize MLflow callback.61 62 Args:63 log_model: Whether to log model64 """65 super().__init__()66 self.log_model = log_model67 68 def on_train_epoch_end(self, trainer: pl.Trainer, pl_module: pl.LightningModule) -> None:69 """Log metrics at end of training epoch."""70 if hasattr(trainer, 'logger') and isinstance(trainer.logger, MLFlowLogger):71 metrics = {f"train_{k}": v for k, v in trainer.callback_metrics.items()}72 trainer.logger.experiment.log_metrics(metrics, step=trainer.current_epoch)73 74 def on_validation_epoch_end(self, trainer: pl.Trainer, pl_module: pl.LightningModule) -> None:75 """Log metrics at end of validation epoch."""76 if hasattr(trainer, 'logger') and isinstance(trainer.logger, MLFlowLogger):77 metrics = {f"val_{k}": v for k, v in trainer.callback_metrics.items()}78 trainer.logger.experiment.log_metrics(metrics, step=trainer.current_epoch)79 80 def on_train_end(self, trainer: pl.Trainer, pl_module: pl.LightningModule) -> None:81 """Log model at end of training."""82 if self.log_model and hasattr(trainer, 'logger'):83 if isinstance(trainer.logger, MLFlowLogger):84 # Log model85 trainer.logger.experiment.log_model(86 pl_module.model,87 artifact_path="model"88 )89 90 91def create_tracking_loggers(92 use_wandb: bool = True,93 use_mlflow: bool = True,94 project_name: str = "russian-news-classification",95 experiment_name: Optional[str] = None,96 **kwargs97) -> tuple[list, list]:98 """99 Create tracking loggers and callbacks.100 101 Args:102 use_wandb: Enable WandB103 use_mlflow: Enable MLflow104 project_name: Project name105 experiment_name: Experiment name106 **kwargs: Additional arguments107 108 Returns:109 Tuple of (loggers, callbacks)110 """111 loggers = []112 callbacks = []113 114 if use_wandb:115 try:116 wandb_logger = WandbLogger(117 project=project_name,118 name=experiment_name,119 **kwargs.get('wandb', {})120 )121 loggers.append(wandb_logger)122 callbacks.append(WandBCallback())123 logger.info("WandB logger created")124 except Exception as e:125 logger.warning(f"Failed to create WandB logger: {e}")126 127 if use_mlflow:128 try:129 mlflow_logger = MLFlowLogger(130 experiment_name=experiment_name or project_name,131 tracking_uri=kwargs.get('mlflow', {}).get('tracking_uri'),132 )133 loggers.append(mlflow_logger)134 callbacks.append(MLflowCallback())135 logger.info("MLflow logger created")136 except Exception as e:137 logger.warning(f"Failed to create MLflow logger: {e}")138 139 return loggers, callbacks140 141 