[ENH]: Move data downloading/handling to junifer-data package. #363
10 changed files with 148 additions and 318 deletions
|
|
@ -48,6 +48,7 @@ def test_get_dependency_information_short() -> None:
|
|||
"lapy",
|
||||
"lazy_loader",
|
||||
"looseversion",
|
||||
"junifer_data",
|
||||
]
|
||||
|
||||
if sys.version_info < (3, 11):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
Loading…
Reference in a new issue