load_data.py 2.6 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879
  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.seeding import seed_everything
  7. from util.ui_logger import PipelineLogger
  8. def load_data_task(
  9. track: ProgressTracker,
  10. log: PipelineLogger,
  11. config: Dict[str, Any],
  12. state: Dict[str, Any],
  13. ):
  14. if config["data"]["seed"] is None:
  15. log.info("Seed is not defined - using default seed of 0")
  16. config["data"]["seed"] = 0
  17. # Seed everything once at the start of the pipeline so data loading, weight
  18. # init, and shuffling are reproducible. Per-member seeds are derived from
  19. # this base inside the training tasks.
  20. base_seed = int(config["data"]["seed"])
  21. deterministic = bool(config.get("training", {}).get("deterministic", False))
  22. applied = seed_everything(base_seed, deterministic=deterministic)
  23. log.info(
  24. f"Seeded RNGs with base seed {applied}"
  25. + (" (deterministic cuDNN)" if deterministic else "")
  26. )
  27. log.info("Loading Files")
  28. mri_files = pl.Path(config["data"]["mri_files_path"]).glob("*.nii")
  29. xls_file = pl.Path(config["data"]["xls_file_path"])
  30. dataset = ds.load_adni_data_from_file(
  31. mri_files,
  32. xls_file,
  33. device=config["training"]["device"],
  34. xls_preprocessor=ds.xls_pre,
  35. )
  36. ptid_df = pd.read_csv(xls_file)
  37. ptid_df.columns = ptid_df.columns.str.strip()
  38. ptid_df = ptid_df[["Image Data ID", "PTID"]].dropna( # type: ignore
  39. subset=["Image Data ID", "PTID"]
  40. )
  41. ptid_df["Image Data ID"] = ptid_df["Image Data ID"].astype(int)
  42. ptid_df["PTID"] = ptid_df["PTID"].astype(str).str.strip()
  43. ptid_df = ptid_df[ptid_df["PTID"] != ""]
  44. ptids = list(zip(ptid_df["Image Data ID"].tolist(), ptid_df["PTID"].tolist()))
  45. # Mapping used later by evaluation to attach patient ids to the sample axis.
  46. state["image_to_ptid"] = {int(iid): str(pid) for iid, pid in ptids}
  47. # Split is grouped by PTID to prevent patient-level leakage across partitions.
  48. datasets = ds.divide_dataset_by_patient_id(
  49. dataset,
  50. ptids,
  51. config["data"]["data_splits"],
  52. seed=config["data"]["seed"],
  53. )
  54. # Initialize the dataloaders. The train shuffle generator is seeded from the
  55. # base seed so shuffling is reproducible; val/test are left unshuffled.
  56. train_loader, val_loader, test_loader = ds.initalize_dataloaders(
  57. datasets,
  58. batch_size=config["training"]["batch_size"],
  59. seed=base_seed,
  60. )
  61. log.info("Dataloaders initalized")
  62. state["train_loader"] = train_loader
  63. state["val_loader"] = val_loader
  64. state["test_loader"] = test_loader