File size: 6,169 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
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
"""
Visual inspection of single transforms on the NIfTI files of one examination.

    python scripts/data_utils/inspect_transforms.py <path/to/one/examination_folder>

Opens matplotlib windows; not a test, despite where this file used to live.
"""

import argparse
import os

import torch
import torchio as tio
import matplotlib.pyplot as plt
import torchvision
import pandas as pd

from auto_detect_breast_mri.config import resolve_path
from auto_detect_breast_mri.data.breast_mri_dataset import BreastMRISubjects
from auto_detect_breast_mri.data.nifti_io import read_nifti


def test_flips():
    flip_3d = tio.Compose([tio.RandomFlip(['LR', 'IS', 'AP'], flip_probability=0.5)])  # NOTE: here AP does not refer to abbreviated protocol (as stated in README) but to Anterior-Posterior
    flip_2d = tio.Compose([tio.RandomFlip(['LR', 'IS'], flip_probability=0.5)])

    print("2D Flip as done until now.")
    flipped3 = flip_3d(dyn0)
    for i in range(0, nifti.shape[0]):
        if flipped3[0, i, :, :].sum() == 0:
            print("slice {} became zero.".format(i))
        if i % step == 0:
            plt.imshow(flipped3[0][i], cmap='gray', interpolation=None)
            plt.show()

    plt.imshow(seperation[0], cmap='gray', interpolation=None)
    plt.show()
    print("3D Flip.")
    flipped2 = flip_2d(dyn0)
    for i in range(0, nifti.shape[0]):
        if flipped2[0, i, :, :].sum() == 0:
            print("slice {} became zero.".format(i))
        if i % step == 0:
            plt.imshow(flipped2[0][i], cmap='gray', interpolation=None)
            plt.show()


def test_rotation():
    rotation = tio.Compose([torchvision.transforms.RandomRotation(35)])

    print("Random Rotation --> correct axis?")
    rotated = rotation(dyn0)
    for i in range(0, nifti.shape[0], step):
        plt.imshow(rotated[0][i], cmap='gray', interpolation=None)
        plt.show()


def test_rescale():
    rescale = tio.Compose([tio.RandomAffine(scales=(0.8, 1.2), degrees=0)])

    print("Check rescale")
    rescaled = rescale(dyn0)
    for i in range(0, nifti.shape[0], step):
        plt.imshow(rescaled[0][i], cmap='gray', interpolation=None)
        plt.show()


def test_normalize():
    norm = tio.ZNormalization()
    normed = norm(dyn0)
    for i in range(0, nifti.shape[0], step):
        plt.imshow(normed[0][i], cmap='gray', interpolation=None)
        plt.show()


def test_read_data(path_base: str, set_folder: str, feature_path: str, pre_image_shape: tuple, transform: torchvision.transforms.Compose, protocol: str):
    fold = 0
    train_set_filename = "stratified_training_set-f{}.csv".format(fold)
    eval_set_filename = "stratified_evaluation_set-f{}.csv".format(fold)
    test_set_filename = "stratified_test_set.csv"
    print("TRAINING DATASET")
    traindata_set = BreastMRISubjects(path_base, set_folder + train_set_filename,
                                    protocol=protocol, transform=transform)
    print("EVALUATION DATASET")
    evaldata_set = BreastMRISubjects(path_base, set_folder + eval_set_filename,
                                   protocol=protocol, transform=transform)
    print("TEST DATASET")
    testdata_set = BreastMRISubjects(path_base, set_folder + test_set_filename,
                                   protocol=protocol, transform=transform)


def test_inspect_summary_files(dataset_root: str, dicom_root: str, metadata_file: str = None):
    """Report which examination IDs of the metadata export are covered by the split files."""
    metadata = pd.read_csv(resolve_path(metadata_file, "metadata_file", "metadata export"))
    fold = 0
    train_set_filename = "stratified_training_set-f{}-0.csv".format(fold)
    eval_set_filename = "stratified_evaluation_set-f{}-0.csv".format(fold)
    test_set_filename = "stratified_test_set-f{}.csv".format(fold)
    train = pd.read_csv(dataset_root + train_set_filename, dtype=str)
    eval = pd.read_csv(dataset_root + eval_set_filename, dtype=str)
    test = pd.read_csv(dataset_root + test_set_filename, dtype=str)

    # remove nans
    all_pids = metadata['AnforderungsNrE'].astype('int64')
    train_read = train['AnforderungsNrE'].drop_duplicates().astype('int64')
    eval_read = eval['AnforderungsNrE'].drop_duplicates().astype('int64')
    test_read = test['AnforderungsNrE'].drop_duplicates().astype('int64')

    all_used = pd.concat([train_read.to_frame(), eval_read.to_frame(), test_read.to_frame()])
    missing = all_pids[~all_pids.isin(all_used['AnforderungsNrE'])]
    if len(missing) + len(all_used) != len(all_pids):
        print("{} PIDs in total. \n{} PIDS in use: {} Training, {} Evaluation, {} Testing \n{} PIDs from all_pids are not found in used PIDs.".format(len(all_pids), len(all_used), len(train), len(eval), len(test), len(missing)))
    # find the ones from all_used that are not in all_pids
    origin_unclear = all_used[~all_used['AnforderungsNrE'].isin(all_pids)]
    print("Found {} PIDs that are from unknown origin.".format(len(origin_unclear)))
    pids_without_path = []
    # check which PIDs have no corresponding folder
    for pid in all_used['AnforderungsNrE']:
        if not os.path.exists(dicom_root + str(pid)):
            pids_without_path.append(pid)
    print("{} PIDs to which no data folder exists.".format(len(pids_without_path)))

parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("examination_folder",
                    help="folder of a single examination whose NIfTI files are inspected")
parser.add_argument("--step", type=int, default=12, help="show every n-th slice. Default: 12")
args = parser.parse_args()

file_root = args.examination_folder.rstrip(os.sep) + os.sep
filenames = os.listdir(file_root)
for file in filenames:
    nifti = read_nifti(file_root + file)
    step = args.step
    seperation = torch.zeros(nifti.shape)
    dyn0 = torch.zeros((1,) + nifti.shape)
    dyn0[0] = nifti

    for i in range(0, nifti.shape[0], step):
        plt.imshow(nifti[i], cmap='gray', interpolation=None)
        plt.show()

    plt.imshow(seperation[0], cmap='gray', interpolation=None)
    plt.show()

    #test_flips()
    test_rotation()
    #test_rescale()
    test_normalize()