|
@@ -0,0 +1,140 @@
|
|
|
|
|
+import pathlib as pl
|
|
|
|
|
+import time
|
|
|
|
|
+import gc
|
|
|
|
|
+from threading import Event
|
|
|
|
|
+from typing import Any, Dict
|
|
|
|
|
+
|
|
|
|
|
+import torch
|
|
|
|
|
+from bayesian_torch.models.dnn_to_bnn import dnn_to_bnn, get_kl_loss # type: ignore[import-untyped]
|
|
|
|
|
+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.ui_logger import PipelineLogger
|
|
|
|
|
+
|
|
|
|
|
+
|
|
|
|
|
+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,
|
|
|
|
|
+ }
|
|
|
|
|
+
|
|
|
|
|
+ for model_num in range(config["training"]["ensemble_size"]):
|
|
|
|
|
+ _check_control_events(stop_event=stop_event, pause_event=pause_event)
|
|
|
|
|
+ log.info(
|
|
|
|
|
+ f"Training model {model_num + 1}/{config['training']['ensemble_size']}..."
|
|
|
|
|
+ )
|
|
|
|
|
+ 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"])
|
|
|
|
|
+ )
|
|
|
|
|
+ dnn_to_bnn(model, bnn_prior_parameters)
|
|
|
|
|
+
|
|
|
|
|
+ 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']} - "
|
|
|
|
|
+ f"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.")
|