import pathlib as pl import time import gc from threading import Event from typing import Any, Dict import torch from bayesian_torch.models.dnn_to_bnn import dnn_to_bnn, get_kl_loss # type: ignore[import-untyped] from torch import nn, optim from torch.utils.data import DataLoader import model.dataset as ds import model.training as tn from model.cnn import CNN3D from util.progress import ProgressTracker from util.ui_logger import PipelineLogger def _release_torch_memory(device: str) -> None: gc.collect() if device.startswith("cuda"): torch.cuda.empty_cache() elif device.startswith("mps") and hasattr(torch, "mps"): torch.mps.empty_cache() def _check_control_events( stop_event: Event | None, pause_event: Event | None, ) -> None: if stop_event is not None and stop_event.is_set(): raise InterruptedError("Pipeline execution stopped by user.") while pause_event is not None and pause_event.is_set(): time.sleep(0.5) if stop_event is not None and stop_event.is_set(): raise InterruptedError("Pipeline execution stopped by user while paused.") def train_bayesian_task( track: ProgressTracker, log: PipelineLogger, config: Dict[str, Any], state: Dict[str, Any], ) -> None: log.info("Starting bayesian model training...") events = state.get("events") stop_event = events.get("stop") if isinstance(events, dict) else None pause_event = events.get("pause") if isinstance(events, dict) else None # Get the loaded datesets from the state train_loader: DataLoader[ds.ADNIDataset] = state["train_loader"] val_loader: DataLoader[ds.ADNIDataset] = state["val_loader"] test_loader: DataLoader[ds.ADNIDataset] = state["test_loader"] bayesian_models_path = pl.Path(config["work_dir"]) / "bayesian_models" # Set up intermediate model directory intermediate_model_dir = bayesian_models_path / "intermediate_models" if not intermediate_model_dir.exists(): intermediate_model_dir.mkdir(parents=True, exist_ok=True) log.info(f"Intermediate models will be saved to {intermediate_model_dir}") train_progress = track.get_sub_tracker("Training Progress") track.update(total=config["training"]["ensemble_size"], advance=0) bnn_prior_parameters = { "prior_mu": 0.0, "prior_sigma": 1.0, "posterior_mu_init": 0.0, "posterior_rho_init": -3.0, "type": "Reparameterization", "moped_enable": False, "moped_delta": 0.5, } for model_num in range(config["training"]["ensemble_size"]): _check_control_events(stop_event=stop_event, pause_event=pause_event) log.info( f"Training model {model_num + 1}/{config['training']['ensemble_size']}..." ) model = ( CNN3D( image_channels=config["data"]["image_channels"], clin_data_channels=config["data"]["clin_data_channels"], num_classes=config["data"]["num_classes"], droprate=config["training"]["droprate"], ) .float() .to(config["training"]["device"]) ) dnn_to_bnn(model, bnn_prior_parameters) optimizer = optim.Adam( model.parameters(), lr=config["training"]["learning_rate"] ) criterion = nn.BCELoss() model, history = tn.train_model_bayesian( log=log, progress=train_progress, model=model, train_loader=train_loader, val_loader=val_loader, optimizer=optimizer, criterion=criterion, num_epochs=config["training"]["num_epochs"], output_path=bayesian_models_path, get_kl_loss=get_kl_loss, stop_event=stop_event, pause_event=pause_event, ) state.setdefault("bayesian_histories", []).append(history) test_loss, test_acc = tn.test_model_bayesian( model=model, test_loader=test_loader, criterion=criterion, get_kl_loss=get_kl_loss, progress=track.get_sub_tracker("Testing Progress"), log=log, stop_event=stop_event, pause_event=pause_event, ) log.info( f"Run {model_num + 1}/{config['training']['ensemble_size']} - " f"Test Loss: {test_loss:.4f}, Test Accuracy: {test_acc:.4f}" ) model_save_path = bayesian_models_path / f"model_{model_num + 1}.pt" torch.save(model.state_dict(), model_save_path) log.info(f"Model saved to {model_save_path}") del history del optimizer del criterion del model _release_torch_memory(config["training"]["device"]) track.update(total=config["training"]["ensemble_size"], advance=1) log.info("All Bayesian models trained and saved successfully.")