import pathlib as pl 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 DEFAULT_BNN_PRIOR_PARAMETERS, dnn_to_bnn_mod, get_kl_loss from util.control import check_control_events as _check_control_events from util.control import release_torch_memory as _release_torch_memory 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 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) # Shared with the model loader (evaluation) so saved state_dicts map onto an # identically-converted model. bnn_prior_parameters = DEFAULT_BNN_PRIOR_PARAMETERS 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.")