| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148 |
- 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.")
|