load_data.py 1.8 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061
  1. import pathlib as pl
  2. from typing import Any, Dict
  3. import pandas as pd
  4. import model.dataset as ds
  5. from util.progress import ProgressTracker
  6. from util.ui_logger import PipelineLogger
  7. def load_data_task(
  8. track: ProgressTracker,
  9. log: PipelineLogger,
  10. config: Dict[str, Any],
  11. state: Dict[str, Any],
  12. ):
  13. log.info("Loading Files")
  14. mri_files = pl.Path(config["data"]["mri_files_path"]).glob("*.nii")
  15. xls_file = pl.Path(config["data"]["xls_file_path"])
  16. dataset = ds.load_adni_data_from_file(
  17. mri_files,
  18. xls_file,
  19. device=config["training"]["device"],
  20. xls_preprocessor=ds.xls_pre,
  21. )
  22. if config["data"]["seed"] is None:
  23. log.info("Seed is not defined - using default seed of 0")
  24. config["data"]["seed"] = 0
  25. ptid_df = pd.read_csv(xls_file)
  26. ptid_df.columns = ptid_df.columns.str.strip()
  27. ptid_df = ptid_df[["Image Data ID", "PTID"]].dropna( # type: ignore
  28. subset=["Image Data ID", "PTID"]
  29. )
  30. ptid_df["Image Data ID"] = ptid_df["Image Data ID"].astype(int)
  31. ptid_df["PTID"] = ptid_df["PTID"].astype(str).str.strip()
  32. ptid_df = ptid_df[ptid_df["PTID"] != ""]
  33. ptids = list(zip(ptid_df["Image Data ID"].tolist(), ptid_df["PTID"].tolist()))
  34. # Split is grouped by PTID to prevent patient-level leakage across partitions.
  35. datasets = ds.divide_dataset_by_patient_id(
  36. dataset,
  37. ptids,
  38. config["data"]["data_splits"],
  39. seed=config["data"]["seed"],
  40. )
  41. # Initialize the dataloaders
  42. train_loader, val_loader, test_loader = ds.initalize_dataloaders(
  43. datasets, batch_size=config["training"]["batch_size"]
  44. )
  45. log.info("Dataloaders initalized")
  46. state["train_loader"] = train_loader
  47. state["val_loader"] = val_loader
  48. state["test_loader"] = test_loader