train_bayesian.py 5.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148
  1. import gc
  2. import pathlib as pl
  3. import time
  4. from threading import Event
  5. from typing import Any, Dict
  6. import torch
  7. from torch import nn, optim
  8. from torch.utils.data import DataLoader
  9. import model.dataset as ds
  10. import model.training as tn
  11. from model.cnn import CNN3D
  12. from model.dnn_mod import dnn_to_bnn_mod, get_kl_loss
  13. from util.progress import ProgressTracker
  14. from util.seeding import derive_seed, seed_everything
  15. from util.ui_logger import PipelineLogger
  16. # RNG stream id for the Bayesian ensemble; distinct from the normal ensemble so
  17. # bayesian-member-N and normal-member-N do not share an RNG stream.
  18. _BAYESIAN_SEED_STREAM = 1
  19. def _release_torch_memory(device: str) -> None:
  20. gc.collect()
  21. if device.startswith("cuda"):
  22. torch.cuda.empty_cache()
  23. elif device.startswith("mps") and hasattr(torch, "mps"):
  24. torch.mps.empty_cache()
  25. def _check_control_events(
  26. stop_event: Event | None,
  27. pause_event: Event | None,
  28. ) -> None:
  29. if stop_event is not None and stop_event.is_set():
  30. raise InterruptedError("Pipeline execution stopped by user.")
  31. while pause_event is not None and pause_event.is_set():
  32. time.sleep(0.5)
  33. if stop_event is not None and stop_event.is_set():
  34. raise InterruptedError("Pipeline execution stopped by user while paused.")
  35. def train_bayesian_task(
  36. track: ProgressTracker,
  37. log: PipelineLogger,
  38. config: Dict[str, Any],
  39. state: Dict[str, Any],
  40. ) -> None:
  41. log.info("Starting bayesian model training...")
  42. events = state.get("events")
  43. stop_event = events.get("stop") if isinstance(events, dict) else None
  44. pause_event = events.get("pause") if isinstance(events, dict) else None
  45. # Get the loaded datesets from the state
  46. train_loader: DataLoader[ds.ADNIDataset] = state["train_loader"]
  47. val_loader: DataLoader[ds.ADNIDataset] = state["val_loader"]
  48. test_loader: DataLoader[ds.ADNIDataset] = state["test_loader"]
  49. bayesian_models_path = pl.Path(config["work_dir"]) / "bayesian_models"
  50. # Set up intermediate model directory
  51. intermediate_model_dir = bayesian_models_path / "intermediate_models"
  52. if not intermediate_model_dir.exists():
  53. intermediate_model_dir.mkdir(parents=True, exist_ok=True)
  54. log.info(f"Intermediate models will be saved to {intermediate_model_dir}")
  55. train_progress = track.get_sub_tracker("Training Progress")
  56. track.update(total=config["training"]["ensemble_size"], advance=0)
  57. bnn_prior_parameters = {
  58. "prior_mu": 0.0,
  59. "prior_sigma": 1.0,
  60. "posterior_mu_init": 0.0,
  61. "posterior_rho_init": -3.0,
  62. "type": "Reparameterization",
  63. "moped_enable": False,
  64. "moped_delta": 0.5,
  65. }
  66. base_seed = int(config["data"]["seed"])
  67. for model_num in range(config["training"]["ensemble_size"]):
  68. _check_control_events(stop_event=stop_event, pause_event=pause_event)
  69. # Seed per member (distinct stream from the normal ensemble).
  70. member_seed = derive_seed(base_seed, model_num, stream=_BAYESIAN_SEED_STREAM)
  71. seed_everything(member_seed)
  72. log.info(
  73. f"Training model {model_num + 1}/{config['training']['ensemble_size']} "
  74. f"(seed {member_seed})..."
  75. )
  76. model = CNN3D(
  77. image_channels=config["data"]["image_channels"],
  78. clin_data_channels=config["data"]["clin_data_channels"],
  79. num_classes=config["data"]["num_classes"],
  80. droprate=config["training"]["droprate"],
  81. ).float()
  82. dnn_to_bnn_mod(model, bnn_prior_parameters)
  83. model.to(config["training"]["device"])
  84. optimizer = optim.Adam(
  85. model.parameters(), lr=config["training"]["learning_rate"]
  86. )
  87. criterion = nn.BCELoss()
  88. model, history = tn.train_model_bayesian(
  89. log=log,
  90. progress=train_progress,
  91. model=model,
  92. train_loader=train_loader,
  93. val_loader=val_loader,
  94. optimizer=optimizer,
  95. criterion=criterion,
  96. num_epochs=config["training"]["num_epochs"],
  97. output_path=bayesian_models_path,
  98. get_kl_loss=get_kl_loss,
  99. stop_event=stop_event,
  100. pause_event=pause_event,
  101. )
  102. state.setdefault("bayesian_histories", []).append(history)
  103. test_loss, test_acc = tn.test_model_bayesian(
  104. model=model,
  105. test_loader=test_loader,
  106. criterion=criterion,
  107. get_kl_loss=get_kl_loss,
  108. progress=track.get_sub_tracker("Testing Progress"),
  109. log=log,
  110. stop_event=stop_event,
  111. pause_event=pause_event,
  112. )
  113. log.info(
  114. f"Run {model_num + 1}/{config['training']['ensemble_size']} - Test Loss: {test_loss:.4f}, Test Accuracy: {test_acc:.4f}"
  115. )
  116. model_save_path = bayesian_models_path / f"model_{model_num + 1}.pt"
  117. torch.save(model.state_dict(), model_save_path)
  118. log.info(f"Model saved to {model_save_path}")
  119. del history
  120. del optimizer
  121. del criterion
  122. del model
  123. _release_torch_memory(config["training"]["device"])
  124. track.update(total=config["training"]["ensemble_size"], advance=1)
  125. log.info("All Bayesian models trained and saved successfully.")