File size: 6,657 Bytes
64fd08f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 | import pandas as pd
import numpy as np
from tqdm import tqdm
from pathlib import Path
import torch
import torchio as tio
import SimpleITK as sitk
from typing import List, Tuple
from utils.transforms import get_image_transforms
def create_subjects_list(df: pd.DataFrame, all_scans: bool = False, finetune_label: str = None, dataset: str = "odelia", debug_subset: bool = False) -> List[tio.Subject]:
"""
Create a list of subjects from a dataframe of file paths and labels.
"""
subjects = []
patient_ids = df["patient_id"].unique().tolist()
# Sample 5 patients from each class in debug mode
if debug_subset:
if finetune_label is not None:
patient_ids = []
for cls in df[finetune_label].dropna().unique().tolist():
cls_patients = df[df[finetune_label] == cls]["patient_id"].unique()
n_select = min(5, len(cls_patients))
patient_ids.extend(np.random.choice(cls_patients, size=n_select, replace=False).tolist())
else:
patient_ids = patient_ids[:10]
if dataset == "odelia":
odelia_mapper = {"normal": 0, "benign": 1, "malignant": 2}
for patient_id in tqdm(patient_ids, total=len(patient_ids)):
pt_df = df[df["patient_id"] == patient_id]
# Get all breast volumes per patient
for side in pt_df["side"].unique().tolist():
side_df = pt_df[pt_df["side"] == side]
side_df = side_df.sort_values(by="filepath")
if finetune_label == "breast_label":
side_label = torch.tensor(odelia_mapper[side_df["breast_label"].values[0]], dtype=torch.float32)
else:
side_label = "n.a."
# Use a single scan per patient
if not all_scans:
side_images = [tio.ScalarImage(side_df["filepath"].values[0])]
side_image_keys = ["image_1"]
side_subject_dict = tio.Subject(name=f"{patient_id}_{side}", **{key: value for key, value in zip(side_image_keys, side_images)}, label=side_label)
subjects.append(side_subject_dict)
return subjects
def run_sanity_check_single_image_dataloader(dataloader: tio.data.SubjectsLoader, num_samples: int, save_dir: Path):
"""
Sanity check the dataloader.
"""
save_dir.mkdir(parents=True, exist_ok=True)
# Extract samples
examples = []
for batch in dataloader:
examples.extend(batch)
if len(examples) == num_samples:
break
# Save them to visually inspect train transforms
for example in examples:
image = sitk.GetImageFromArray(example["image_1"][tio.DATA].squeeze(0).numpy())
label = int(example["label"].numpy())
sitk.WriteImage(image, save_dir.joinpath(f"{example['name']}_(gt={label}).nii.gz"))
return
def create_balanced_sampler(subjects, num_classes: int = 3):
"""
Create a balanced sampler to handle class imbalance.
"""
labels = [subject["label"].item() for subject in subjects]
unique, counts = np.unique(labels, return_counts=True)
if len(unique) != num_classes:
print(f"WARNING: Number of classes ({len(unique)}) does not match expected number of classes ({num_classes}) in balanced sampler")
# Calculate weights for each sample
class_weights = {label: 1.0 / count for label, count in zip(unique, counts)}
sample_weights = [class_weights[label] for label in labels]
sampler = torch.utils.data.WeightedRandomSampler(
weights=sample_weights,
num_samples=len(sample_weights),
replacement=True
)
return sampler
def prepare_batch_single_scan(batch: List[tio.Subject]) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Helper function to prepare a batch of subjects for training.
"""
x = batch["image_1"][tio.DATA].to(device='cuda', dtype=torch.float16, non_blocking=True)
y = batch["label"].to(device='cuda', dtype=torch.long, non_blocking=True)
return x, y
def prepare_loaders(
df: pd.DataFrame,
do_augmentation: bool,
mode: str,
batch_size: int,
all_scans: bool,
max_num_scans: int,
finetune_label: str = None,
debug_subset: bool = False,
balanced_sampling: bool = False,
num_classes: int = None,
val_fold: int = 0
):
"""
Helper function to prepare data loaders.
"""
# Get transforms
train_transforms, val_transforms = get_image_transforms(do_augmentation=do_augmentation)
# Get data loaders
print(f"Preparing data loaders...")
all_folds = [0, 1, 2, 3, 4]
train_folds = [i for i in all_folds if i != val_fold]
train_subjects = create_subjects_list(df[df["fold"].isin(train_folds)], all_scans=all_scans, finetune_label=finetune_label, debug_subset=debug_subset)
train_subjects_names = [i["name"] for i in train_subjects]
train_dataset = tio.data.SubjectsDataset(train_subjects, transform=train_transforms, load_getitem=False)
if balanced_sampling and num_classes is not None:
train_sampler = create_balanced_sampler(train_subjects, num_classes=num_classes)
train_loader = tio.data.SubjectsLoader(train_dataset, batch_size=batch_size, sampler=train_sampler, num_workers=6, prefetch_factor=4, persistent_workers=True)
else:
train_loader = tio.data.SubjectsLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=6, prefetch_factor=4, persistent_workers=True)
val_subjects = create_subjects_list(df[df["fold"].isin([val_fold])], all_scans=all_scans, finetune_label=finetune_label, debug_subset=debug_subset)
val_dataset = tio.data.SubjectsDataset(val_subjects, transform=val_transforms, load_getitem=False)
val_loader = tio.data.SubjectsLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=6, prefetch_factor=4, persistent_workers=True)
# Sanity loader to fetch attention weights from some training and val cases
sanity_subjects = np.random.choice(train_subjects, size=5, replace=False).tolist() + np.random.choice(val_subjects, size=5, replace=False).tolist()
sanity_dataset = tio.data.SubjectsDataset(sanity_subjects, transform=train_transforms, load_getitem=False)
sanity_loader = tio.data.SubjectsLoader(sanity_dataset, batch_size=batch_size, shuffle=False, num_workers=6, prefetch_factor=4, persistent_workers=True)
return train_loader, val_loader, sanity_loader, train_subjects_names
|