| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121 |
- import pathlib as pl
- from typing import Any, Dict
- import torch
- from torch import nn, optim
- from torch.utils.data import DataLoader
- import model.dataset as ds
- import model.training as tn
- from model.cnn import CNN3D
- from model.dnn_mod import DEFAULT_BNN_PRIOR_PARAMETERS, dnn_to_bnn_mod, get_kl_loss
- from util.control import check_control_events as _check_control_events
- from util.control import release_torch_memory as _release_torch_memory
- from util.progress import ProgressTracker
- from util.seeding import derive_seed, seed_everything
- from util.ui_logger import PipelineLogger
- # RNG stream id for the Bayesian ensemble; distinct from the normal ensemble so
- # bayesian-member-N and normal-member-N do not share an RNG stream.
- _BAYESIAN_SEED_STREAM = 1
- def train_bayesian_task(
- track: ProgressTracker,
- log: PipelineLogger,
- config: Dict[str, Any],
- state: Dict[str, Any],
- ) -> None:
- log.info("Starting bayesian model training...")
- events = state.get("events")
- stop_event = events.get("stop") if isinstance(events, dict) else None
- pause_event = events.get("pause") if isinstance(events, dict) else None
- # Get the loaded datesets from the state
- train_loader: DataLoader[ds.ADNIDataset] = state["train_loader"]
- val_loader: DataLoader[ds.ADNIDataset] = state["val_loader"]
- test_loader: DataLoader[ds.ADNIDataset] = state["test_loader"]
- bayesian_models_path = pl.Path(config["work_dir"]) / "bayesian_models"
- # Set up intermediate model directory
- intermediate_model_dir = bayesian_models_path / "intermediate_models"
- if not intermediate_model_dir.exists():
- intermediate_model_dir.mkdir(parents=True, exist_ok=True)
- log.info(f"Intermediate models will be saved to {intermediate_model_dir}")
- train_progress = track.get_sub_tracker("Training Progress")
- track.update(total=config["training"]["ensemble_size"], advance=0)
- # Shared with the model loader (evaluation) so saved state_dicts map onto an
- # identically-converted model.
- bnn_prior_parameters = DEFAULT_BNN_PRIOR_PARAMETERS
- base_seed = int(config["data"]["seed"])
- for model_num in range(config["training"]["ensemble_size"]):
- _check_control_events(stop_event=stop_event, pause_event=pause_event)
- # Seed per member (distinct stream from the normal ensemble).
- member_seed = derive_seed(base_seed, model_num, stream=_BAYESIAN_SEED_STREAM)
- seed_everything(member_seed)
- log.info(
- f"Training model {model_num + 1}/{config['training']['ensemble_size']} "
- f"(seed {member_seed})..."
- )
- model = CNN3D(
- image_channels=config["data"]["image_channels"],
- clin_data_channels=config["data"]["clin_data_channels"],
- num_classes=config["data"]["num_classes"],
- droprate=config["training"]["droprate"],
- ).float()
- dnn_to_bnn_mod(model, bnn_prior_parameters)
- model.to(config["training"]["device"])
- optimizer = optim.Adam(
- model.parameters(), lr=config["training"]["learning_rate"]
- )
- criterion = nn.BCELoss()
- model, history = tn.train_model_bayesian(
- log=log,
- progress=train_progress,
- model=model,
- train_loader=train_loader,
- val_loader=val_loader,
- optimizer=optimizer,
- criterion=criterion,
- num_epochs=config["training"]["num_epochs"],
- output_path=bayesian_models_path,
- get_kl_loss=get_kl_loss,
- stop_event=stop_event,
- pause_event=pause_event,
- )
- state.setdefault("bayesian_histories", []).append(history)
- test_loss, test_acc = tn.test_model_bayesian(
- model=model,
- test_loader=test_loader,
- criterion=criterion,
- get_kl_loss=get_kl_loss,
- progress=track.get_sub_tracker("Testing Progress"),
- log=log,
- stop_event=stop_event,
- pause_event=pause_event,
- )
- log.info(
- f"Run {model_num + 1}/{config['training']['ensemble_size']} - Test Loss: {test_loss:.4f}, Test Accuracy: {test_acc:.4f}"
- )
- model_save_path = bayesian_models_path / f"model_{model_num + 1}.pt"
- torch.save(model.state_dict(), model_save_path)
- log.info(f"Model saved to {model_save_path}")
- del history
- del optimizer
- del criterion
- del model
- _release_torch_memory(config["training"]["device"])
- track.update(total=config["training"]["ensemble_size"], advance=1)
- log.info("All Bayesian models trained and saved successfully.")
|