AI_MRI / scripts /data_utils /inspect_transforms.py
DeboraJ1's picture
init
a23d562
Raw History Blame Contribute Delete
6.17 kB
"""
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()