""" Visual inspection of single transforms on the NIfTI files of one examination. python scripts/data_utils/inspect_transforms.py 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()