Feat/expose parc merge #202
6 changed files with 270 additions and 115 deletions
|
|
@ -67,6 +67,8 @@ Enhancements
|
||||||
|
|
||||||
- Add more documentation on registering parcellations, coordinates, and masks (:gh:`166` by `Leonard Sasse`_)
|
- Add more documentation on registering parcellations, coordinates, and masks (:gh:`166` by `Leonard Sasse`_)
|
||||||
|
|
||||||
|
- Expose a :func:`junifer.data.parcellations.merge_parcellations` function to merge a list of parcellations (:gh:`146` by Leonard Sasse`_).
|
||||||
|
|
||||||
Bugs
|
Bugs
|
||||||
~~~~
|
~~~~
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -13,6 +13,7 @@ from .parcellations import (
|
||||||
list_parcellations,
|
list_parcellations,
|
||||||
load_parcellation,
|
load_parcellation,
|
||||||
register_parcellation,
|
register_parcellation,
|
||||||
|
merge_parcellations,
|
||||||
)
|
)
|
||||||
|
|
||||||
from .masks import (
|
from .masks import (
|
||||||
|
|
|
||||||
|
|
@ -16,9 +16,9 @@ import nibabel as nib
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
import requests
|
import requests
|
||||||
from nilearn import datasets
|
from nilearn import datasets, image
|
||||||
|
|
||||||
from ..utils.logging import logger, raise_error
|
from ..utils.logging import logger, raise_error, warn_with_log
|
||||||
from .utils import closest_resolution
|
from .utils import closest_resolution
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -700,3 +700,84 @@ def _retrieve_suit(
|
||||||
].to_list()
|
].to_list()
|
||||||
|
no, it should be no, it should be `List["Nifti1Image"]` i think, should i put it?
makes sense, for the Tuple it works as makes sense, for the Tuple it works as `Tuple["Nifti1Image", List[str]]`?
Yeah. Yeah.
you mean you mean `parcellations_list : list of Nifti1Image`?
Yeah exactly. Yeah exactly.
Yeah but I think it should be Yeah but I think it should be `... niimg-like object` as `nibabel` calls it.
ok will try that, i think i put ok will try that, i think i put `niimg` before which the docs didn't recognise, but `niimg-like object` makes sense
|
|||||||
|
|
||||||
return parcellation_fname, labels
|
return parcellation_fname, labels
|
||||||
|
|
||||||
|
|
||||||
|
def merge_parcellations(
|
||||||
|
parcellations_list: List["Nifti1Image"],
|
||||||
|
parcellations_names: List[str],
|
||||||
|
labels_lists: List[List[str]],
|
||||||
|
) -> Tuple["Nifti1Image", List[str]]:
|
||||||
|
"""Merge all parcellations from a list into one parcellation.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
parcellations_list : list of niimg-like object
|
||||||
|
List of parcellations to merge.
|
||||||
|
parcellations_names: list of str
|
||||||
|
List of names for parcellations at the corresponding indices.
|
||||||
|
labels_lists : list of list of str
|
||||||
|
A list of lists. Each list in the list contains the labels for the
|
||||||
|
parcellation at the corresponding index.
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
parcellation : niimg-like object
|
||||||
|
The parcellation that results from merging the list of input
|
||||||
|
parcellations.
|
||||||
|
labels : list of str
|
||||||
|
List of labels for the resultant parcellation.
|
||||||
|
|
||||||
|
"""
|
||||||
|
# Check for duplicated labels
|
||||||
|
labels_lists_flat = [item for sublist in labels_lists for item in sublist]
|
||||||
|
if len(labels_lists_flat) != len(set(labels_lists_flat)):
|
||||||
|
warn_with_log(
|
||||||
|
"The parcellations have duplicated labels. "
|
||||||
|
"Each label will be prefixed with the parcellation name."
|
||||||
|
)
|
||||||
|
for i_parcellation, t_labels in enumerate(labels_lists):
|
||||||
|
labels_lists[i_parcellation] = [
|
||||||
|
f"{parcellations_names[i_parcellation]}_{t_label}"
|
||||||
|
for t_label in t_labels
|
||||||
|
]
|
||||||
|
overlapping_voxels = False
|
||||||
|
ref_parc = parcellations_list[0]
|
||||||
|
parc_data = ref_parc.get_fdata()
|
||||||
|
|
||||||
|
labels = labels_lists[0]
|
||||||
|
|
||||||
|
for t_parc, t_labels in zip(parcellations_list[1:], labels_lists[1:]):
|
||||||
|
if t_parc.shape != ref_parc.shape:
|
||||||
|
warn_with_log(
|
||||||
|
"The parcellations have different resolutions!"
|
||||||
|
"Resampling all parcellations to the first one in the list."
|
||||||
|
)
|
||||||
|
t_parc = image.resample_to_img(
|
||||||
|
t_parc, ref_parc, interpolation="nearest", copy=True
|
||||||
|
)
|
||||||
|
|
||||||
|
# Get the data from this parcellation
|
||||||
|
t_parc_data = t_parc.get_fdata().copy() # must be copied
|
||||||
|
# Increase the values of each ROI to match the labels
|
||||||
|
t_parc_data[t_parc_data != 0] += len(labels)
|
||||||
|
|
||||||
|
# Only set new values for the voxels that are 0
|
||||||
|
# This makes sure that the voxels that are in multiple
|
||||||
|
# parcellations are assigned to the parcellation that was
|
||||||
|
# first in the list.
|
||||||
|
if np.any(parc_data[t_parc_data != 0] != 0):
|
||||||
|
overlapping_voxels = True
|
||||||
|
|
||||||
|
parc_data[parc_data == 0] += t_parc_data[parc_data == 0]
|
||||||
|
labels.extend(t_labels)
|
||||||
|
|
||||||
|
if overlapping_voxels:
|
||||||
|
warn_with_log(
|
||||||
|
"The parcellations have overlapping voxels. "
|
||||||
|
"The overlapping voxels will be assigned to the "
|
||||||
|
"parcellation that was first in the list."
|
||||||
|
)
|
||||||
|
|
||||||
|
parcellation_img_res = image.new_img_like(parcellations_list[0], parc_data)
|
||||||
|
|
||||||
|
return parcellation_img_res, labels
|
||||||
|
|
|
||||||
|
|
@ -9,6 +9,7 @@ from pathlib import Path
|
||||||
from typing import List
|
from typing import List
|
||||||
|
|
||||||
import nibabel as nib
|
import nibabel as nib
|
||||||
|
import numpy as np
|
||||||
import pytest
|
import pytest
|
||||||
from nilearn.image import new_img_like
|
from nilearn.image import new_img_like
|
||||||
from numpy.testing import assert_array_almost_equal, assert_array_equal
|
from numpy.testing import assert_array_almost_equal, assert_array_equal
|
||||||
|
|
@ -20,6 +21,7 @@ from junifer.data.parcellations import (
|
||||||
_retrieve_tian,
|
_retrieve_tian,
|
||||||
list_parcellations,
|
list_parcellations,
|
||||||
load_parcellation,
|
load_parcellation,
|
||||||
|
merge_parcellations,
|
||||||
register_parcellation,
|
register_parcellation,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -236,9 +238,7 @@ def test_schaefer_parcellation(tmp_path: Path) -> None:
|
||||||
)
|
)
|
||||||
# Load parcellation
|
# Load parcellation
|
||||||
img2, lbl, fname = load_parcellation(
|
img2, lbl, fname = load_parcellation(
|
||||||
name="Schaefer100x7",
|
name="Schaefer100x7", parcellations_dir=tmp_path, resolution=3
|
||||||
parcellations_dir=tmp_path,
|
|
||||||
resolution=3,
|
|
||||||
)
|
)
|
||||||
# Check parcellation values
|
# Check parcellation values
|
||||||
assert fname.name == fname2
|
assert fname.name == fname2
|
||||||
|
|
@ -247,9 +247,7 @@ def test_schaefer_parcellation(tmp_path: Path) -> None:
|
||||||
assert_array_equal(img2.header["pixdim"][1:4], [2, 2, 2]) # type: ignore
|
assert_array_equal(img2.header["pixdim"][1:4], [2, 2, 2]) # type: ignore
|
||||||
# Load parcellation
|
# Load parcellation
|
||||||
img2, lbl, fname = load_parcellation(
|
img2, lbl, fname = load_parcellation(
|
||||||
"Schaefer100x7",
|
"Schaefer100x7", parcellations_dir=tmp_path, resolution=2.1
|
||||||
parcellations_dir=tmp_path,
|
|
||||||
resolution=2.1,
|
|
||||||
)
|
)
|
||||||
# Check parcellation values
|
# Check parcellation values
|
||||||
assert fname.name == fname2
|
assert fname.name == fname2
|
||||||
|
|
@ -258,9 +256,7 @@ def test_schaefer_parcellation(tmp_path: Path) -> None:
|
||||||
assert_array_equal(img2.header["pixdim"][1:4], [2, 2, 2]) # type: ignore
|
assert_array_equal(img2.header["pixdim"][1:4], [2, 2, 2]) # type: ignore
|
||||||
# Load parcellation
|
# Load parcellation
|
||||||
img2, lbl, fname = load_parcellation(
|
img2, lbl, fname = load_parcellation(
|
||||||
"Schaefer100x7",
|
"Schaefer100x7", parcellations_dir=tmp_path, resolution=1.99
|
||||||
parcellations_dir=tmp_path,
|
|
||||||
resolution=1.99,
|
|
||||||
)
|
)
|
||||||
# Check parcellation values
|
# Check parcellation values
|
||||||
assert fname.name == fname1
|
assert fname.name == fname1
|
||||||
|
|
@ -269,9 +265,7 @@ def test_schaefer_parcellation(tmp_path: Path) -> None:
|
||||||
assert_array_equal(img2.header["pixdim"][1:4], [1, 1, 1]) # type: ignore
|
assert_array_equal(img2.header["pixdim"][1:4], [1, 1, 1]) # type: ignore
|
||||||
# Load parcellation
|
# Load parcellation
|
||||||
img2, lbl, fname = load_parcellation(
|
img2, lbl, fname = load_parcellation(
|
||||||
"Schaefer100x7",
|
"Schaefer100x7", parcellations_dir=tmp_path, resolution=0.5
|
||||||
parcellations_dir=tmp_path,
|
|
||||||
resolution=0.5,
|
|
||||||
)
|
)
|
||||||
# Check parcellation values
|
# Check parcellation values
|
||||||
assert fname.name == fname1
|
assert fname.name == fname1
|
||||||
|
|
@ -383,18 +377,10 @@ def test_retrieve_suit_incorrect_space(tmp_path: Path) -> None:
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
"scale, n_label",
|
"scale, n_label", [(1, 16), (2, 32), (3, 50), (4, 54)]
|
||||||
[
|
|
||||||
(1, 16),
|
|
||||||
(2, 32),
|
|
||||||
(3, 50),
|
|
||||||
(4, 54),
|
|
||||||
],
|
|
||||||
)
|
)
|
||||||
def test_tian_3T_6thgeneration(
|
def test_tian_3T_6thgeneration(
|
||||||
tmp_path: Path,
|
tmp_path: Path, scale: int, n_label: int
|
||||||
scale: int,
|
|
||||||
n_label: int,
|
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Test Tian parcellation.
|
"""Test Tian parcellation.
|
||||||
|
|
||||||
|
|
@ -415,8 +401,7 @@ def test_tian_3T_6thgeneration(
|
||||||
assert "TianxS4x3TxMNI6thgeneration" in parcellations
|
assert "TianxS4x3TxMNI6thgeneration" in parcellations
|
||||||
# Load parcellation
|
# Load parcellation
|
||||||
img, lbl, fname = load_parcellation(
|
img, lbl, fname = load_parcellation(
|
||||||
name=f"TianxS{scale}x3TxMNI6thgeneration",
|
name=f"TianxS{scale}x3TxMNI6thgeneration", parcellations_dir=tmp_path
|
||||||
parcellations_dir=tmp_path,
|
|
||||||
)
|
)
|
||||||
fname1 = f"Tian_Subcortex_S{scale}_3T_1mm.nii.gz"
|
fname1 = f"Tian_Subcortex_S{scale}_3T_1mm.nii.gz"
|
||||||
assert img is not None
|
assert img is not None
|
||||||
|
|
@ -437,18 +422,10 @@ def test_tian_3T_6thgeneration(
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
"scale, n_label",
|
"scale, n_label", [(1, 16), (2, 32), (3, 50), (4, 54)]
|
||||||
[
|
|
||||||
(1, 16),
|
|
||||||
(2, 32),
|
|
||||||
(3, 50),
|
|
||||||
(4, 54),
|
|
||||||
],
|
|
||||||
)
|
)
|
||||||
def test_tian_3T_nonlinear2009cAsym(
|
def test_tian_3T_nonlinear2009cAsym(
|
||||||
tmp_path: Path,
|
tmp_path: Path, scale: int, n_label: int
|
||||||
scale: int,
|
|
||||||
n_label: int,
|
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Test Tian parcellation.
|
"""Test Tian parcellation.
|
||||||
|
|
||||||
|
|
@ -480,18 +457,10 @@ def test_tian_3T_nonlinear2009cAsym(
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
"scale, n_label",
|
"scale, n_label", [(1, 16), (2, 34), (3, 54), (4, 62)]
|
||||||
[
|
|
||||||
(1, 16),
|
|
||||||
(2, 34),
|
|
||||||
(3, 54),
|
|
||||||
(4, 62),
|
|
||||||
],
|
|
||||||
)
|
)
|
||||||
def test_tian_7T_6thgeneration(
|
def test_tian_7T_6thgeneration(
|
||||||
tmp_path: Path,
|
tmp_path: Path, scale: int, n_label: int
|
||||||
scale: int,
|
|
||||||
n_label: int,
|
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Test Tian parcellation.
|
"""Test Tian parcellation.
|
||||||
|
|
||||||
|
|
@ -534,10 +503,7 @@ def test_retrieve_tian_incorrect_space(tmp_path: Path) -> None:
|
||||||
"""
|
"""
|
||||||
with pytest.raises(ValueError, match=r"The parameter `space`"):
|
with pytest.raises(ValueError, match=r"The parameter `space`"):
|
||||||
_retrieve_tian(
|
_retrieve_tian(
|
||||||
parcellations_dir=tmp_path,
|
parcellations_dir=tmp_path, resolution=1, scale=1, space="wrong"
|
||||||
resolution=1,
|
|
||||||
scale=1,
|
|
||||||
space="wrong",
|
|
||||||
)
|
)
|
||||||
|
|
||||||
with pytest.raises(ValueError, match=r"MNI6thgeneration"):
|
with pytest.raises(ValueError, match=r"MNI6thgeneration"):
|
||||||
|
|
@ -584,3 +550,162 @@ def test_retrieve_tian_incorrect_scale(tmp_path: Path) -> None:
|
||||||
scale=5,
|
scale=5,
|
||||||
space="MNI6thgeneration",
|
space="MNI6thgeneration",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_merge_parcellations() -> None:
|
||||||
|
"""Test merging parcellations."""
|
||||||
|
# load some parcellations for testing
|
||||||
|
schaefer_parcellation, schaefer_labels, _ = load_parcellation(
|
||||||
|
"Schaefer100x17"
|
||||||
|
)
|
||||||
|
tian_parcellation, tian_labels, _ = load_parcellation(
|
||||||
|
"TianxS2x3TxMNInonlinear2009cAsym"
|
||||||
|
)
|
||||||
|
# prepare the list of the actual parcellations
|
||||||
|
parcellation_list = [schaefer_parcellation, tian_parcellation]
|
||||||
|
# prepare a list of names
|
||||||
|
names = ["Schaefer100x17", "TianxS2x3TxMNInonlinear2009cAsym"]
|
||||||
|
# prepare a list of label lists
|
||||||
|
labels_lists = [schaefer_labels, tian_labels]
|
||||||
|
# merge the parcellations
|
||||||
|
merged_parc, labels = merge_parcellations(
|
||||||
|
parcellation_list, names, labels_lists
|
||||||
|
)
|
||||||
|
|
||||||
|
# we should have 132 integer labels plus 1 for background
|
||||||
|
parc_data = merged_parc.get_fdata()
|
||||||
|
assert len(np.unique(parc_data)) == 133
|
||||||
|
# no background label, so labels is one less
|
||||||
|
assert len(labels) == 132
|
||||||
|
|
||||||
|
|
||||||
|
def test_merge_parcellations_3D_multiple_non_overlapping(
|
||||||
|
tmp_path: Path,
|
||||||
|
) -> None:
|
||||||
|
"""Test merge_parcellations with multiple non-overlapping parcellations.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
tmp_path : pathlib.Path
|
||||||
|
The path to the test directory.
|
||||||
|
|
||||||
|
"""
|
||||||
|
# Get the testing parcellation
|
||||||
|
parcellation, labels, _ = load_parcellation("Schaefer100x7")
|
||||||
|
|
||||||
|
assert parcellation is not None
|
||||||
|
|
||||||
|
# Create two parcellations from it
|
||||||
|
parcellation_data = parcellation.get_fdata()
|
||||||
|
parcellation1_data = parcellation_data.copy()
|
||||||
|
parcellation1_data[parcellation1_data > 50] = 0
|
||||||
|
parcellation2_data = parcellation_data.copy()
|
||||||
|
parcellation2_data[parcellation2_data <= 50] = 0
|
||||||
|
parcellation2_data[parcellation2_data > 0] -= 50
|
||||||
|
labels1 = labels[:50]
|
||||||
|
labels2 = labels[50:]
|
||||||
|
|
||||||
|
parcellation1_img = new_img_like(parcellation, parcellation1_data)
|
||||||
|
parcellation2_img = new_img_like(parcellation, parcellation2_data)
|
||||||
|
|
||||||
|
parcellation_list = [parcellation1_img, parcellation2_img]
|
||||||
|
names = ["high", "low"]
|
||||||
|
labels_lists = [labels1, labels2]
|
||||||
|
|
||||||
|
merged_parc, merged_labels = merge_parcellations(
|
||||||
|
parcellation_list, names, labels_lists
|
||||||
|
)
|
||||||
|
|
||||||
|
parc_data = parcellation.get_fdata()
|
||||||
|
assert_array_equal(parc_data, merged_parc.get_fdata())
|
||||||
|
assert len(labels) == 100
|
||||||
|
assert len(np.unique(parc_data)) == 101 # 100 + 1 because background 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_merge_parcellations_3D_multiple_overlapping(tmp_path: Path) -> None:
|
||||||
|
"""Test merge_parcellations with multiple overlapping parcellations.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
tmp_path : pathlib.Path
|
||||||
|
The path to the test directory.
|
||||||
|
|
||||||
|
"""
|
||||||
|
|
||||||
|
# Get the testing parcellation
|
||||||
|
parcellation, labels, _ = load_parcellation("Schaefer100x7")
|
||||||
|
|
||||||
|
assert parcellation is not None
|
||||||
|
|
||||||
|
# Create two parcellations from it
|
||||||
|
parcellation_data = parcellation.get_fdata()
|
||||||
|
parcellation1_data = parcellation_data.copy()
|
||||||
|
parcellation1_data[parcellation1_data > 50] = 0
|
||||||
|
parcellation2_data = parcellation_data.copy()
|
||||||
|
|
||||||
|
# Make the second parcellation overlap with the first
|
||||||
|
parcellation2_data[parcellation2_data <= 45] = 0
|
||||||
|
parcellation2_data[parcellation2_data > 0] -= 45
|
||||||
|
labels1 = [f"low_{x}" for x in labels[:50]] # Change the labels
|
||||||
|
labels2 = [f"high_{x}" for x in labels[45:]] # Change the labels
|
||||||
|
|
||||||
|
parcellation1_img = new_img_like(parcellation, parcellation1_data)
|
||||||
|
parcellation2_img = new_img_like(parcellation, parcellation2_data)
|
||||||
|
|
||||||
|
parcellation_list = [parcellation1_img, parcellation2_img]
|
||||||
|
names = ["high", "low"]
|
||||||
|
labels_lists = [labels1, labels2]
|
||||||
|
|
||||||
|
with pytest.warns(RuntimeWarning, match="overlapping voxels"):
|
||||||
|
merged_parc, merged_labels = merge_parcellations(
|
||||||
|
parcellation_list, names, labels_lists
|
||||||
|
)
|
||||||
|
|
||||||
|
parc_data = parcellation.get_fdata()
|
||||||
|
assert len(labels) == 100
|
||||||
|
assert len(np.unique(parc_data)) == 101 # 100 + 1 because background 0
|
||||||
|
|
||||||
|
|
||||||
|
def test_merge_parcellations_3D_multiple_duplicated_labels(
|
||||||
|
tmp_path: Path,
|
||||||
|
) -> None:
|
||||||
|
"""Test merge_parcellations with two parcellations with duplicated labels.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
tmp_path : pathlib.Path
|
||||||
|
The path to the test directory.
|
||||||
|
|
||||||
|
"""
|
||||||
|
|
||||||
|
# Get the testing parcellation
|
||||||
|
parcellation, labels, _ = load_parcellation("Schaefer100x7")
|
||||||
|
|
||||||
|
assert parcellation is not None
|
||||||
|
|
||||||
|
# Create two parcellations from it
|
||||||
|
parcellation_data = parcellation.get_fdata()
|
||||||
|
parcellation1_data = parcellation_data.copy()
|
||||||
|
parcellation1_data[parcellation1_data > 50] = 0
|
||||||
|
parcellation2_data = parcellation_data.copy()
|
||||||
|
parcellation2_data[parcellation2_data <= 50] = 0
|
||||||
|
parcellation2_data[parcellation2_data > 0] -= 50
|
||||||
|
labels1 = labels[:50]
|
||||||
|
labels2 = labels[49:-1] # One label is duplicated
|
||||||
|
|
||||||
|
parcellation1_img = new_img_like(parcellation, parcellation1_data)
|
||||||
|
parcellation2_img = new_img_like(parcellation, parcellation2_data)
|
||||||
|
|
||||||
|
parcellation_list = [parcellation1_img, parcellation2_img]
|
||||||
|
names = ["high", "low"]
|
||||||
|
labels_lists = [labels1, labels2]
|
||||||
|
|
||||||
|
with pytest.warns(RuntimeWarning, match="duplicated labels."):
|
||||||
|
merged_parc, merged_labels = merge_parcellations(
|
||||||
|
parcellation_list, names, labels_lists
|
||||||
|
)
|
||||||
|
|
||||||
|
parc_data = parcellation.get_fdata()
|
||||||
|
assert_array_equal(parc_data, merged_parc.get_fdata())
|
||||||
|
assert len(labels) == 100
|
||||||
|
assert len(np.unique(parc_data)) == 101 # 100 + 1 because background 0
|
||||||
|
|
|
||||||
|
|
@ -178,9 +178,7 @@ class DataladDataGrabber(BaseDataGrabber):
|
||||||
try:
|
try:
|
||||||
dl_out = self._dataset.get(to_get, result_renderer="disabled")
|
dl_out = self._dataset.get(to_get, result_renderer="disabled")
|
||||||
except IncompleteResultsError as e:
|
except IncompleteResultsError as e:
|
||||||
raise_error(
|
raise_error(f"Failed to get from dataset: {e.failed}")
|
||||||
f"Failed to get from dataset: {e.failed}"
|
|
||||||
)
|
|
||||||
if not self._was_cloned:
|
if not self._was_cloned:
|
||||||
# If the dataset was already installed, check that the
|
# If the dataset was already installed, check that the
|
||||||
# file was actually downloaded to avoid removing a
|
# file was actually downloaded to avoid removing a
|
||||||
|
|
|
||||||
|
|
@ -7,13 +7,13 @@
|
||||||
from typing import Any, Dict, List, Optional, Union
|
from typing import Any, Dict, List, Optional, Union
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
from nilearn.image import math_img, new_img_like, resample_to_img
|
from nilearn.image import math_img, resample_to_img
|
||||||
from nilearn.maskers import NiftiMasker
|
from nilearn.maskers import NiftiMasker
|
||||||
|
|
||||||
from ..api.decorators import register_marker
|
from ..api.decorators import register_marker
|
||||||
from ..data import get_mask, load_parcellation
|
from ..data import get_mask, load_parcellation, merge_parcellations
|
||||||
from ..stats import get_aggfunc_by_name
|
from ..stats import get_aggfunc_by_name
|
||||||
from ..utils import logger, warn_with_log
|
from ..utils import logger
|
||||||
from .base import BaseMarker
|
from .base import BaseMarker
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -98,9 +98,7 @@ class ParcelAggregation(BaseMarker):
|
||||||
raise ValueError(f"Unknown input kind for {input_type}")
|
raise ValueError(f"Unknown input kind for {input_type}")
|
||||||
|
|
||||||
def compute(
|
def compute(
|
||||||
self,
|
self, input: Dict[str, Any], extra_input: Optional[Dict] = None
|
||||||
input: Dict[str, Any],
|
|
||||||
extra_input: Optional[Dict] = None,
|
|
||||||
) -> Dict:
|
) -> Dict:
|
||||||
"""Compute.
|
"""Compute.
|
||||||
|
|
||||||
|
|
@ -129,8 +127,7 @@ class ParcelAggregation(BaseMarker):
|
||||||
t_input_img = input["data"]
|
t_input_img = input["data"]
|
||||||
logger.debug(f"Parcel aggregation using {self.method}")
|
logger.debug(f"Parcel aggregation using {self.method}")
|
||||||
agg_func = get_aggfunc_by_name(
|
agg_func = get_aggfunc_by_name(
|
||||||
name=self.method,
|
name=self.method, func_params=self.method_params
|
||||||
func_params=self.method_params,
|
|
||||||
)
|
)
|
||||||
# Get the min of the voxels sizes and use it as the resolution
|
# Get the min of the voxels sizes and use it as the resolution
|
||||||
resolution = np.min(t_input_img.header.get_zooms()[:3])
|
resolution = np.min(t_input_img.header.get_zooms()[:3])
|
||||||
|
|
@ -140,15 +137,11 @@ class ParcelAggregation(BaseMarker):
|
||||||
all_labels = []
|
all_labels = []
|
||||||
for t_parc_name in self.parcellation:
|
for t_parc_name in self.parcellation:
|
||||||
t_parcellation, t_labels, _ = load_parcellation(
|
t_parcellation, t_labels, _ = load_parcellation(
|
||||||
name=t_parc_name,
|
name=t_parc_name, resolution=resolution
|
||||||
resolution=resolution,
|
|
||||||
)
|
)
|
||||||
# Resample all of them to the image
|
# Resample all of them to the image
|
||||||
t_parcellation_img_res = resample_to_img(
|
t_parcellation_img_res = resample_to_img(
|
||||||
t_parcellation,
|
t_parcellation, t_input_img, interpolation="nearest", copy=True
|
||||||
t_input_img,
|
|
||||||
interpolation="nearest",
|
|
||||||
copy=True,
|
|
||||||
)
|
)
|
||||||
all_parcelations.append(t_parcellation_img_res)
|
all_parcelations.append(t_parcellation_img_res)
|
||||||
all_labels.append(t_labels)
|
all_labels.append(t_labels)
|
||||||
|
|
@ -159,56 +152,11 @@ class ParcelAggregation(BaseMarker):
|
||||||
labels = all_labels[0]
|
labels = all_labels[0]
|
||||||
else:
|
else:
|
||||||
# Merge the parcellations
|
# Merge the parcellations
|
||||||
|
parcellation_img_res, labels = merge_parcellations(
|
||||||
# Check for duplicated labels
|
all_parcelations, self.parcellation, all_labels
|
||||||
all_labels_flat = [
|
|
||||||
item for sublist in all_labels for item in sublist
|
|
||||||
]
|
|
||||||
if len(all_labels_flat) != len(set(all_labels_flat)):
|
|
||||||
warn_with_log(
|
|
||||||
"The parcellations have duplicated labels. "
|
|
||||||
"Each label will be prefixed with the parcellation name."
|
|
||||||
)
|
|
||||||
for i_parcellation, t_labels in enumerate(all_labels):
|
|
||||||
all_labels[i_parcellation] = [
|
|
||||||
f"{self.parcellation[i_parcellation]}_{t_label}"
|
|
||||||
for t_label in t_labels
|
|
||||||
]
|
|
||||||
overlapping_voxels = False
|
|
||||||
parc_data = all_parcelations[0].get_fdata()
|
|
||||||
labels = all_labels[0]
|
|
||||||
for t_parc, t_labels in zip(all_parcelations[1:], all_labels[1:]):
|
|
||||||
# Get the data from this parcellation
|
|
||||||
t_parc_data = t_parc.get_fdata().copy() # must be copied
|
|
||||||
# Increase the values of each ROI to match the labels
|
|
||||||
t_parc_data[t_parc_data != 0] += len(labels)
|
|
||||||
|
|
||||||
# Only set new values for the voxels that are 0
|
|
||||||
# This makes sure that the voxels that are in multiple
|
|
||||||
# parcellations are assigned to the parcellation that was
|
|
||||||
# first in the list.
|
|
||||||
if np.any(parc_data[t_parc_data != 0] != 0):
|
|
||||||
overlapping_voxels = True
|
|
||||||
|
|
||||||
parc_data[parc_data == 0] += t_parc_data[parc_data == 0]
|
|
||||||
labels.extend(t_labels)
|
|
||||||
|
|
||||||
if overlapping_voxels:
|
|
||||||
warn_with_log(
|
|
||||||
"The parcellations have overlapping voxels. "
|
|
||||||
"The overlapping voxels will be assigned to the "
|
|
||||||
"parcellation that was first in the list."
|
|
||||||
)
|
)
|
||||||
|
|
||||||
parcellation_img_res = new_img_like(
|
parcellation_bin = math_img("img != 0", img=parcellation_img_res)
|
||||||
all_parcelations[0],
|
|
||||||
parc_data,
|
|
||||||
)
|
|
||||||
|
|
||||||
parcellation_bin = math_img(
|
|
||||||
"img != 0",
|
|
||||||
img=parcellation_img_res,
|
|
||||||
)
|
|
||||||
|
|
||||||
if self.masks is not None:
|
if self.masks is not None:
|
||||||
logger.debug(f"Masking with {self.masks}")
|
logger.debug(f"Masking with {self.masks}")
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue
Is it
List[str]?I think a bit of explicit types for
ListandTuplewould be nice.The basic types like str can be documented here as well for ease.