| 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() |
| |
| |
| 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] |
| |
| |
| 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." |
| |
| |
| 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) |
| |
| |
| examples = [] |
| |
| for batch in dataloader: |
| examples.extend(batch) |
| if len(examples) == num_samples: |
| break |
| |
| |
| 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") |
| |
| |
| 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. |
| """ |
| |
| |
| train_transforms, val_transforms = get_image_transforms(do_augmentation=do_augmentation) |
| |
| |
| 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_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 |
| |