[ENH]: Move data downloading/handling to junifer-data package. #363

Merged
synchon merged 13 commits from refactor/junifer-data-api into main 2025-02-13 14:11:39 +00:00
10 changed files with 148 additions and 318 deletions

View file

@ -48,6 +48,7 @@ def test_get_dependency_information_short() -> None:
"lapy", "lapy",
"lazy_loader", "lazy_loader",
"looseversion", "looseversion",
"junifer_data",
] ]
if sys.version_info < (3, 11): if sys.version_info < (3, 11):

View file

@ -4,16 +4,18 @@
# Synchon Mandal <s.mandal@fz-juelich.de> # Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from pathlib import Path
from typing import Any, Optional from typing import Any, Optional
import numpy as np import numpy as np
import pandas as pd import pandas as pd
from junifer_data import get
from numpy.typing import ArrayLike from numpy.typing import ArrayLike
from ...utils import logger, raise_error from ...utils import logger, raise_error
from ...utils.singleton import Singleton from ...utils.singleton import Singleton
from ..pipeline_data_registry_base import BasePipelineDataRegistry 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 ._ants_coordinates_warper import ANTsCoordinatesWarper
from ._fsl_coordinates_warper import FSLCoordinatesWarper from ._fsl_coordinates_warper import FSLCoordinatesWarper
@ -273,23 +275,17 @@ class CoordinatesRegistry(BasePipelineDataRegistry, metaclass=Singleton):
# Load data for in-built ones # Load data for in-built ones
if t_coord.get("file_path_suffix") is not None: if t_coord.get("file_path_suffix") is not None:
# Get dataset
dataset = check_dataset()
# Set file path to retrieve # Set file path to retrieve
coords_file_path = ( coords_file_path = Path(
dataset.pathobj f"coordinates/{name}/{t_coord['file_path_suffix']}"
/ "coordinates"
/ name
/ t_coord["file_path_suffix"]
)
logger.debug(
f"Loading coordinates `{name}` from: "
f"{coords_file_path.absolute()!s}"
) )
logger.debug(f"Loading coordinates: `{name}`")
# Load via pandas # Load via pandas
df_coords = pd.read_csv( df_coords = pd.read_csv(
fetch_file_via_datalad( get(
dataset=dataset, file_path=coords_file_path file_path=coords_file_path,
dataset_path=get_dataset_path(),
tag=JUNIFER_DATA_VERSION,
), ),
sep="\t", sep="\t",
header=None, header=None,

View file

@ -16,6 +16,7 @@ from typing import (
import nibabel as nib import nibabel as nib
import nilearn.image as nimg import nilearn.image as nimg
import numpy as np import numpy as np
from junifer_data import get
from nilearn.masking import ( from nilearn.masking import (
compute_background_mask, compute_background_mask,
compute_epi_mask, compute_epi_mask,
@ -27,9 +28,9 @@ from ...utils.singleton import Singleton
from ..pipeline_data_registry_base import BasePipelineDataRegistry from ..pipeline_data_registry_base import BasePipelineDataRegistry
from ..template_spaces import get_template from ..template_spaces import get_template
from ..utils import ( from ..utils import (
check_dataset, JUNIFER_DATA_VERSION,
closest_resolution, closest_resolution,
fetch_file_via_datalad, get_dataset_path,
get_native_warper, get_native_warper,
) )
from ._ants_mask_warper import ANTsMaskWarper from ._ants_mask_warper import ANTsMaskWarper
@ -37,7 +38,6 @@ from ._fsl_mask_warper import FSLMaskWarper
if TYPE_CHECKING: if TYPE_CHECKING:
from datalad.api import Dataset
from nibabel.nifti1 import Nifti1Image from nibabel.nifti1 import Nifti1Image
@ -406,17 +406,14 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
mask_img = mask_definition["func"] mask_img = mask_definition["func"]
mask_fname = None mask_fname = None
elif t_family in ["Vickery-Patil", "UKB"]: elif t_family in ["Vickery-Patil", "UKB"]:
# Get dataset
dataset = check_dataset()
# Load mask # Load mask
if t_family == "Vickery-Patil": if t_family == "Vickery-Patil":
mask_fname = _load_vickery_patil_mask( mask_fname = _load_vickery_patil_mask(
dataset=dataset,
name=name, name=name,
resolution=resolution, resolution=resolution,
) )
elif t_family == "UKB": elif t_family == "UKB":
mask_fname = _load_ukb_mask(dataset=dataset, name=name) mask_fname = _load_ukb_mask(name=name)
else: else:
raise_error(f"Unknown mask family: {t_family}") raise_error(f"Unknown mask family: {t_family}")
@ -698,7 +695,6 @@ class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
def _load_vickery_patil_mask( def _load_vickery_patil_mask(
dataset: "Dataset",
name: str, name: str,
resolution: Optional[float] = None, resolution: Optional[float] = None,
) -> Path: ) -> Path:
@ -706,8 +702,6 @@ def _load_vickery_patil_mask(
Parameters Parameters
---------- ----------
dataset : datalad.api.Dataset
The datalad dataset to fetch mask from.
name : {"GM_prob0.2", "GM_prob0.2_cortex"} name : {"GM_prob0.2", "GM_prob0.2_cortex"}
The name of the mask. The name of the mask.
resolution : float, optional resolution : float, optional
@ -748,19 +742,18 @@ def _load_vickery_patil_mask(
raise_error(f"Cannot find a Vickery-Patil mask called {name}") raise_error(f"Cannot find a Vickery-Patil mask called {name}")
# Fetch file # Fetch file
return fetch_file_via_datalad( return get(
dataset=dataset, file_path=Path(f"masks/Vickery-Patil/{mask_fname}"),
file_path=dataset.pathobj / "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. """Load UKB mask.
Parameters Parameters
---------- ----------
dataset : datalad.api.Dataset
The datalad dataset to fetch mask from.
name : {"UKB_15K_GM"} name : {"UKB_15K_GM"}
The name of the mask. 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}") raise_error(f"Cannot find a UKB mask called {name}")
# Fetch file # Fetch file
return fetch_file_via_datalad( return get(
dataset=dataset, file_path=Path(f"masks/UKB/{mask_fname}"),
file_path=dataset.pathobj / "masks" / "UKB" / mask_fname, dataset_path=get_dataset_path(),
tag=JUNIFER_DATA_VERSION,
) )

View file

@ -26,7 +26,6 @@ from junifer.data.masks._masks import (
_load_ukb_mask, _load_ukb_mask,
_load_vickery_patil_mask, _load_vickery_patil_mask,
) )
from junifer.data.utils import check_dataset
from junifer.datagrabber import DMCC13Benchmark from junifer.datagrabber import DMCC13Benchmark
from junifer.datareader import DefaultDataReader from junifer.datareader import DefaultDataReader
from junifer.testing.datagrabbers import ( from junifer.testing.datagrabbers import (
@ -283,9 +282,7 @@ def test_vickery_patil(
def test_vickery_patil_error() -> None: def test_vickery_patil_error() -> None:
"""Test error for Vickery-Patil mask.""" """Test error for Vickery-Patil mask."""
with pytest.raises(ValueError, match=r"find a Vickery-Patil mask "): with pytest.raises(ValueError, match=r"find a Vickery-Patil mask "):
_load_vickery_patil_mask( _load_vickery_patil_mask(name="wrong", resolution=2.0)
dataset=check_dataset(), name="wrong", resolution=2.0
)
def test_ukb() -> None: def test_ukb() -> None:
@ -300,7 +297,7 @@ def test_ukb() -> None:
def test_ukb_error() -> None: def test_ukb_error() -> None:
"""Test error for UKB mask.""" """Test error for UKB mask."""
with pytest.raises(ValueError, match=r"find a 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: def test_get() -> None:

View file

@ -13,14 +13,15 @@ import nibabel as nib
import nilearn.image as nimg import nilearn.image as nimg
import numpy as np import numpy as np
import pandas as pd import pandas as pd
from junifer_data import get
from ...utils import logger, raise_error, warn_with_log from ...utils import logger, raise_error, warn_with_log
from ...utils.singleton import Singleton from ...utils.singleton import Singleton
from ..pipeline_data_registry_base import BasePipelineDataRegistry from ..pipeline_data_registry_base import BasePipelineDataRegistry
from ..utils import ( from ..utils import (
check_dataset, JUNIFER_DATA_VERSION,
closest_resolution, closest_resolution,
fetch_file_via_datalad, get_dataset_path,
get_native_warper, get_native_warper,
) )
from ._ants_parcellation_warper import ANTsParcellationWarper from ._ants_parcellation_warper import ANTsParcellationWarper
@ -28,7 +29,6 @@ from ._fsl_parcellation_warper import FSLParcellationWarper
if TYPE_CHECKING: if TYPE_CHECKING:
from datalad.api import Dataset
from nibabel.nifti1 import Nifti1Image from nibabel.nifti1 import Nifti1Image
@ -357,49 +357,40 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton):
"Yan2023", "Yan2023",
"Brainnetome", "Brainnetome",
]: ]:
# Get dataset
dataset = check_dataset()
# Load parcellation and labels # Load parcellation and labels
if t_family == "Schaefer2018": if t_family == "Schaefer2018":
parcellation_fname, parcellation_labels = _retrieve_schaefer( parcellation_fname, parcellation_labels = _retrieve_schaefer(
dataset=dataset,
resolution=resolution, resolution=resolution,
**parcellation_definition, **parcellation_definition,
) )
elif t_family == "SUIT": elif t_family == "SUIT":
parcellation_fname, parcellation_labels = _retrieve_suit( parcellation_fname, parcellation_labels = _retrieve_suit(
dataset=dataset,
resolution=resolution, resolution=resolution,
**parcellation_definition, **parcellation_definition,
) )
elif t_family == "Melbourne": elif t_family == "Melbourne":
parcellation_fname, parcellation_labels = _retrieve_tian( parcellation_fname, parcellation_labels = _retrieve_tian(
dataset=dataset,
resolution=resolution, resolution=resolution,
**parcellation_definition, **parcellation_definition,
) )
elif t_family == "AICHA": elif t_family == "AICHA":
parcellation_fname, parcellation_labels = _retrieve_aicha( parcellation_fname, parcellation_labels = _retrieve_aicha(
dataset=dataset,
resolution=resolution, resolution=resolution,
**parcellation_definition, **parcellation_definition,
) )
elif t_family == "Shen": elif t_family == "Shen":
parcellation_fname, parcellation_labels = _retrieve_shen( parcellation_fname, parcellation_labels = _retrieve_shen(
dataset=dataset,
resolution=resolution, resolution=resolution,
**parcellation_definition, **parcellation_definition,
) )
elif t_family == "Yan2023": elif t_family == "Yan2023":
parcellation_fname, parcellation_labels = _retrieve_yan( parcellation_fname, parcellation_labels = _retrieve_yan(
dataset=dataset,
resolution=resolution, resolution=resolution,
**parcellation_definition, **parcellation_definition,
) )
elif t_family == "Brainnetome": elif t_family == "Brainnetome":
parcellation_fname, parcellation_labels = ( parcellation_fname, parcellation_labels = (
_retrieve_brainnetome( _retrieve_brainnetome(
dataset=dataset,
resolution=resolution, resolution=resolution,
**parcellation_definition, **parcellation_definition,
) )
@ -585,7 +576,6 @@ class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton):
def _retrieve_schaefer( def _retrieve_schaefer(
dataset: "Dataset",
resolution: Optional[float] = None, resolution: Optional[float] = None,
n_rois: Optional[int] = None, n_rois: Optional[int] = None,
yeo_networks: int = 7, yeo_networks: int = 7,
@ -594,8 +584,6 @@ def _retrieve_schaefer(
Parameters Parameters
---------- ----------
dataset : datalad.api.Dataset
The datalad dataset to fetch parcellation from.
resolution : float, optional resolution : float, optional
The desired resolution of the parcellation to load. If it is not The desired resolution of the parcellation to load. If it is not
available, the closest resolution will be loaded. Preferably, use a available, the closest resolution will be loaded. Preferably, use a
@ -646,24 +634,18 @@ def _retrieve_schaefer(
resolution = closest_resolution(resolution, _valid_resolutions) resolution = closest_resolution(resolution, _valid_resolutions)
# Fetch file paths # Fetch file paths
parcellation_img_path = fetch_file_via_datalad( path_prefix = Path("parcellations/Schaefer2018/Yeo2011")
dataset=dataset, parcellation_img_path = get(
file_path=dataset.pathobj file_path=path_prefix / f"Schaefer2018_{n_rois}Parcels_{yeo_networks}"
/ "parcellations" f"Networks_order_FSLMNI152_{resolution}mm.nii.gz",
/ "Schaefer2018" dataset_path=get_dataset_path(),
/ "Yeo2011" tag=JUNIFER_DATA_VERSION,
/ (
f"Schaefer2018_{n_rois}Parcels_{yeo_networks}Networks_order_"
f"FSLMNI152_{resolution}mm.nii.gz"
),
) )
parcellation_label_path = fetch_file_via_datalad( parcellation_label_path = get(
dataset=dataset, file_path=path_prefix
file_path=dataset.pathobj / f"Schaefer2018_{n_rois}Parcels_{yeo_networks}Networks_order.txt",
/ "parcellations" dataset_path=get_dataset_path(),
/ "Schaefer2018" tag=JUNIFER_DATA_VERSION,
/ "Yeo2011"
/ (f"Schaefer2018_{n_rois}Parcels_{yeo_networks}Networks_order.txt"),
) )
# Load labels # Load labels
@ -678,7 +660,6 @@ def _retrieve_schaefer(
def _retrieve_tian( def _retrieve_tian(
dataset: "Dataset",
resolution: Optional[float] = None, resolution: Optional[float] = None,
scale: Optional[int] = None, scale: Optional[int] = None,
space: str = "MNI152NLin6Asym", space: str = "MNI152NLin6Asym",
@ -688,8 +669,6 @@ def _retrieve_tian(
Parameters Parameters
---------- ----------
dataset : datalad.api.Dataset
The datalad dataset to fetch parcellation from.
resolution : float, optional resolution : float, optional
The desired resolution of the parcellation to load. If it is not The desired resolution of the parcellation to load. If it is not
available, the closest resolution will be loaded. Preferably, use a available, the closest resolution will be loaded. Preferably, use a
@ -758,13 +737,8 @@ def _retrieve_tian(
# Fetch file paths # Fetch file paths
if magneticfield == "3T": if magneticfield == "3T":
parcellation_fname_base_3T = ( parcellation_fname_base_3T = Path(
dataset.pathobj "parcellations/Melbourne/v1.4/3T/Subcortex-Only"
/ "parcellations"
/ "Melbourne"
/ "v1.4"
/ "3T"
/ "Subcortex-Only"
) )
if space == "MNI152NLin6Asym": if space == "MNI152NLin6Asym":
if resolution == 1: if resolution == 1:
@ -787,28 +761,29 @@ def _retrieve_tian(
f"Tian_Subcortex_S{scale}_{magneticfield}_{space}.nii.gz" f"Tian_Subcortex_S{scale}_{magneticfield}_{space}.nii.gz"
) )
parcellation_img_path = fetch_file_via_datalad( parcellation_img_path = get(
dataset=dataset,
file_path=parcellation_fname, file_path=parcellation_fname,
dataset_path=get_dataset_path(),
tag=JUNIFER_DATA_VERSION,
) )
parcellation_label_path = fetch_file_via_datalad( parcellation_label_path = get(
dataset=dataset,
file_path=parcellation_fname_base_3T file_path=parcellation_fname_base_3T
/ f"Tian_Subcortex_S{scale}_3T_label.txt", / f"Tian_Subcortex_S{scale}_3T_label.txt",
dataset_path=get_dataset_path(),
tag=JUNIFER_DATA_VERSION,
) )
# Load labels # Load labels
labels = pd.read_csv(parcellation_label_path, sep=" ", header=None)[ labels = pd.read_csv(parcellation_label_path, sep=" ", header=None)[
0 0
].to_list() ].to_list()
elif magneticfield == "7T": elif magneticfield == "7T":
parcellation_img_path = fetch_file_via_datalad( parcellation_img_path = get(
dataset=dataset, file_path=Path(
file_path=dataset.pathobj "parcellations/Melbourne/v1.4/7T/"
/ "parcellations" f"Tian_Subcortex_S{scale}_{magneticfield}.nii.gz"
/ "Melbourne" ),
/ "v1.4" dataset_path=get_dataset_path(),
/ "7T" tag=JUNIFER_DATA_VERSION,
/ f"Tian_Subcortex_S{scale}_{magneticfield}.nii.gz",
) )
# define 7T labels (b/c currently no labels file available for 7T) # define 7T labels (b/c currently no labels file available for 7T)
scale7Trois = {1: 16, 2: 34, 3: 54, 4: 62} scale7Trois = {1: 16, 2: 34, 3: 54, 4: 62}
@ -825,7 +800,6 @@ def _retrieve_tian(
def _retrieve_suit( def _retrieve_suit(
dataset: "Dataset",
resolution: Optional[float], resolution: Optional[float],
space: str = "MNI152NLin6Asym", space: str = "MNI152NLin6Asym",
) -> tuple[Path, list[str]]: ) -> tuple[Path, list[str]]:
@ -833,8 +807,6 @@ def _retrieve_suit(
Parameters Parameters
---------- ----------
dataset : datalad.api.Dataset
The datalad dataset to fetch parcellation from.
resolution : float, optional resolution : float, optional
The desired resolution of the parcellation to load. If it is not The desired resolution of the parcellation to load. If it is not
available, the closest resolution will be loaded. Preferably, use a available, the closest resolution will be loaded. Preferably, use a
@ -879,19 +851,16 @@ def _retrieve_suit(
space = "MNI" space = "MNI"
# Fetch file paths # Fetch file paths
parcellation_img_path = fetch_file_via_datalad( path_prefix = Path("parcellations/SUIT")
dataset=dataset, parcellation_img_path = get(
file_path=dataset.pathobj file_path=path_prefix / f"SUIT_{space}Space_{resolution}mm.nii",
/ "parcellations" dataset_path=get_dataset_path(),
/ "SUIT" tag=JUNIFER_DATA_VERSION,
/ f"SUIT_{space}Space_{resolution}mm.nii",
) )
parcellation_label_path = fetch_file_via_datalad( parcellation_label_path = get(
dataset=dataset, file_path=path_prefix / f"SUIT_{space}Space_{resolution}mm.tsv",
file_path=dataset.pathobj dataset_path=get_dataset_path(),
/ "parcellations" tag=JUNIFER_DATA_VERSION,
/ "SUIT"
/ f"SUIT_{space}Space_{resolution}mm.tsv",
) )
# Load labels # Load labels
@ -903,7 +872,6 @@ def _retrieve_suit(
def _retrieve_aicha( def _retrieve_aicha(
dataset: "Dataset",
resolution: Optional[float] = None, resolution: Optional[float] = None,
version: int = 2, version: int = 2,
) -> tuple[Path, list[str]]: ) -> tuple[Path, list[str]]:
@ -911,8 +879,6 @@ def _retrieve_aicha(
Parameters Parameters
---------- ----------
dataset : datalad.api.Dataset
The datalad dataset to fetch parcellation from.
resolution : float, optional resolution : float, optional
The desired resolution of the parcellation to load. If it is not The desired resolution of the parcellation to load. If it is not
available, the closest resolution will be loaded. Preferably, use a available, the closest resolution will be loaded. Preferably, use a
@ -968,32 +934,24 @@ def _retrieve_aicha(
resolution = closest_resolution(resolution, _valid_resolutions) resolution = closest_resolution(resolution, _valid_resolutions)
# Fetch file paths # Fetch file paths
parcellation_img_path = fetch_file_via_datalad( path_prefix = Path(f"parcellations/AICHA/v{version}")
dataset=dataset, parcellation_img_path = get(
file_path=dataset.pathobj file_path=path_prefix / "AICHA.nii",
/ "parcellations" dataset_path=get_dataset_path(),
/ "AICHA" tag=JUNIFER_DATA_VERSION,
/ f"v{version}"
/ "AICHA.nii",
) )
# Conditional label file fetch # Conditional label file fetch
if version == 1: if version == 1:
parcellation_label_path = fetch_file_via_datalad( parcellation_label_path = get(
dataset=dataset, file_path=path_prefix / "AICHA_vol1.txt",
file_path=dataset.pathobj dataset_path=get_dataset_path(),
/ "parcellations" tag=JUNIFER_DATA_VERSION,
/ "AICHA"
/ f"v{version}"
/ "AICHA_vol1.txt",
) )
elif version == 2: elif version == 2:
parcellation_label_path = fetch_file_via_datalad( parcellation_label_path = get(
dataset=dataset, file_path=path_prefix / "AICHA_vol3.txt",
file_path=dataset.pathobj dataset_path=get_dataset_path(),
/ "parcellations" tag=JUNIFER_DATA_VERSION,
/ "AICHA"
/ f"v{version}"
/ "AICHA_vol3.txt",
) )
# Load labels # Load labels
@ -1008,7 +966,6 @@ def _retrieve_aicha(
def _retrieve_shen( def _retrieve_shen(
dataset: "Dataset",
resolution: Optional[float] = None, resolution: Optional[float] = None,
year: int = 2015, year: int = 2015,
n_rois: int = 268, n_rois: int = 268,
@ -1017,8 +974,6 @@ def _retrieve_shen(
Parameters Parameters
---------- ----------
dataset : datalad.api.Dataset
The datalad dataset to fetch parcellation from.
resolution : float, optional resolution : float, optional
The desired resolution of the parcellation to load. If it is not The desired resolution of the parcellation to load. If it is not
available, the closest resolution will be loaded. Preferably, use a available, the closest resolution will be loaded. Preferably, use a
@ -1095,22 +1050,17 @@ def _retrieve_shen(
) )
# Fetch file paths based on year # Fetch file paths based on year
path_prefix = Path(f"parcellations/Shen/{year}")
if year == 2013: if year == 2013:
parcellation_img_path = fetch_file_via_datalad( parcellation_img_path = get(
dataset=dataset, file_path=path_prefix / f"fconn_atlas_{n_rois}_{resolution}mm.nii",
file_path=dataset.pathobj dataset_path=get_dataset_path(),
/ "parcellations" tag=JUNIFER_DATA_VERSION,
/ "Shen"
/ "2013"
/ f"fconn_atlas_{n_rois}_{resolution}mm.nii",
) )
parcellation_label_path = fetch_file_via_datalad( parcellation_label_path = get(
dataset=dataset, file_path=path_prefix / f"Group_seg{n_rois}_BAindexing_setA.txt",
file_path=dataset.pathobj dataset_path=get_dataset_path(),
/ "parcellations" tag=JUNIFER_DATA_VERSION,
/ "Shen"
/ "2013"
/ f"Group_seg{n_rois}_BAindexing_setA.txt",
) )
labels = ( labels = (
pd.read_csv( pd.read_csv(
@ -1123,23 +1073,18 @@ def _retrieve_shen(
.to_list() .to_list()
) )
elif year == 2015: elif year == 2015:
parcellation_img_path = fetch_file_via_datalad( parcellation_img_path = get(
dataset=dataset, file_path=path_prefix
file_path=dataset.pathobj
/ "parcellations"
/ "Shen"
/ "2015"
/ f"shen_{resolution}mm_268_parcellation.nii.gz", / f"shen_{resolution}mm_268_parcellation.nii.gz",
dataset_path=get_dataset_path(),
tag=JUNIFER_DATA_VERSION,
) )
labels = list(range(1, 269)) labels = list(range(1, 269))
elif year == 2019: elif year == 2019:
parcellation_img_path = fetch_file_via_datalad( parcellation_img_path = get(
dataset=dataset, file_path=path_prefix / "Shen_1mm_368_parcellation.nii.gz",
file_path=dataset.pathobj dataset_path=get_dataset_path(),
/ "parcellations" tag=JUNIFER_DATA_VERSION,
/ "Shen"
/ "2019"
/ "Shen_1mm_368_parcellation.nii.gz",
) )
labels = list(range(1, 369)) labels = list(range(1, 369))
@ -1147,7 +1092,6 @@ def _retrieve_shen(
def _retrieve_yan( def _retrieve_yan(
dataset: "Dataset",
resolution: Optional[float] = None, resolution: Optional[float] = None,
n_rois: Optional[int] = None, n_rois: Optional[int] = None,
yeo_networks: Optional[int] = None, yeo_networks: Optional[int] = None,
@ -1157,8 +1101,6 @@ def _retrieve_yan(
Parameters Parameters
---------- ----------
dataset : datalad.api.Dataset
The datalad dataset to fetch parcellation from.
resolution : float, optional resolution : float, optional
The desired resolution of the parcellation to load. If it is not The desired resolution of the parcellation to load. If it is not
available, the closest resolution will be loaded. Preferably, use a available, the closest resolution will be loaded. Preferably, use a
@ -1214,6 +1156,7 @@ def _retrieve_yan(
) )
# Fetch file paths based on networks # Fetch file paths based on networks
pre_path_prefix = Path("parcellations/Yan2023")
if yeo_networks: if yeo_networks:
# Check yeo_networks value # Check yeo_networks value
_valid_yeo_networks = [7, 17] _valid_yeo_networks = [7, 17]
@ -1223,24 +1166,21 @@ def _retrieve_yan(
f"one of the following: {_valid_yeo_networks}" f"one of the following: {_valid_yeo_networks}"
) )
parcellation_img_path = fetch_file_via_datalad( path_prefix = pre_path_prefix / "Yeo2011"
dataset=dataset, parcellation_img_path = get(
file_path=dataset.pathobj file_path=path_prefix
/ "parcellations"
/ "Yan2023"
/ "Yeo2011"
/ ( / (
f"{n_rois}Parcels_Yeo2011_{yeo_networks}Networks_FSLMNI152_" f"{n_rois}Parcels_Yeo2011_{yeo_networks}Networks_FSLMNI152_"
f"{resolution}mm.nii.gz" f"{resolution}mm.nii.gz"
), ),
dataset_path=get_dataset_path(),
tag=JUNIFER_DATA_VERSION,
) )
parcellation_label_path = fetch_file_via_datalad( parcellation_label_path = get(
dataset=dataset, file_path=path_prefix
file_path=dataset.pathobj
/ "parcellations"
/ "Yan2023"
/ "Yeo2011"
/ f"{n_rois}Parcels_Yeo2011_{yeo_networks}Networks_LUT.txt", / f"{n_rois}Parcels_Yeo2011_{yeo_networks}Networks_LUT.txt",
dataset_path=get_dataset_path(),
tag=JUNIFER_DATA_VERSION,
) )
elif kong_networks: elif kong_networks:
# Check kong_networks value # Check kong_networks value
@ -1251,24 +1191,21 @@ def _retrieve_yan(
f"one of the following: {_valid_kong_networks}" f"one of the following: {_valid_kong_networks}"
) )
parcellation_img_path = fetch_file_via_datalad( path_prefix = pre_path_prefix / "Kong2022"
dataset=dataset, parcellation_img_path = get(
file_path=dataset.pathobj file_path=path_prefix
/ "parcellations"
/ "Yan2023"
/ "Kong2022"
/ ( / (
f"{n_rois}Parcels_Kong2022_{kong_networks}Networks_FSLMNI152_" f"{n_rois}Parcels_Kong2022_{kong_networks}Networks_FSLMNI152_"
f"{resolution}mm.nii.gz" f"{resolution}mm.nii.gz"
), ),
dataset_path=get_dataset_path(),
tag=JUNIFER_DATA_VERSION,
) )
parcellation_label_path = fetch_file_via_datalad( parcellation_label_path = get(
dataset=dataset, file_path=path_prefix
file_path=dataset.pathobj
/ "parcellations"
/ "Yan2023"
/ "Kong2022"
/ f"{n_rois}Parcels_Kong2022_{kong_networks}Networks_LUT.txt", / f"{n_rois}Parcels_Kong2022_{kong_networks}Networks_LUT.txt",
dataset_path=get_dataset_path(),
tag=JUNIFER_DATA_VERSION,
) )
# Load label file # Load label file
@ -1280,7 +1217,6 @@ def _retrieve_yan(
def _retrieve_brainnetome( def _retrieve_brainnetome(
dataset: "Dataset",
resolution: Optional[float] = None, resolution: Optional[float] = None,
threshold: Optional[int] = None, threshold: Optional[int] = None,
) -> tuple[Path, list[str]]: ) -> tuple[Path, list[str]]:
@ -1288,8 +1224,6 @@ def _retrieve_brainnetome(
Parameters Parameters
---------- ----------
dataset : datalad.api.Dataset
The datalad dataset to fetch parcellation from.
resolution : {1.0, 1.25, 2.0}, optional resolution : {1.0, 1.25, 2.0}, optional
The desired resolution of the parcellation to load. If it is not The desired resolution of the parcellation to load. If it is not
available, the closest resolution will be loaded. Preferably, use a available, the closest resolution will be loaded. Preferably, use a
@ -1332,12 +1266,13 @@ def _retrieve_brainnetome(
resolution = int(resolution) resolution = int(resolution)
# Fetch file path # Fetch file path
parcellation_img_path = fetch_file_via_datalad( parcellation_img_path = get(
dataset=dataset, file_path=Path(
file_path=dataset.pathobj "parcellations/Brainnetome/"
/ "parcellations" f"BNA-maxprob-thr{threshold}-{resolution}mm.nii.gz"
/ "Brainnetome" ),
/ f"BNA-maxprob-thr{threshold}-{resolution}mm.nii.gz", dataset_path=get_dataset_path(),
tag=JUNIFER_DATA_VERSION,
) )
# Load labels # Load labels

View file

@ -24,7 +24,6 @@ from junifer.data.parcellations._parcellations import (
_retrieve_tian, _retrieve_tian,
_retrieve_yan, _retrieve_yan,
) )
from junifer.data.utils import check_dataset
from junifer.datareader import DefaultDataReader from junifer.datareader import DefaultDataReader
from junifer.pipeline.utils import _check_ants from junifer.pipeline.utils import _check_ants
from junifer.testing.datagrabbers import ( from junifer.testing.datagrabbers import (
@ -335,7 +334,6 @@ def test_retrieve_schaefer_incorrect_n_rois() -> None:
"""Test retrieve Schaefer with incorrect ROIs.""" """Test retrieve Schaefer with incorrect ROIs."""
with pytest.raises(ValueError, match=r"The parameter `n_rois`"): with pytest.raises(ValueError, match=r"The parameter `n_rois`"):
_retrieve_schaefer( _retrieve_schaefer(
dataset=check_dataset(),
resolution=1, resolution=1,
n_rois=101, n_rois=101,
yeo_networks=7, yeo_networks=7,
@ -346,7 +344,6 @@ def test_retrieve_schaefer_incorrect_yeo_networks() -> None:
"""Test retrieve Schaefer with incorrect Yeo networks.""" """Test retrieve Schaefer with incorrect Yeo networks."""
with pytest.raises(ValueError, match=r"The parameter `yeo_networks`"): with pytest.raises(ValueError, match=r"The parameter `yeo_networks`"):
_retrieve_schaefer( _retrieve_schaefer(
dataset=check_dataset(),
resolution=1, resolution=1,
n_rois=100, n_rois=100,
yeo_networks=8, yeo_networks=8,
@ -384,7 +381,7 @@ def test_suit(space_key: str, space: str) -> None:
def test_retrieve_suit_incorrect_space() -> None: def test_retrieve_suit_incorrect_space() -> None:
"""Test retrieve SUIT with incorrect space.""" """Test retrieve SUIT with incorrect space."""
with pytest.raises(ValueError, match=r"The parameter `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( @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: def test_retrieve_tian_incorrect_space() -> None:
"""Test retrieve tian with incorrect space.""" """Test retrieve tian with incorrect space."""
with pytest.raises(ValueError, match=r"The parameter `space`"): with pytest.raises(ValueError, match=r"The parameter `space`"):
_retrieve_tian( _retrieve_tian(resolution=1, scale=1, space="wrong")
dataset=check_dataset(), resolution=1, scale=1, space="wrong"
)
with pytest.raises(ValueError, match=r"MNI152NLin6Asym"): with pytest.raises(ValueError, match=r"MNI152NLin6Asym"):
_retrieve_tian( _retrieve_tian(
dataset=check_dataset(),
resolution=1, resolution=1,
scale=1, scale=1,
magneticfield="7T", magneticfield="7T",
@ -530,7 +524,6 @@ def test_retrieve_tian_incorrect_magneticfield() -> None:
"""Test retrieve tian with incorrect magneticfield.""" """Test retrieve tian with incorrect magneticfield."""
with pytest.raises(ValueError, match=r"The parameter `magneticfield`"): with pytest.raises(ValueError, match=r"The parameter `magneticfield`"):
_retrieve_tian( _retrieve_tian(
dataset=check_dataset(),
resolution=1, resolution=1,
scale=1, scale=1,
magneticfield="wrong", magneticfield="wrong",
@ -541,7 +534,6 @@ def test_retrieve_tian_incorrect_scale(tmp_path: Path) -> None:
"""Test retrieve tian with incorrect scale.""" """Test retrieve tian with incorrect scale."""
with pytest.raises(ValueError, match=r"The parameter `scale`"): with pytest.raises(ValueError, match=r"The parameter `scale`"):
_retrieve_tian( _retrieve_tian(
dataset=check_dataset(),
resolution=1, resolution=1,
scale=5, scale=5,
space="MNI152NLin6Asym", space="MNI152NLin6Asym",
@ -577,7 +569,6 @@ def test_retrieve_aicha_incorrect_version() -> None:
"""Test retrieve AICHA with incorrect version.""" """Test retrieve AICHA with incorrect version."""
with pytest.raises(ValueError, match="The parameter `version`"): with pytest.raises(ValueError, match="The parameter `version`"):
_retrieve_aicha( _retrieve_aicha(
dataset=check_dataset(),
version=100, version=100,
) )
@ -639,7 +630,6 @@ def test_retrieve_shen_incorrect_year() -> None:
"""Test retrieve Shen with incorrect year.""" """Test retrieve Shen with incorrect year."""
with pytest.raises(ValueError, match="The parameter `year`"): with pytest.raises(ValueError, match="The parameter `year`"):
_retrieve_shen( _retrieve_shen(
dataset=check_dataset(),
year=1969, year=1969,
) )
@ -648,7 +638,6 @@ def test_retrieve_shen_incorrect_n_rois() -> None:
"""Test retrieve Shen with incorrect ROIs.""" """Test retrieve Shen with incorrect ROIs."""
with pytest.raises(ValueError, match="The parameter `n_rois`"): with pytest.raises(ValueError, match="The parameter `n_rois`"):
_retrieve_shen( _retrieve_shen(
dataset=check_dataset(),
year=2015, year=2015,
n_rois=10, n_rois=10,
) )
@ -691,7 +680,6 @@ def test_retrieve_shen_incorrect_param_combo(
""" """
with pytest.raises(ValueError, match="The parameter combination"): with pytest.raises(ValueError, match="The parameter combination"):
_retrieve_shen( _retrieve_shen(
dataset=check_dataset(),
resolution=resolution, resolution=resolution,
year=year, year=year,
n_rois=n_rois, 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`" ValueError, match="Either one of `yeo_networks` or `kong_networks`"
): ):
_retrieve_yan( _retrieve_yan(
dataset=check_dataset(),
n_rois=31418, n_rois=31418,
yeo_networks=100, yeo_networks=100,
kong_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`" ValueError, match="Either one of `yeo_networks` or `kong_networks`"
): ):
_retrieve_yan( _retrieve_yan(
dataset=check_dataset(),
n_rois=31418, n_rois=31418,
yeo_networks=None, yeo_networks=None,
kong_networks=None, kong_networks=None,
@ -840,7 +826,6 @@ def test_retrieve_yan_incorrect_n_rois() -> None:
"""Test retrieve Yan with incorrect ROIs.""" """Test retrieve Yan with incorrect ROIs."""
with pytest.raises(ValueError, match="The parameter `n_rois`"): with pytest.raises(ValueError, match="The parameter `n_rois`"):
_retrieve_yan( _retrieve_yan(
dataset=check_dataset(),
n_rois=31418, n_rois=31418,
yeo_networks=7, yeo_networks=7,
) )
@ -850,7 +835,6 @@ def test_retrieve_yan_incorrect_yeo_networks() -> None:
"""Test retrieve Yan with incorrect Yeo networks.""" """Test retrieve Yan with incorrect Yeo networks."""
with pytest.raises(ValueError, match="The parameter `yeo_networks`"): with pytest.raises(ValueError, match="The parameter `yeo_networks`"):
_retrieve_yan( _retrieve_yan(
dataset=check_dataset(),
n_rois=100, n_rois=100,
yeo_networks=27, yeo_networks=27,
) )
@ -860,7 +844,6 @@ def test_retrieve_yan_incorrect_kong_networks() -> None:
"""Test retrieve Yan with incorrect Kong networks.""" """Test retrieve Yan with incorrect Kong networks."""
with pytest.raises(ValueError, match="The parameter `kong_networks`"): with pytest.raises(ValueError, match="The parameter `kong_networks`"):
_retrieve_yan( _retrieve_yan(
dataset=check_dataset(),
n_rois=100, n_rois=100,
kong_networks=27, kong_networks=27,
) )
@ -922,7 +905,6 @@ def test_retrieve_brainnetome_incorrect_threshold() -> None:
"""Test retrieve Brainnetome with incorrect threshold.""" """Test retrieve Brainnetome with incorrect threshold."""
with pytest.raises(ValueError, match="The parameter `threshold`"): with pytest.raises(ValueError, match="The parameter `threshold`"):
_retrieve_brainnetome( _retrieve_brainnetome(
dataset=check_dataset(),
threshold=100, threshold=100,
) )

View file

@ -8,10 +8,11 @@ from typing import Any, Optional, Union
import nibabel as nib import nibabel as nib
import numpy as np import numpy as np
from junifer_data import get
from templateflow import api as tflow from templateflow import api as tflow
from ..utils import logger, raise_error 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"] __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. The path to the transformation file.
""" """
# Get dataset
dataset = check_dataset()
# Set file path to retrieve # Set file path to retrieve
xfm_file_path = ( xfm_file_path = Path(f"xfms/{src}_to_{dst}/{src}_to_{dst}_Composite.h5")
dataset.pathobj
/ "xfms"
/ f"{src}_to_{dst}"
/ f"{src}_to_{dst}_Composite.h5"
)
# Retrieve file # 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( def get_template(

View file

@ -8,21 +8,23 @@ from collections.abc import MutableMapping
from pathlib import Path from pathlib import Path
from typing import Optional, Union from typing import Optional, Union
import datalad.api as dl
import numpy as np import numpy as np
from datalad.support.exceptions import IncompleteResultsError
from ..utils import config, logger, raise_error from ..utils import config, logger, raise_error
__all__ = [ __all__ = [
"check_dataset", "JUNIFER_DATA_VERSION",
"closest_resolution", "closest_resolution",
"fetch_file_via_datalad", "get_dataset_path",
"get_native_warper", "get_native_warper",
] ]
# junifer-data version constant
JUNIFER_DATA_VERSION = "1"
def closest_resolution( def closest_resolution(
resolution: Optional[Union[float, int]], resolution: Optional[Union[float, int]],
valid_resolution: Union[list[float], list[int], np.ndarray], valid_resolution: Union[list[float], list[int], np.ndarray],
@ -124,94 +126,17 @@ def get_native_warper(
return possible_warpers[0] return possible_warpers[0]
def check_dataset() -> dl.Dataset: def get_dataset_path() -> Optional[Path]:
"""Get or install junifer-data dataset. """Get junifer-data dataset path.
Returns Returns
------- -------
datalad.api.Dataset pathlib.Path or None
The junifer-data dataset. Path to the dataset or None.
Raises
------
RuntimeError
If there is a problem cloning the dataset.
""" """
# Check config and set default if not passed return (
data_dir = config.get("data.location") Path(config.get("data.location"))
if data_dir is not None: if config.get("data.location") is not None
data_dir = Path(data_dir) else None
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,
)

View file

@ -63,6 +63,6 @@ MarkerInOutMappings = MutableMapping[str, MutableMapping[str, str]]
DataGrabberPatterns = dict[ DataGrabberPatterns = dict[
str, Union[dict[str, str], Sequence[dict[str, str]]] 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, ...]] Element = Union[str, tuple[str, ...]]
Elements = Sequence[Element] Elements = Sequence[Element]

View file

@ -52,6 +52,7 @@ dependencies = [
"lazy_loader==0.4", "lazy_loader==0.4",
"importlib_metadata; python_version<'3.9'", "importlib_metadata; python_version<'3.9'",
"looseversion==1.3.0; python_version>='3.12'", "looseversion==1.3.0; python_version>='3.12'",
"junifer_data==1.1.0",
] ]
dynamic = ["version"] dynamic = ["version"]
@ -212,6 +213,7 @@ known-third-party = [
"brainprint", "brainprint",
"lapy", "lapy",
"pytest", "pytest",
"junifer_data",
] ]
[tool.ruff.lint.mccabe] [tool.ruff.lint.mccabe]