Feat/expose parc merge #202

Merged
LeSasse merged 8 commits from feat/expose_parc_merge into main 2023-03-27 14:55:15 +00:00
6 changed files with 270 additions and 115 deletions

View file

@ -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
~~~~ ~~~~

View file

@ -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 (

View file

@ -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()
synchon commented 2023-03-27 08:45:32 +00:00 (Migrated from github.com)

Is it List[str]?

Is it `List[str]`?
synchon commented 2023-03-27 08:51:38 +00:00 (Migrated from github.com)

I think a bit of explicit types for List and Tuple would be nice.

I think a bit of explicit types for `List` and `Tuple` would be nice.
synchon commented 2023-03-27 08:52:23 +00:00 (Migrated from github.com)

The basic types like str can be documented here as well for ease.

The basic types like str can be documented here as well for ease.
LeSasse commented 2023-03-27 10:42:09 +00:00 (Migrated from github.com)

no, it should be List["Nifti1Image"] i think, should i put it?

no, it should be `List["Nifti1Image"]` i think, should i put it?
LeSasse commented 2023-03-27 10:45:26 +00:00 (Migrated from github.com)

makes sense, for the Tuple it works as Tuple["Nifti1Image", List[str]]?

makes sense, for the Tuple it works as `Tuple["Nifti1Image", List[str]]`?
synchon commented 2023-03-27 10:45:36 +00:00 (Migrated from github.com)

Yeah.

Yeah.
LeSasse commented 2023-03-27 10:46:05 +00:00 (Migrated from github.com)

you mean parcellations_list : list of Nifti1Image?

you mean `parcellations_list : list of Nifti1Image`?
synchon commented 2023-03-27 10:46:16 +00:00 (Migrated from github.com)

Yeah exactly.

Yeah exactly.
synchon commented 2023-03-27 10:47:47 +00:00 (Migrated from github.com)

Yeah but I think it should be ... niimg-like object as nibabel calls it.

Yeah but I think it should be `... niimg-like object` as `nibabel` calls it.
LeSasse commented 2023-03-27 10:51:47 +00:00 (Migrated from github.com)

ok will try that, i think i put niimg before which the docs didn't recognise, but niimg-like object makes sense

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

View file

@ -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

View file

@ -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

View file

@ -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}")