Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
166 changes: 149 additions & 17 deletions src/icon_registration/unicarl/dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,48 +30,88 @@ def __init__(
image_glob: str,
cache_filename=None,
maximum_images=None,
shuffle=False
shuffle=False,
world_size=1,
world_rank=0,
shard_threshold=5000
):
print(name)
self.name = name
self.image_glob = image_glob
self.input_shape = input_shape
self.world_size = world_size
self.world_rank = world_rank
self.shard_threshold = shard_threshold

if not cache_filename:
self.store = {}
paths = self.get_image_paths()
all_paths = self.get_image_paths()
if shuffle:
random.shuffle(paths)
random.shuffle(all_paths)
if maximum_images:
paths = paths[:maximum_images]
all_paths = all_paths[:maximum_images]

# Let subclass decide which paths to keep for this rank
paths = self._select_paths_for_rank(all_paths)

for path in tqdm.tqdm(paths):
try:
self.store[path] = self.preprocess_itk_image(path)
except Exception as e: # (IndexError, ValueError, itk.TemplateTypeError) as e:
print(e)

cache_suffix = self._get_cache_suffix()
torch.save(
{
"name": self.name,
"image_glob": self.image_glob,
"maximum_images": maximum_images,
"store": self.store,
"world_size": self.world_size,
"world_rank": self.world_rank,
"shard_threshold": self.shard_threshold,
},
footsteps.output_dir + self.name + "_cached_dataset.trch",
footsteps.output_dir + self.name + cache_suffix + "_cached_dataset.trch",
)
else:
loaded_cache = torch.load(
cache_filename + "/" + self.name + "_cached_dataset.trch",
weights_only = False
)
assert self.name == loaded_cache["name"]
assert maximum_images == loaded_cache["maximum_images"]
assert self.image_glob == loaded_cache["image_glob"]
paths = self.get_image_paths()
cache_suffix = self._get_cache_suffix()
cache_path = cache_filename + "/" + self.name + cache_suffix + "_cached_dataset.trch"
loaded_cache = torch.load(cache_path, weights_only=False)

# Validate cache metadata
assert self.name == loaded_cache["name"], f"Dataset name mismatch: {self.name} != {loaded_cache['name']}"
assert self.image_glob == loaded_cache["image_glob"], f"Image glob mismatch for {self.name}"

# Validate distributed configuration matches
cached_world_size = loaded_cache.get("world_size")
cached_world_rank = loaded_cache.get("world_rank")
if cached_world_size != self.world_size or cached_world_rank != self.world_rank:
raise ValueError(
f"Cache distributed config mismatch for {self.name}: "
f"expected (world_size={self.world_size}, rank={self.world_rank}), "
f"got (world_size={cached_world_size}, rank={cached_world_rank})"
)

# Load the preprocessed store (no need to re-glob original files)
self.store = loaded_cache["store"]
# assert(paths[0] in self.store) # sanity check
self.keys = list(self.store.keys())
print("Image count: ", len(self.keys))
if self.world_size > 1:
print(f"Image count: {len(self.keys)} (rank {self.world_rank}/{self.world_size})")
else:
print(f"Image count: {len(self.keys)}")

def _get_cache_suffix(self):
"""Generate cache filename suffix based on distributed config"""
if self.world_size <= 1:
return ""
return f"_rank{self.world_rank}_of_{self.world_size}"

def _select_paths_for_rank(self, all_paths):
"""Select which paths this rank should process. Override in subclasses for custom sharding logic."""
should_shard = self.world_size > 1 and len(all_paths) > self.shard_threshold
if should_shard:
return all_paths[self.world_rank::self.world_size]
return all_paths

def get_image_paths(self) -> [str]:
return list(glob.glob(self.image_glob))
Expand Down Expand Up @@ -155,17 +195,29 @@ def __init__(
cache_filename=None,
maximum_images=None,
match_regex=None,
world_size=1,
world_rank=0,
shard_threshold=5000,
):
if match_regex == None:
raise NotImplementedError()

# Store match_regex BEFORE calling super().__init__()
# so it's available in _select_paths_for_rank()
self.match_regex = match_regex

super().__init__(
input_shape,
name,
image_glob,
cache_filename=cache_filename,
maximum_images=maximum_images,
world_size=world_size,
world_rank=world_rank,
shard_threshold=shard_threshold,
)
if match_regex == None:
raise NotImplementedError()

# Build pair lookup from store
self.pair_lookup = collections.defaultdict(lambda: [])
self.pair_keys = {}

Expand All @@ -176,6 +228,25 @@ def __init__(
# for pair_key in self.pair_lookup.keys():
# assert len(self.pair_lookup[pair_key]) != 1 , f"{self.pair_lookup[pair_key]}"

def _select_paths_for_rank(self, all_paths):
"""Shard by groups/patients instead of individual images to keep pairs together."""
should_shard = self.world_size > 1 and len(all_paths) > self.shard_threshold
if not should_shard:
return all_paths

groups = collections.defaultdict(list)
for path in all_paths:
group_id = regex.search(self.match_regex, path).group(1)
groups[group_id].append(path)

group_ids = sorted(groups.keys())
my_groups = group_ids[self.world_rank::self.world_size]

my_paths = []
for group_id in my_groups:
my_paths.extend(groups[group_id])
return my_paths

def get_key_pair(self):
image_key_1 = random.choice(self.keys)
image_key_2 = image_key_1
Expand Down Expand Up @@ -328,3 +399,64 @@ def read_image(self, path:str):
spacing = np.array((1, 1, 1))
print(spacing)
return image[0], spacing


# With apologies, someone better at OOP can refactor this hiearchy into a DAG or some cursed thing to prevent this
# from being a copy paste of PairedDICOMDataset
class DICOMDataset(Dataset):
def read_image(self, path: str):
print(path)
"""
Reads a DICOM series from a directory path and returns it as a tensor.

Args:
path (str): Directory containing DICOM files
e.g., "files/image342/"

Returns:
torch.Tensor: 3D tensor containing the DICOM volume
"""
# import SimpleITK as sitk
import os

namesGenerator = itk.GDCMSeriesFileNames.New()
namesGenerator.SetUseSeriesDetails(True)
namesGenerator.SetDirectory(path)
seriesUID = namesGenerator.GetSeriesUIDs()

dicom_files = namesGenerator.GetFileNames(seriesUID[0])

# Read the DICOM series as a 3D image
reader = itk.ImageSeriesReader[itk.Image[itk.SS, 3]].New()
dicomIO = itk.GDCMImageIO.New()
reader.SetImageIO(dicomIO)
reader.SetFileNames(dicom_files)
reader.Update()
image = reader.GetOutput()
image = reorient(image)

if (
"ITK_non_uniform_sampling_deviation"
in image.GetMetaDataDictionary().GetKeys()
):
spacing_deviation = image.GetMetaDataDictionary().Get(
"ITK_non_uniform_sampling_deviation"
)
spacing_deviation = (
itk.MetaDataObject[itk.D]
.cast(spacing_deviation)
.GetMetaDataObjectValue()
)

if spacing_deviation > 5:
raise ValueError("image has non-uniform-spacing: likely a mish-mash")

# Convert to tensor
image_array = itk.GetArrayFromImage(image)

image_tensor = torch.tensor(image_array)

if np.any(np.array(image_array.shape) < 20):
raise ValueError("image too low resolution")

return image_tensor, np.array(image.GetSpacing())[::-1]
76 changes: 43 additions & 33 deletions training_scripts/unicarl/datasets.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,55 +2,65 @@
import torch
import footsteps
import matplotlib.pyplot as plt
import os

# If you are in uncbiag running this script to preprocess the unicarl Dataset from scratch
# First, sorry.
# Second:
# run it on biag-w05 where the autoPET dataset is
# run it on biag-w05 where the autoPET dataset is
# do a sshfs biag-lambda2:/data /data to access the abdomen8k data on that maching

input_shape = (1, 1, 160, 160, 160)

public = False

public = True

datasets_ = []

#datasets = lambda: None
#datasets.append = lambda x: None

maximum_images=1000

cache_filename = "results/unicarl_private/"
# Common kwargs for all datasets
dataset_kwargs = {
'maximum_images': 10000000000,
'cache_filename': 'unicarl_public_train/',
'world_size': int(os.environ.get("WORLD_SIZE", 1)),
'world_rank': int(os.environ.get("RANK", 0)),
'shard_threshold': 500,
}

datasets_.append(dataset.PairedDataset(input_shape, "HAN-Seg",
"/playpen-raid1/Data/HaN-Seg/HaN-Seg/set_1/case_??/case_??_IMG_*.nrrd", match_regex=r"/(case_[0-9]*)/", maximum_images=maximum_images, cache_filename=cache_filename) )
datasets_.append(dataset.DiffusionDataset( input_shape, "ebrahim-diffusion", "/playpen-raid1/tgreer/ebrahim_brains/data/degree_powers_normalized_dipy/degree_power_images/*", maximum_images=maximum_images, cache_filename=cache_filename))
datasets_.append(dataset.PairedDICOMDataset(input_shape, "CPTAC-UCEC", "/playpen-raid1/Data/TCIA_CPTAC-UCEC/manifest-1712342731330/CPTAC-UCEC/*/*/*/", match_regex=r"/(C3[NL]-[0-9]*)/", maximum_images=maximum_images * 4, cache_filename=cache_filename) )
datasets_.append(dataset.PairedDICOMDataset(input_shape, "TCIA-hastings", "/playpen-raid1/Data/TCIA_Hastings_custom_mrct/manifest-1743108366953/*/*/*/*/", match_regex=r"66953/[a-zA-Z\-]*/([A-Z0-9\-]+)/", maximum_images=maximum_images * 4, cache_filename=cache_filename))
datasets_.append(dataset.PairedDICOMDataset(input_shape, "CPTAC-Sarcoma",
"/playpen-raid1/Data/TCIA-Sarcoma/manifest-MjbMt99Q1553106146386120388/Soft-tissue-Sarcoma/*/*/?.*/", match_regex=r"/(STS_[0-9]*)/", maximum_images=maximum_images*4, cache_filename=cache_filename) )
datasets_.append(dataset.PairedDataset(input_shape, "anatomixX",
"/playpen-raid1/tgreer/anatomix/anatomix/synthetic-data-generation/synthesized_views/view*/*X.nii.gz", match_regex=r"([0-9A-Z]+).nii.gz", **dataset_kwargs) )
datasets_.append(dataset.PairedDataset(input_shape, "anatomixY",
"/playpen-raid1/tgreer/anatomix/anatomix/synthetic-data-generation/synthesized_views/view*/*Y.nii.gz", match_regex=r"([0-9A-Z]+).nii.gz", **dataset_kwargs) )
datasets_.append(dataset.PairedDataset(input_shape, "anatomixZ",
"/playpen-raid1/tgreer/anatomix/anatomix/synthetic-data-generation/synthesized_views/view*/*Z.nii.gz", match_regex=r"([0-9A-Z]+).nii.gz", **dataset_kwargs) )
datasets_.append(dataset.DICOMDataset(input_shape, "Pediatric-CT-Seg", "/playpen-raid1/Data/Pediatric-CT-SEG/*/*/*-CT*/", **dataset_kwargs))
datasets_.append(dataset.PairedDataset(input_shape, "HAN-Seg",
"/playpen-raid1/Data/HaN-Seg/HaN-Seg/set_1/case_??/case_??_IMG_*.nrrd", match_regex=r"/(case_[0-9]*)/", **dataset_kwargs) )
datasets_.append(dataset.DiffusionDataset( input_shape, "ebrahim-diffusion", "/playpen-raid1/tgreer/ebrahim_brains/data/degree_powers_normalized_dipy/degree_power_images/*", **dataset_kwargs))
datasets_.append(dataset.PairedDICOMDataset(input_shape, "CPTAC-UCEC", "/playpen-raid1/Data/TCIA_CPTAC-UCEC/manifest-1712342731330/CPTAC-UCEC/*/*/*/", match_regex=r"/(C3[NL]-[0-9]*)/", **dataset_kwargs) )
datasets_.append(dataset.PairedDICOMDataset(input_shape, "TCIA-hastings", "/playpen-raid1/Data/TCIA_Hastings_custom_mrct/manifest-1743108366953/*/*/*/*/", match_regex=r"66953/[a-zA-Z\-]*/([A-Z0-9\-]+)/", **dataset_kwargs))
datasets_.append(dataset.PairedDICOMDataset(input_shape, "CPTAC-Sarcoma",
"/playpen-raid1/Data/TCIA-Sarcoma/manifest-MjbMt99Q1553106146386120388/Soft-tissue-Sarcoma/*/*/?.*/", match_regex=r"/(STS_[0-9]*)/", **dataset_kwargs) )
if (not public):
datasets_.append(dataset.PairedDataset( input_shape, "pancreas", "/playpen-raid1/tgreer/pancreatic_cancer_registration/data/*/Processed/*/original_image.nii.gz", maximum_images=maximum_images, cache_filename=cache_filename, match_regex=r"data/([0-9]+)/Processed/"))
datasets_.append(dataset.PairedDataset( input_shape, "dirlab_clamped", "/playpen-raid2/Data/Lung_Registration_clamp_normal_transposed/*/*_img.nii.gz", maximum_images=maximum_images, cache_filename=cache_filename, match_regex=r'transposed/([a-zA-Z0-9]+)/'))
datasets_.append(dataset.PairedDataset( input_shape, "dirlab", "/playpen-raid2/Data/Lung_Registration_transposed/*/*_img.nii.gz", maximum_images=maximum_images, cache_filename=cache_filename, match_regex=r'transposed/([a-zA-Z0-9]+)/'))
datasets_.append(dataset.Dataset( input_shape, "HCP_t1_stripped", "/playpen-raid2/Data/HCP/HCP_1200/*/T1w/T1w_acpc_dc_restore_brain.nii.gz", maximum_images=maximum_images, cache_filename=cache_filename))
datasets_.append(dataset.Dataset( input_shape, "HCP_t1", "/playpen-raid2/Data/HCP/HCP_1200/*/T1w/T1w_acpc_dc_restore.nii.gz", maximum_images=maximum_images, cache_filename=cache_filename))
datasets_.append(dataset.Dataset( input_shape, "HCP_t2_stripped", "/playpen-raid2/Data/HCP/HCP_1200/*/T1w/T2w_acpc_dc_restore_brain.nii.gz", maximum_images=maximum_images, cache_filename=cache_filename))
datasets_.append(dataset.Dataset( input_shape, "HCP_t2", "/playpen-raid2/Data/HCP/HCP_1200/*/T1w/T2w_acpc_dc_restore.nii.gz", maximum_images=maximum_images, cache_filename=cache_filename))
datasets_.append(dataset.Dataset( input_shape, "OAI", "/playpen-raid/zhenlinx/Data/OAI_segmentation/Nifti_rescaled_LEFT/*_image.nii.gz", maximum_images=maximum_images, cache_filename=cache_filename))
datasets_.append(dataset.Dataset( input_shape, "translucence", "/playpen-raid1/tgreer/mouse_brain_translucence/data/auto_files_resampled/*", cache_filename=cache_filename, maximum_images=maximum_images))
datasets_.append(dataset.Dataset( input_shape, "abdomen8k", "/data/hastings/Abdomen8k/AbdomenAtlas/*/ct.nii.gz", maximum_images=maximum_images, cache_filename=cache_filename, shuffle=True))
datasets_.append(dataset.PairedDICOMDataset(input_shape, "DukeLivers", "/playpen-raid1/Data/DukeLivers/Segmentation/Segmentation/*/*/images/", maximum_images=maximum_images, match_regex=r"Segmentation/([0-9]+)/", cache_filename=cache_filename))
#datasets_.append(dataset.Dataset( input_shape, "TotalSegmentatorMRI", "/playpen-raid1/soumitri/data/TotalSegMRI/*/*/mri.nii.gz", maximum_images=maximum_images, cache_filename=cache_filename))
datasets_.append(dataset.PairedDataset( input_shape, "bratsreg", "/playpen-raid2/Data/BraTS-Reg/BraTSReg_Training_Data_v3/*/*.nii.gz", maximum_images=maximum_images, cache_filename=cache_filename, match_regex=r"v3/(BraTSReg_[0-9]+)/"))
datasets_.append(dataset.Dataset( input_shape, "abdomen1k", "/playpen-raid2/Data/AbdomenCT-1K/AbdomenCT-1K-ImagePart*/Case_*", maximum_images=maximum_images, cache_filename=cache_filename, shuffle=True))
datasets_.append(dataset.Dataset( input_shape, "fmost", "/playpen-raid2/Data/fMost/subject/*_red_mm_RSA.nii.gz", maximum_images=maximum_images, cache_filename=cache_filename))
datasets_.append(dataset.Dataset( input_shape, "oasis", "/playpen-raid2/Data/oasis/OASIS_OAS1_*_MR1/orig.nii.gz", maximum_images=maximum_images, cache_filename=cache_filename))
datasets_.append(dataset.Dataset( input_shape, "lumir", "/playpen-raid1/Data/LUMIR/imagesTr/*", maximum_images=maximum_images, cache_filename=cache_filename))
datasets_.append(dataset.PairedDataset( input_shape, "autoPET", "/playpen1/tgreer/PET/FDG-PET-CT-Lesions/*/*/[PC][TE][Tr]*.nii.gz", match_regex=r"(/PETCT_[0-9a-z]+/)", maximum_images=maximum_images, cache_filename=cache_filename))
datasets_.append(dataset.PairedDataset(input_shape, "anatomix",
"/playpen-raid1/tgreer/anatomix/anatomix/synthetic-data-generation/synthesized_views/view*/*Z.nii.gz", match_regex=r"([0-9A-Z]+).nii.gz", maximum_images=8 * maximum_images, cache_filename=cache_filename) )
datasets_.append(dataset.PairedDataset( input_shape, "pancreas", "/playpen-raid1/tgreer/pancreatic_cancer_registration/data/*/Processed/*/original_image.nii.gz", **dataset_kwargs, match_regex=r"data/([0-9]+)/Processed/"))
datasets_.append(dataset.PairedDataset( input_shape, "dirlab_clamped", "/playpen-raid2/Data/Lung_Registration_clamp_normal_transposed/*/*_img.nii.gz", **dataset_kwargs, match_regex=r'transposed/([a-zA-Z0-9]+)/'))
datasets_.append(dataset.PairedDataset( input_shape, "dirlab", "/playpen-raid2/Data/Lung_Registration_transposed/*/*_img.nii.gz", **dataset_kwargs, match_regex=r'transposed/([a-zA-Z0-9]+)/'))
datasets_.append(dataset.Dataset( input_shape, "HCP_t1_stripped", "/playpen-raid2/Data/HCP/HCP_1200/*/T1w/T1w_acpc_dc_restore_brain.nii.gz", **dataset_kwargs))
datasets_.append(dataset.Dataset( input_shape, "HCP_t1", "/playpen-raid2/Data/HCP/HCP_1200/*/T1w/T1w_acpc_dc_restore.nii.gz", **dataset_kwargs))
datasets_.append(dataset.Dataset( input_shape, "HCP_t2_stripped", "/playpen-raid2/Data/HCP/HCP_1200/*/T1w/T2w_acpc_dc_restore_brain.nii.gz", **dataset_kwargs))
datasets_.append(dataset.Dataset( input_shape, "HCP_t2", "/playpen-raid2/Data/HCP/HCP_1200/*/T1w/T2w_acpc_dc_restore.nii.gz", **dataset_kwargs))
datasets_.append(dataset.Dataset( input_shape, "OAI", "/playpen-raid/zhenlinx/Data/OAI_segmentation/Nifti_rescaled_LEFT/*_image.nii.gz", **dataset_kwargs))
datasets_.append(dataset.Dataset( input_shape, "translucence", "/playpen-raid1/tgreer/mouse_brain_translucence/data/auto_files_resampled/*", **dataset_kwargs))
datasets_.append(dataset.Dataset( input_shape, "abdomen8k", "/data/hastings/Abdomen8k/AbdomenAtlas/*/ct.nii.gz", **dataset_kwargs, shuffle=True))
datasets_.append(dataset.PairedDICOMDataset(input_shape, "DukeLivers", "/playpen-raid1/Data/DukeLivers/Segmentation/Segmentation/*/*/images/", match_regex=r"Segmentation/([0-9]+)/", **dataset_kwargs))
#datasets_.append(dataset.Dataset( input_shape, "TotalSegmentatorMRI", "/playpen-raid1/soumitri/data/TotalSegMRI/*/*/mri.nii.gz", **dataset_kwargs))
datasets_.append(dataset.PairedDataset( input_shape, "bratsreg", "/playpen-raid2/Data/BraTS-Reg/BraTSReg_Training_Data_v3/*/*.nii.gz", **dataset_kwargs, match_regex=r"v3/(BraTSReg_[0-9]+)/"))
datasets_.append(dataset.Dataset( input_shape, "abdomen1k", "/playpen-raid2/Data/AbdomenCT-1K/AbdomenCT-1K-ImagePart*/Case_*", **dataset_kwargs, shuffle=True))
datasets_.append(dataset.Dataset( input_shape, "fmost", "/playpen-raid2/Data/fMost/subject/*_red_mm_RSA.nii.gz", **dataset_kwargs))
datasets_.append(dataset.Dataset( input_shape, "oasis", "/playpen-raid2/Data/oasis/OASIS_OAS1_*_MR1/orig.nii.gz", **dataset_kwargs))
datasets_.append(dataset.Dataset( input_shape, "lumir", "/playpen-raid1/Data/LUMIR/imagesTr/*", **dataset_kwargs))
datasets_.append(dataset.PairedDataset( input_shape, "autoPET", "/playpen1/tgreer/PET/FDG-PET-CT-Lesions/*/*/[PC][TE][Tr]*.nii.gz", match_regex=r"(/PETCT_[0-9a-z]+/)", **dataset_kwargs))

import torch.nn.functional as F

Expand Down
Loading