load_data.py 3.3 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192939495
  1. import pathlib as pl
  2. from typing import Any, Dict, List, Tuple
  3. import pandas as pd
  4. import model.dataset as ds
  5. from util.progress import ProgressTracker
  6. from util.seeding import seed_everything
  7. from util.ui_logger import PipelineLogger
  8. def seed_pipeline(config: Dict[str, Any], log: PipelineLogger) -> int:
  9. """Resolve the base seed (default 0) and seed all RNGs once. Returns it."""
  10. if config["data"].get("seed") is None:
  11. log.info("Seed is not defined - using default seed of 0")
  12. config["data"]["seed"] = 0
  13. base_seed = int(config["data"]["seed"])
  14. deterministic = bool(config.get("training", {}).get("deterministic", False))
  15. applied = seed_everything(base_seed, deterministic=deterministic)
  16. log.info(
  17. f"Seeded RNGs with base seed {applied}"
  18. + (" (deterministic cuDNN)" if deterministic else "")
  19. )
  20. return base_seed
  21. def load_full_dataset(
  22. config: Dict[str, Any], log: PipelineLogger
  23. ) -> Tuple["ds.ADNIDataset", List[Tuple[int, str]], Dict[int, str]]:
  24. """Load the full ADNI dataset and its patient mapping (no split).
  25. Shared by the single-split loader and the k-fold driver, which partition the
  26. returned dataset differently.
  27. """
  28. log.info("Loading Files")
  29. mri_files = pl.Path(config["data"]["mri_files_path"]).glob("*.nii")
  30. xls_file = pl.Path(config["data"]["xls_file_path"])
  31. dataset = ds.load_adni_data_from_file(
  32. mri_files,
  33. xls_file,
  34. device=config["training"]["device"],
  35. xls_preprocessor=ds.xls_pre,
  36. )
  37. ptid_df = pd.read_csv(xls_file)
  38. ptid_df.columns = ptid_df.columns.str.strip()
  39. ptid_df = ptid_df[["Image Data ID", "PTID"]].dropna( # type: ignore
  40. subset=["Image Data ID", "PTID"]
  41. )
  42. ptid_df["Image Data ID"] = ptid_df["Image Data ID"].astype(int)
  43. ptid_df["PTID"] = ptid_df["PTID"].astype(str).str.strip()
  44. ptid_df = ptid_df[ptid_df["PTID"] != ""]
  45. ptids = list(zip(ptid_df["Image Data ID"].tolist(), ptid_df["PTID"].tolist()))
  46. image_to_ptid = {int(iid): str(pid) for iid, pid in ptids}
  47. return dataset, ptids, image_to_ptid
  48. def load_data_task(
  49. track: ProgressTracker,
  50. log: PipelineLogger,
  51. config: Dict[str, Any],
  52. state: Dict[str, Any],
  53. ):
  54. # Seed everything once at the start of the pipeline so data loading, weight
  55. # init, and shuffling are reproducible. Per-member seeds are derived from
  56. # this base inside the training tasks.
  57. base_seed = seed_pipeline(config, log)
  58. dataset, ptids, image_to_ptid = load_full_dataset(config, log)
  59. # Mapping used later by evaluation to attach patient ids to the sample axis.
  60. state["image_to_ptid"] = image_to_ptid
  61. # Split is grouped by PTID to prevent patient-level leakage across partitions.
  62. datasets = ds.divide_dataset_by_patient_id(
  63. dataset,
  64. ptids,
  65. config["data"]["data_splits"],
  66. seed=config["data"]["seed"],
  67. )
  68. # Initialize the dataloaders. The train shuffle generator is seeded from the
  69. # base seed so shuffling is reproducible; val/test are left unshuffled.
  70. train_loader, val_loader, test_loader = ds.initalize_dataloaders(
  71. datasets,
  72. batch_size=config["training"]["batch_size"],
  73. seed=base_seed,
  74. )
  75. log.info("Dataloaders initalized")
  76. state["train_loader"] = train_loader
  77. state["val_loader"] = val_loader
  78. state["test_loader"] = test_loader