[BUG]: allow for empty parcels in ParcelAggregation #194

Merged
fraimondo merged 3 commits from fix/empty_parcel into main 2023-03-20 14:58:55 +00:00
3 changed files with 148 additions and 14 deletions

View file

@ -59,6 +59,7 @@ Enhancements
- Modify the behaviour of the ``collect`` parameter in HTCondor ``queue`` function to run a collect job even if some of the previous jobs fail. This is useful to collect the results of a pipeline even if some of the jobs fail (:gh:`190` by `Fede Raimondo`_).
- Allow for empty parcels in :class:`junifer.markers.ParcelAggregation`, that will result in NaNs (:gh:`194` by `Fede Raimondo`_).
Bugs
~~~~
@ -76,6 +77,8 @@ Bugs
- Fix an issue with datalad cache and locks in which the overriden settings in Junifer were not propagated to subprocesses, resulting in using the default settings (:gh:`199` by `Fede Raimondo`_).
- Fix a bug in which :class:`junifer.markers.ParcelAggregation` could yield duplicated column names if two or more parcels were used and label names were not unique (:gh:`194` by `Fede Raimondo`_).
API changes
~~~~~~~~~~~

View file

@ -13,7 +13,7 @@ from nilearn.maskers import NiftiMasker
from ..api.decorators import register_marker
from ..data import get_mask, load_parcellation
from ..stats import get_aggfunc_by_name
from ..utils import logger
from ..utils import logger, warn_with_log
from .base import BaseMarker
@ -159,6 +159,22 @@ class ParcelAggregation(BaseMarker):
labels = all_labels[0]
else:
# Merge the parcellations
# Check for duplicated 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:]):
@ -171,9 +187,19 @@ class ParcelAggregation(BaseMarker):
# 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(
all_parcelations[0],
parc_data,
@ -208,17 +234,14 @@ class ParcelAggregation(BaseMarker):
# Get the values for each parcel and apply agg function
logger.debug("Computing ROI means")
parcellation_roi_vals = sorted(np.unique(parcellation_values))
out_labels = []
out_values = []
# Iterate over the parcels (existing)
for t_v in parcellation_roi_vals:
for t_v in range(1, len(labels) + 1):
t_values = agg_func(data[:, parcellation_values == t_v], axis=-1)
out_values.append(t_values)
# Update the labels just in case a parcel has no voxels
# in it
out_labels.append(labels[t_v - 1])
out_values = np.array(out_values).T
out = {"data": out_values, "col_names": out_labels}
out = {"data": out_values, "col_names": labels}
return out

View file

@ -4,6 +4,7 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
import warnings
from pathlib import Path
import nibabel as nib
@ -328,8 +329,12 @@ def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None:
nib.save(parcellation1_img, parcellation1_path)
nib.save(parcellation2_img, parcellation2_path)
register_parcellation("Schaefer100x7_low", parcellation1_path, labels1)
register_parcellation("Schaefer100x7_high", parcellation2_path, labels2)
register_parcellation(
"Schaefer100x7_low", parcellation1_path, labels1, overwrite=True
)
register_parcellation(
"Schaefer100x7_high", parcellation2_path, labels2, overwrite=True
)
# Use the ParcelAggregation object on the original parcellation
marker_original = ParcelAggregation(
@ -355,7 +360,11 @@ def test_ParcelAggregation_3D_multiple_non_overlapping(tmp_path: Path) -> None:
on="VBM_GM",
) # Test passing "on" as a keyword argument
input = {"VBM_GM": {"data": img, "meta": {}}}
split_mean = marker_split.fit_transform(input)["VBM_GM"]
# No warnings should be raised
with warnings.catch_warnings():
warnings.simplefilter("error", category=UserWarning)
split_mean = marker_split.fit_transform(input)["VBM_GM"]
split_mean_data = split_mean["data"]
assert split_mean_data.ndim == 2
@ -408,8 +417,12 @@ def test_ParcelAggregation_3D_multiple_overlapping(tmp_path: Path) -> None:
nib.save(parcellation1_img, parcellation1_path)
nib.save(parcellation2_img, parcellation2_path)
register_parcellation("Schaefer100x7_low2", parcellation1_path, labels1)
register_parcellation("Schaefer100x7_high2", parcellation2_path, labels2)
register_parcellation(
"Schaefer100x7_low2", parcellation1_path, labels1, overwrite=True
)
register_parcellation(
"Schaefer100x7_high2", parcellation2_path, labels2, overwrite=True
)
# Use the ParcelAggregation object on the original parcellation
marker_original = ParcelAggregation(
@ -435,7 +448,101 @@ def test_ParcelAggregation_3D_multiple_overlapping(tmp_path: Path) -> None:
on="VBM_GM",
) # Test passing "on" as a keyword argument
input = {"VBM_GM": {"data": img, "meta": {}}}
split_mean = marker_split.fit_transform(input)["VBM_GM"]
with pytest.warns(RuntimeWarning, match="overlapping voxels"):
split_mean = marker_split.fit_transform(input)["VBM_GM"]
split_mean_data = split_mean["data"]
assert split_mean_data.ndim == 2
assert split_mean_data.shape[0] == 1
assert split_mean_data.shape[1] == 105
# Overlapping voxels should be NaN
assert np.isnan(split_mean_data[:, 50:55]).all()
non_nan = split_mean_data[~np.isnan(split_mean_data)]
# Data should be the same
assert_array_equal(orig_mean_data, non_nan[None, :])
# Labels should be "low" for the first 50 and "high" for the second 50
assert all(x.startswith("low") for x in split_mean["col_names"][:50])
assert all(x.startswith("high") for x in split_mean["col_names"][50:])
def test_ParcelAggregation_3D_multiple_duplicated_labels(
tmp_path: Path,
) -> None:
"""Test ParcelAggregation 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
# Get the oasis VBM data
oasis_dataset = datasets.fetch_oasis_vbm(n_subjects=1)
vbm = oasis_dataset.gray_matter_maps[0]
img = nib.load(vbm)
# 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)
parcellation1_path = tmp_path / "parcellation1.nii.gz"
parcellation2_path = tmp_path / "parcellation2.nii.gz"
nib.save(parcellation1_img, parcellation1_path)
nib.save(parcellation2_img, parcellation2_path)
register_parcellation(
"Schaefer100x7_low", parcellation1_path, labels1, overwrite=True
)
register_parcellation(
"Schaefer100x7_high", parcellation2_path, labels2, overwrite=True
)
# Use the ParcelAggregation object on the original parcellation
marker_original = ParcelAggregation(
parcellation="Schaefer100x7",
method="mean",
name="gmd_schaefer100x7_mean",
on="VBM_GM",
) # Test passing "on" as a keyword argument
input = {"VBM_GM": {"data": img, "meta": {}}}
orig_mean = marker_original.fit_transform(input)["VBM_GM"]
orig_mean_data = orig_mean["data"]
assert orig_mean_data.ndim == 2
assert orig_mean_data.shape[0] == 1
assert orig_mean_data.shape[1] == 100
# assert_array_almost_equal(auto, jun_values3d_mean)
# Use the ParcelAggregation object on the two parcellations
marker_split = ParcelAggregation(
parcellation=["Schaefer100x7_low", "Schaefer100x7_high"],
method="mean",
name="gmd_schaefer100x7_mean",
on="VBM_GM",
) # Test passing "on" as a keyword argument
input = {"VBM_GM": {"data": img, "meta": {}}}
with pytest.warns(RuntimeWarning, match="duplicated labels."):
split_mean = marker_split.fit_transform(input)["VBM_GM"]
split_mean_data = split_mean["data"]
assert split_mean_data.ndim == 2
@ -445,6 +552,7 @@ def test_ParcelAggregation_3D_multiple_overlapping(tmp_path: Path) -> None:
# Data should be the same
assert_array_equal(orig_mean_data, split_mean_data)
# Labels should be "low" for the first 50 and "high" for the second 50
assert all(x.startswith("low") for x in split_mean["col_names"][:50])
assert all(x.startswith("high") for x in split_mean["col_names"][50:])
# Labels should be prefixed with the parcellation name
col_names = [f"Schaefer100x7_low_{x}" for x in labels1]
col_names += [f"Schaefer100x7_high_{x}" for x in labels2]
assert col_names == split_mean["col_names"]