mrid-python 0.1.3__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- mrid/__init__.py +5 -0
- mrid/atlas/MNI152/__init__.py +80 -0
- mrid/atlas/SRI24/__init__.py +77 -0
- mrid/atlas/__init__.py +7 -0
- mrid/loading/__init__.py +1 -0
- mrid/loading/convert.py +68 -0
- mrid/preprocessing/__init__.py +12 -0
- mrid/preprocessing/bias_field_correction.py +28 -0
- mrid/preprocessing/cropping.py +36 -0
- mrid/preprocessing/registration.py +251 -0
- mrid/preprocessing/skullstripping.py +185 -0
- mrid/study.py +442 -0
- mrid/utils/__init__.py +3 -0
- mrid/utils/dcm2niix.py +90 -0
- mrid/utils/dicom_uid_fixer.py +86 -0
- mrid/utils/plotting.py +102 -0
- mrid/utils/python_utils.py +48 -0
- mrid/utils/stl_utils.py +70 -0
- mrid/utils/torch_utils.py +16 -0
- mrid_python-0.1.3.dist-info/METADATA +140 -0
- mrid_python-0.1.3.dist-info/RECORD +27 -0
- mrid_python-0.1.3.dist-info/WHEEL +5 -0
- mrid_python-0.1.3.dist-info/top_level.txt +2 -0
- tests/test_loading.py +82 -0
- tests/test_preprocessing.py +43 -0
- tests/test_study.py +136 -0
- tests/test_utils.py +16 -0
mrid/utils/dcm2niix.py
ADDED
|
@@ -0,0 +1,90 @@
|
|
|
1
|
+
import warnings
|
|
2
|
+
import os
|
|
3
|
+
import shutil
|
|
4
|
+
import subprocess
|
|
5
|
+
import tempfile
|
|
6
|
+
|
|
7
|
+
import SimpleITK as sitk
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def run_dcm2niix(
|
|
11
|
+
inpath: str | os.PathLike,
|
|
12
|
+
outfolder: str | os.PathLike,
|
|
13
|
+
outfname: str,
|
|
14
|
+
mkdirs=True,
|
|
15
|
+
save_BIDS=False,
|
|
16
|
+
allow_stacking=True,
|
|
17
|
+
) -> str:
|
|
18
|
+
"""Convert dicom folder to NIfTI format and return path to the output ``nii.gz`` file,
|
|
19
|
+
uses dcm2niix (https://github.com/rordenlab/dcm2niix) which needs to be installed.
|
|
20
|
+
|
|
21
|
+
This is a simple wrapper around dcm2niix command line interface using subprocess, it also handles non-ascii paths.
|
|
22
|
+
|
|
23
|
+
Args:
|
|
24
|
+
inpath (str):
|
|
25
|
+
Path to the dicom files of a single study and single modality.
|
|
26
|
+
Software like Weasis has DICOM export functionality that can be
|
|
27
|
+
used to organize DICOM files into folders by patients studies and modalities.
|
|
28
|
+
|
|
29
|
+
outfolder (str): Path to the output folder (e.g. ``D:/MRI/patient001/0``).
|
|
30
|
+
|
|
31
|
+
outname (str):
|
|
32
|
+
Output filename, excluding ``.nii.gz`` because it will be added by ``dcm2niix``.
|
|
33
|
+
Can use modifiers (e.g. `%d` will be replaced with series description string from DICOM metadata),
|
|
34
|
+
as explained here https://www.nitrc.org/plugins/mwiki/index.php/dcm2nii:MainPage#General_Usage
|
|
35
|
+
|
|
36
|
+
mkdirs (bool, optional): Whether to create ``outfolder`` if it doesn't exist, otherwise throws an error. Defaults to True.
|
|
37
|
+
|
|
38
|
+
save_BIDS (bool, optional):
|
|
39
|
+
Whether to save extra BIDS sidecar - extra info in JSON format that can't be saved into nifti. Defaults to False.
|
|
40
|
+
|
|
41
|
+
allow_stacking (bool, optional):
|
|
42
|
+
Whether to allow stacking different studies into a single file.
|
|
43
|
+
Sometimes this may help with malformed DICOMs that are recognized as separate studies.
|
|
44
|
+
"""
|
|
45
|
+
# dicom2niix doesnt support non-ascii paths, so convert to temporary directory
|
|
46
|
+
with tempfile.TemporaryDirectory() as tmpdir:
|
|
47
|
+
|
|
48
|
+
# create temporary folders
|
|
49
|
+
tmp_input_dir = os.path.join(tmpdir, "mrid_dcm2niix_input")
|
|
50
|
+
tmp_output_dir = os.path.join(tmpdir, "mrid_dcm2niix_output")
|
|
51
|
+
if os.path.exists(tmp_output_dir): shutil.rmtree(tmp_output_dir)
|
|
52
|
+
os.mkdir(tmp_output_dir)
|
|
53
|
+
shutil.copytree(inpath, tmp_input_dir)
|
|
54
|
+
|
|
55
|
+
# create output dir if it doesn't exist
|
|
56
|
+
if outfolder != '':
|
|
57
|
+
if not os.path.exists(outfolder):
|
|
58
|
+
if mkdirs: os.makedirs(outfolder)
|
|
59
|
+
else: raise NotADirectoryError(f"Output path {outfolder} doesn't exist")
|
|
60
|
+
|
|
61
|
+
# run dcm2niix
|
|
62
|
+
subprocess.run(["dcm2niix",
|
|
63
|
+
"-z", "y", # compression
|
|
64
|
+
"-m", "y" if allow_stacking else 'n', # disable stacking images from different studies
|
|
65
|
+
"-b", 'y' if save_BIDS else 'n', # save additional JSON info that can't be saved into nifti (https://bids.neuroimaging.io/ BIDS sidecar format)
|
|
66
|
+
"-o", os.path.normpath(tmp_output_dir), # output folder
|
|
67
|
+
"-f", outfname, # output filename
|
|
68
|
+
os.path.normpath(tmp_input_dir)], # input folder
|
|
69
|
+
check=True)
|
|
70
|
+
|
|
71
|
+
# find what new nifti files were created
|
|
72
|
+
out_files = [i for i in os.listdir(tmp_output_dir) if i.lower().strip().endswith('.nii.gz')]
|
|
73
|
+
|
|
74
|
+
# move them to output folder
|
|
75
|
+
shutil.copytree(tmp_output_dir, outfolder)
|
|
76
|
+
|
|
77
|
+
if len(out_files) > 1:
|
|
78
|
+
warnings.warn(f"More than one NIfTI file was created in {outfolder}, path to the first one will be returned. Something may be wrong.")
|
|
79
|
+
|
|
80
|
+
if len(out_files) == 0:
|
|
81
|
+
raise RuntimeError(f"No nifti files were created in {outfolder}")
|
|
82
|
+
|
|
83
|
+
# return path to the created nifti file
|
|
84
|
+
return os.path.join(outfolder, out_files[0])
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
def dcm2sitk(inpath:str | os.PathLike) -> sitk.Image:
|
|
88
|
+
with tempfile.TemporaryDirectory() as tmpdir:
|
|
89
|
+
nifti_path = run_dcm2niix(inpath=inpath, outfolder=tmpdir, outfname='temp', mkdirs=False, save_BIDS=False)
|
|
90
|
+
return sitk.ReadImage(nifti_path)
|
|
@@ -0,0 +1,86 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import warnings
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
def fix_dicom_uids(input_folder: str | os.PathLike, output_folder: str | os.PathLike):
|
|
6
|
+
"""
|
|
7
|
+
Reads DICOM files from input_folder, assigns a new common SeriesInstanceUID,
|
|
8
|
+
generates new unique SOPInstanceUIDs and sequential InstanceNumbers,
|
|
9
|
+
and saves the modified files to output_folder.
|
|
10
|
+
|
|
11
|
+
This can fix malformed DICOMs that are loaded as separate series by some software.
|
|
12
|
+
|
|
13
|
+
Args:
|
|
14
|
+
input_folder (str): Path to the folder containing the problematic DICOM files.
|
|
15
|
+
output_folder (str): Path to the folder where fixed DICOM files will be saved.
|
|
16
|
+
It will be created if it doesn't exist.
|
|
17
|
+
"""
|
|
18
|
+
import pydicom
|
|
19
|
+
from pydicom.errors import InvalidDicomError
|
|
20
|
+
from pydicom.uid import generate_uid
|
|
21
|
+
|
|
22
|
+
os.makedirs(output_folder, exist_ok=True)
|
|
23
|
+
|
|
24
|
+
dicom_files_info = []
|
|
25
|
+
|
|
26
|
+
# -------------------------------- load DICOMs ------------------------------- #
|
|
27
|
+
for filename in os.listdir(input_folder):
|
|
28
|
+
input_filepath = os.path.join(input_folder, filename)
|
|
29
|
+
if os.path.isfile(input_filepath):
|
|
30
|
+
try:
|
|
31
|
+
ds = pydicom.dcmread(input_filepath, defer_size="512 KB", stop_before_pixels=False)
|
|
32
|
+
|
|
33
|
+
if 'PixelData' in ds:
|
|
34
|
+
dicom_files_info.append({'dataset': ds, 'original_filename': filename})
|
|
35
|
+
else:
|
|
36
|
+
warnings.warn(f"Skipping DICOM file with no PixelData: {filename}")
|
|
37
|
+
|
|
38
|
+
# except InvalidDicomError as e:
|
|
39
|
+
# warnings.warn(f"Failed to load {filename}, skipping, exception was: {e}")
|
|
40
|
+
except Exception as e:
|
|
41
|
+
warnings.warn(f"Failed to load {filename}, skipping, exception was: {e}")
|
|
42
|
+
|
|
43
|
+
if not dicom_files_info:
|
|
44
|
+
raise FileNotFoundError(f'No valid DICOM files with PixelData found in the "{input_folder}"')
|
|
45
|
+
|
|
46
|
+
# ------------------------- sort by instance numbers ------------------------- #
|
|
47
|
+
dicom_files_info.sort(key=lambda info: int(info['dataset'].get('InstanceNumber', 0)))
|
|
48
|
+
|
|
49
|
+
# ---------------------- make new dataset with new UIDs ---------------------- #
|
|
50
|
+
new_series_uid = generate_uid()
|
|
51
|
+
|
|
52
|
+
for i, file_info in enumerate(dicom_files_info):
|
|
53
|
+
ds = file_info['dataset']
|
|
54
|
+
original_filename = file_info['original_filename']
|
|
55
|
+
instance_number = i + 1 # Generate sequential instance number (1-based)
|
|
56
|
+
|
|
57
|
+
# Generate a NEW, UNIQUE SOP Instance UID for this specific file
|
|
58
|
+
new_sop_instance_uid = generate_uid()
|
|
59
|
+
|
|
60
|
+
# Update main DICOM tags
|
|
61
|
+
ds.SeriesInstanceUID = new_series_uid
|
|
62
|
+
ds.SOPInstanceUID = new_sop_instance_uid
|
|
63
|
+
ds.InstanceNumber = str(instance_number) # VR 'IS' (Integer String)
|
|
64
|
+
|
|
65
|
+
# Update file_meta (Group 0002)
|
|
66
|
+
# Check if file_meta exists (it should for standard DICOM files)
|
|
67
|
+
if hasattr(ds, 'file_meta') and ds.file_meta:
|
|
68
|
+
ds.file_meta.MediaStorageSOPInstanceUID = new_sop_instance_uid
|
|
69
|
+
|
|
70
|
+
# update Implementation Class UID and Version Name
|
|
71
|
+
# ds.file_meta.ImplementationClassUID = pydicom.uid.PYDICOM_IMPLEMENTATION_UID
|
|
72
|
+
# ds.file_meta.ImplementationVersionName = f"PYDICOM {pydicom.__version__}"
|
|
73
|
+
else:
|
|
74
|
+
warnings.warn(f"Warning: File Meta Information (Group 0002) missing or empty in {original_filename}. Cannot update MediaStorageSOPInstanceUID.")
|
|
75
|
+
|
|
76
|
+
# other potentially helpful tags (but avoid geometry if unsure)
|
|
77
|
+
# ds.add_new(0x00080013, 'TM', datetime.datetime.now().strftime('%H%M%S.%f')[:16]) # Instance Creation Time (dummy)
|
|
78
|
+
# ds.add_new(0x00080033, 'TM', datetime.datetime.now().strftime('%H%M%S.%f')[:16]) # Content Time (dummy)
|
|
79
|
+
|
|
80
|
+
# ----------------------------- save new dataset ----------------------------- #
|
|
81
|
+
for file_info in dicom_files_info:
|
|
82
|
+
ds = file_info['dataset']
|
|
83
|
+
output_filepath = os.path.join(output_folder, file_info['original_filename'])
|
|
84
|
+
|
|
85
|
+
os.makedirs(os.path.dirname(output_filepath), exist_ok=True)
|
|
86
|
+
ds.save_as(output_filepath, write_like_original=True)
|
mrid/utils/plotting.py
ADDED
|
@@ -0,0 +1,102 @@
|
|
|
1
|
+
from collections.abc import Mapping
|
|
2
|
+
|
|
3
|
+
import matplotlib.gridspec as gridspec
|
|
4
|
+
import matplotlib.pyplot as plt
|
|
5
|
+
import numpy as np
|
|
6
|
+
|
|
7
|
+
def visualize_3d_arrays(data: Mapping[str, np.ndarray]):
|
|
8
|
+
n_vals = len(data)
|
|
9
|
+
|
|
10
|
+
# 1. Determine layout for the Outer Grid (Modalities)
|
|
11
|
+
# We try to make it roughly square (e.g., 4 items -> 2x2, 5 items -> 2x3)
|
|
12
|
+
n_cols_outer = int(np.ceil(np.sqrt(n_vals)))
|
|
13
|
+
n_rows_outer = int(np.ceil(n_vals / n_cols_outer))
|
|
14
|
+
|
|
15
|
+
# Figure size scaling
|
|
16
|
+
fig = plt.figure(figsize=(n_cols_outer * 5, n_rows_outer * 5))
|
|
17
|
+
|
|
18
|
+
# Create the Outer Grid
|
|
19
|
+
outer_grid = gridspec.GridSpec(n_rows_outer, n_cols_outer, figure=fig, wspace=0.3, hspace=0.3)
|
|
20
|
+
|
|
21
|
+
# 2. Iterate through each modality
|
|
22
|
+
for i, (modality_name, volume) in enumerate(data.items()):
|
|
23
|
+
# Create an Inner Grid (3x3) inside the specific cell of the Outer Grid
|
|
24
|
+
inner_grid = gridspec.GridSpecFromSubplotSpec(
|
|
25
|
+
3, 3,
|
|
26
|
+
subplot_spec=outer_grid[i],
|
|
27
|
+
wspace=0.05, hspace=0.05
|
|
28
|
+
)
|
|
29
|
+
|
|
30
|
+
# Calculate slicing indices for 25%, 50%, 75%
|
|
31
|
+
shapes = volume.shape
|
|
32
|
+
percentages = [0.25, 0.50, 0.75]
|
|
33
|
+
|
|
34
|
+
# Loop through dimensions (Rows of the 3x3)
|
|
35
|
+
for dim_idx in range(3):
|
|
36
|
+
|
|
37
|
+
# Get the exact indices for this dimension
|
|
38
|
+
slice_indices = [int(shapes[dim_idx] * p) for p in percentages]
|
|
39
|
+
|
|
40
|
+
# Loop through the slices (Columns of the 3x3)
|
|
41
|
+
for col_idx, slice_loc in enumerate(slice_indices):
|
|
42
|
+
|
|
43
|
+
# Create the subplot in the inner grid
|
|
44
|
+
ax = fig.add_subplot(inner_grid[dim_idx, col_idx])
|
|
45
|
+
|
|
46
|
+
# 3. Extract the 2D Slice
|
|
47
|
+
if dim_idx == 0:
|
|
48
|
+
img_slice = volume[slice_loc, :, :]
|
|
49
|
+
row_label = f"Dim 0\n(Slice {slice_loc})"
|
|
50
|
+
elif dim_idx == 1:
|
|
51
|
+
img_slice = volume[:, slice_loc, :]
|
|
52
|
+
row_label = f"Dim 1\n(Slice {slice_loc})"
|
|
53
|
+
else: # dim_idx == 2
|
|
54
|
+
img_slice = volume[:, :, slice_loc]
|
|
55
|
+
row_label = f"Dim 2\n(Slice {slice_loc})"
|
|
56
|
+
|
|
57
|
+
# Plot image
|
|
58
|
+
ax.imshow(img_slice, cmap='gray', aspect='auto')
|
|
59
|
+
|
|
60
|
+
# Remove ticks for cleanliness
|
|
61
|
+
ax.set_xticks([])
|
|
62
|
+
ax.set_yticks([])
|
|
63
|
+
|
|
64
|
+
# Labels: Only add row labels to the first column
|
|
65
|
+
if col_idx == 0:
|
|
66
|
+
ax.set_ylabel(row_label, fontsize=9, rotation=90)
|
|
67
|
+
|
|
68
|
+
# Labels: Only add col labels to the first row
|
|
69
|
+
if dim_idx == 0:
|
|
70
|
+
ax.set_title(f"{int(percentages[col_idx]*100)}%", fontsize=10)
|
|
71
|
+
|
|
72
|
+
# Add the main Modality Name on top of the grid block
|
|
73
|
+
# We fetch the geometric center of the top row of the inner grid
|
|
74
|
+
box = outer_grid[i].get_position(fig)
|
|
75
|
+
fig.text(box.x0 + box.width/2, box.y1 + 0.02, modality_name,
|
|
76
|
+
ha='center', va='bottom', fontsize=14, fontweight='bold')
|
|
77
|
+
|
|
78
|
+
plt.show()
|
|
79
|
+
|
|
80
|
+
if __name__ == "__main__":
|
|
81
|
+
|
|
82
|
+
def dummy_data(shape, type='sphere'):
|
|
83
|
+
x, y, z = np.indices(shape)
|
|
84
|
+
cx, cy, cz = shape[0]//2, shape[1]//2, shape[2]//2
|
|
85
|
+
mask = np.zeros(shape)
|
|
86
|
+
if type == 'sphere':
|
|
87
|
+
r = min(shape)//3
|
|
88
|
+
mask[(x-cx)**2 + (y-cy)**2 + (z-cz)**2 < r**2] = 1
|
|
89
|
+
elif type == 'cube':
|
|
90
|
+
r = min(shape)//4
|
|
91
|
+
mask[cx-r:cx+r, cy-r:cy+r, cz-r:cz+r] = 1
|
|
92
|
+
elif type == 'noise':
|
|
93
|
+
mask = np.random.rand(*shape)
|
|
94
|
+
|
|
95
|
+
return mask
|
|
96
|
+
|
|
97
|
+
visualize_3d_arrays({
|
|
98
|
+
"T1 Weighted (Sphere)": dummy_data((60, 60, 60), 'sphere'),
|
|
99
|
+
"T2 Weighted (Cube)": dummy_data((60, 60, 60), 'cube'),
|
|
100
|
+
"Proton Density (Noise)": dummy_data((60, 60, 60), 'noise'),
|
|
101
|
+
"FLAIR (Sphere)": dummy_data((60, 60, 60), 'sphere'),
|
|
102
|
+
})
|
|
@@ -0,0 +1,48 @@
|
|
|
1
|
+
import importlib.util
|
|
2
|
+
from collections.abc import Mapping, Sequence
|
|
3
|
+
from typing import Any
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
# lazy loader from https://stackoverflow.com/a/78312674/15673832
|
|
7
|
+
class LazyLoader:
|
|
8
|
+
'thin shell class to wrap modules. load real module on first access and pass thru'
|
|
9
|
+
|
|
10
|
+
def __init__(self, modname):
|
|
11
|
+
self._modname = modname
|
|
12
|
+
self._mod = None
|
|
13
|
+
|
|
14
|
+
def __getattr__(self, attr):
|
|
15
|
+
'import module on first attribute access'
|
|
16
|
+
|
|
17
|
+
try:
|
|
18
|
+
return getattr(self._mod, attr)
|
|
19
|
+
|
|
20
|
+
except Exception as e :
|
|
21
|
+
if self._mod is None :
|
|
22
|
+
# module is unset, load it
|
|
23
|
+
self._mod = importlib.import_module (self._modname)
|
|
24
|
+
else :
|
|
25
|
+
# module is set, got different exception from getattr (). reraise it
|
|
26
|
+
raise e
|
|
27
|
+
|
|
28
|
+
# retry getattr if module was just loaded for first time
|
|
29
|
+
# call this outside exception handler in case it raises new exception
|
|
30
|
+
return getattr (self._mod, attr)
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
# this allows transforms to support any kind of container
|
|
34
|
+
# class _Packer[T: Sequence | Mapping]:
|
|
35
|
+
# def __init__(self, type: type[T], keys: list | None = None):
|
|
36
|
+
# self.type: Any = type
|
|
37
|
+
# self.keys = keys
|
|
38
|
+
|
|
39
|
+
# def pack(self, unpacked: Sequence) -> T:
|
|
40
|
+
# if self.keys is not None: return self.type(dict(zip(self.keys, unpacked)))
|
|
41
|
+
# return self.type(unpacked)
|
|
42
|
+
|
|
43
|
+
# def unpack_struct[T: Sequence | Mapping](struct: T) -> tuple[Any, _Packer[T]]:
|
|
44
|
+
# if isinstance(struct, Sequence):
|
|
45
|
+
# return list(struct), _Packer(type(struct))
|
|
46
|
+
# if isinstance(struct, Mapping):
|
|
47
|
+
# return list(struct.values()), _Packer(type(struct), list(struct.values()))
|
|
48
|
+
# raise TypeError(f"Transformation functions accept lists and dictionaries, but received {type(struct)}")
|
mrid/utils/stl_utils.py
ADDED
|
@@ -0,0 +1,70 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import warnings
|
|
3
|
+
|
|
4
|
+
import numpy as np
|
|
5
|
+
import SimpleITK as sitk
|
|
6
|
+
|
|
7
|
+
from ..loading.convert import tositk
|
|
8
|
+
|
|
9
|
+
def stl2sitk(
|
|
10
|
+
stl_path: str | os.PathLike,
|
|
11
|
+
reference: str | os.PathLike | sitk.Image,
|
|
12
|
+
fix_holes: bool = False,
|
|
13
|
+
):
|
|
14
|
+
"""
|
|
15
|
+
Loads an STL file under ``stl_path`` and converts it to ``sitk.Image`` aligned with ``reference``.
|
|
16
|
+
Note that this might take a few minutes.
|
|
17
|
+
|
|
18
|
+
Args:
|
|
19
|
+
stl_path (str): Path to the STL segmentation file. The STL coordinates
|
|
20
|
+
MUST be in the same coordinate system as the ``reference``.
|
|
21
|
+
reference (str | os.PathLike | sitk.Image): path to a directory of DICOM files or a NIfTI file, or a ``sitk.Image``.
|
|
22
|
+
fix_holes (bool, optional): whether to try to fix holes in STL if they are detected (this can be very slow).
|
|
23
|
+
"""
|
|
24
|
+
import trimesh
|
|
25
|
+
|
|
26
|
+
# ------------------------------ load reference ------------------------------ #
|
|
27
|
+
reference = tositk(reference)
|
|
28
|
+
|
|
29
|
+
origin = np.array(reference.GetOrigin())
|
|
30
|
+
spacing = np.array(reference.GetSpacing())
|
|
31
|
+
# ct_direction = np.array(ct_image.GetDirection()).reshape(3, 3)
|
|
32
|
+
size = np.array(reference.GetSize()) # Order: x, y, z
|
|
33
|
+
shape_xyz = size
|
|
34
|
+
shape_zyx = size[::-1]
|
|
35
|
+
|
|
36
|
+
# --------------------------------- load STL --------------------------------- #
|
|
37
|
+
mesh = trimesh.load_mesh(stl_path)
|
|
38
|
+
|
|
39
|
+
if not mesh.is_watertight:
|
|
40
|
+
warnings.warn(f"Warning: STL mesh '{stl_path}' is not watertight. Voxelization using 'contains' might be inaccurate.")
|
|
41
|
+
if fix_holes:
|
|
42
|
+
mesh.fill_holes()
|
|
43
|
+
if not mesh.is_watertight:
|
|
44
|
+
print("Warning: Failed to make mesh watertight after filling holes.")
|
|
45
|
+
|
|
46
|
+
# ------------------------------- voxelize STL ------------------------------- #
|
|
47
|
+
x_coords = origin[0] + np.arange(shape_xyz[0]) * spacing[0]
|
|
48
|
+
y_coords = origin[1] + np.arange(shape_xyz[1]) * spacing[1]
|
|
49
|
+
z_coords = origin[2] + np.arange(shape_xyz[2]) * spacing[2]
|
|
50
|
+
|
|
51
|
+
# Use meshgrid to create a grid of coordinates
|
|
52
|
+
# Note the 'ij' indexing to match the z, y, x array structure
|
|
53
|
+
zz, yy, xx = np.meshgrid(z_coords, y_coords, x_coords, indexing='ij')
|
|
54
|
+
|
|
55
|
+
# Stack coordinates into a (N, 3) array where N = Z*Y*X
|
|
56
|
+
voxel_centers_xyz = np.stack([xx.ravel(), yy.ravel(), zz.ravel()], axis=-1)
|
|
57
|
+
|
|
58
|
+
# This checks which voxel center points fall inside the mesh volume
|
|
59
|
+
voxel_mask_flat = mesh.contains(voxel_centers_xyz)
|
|
60
|
+
|
|
61
|
+
# Reshape the flat boolean mask back into the 3D CT shape (z, y, x)
|
|
62
|
+
stl_array = voxel_mask_flat.reshape(shape_zyx).astype(np.uint8) # Use uint8 for masks
|
|
63
|
+
|
|
64
|
+
# ------------------------------ make sitk.Image ----------------------------- #
|
|
65
|
+
stl_sitk = sitk.GetImageFromArray(stl_array)
|
|
66
|
+
|
|
67
|
+
stl_sitk.SetOrigin(reference.GetOrigin())
|
|
68
|
+
stl_sitk.SetSpacing(reference.GetSpacing())
|
|
69
|
+
stl_sitk.SetDirection(reference.GetDirection())
|
|
70
|
+
return stl_sitk
|
|
@@ -0,0 +1,16 @@
|
|
|
1
|
+
import importlib.util
|
|
2
|
+
from typing import TYPE_CHECKING, cast
|
|
3
|
+
|
|
4
|
+
from .python_utils import LazyLoader
|
|
5
|
+
|
|
6
|
+
lazy_torch = LazyLoader("torch")
|
|
7
|
+
if TYPE_CHECKING:
|
|
8
|
+
import torch
|
|
9
|
+
lazy_torch = cast(torch, lazy_torch)
|
|
10
|
+
|
|
11
|
+
TORCH_INSTALLED = importlib.util.find_spec("torch") is not None
|
|
12
|
+
|
|
13
|
+
if TORCH_INSTALLED:
|
|
14
|
+
CUDA_IF_AVAILABLE = 'cuda' if lazy_torch.cuda.is_available() else 'cpu'
|
|
15
|
+
else:
|
|
16
|
+
CUDA_IF_AVAILABLE = 'cpu'
|
|
@@ -0,0 +1,140 @@
|
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
|
+
Name: mrid-python
|
|
3
|
+
Version: 0.1.3
|
|
4
|
+
Summary: Tools for working with 3D medical images and segmentations - registration, brain skull-stripping, etc.
|
|
5
|
+
Author-email: Ivan Nikishev <nkshv2@gmail.com>
|
|
6
|
+
Project-URL: Homepage, https://github.com/inikishev/mrid
|
|
7
|
+
Project-URL: Repository, https://github.com/inikishev/mrid
|
|
8
|
+
Project-URL: Issues, https://github.com/inikishev/mrid/isses
|
|
9
|
+
Keywords: MRI,medical imaging
|
|
10
|
+
Requires-Python: >=3.10
|
|
11
|
+
Description-Content-Type: text/markdown
|
|
12
|
+
Requires-Dist: numpy
|
|
13
|
+
Requires-Dist: SimpleITK
|
|
14
|
+
|
|
15
|
+
<h1 align='center'>mrid</h1>
|
|
16
|
+
|
|
17
|
+
mrid is a library for preprocessing of 3D images, particularly medical images.
|
|
18
|
+
|
|
19
|
+
`mrid.Study` provides a convenient way to work with scans and provides all the tools (skull-stripping, registration, etc). See [this notebook](https://nbviewer.org/github/inikishev/mrid/blob/main/mrid_tutorial.ipynb) for an example of how to use it ([Google Colab](https://colab.research.google.com/drive/1d1MWPaGAqfK6q_5879LpwwxPCB-tbExe?usp=sharing)).
|
|
20
|
+
|
|
21
|
+
All methods are also available as separate functions, see below.
|
|
22
|
+
|
|
23
|
+
## Registering a scan
|
|
24
|
+
|
|
25
|
+
All methods accept scans that can be a path to NIfTI, DICOM dir, sitk Image, numpy array or torch tensor, and return sitk.Image format.
|
|
26
|
+
|
|
27
|
+
For registering you need to have SimpleITK-SimpleElastix installed (<https://pypi.org/project/SimpleITK-SimpleElastix/>)
|
|
28
|
+
|
|
29
|
+
```python
|
|
30
|
+
registered = mrid.register(moving, fixed) # returns sitk.Image
|
|
31
|
+
```
|
|
32
|
+
|
|
33
|
+
mrid includes downloaders for commonly used templates - MNI152 and SRI24: `mrid.get_mni152`, and `mrid.get_sri24`.
|
|
34
|
+
|
|
35
|
+
```python
|
|
36
|
+
t1_sri = mrid.register(t1, mrid.get_sri24("T1")) # returns sitk.Image
|
|
37
|
+
```
|
|
38
|
+
|
|
39
|
+
note you can always convert anything to numpy array using `mrid.tonumpy(src)`:
|
|
40
|
+
|
|
41
|
+
```python
|
|
42
|
+
t1_sri_np = mrid.tonumpy(t1_sri) # np.array
|
|
43
|
+
```
|
|
44
|
+
|
|
45
|
+
## Registering multiple scans
|
|
46
|
+
|
|
47
|
+
If you have multiple scans that are already aligned, you can register one of them and use the same transformation to transform the rest. This can also be used to register segmentations.
|
|
48
|
+
|
|
49
|
+
In this example we register T1 to SRI24 T1 template, and transform the rest
|
|
50
|
+
|
|
51
|
+
```python
|
|
52
|
+
data = {"t1": "t1c.nii.gz", "t2": "t2w.nii.gz", "flair": "t2f.nii.gz"}
|
|
53
|
+
registered_data = mrid.register_D(data, key="t1", to=mrid.get_sri24("T1"))
|
|
54
|
+
# dict of sitk.Image registered to SRI24
|
|
55
|
+
```
|
|
56
|
+
|
|
57
|
+
If you have multiple scans that are not aligned (e.g. have different shapes, orientations), you can use `register_each`. This registers one scan to the target, and then registers all other scans to the first scan. In the example it registers T1 to SRI24, and then registers T2 and FLAIR to T1.
|
|
58
|
+
|
|
59
|
+
```python
|
|
60
|
+
data = {"t1": "t1c.nii.gz", "t2": "t2w.nii.gz", "flair": "t2f.nii.gz"}
|
|
61
|
+
registered_data = mrid.register_each(data, key="t1", to=mrid.get_sri24("T1"))
|
|
62
|
+
# dict of sitk.Image registered to SRI24
|
|
63
|
+
```
|
|
64
|
+
|
|
65
|
+
## Inverse transform
|
|
66
|
+
|
|
67
|
+
mrid provides a simple way to compute an inverse of registration, this is also another way to perform registration in general.
|
|
68
|
+
|
|
69
|
+
```python
|
|
70
|
+
reg = mrid.Registration()
|
|
71
|
+
|
|
72
|
+
# first we need to fit a transform
|
|
73
|
+
reg.find_transform(t1, mrid.get_sri24("T1"))
|
|
74
|
+
|
|
75
|
+
# now we can use `reg.apply_transform` to transform any scan (assuming it is aligned to `t1`)
|
|
76
|
+
t2_sri24 = reg.apply_transform(t2) # returns sitk.Image
|
|
77
|
+
|
|
78
|
+
# use `apply_inverse_transform` to transform back to original space
|
|
79
|
+
t2_recovered = reg.apply_inverse_transform(t2_sri24) # returns sitk.Image
|
|
80
|
+
|
|
81
|
+
# when transforming segmentations, use nearest
|
|
82
|
+
# neigbour interpolation to preserve edges
|
|
83
|
+
seg = (mrid.tonumpy(t2_sri24) > 100).astype(np.uint8)
|
|
84
|
+
seg_orig = reg.apply_inverse_transform(seg, use_nearest_interpolation=True) # returns sitk.Image
|
|
85
|
+
```
|
|
86
|
+
|
|
87
|
+
## Skullstripping
|
|
88
|
+
|
|
89
|
+
For skullstripping mrid uses HD-BET model (<https://github.com/MIC-DKFZ/HD-BET>), and you need to have it installed (it uses torch). It works with pre/post contrast T1, T2 and FLAIR MRIs, but it expects them to be in MNI152 space (while many new datasets like BraTS are in SRI24). If your scans aren't in MNI152, the `skullstrip` function can register specified scan to MNI152, pass it to HD-BET, and then register back to original space, if `register_to_mni152` argument is specified.
|
|
90
|
+
|
|
91
|
+
```python
|
|
92
|
+
t1_skullstripped = mrid.skullstrip(t1) # if t1 is in MNI152
|
|
93
|
+
t1_skullstripped = mrid.skullstrip(t1, register_to_mni152="T1") # if t1 is not in MNI152
|
|
94
|
+
```
|
|
95
|
+
|
|
96
|
+
HD-BET is generally very robust but it can remove parts of tumors adjascent to the skull, but this is true for all skullstripping tools.
|
|
97
|
+
|
|
98
|
+
## Skullstripping multiple scans
|
|
99
|
+
|
|
100
|
+
If you have multiple scans that are already aligned, you can compute brain mask on one of them (ideally post-contrast T1) and apply it to all scans.
|
|
101
|
+
|
|
102
|
+
```python
|
|
103
|
+
data = {"t1": "t1c.nii.gz", "t2": "t2w.nii.gz", "flair": "t2f.nii.gz"}
|
|
104
|
+
skullstripped_data = mrid.skullstrip_D(data, key="t1", to=mrid.get_sri24("T1"))
|
|
105
|
+
# dict of skullstripped sitk.Image
|
|
106
|
+
```
|
|
107
|
+
|
|
108
|
+
## dcm2niix
|
|
109
|
+
|
|
110
|
+
mrid provides a simple wrapper around [dcm2niix](https://github.com/rordenlab/dcm2niix), which is a very robust tool for converting DICOMs to NIfTI.
|
|
111
|
+
|
|
112
|
+
```python
|
|
113
|
+
# this saves output.nii.gz to output_path. It can save multiple files if there are multiple cases in "path/to/dicoms".
|
|
114
|
+
mrid.utils.run_dcm2niix("path/to/dicoms", outfolder="output_path", outfname="output")
|
|
115
|
+
```
|
|
116
|
+
|
|
117
|
+
If you want to use it convert DICOMs directly to SimpleITK, use `dcm2sitk`:
|
|
118
|
+
|
|
119
|
+
```python
|
|
120
|
+
t1 = mrid.utils.dcm2sitk("path/to/dicoms") # returns sitk.Image with correct spatial information
|
|
121
|
+
```
|
|
122
|
+
dcm2niix is much slower compared to SimpleITK's series reader which mrid uses by default, but it's more careful about splitting different studies to different files, although it might fail with malformed DICOMs which are surprisingly common (where it would think each slice is a separate study).
|
|
123
|
+
|
|
124
|
+
## Loading STL segmentations
|
|
125
|
+
|
|
126
|
+
`stl2sitk` converts segmentations in STL format to `sitk.Image` given a reference `sitk.Image` that the STL should be aligned to. Basically something I encountered so I had to make this.
|
|
127
|
+
|
|
128
|
+
```python
|
|
129
|
+
seg = mrid.utils.stl2sitk("path/to/file.stl", reference) # returns sitk.Image (hopefully) aligned with `reference`
|
|
130
|
+
```
|
|
131
|
+
|
|
132
|
+
## How to install
|
|
133
|
+
Either run
|
|
134
|
+
```
|
|
135
|
+
pip install mrid-python
|
|
136
|
+
```
|
|
137
|
+
or
|
|
138
|
+
```
|
|
139
|
+
pip install git+https://github.com/inikishev/mrid
|
|
140
|
+
```
|
|
@@ -0,0 +1,27 @@
|
|
|
1
|
+
mrid/__init__.py,sha256=d-h6l6e-sfPUKDo27_dXcx8IRdZ5ByDC1YtdkM6y39o,138
|
|
2
|
+
mrid/study.py,sha256=fa_sipmUlDvYz9QjVvuMUj5eSc-kMjlM2PYRyJ12GAQ,19595
|
|
3
|
+
mrid/atlas/__init__.py,sha256=ecOAVXUOt0TUI_wK19Ed05arF6b_G61wucVq2c8U4ZY,110
|
|
4
|
+
mrid/atlas/MNI152/__init__.py,sha256=QHm0ok5EkxrWCVV5uGmMvFmlJoXxJtqr6qkGB9H-9Fo,2833
|
|
5
|
+
mrid/atlas/SRI24/__init__.py,sha256=jFFY6LeJcWwCtxVJGbdWPEUGFKZjqzxu01ZTbbSAwEo,2939
|
|
6
|
+
mrid/loading/__init__.py,sha256=h5uDRBxjJBuz9w95BEJIUbBpV2rYswTr12B_83GfVfw,57
|
|
7
|
+
mrid/loading/convert.py,sha256=NtXRmYKdjnms7GUyHLoyV0H6ntHxvTj615PIdTUlCQY,2799
|
|
8
|
+
mrid/preprocessing/__init__.py,sha256=uA2FDGBUZ2MvLqmpRAYsfQxlFOzJh_U87G2H8vXgcJE,525
|
|
9
|
+
mrid/preprocessing/bias_field_correction.py,sha256=K2w75JkVIV_Gj7W-9gpHUwz2teEvgSRt5wX-CqQeRik,1143
|
|
10
|
+
mrid/preprocessing/cropping.py,sha256=I26hxo7RuWqdjx8ApJ2t9FBPdi3PtRsBdMlLC2KgKTk,1336
|
|
11
|
+
mrid/preprocessing/registration.py,sha256=un_dLJI0jO--huUrdRkr5t8Wsczi9erEqi91QxkACsw,11901
|
|
12
|
+
mrid/preprocessing/skullstripping.py,sha256=-M3SmeeeXO3hQYur5tJhWvBirbeWE239pnMwA5Klo7Y,8294
|
|
13
|
+
mrid/utils/__init__.py,sha256=9pbvafeDUmkFQghV3OPXGxg-004yO9BzmrbcgF_vgyo,121
|
|
14
|
+
mrid/utils/dcm2niix.py,sha256=KHU7dgUVsU7WOEa82HFO-G0SsGIxYjeUFEBcX27mkWQ,4124
|
|
15
|
+
mrid/utils/dicom_uid_fixer.py,sha256=xjjfvAF13Ni5CEgwPcr7oDT7excdEoeRGQetvjGkSRM,4046
|
|
16
|
+
mrid/utils/plotting.py,sha256=lU1UqygEEAxp7hUpNHKVi5oiWCRDX_3llFqQZt4zh1I,3831
|
|
17
|
+
mrid/utils/python_utils.py,sha256=IYB8-KN4d6l-5ILgUNYux0vNrK0D0vh0AbSRqWhs-fo,1793
|
|
18
|
+
mrid/utils/stl_utils.py,sha256=T4qT11NgvFCERfaB3duA_O1Qi9FBW7-b3B5aZKS5UU4,2888
|
|
19
|
+
mrid/utils/torch_utils.py,sha256=nKby_tPWzCGdEASkDUq7xNo-iUe9zjwmhq_chV0RCGU,407
|
|
20
|
+
tests/test_loading.py,sha256=CciTvnqQ7hY-Vk2GgH2SWSu8nyS8lLeBqHu3XzClBLY,2141
|
|
21
|
+
tests/test_preprocessing.py,sha256=HcoubITfmlJkgPaBbkxZiY04fJv9URBGMWDJWivDqLU,1300
|
|
22
|
+
tests/test_study.py,sha256=pK__azqnMz8DSyOYaWQmBkpAqLntCIIjEz0m8oBS4RQ,3984
|
|
23
|
+
tests/test_utils.py,sha256=Q4uf5dog673RDofOC6RPLYnxVN5NAIVRTiyTbf9Az74,359
|
|
24
|
+
mrid_python-0.1.3.dist-info/METADATA,sha256=rUfrJuDkvNA9GaU60dsapfC4aGafIzcv6K4R8bFlV6s,6099
|
|
25
|
+
mrid_python-0.1.3.dist-info/WHEEL,sha256=_zCd3N1l69ArxyTb8rzEoP9TpbYXkqRFSNOD5OuxnTs,91
|
|
26
|
+
mrid_python-0.1.3.dist-info/top_level.txt,sha256=lBv75ms7UoIM4elDVX3CbKiSlh-X0vABtaUwj4TGx4o,11
|
|
27
|
+
mrid_python-0.1.3.dist-info/RECORD,,
|
tests/test_loading.py
ADDED
|
@@ -0,0 +1,82 @@
|
|
|
1
|
+
import numpy as np
|
|
2
|
+
import pytest
|
|
3
|
+
import SimpleITK as sitk
|
|
4
|
+
from mrid.loading.convert import tositk, tonumpy, totensor
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
def test_tositk_with_numpy():
|
|
8
|
+
data = np.random.rand(10, 20, 30).astype(np.float32)
|
|
9
|
+
sitk_img = tositk(data)
|
|
10
|
+
|
|
11
|
+
assert isinstance(sitk_img, sitk.Image)
|
|
12
|
+
assert sitk.GetArrayFromImage(sitk_img).shape == (10, 20, 30)
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def test_tositk_with_sitk_image():
|
|
16
|
+
original_img = sitk.Image(10, 20, 30, sitk.sitkFloat32)
|
|
17
|
+
result_img = tositk(original_img)
|
|
18
|
+
|
|
19
|
+
assert result_img is original_img
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def test_tositk_with_pathlike():
|
|
23
|
+
with pytest.raises(FileNotFoundError):
|
|
24
|
+
tositk("/nonexistent/path.nii.gz")
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def test_tonumpy_with_numpy():
|
|
28
|
+
data = np.random.rand(10, 20, 30).astype(np.float32)
|
|
29
|
+
result = tonumpy(data)
|
|
30
|
+
|
|
31
|
+
assert isinstance(result, np.ndarray)
|
|
32
|
+
assert np.array_equal(data, result)
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def test_tonumpy_with_sitk():
|
|
36
|
+
sitk_img = sitk.Image(10, 20, 30, sitk.sitkFloat32)
|
|
37
|
+
array = sitk.GetArrayFromImage(sitk_img)
|
|
38
|
+
|
|
39
|
+
result = tonumpy(sitk_img)
|
|
40
|
+
|
|
41
|
+
assert isinstance(result, np.ndarray)
|
|
42
|
+
assert result.shape == array.shape
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def test_tonumpy_with_pathlike():
|
|
46
|
+
with pytest.raises(FileNotFoundError):
|
|
47
|
+
tonumpy("/nonexistent/path.nii.gz")
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
def test_totensor_with_numpy():
|
|
51
|
+
try:
|
|
52
|
+
import torch
|
|
53
|
+
data = np.random.rand(10, 20, 30).astype(np.float32)
|
|
54
|
+
tensor = totensor(data)
|
|
55
|
+
|
|
56
|
+
assert isinstance(tensor, torch.Tensor)
|
|
57
|
+
assert tensor.shape == torch.Size([10, 20, 30])
|
|
58
|
+
except ImportError:
|
|
59
|
+
pytest.skip("no torch")
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
def test_totensor_with_sitk():
|
|
63
|
+
try:
|
|
64
|
+
import torch
|
|
65
|
+
sitk_img = sitk.Image(30, 20, 10, sitk.sitkFloat32) # note that dims are reversed in sitk
|
|
66
|
+
tensor = totensor(sitk_img)
|
|
67
|
+
|
|
68
|
+
assert isinstance(tensor, torch.Tensor)
|
|
69
|
+
assert tensor.shape == torch.Size([10, 20, 30])
|
|
70
|
+
except ImportError:
|
|
71
|
+
pytest.skip("no torch")
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
def test_totensor_with_tensor():
|
|
75
|
+
try:
|
|
76
|
+
import torch
|
|
77
|
+
original_tensor = torch.rand(10, 20, 30)
|
|
78
|
+
result = totensor(original_tensor)
|
|
79
|
+
|
|
80
|
+
assert result is original_tensor
|
|
81
|
+
except ImportError:
|
|
82
|
+
pytest.skip("no torch")
|