File size: 2,001 Bytes
a23d562
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from matplotlib import pyplot as plt
import torchio as tio

from auto_detect_breast_mri.data.metadata import get_uka_metatensor
from auto_detect_breast_mri.data.transforms import RandomCropOrPad, ZNormalization, ImageToTensor
from auto_detect_breast_mri.data.breast_mri_dataset import BreastMRISubjects
from auto_detect_breast_mri.data.loaders import get_multiple_subjects_dataloader
from auto_detect_breast_mri.config import resolve_path

# Paths come from the site config (see config.example.yaml).
path_base = resolve_path(None, "data_root", "root folder of the NIfTI data")
subset_path = resolve_path(None, "split_root", "folder holding the split files")
subset_file = "fold0/stratified_training_set-f0-0.csv"
feature_path = resolve_path(None, "metadata_file", "metadata export")
pre_image_shape = (32, 512, 512)
protocol = ['Sub_1']
train_prop = 0.7
train_fraction = 0.5
batch_size = 4
transform = tio.Compose([
    tio.RandomFlip((0,1,2), flip_probability=0.5),
    RandomCropOrPad(pre_image_shape),
    ZNormalization(per_channel=True, percentiles=(0.5, 99.5), masking_method=lambda x:x>0),
    tio.RandomNoise(std=(0.25,0.5)),
    ImageToTensor()
])

data_set = BreastMRISubjects(path_base, subset_path + subset_file, protocol=protocol, transform=transform)

feature_dataframe = get_uka_metatensor(0, feature_path)
train_loader, eval_loader, test_loader = get_multiple_subjects_dataloader(path_base, feature_dataframe,
                                                                         pre_image_shape, transform, protocol,
                                                                         subset_path, batch_size, stratified=True,
                                                                         fraction=train_fraction, fold=0, subfold=0)

for i, batch in enumerate(train_loader):
    image = batch['image']['data']
    label = batch['label']
    plt.imshow(image[0, 0, :, :, 14], cmap='gray', interpolation=None)
    plt.show()
    print(label[0])
    print("-------")