Selaa lähdekoodia

New files for previous commit

Nicholas Schense 2 viikkoa sitten
vanhempi
commit
4aa417a447
6 muutettua tiedostoa jossa 517 lisäystä ja 0 poistoa
  1. 23 0
      config.toml
  2. 98 0
      tasks/__init__.py
  3. 61 0
      tasks/load_data.py
  4. 140 0
      tasks/train_bayesian.py
  5. 131 0
      tasks/train_normal.py
  6. 64 0
      util/config_manager.py

+ 23 - 0
config.toml

@@ -0,0 +1,23 @@
+# Pipline Scenario to run
+# "scen_train_all" - Train and evaluate model, with noise
+# "scen_load_all" - Load and evaluate model, with noise
+# "scen_load_eval" - Load and evaluate model, without noise
+scenario = "scen_train_all"
+
+[data]
+mri_files_path = "../data/PET_volumes_customtemplate_float32/"
+xls_file_path = "../data/LP_ADNIMERGE.csv"
+seed = 42
+data_splits = [0.7, 0.2, 0.1] # train, validation, test
+image_channels = 1
+clin_data_channels = 2
+num_classes = 2 # AD, NL
+
+
+[training]
+device = "cuda:0" # "cpu", "cuda", "mps"
+batch_size = 32
+ensemble_size = 50
+droprate = 0.05
+learning_rate = 0.0001
+num_epochs = 30

+ 98 - 0
tasks/__init__.py

@@ -0,0 +1,98 @@
+import time
+from typing import Any, Dict
+
+from util.progress import ProgressTracker
+from util.ui_logger import PipelineLogger
+
+from . import load_data
+from . import train_bayesian
+from . import train_normal
+
+
+def dummy_task(
+    tracker: ProgressTracker,
+    logger: PipelineLogger,
+    config: Dict[str, Any],
+    state: Dict[str, Any],
+):
+    stop_event = state["events"]["stop"]
+    pause_event = state["events"]["pause"]
+
+    steps = 5
+    # The title is now set gracefully by the parent injecting the sub_tracker,
+    # so we just initialize the total.
+    tracker.update(total=steps, advance=0)
+
+    logger.info("Initializing process...")
+
+    for i in range(steps):
+        if stop_event.is_set():
+            raise InterruptedError("Pipeline execution stopped by user.")
+
+        while pause_event.is_set():
+            time.sleep(0.5)
+            if stop_event.is_set():
+                raise InterruptedError(
+                    "Pipeline execution stopped by user while paused."
+                )
+
+        # Child Process (Sub Progress Tracker)
+        sub_steps = 10
+        sub_tracker = tracker.get_sub_tracker(f"Batch {i + 1}")
+        sub_tracker.update(total=sub_steps, advance=0)
+
+        for j in range(sub_steps):
+            if stop_event.is_set():
+                raise InterruptedError("Pipeline execution stopped by user.")
+            time.sleep(0.05)
+            sub_tracker.update(advance=1)  # Advance sub task
+
+        tracker.update(advance=1)  # Advance main task (clears sub task)
+
+        if i == 2:
+            logger.info("Halfway through current task execution...")
+
+
+PIPELINE_TASKS = {
+    "load_data": {
+        "task_name": "Load Image and ADNIMERGE",
+        "task_func": load_data.load_data_task,
+    },
+    "train_regular": {
+        "task_name": "Train Regular Models",
+        "task_func": train_normal.train_normal_task,
+    },
+    "train_bayesian": {
+        "task_name": "Train Bayesian Models",
+        "task_func": train_bayesian.train_bayesian_task,
+    },
+    "evaluate_regular": {
+        "task_name": "Evaluate Regular Models",
+        "task_func": dummy_task,
+    },
+    "evaluate_bayesian": {
+        "task_name": "Evaluate Bayesian Models",
+        "task_func": dummy_task,
+    },
+}
+
+SCENARIOS = {
+    "scen_train_all": {
+        "label": "1. Train, Evaluate, & Noise Analysis",
+        "tasks": [
+            "load_data",
+            "train_regular",
+            "train_bayesian",
+            "evaluate_regular",
+            "evaluate_bayesian",
+        ],
+    },
+    "scen_load_all": {
+        "label": "2. Load, Evaluate, & Noise Analysis",
+        "tasks": ["load_data", "evaluate_regular", "evaluate_bayesian"],
+    },
+    "scen_load_eval": {
+        "label": "3. Load & Evaluate (Skip Noise)",
+        "tasks": ["load_data", "evaluate_regular", "evaluate_bayesian"],
+    },
+}

+ 61 - 0
tasks/load_data.py

@@ -0,0 +1,61 @@
+import pathlib as pl
+from typing import Any, Dict
+
+import pandas as pd
+
+import model.dataset as ds
+from util.progress import ProgressTracker
+from util.ui_logger import PipelineLogger
+
+
+def load_data_task(
+    track: ProgressTracker,
+    log: PipelineLogger,
+    config: Dict[str, Any],
+    state: Dict[str, Any],
+):
+
+    log.info("Loading Files")
+    mri_files = pl.Path(config["data"]["mri_files_path"]).glob("*.nii")
+    xls_file = pl.Path(config["data"]["xls_file_path"])
+
+    dataset = ds.load_adni_data_from_file(
+        mri_files,
+        xls_file,
+        device=config["training"]["device"],
+        xls_preprocessor=ds.xls_pre,
+    )
+
+    if config["data"]["seed"] is None:
+        log.info("Seed is not defined - using default seed of 0")
+        config["data"]["seed"] = 0
+
+    ptid_df = pd.read_csv(xls_file)
+    ptid_df.columns = ptid_df.columns.str.strip()
+
+    ptid_df = ptid_df[["Image Data ID", "PTID"]].dropna(  # type: ignore
+        subset=["Image Data ID", "PTID"]
+    )
+    ptid_df["Image Data ID"] = ptid_df["Image Data ID"].astype(int)
+    ptid_df["PTID"] = ptid_df["PTID"].astype(str).str.strip()
+    ptid_df = ptid_df[ptid_df["PTID"] != ""]
+
+    ptids = list(zip(ptid_df["Image Data ID"].tolist(), ptid_df["PTID"].tolist()))
+
+    # Split is grouped by PTID to prevent patient-level leakage across partitions.
+    datasets = ds.divide_dataset_by_patient_id(
+        dataset,
+        ptids,
+        config["data"]["data_splits"],
+        seed=config["data"]["seed"],
+    )
+
+    # Initialize the dataloaders
+    train_loader, val_loader, test_loader = ds.initalize_dataloaders(
+        datasets, batch_size=config["training"]["batch_size"]
+    )
+
+    log.info("Dataloaders initalized")
+    state["train_loader"] = train_loader
+    state["val_loader"] = val_loader
+    state["test_loader"] = test_loader

+ 140 - 0
tasks/train_bayesian.py

@@ -0,0 +1,140 @@
+import pathlib as pl
+import time
+import gc
+from threading import Event
+from typing import Any, Dict
+
+import torch
+from bayesian_torch.models.dnn_to_bnn import dnn_to_bnn, get_kl_loss  # type: ignore[import-untyped]
+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_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)
+
+    bnn_prior_parameters = {
+        "prior_mu": 0.0,
+        "prior_sigma": 1.0,
+        "posterior_mu_init": 0.0,
+        "posterior_rho_init": -3.0,
+        "type": "Reparameterization",
+        "moped_enable": False,
+        "moped_delta": 0.5,
+    }
+
+    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']}..."
+        )
+        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"])
+        )
+        dnn_to_bnn(model, bnn_prior_parameters)
+
+        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']} - "
+            f"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.")

+ 131 - 0
tasks/train_normal.py

@@ -0,0 +1,131 @@
+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.")

+ 64 - 0
util/config_manager.py

@@ -0,0 +1,64 @@
+import os
+import shutil
+from typing import Any, Dict, Tuple
+
+import toml
+
+DEFAULT_CONFIG_PATH = "./config.toml"
+
+
+def load_config(filepath: str) -> Dict[str, Any]:
+    """Reads and parses a TOML configuration file."""
+    with open(filepath, "r") as f:
+        return toml.load(f)
+
+
+def save_config(filepath: str, config: Dict[str, Any]) -> None:
+    """Saves a dictionary to a TOML configuration file."""
+    os.makedirs(os.path.dirname(filepath), exist_ok=True)
+    with open(filepath, "w") as f:
+        toml.dump(config, f)
+
+
+def handle_directory_config(work_dir: str) -> Tuple[bool, str, Dict[str, Any]]:
+    """
+    Checks the directory, loads an existing config, or copies the default one.
+
+    Returns:
+        Tuple containing: (Success boolean, Status message, Configuration Dictionary)
+    """
+    if not work_dir:
+        return False, "Please specify a working directory.", {}
+
+    config_path = os.path.join(work_dir, "config.toml")
+
+    # 1. If config exists, load it
+    if os.path.exists(config_path):
+        try:
+            conf = load_config(config_path)
+            return True, f"Successfully loaded config from {config_path}", conf
+        except Exception as e:
+            return False, f"Failed to load config: {str(e)}", {}
+
+    # 2. If it doesn't exist, create directory and copy default config
+    else:
+        try:
+            os.makedirs(work_dir, exist_ok=True)
+
+            if not os.path.exists(DEFAULT_CONFIG_PATH):
+                return (
+                    False,
+                    f"Default config not found at {DEFAULT_CONFIG_PATH} to copy!",
+                    {},
+                )
+
+            shutil.copy(DEFAULT_CONFIG_PATH, config_path)
+            conf = load_config(config_path)
+            return (
+                True,
+                f"Copied default config to {config_path}. Ready for editing!",
+                conf,
+            )
+
+        except Exception as e:
+            return False, f"Failed to initialize directory/config: {str(e)}", {}