train_normal.py 4.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119
  1. import pathlib as pl
  2. from typing import Any, Dict
  3. import torch
  4. from torch import nn, optim
  5. from torch.utils.data import DataLoader
  6. import model.dataset as ds
  7. import model.training as tn
  8. from model.cnn import CNN3D
  9. from util.control import check_control_events as _check_control_events
  10. from util.control import release_torch_memory as _release_torch_memory
  11. from util.progress import ProgressTracker
  12. from util.seeding import derive_seed, seed_everything
  13. from util.ui_logger import PipelineLogger
  14. # RNG stream id for the normal ensemble, keeps its seeds distinct from the
  15. # Bayesian ensemble's (see util.seeding.derive_seed).
  16. _NORMAL_SEED_STREAM = 0
  17. def train_normal_task(
  18. track: ProgressTracker,
  19. log: PipelineLogger,
  20. config: Dict[str, Any],
  21. state: Dict[str, Any],
  22. ) -> None:
  23. log.info("Starting normal model training...")
  24. events = state.get("events")
  25. stop_event = events.get("stop") if isinstance(events, dict) else None
  26. pause_event = events.get("pause") if isinstance(events, dict) else None
  27. # Get the loaded datesets from the state
  28. train_loader: DataLoader[ds.ADNIDataset] = state["train_loader"]
  29. val_loader: DataLoader[ds.ADNIDataset] = state["val_loader"]
  30. test_loader: DataLoader[ds.ADNIDataset] = state["test_loader"]
  31. normal_models_path = pl.Path(config["work_dir"]) / "normal_models"
  32. # Set up intermediate model directory
  33. intermediate_model_dir = normal_models_path / "intermediate_models"
  34. if not intermediate_model_dir.exists():
  35. intermediate_model_dir.mkdir(parents=True, exist_ok=True)
  36. log.info(f"Intermediate models will be saved to {intermediate_model_dir}")
  37. base_seed = int(config["data"]["seed"])
  38. train_progress = track.get_sub_tracker("Training Progress")
  39. track.update(total=config["training"]["ensemble_size"], advance=0)
  40. for model_num in range(config["training"]["ensemble_size"]):
  41. _check_control_events(stop_event=stop_event, pause_event=pause_event)
  42. # Seed per member so weight init / dropout are reproducible AND distinct
  43. # across ensemble members.
  44. member_seed = derive_seed(base_seed, model_num, stream=_NORMAL_SEED_STREAM)
  45. seed_everything(member_seed)
  46. log.info(
  47. f"Training model {model_num + 1}/{config['training']['ensemble_size']} "
  48. f"(seed {member_seed})..."
  49. )
  50. # Train the model
  51. model = (
  52. CNN3D(
  53. image_channels=config["data"]["image_channels"],
  54. clin_data_channels=config["data"]["clin_data_channels"],
  55. num_classes=config["data"]["num_classes"],
  56. droprate=config["training"]["droprate"],
  57. )
  58. .float()
  59. .to(config["training"]["device"])
  60. )
  61. optimizer = optim.Adam(
  62. model.parameters(), lr=config["training"]["learning_rate"]
  63. )
  64. criterion = nn.BCELoss()
  65. # Train model -
  66. model, history = tn.train_model(
  67. log=log,
  68. progress=train_progress,
  69. model=model,
  70. train_loader=train_loader,
  71. val_loader=val_loader,
  72. optimizer=optimizer,
  73. criterion=criterion,
  74. num_epochs=config["training"]["num_epochs"],
  75. output_path=normal_models_path,
  76. stop_event=stop_event,
  77. pause_event=pause_event,
  78. )
  79. # Test model
  80. test_loss, test_acc = tn.test_model(
  81. model=model,
  82. test_loader=test_loader,
  83. criterion=criterion,
  84. progress=track.get_sub_tracker("Testing Progress"),
  85. log=log,
  86. stop_event=stop_event,
  87. pause_event=pause_event,
  88. )
  89. log.info(
  90. f"Run {model_num + 1}/{config['training']['ensemble_size']} - "
  91. f"Test Loss: {test_loss:.4f}, Test Accuracy: {test_acc:.4f}"
  92. )
  93. # Save the model
  94. model_save_path = normal_models_path / f"model_{model_num + 1}.pt"
  95. torch.save(model.state_dict(), model_save_path)
  96. log.info(f"Model saved to {model_save_path}")
  97. del history
  98. del optimizer
  99. del criterion
  100. del model
  101. _release_torch_memory(config["training"]["device"])
  102. track.update(total=config["training"]["ensemble_size"], advance=1)
  103. log.info("All models trained and saved successfully.")