[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",
"lazy_loader",
"looseversion",
"junifer_data",
]
if sys.version_info < (3, 11):

View file

@ -4,16 +4,18 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# 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,

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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