train_bayesian.py 4.8 KB

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