diff --git a/junifer/cli/tests/test_cli_utils.py b/junifer/cli/tests/test_cli_utils.py index d8a18da8f..662f54479 100644 --- a/junifer/cli/tests/test_cli_utils.py +++ b/junifer/cli/tests/test_cli_utils.py @@ -48,6 +48,7 @@ def test_get_dependency_information_short() -> None: "lapy", "lazy_loader", "looseversion", + "junifer_data", ] if sys.version_info < (3, 11): diff --git a/junifer/data/coordinates/_coordinates.py b/junifer/data/coordinates/_coordinates.py index 6499ebc47..fa8287ef6 100644 --- a/junifer/data/coordinates/_coordinates.py +++ b/junifer/data/coordinates/_coordinates.py @@ -4,16 +4,18 @@ # Synchon Mandal # License: AGPL +from pathlib import Path from typing import Any, Optional import numpy as np import pandas as pd +from junifer_data import get from numpy.typing import ArrayLike from ...utils import logger, raise_error from ...utils.singleton import Singleton from ..pipeline_data_registry_base import BasePipelineDataRegistry -from ..utils import check_dataset, fetch_file_via_datalad, get_native_warper +from ..utils import JUNIFER_DATA_VERSION, get_dataset_path, get_native_warper from ._ants_coordinates_warper import ANTsCoordinatesWarper from ._fsl_coordinates_warper import FSLCoordinatesWarper @@ -273,23 +275,17 @@ class CoordinatesRegistry(BasePipelineDataRegistry, metaclass=Singleton): # Load data for in-built ones if t_coord.get("file_path_suffix") is not None: - # Get dataset - dataset = check_dataset() # Set file path to retrieve - coords_file_path = ( - dataset.pathobj - / "coordinates" - / name - / t_coord["file_path_suffix"] - ) - logger.debug( - f"Loading coordinates `{name}` from: " - f"{coords_file_path.absolute()!s}" + coords_file_path = Path( + f"coordinates/{name}/{t_coord['file_path_suffix']}" ) + logger.debug(f"Loading coordinates: `{name}`") # Load via pandas df_coords = pd.read_csv( - fetch_file_via_datalad( - dataset=dataset, file_path=coords_file_path + get( + file_path=coords_file_path, + dataset_path=get_dataset_path(), + tag=JUNIFER_DATA_VERSION, ), sep="\t", header=None, diff --git a/junifer/data/masks/_masks.py b/junifer/data/masks/_masks.py index d68a76ac4..3eeba98bc 100644 --- a/junifer/data/masks/_masks.py +++ b/junifer/data/masks/_masks.py @@ -16,6 +16,7 @@ from typing import ( import nibabel as nib import nilearn.image as nimg import numpy as np +from junifer_data import get from nilearn.masking import ( compute_background_mask, compute_epi_mask, @@ -27,9 +28,9 @@ from ...utils.singleton import Singleton from ..pipeline_data_registry_base import BasePipelineDataRegistry from ..template_spaces import get_template from ..utils import ( - check_dataset, + JUNIFER_DATA_VERSION, closest_resolution, - fetch_file_via_datalad, + get_dataset_path, get_native_warper, ) from ._ants_mask_warper import ANTsMaskWarper @@ -37,7 +38,6 @@ from ._fsl_mask_warper import FSLMaskWarper if TYPE_CHECKING: - from datalad.api import Dataset from nibabel.nifti1 import Nifti1Image @@ -406,17 +406,14 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton): mask_img = mask_definition["func"] mask_fname = None elif t_family in ["Vickery-Patil", "UKB"]: - # Get dataset - dataset = check_dataset() # Load mask if t_family == "Vickery-Patil": mask_fname = _load_vickery_patil_mask( - dataset=dataset, name=name, resolution=resolution, ) elif t_family == "UKB": - mask_fname = _load_ukb_mask(dataset=dataset, name=name) + mask_fname = _load_ukb_mask(name=name) else: raise_error(f"Unknown mask family: {t_family}") @@ -698,7 +695,6 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton): def _load_vickery_patil_mask( - dataset: "Dataset", name: str, resolution: Optional[float] = None, ) -> Path: @@ -706,8 +702,6 @@ def _load_vickery_patil_mask( Parameters ---------- - dataset : datalad.api.Dataset - The datalad dataset to fetch mask from. name : {"GM_prob0.2", "GM_prob0.2_cortex"} The name of the mask. resolution : float, optional @@ -748,19 +742,18 @@ def _load_vickery_patil_mask( raise_error(f"Cannot find a Vickery-Patil mask called {name}") # Fetch file - return fetch_file_via_datalad( - dataset=dataset, - file_path=dataset.pathobj / "masks" / "Vickery-Patil" / mask_fname, + return get( + file_path=Path(f"masks/Vickery-Patil/{mask_fname}"), + dataset_path=get_dataset_path(), + tag=JUNIFER_DATA_VERSION, ) -def _load_ukb_mask(dataset: "Dataset", name: str) -> Path: +def _load_ukb_mask(name: str) -> Path: """Load UKB mask. Parameters ---------- - dataset : datalad.api.Dataset - The datalad dataset to fetch mask from. name : {"UKB_15K_GM"} The name of the mask. @@ -782,9 +775,10 @@ def _load_ukb_mask(dataset: "Dataset", name: str) -> Path: raise_error(f"Cannot find a UKB mask called {name}") # Fetch file - return fetch_file_via_datalad( - dataset=dataset, - file_path=dataset.pathobj / "masks" / "UKB" / mask_fname, + return get( + file_path=Path(f"masks/UKB/{mask_fname}"), + dataset_path=get_dataset_path(), + tag=JUNIFER_DATA_VERSION, ) diff --git a/junifer/data/masks/tests/test_masks.py b/junifer/data/masks/tests/test_masks.py index 0a8a2dc4e..a51380062 100644 --- a/junifer/data/masks/tests/test_masks.py +++ b/junifer/data/masks/tests/test_masks.py @@ -26,7 +26,6 @@ from junifer.data.masks._masks import ( _load_ukb_mask, _load_vickery_patil_mask, ) -from junifer.data.utils import check_dataset from junifer.datagrabber import DMCC13Benchmark from junifer.datareader import DefaultDataReader from junifer.testing.datagrabbers import ( @@ -283,9 +282,7 @@ def test_vickery_patil( def test_vickery_patil_error() -> None: """Test error for Vickery-Patil mask.""" with pytest.raises(ValueError, match=r"find a Vickery-Patil mask "): - _load_vickery_patil_mask( - dataset=check_dataset(), name="wrong", resolution=2.0 - ) + _load_vickery_patil_mask(name="wrong", resolution=2.0) def test_ukb() -> None: @@ -300,7 +297,7 @@ def test_ukb() -> None: def test_ukb_error() -> None: """Test error for UKB mask.""" with pytest.raises(ValueError, match=r"find a UKB mask "): - _load_ukb_mask(dataset=check_dataset(), name="wrong") + _load_ukb_mask(name="wrong") def test_get() -> None: diff --git a/junifer/data/parcellations/_parcellations.py b/junifer/data/parcellations/_parcellations.py index 1c4a8adff..9ad4ff80f 100644 --- a/junifer/data/parcellations/_parcellations.py +++ b/junifer/data/parcellations/_parcellations.py @@ -13,14 +13,15 @@ import nibabel as nib import nilearn.image as nimg import numpy as np import pandas as pd +from junifer_data import get from ...utils import logger, raise_error, warn_with_log from ...utils.singleton import Singleton from ..pipeline_data_registry_base import BasePipelineDataRegistry from ..utils import ( - check_dataset, + JUNIFER_DATA_VERSION, closest_resolution, - fetch_file_via_datalad, + get_dataset_path, get_native_warper, ) from ._ants_parcellation_warper import ANTsParcellationWarper @@ -28,7 +29,6 @@ from ._fsl_parcellation_warper import FSLParcellationWarper if TYPE_CHECKING: - from datalad.api import Dataset from nibabel.nifti1 import Nifti1Image @@ -357,49 +357,40 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): "Yan2023", "Brainnetome", ]: - # Get dataset - dataset = check_dataset() # Load parcellation and labels if t_family == "Schaefer2018": parcellation_fname, parcellation_labels = _retrieve_schaefer( - dataset=dataset, resolution=resolution, **parcellation_definition, ) elif t_family == "SUIT": parcellation_fname, parcellation_labels = _retrieve_suit( - dataset=dataset, resolution=resolution, **parcellation_definition, ) elif t_family == "Melbourne": parcellation_fname, parcellation_labels = _retrieve_tian( - dataset=dataset, resolution=resolution, **parcellation_definition, ) elif t_family == "AICHA": parcellation_fname, parcellation_labels = _retrieve_aicha( - dataset=dataset, resolution=resolution, **parcellation_definition, ) elif t_family == "Shen": parcellation_fname, parcellation_labels = _retrieve_shen( - dataset=dataset, resolution=resolution, **parcellation_definition, ) elif t_family == "Yan2023": parcellation_fname, parcellation_labels = _retrieve_yan( - dataset=dataset, resolution=resolution, **parcellation_definition, ) elif t_family == "Brainnetome": parcellation_fname, parcellation_labels = ( _retrieve_brainnetome( - dataset=dataset, resolution=resolution, **parcellation_definition, ) @@ -585,7 +576,6 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton): def _retrieve_schaefer( - dataset: "Dataset", resolution: Optional[float] = None, n_rois: Optional[int] = None, yeo_networks: int = 7, @@ -594,8 +584,6 @@ def _retrieve_schaefer( Parameters ---------- - dataset : datalad.api.Dataset - The datalad dataset to fetch parcellation from. resolution : float, optional The desired resolution of the parcellation to load. If it is not available, the closest resolution will be loaded. Preferably, use a @@ -646,24 +634,18 @@ def _retrieve_schaefer( resolution = closest_resolution(resolution, _valid_resolutions) # Fetch file paths - parcellation_img_path = fetch_file_via_datalad( - dataset=dataset, - file_path=dataset.pathobj - / "parcellations" - / "Schaefer2018" - / "Yeo2011" - / ( - f"Schaefer2018_{n_rois}Parcels_{yeo_networks}Networks_order_" - f"FSLMNI152_{resolution}mm.nii.gz" - ), + path_prefix = Path("parcellations/Schaefer2018/Yeo2011") + parcellation_img_path = get( + file_path=path_prefix / f"Schaefer2018_{n_rois}Parcels_{yeo_networks}" + f"Networks_order_FSLMNI152_{resolution}mm.nii.gz", + dataset_path=get_dataset_path(), + tag=JUNIFER_DATA_VERSION, ) - parcellation_label_path = fetch_file_via_datalad( - dataset=dataset, - file_path=dataset.pathobj - / "parcellations" - / "Schaefer2018" - / "Yeo2011" - / (f"Schaefer2018_{n_rois}Parcels_{yeo_networks}Networks_order.txt"), + parcellation_label_path = get( + file_path=path_prefix + / f"Schaefer2018_{n_rois}Parcels_{yeo_networks}Networks_order.txt", + dataset_path=get_dataset_path(), + tag=JUNIFER_DATA_VERSION, ) # Load labels @@ -678,7 +660,6 @@ def _retrieve_schaefer( def _retrieve_tian( - dataset: "Dataset", resolution: Optional[float] = None, scale: Optional[int] = None, space: str = "MNI152NLin6Asym", @@ -688,8 +669,6 @@ def _retrieve_tian( Parameters ---------- - dataset : datalad.api.Dataset - The datalad dataset to fetch parcellation from. resolution : float, optional The desired resolution of the parcellation to load. If it is not available, the closest resolution will be loaded. Preferably, use a @@ -758,13 +737,8 @@ def _retrieve_tian( # Fetch file paths if magneticfield == "3T": - parcellation_fname_base_3T = ( - dataset.pathobj - / "parcellations" - / "Melbourne" - / "v1.4" - / "3T" - / "Subcortex-Only" + parcellation_fname_base_3T = Path( + "parcellations/Melbourne/v1.4/3T/Subcortex-Only" ) if space == "MNI152NLin6Asym": if resolution == 1: @@ -787,28 +761,29 @@ def _retrieve_tian( f"Tian_Subcortex_S{scale}_{magneticfield}_{space}.nii.gz" ) - parcellation_img_path = fetch_file_via_datalad( - dataset=dataset, + parcellation_img_path = get( file_path=parcellation_fname, + dataset_path=get_dataset_path(), + tag=JUNIFER_DATA_VERSION, ) - parcellation_label_path = fetch_file_via_datalad( - dataset=dataset, + parcellation_label_path = get( file_path=parcellation_fname_base_3T / f"Tian_Subcortex_S{scale}_3T_label.txt", + dataset_path=get_dataset_path(), + tag=JUNIFER_DATA_VERSION, ) # Load labels labels = pd.read_csv(parcellation_label_path, sep=" ", header=None)[ 0 ].to_list() elif magneticfield == "7T": - parcellation_img_path = fetch_file_via_datalad( - dataset=dataset, - file_path=dataset.pathobj - / "parcellations" - / "Melbourne" - / "v1.4" - / "7T" - / f"Tian_Subcortex_S{scale}_{magneticfield}.nii.gz", + parcellation_img_path = get( + file_path=Path( + "parcellations/Melbourne/v1.4/7T/" + f"Tian_Subcortex_S{scale}_{magneticfield}.nii.gz" + ), + dataset_path=get_dataset_path(), + tag=JUNIFER_DATA_VERSION, ) # define 7T labels (b/c currently no labels file available for 7T) scale7Trois = {1: 16, 2: 34, 3: 54, 4: 62} @@ -825,7 +800,6 @@ def _retrieve_tian( def _retrieve_suit( - dataset: "Dataset", resolution: Optional[float], space: str = "MNI152NLin6Asym", ) -> tuple[Path, list[str]]: @@ -833,8 +807,6 @@ def _retrieve_suit( Parameters ---------- - dataset : datalad.api.Dataset - The datalad dataset to fetch parcellation from. resolution : float, optional The desired resolution of the parcellation to load. If it is not available, the closest resolution will be loaded. Preferably, use a @@ -879,19 +851,16 @@ def _retrieve_suit( space = "MNI" # Fetch file paths - parcellation_img_path = fetch_file_via_datalad( - dataset=dataset, - file_path=dataset.pathobj - / "parcellations" - / "SUIT" - / f"SUIT_{space}Space_{resolution}mm.nii", + path_prefix = Path("parcellations/SUIT") + parcellation_img_path = get( + file_path=path_prefix / f"SUIT_{space}Space_{resolution}mm.nii", + dataset_path=get_dataset_path(), + tag=JUNIFER_DATA_VERSION, ) - parcellation_label_path = fetch_file_via_datalad( - dataset=dataset, - file_path=dataset.pathobj - / "parcellations" - / "SUIT" - / f"SUIT_{space}Space_{resolution}mm.tsv", + parcellation_label_path = get( + file_path=path_prefix / f"SUIT_{space}Space_{resolution}mm.tsv", + dataset_path=get_dataset_path(), + tag=JUNIFER_DATA_VERSION, ) # Load labels @@ -903,7 +872,6 @@ def _retrieve_suit( def _retrieve_aicha( - dataset: "Dataset", resolution: Optional[float] = None, version: int = 2, ) -> tuple[Path, list[str]]: @@ -911,8 +879,6 @@ def _retrieve_aicha( Parameters ---------- - dataset : datalad.api.Dataset - The datalad dataset to fetch parcellation from. resolution : float, optional The desired resolution of the parcellation to load. If it is not available, the closest resolution will be loaded. Preferably, use a @@ -968,32 +934,24 @@ def _retrieve_aicha( resolution = closest_resolution(resolution, _valid_resolutions) # Fetch file paths - parcellation_img_path = fetch_file_via_datalad( - dataset=dataset, - file_path=dataset.pathobj - / "parcellations" - / "AICHA" - / f"v{version}" - / "AICHA.nii", + path_prefix = Path(f"parcellations/AICHA/v{version}") + parcellation_img_path = get( + file_path=path_prefix / "AICHA.nii", + dataset_path=get_dataset_path(), + tag=JUNIFER_DATA_VERSION, ) # Conditional label file fetch if version == 1: - parcellation_label_path = fetch_file_via_datalad( - dataset=dataset, - file_path=dataset.pathobj - / "parcellations" - / "AICHA" - / f"v{version}" - / "AICHA_vol1.txt", + parcellation_label_path = get( + file_path=path_prefix / "AICHA_vol1.txt", + dataset_path=get_dataset_path(), + tag=JUNIFER_DATA_VERSION, ) elif version == 2: - parcellation_label_path = fetch_file_via_datalad( - dataset=dataset, - file_path=dataset.pathobj - / "parcellations" - / "AICHA" - / f"v{version}" - / "AICHA_vol3.txt", + parcellation_label_path = get( + file_path=path_prefix / "AICHA_vol3.txt", + dataset_path=get_dataset_path(), + tag=JUNIFER_DATA_VERSION, ) # Load labels @@ -1008,7 +966,6 @@ def _retrieve_aicha( def _retrieve_shen( - dataset: "Dataset", resolution: Optional[float] = None, year: int = 2015, n_rois: int = 268, @@ -1017,8 +974,6 @@ def _retrieve_shen( Parameters ---------- - dataset : datalad.api.Dataset - The datalad dataset to fetch parcellation from. resolution : float, optional The desired resolution of the parcellation to load. If it is not available, the closest resolution will be loaded. Preferably, use a @@ -1095,22 +1050,17 @@ def _retrieve_shen( ) # Fetch file paths based on year + path_prefix = Path(f"parcellations/Shen/{year}") if year == 2013: - parcellation_img_path = fetch_file_via_datalad( - dataset=dataset, - file_path=dataset.pathobj - / "parcellations" - / "Shen" - / "2013" - / f"fconn_atlas_{n_rois}_{resolution}mm.nii", + parcellation_img_path = get( + file_path=path_prefix / f"fconn_atlas_{n_rois}_{resolution}mm.nii", + dataset_path=get_dataset_path(), + tag=JUNIFER_DATA_VERSION, ) - parcellation_label_path = fetch_file_via_datalad( - dataset=dataset, - file_path=dataset.pathobj - / "parcellations" - / "Shen" - / "2013" - / f"Group_seg{n_rois}_BAindexing_setA.txt", + parcellation_label_path = get( + file_path=path_prefix / f"Group_seg{n_rois}_BAindexing_setA.txt", + dataset_path=get_dataset_path(), + tag=JUNIFER_DATA_VERSION, ) labels = ( pd.read_csv( @@ -1123,23 +1073,18 @@ def _retrieve_shen( .to_list() ) elif year == 2015: - parcellation_img_path = fetch_file_via_datalad( - dataset=dataset, - file_path=dataset.pathobj - / "parcellations" - / "Shen" - / "2015" + parcellation_img_path = get( + file_path=path_prefix / f"shen_{resolution}mm_268_parcellation.nii.gz", + dataset_path=get_dataset_path(), + tag=JUNIFER_DATA_VERSION, ) labels = list(range(1, 269)) elif year == 2019: - parcellation_img_path = fetch_file_via_datalad( - dataset=dataset, - file_path=dataset.pathobj - / "parcellations" - / "Shen" - / "2019" - / "Shen_1mm_368_parcellation.nii.gz", + parcellation_img_path = get( + file_path=path_prefix / "Shen_1mm_368_parcellation.nii.gz", + dataset_path=get_dataset_path(), + tag=JUNIFER_DATA_VERSION, ) labels = list(range(1, 369)) @@ -1147,7 +1092,6 @@ def _retrieve_shen( def _retrieve_yan( - dataset: "Dataset", resolution: Optional[float] = None, n_rois: Optional[int] = None, yeo_networks: Optional[int] = None, @@ -1157,8 +1101,6 @@ def _retrieve_yan( Parameters ---------- - dataset : datalad.api.Dataset - The datalad dataset to fetch parcellation from. resolution : float, optional The desired resolution of the parcellation to load. If it is not available, the closest resolution will be loaded. Preferably, use a @@ -1214,6 +1156,7 @@ def _retrieve_yan( ) # Fetch file paths based on networks + pre_path_prefix = Path("parcellations/Yan2023") if yeo_networks: # Check yeo_networks value _valid_yeo_networks = [7, 17] @@ -1223,24 +1166,21 @@ def _retrieve_yan( f"one of the following: {_valid_yeo_networks}" ) - parcellation_img_path = fetch_file_via_datalad( - dataset=dataset, - file_path=dataset.pathobj - / "parcellations" - / "Yan2023" - / "Yeo2011" + path_prefix = pre_path_prefix / "Yeo2011" + parcellation_img_path = get( + file_path=path_prefix / ( f"{n_rois}Parcels_Yeo2011_{yeo_networks}Networks_FSLMNI152_" f"{resolution}mm.nii.gz" ), + dataset_path=get_dataset_path(), + tag=JUNIFER_DATA_VERSION, ) - parcellation_label_path = fetch_file_via_datalad( - dataset=dataset, - file_path=dataset.pathobj - / "parcellations" - / "Yan2023" - / "Yeo2011" + parcellation_label_path = get( + file_path=path_prefix / f"{n_rois}Parcels_Yeo2011_{yeo_networks}Networks_LUT.txt", + dataset_path=get_dataset_path(), + tag=JUNIFER_DATA_VERSION, ) elif kong_networks: # Check kong_networks value @@ -1251,24 +1191,21 @@ def _retrieve_yan( f"one of the following: {_valid_kong_networks}" ) - parcellation_img_path = fetch_file_via_datalad( - dataset=dataset, - file_path=dataset.pathobj - / "parcellations" - / "Yan2023" - / "Kong2022" + path_prefix = pre_path_prefix / "Kong2022" + parcellation_img_path = get( + file_path=path_prefix / ( f"{n_rois}Parcels_Kong2022_{kong_networks}Networks_FSLMNI152_" f"{resolution}mm.nii.gz" ), + dataset_path=get_dataset_path(), + tag=JUNIFER_DATA_VERSION, ) - parcellation_label_path = fetch_file_via_datalad( - dataset=dataset, - file_path=dataset.pathobj - / "parcellations" - / "Yan2023" - / "Kong2022" + parcellation_label_path = get( + file_path=path_prefix / f"{n_rois}Parcels_Kong2022_{kong_networks}Networks_LUT.txt", + dataset_path=get_dataset_path(), + tag=JUNIFER_DATA_VERSION, ) # Load label file @@ -1280,7 +1217,6 @@ def _retrieve_yan( def _retrieve_brainnetome( - dataset: "Dataset", resolution: Optional[float] = None, threshold: Optional[int] = None, ) -> tuple[Path, list[str]]: @@ -1288,8 +1224,6 @@ def _retrieve_brainnetome( Parameters ---------- - dataset : datalad.api.Dataset - The datalad dataset to fetch parcellation from. resolution : {1.0, 1.25, 2.0}, optional The desired resolution of the parcellation to load. If it is not available, the closest resolution will be loaded. Preferably, use a @@ -1332,12 +1266,13 @@ def _retrieve_brainnetome( resolution = int(resolution) # Fetch file path - parcellation_img_path = fetch_file_via_datalad( - dataset=dataset, - file_path=dataset.pathobj - / "parcellations" - / "Brainnetome" - / f"BNA-maxprob-thr{threshold}-{resolution}mm.nii.gz", + parcellation_img_path = get( + file_path=Path( + "parcellations/Brainnetome/" + f"BNA-maxprob-thr{threshold}-{resolution}mm.nii.gz" + ), + dataset_path=get_dataset_path(), + tag=JUNIFER_DATA_VERSION, ) # Load labels diff --git a/junifer/data/parcellations/tests/test_parcellations.py b/junifer/data/parcellations/tests/test_parcellations.py index 247956622..6765aafe4 100644 --- a/junifer/data/parcellations/tests/test_parcellations.py +++ b/junifer/data/parcellations/tests/test_parcellations.py @@ -24,7 +24,6 @@ from junifer.data.parcellations._parcellations import ( _retrieve_tian, _retrieve_yan, ) -from junifer.data.utils import check_dataset from junifer.datareader import DefaultDataReader from junifer.pipeline.utils import _check_ants from junifer.testing.datagrabbers import ( @@ -335,7 +334,6 @@ def test_retrieve_schaefer_incorrect_n_rois() -> None: """Test retrieve Schaefer with incorrect ROIs.""" with pytest.raises(ValueError, match=r"The parameter `n_rois`"): _retrieve_schaefer( - dataset=check_dataset(), resolution=1, n_rois=101, yeo_networks=7, @@ -346,7 +344,6 @@ def test_retrieve_schaefer_incorrect_yeo_networks() -> None: """Test retrieve Schaefer with incorrect Yeo networks.""" with pytest.raises(ValueError, match=r"The parameter `yeo_networks`"): _retrieve_schaefer( - dataset=check_dataset(), resolution=1, n_rois=100, yeo_networks=8, @@ -384,7 +381,7 @@ def test_suit(space_key: str, space: str) -> None: def test_retrieve_suit_incorrect_space() -> None: """Test retrieve SUIT with incorrect space.""" with pytest.raises(ValueError, match=r"The parameter `space`"): - _retrieve_suit(dataset=check_dataset(), resolution=1.0, space="wrong") + _retrieve_suit(resolution=1.0, space="wrong") @pytest.mark.parametrize( @@ -512,13 +509,10 @@ def test_tian_7T_6thgeneration(scale: int, n_label: int) -> None: def test_retrieve_tian_incorrect_space() -> None: """Test retrieve tian with incorrect space.""" with pytest.raises(ValueError, match=r"The parameter `space`"): - _retrieve_tian( - dataset=check_dataset(), resolution=1, scale=1, space="wrong" - ) + _retrieve_tian(resolution=1, scale=1, space="wrong") with pytest.raises(ValueError, match=r"MNI152NLin6Asym"): _retrieve_tian( - dataset=check_dataset(), resolution=1, scale=1, magneticfield="7T", @@ -530,7 +524,6 @@ def test_retrieve_tian_incorrect_magneticfield() -> None: """Test retrieve tian with incorrect magneticfield.""" with pytest.raises(ValueError, match=r"The parameter `magneticfield`"): _retrieve_tian( - dataset=check_dataset(), resolution=1, scale=1, magneticfield="wrong", @@ -541,7 +534,6 @@ def test_retrieve_tian_incorrect_scale(tmp_path: Path) -> None: """Test retrieve tian with incorrect scale.""" with pytest.raises(ValueError, match=r"The parameter `scale`"): _retrieve_tian( - dataset=check_dataset(), resolution=1, scale=5, space="MNI152NLin6Asym", @@ -577,7 +569,6 @@ def test_retrieve_aicha_incorrect_version() -> None: """Test retrieve AICHA with incorrect version.""" with pytest.raises(ValueError, match="The parameter `version`"): _retrieve_aicha( - dataset=check_dataset(), version=100, ) @@ -639,7 +630,6 @@ def test_retrieve_shen_incorrect_year() -> None: """Test retrieve Shen with incorrect year.""" with pytest.raises(ValueError, match="The parameter `year`"): _retrieve_shen( - dataset=check_dataset(), year=1969, ) @@ -648,7 +638,6 @@ def test_retrieve_shen_incorrect_n_rois() -> None: """Test retrieve Shen with incorrect ROIs.""" with pytest.raises(ValueError, match="The parameter `n_rois`"): _retrieve_shen( - dataset=check_dataset(), year=2015, n_rois=10, ) @@ -691,7 +680,6 @@ def test_retrieve_shen_incorrect_param_combo( """ with pytest.raises(ValueError, match="The parameter combination"): _retrieve_shen( - dataset=check_dataset(), resolution=resolution, year=year, n_rois=n_rois, @@ -819,7 +807,6 @@ def test_retrieve_yan_incorrect_networks() -> None: ValueError, match="Either one of `yeo_networks` or `kong_networks`" ): _retrieve_yan( - dataset=check_dataset(), n_rois=31418, yeo_networks=100, kong_networks=100, @@ -829,7 +816,6 @@ def test_retrieve_yan_incorrect_networks() -> None: ValueError, match="Either one of `yeo_networks` or `kong_networks`" ): _retrieve_yan( - dataset=check_dataset(), n_rois=31418, yeo_networks=None, kong_networks=None, @@ -840,7 +826,6 @@ def test_retrieve_yan_incorrect_n_rois() -> None: """Test retrieve Yan with incorrect ROIs.""" with pytest.raises(ValueError, match="The parameter `n_rois`"): _retrieve_yan( - dataset=check_dataset(), n_rois=31418, yeo_networks=7, ) @@ -850,7 +835,6 @@ def test_retrieve_yan_incorrect_yeo_networks() -> None: """Test retrieve Yan with incorrect Yeo networks.""" with pytest.raises(ValueError, match="The parameter `yeo_networks`"): _retrieve_yan( - dataset=check_dataset(), n_rois=100, yeo_networks=27, ) @@ -860,7 +844,6 @@ def test_retrieve_yan_incorrect_kong_networks() -> None: """Test retrieve Yan with incorrect Kong networks.""" with pytest.raises(ValueError, match="The parameter `kong_networks`"): _retrieve_yan( - dataset=check_dataset(), n_rois=100, kong_networks=27, ) @@ -922,7 +905,6 @@ def test_retrieve_brainnetome_incorrect_threshold() -> None: """Test retrieve Brainnetome with incorrect threshold.""" with pytest.raises(ValueError, match="The parameter `threshold`"): _retrieve_brainnetome( - dataset=check_dataset(), threshold=100, ) diff --git a/junifer/data/template_spaces.py b/junifer/data/template_spaces.py index f0705a090..65c74c7c2 100644 --- a/junifer/data/template_spaces.py +++ b/junifer/data/template_spaces.py @@ -8,10 +8,11 @@ from typing import Any, Optional, Union import nibabel as nib import numpy as np +from junifer_data import get from templateflow import api as tflow from ..utils import logger, raise_error -from .utils import check_dataset, closest_resolution, fetch_file_via_datalad +from .utils import JUNIFER_DATA_VERSION, closest_resolution, get_dataset_path __all__ = ["get_template", "get_xfm"] @@ -33,17 +34,14 @@ def get_xfm(src: str, dst: str) -> Path: # pragma: no cover The path to the transformation file. """ - # Get dataset - dataset = check_dataset() # Set file path to retrieve - xfm_file_path = ( - dataset.pathobj - / "xfms" - / f"{src}_to_{dst}" - / f"{src}_to_{dst}_Composite.h5" - ) + xfm_file_path = Path(f"xfms/{src}_to_{dst}/{src}_to_{dst}_Composite.h5") # Retrieve file - return fetch_file_via_datalad(dataset=dataset, file_path=xfm_file_path) + return get( + file_path=xfm_file_path, + dataset_path=get_dataset_path(), + tag=JUNIFER_DATA_VERSION, + ) def get_template( diff --git a/junifer/data/utils.py b/junifer/data/utils.py index 29f3d0dac..2b12d8033 100644 --- a/junifer/data/utils.py +++ b/junifer/data/utils.py @@ -8,21 +8,23 @@ from collections.abc import MutableMapping from pathlib import Path from typing import Optional, Union -import datalad.api as dl import numpy as np -from datalad.support.exceptions import IncompleteResultsError from ..utils import config, logger, raise_error __all__ = [ - "check_dataset", + "JUNIFER_DATA_VERSION", "closest_resolution", - "fetch_file_via_datalad", + "get_dataset_path", "get_native_warper", ] +# junifer-data version constant +JUNIFER_DATA_VERSION = "1" + + def closest_resolution( resolution: Optional[Union[float, int]], valid_resolution: Union[list[float], list[int], np.ndarray], @@ -124,94 +126,17 @@ def get_native_warper( return possible_warpers[0] -def check_dataset() -> dl.Dataset: - """Get or install junifer-data dataset. +def get_dataset_path() -> Optional[Path]: + """Get junifer-data dataset path. Returns ------- - datalad.api.Dataset - The junifer-data dataset. - - Raises - ------ - RuntimeError - If there is a problem cloning the dataset. + pathlib.Path or None + Path to the dataset or None. """ - # Check config and set default if not passed - data_dir = config.get("data.location") - if data_dir is not None: - data_dir = Path(data_dir) - else: - data_dir = Path().home() / "junifer_data" - - # Check if the dataset is installed at storage path; - # else clone a fresh copy - if dl.Dataset(data_dir).is_installed(): - logger.debug(f"Found existing junifer-data at: {data_dir.resolve()}") - return dl.Dataset(data_dir) - else: - logger.debug(f"Cloning junifer-data to: {data_dir.resolve()}") - # Clone dataset - try: - dataset = dl.clone( - "https://github.com/juaml/junifer-data.git", - path=data_dir, - result_renderer="disabled", - ) - except IncompleteResultsError as e: - raise_error( - msg=f"Failed to clone junifer-data: {e.failed}", - klass=RuntimeError, - ) - else: - logger.debug( - f"Successfully cloned junifer-data to: " - f"{data_dir.resolve()}" - ) - return dataset - - -def fetch_file_via_datalad(dataset: dl.Dataset, file_path: Path) -> Path: - """Fetch `file_path` from `dataset` via datalad. - - Parameters - ---------- - dataset : datalad.api.Dataset - The datalad dataset to fetch files from. - file_path : pathlib.Path - The file path to fetch. - - Returns - ------- - pathlib.Path - Resolved fetched file path. - - Raises - ------ - RuntimeError - If there is a problem fetching the file. - - """ - try: - got = dataset.get(file_path, result_renderer="disabled") - except IncompleteResultsError as e: - raise_error( - msg=f"Failed to get file from dataset: {e.failed}", - klass=RuntimeError, - ) - else: - got_path = Path(got[0]["path"]) - # Conditional logging based on file fetch - status = got[0]["status"] - if status == "ok": - logger.info(f"Successfully fetched file: {got_path.resolve()}") - return got_path - elif status == "notneeded": - logger.debug(f"Found existing file: {got_path.resolve()}") - return got_path - else: - raise_error( - msg=f"Failed to fetch file: {got_path.resolve()}", - klass=RuntimeError, - ) + return ( + Path(config.get("data.location")) + if config.get("data.location") is not None + else None + ) diff --git a/junifer/typing/_typing.py b/junifer/typing/_typing.py index b0887241a..cd2311320 100644 --- a/junifer/typing/_typing.py +++ b/junifer/typing/_typing.py @@ -63,6 +63,6 @@ MarkerInOutMappings = MutableMapping[str, MutableMapping[str, str]] DataGrabberPatterns = dict[ str, Union[dict[str, str], Sequence[dict[str, str]]] ] -ConfigVal = Union[bool, int, float] +ConfigVal = Union[bool, int, float, str] Element = Union[str, tuple[str, ...]] Elements = Sequence[Element] diff --git a/pyproject.toml b/pyproject.toml index d0b4f6274..4d7851ddc 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -52,6 +52,7 @@ dependencies = [ "lazy_loader==0.4", "importlib_metadata; python_version<'3.9'", "looseversion==1.3.0; python_version>='3.12'", + "junifer_data==1.1.0", ] dynamic = ["version"] @@ -212,6 +213,7 @@ known-third-party = [ "brainprint", "lapy", "pytest", + "junifer_data", ] [tool.ruff.lint.mccabe]