CoolFace
Apppublic

solarevat/multilabel-news-classifier

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
lightning_module_tracking.py141 linesDownload Raw Back to models
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