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