diff --git a/junifer/data/__init__.py b/junifer/data/__init__.py index 4e8ab386c..79e23f86f 100644 --- a/junifer/data/__init__.py +++ b/junifer/data/__init__.py @@ -1 +1,7 @@ -from .atlases import list_atlases, register_atlas, load_atlas \ No newline at end of file +"""Provide imports for data sub-package.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from .atlases import list_atlases, register_atlas, load_atlas diff --git a/junifer/data/atlases.py b/junifer/data/atlases.py index d429bbf82..7fe70a46a 100644 --- a/junifer/data/atlases.py +++ b/junifer/data/atlases.py @@ -5,20 +5,26 @@ # Synchon Mandal # License: AGPL -from pathlib import Path import io -import tempfile -import requests import shutil +import tempfile import zipfile -import numpy as np -import pandas as pd +from pathlib import Path +from typing import TYPE_CHECKING, List, Optional, Tuple, Union import nibabel as nib +import numpy as np +import pandas as pd +import requests from nilearn import datasets from ..utils.logging import logger, raise_error + +if TYPE_CHECKING: + from nibabel import Nifti1Image + + """ A dictionary containing all supported atlases and their respective valid parameters. @@ -30,97 +36,104 @@ Optional keys: * 'valid_resolutions': a list of valid resolutions for the atlas (e.g. [1, 2]) """ +# TODO: have separate dictionary for built-in _available_atlases = { - 'SUITxSUIT': { - 'family': 'SUIT', - 'space': 'SUIT' - }, - 'SUITxMNI': { - 'family': 'SUIT', - 'space': 'MNI' - }, - + "SUITxSUIT": {"family": "SUIT", "space": "SUIT"}, + "SUITxMNI": {"family": "SUIT", "space": "MNI"}, } - +# Add Schaefer atlas info for n_rois in range(100, 1001, 100): for t_net in [7, 17]: - t_name = f'Schaefer{n_rois}x{t_net}' + t_name = f"Schaefer{n_rois}x{t_net}" _available_atlases[t_name] = { - 'family': 'Schaefer', - 'n_rois': n_rois, - 'yeo_networks': t_net, + "family": "Schaefer", + "n_rois": n_rois, + "yeo_networks": t_net, } - +# Add Tian atlas info for scale in range(1, 5): - t_name = f'TianxS{scale}x7TxMNI6thgeneration' + t_name = f"TianxS{scale}x7TxMNI6thgeneration" _available_atlases[t_name] = { - 'family': 'Tian', - 'scale': scale, - 'magneticfield': '7T', - 'space': 'MNI6thgeneration' + "family": "Tian", + "scale": scale, + "magneticfield": "7T", + "space": "MNI6thgeneration", } - t_name = f'TianxS{scale}x3TxMNI6thgeneration' + t_name = f"TianxS{scale}x3TxMNI6thgeneration" _available_atlases[t_name] = { - 'family': 'Tian', - 'scale': scale, - 'magneticfield': '3T', - 'space': 'MNI6thgeneration' + "family": "Tian", + "scale": scale, + "magneticfield": "3T", + "space": "MNI6thgeneration", } - t_name = f'TianxS{scale}x3TxMNInonlinear2009cAsym' + t_name = f"TianxS{scale}x3TxMNInonlinear2009cAsym" _available_atlases[t_name] = { - 'family': 'Tian', - 'scale': scale, - 'magneticfield': '3T', - 'space': 'MNInonlinear2009cAsym' + "family": "Tian", + "scale": scale, + "magneticfield": "3T", + "space": "MNInonlinear2009cAsym", } -def register_atlas(name, atlas_path, atl_labels, overwrite=False): +def register_atlas( + name: str, + atlas_path: Union[str, Path], + atl_labels: List[str], + overwrite: bool = False, +) -> None: """Register a custom user atlas. Parameters ---------- name : str The name of the atlas. - atlas_path : str + atlas_path : str or pathlib.Path The path to the atlas file. - atl_labels : list(str) + atl_labels : list of str The list of labels for the atlas. - overwrite : bool - If True, overwrite an existing atlas with the same name. Defaults to - False. + overwrite : bool, optional + If True, overwrite an existing atlas with the same name. + Does not apply to built-in atlases (default False). Raises ------ ValueError If the atlas name is already registered and overwrite is set to False or if the atlas name is a built-in atlas. + """ + # Check for attempt of overwriting built-in atlases if name in _available_atlases: if overwrite is True: - logger.info(f'Overwritting {name} atlas') - if _available_atlases[name]['family'] != 'CustomUserAtlas': + logger.info(f"Overwriting {name} atlas") + if _available_atlases[name]["family"] != "CustomUserAtlas": raise_error( - f'Cannot overwrite {name} atlas. It is a built-in atlas.') + f"Cannot overwrite {name} atlas. It is a built-in atlas." + ) else: raise_error( - f'Atlas {name} already registered. Set `overwrite=True` to ' - 'update its value.') + f"Atlas {name} already registered. Set `overwrite=True` to " + "update its value." + ) + # Convert str to Path if not isinstance(atlas_path, Path): atlas_path = Path(atlas_path) + # Add user atlas info _available_atlases[name] = { - 'path': atlas_path, 'labels': atl_labels, - 'family': 'CustomUserAtlas'} + "path": str(atlas_path.absolute()), + "labels": atl_labels, + "family": "CustomUserAtlas", + } -def list_atlases(): - """ - List all the available atlases. +def list_atlases() -> List[str]: + """List all the available atlases. Returns ------- - out : list(str) or dict - A list or dict with all available atlases. + list of str + A list with all available atlases. + """ return sorted(_available_atlases.keys()) @@ -133,81 +146,87 @@ def list_atlases(): # return resolution -def load_atlas(name, atlas_dir=None, resolution=None, path_only=False): +# TODO: keyword arguments are not passed, check +def load_atlas( + name: str, + atlas_dir: Union[str, Path, None] = None, + resolution: Optional[int] = None, + path_only: bool = False, +) -> Tuple[Optional["Nifti1Image"], List[str], Path]: """Load a brain atlas (including a label file). - If it is built-in atlas and file is not present in the `atlas_dir` + If it is a built-in atlas and file is not present in the `atlas_dir` directory, it will be downloaded. Parameters ---------- name : str - The name of the atlas. - Check valid options by calling `list_atlases`. - atlas_dir: path - Path where the atlas files are stored. - Defaults to: $HOME/junifer/data/atlas - resolution : int - The (desired) resolution of the atlas to load. If its not available, + The name of the atlas. Check valid options by calling `list_atlases`. + atlas_dir : str or pathlib.Path, optional + Path where the atlas files are stored. The default location is + "$HOME/junifer/data/atlas" (default None). + resolution : int, optional + The desired resolution of the atlas to load. If it is not available, the closest resolution will be loaded. Preferably, use a resolution - higher than the desired one. Defaults to None (load the highest one). - path_only : bool - If True, the atlas image will not be loaded. + higher than the desired one. By default, will load the highest one + (default None). + path_only : bool, optional + If True, the atlas image will not be loaded (default False). - Parameters (optional, atlas dependent) - -------------------------------------- - Use to specify atlas specific keyword arguments. . + Extra Parameters + ---------------- + Use to specify atlas specific keyword arguments. - Schaefer : - n_rois (required) : int - Granularity of atlas to be used. Valid values: between 100 and 1000 - (included) in steps of 100. - yeo_network (optional) : int - Number of yeo networks to use. Valid values: 7, 17. Defaults to 7. - - Tian : - scale (required) : int - Scale of atlas between 1 and 4 (defines granularity) - space (optional) : str - Space of atlas can be either 'MNI6thgeneration' or - 'MNInonlinear2009cAsym' (for some cases). - Defaults to 'MNI6thgeneration'. (For more information see - https://github.com/yetianmed/subcortex) - magneticfield (optional) : str - Options are 3T and 7T, defaults to 3T. - - SUIT : - space (optional) : str - Space of atlas can be either 'MNI' or 'SUIT' (for more information - see http://www.diedrichsenlab.org/imaging/suit.htm). Defaults to - 'MNI'. + - Schaefer : + n_rois : {100, 200, 300, 400, 500, 600, 700, 800, 900, 1000} + Granularity of atlas to be used. + yeo_network : {7, 17}, optional + Number of yeo networks to use (default 7). + - Tian : + scale : {1, 2, 3, 4} + Scale of atlas (defines granularity). + space : {"MNI6thgeneration", "MNInonlinear2009cAsym"}, optional + Space of atlas (default "MNI6thgeneration"). (For more information + see https://github.com/yetianmed/subcortex) + magneticfield : {"3T", "7T"}, optional + Magnetic field (default "3T"). + - SUIT : + space : {"MNI", "SUIT"}, optional + Space of atlas (default "MNI"). (For more information + see http://www.diedrichsenlab.org/imaging/suit.htm). Returns ------- - atlas_img : niimg-like object or None + niimg-like object or None Loaded atlas image. - atlas_labels : List of str + list of str Atlas labels. - atlas_fname : Path + pathlib.Path File path to the atlas image. - """ - if name not in _available_atlases: - raise_error(f'Atlas {name} not found. ' - f'Valid options are: {list_atlases()}') - atlas_definition = _available_atlases[name].copy() - t_family = atlas_definition.pop('family') - if t_family == 'CustomUserAtlas': - atlas_fname = atlas_definition['path'] - atlas_labels = atlas_definition['labels'] + """ + # Invalid atlas name + if name not in _available_atlases: + raise_error( + f"Atlas {name} not found. Valid options are: {list_atlases()}" + ) + + atlas_definition = _available_atlases[name].copy() + t_family = atlas_definition.pop("family") + + if t_family == "CustomUserAtlas": + atlas_fname = Path(atlas_definition["path"]) + atlas_labels = atlas_definition["labels"] else: # retrieve atlases by passing arguments on to _retrieve_atlas() atlas_fname, atlas_labels = _retrieve_atlas( - t_family, resolution=resolution, atlas_dir=atlas_dir, - **atlas_definition) + family=t_family, + atlas_dir=atlas_dir, + resolution=resolution, + **atlas_definition, + ) - logger.info( - f'Loading atlas {atlas_fname.as_posix()}') # type: ignore + logger.info(f"Loading atlas {str(atlas_fname.absolute())}") atlas_img = None if path_only is False: @@ -216,89 +235,121 @@ def load_atlas(name, atlas_dir=None, resolution=None, path_only=False): return atlas_img, atlas_labels, atlas_fname -def _retrieve_atlas(family, atlas_dir=None, resolution=None, **kwargs): +def _retrieve_atlas( + family: str, + atlas_dir: Union[str, Path, None] = None, + resolution: Optional[int] = None, + **kwargs, +) -> Tuple[Path, List[str]]: """Retrieve a brain atlas object from nilearn or a specified online source. Only returns one atlas per call. Call function multiple times for different parameter specifications. Only retrieves atlas if it is not yet in atlas_dir. - Parameters (required) - --------------------- + Parameters + ---------- family : str - Specify by name of atlas family, e.g. 'Schaefer'. - atlas_dir: str or Path - Path to where to store the retrieved atlas file. - Defaults to: $HOME/junifer/data/atlas - resolution : int - The (desired) resolution of the atlas to load. If its not available, + The name of the atlas family, e.g. 'Schaefer'. + atlas_dir : str or pathlib.Path, optional + Path where the retrieved atlas file is stored. The default location is + "$HOME/junifer/data/atlas" (default None). + resolution : int, optional + The desired resolution of the atlas to load. If it is not available, the closest resolution will be loaded. Preferably, use a resolution - higher than the desired one. Defaults to None (load the highest one). + higher than the desired one. By default, will load the highest one + (default None). - Parameters (optional, atlas dependent) - -------------------------------------- - Use to specify atlas specific keyword arguments + Extra Parameters + ---------------- + **kwargs + Use to specify atlas specific keyword arguments. - Schaefer : - n_rois (required) : int - Granularity of atlas to be used. Valid values: between 100 and 1000 - (included) in steps of 100. - yeo_network (optional) : int - Number of yeo networks to use [7 or 17]. Defaults to 7. - Tian : - scale (required) : int - Scale of atlas between 1 and 4 (defines granularity) - space (optional) : str - Space of atlas can be either 'MNI6thgeneration' or - 'MNInonlinear2009cAsym' (for some cases). - Defaults to 'MNI6thgeneration'. (For more information see - https://github.com/yetianmed/subcortex) - magneticfield (optional) : str - Options are 3T and 7T, defaults to 3T. - SUIT : - space (optional) : str - Space of atlas can be either 'MNI' or 'SUIT' (for more information - see http://www.diedrichsenlab.org/imaging/suit.htm). Defaults to - 'MNI'. + - Schaefer : + n_rois : {100, 200, 300, 400, 500, 600, 700, 800, 900, 1000} + Granularity of atlas to be used. + yeo_network : {7, 17}, optional + Number of yeo networks to use (default 7). + - Tian : + scale : {1, 2, 3, 4} + Scale of atlas (defines granularity). + space : {"MNI6thgeneration", "MNInonlinear2009cAsym"}, optional + Space of atlas (default "MNI6thgeneration"). (For more + information see https://github.com/yetianmed/subcortex) + magneticfield : {"3T", "7T"}, optional + Magnetic field (default "3T"). + - SUIT : + space : {"MNI", "SUIT"}, optional + Space of atlas (default "MNI"). (For more information + see http://www.diedrichsenlab.org/imaging/suit.htm). Returns ------- - atlas_fname : Path + pathlib.Path File path to the atlas image. - atlas_labels : List of str + list of str Atlas labels. + + Raises + ------ + ValueError + If the atlas name is invalid. + """ if atlas_dir is None: - atlas_dir = Path().home() / 'junifer' / 'data' / 'atlas' + atlas_dir = Path().home() / "junifer" / "data" / "atlas" + # Create default junifer data directory if not present atlas_dir.mkdir(exist_ok=True, parents=True) + # Convert str to Path elif not isinstance(atlas_dir, Path): atlas_dir = Path(atlas_dir) logger.info(f"Fetching one of {family} atlas.") - # retrieval details per atlas - if family == 'Schaefer': - atlas_fname, atl_labels = \ - _retrieve_schaefer(atlas_dir, resolution=resolution, **kwargs) - elif family == 'SUIT': - atlas_fname, atl_labels = \ - _retrieve_suit(atlas_dir, resolution=resolution, **kwargs) - elif family == 'Tian': - atlas_fname, atl_labels = \ - _retrieve_tian(atlas_dir, resolution=resolution, **kwargs) + # Retrieval details per atlas + if family == "Schaefer": + atlas_fname, atl_labels = _retrieve_schaefer( + atlas_dir=atlas_dir, resolution=resolution, **kwargs + ) + elif family == "SUIT": + atlas_fname, atl_labels = _retrieve_suit( + atlas_dir=atlas_dir, resolution=resolution, **kwargs + ) + elif family == "Tian": + atlas_fname, atl_labels = _retrieve_tian( + atlas_dir=atlas_dir, resolution=resolution, **kwargs + ) else: - raise_error( - f"The provided atlas name {family} cannot be retrieved. ") + raise_error(f"The provided atlas name {family} cannot be retrieved.") return atlas_fname, atl_labels -def _closest_resolution(resolution, valid_resolution): - closest = None +def _closest_resolution( + resolution: int, + valid_resolution: Union[List[int], np.ndarray], +) -> int: + """Find the closest resolution. + + Parameters + ---------- + resolution : int + The given resolution. + valid_resolution : list of int or np.ndarray + The array of valid resolutions. + + Returns + ------- + int + The closest valid resolution. + + """ + # Convert list of int to numpy.ndarray if not isinstance(valid_resolution, np.ndarray): valid_resolution = np.array(valid_resolution) + if resolution is None: - logger.info('Resolution set to None, using highest resolution.') + logger.info("Resolution set to None, using highest resolution.") closest = np.min(valid_resolution) elif any(x <= resolution for x in valid_resolution): # Case 1: get the highest closest resolution @@ -310,11 +361,43 @@ def _closest_resolution(resolution, valid_resolution): return closest -def _retrieve_schaefer(atlas_dir, resolution, n_rois=None, yeo_networks=7): - logger.info('Atlas parameters:') - logger.info(f'\tn_rois: {n_rois}') - logger.info(f'\tyeo_networks: {yeo_networks}') - logger.info(f'\tresolution: {resolution}') +def _retrieve_schaefer( + atlas_dir: Path, + resolution: int, + n_rois: Optional[int] = None, + yeo_networks: int = 7, +) -> Tuple[Path, List[str]]: + """Retrieve Schaefer atlas. + + Parameters + ---------- + atlas_dir : pathlib.Path + The path to the atlas data directory. + resolution : {1, 2} + The resolution of the atlas to load. + n_rois : {100, 200, 300, 400, 500, 600, 700, 800, 900, 1000}, optional + Granularity of the atlas to be used (default None). + yeo_networks : {7, 17}, optional + Number of yeo networks to use (default 7). + + Returns + ------- + pathlib.Path + File path to the atlas image. + list of str + Atlas labels. + + Raises + ------ + ValueError + If invalid value is provided for `n_rois` or `yeo_networks` or if + there is a problem fetching the atlas. + + """ + logger.info("Atlas parameters:") + logger.info(f"\tn_rois: {n_rois}") + logger.info(f"\tyeo_networks: {yeo_networks}") + logger.info(f"\tresolution: {resolution}") _valid_n_rois = [100, 200, 300, 400, 500, 600, 700, 800, 900, 1000] _valid_networks = [7, 17] @@ -322,160 +405,253 @@ def _retrieve_schaefer(atlas_dir, resolution, n_rois=None, yeo_networks=7): if n_rois not in _valid_n_rois: raise_error( - f'The parameter `n_rois` ({n_rois}) needs to be one of the ' - f'following: {_valid_n_rois}') + f"The parameter `n_rois` ({n_rois}) needs to be one of the " + f"following: {_valid_n_rois}" + ) if yeo_networks not in _valid_networks: raise_error( - f'The parameter `yeo_networks` ({yeo_networks}) needs to be one of' - f' the following: {_valid_networks}') + f"The parameter `yeo_networks` ({yeo_networks}) needs to be one " + f"of the following: {_valid_networks}" + ) resolution = _closest_resolution(resolution, _valid_resolutions) # define file names - atlas_fname = atlas_dir / 'schaefer_2018' / ( - f'Schaefer2018_{n_rois}Parcels_{yeo_networks}Networks_order_' - f'FSLMNI152_{resolution}mm.nii.gz') - atlas_lname = atlas_dir / 'schaefer_2018' / ( - f'Schaefer2018_{n_rois}Parcels_{yeo_networks}Networks_order.txt') + atlas_fname = ( + atlas_dir + / "schaefer_2018" + / ( + f"Schaefer2018_{n_rois}Parcels_{yeo_networks}Networks_order_" + f"FSLMNI152_{resolution}mm.nii.gz" + ) + ) + atlas_lname = ( + atlas_dir + / "schaefer_2018" + / (f"Schaefer2018_{n_rois}Parcels_{yeo_networks}Networks_order.txt") + ) # check existence of atlas if not (atlas_fname.exists() and atlas_lname.exists()): logger.info( - 'At least one of the atlas files is missing. ' - 'Fetching using nilearn.') + "At least one of the atlas files is missing. " + "Fetching using nilearn." + ) datasets.fetch_atlas_schaefer_2018( - n_rois=n_rois, # type: ignore + n_rois=n_rois, yeo_networks=yeo_networks, resolution_mm=resolution, - data_dir=atlas_dir.as_posix()) + data_dir=str(atlas_dir.absolute()), + ) - if not (atlas_fname.exists() and - atlas_lname.exists()): # pragma: no cover - raise_error('There was a problem fetching the atlases.') + if not ( + atlas_fname.exists() and atlas_lname.exists() + ): # pragma: no cover + raise_error("There was a problem fetching the atlases.") # Load labels labels = [ - '_'.join(x.split('_')[1:]) - for x in pd.read_csv( - atlas_lname, sep='\t', header=None).iloc[:, 1].to_list() + "_".join(x.split("_")[1:]) + for x in pd.read_csv(atlas_lname, sep="\t", header=None) + .iloc[:, 1] + .to_list() ] return atlas_fname, labels def _retrieve_tian( - atlas_dir, resolution, scale=None, space='MNI6thgeneration', - magneticfield='3T'): + atlas_dir: Path, + resolution: int, + scale: Optional[int] = None, + space: str = "MNI6thgeneration", + magneticfield: str = "3T", +) -> Tuple[Path, List[str]]: + """Retrieve Tian atlas. + + Parameters + ---------- + atlas_dir : pathlib.Path + The path to the atlas data directory. + resolution : {1, 2} + The resolution of the atlas to load. + scale : {1, 2, 3, 4}, optional + Scale of atlas (defines granularity) (default None). + space : {"MNI6thgeneration", "MNInonlinear2009cAsym"}, optional + Space of atlas (default "MNI6thgeneration"). (For more + information see https://github.com/yetianmed/subcortex) + magneticfield : {"3T", "7T"}, optional + Magnetic field (default "3T"). + + Returns + ------- + pathlib.Path + File path to the atlas image. + list of str + Atlas labels. + + Raises + ------ + ValueError + If invalid value is provided for `scale` or `magneticfield` or `space` + or if there is a problem fetching the atlas. + + """ # show atlas parameters to user - logger.info('Atlas parameters:') - logger.info(f'\tscale: {scale}') - logger.info(f'\tspace: {space}') - logger.info(f'\tmagneticfield: {magneticfield}') - logger.info(f'\tresolution: {resolution}') + logger.info("Atlas parameters:") + logger.info(f"\tscale: {scale}") + logger.info(f"\tspace: {space}") + logger.info(f"\tmagneticfield: {magneticfield}") + logger.info(f"\tresolution: {resolution}") # check validity of atlas parameters _valid_scales = [1, 2, 3, 4] - _valid_fields = ['3T', '7T'] + _valid_fields = ["3T", "7T"] if scale not in _valid_scales: raise_error( - f'The parameter `scale` ({scale}) needs to be one of the ' - f'following: {_valid_scales}') + f"The parameter `scale` ({scale}) needs to be one of the " + f"following: {_valid_scales}" + ) if magneticfield not in _valid_fields: raise_error( - f'The parameter `magneticfield` ({magneticfield}) needs to be ' - f'one of the following: {_valid_fields}') + f"The parameter `magneticfield` ({magneticfield}) needs to be " + f"one of the following: {_valid_fields}" + ) - if magneticfield == '3T': - _valid_spaces = ['MNI6thgeneration', 'MNInonlinear2009cAsym'] - if space == 'MNI6thgeneration': + if magneticfield == "3T": + _valid_spaces = ["MNI6thgeneration", "MNInonlinear2009cAsym"] + if space == "MNI6thgeneration": _valid_resolutions = [1, 2] - else: # space == 'MNInonlinear2009cAsym': + elif space == "MNInonlinear2009cAsym": _valid_resolutions = [2] - else: # magneticfield == '7T': - _valid_spaces = ['MNI6thgeneration'] + elif magneticfield == "7T": + _valid_spaces = ["MNI6thgeneration"] _valid_resolutions = [1.6] + if space not in _valid_spaces: raise_error( - f'The parameter `space` ({space}) needs to be one of ' - f'the following: {_valid_spaces}') + f"The parameter `space` ({space}) needs to be one of " + f"the following: {_valid_spaces}" + ) resolution = _closest_resolution(resolution, _valid_resolutions) # define file names - if magneticfield == '3T': + if magneticfield == "3T": atlas_fname_base_3T = ( - atlas_dir / 'Tian2020MSA_v1.1' / '3T' / 'Subcortex-Only') + atlas_dir / "Tian2020MSA_v1.1" / "3T" / "Subcortex-Only" + ) atlas_lname = atlas_fname_base_3T / ( - f'Tian_Subcortex_S{scale}_3T_label.txt') - if space == 'MNI6thgeneration': + f"Tian_Subcortex_S{scale}_3T_label.txt" + ) + if space == "MNI6thgeneration": atlas_fname = atlas_fname_base_3T / ( - f'Tian_Subcortex_S{scale}_{magneticfield}.nii.gz') + f"Tian_Subcortex_S{scale}_{magneticfield}.nii.gz" + ) if resolution == 1: - atlas_fname = atlas_fname_base_3T / \ - f'Tian_Subcortex_S{scale}_{magneticfield}_1mm.nii.gz' - else: # space == 'MNInonlinear2009cAsym': - space = '2009cAsym' + atlas_fname = ( + atlas_fname_base_3T + / f"Tian_Subcortex_S{scale}_{magneticfield}_1mm.nii.gz" + ) + elif space == "MNInonlinear2009cAsym": + space = "2009cAsym" atlas_fname = atlas_fname_base_3T / ( - f'Tian_Subcortex_S{scale}_{magneticfield}_{space}.nii.gz') - else: # magneticfield == '7T': - atlas_fname_base_7T = ( - atlas_dir / 'Tian2020MSA_v1.1' / '7T') - atlas_fname = atlas_dir / 'Tian2020MSA_v1.1' / f'{magneticfield}' / ( - f'Tian_Subcortex_S{scale}_{magneticfield}.nii.gz') + f"Tian_Subcortex_S{scale}_{magneticfield}_{space}.nii.gz" + ) + elif magneticfield == "7T": + atlas_fname_base_7T = atlas_dir / "Tian2020MSA_v1.1" / "7T" + atlas_fname_base_7T.mkdir(exist_ok=True, parents=True) + atlas_fname = ( + atlas_dir + / "Tian2020MSA_v1.1" + / f"{magneticfield}" + / (f"Tian_Subcortex_S{scale}_{magneticfield}.nii.gz") + ) # define 7T labels (b/c currently no labels file available for 7T) scale7Trois = {1: 16, 2: 34, 3: 54, 4: 62} labels = [ - ('parcel_' + str(x)) for x in np.arange(1, scale7Trois[scale] + 1)] + ("parcel_" + str(x)) for x in np.arange(1, scale7Trois[scale] + 1) + ] atlas_lname = atlas_fname_base_7T / ( - f'Tian_Subcortex_S{scale}_7T_labelnumbering.txt') - with open(atlas_lname, 'w') as filehandle: + f"Tian_Subcortex_S{scale}_7T_labelnumbering.txt" + ) + with open(atlas_lname, "w") as filehandle: for listitem in labels: - filehandle.write('%s\n' % listitem) + filehandle.write("%s\n" % listitem) logger.info( - 'Currently there are no labels provided for the 7T Tian atlas. A ' - 'simple numbering scheme for distinction was therefore used.') + "Currently there are no labels provided for the 7T Tian atlas. " + "A simple numbering scheme for distinction was therefore used." + ) # check existence of atlas if not (atlas_fname.exists() and atlas_lname.exists()): - logger.info( - 'At least one of the atlas files is missing. ' - 'Fetching.') + logger.info("At least one of the atlas files is missing, fetching.") - url_basis = \ - 'https://www.nitrc.org/frs/download.php/12012/Tian2020MSA_v1.1.zip' + url_basis = ( + "https://www.nitrc.org/frs/download.php/12012/Tian2020MSA_v1.1.zip" + ) - logger.info(f'Downloading TIAN from {url_basis}') + logger.info(f"Downloading TIAN from {url_basis}") with tempfile.TemporaryDirectory() as tmpdir: atlas_download = requests.get(url_basis) - atlas_zip_fname = Path(tmpdir) / 'Tian2020MSA_v1.1.zip' - with open(atlas_zip_fname, 'wb') as f: + atlas_zip_fname = Path(tmpdir) / "Tian2020MSA_v1.1.zip" + with open(atlas_zip_fname, "wb") as f: f.write(atlas_download.content) - with zipfile.ZipFile(atlas_zip_fname, 'r') as zip_ref: + with zipfile.ZipFile(atlas_zip_fname, "r") as zip_ref: zip_ref.extractall(atlas_dir.as_posix()) # clean after unzipping - if (atlas_dir / '__MACOSX').exists(): - shutil.rmtree((atlas_dir / '__MACOSX').as_posix()) + if (atlas_dir / "__MACOSX").exists(): + shutil.rmtree((atlas_dir / "__MACOSX").as_posix()) labels = pd.read_csv(atlas_lname, sep=" ", header=None)[0].to_list() if not (atlas_fname.exists() and atlas_lname.exists()): - raise_error('There was a problem fetching the atlases.') + raise_error("There was a problem fetching the atlases.") labels = pd.read_csv(atlas_lname, sep=" ", header=None)[0].to_list() return atlas_fname, labels -def _retrieve_suit(atlas_path, resolution, space='MNI'): - logger.info('Atlas parameters:') - logger.info(f'\tspace: {space}') +def _retrieve_suit( + atlas_dir: Path, resolution: int, space: str = "MNI" +) -> Tuple[Path, List[str]]: + """Retrieve SUIT atlas. - _valid_spaces = ['MNI', 'SUIT'] + Parameters + ---------- + atlas_dir : pathlib.Path + The path to the atlas data directory. + resolution : {1, 2} + The resolution of the atlas to load. + space : {"MNI", "SUIT"}, optional + Space of atlas (default "MNI"). (For more information + see http://www.diedrichsenlab.org/imaging/suit.htm). + + Returns + ------ + pathlib.Path + File path to the atlas image. + list of str + Atlas labels. + + Raises + ------ + ValueError + If invalid value is provided for `space` or if there is a problem + fetching the atlas. + + """ + logger.info("Atlas parameters:") + logger.info(f"\tspace: {space}") + + _valid_spaces = ["MNI", "SUIT"] # check validity of atlas parameters if space not in _valid_spaces: raise_error( - f'The parameter `space` ({space}) needs to be one of the ' - f'following: {_valid_spaces}') + f"The parameter `space` ({space}) needs to be one of the " + f"following: {_valid_spaces}" + ) # TODO: Validate this with Vera _valid_resolutions = [1] @@ -483,47 +659,52 @@ def _retrieve_suit(atlas_path, resolution, space='MNI'): resolution = _closest_resolution(resolution, _valid_resolutions) # define file names - atlas_fname = atlas_path / 'SUIT' / ( - f'SUIT_{space}Space_{resolution}mm.nii') - atlas_lname = atlas_path / 'SUIT' / ( - f'SUIT_{space}Space_{resolution}mm.tsv') + atlas_fname = ( + atlas_dir / "SUIT" / (f"SUIT_{space}Space_{resolution}mm.nii") + ) + atlas_lname = ( + atlas_dir / "SUIT" / (f"SUIT_{space}Space_{resolution}mm.tsv") + ) # check existence of atlas if not (atlas_fname.exists() and atlas_lname.exists()): atlas_fname.parent.mkdir(exist_ok=True, parents=True) - logger.info( - 'At least one of the atlas files is missing. ' - 'Fetching.') + logger.info("At least one of the atlas files is missing, fetching.") url_basis = ( - 'https://github.com/DiedrichsenLab/cerebellar_atlases/raw' - '/master/Diedrichsen_2009/') - url_MNI = url_basis + 'atl-Anatom_space-MNI_dseg.nii' - url_SUIT = url_basis + 'atl-Anatom_space-SUIT_dseg.nii' - url_labels = url_basis + 'atl-Anatom.tsv' + "https://github.com/DiedrichsenLab/cerebellar_atlases/raw" + "/master/Diedrichsen_2009/" + ) + url_MNI = url_basis + "atl-Anatom_space-MNI_dseg.nii" + url_SUIT = url_basis + "atl-Anatom_space-SUIT_dseg.nii" + url_labels = url_basis + "atl-Anatom.tsv" - if space == 'MNI': - logger.info(f'Downloading {url_MNI}') + if space == "MNI": + logger.info(f"Downloading {url_MNI}") atlas_download = requests.get(url_MNI) - with open(atlas_fname, 'wb') as f: + with open(atlas_fname, "wb") as f: f.write(atlas_download.content) else: # if not MNI, then SUIT - logger.info(f'Downloading {url_SUIT}') + logger.info(f"Downloading {url_SUIT}") atlas_download = requests.get(url_SUIT) - with open(atlas_fname, 'wb') as f: + with open(atlas_fname, "wb") as f: f.write(atlas_download.content) labels_download = requests.get(url_labels) labels = pd.read_csv( io.StringIO(labels_download.content.decode("utf-8")), - sep='\t', usecols=['name']) + sep="\t", + usecols=["name"], + ) - labels.to_csv(atlas_lname, sep='\t', index=False) - if not atlas_fname.exists() and \ - atlas_lname.exists(): # pragma: no cover - raise_error('There was a problem fetching the atlases.') + labels.to_csv(atlas_lname, sep="\t", index=False) + if ( + not atlas_fname.exists() and atlas_lname.exists() + ): # pragma: no cover + raise_error("There was a problem fetching the atlases.") - labels = pd.read_csv( - atlas_lname, sep='\t', usecols=['name'])['name'].to_list() + labels = pd.read_csv(atlas_lname, sep="\t", usecols=["name"])[ + "name" + ].to_list() return atlas_fname, labels diff --git a/junifer/data/tests/test_atlas.py b/junifer/data/tests/test_atlas.py deleted file mode 100644 index 0e4cd17e5..000000000 --- a/junifer/data/tests/test_atlas.py +++ /dev/null @@ -1,223 +0,0 @@ -"""Provide tests for atlas.""" - -import tempfile -import pytest -from pathlib import Path -from numpy.testing import assert_array_equal, assert_array_almost_equal - -from junifer.data.atlases import ( - register_atlas, list_atlases, - load_atlas, - _retrieve_schaefer, - _retrieve_suit, - _retrieve_atlas, - _retrieve_tian, -) - - -def test_register_atlas(): - """Test atlas registration.""" - - atlases = list_atlases() - assert 'testatlas' not in atlases - - register_atlas('testatlas', 'testatlas.nii.gz', ['1', '2', '3']) - - atlases = list_atlases() - assert 'testatlas' in atlases - - _, lbl, fname = load_atlas('testatlas', path_only=True) - - assert lbl == ['1', '2', '3'] - assert fname.name == 'testatlas.nii.gz' # type: ignore - - with pytest.raises(ValueError, match=r"already registered."): - register_atlas('testatlas', 'testatlas.nii.gz', ['1', '2', '3']) - - with pytest.raises(ValueError, match=r"built-in atlas"): - register_atlas('SUITxSUIT', 'testatlas.nii.gz', ['1', '2', '3'], - overwrite=True) - - register_atlas('testatlas', 'testatlas_2.nii.gz', ['1', '2', '6'], - overwrite=True) - - register_atlas('testatlas', Path('testatlas_2.nii.gz'), ['1', '2', '6'], - overwrite=True) - - _, lbl, fname = load_atlas('testatlas', path_only=True) - - assert lbl == ['1', '2', '6'] - assert fname.name == 'testatlas_2.nii.gz' # type: ignore - - -def test_wrong_atlas(): - """Test invalid atlas.""" - - with pytest.raises(ValueError, match=r"not found"): - load_atlas('wrongatlas') - with pytest.raises(ValueError, match=r"provided atlas name"): - _retrieve_atlas('wrongatlas') - - -def test_schaefer_atlas(): - """Test Schaefer atlas.""" - - atlases = list_atlases() - - for n_rois in range(100, 1001, 100): - for t_net in [7, 17]: - t_name = f'Schaefer{n_rois}x{t_net}' - assert t_name in atlases - - with tempfile.TemporaryDirectory() as tmpdir: - fname1 = 'Schaefer2018_100Parcels_7Networks_order_FSLMNI152_1mm.nii.gz' - fname2 = 'Schaefer2018_100Parcels_7Networks_order_FSLMNI152_2mm.nii.gz' - - img, lbl, fname = load_atlas('Schaefer100x7', atlas_dir=tmpdir) - - assert img is not None - assert fname.name == fname1 # type: ignore - - assert len(lbl) == 100 - assert_array_equal(img.header['pixdim'][1:4], [1, 1, 1]) - - # test with Path - img, lbl, fname = load_atlas('Schaefer100x7', atlas_dir=Path(tmpdir)) - - img2, lbl, fname = load_atlas( - 'Schaefer100x7', atlas_dir=tmpdir, resolution=3) - assert fname.name == fname2 # type: ignore - assert len(lbl) == 100 - assert img2 is not None - assert_array_equal(img2.header['pixdim'][1:4], [2, 2, 2]) - - img2, lbl, fname = load_atlas( - 'Schaefer100x7', atlas_dir=tmpdir, resolution=2.1) - assert fname.name == fname2 # type: ignore - assert len(lbl) == 100 - assert img2 is not None - assert_array_equal(img2.header['pixdim'][1:4], [2, 2, 2]) - - img2, lbl, fname = load_atlas( - 'Schaefer100x7', atlas_dir=tmpdir, resolution=1.99) - assert fname.name == fname1 # type: ignore - assert len(lbl) == 100 - assert img2 is not None - assert_array_equal(img2.header['pixdim'][1:4], [1, 1, 1]) - - img2, lbl, fname = load_atlas( - 'Schaefer100x7', atlas_dir=tmpdir, resolution=0.5) - assert fname.name == fname1 # type: ignore - assert len(lbl) == 100 - assert img2 is not None - assert_array_equal(img2.header['pixdim'][1:4], [1, 1, 1]) - - with pytest.raises(ValueError, match=r"The parameter `n_rois`"): - _retrieve_schaefer(tmpdir, 1, 101, 7) - - with pytest.raises(ValueError, match=r"The parameter `yeo_networks`"): - _retrieve_schaefer(tmpdir, 1, 100, 8) - - # Test without a dir - - img, lbl, fname = load_atlas('Schaefer100x7') - assert img is not None - home_dir = Path().home() / 'junifer' / 'data' / 'atlas' - assert home_dir in fname.parents # type: ignore - - -def test_suit(): - """Test SUIT atlas.""" - - atlases = list_atlases() - assert 'SUITxSUIT' in atlases - assert 'SUITxMNI' in atlases - - with tempfile.TemporaryDirectory() as tmpdir: - img, lbl, fname = load_atlas('SUITxSUIT', atlas_dir=tmpdir) - fname1 = 'SUIT_SUITSpace_1mm.nii' - assert img is not None - assert fname.name == fname1 # type: ignore - assert len(lbl) == 34 - assert_array_equal(img.header['pixdim'][1:4], [1, 1, 1]) - - img, lbl, fname = load_atlas('SUITxSUIT', atlas_dir=tmpdir) - fname1 = 'SUIT_SUITSpace_1mm.nii' - assert img is not None - assert fname.name == fname1 # type: ignore - assert len(lbl) == 34 - assert_array_equal(img.header['pixdim'][1:4], [1, 1, 1]) - - img, lbl, fname = load_atlas('SUITxMNI', atlas_dir=tmpdir) - fname1 = 'SUIT_MNISpace_1mm.nii' - assert img is not None - assert fname.name == fname1 # type: ignore - assert len(lbl) == 34 - assert_array_equal(img.header['pixdim'][1:4], [1, 1, 1]) - - with pytest.raises(ValueError, match=r"The parameter `space`"): - _retrieve_suit(tmpdir, 1, space='wrong') - - -def test_tian(): - """Test TIAN atlas.""" - - atlases = list_atlases() - assert 'TianxS1x3TxMNI6thgeneration' in atlases - assert 'TianxS2x3TxMNI6thgeneration' in atlases - assert 'TianxS3x3TxMNI6thgeneration' in atlases - assert 'TianxS4x3TxMNI6thgeneration' in atlases - - assert 'TianxS1x3TxMNInonlinear2009cAsym' in atlases - assert 'TianxS2x3TxMNInonlinear2009cAsym' in atlases - assert 'TianxS3x3TxMNInonlinear2009cAsym' in atlases - assert 'TianxS4x3TxMNInonlinear2009cAsym' in atlases - - assert 'TianxS1x7TxMNI6thgeneration' in atlases - assert 'TianxS2x7TxMNI6thgeneration' in atlases - assert 'TianxS3x7TxMNI6thgeneration' in atlases - assert 'TianxS4x7TxMNI6thgeneration' in atlases - - with tempfile.TemporaryDirectory() as tmpdir: - for scale, n_lbl in zip([1, 2, 3, 4], [16, 32, 50, 54]): - img, lbl, fname = load_atlas( - f'TianxS{scale}x3TxMNI6thgeneration', atlas_dir=tmpdir) - fname1 = f'Tian_Subcortex_S{scale}_3T_1mm.nii.gz' - assert img is not None - assert fname.name == fname1 # type: ignore - assert len(lbl) == n_lbl - assert_array_equal(img.header['pixdim'][1:4], [1, 1, 1]) - - img, lbl, fname = load_atlas( - f'TianxS{scale}x3TxMNI6thgeneration', atlas_dir=tmpdir, - resolution=2) - fname1 = f'Tian_Subcortex_S{scale}_3T.nii.gz' - assert img is not None - assert fname.name == fname1 # type: ignore - assert len(lbl) == n_lbl - assert_array_equal(img.header['pixdim'][1:4], [2, 2, 2]) - - img, lbl, fname = load_atlas( - f'TianxS{scale}x3TxMNInonlinear2009cAsym', atlas_dir=tmpdir) - fname1 = f'Tian_Subcortex_S{scale}_3T_2009cAsym.nii.gz' - assert img is not None - assert fname.name == fname1 # type: ignore - assert len(lbl) == n_lbl - assert_array_equal(img.header['pixdim'][1:4], [2, 2, 2]) - - for scale, n_lbl in zip([1, 2, 3, 4], [16, 34, 54, 62]): - img, lbl, fname = load_atlas( - f'TianxS{scale}x7TxMNI6thgeneration', atlas_dir=tmpdir) - fname1 = f'Tian_Subcortex_S{scale}_7T.nii.gz' - assert img is not None - assert fname.name == fname1 # type: ignore - assert len(lbl) == n_lbl - assert_array_almost_equal( - img.header['pixdim'][1:4], [1.6, 1.6, 1.6]) - - with pytest.raises(ValueError, match=r"The parameter `space`"): - _retrieve_tian(tmpdir, resolution=1, scale=1, space='wrong') - - with pytest.raises(ValueError, match=r"The parameter `magneticfield`"): - _retrieve_tian( - tmpdir, resolution=1, scale=1, magneticfield='wrong') diff --git a/junifer/data/tests/test_atlases.py b/junifer/data/tests/test_atlases.py new file mode 100644 index 000000000..6b9eea365 --- /dev/null +++ b/junifer/data/tests/test_atlases.py @@ -0,0 +1,481 @@ +"""Provide tests for atlas.""" + +# Authors: Federico Raimondo +# Vera Komeyer +# Synchon Mandal +# License: AGPL + +from pathlib import Path +from typing import List + +import pytest +from numpy.testing import assert_array_almost_equal, assert_array_equal + +from junifer.data.atlases import ( + _retrieve_atlas, + _retrieve_schaefer, + _retrieve_suit, + _retrieve_tian, + list_atlases, + load_atlas, + register_atlas, +) + + +def test_register_atlas_built_in_check() -> None: + """Test atlas registration check for built-in atlas.""" + with pytest.raises(ValueError, match=r"built-in atlas"): + register_atlas( + name="SUITxSUIT", + atlas_path="testatlas.nii.gz", + atl_labels=["1", "2", "3"], + overwrite=True, + ) + + +def test_list_atlases_incorrect() -> None: + """Test incorrect information check for list atlases.""" + atlases = list_atlases() + assert "testatlas" not in atlases + + +def test_register_atlas_already_registered() -> None: + """Test atlas registration check for already registered atlas.""" + # Register custom atlas + register_atlas( + name="testatlas", + atlas_path="testatlas.nii.gz", + atl_labels=["1", "2", "3"], + ) + # Try registering again + with pytest.raises(ValueError, match=r"already registered."): + register_atlas( + name="testatlas", + atlas_path="testatlas.nii.gz", + atl_labels=["1", "2", "3"], + ) + + +@pytest.mark.parametrize( + "name, atlas_path, atlas_labels, overwrite", + [ + ("testatlas_1", "testatlas_1.nii.gz", ["1", "2", "3"], True), + ("testatlas_2", "testatlas_2.nii.gz", ["1", "2", "6"], True), + ("testatlas_3", Path("testatlas_3.nii.gz"), ["1", "2", "6"], True), + ], +) +def test_register_atlas( + name: str, + atlas_path: str, + atlas_labels: List[str], + overwrite: bool, +) -> None: + """Test atlas registration. + + Parameters + ---------- + name : str + The parametrized atlas name. + atlas_path : str or pathlib.Path + The parametrized atlas path. + atlas_labels : list of str + The parametrized atlas labels. + overwrite : bool + The parametrized atlas overwrite value. + + """ + # Register custom atlas + register_atlas( + name=name, + atlas_path=atlas_path, + atl_labels=atlas_labels, + overwrite=overwrite, + ) + # List available atlas and check registration + atlases = list_atlases() + assert name in atlases + # Load registered atlas + _, lbl, fname = load_atlas(name=name, path_only=True) + # Check values for registered atlas + assert lbl == atlas_labels + assert fname.name == f"{name}.nii.gz" + + +@pytest.mark.parametrize( + "atlas_name", + [ + "SUITxSUIT", + "SUITxMNI", + "Schaefer100x7", + "Schaefer100x17", + "TianxS1x7TxMNI6thgeneration", + "TianxS3x3TxMNI6thgeneration", + "TianxS4x3TxMNInonlinear2009cAsym", + ], +) +def test_list_atlases_correct(atlas_name: str) -> None: + """Test correct information check for list atlases. + + Parameters + ---------- + atlas_name : str + The parametrized atlas name. + + """ + atlases = list_atlases() + assert atlas_name in atlases + + +def test_load_atlas_incorrect() -> None: + """Test loading of invalid atlas.""" + with pytest.raises(ValueError, match=r"not found"): + load_atlas("wrongatlas") + + +def test_retrieve_atlas_incorrect() -> None: + """Test retrieval of invalid atlas.""" + with pytest.raises(ValueError, match=r"provided atlas name"): + _retrieve_atlas("wrongatlas") + + +# TODO: paramdtrize test +def test_schaefer_atlas(tmp_path: Path) -> None: + """Test Schaefer atlas. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + atlases = list_atlases() + for n_rois in range(100, 1001, 100): + for t_net in [7, 17]: + t_name = f"Schaefer{n_rois}x{t_net}" + assert t_name in atlases + + # Define atlas file names + fname1 = "Schaefer2018_100Parcels_7Networks_order_FSLMNI152_1mm.nii.gz" + fname2 = "Schaefer2018_100Parcels_7Networks_order_FSLMNI152_2mm.nii.gz" + + # Load atlas + img, lbl, fname = load_atlas( + name="Schaefer100x7", atlas_dir=str(tmp_path.absolute()) + ) + # Check atlas values + assert img is not None + assert fname.name == fname1 + assert len(lbl) == 100 + assert_array_equal(img.header["pixdim"][1:4], [1, 1, 1]) + + # Test with Path + img, lbl, fname = load_atlas(name="Schaefer100x7", atlas_dir=tmp_path) + # Load atlas + img2, lbl, fname = load_atlas( + name="Schaefer100x7", + atlas_dir=tmp_path, + resolution=3, + ) + # Check atlas values + assert fname.name == fname2 + assert len(lbl) == 100 + assert img2 is not None + assert_array_equal(img2.header["pixdim"][1:4], [2, 2, 2]) + # Load atlas + img2, lbl, fname = load_atlas( + "Schaefer100x7", + atlas_dir=tmp_path, + resolution=2.1, + ) + # Check atlas values + assert fname.name == fname2 + assert len(lbl) == 100 + assert img2 is not None + assert_array_equal(img2.header["pixdim"][1:4], [2, 2, 2]) + # Load atlas + img2, lbl, fname = load_atlas( + "Schaefer100x7", + atlas_dir=tmp_path, + resolution=1.99, + ) + # Check atlas values + assert fname.name == fname1 + assert len(lbl) == 100 + assert img2 is not None + assert_array_equal(img2.header["pixdim"][1:4], [1, 1, 1]) + # Load atlas + img2, lbl, fname = load_atlas( + "Schaefer100x7", + atlas_dir=tmp_path, + resolution=0.5, + ) + # Check atlas values + assert fname.name == fname1 + assert len(lbl) == 100 + assert img2 is not None + assert_array_equal(img2.header["pixdim"][1:4], [1, 1, 1]) + + +def test_load_atlas_schaefer() -> None: + """Test Schaefer atlas loading.""" + img, lbl, fname = load_atlas(name="Schaefer100x7") + assert img is not None + home_dir = Path().home() / "junifer" / "data" / "atlas" + assert home_dir in fname.parents + + +def test_retrieve_schaefer_incorrect_n_rois(tmp_path: Path) -> None: + """Test retrieve schaefer with incorrect n_rois. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + with pytest.raises(ValueError, match=r"The parameter `n_rois`"): + _retrieve_schaefer( + atlas_dir=tmp_path, resolution=1, n_rois=101, yeo_networks=7 + ) + + +def test_retrieve_schaefer_incorrect_yeo_networks(tmp_path: Path) -> None: + """Test retrieve schaefer with incorrect yeo_networks. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + with pytest.raises(ValueError, match=r"The parameter `yeo_networks`"): + _retrieve_schaefer( + atlas_dir=tmp_path, resolution=1, n_rois=100, yeo_networks=8 + ) + + +# TODO: parametrize test +def test_suit(tmp_path: Path) -> None: + """Test SUIT atlas. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + atlases = list_atlases() + assert "SUITxSUIT" in atlases + assert "SUITxMNI" in atlases + + # Load atlas + img, lbl, fname = load_atlas(name="SUITxSUIT", atlas_dir=tmp_path) + fname1 = "SUIT_SUITSpace_1mm.nii" + assert img is not None + assert fname.name == fname1 + assert len(lbl) == 34 + assert_array_equal(img.header["pixdim"][1:4], [1, 1, 1]) + + # Load atlas + img, lbl, fname = load_atlas(name="SUITxSUIT", atlas_dir=tmp_path) + fname1 = "SUIT_SUITSpace_1mm.nii" + assert img is not None + assert fname.name == fname1 + assert len(lbl) == 34 + assert_array_equal(img.header["pixdim"][1:4], [1, 1, 1]) + + # Load atlas + img, lbl, fname = load_atlas(name="SUITxMNI", atlas_dir=tmp_path) + fname1 = "SUIT_MNISpace_1mm.nii" + assert img is not None + assert fname.name == fname1 + assert len(lbl) == 34 + assert_array_equal(img.header["pixdim"][1:4], [1, 1, 1]) + + +def test_retrieve_suit_incorrect_space(tmp_path: Path) -> None: + """Test retrieve suit with incorrect space. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + with pytest.raises(ValueError, match=r"The parameter `space`"): + _retrieve_suit(atlas_dir=tmp_path, resolution=1, space="wrong") + + +@pytest.mark.parametrize( + "scale, n_label", + [ + (1, 16), + (2, 32), + (3, 50), + (4, 54), + ], +) +def test_tian_3T_6thgeneration( + tmp_path: Path, + scale: int, + n_label: int, +) -> None: + """Test Tian atlas. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + scale : int + The parametrized scale values. + n_label : int + The parametrized n_label values. + + """ + atlases = list_atlases() + assert "TianxS1x3TxMNI6thgeneration" in atlases + assert "TianxS2x3TxMNI6thgeneration" in atlases + assert "TianxS3x3TxMNI6thgeneration" in atlases + assert "TianxS4x3TxMNI6thgeneration" in atlases + # Load atlas + img, lbl, fname = load_atlas( + name=f"TianxS{scale}x3TxMNI6thgeneration", + atlas_dir=tmp_path, + ) + fname1 = f"Tian_Subcortex_S{scale}_3T_1mm.nii.gz" + assert img is not None + assert fname.name == fname1 + assert len(lbl) == n_label + assert_array_equal(img.header["pixdim"][1:4], [1, 1, 1]) + # Load atlas + img, lbl, fname = load_atlas( + name=f"TianxS{scale}x3TxMNI6thgeneration", + atlas_dir=tmp_path, + resolution=2, + ) + fname1 = f"Tian_Subcortex_S{scale}_3T.nii.gz" + assert img is not None + assert fname.name == fname1 + assert len(lbl) == n_label + assert_array_equal(img.header["pixdim"][1:4], [2, 2, 2]) + + +@pytest.mark.parametrize( + "scale, n_label", + [ + (1, 16), + (2, 32), + (3, 50), + (4, 54), + ], +) +def test_tian_3T_nonlinear2009cAsym( + tmp_path: Path, + scale: int, + n_label: int, +) -> None: + """Test Tian atlas. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + scale : int + The parametrized scale values. + n_label : int + The parametrized n_label values. + + """ + atlases = list_atlases() + assert "TianxS1x3TxMNInonlinear2009cAsym" in atlases + assert "TianxS2x3TxMNInonlinear2009cAsym" in atlases + assert "TianxS3x3TxMNInonlinear2009cAsym" in atlases + assert "TianxS4x3TxMNInonlinear2009cAsym" in atlases + # Load atlas + img, lbl, fname = load_atlas( + name=f"TianxS{scale}x3TxMNInonlinear2009cAsym", + atlas_dir=tmp_path, + ) + fname1 = f"Tian_Subcortex_S{scale}_3T_2009cAsym.nii.gz" + assert img is not None + assert fname.name == fname1 + assert len(lbl) == n_label + assert_array_equal(img.header["pixdim"][1:4], [2, 2, 2]) + + +@pytest.mark.parametrize( + "scale, n_label", + [ + (1, 16), + (2, 34), + (3, 54), + (4, 62), + ], +) +def test_tian_7T_6thgeneration( + tmp_path: Path, + scale: int, + n_label: int, +) -> None: + """Test Tian atlas. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + scale : int + The parametrized scale values. + n_label : int + The parametrized n_label values. + + """ + atlases = list_atlases() + assert "TianxS1x7TxMNI6thgeneration" in atlases + assert "TianxS2x7TxMNI6thgeneration" in atlases + assert "TianxS3x7TxMNI6thgeneration" in atlases + assert "TianxS4x7TxMNI6thgeneration" in atlases + # Load atlas + img, lbl, fname = load_atlas( + name=f"TianxS{scale}x7TxMNI6thgeneration", atlas_dir=tmp_path + ) + fname1 = f"Tian_Subcortex_S{scale}_7T.nii.gz" + assert img is not None + assert fname.name == fname1 + assert len(lbl) == n_label + assert_array_almost_equal(img.header["pixdim"][1:4], [1.6, 1.6, 1.6]) + + +def test_retrieve_tian_incorrect_space(tmp_path: Path) -> None: + """Test retrieve tian with incorrect space. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + with pytest.raises(ValueError, match=r"The parameter `space`"): + _retrieve_tian( + atlas_dir=tmp_path, + resolution=1, + scale=1, + space="wrong", + ) + + +def test_retrieve_tian_incorrect_magneticfield(tmp_path: Path) -> None: + """Test retrieve tian with incorrect magneticfield. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + with pytest.raises(ValueError, match=r"The parameter `magneticfield`"): + _retrieve_tian( + atlas_dir=tmp_path, + resolution=1, + scale=1, + magneticfield="wrong", + ) diff --git a/junifer/testing/__init__.py b/junifer/testing/__init__.py index bae0e931f..a1e573655 100644 --- a/junifer/testing/__init__.py +++ b/junifer/testing/__init__.py @@ -4,4 +4,4 @@ # Synchon Mandal # License: AGPL -from .datagrabbers import datagrabbers +from . import datagrabbers