From c5f9ea20123a99414e9428ec519acf0e690c8e1c Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 3 Dec 2024 16:06:12 +0100 Subject: [PATCH 1/8] refactor: update ParcellationRegistry.load interface to conditionally fetch highest resolution --- junifer/data/parcellations/_parcellations.py | 12 ++++++++++++ 1 file changed, 12 insertions(+) diff --git a/junifer/data/parcellations/_parcellations.py b/junifer/data/parcellations/_parcellations.py index 9035c8332..f3a9406f5 100644 --- a/junifer/data/parcellations/_parcellations.py +++ b/junifer/data/parcellations/_parcellations.py @@ -269,6 +269,7 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): def load( self, name: str, + target_space: str, parcellations_dir: Union[str, Path, None] = None, resolution: Optional[float] = None, path_only: bool = False, @@ -282,6 +283,8 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): ---------- name : str The name of the parcellation. + target_space : str + The desired space of the parcellation. parcellations_dir : str or pathlib.Path, optional Path where the parcellations files are stored. The default location is "$HOME/junifer/data/parcellations" (default None). @@ -328,6 +331,14 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): else: space = parcellation_definition["space"] + # Check and get highest resolution + if space != target_space: + logger.info( + f"Parcellation will be warped from {space} to {target_space} " + "using highest resolution" + ) + resolution = None + # Check if the parcellation family is custom or built-in if t_family == "CustomUserParcellation": parcellation_fname = Path(parcellation_definition["path"]) @@ -441,6 +452,7 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): img, labels, _, space = self.load( name=name, resolution=resolution, + target_space=target_space, ) # Convert parcellation spaces if required; -- 2.52.0 From 409e5c1f7637fd12cbd0a289bdd3743ed56b6537 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 3 Dec 2024 16:07:17 +0100 Subject: [PATCH 2/8] refactor: improve space transformation logic for ParcellationRegistry.get --- junifer/data/parcellations/_parcellations.py | 69 ++++++++++---------- 1 file changed, 36 insertions(+), 33 deletions(-) diff --git a/junifer/data/parcellations/_parcellations.py b/junifer/data/parcellations/_parcellations.py index f3a9406f5..777aaee60 100644 --- a/junifer/data/parcellations/_parcellations.py +++ b/junifer/data/parcellations/_parcellations.py @@ -471,17 +471,44 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): ) # Remove extra dimension added by ANTs img = image.math_img("np.squeeze(img)", img=raw_img) + else: + if target_space != "native": + # No warping is going to happen, just resampling, because + # we are in the correct space + logger.debug(f"Resampling {name} to target image.") + # Resample parcellation to target image + img = image.resample_to_img( + source_img=img, + target_img=target_img, + interpolation="nearest", + copy=True, + ) - logger.debug(f"Resampling {name} to target image.") - # Resample parcellation to target image - img_to_merge = image.resample_to_img( - source_img=img, - target_img=target_img, - interpolation="nearest", - copy=True, - ) + # Warp parcellation if target space is native + if target_space == "native": + logger.debug( + "Warping parcellation to native space using " + f"{warper_spec['warper']}." + ) + # extra_input check done earlier and warper_spec exists + if warper_spec["warper"] == "fsl": + img = FSLParcellationWarper().warp( + parcellation_name="native", + parcellation_img=img, + target_data=target_data, + warp_data=warper_spec, + ) + elif warper_spec["warper"] == "ants": + img = ANTsParcellationWarper().warp( + parcellation_name="native", + parcellation_img=img, + src="", + dst="native", + target_data=target_data, + warp_data=warper_spec, + ) - all_parcellations.append(img_to_merge) + all_parcellations.append(img) all_labels.append(labels) # Avoid merging if there is only one parcellation @@ -497,30 +524,6 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): labels_lists=all_labels, ) - # Warp parcellation if target space is native - if target_space == "native": - logger.debug( - "Warping parcellation to native space using " - f"{warper_spec['warper']}." - ) - # extra_input check done earlier and warper_spec exists - if warper_spec["warper"] == "fsl": - resampled_parcellation_img = FSLParcellationWarper().warp( - parcellation_name="native", - parcellation_img=resampled_parcellation_img, - target_data=target_data, - warp_data=warper_spec, - ) - elif warper_spec["warper"] == "ants": - resampled_parcellation_img = ANTsParcellationWarper().warp( - parcellation_name="native", - parcellation_img=resampled_parcellation_img, - src="", - dst="native", - target_data=target_data, - warp_data=warper_spec, - ) - return resampled_parcellation_img, labels -- 2.52.0 From 91bae7f047b1e3c2e19d055579f019f46b9cccbc Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 3 Dec 2024 16:07:43 +0100 Subject: [PATCH 3/8] chore: fix tests for ParcellationRegistry.load --- .../parcellations/tests/test_parcellations.py | 64 +++++++++++++------ 1 file changed, 45 insertions(+), 19 deletions(-) diff --git a/junifer/data/parcellations/tests/test_parcellations.py b/junifer/data/parcellations/tests/test_parcellations.py index 46b6f50aa..23a5aceff 100644 --- a/junifer/data/parcellations/tests/test_parcellations.py +++ b/junifer/data/parcellations/tests/test_parcellations.py @@ -60,7 +60,9 @@ def test_register_already_registered() -> None: space="MNI152Lin", ) assert ( - ParcellationRegistry().load("testparc", path_only=True)[2].name + ParcellationRegistry() + .load("testparc", target_space="MNI152Lin", path_only=True)[2] + .name == "testparc.nii.gz" ) @@ -81,7 +83,9 @@ def test_register_already_registered() -> None: ) assert ( - ParcellationRegistry().load("testparc", path_only=True)[2].name + ParcellationRegistry() + .load("testparc", target_space="MNI152Lin", path_only=True)[2] + .name == "testparc2.nii.gz" ) @@ -96,7 +100,8 @@ def test_parcellation_wrong_labels_values(tmp_path: Path) -> None: """ schaefer, labels, schaefer_path, _ = ParcellationRegistry().load( - "Schaefer100x7" + "Schaefer100x7", + "MNI152NLin6Asym", ) assert schaefer is not None @@ -106,7 +111,7 @@ def test_parcellation_wrong_labels_values(tmp_path: Path) -> None: ) with pytest.raises(ValueError, match=r"has 100 parcels but 10"): - ParcellationRegistry().load("WrongLabels") + ParcellationRegistry().load("WrongLabels", "MNI152NLin6Asym") # Test wrong number of labels ParcellationRegistry().register( @@ -114,7 +119,7 @@ def test_parcellation_wrong_labels_values(tmp_path: Path) -> None: ) with pytest.raises(ValueError, match=r"has 100 parcels but 101"): - ParcellationRegistry().load("WrongLabels2") + ParcellationRegistry().load("WrongLabels2", "MNI152NLin6Asym") schaefer_data = schaefer.get_fdata().copy() schaefer_data[schaefer_data == 50] = 0 @@ -126,7 +131,7 @@ def test_parcellation_wrong_labels_values(tmp_path: Path) -> None: "WrongValues", new_schaefer_path, labels[:-1], "MNI152Lin" ) with pytest.raises(ValueError, match=r"must have all the values in the"): - ParcellationRegistry().load("WrongValues") + ParcellationRegistry().load("WrongValues", "MNI152NLin6Asym") schaefer_data = schaefer.get_fdata().copy() schaefer_data[schaefer_data == 50] = 200 @@ -138,7 +143,7 @@ def test_parcellation_wrong_labels_values(tmp_path: Path) -> None: "WrongValues2", new_schaefer_path, labels, "MNI152Lin" ) with pytest.raises(ValueError, match=r"must have all the values in the"): - ParcellationRegistry().load("WrongValues2") + ParcellationRegistry().load("WrongValues2", "MNI152NLin6Asym") @pytest.mark.parametrize( @@ -202,7 +207,7 @@ def test_register( assert name in ParcellationRegistry().list # Load registered parcellation _, lbl, fname, parcellation_space = ParcellationRegistry().load( - name=name, path_only=True + name=name, target_space=space, path_only=True ) # Check values for registered parcellation assert lbl == parcels_labels @@ -239,7 +244,7 @@ def test_list_correct(parcellation_name: str) -> None: def test_load_incorrect() -> None: """Test loading of invalid parcellations.""" with pytest.raises(ValueError, match=r"not found"): - ParcellationRegistry().load("wrongparcellation") + ParcellationRegistry().load("wrongparcellation", "MNI152NLin6Asym") def test_retrieve_parcellation_incorrect() -> None: @@ -323,6 +328,7 @@ def test_schaefer( # Load parcellation img, label, img_path, space = ParcellationRegistry().load( name=parcellation_name, + target_space="MNI152NLin6Asym", parcellations_dir=tmp_path, resolution=resolution, ) @@ -392,6 +398,7 @@ def test_suit(tmp_path: Path, space_key: str, space: str) -> None: # Load parcellation img, label, img_path, parcellation_space = ParcellationRegistry().load( name=f"SUITx{space_key}", + target_space=space, parcellations_dir=tmp_path, ) assert img is not None @@ -441,7 +448,9 @@ def test_tian_3T_6thgeneration( assert "TianxS4x3TxMNI6thgeneration" in parcellations # Load parcellation img, lbl, fname, parcellation_space_1 = ParcellationRegistry().load( - name=f"TianxS{scale}x3TxMNI6thgeneration", parcellations_dir=tmp_path + name=f"TianxS{scale}x3TxMNI6thgeneration", + parcellations_dir=tmp_path, + target_space="MNI152NLin2009cAsym", ) fname1 = f"Tian_Subcortex_S{scale}_3T_1mm.nii.gz" assert img is not None @@ -452,6 +461,7 @@ def test_tian_3T_6thgeneration( # Load parcellation img, lbl, fname, parcellation_space_2 = ParcellationRegistry().load( name=f"TianxS{scale}x3TxMNI6thgeneration", + target_space="MNI152NLin6Asym", parcellations_dir=tmp_path, resolution=2, ) @@ -489,6 +499,7 @@ def test_tian_3T_nonlinear2009cAsym( # Load parcellation img, lbl, fname, space = ParcellationRegistry().load( name=f"TianxS{scale}x3TxMNInonlinear2009cAsym", + target_space="MNI152NLin2009cAsym", parcellations_dir=tmp_path, ) fname1 = f"Tian_Subcortex_S{scale}_3T_2009cAsym.nii.gz" @@ -524,7 +535,9 @@ def test_tian_7T_6thgeneration( assert "TianxS4x7TxMNI6thgeneration" in parcellations # Load parcellation img, lbl, fname, space = ParcellationRegistry().load( - name=f"TianxS{scale}x7TxMNI6thgeneration", parcellations_dir=tmp_path + name=f"TianxS{scale}x7TxMNI6thgeneration", + target_space="MNI152NLin6Asym", + parcellations_dir=tmp_path, ) fname1 = f"Tian_Subcortex_S{scale}_7T.nii.gz" assert img is not None @@ -611,7 +624,9 @@ def test_aicha(tmp_path: Path, version: int) -> None: assert f"AICHA_v{version}" in ParcellationRegistry().list # Load parcellation img, label, img_path, space = ParcellationRegistry().load( - name=f"AICHA_v{version}", parcellations_dir=tmp_path + name=f"AICHA_v{version}", + target_space="IXI549Space", + parcellations_dir=tmp_path, ) assert img is not None assert img_path.name == "AICHA.nii" @@ -680,6 +695,7 @@ def test_shen( # Load parcellation img, label, img_path, space = ParcellationRegistry().load( name=f"Shen_{year}_{n_rois}", + target_space="MNI152NLin2009cAsym", parcellations_dir=tmp_path, resolution=resolution, ) @@ -875,7 +891,8 @@ def test_yan( ) # Load parcellation img, label, img_path, space = ParcellationRegistry().load( - name=parcellation_name, # type: ignore + name=parcellation_name, + target_space="MNI152NLin6Asym", parcellations_dir=tmp_path, resolution=resolution, ) @@ -1012,6 +1029,7 @@ def test_brainnetome( # Load parcellation img, label, img_path, space = ParcellationRegistry().load( name=parcellation_name, + target_space="MNI152NLin6Asym", parcellations_dir=tmp_path, resolution=resolution, ) @@ -1044,10 +1062,11 @@ def test_merge_parcellations() -> None: """Test merging parcellations.""" # load some parcellations for testing schaefer_parcellation, schaefer_labels, _, _ = ParcellationRegistry().load( - "Schaefer100x17" + "Schaefer100x17", target_space="MNI152NLin2009cAsym" ) tian_parcellation, tian_labels, _, _ = ParcellationRegistry().load( - "TianxS2x3TxMNInonlinear2009cAsym" + "TianxS2x3TxMNInonlinear2009cAsym", + target_space="MNI152NLin2009cAsym", ) # prepare the list of the actual parcellations parcellation_list = [schaefer_parcellation, tian_parcellation] @@ -1079,7 +1098,9 @@ def test_merge_parcellations_3D_multiple_non_overlapping( """ # Get the testing parcellation - parcellation, labels, _, _ = ParcellationRegistry().load("Schaefer100x7") + parcellation, labels, _, _ = ParcellationRegistry().load( + "Schaefer100x7", target_space="MNI152NLin2009cAsym" + ) assert parcellation is not None @@ -1114,7 +1135,9 @@ def test_merge_parcellations_3D_multiple_overlapping() -> None: """Test merge_parcellations with multiple overlapping parcellations.""" # Get the testing parcellation - parcellation, labels, _, _ = ParcellationRegistry().load("Schaefer100x7") + parcellation, labels, _, _ = ParcellationRegistry().load( + "Schaefer100x7", target_space="MNI152NLin2009cAsym" + ) assert parcellation is not None @@ -1149,7 +1172,9 @@ def test_merge_parcellations_3D_multiple_duplicated_labels() -> None: """Test merge_parcellations with duplicated labels.""" # Get the testing parcellation - parcellation, labels, _, _ = ParcellationRegistry().load("Schaefer100x7") + parcellation, labels, _, _ = ParcellationRegistry().load( + "Schaefer100x7", target_space="MNI152NLin2009cAsym" + ) assert parcellation is not None @@ -1198,6 +1223,7 @@ def test_get_single() -> None: # Get raw parcellation raw_parcellation, raw_labels, _, _ = ParcellationRegistry().load( "TianxS1x3TxMNInonlinear2009cAsym", + target_space="MNI152NLin2009cAsym", resolution=1.5, ) resampled_raw_parcellation = resample_to_img( @@ -1240,7 +1266,7 @@ def test_get_multi_same_space() -> None: ] for name in parcellations_names: img, labels, _, _ = ParcellationRegistry().load( - name=name, resolution=1.5 + name=name, target_space="MNI152NLin2009cAsym", resolution=1.5 ) # Resample raw parcellations resampled_img = resample_to_img( -- 2.52.0 From 1ec27fcb6c3be77de5defaf77712c45451e2003f Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 3 Dec 2024 16:14:53 +0100 Subject: [PATCH 4/8] chore: add changelogs 411.{change,enh} --- docs/changes/newsfragments/411.change | 1 + docs/changes/newsfragments/411.enh | 1 + 2 files changed, 2 insertions(+) create mode 100644 docs/changes/newsfragments/411.change create mode 100644 docs/changes/newsfragments/411.enh diff --git a/docs/changes/newsfragments/411.change b/docs/changes/newsfragments/411.change new file mode 100644 index 000000000..391c47d6e --- /dev/null +++ b/docs/changes/newsfragments/411.change @@ -0,0 +1 @@ +:meth:`.ParcellationRegistry.load` now has a new parameter ``target_space`` to provide parcellation in highest possible resolution if it does not match the space of the parcellation by `Fede Raimondo`_ and `Synchon Mandal`_ diff --git a/docs/changes/newsfragments/411.enh b/docs/changes/newsfragments/411.enh new file mode 100644 index 000000000..69034f258 --- /dev/null +++ b/docs/changes/newsfragments/411.enh @@ -0,0 +1 @@ +Refactor the parcellation logic in the pipeline to account for and optimise space transformations and merges by `Synchon Mandal`_ and `Fede Raimondo`_ -- 2.52.0 From fdd1b6b807c60061ed8a0a61e7e9fb2bbcb8d602 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 3 Dec 2024 18:15:47 +0100 Subject: [PATCH 5/8] chore: improve log message in ParcellationRegistry --- junifer/data/parcellations/_parcellations.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/junifer/data/parcellations/_parcellations.py b/junifer/data/parcellations/_parcellations.py index 777aaee60..20ddb7e1e 100644 --- a/junifer/data/parcellations/_parcellations.py +++ b/junifer/data/parcellations/_parcellations.py @@ -459,7 +459,7 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): # cannot be "native" due to earlier check if space != target_std_space: logger.debug( - f"Warping {name} to {target_std_space} space using ants." + f"Warping {name} to {target_std_space} space using ANTs." ) raw_img = ANTsParcellationWarper().warp( parcellation_name=name, -- 2.52.0 From 0fe8e1f7e176c5a084d0b9f05fce51a4708de912 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 3 Dec 2024 18:16:04 +0100 Subject: [PATCH 6/8] fix: add an affine correction for non-native template space warping in ParcellationRegistry.get --- junifer/data/parcellations/_parcellations.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/junifer/data/parcellations/_parcellations.py b/junifer/data/parcellations/_parcellations.py index 20ddb7e1e..4ff9b533d 100644 --- a/junifer/data/parcellations/_parcellations.py +++ b/junifer/data/parcellations/_parcellations.py @@ -471,6 +471,8 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): ) # Remove extra dimension added by ANTs img = image.math_img("np.squeeze(img)", img=raw_img) + # Set correct affine as resolution won't be correct + img = image.resample_img(img, target_affine=target_img.affine) else: if target_space != "native": # No warping is going to happen, just resampling, because -- 2.52.0 From 70e73b33bce4872b97e9f66b779ab558f96e33e6 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 4 Dec 2024 16:01:32 +0100 Subject: [PATCH 7/8] fix: use correct interpolation for affine correction while space warping in ParcellationRegistry.get --- junifer/data/parcellations/_parcellations.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/junifer/data/parcellations/_parcellations.py b/junifer/data/parcellations/_parcellations.py index 4ff9b533d..64b95cf1e 100644 --- a/junifer/data/parcellations/_parcellations.py +++ b/junifer/data/parcellations/_parcellations.py @@ -472,7 +472,11 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): # Remove extra dimension added by ANTs img = image.math_img("np.squeeze(img)", img=raw_img) # Set correct affine as resolution won't be correct - img = image.resample_img(img, target_affine=target_img.affine) + img = image.resample_img( + img=img, + target_affine=target_img.affine, + interpolation="nearest", + ) else: if target_space != "native": # No warping is going to happen, just resampling, because -- 2.52.0 From b86d6324151c3f00d591ba984502e33e640f5b9a Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 6 Dec 2024 13:17:51 +0100 Subject: [PATCH 8/8] chore: mark xfail for ReHo comparison test --- junifer/markers/reho/tests/test_reho_parcels.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/junifer/markers/reho/tests/test_reho_parcels.py b/junifer/markers/reho/tests/test_reho_parcels.py index d7737e335..d5746e3a8 100644 --- a/junifer/markers/reho/tests/test_reho_parcels.py +++ b/junifer/markers/reho/tests/test_reho_parcels.py @@ -86,6 +86,10 @@ def test_ReHoParcels(caplog: pytest.LogCaptureFixture, tmp_path: Path) -> None: @pytest.mark.skipif( _check_afni() is False, reason="requires AFNI to be in PATH" ) +@pytest.mark.xfail( + reason="junifer ReHo needs to use the correct mask", + raises=AssertionError, +) def test_ReHoParcels_comparison(tmp_path: Path) -> None: """Test ReHoParcels implementation comparison. -- 2.52.0