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.seeding import derive_seed, seed_everything from util.ui_logger import PipelineLogger # RNG stream id for the Bayesian ensemble; distinct from the normal ensemble so # bayesian-member-N and normal-member-N do not share an RNG stream. _BAYESIAN_SEED_STREAM = 1 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, } base_seed = int(config["data"]["seed"]) for model_num in range(config["training"]["ensemble_size"]): _check_control_events(stop_event=stop_event, pause_event=pause_event) # Seed per member (distinct stream from the normal ensemble). member_seed = derive_seed(base_seed, model_num, stream=_BAYESIAN_SEED_STREAM) seed_everything(member_seed) log.info( f"Training model {model_num + 1}/{config['training']['ensemble_size']} " f"(seed {member_seed})..." ) 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.")