| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131 |
- import pathlib as pl
- import time
- import gc
- from threading import Event
- 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.progress import ProgressTracker
- from util.ui_logger import PipelineLogger
- def _release_torch_memory(device: str) -> None:
- gc.collect()
- if device.startswith("cuda"):
- torch.cuda.empty_cache()
- elif device.startswith("mps") and hasattr(torch, "mps"):
- torch.mps.empty_cache()
- def _check_control_events(
- stop_event: Event | None,
- pause_event: Event | None,
- ) -> None:
- if stop_event is not None and stop_event.is_set():
- raise InterruptedError("Pipeline execution stopped by user.")
- while pause_event is not None and pause_event.is_set():
- time.sleep(0.5)
- if stop_event is not None and stop_event.is_set():
- raise InterruptedError("Pipeline execution stopped by user while paused.")
- 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}")
- models = []
- 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)
- log.info(
- f"Training model {model_num + 1}/{config['training']['ensemble_size']}..."
- )
- # 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"])
- )
- models.append(model)
- 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.")
|