| 12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879 |
- 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.seeding import seed_everything
- from util.ui_logger import PipelineLogger
- def load_data_task(
- track: ProgressTracker,
- log: PipelineLogger,
- config: Dict[str, Any],
- state: Dict[str, Any],
- ):
- if config["data"]["seed"] is None:
- log.info("Seed is not defined - using default seed of 0")
- config["data"]["seed"] = 0
- # Seed everything once at the start of the pipeline so data loading, weight
- # init, and shuffling are reproducible. Per-member seeds are derived from
- # this base inside the training tasks.
- base_seed = int(config["data"]["seed"])
- deterministic = bool(config.get("training", {}).get("deterministic", False))
- applied = seed_everything(base_seed, deterministic=deterministic)
- log.info(
- f"Seeded RNGs with base seed {applied}"
- + (" (deterministic cuDNN)" if deterministic else "")
- )
- 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,
- )
- 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()))
- # Mapping used later by evaluation to attach patient ids to the sample axis.
- state["image_to_ptid"] = {int(iid): str(pid) for iid, pid in ptids}
- # 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. The train shuffle generator is seeded from the
- # base seed so shuffling is reproducible; val/test are left unshuffled.
- train_loader, val_loader, test_loader = ds.initalize_dataloaders(
- datasets,
- batch_size=config["training"]["batch_size"],
- seed=base_seed,
- )
- log.info("Dataloaders initalized")
- state["train_loader"] = train_loader
- state["val_loader"] = val_loader
- state["test_loader"] = test_loader
|