train_bayesian.py 4.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137
  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.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 = CNN3D(
  68. image_channels=config["data"]["image_channels"],
  69. clin_data_channels=config["data"]["clin_data_channels"],
  70. num_classes=config["data"]["num_classes"],
  71. droprate=config["training"]["droprate"],
  72. ).float()
  73. dnn_to_bnn_mod(model, bnn_prior_parameters)
  74. model.to(config["training"]["device"])
  75. optimizer = optim.Adam(
  76. model.parameters(), lr=config["training"]["learning_rate"]
  77. )
  78. criterion = nn.BCELoss()
  79. model, history = tn.train_model_bayesian(
  80. log=log,
  81. progress=train_progress,
  82. model=model,
  83. train_loader=train_loader,
  84. val_loader=val_loader,
  85. optimizer=optimizer,
  86. criterion=criterion,
  87. num_epochs=config["training"]["num_epochs"],
  88. output_path=bayesian_models_path,
  89. get_kl_loss=get_kl_loss,
  90. stop_event=stop_event,
  91. pause_event=pause_event,
  92. )
  93. state.setdefault("bayesian_histories", []).append(history)
  94. test_loss, test_acc = tn.test_model_bayesian(
  95. model=model,
  96. test_loader=test_loader,
  97. criterion=criterion,
  98. get_kl_loss=get_kl_loss,
  99. progress=track.get_sub_tracker("Testing Progress"),
  100. log=log,
  101. stop_event=stop_event,
  102. pause_event=pause_event,
  103. )
  104. log.info(
  105. f"Run {model_num + 1}/{config['training']['ensemble_size']} - Test Loss: {test_loss:.4f}, Test Accuracy: {test_acc:.4f}"
  106. )
  107. model_save_path = bayesian_models_path / f"model_{model_num + 1}.pt"
  108. torch.save(model.state_dict(), model_save_path)
  109. log.info(f"Model saved to {model_save_path}")
  110. del history
  111. del optimizer
  112. del criterion
  113. del model
  114. _release_torch_memory(config["training"]["device"])
  115. track.update(total=config["training"]["ensemble_size"], advance=1)
  116. log.info("All Bayesian models trained and saved successfully.")