train_bayesian.py 4.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121
  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 model.dnn_mod import DEFAULT_BNN_PRIOR_PARAMETERS, dnn_to_bnn_mod, get_kl_loss
  10. from util.control import check_control_events as _check_control_events
  11. from util.control import release_torch_memory as _release_torch_memory
  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 Bayesian ensemble; distinct from the normal ensemble so
  16. # bayesian-member-N and normal-member-N do not share an RNG stream.
  17. _BAYESIAN_SEED_STREAM = 1
  18. def train_bayesian_task(
  19. track: ProgressTracker,
  20. log: PipelineLogger,
  21. config: Dict[str, Any],
  22. state: Dict[str, Any],
  23. ) -> None:
  24. log.info("Starting bayesian model training...")
  25. events = state.get("events")
  26. stop_event = events.get("stop") if isinstance(events, dict) else None
  27. pause_event = events.get("pause") if isinstance(events, dict) else None
  28. # Get the loaded datesets from the state
  29. train_loader: DataLoader[ds.ADNIDataset] = state["train_loader"]
  30. val_loader: DataLoader[ds.ADNIDataset] = state["val_loader"]
  31. test_loader: DataLoader[ds.ADNIDataset] = state["test_loader"]
  32. bayesian_models_path = pl.Path(config["work_dir"]) / "bayesian_models"
  33. # Set up intermediate model directory
  34. intermediate_model_dir = bayesian_models_path / "intermediate_models"
  35. if not intermediate_model_dir.exists():
  36. intermediate_model_dir.mkdir(parents=True, exist_ok=True)
  37. log.info(f"Intermediate models will be saved to {intermediate_model_dir}")
  38. train_progress = track.get_sub_tracker("Training Progress")
  39. track.update(total=config["training"]["ensemble_size"], advance=0)
  40. # Shared with the model loader (evaluation) so saved state_dicts map onto an
  41. # identically-converted model.
  42. bnn_prior_parameters = DEFAULT_BNN_PRIOR_PARAMETERS
  43. base_seed = int(config["data"]["seed"])
  44. for model_num in range(config["training"]["ensemble_size"]):
  45. _check_control_events(stop_event=stop_event, pause_event=pause_event)
  46. # Seed per member (distinct stream from the normal ensemble).
  47. member_seed = derive_seed(base_seed, model_num, stream=_BAYESIAN_SEED_STREAM)
  48. seed_everything(member_seed)
  49. log.info(
  50. f"Training model {model_num + 1}/{config['training']['ensemble_size']} "
  51. f"(seed {member_seed})..."
  52. )
  53. model = CNN3D(
  54. image_channels=config["data"]["image_channels"],
  55. clin_data_channels=config["data"]["clin_data_channels"],
  56. num_classes=config["data"]["num_classes"],
  57. droprate=config["training"]["droprate"],
  58. ).float()
  59. dnn_to_bnn_mod(model, bnn_prior_parameters)
  60. model.to(config["training"]["device"])
  61. optimizer = optim.Adam(
  62. model.parameters(), lr=config["training"]["learning_rate"]
  63. )
  64. criterion = nn.BCELoss()
  65. model, history = tn.train_model_bayesian(
  66. log=log,
  67. progress=train_progress,
  68. model=model,
  69. train_loader=train_loader,
  70. val_loader=val_loader,
  71. optimizer=optimizer,
  72. criterion=criterion,
  73. num_epochs=config["training"]["num_epochs"],
  74. output_path=bayesian_models_path,
  75. get_kl_loss=get_kl_loss,
  76. stop_event=stop_event,
  77. pause_event=pause_event,
  78. )
  79. state.setdefault("bayesian_histories", []).append(history)
  80. test_loss, test_acc = tn.test_model_bayesian(
  81. model=model,
  82. test_loader=test_loader,
  83. criterion=criterion,
  84. get_kl_loss=get_kl_loss,
  85. progress=track.get_sub_tracker("Testing Progress"),
  86. log=log,
  87. stop_event=stop_event,
  88. pause_event=pause_event,
  89. )
  90. log.info(
  91. f"Run {model_num + 1}/{config['training']['ensemble_size']} - Test Loss: {test_loss:.4f}, Test Accuracy: {test_acc:.4f}"
  92. )
  93. model_save_path = bayesian_models_path / f"model_{model_num + 1}.pt"
  94. torch.save(model.state_dict(), model_save_path)
  95. log.info(f"Model saved to {model_save_path}")
  96. del history
  97. del optimizer
  98. del criterion
  99. del model
  100. _release_torch_memory(config["training"]["device"])
  101. track.update(total=config["training"]["ensemble_size"], advance=1)
  102. log.info("All Bayesian models trained and saved successfully.")