import pathlib as pl import time import gc 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 util.progress import ProgressTracker from util.seeding import derive_seed, seed_everything from util.ui_logger import PipelineLogger # RNG stream id for the normal ensemble, keeps its seeds distinct from the # Bayesian ensemble's (see util.seeding.derive_seed). _NORMAL_SEED_STREAM = 0 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_normal_task( track: ProgressTracker, log: PipelineLogger, config: Dict[str, Any], state: Dict[str, Any], ) -> None: log.info("Starting normal 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"] normal_models_path = pl.Path(config["work_dir"]) / "normal_models" # Set up intermediate model directory intermediate_model_dir = normal_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}") base_seed = int(config["data"]["seed"]) train_progress = track.get_sub_tracker("Training Progress") track.update(total=config["training"]["ensemble_size"], advance=0) for model_num in range(config["training"]["ensemble_size"]): _check_control_events(stop_event=stop_event, pause_event=pause_event) # Seed per member so weight init / dropout are reproducible AND distinct # across ensemble members. member_seed = derive_seed(base_seed, model_num, stream=_NORMAL_SEED_STREAM) seed_everything(member_seed) log.info( f"Training model {model_num + 1}/{config['training']['ensemble_size']} " f"(seed {member_seed})..." ) # Train the model 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"]) ) optimizer = optim.Adam( model.parameters(), lr=config["training"]["learning_rate"] ) criterion = nn.BCELoss() # Train model - model, history = tn.train_model( 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=normal_models_path, stop_event=stop_event, pause_event=pause_event, ) # Test model test_loss, test_acc = tn.test_model( model=model, test_loader=test_loader, criterion=criterion, 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}" ) # Save the model model_save_path = normal_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 models trained and saved successfully.")