| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119 |
- 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 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 normal ensemble, keeps its seeds distinct from the
- # Bayesian ensemble's (see util.seeding.derive_seed).
- _NORMAL_SEED_STREAM = 0
- def train_normal_task(
- track: ProgressTracker,
- log: PipelineLogger,
- config: Dict[str, Any],
- state: Dict[str, Any],
- ) -> None:
- log.info("Starting normal 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"]
- normal_models_path = pl.Path(config["work_dir"]) / "normal_models"
- # Set up intermediate model directory
- intermediate_model_dir = normal_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}")
- base_seed = int(config["data"]["seed"])
- train_progress = track.get_sub_tracker("Training Progress")
- track.update(total=config["training"]["ensemble_size"], advance=0)
- for model_num in range(config["training"]["ensemble_size"]):
- _check_control_events(stop_event=stop_event, pause_event=pause_event)
- # Seed per member so weight init / dropout are reproducible AND distinct
- # across ensemble members.
- member_seed = derive_seed(base_seed, model_num, stream=_NORMAL_SEED_STREAM)
- seed_everything(member_seed)
- log.info(
- f"Training model {model_num + 1}/{config['training']['ensemble_size']} "
- f"(seed {member_seed})..."
- )
- # Train the model
- 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()
- .to(config["training"]["device"])
- )
- optimizer = optim.Adam(
- model.parameters(), lr=config["training"]["learning_rate"]
- )
- criterion = nn.BCELoss()
- # Train model -
- model, history = tn.train_model(
- 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=normal_models_path,
- stop_event=stop_event,
- pause_event=pause_event,
- )
- # Test model
- test_loss, test_acc = tn.test_model(
- model=model,
- test_loader=test_loader,
- criterion=criterion,
- 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']} - "
- f"Test Loss: {test_loss:.4f}, Test Accuracy: {test_acc:.4f}"
- )
- # Save the model
- model_save_path = normal_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 models trained and saved successfully.")
|