train_normal.py 4.4 KB

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