load_data.py 2.5 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576
  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. # Split is grouped by PTID to prevent patient-level leakage across partitions.
  46. datasets = ds.divide_dataset_by_patient_id(
  47. dataset,
  48. ptids,
  49. config["data"]["data_splits"],
  50. seed=config["data"]["seed"],
  51. )
  52. # Initialize the dataloaders. The train shuffle generator is seeded from the
  53. # base seed so shuffling is reproducible; val/test are left unshuffled.
  54. train_loader, val_loader, test_loader = ds.initalize_dataloaders(
  55. datasets,
  56. batch_size=config["training"]["batch_size"],
  57. seed=base_seed,
  58. )
  59. log.info("Dataloaders initalized")
  60. state["train_loader"] = train_loader
  61. state["val_loader"] = val_loader
  62. state["test_loader"] = test_loader