import gc import pathlib as pl import time from threading import Event from typing import Any, Dict import torch 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 model.dnn_mod import dnn_to_bnn_mod, get_kl_loss 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() dnn_to_bnn_mod(model, bnn_prior_parameters) model.to(config["training"]["device"]) 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']} - 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.")