train_normal.py 4.9 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 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 util.progress import ProgressTracker
  13. from util.seeding import derive_seed, seed_everything
  14. from util.ui_logger import PipelineLogger
  15. # RNG stream id for the normal ensemble, keeps its seeds distinct from the
  16. # Bayesian ensemble's (see util.seeding.derive_seed).
  17. _NORMAL_SEED_STREAM = 0
  18. def _release_torch_memory(device: str) -> None:
  19. gc.collect()
  20. if device.startswith("cuda"):
  21. torch.cuda.empty_cache()
  22. elif device.startswith("mps") and hasattr(torch, "mps"):
  23. torch.mps.empty_cache()
  24. def _check_control_events(
  25. stop_event: Event | None,
  26. pause_event: Event | None,
  27. ) -> None:
  28. if stop_event is not None and stop_event.is_set():
  29. raise InterruptedError("Pipeline execution stopped by user.")
  30. while pause_event is not None and pause_event.is_set():
  31. time.sleep(0.5)
  32. if stop_event is not None and stop_event.is_set():
  33. raise InterruptedError("Pipeline execution stopped by user while paused.")
  34. def train_normal_task(
  35. track: ProgressTracker,
  36. log: PipelineLogger,
  37. config: Dict[str, Any],
  38. state: Dict[str, Any],
  39. ) -> None:
  40. log.info("Starting normal model training...")
  41. events = state.get("events")
  42. stop_event = events.get("stop") if isinstance(events, dict) else None
  43. pause_event = events.get("pause") if isinstance(events, dict) else None
  44. # Get the loaded datesets from the state
  45. train_loader: DataLoader[ds.ADNIDataset] = state["train_loader"]
  46. val_loader: DataLoader[ds.ADNIDataset] = state["val_loader"]
  47. test_loader: DataLoader[ds.ADNIDataset] = state["test_loader"]
  48. normal_models_path = pl.Path(config["work_dir"]) / "normal_models"
  49. # Set up intermediate model directory
  50. intermediate_model_dir = normal_models_path / "intermediate_models"
  51. if not intermediate_model_dir.exists():
  52. intermediate_model_dir.mkdir(parents=True, exist_ok=True)
  53. log.info(f"Intermediate models will be saved to {intermediate_model_dir}")
  54. base_seed = int(config["data"]["seed"])
  55. train_progress = track.get_sub_tracker("Training Progress")
  56. track.update(total=config["training"]["ensemble_size"], advance=0)
  57. for model_num in range(config["training"]["ensemble_size"]):
  58. _check_control_events(stop_event=stop_event, pause_event=pause_event)
  59. # Seed per member so weight init / dropout are reproducible AND distinct
  60. # across ensemble members.
  61. member_seed = derive_seed(base_seed, model_num, stream=_NORMAL_SEED_STREAM)
  62. seed_everything(member_seed)
  63. log.info(
  64. f"Training model {model_num + 1}/{config['training']['ensemble_size']} "
  65. f"(seed {member_seed})..."
  66. )
  67. # Train the model
  68. model = (
  69. CNN3D(
  70. image_channels=config["data"]["image_channels"],
  71. clin_data_channels=config["data"]["clin_data_channels"],
  72. num_classes=config["data"]["num_classes"],
  73. droprate=config["training"]["droprate"],
  74. )
  75. .float()
  76. .to(config["training"]["device"])
  77. )
  78. optimizer = optim.Adam(
  79. model.parameters(), lr=config["training"]["learning_rate"]
  80. )
  81. criterion = nn.BCELoss()
  82. # Train model -
  83. model, history = tn.train_model(
  84. log=log,
  85. progress=train_progress,
  86. model=model,
  87. train_loader=train_loader,
  88. val_loader=val_loader,
  89. optimizer=optimizer,
  90. criterion=criterion,
  91. num_epochs=config["training"]["num_epochs"],
  92. output_path=normal_models_path,
  93. stop_event=stop_event,
  94. pause_event=pause_event,
  95. )
  96. # Test model
  97. test_loss, test_acc = tn.test_model(
  98. model=model,
  99. test_loader=test_loader,
  100. criterion=criterion,
  101. progress=track.get_sub_tracker("Testing Progress"),
  102. log=log,
  103. stop_event=stop_event,
  104. pause_event=pause_event,
  105. )
  106. log.info(
  107. f"Run {model_num + 1}/{config['training']['ensemble_size']} - "
  108. f"Test Loss: {test_loss:.4f}, Test Accuracy: {test_acc:.4f}"
  109. )
  110. # Save the model
  111. model_save_path = normal_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 models trained and saved successfully.")