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 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 normal ensemble, keeps its seeds distinct from the # Bayesian ensemble's (see util.seeding.derive_seed). _NORMAL_SEED_STREAM = 0 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.")