From ac6a0c8bd651028a19312d7e73935e0128f070a0 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 25 Feb 2022 08:42:01 +0100 Subject: [PATCH 001/287] Parcel marker WIP --- docs/data.rst | 32 +++++++++++++++ examples/norun_ukbvm_gmd.py | 3 +- junifer/data/atlases.py | 34 ++++++++-------- junifer/markers/parcel.py | 39 ++++++++++++++++++ junifer/stats.py | 80 +++++++++++++++++++++++++++++++++++++ scratch/test_niftimasker.py | 38 ++++++++++++++++++ 6 files changed, 207 insertions(+), 19 deletions(-) create mode 100644 docs/data.rst create mode 100644 junifer/markers/parcel.py create mode 100644 junifer/stats.py create mode 100644 scratch/test_niftimasker.py diff --git a/docs/data.rst b/docs/data.rst new file mode 100644 index 000000000..080c5e4d2 --- /dev/null +++ b/docs/data.rst @@ -0,0 +1,32 @@ + +.. include:: links.inc + +Data Object +=========== + + +Introduction +^^^^^^^^^^^^ + +Data types +^^^^^^^^^^ + +.. list-table:: Built-in data types + :widths: 30 80 40 + :header-rows: 1 + + * - Name + - Description + - Example + * - `T1w` + - T1w image (3D) + - Preprocessed or Raw T1w image + * - `BOLD` + - BOLD image (4D) + - Preprocessed/Denoised BOLD image (fmriprep output) + * - `VBM_GM` + - VBM Gray Matter segmentation (3D) + - CAT output (`m0wp1` images) + * - `VBM_WM` + - VBM White Matter segmentation (3D) + - CAT output (`m0wp2` images) diff --git a/examples/norun_ukbvm_gmd.py b/examples/norun_ukbvm_gmd.py index 168478d4c..ef1da3544 100644 --- a/examples/norun_ukbvm_gmd.py +++ b/examples/norun_ukbvm_gmd.py @@ -14,7 +14,8 @@ markers = [ {'name': 'Schaefer1000x7_TrimMean80', 'kind': 'ParcelAggregation', 'atlas': 'Schaefer1000x7', - 'method': 'trimmean80'}, + 'method': 'trim_mean', + 'method_params': {'proportiontocut': 0.2}}, {'name': 'Schaefer1000x7_Mean', 'kind': 'ParcelAggregation', 'atlas': 'Schaefer1000x7', diff --git a/junifer/data/atlases.py b/junifer/data/atlases.py index 64e999e77..395517f50 100644 --- a/junifer/data/atlases.py +++ b/junifer/data/atlases.py @@ -42,7 +42,6 @@ for n_rois in range(100, 1001, 100): 'family': 'Schaefer', 'n_rois': n_rois, 'yeo_networks': t_net, - 'valid_resolutions': [1, 2] } @@ -102,8 +101,7 @@ def _check_resolution(resolution, valid_resolution): return resolution -def load_atlas(name, atlas_dir=None, resolution=None, path_only=False, - **kwargs): +def load_atlas(name, atlas_dir=None, resolution=None, path_only=False): """ Loads a brain atlas (including a label file). If it is built-in atlas and file is not present in the `atlas_dir` @@ -154,7 +152,7 @@ def load_atlas(name, atlas_dir=None, resolution=None, path_only=False, File path to the atlas image. """ - atlas_definition = _available_atlases[name] + atlas_definition = _available_atlases[name].copy() t_family = atlas_definition.pop('family') if t_family == 'CustomUserAtlas': @@ -163,7 +161,7 @@ def load_atlas(name, atlas_dir=None, resolution=None, path_only=False, else: # retrieve atlases by passing arguments on to _retrieve_atlas() atlas_fname, atlas_labels = _retrieve_atlas( - t_family, out_dir=atlas_dir, **kwargs) + t_family, atlas_dir=atlas_dir, **atlas_definition) logger.info( f'Loading atlas {atlas_fname.as_posix()}') # type: ignore @@ -226,10 +224,10 @@ def _retrieve_atlas(family, atlas_dir=None, resolution=None, **kwargs): # retrieval details per atlas if family == 'Schaefer': atlas_fname, atl_labels = \ - _retrieve_schaefer(atlas_dir, **kwargs) + _retrieve_schaefer(atlas_dir, resolution=resolution, **kwargs) elif family == 'SUIT': atlas_fname, atl_labels = \ - _retrieve_suit(atlas_dir, **kwargs) + _retrieve_suit(atlas_dir, resolution=resolution, **kwargs) else: raise_error( f"The provided atlas name {family} cannot be retrieved. ") @@ -254,10 +252,10 @@ def _closest_resolution(resolution, valid_resolution): return closest -def _retrieve_schaefer(atlas_dir, resolution, n_rois=None, yeo_network=7): +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_network: {yeo_network}') + 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] @@ -268,19 +266,19 @@ def _retrieve_schaefer(atlas_dir, resolution, n_rois=None, yeo_network=7): raise_error( f'The parameter `n_rois` ({n_rois}) needs to be one of the ' f'following: {_valid_n_rois}') - if yeo_network not in _valid_networks: + if yeo_networks not in _valid_networks: raise_error( - f'The parameter `yeo_network` ({yeo_network}) needs to be one of ' - f'the following: {_valid_networks}') + f'The parameter `yeo_networks` ({yeo_networks}) needs to be one of' + f' 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_network}Networks_order_' + 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_network}Networks_order.txt') + f'Schaefer2018_{n_rois}Parcels_{yeo_networks}Networks_order.txt') # check existance of atlas if not (atlas_fname.exists() and atlas_lname.exists()): @@ -289,7 +287,7 @@ def _retrieve_schaefer(atlas_dir, resolution, n_rois=None, yeo_network=7): 'Fetching using nilearn.') datasets.fetch_atlas_schaefer_2018( n_rois=n_rois, - yeo_networks=yeo_network, + yeo_networks=yeo_networks, resolution_mm=resolution, data_dir=atlas_dir.as_posix()) @@ -306,7 +304,7 @@ def _retrieve_schaefer(atlas_dir, resolution, n_rois=None, yeo_network=7): return atlas_fname, labels -def _retrieve_suit(out_dir, resolution, space='MNI'): +def _retrieve_suit(atlas_path, resolution, space='MNI'): logger.info('Atlas parameters:') logger.info(f'\tspace: {space}') @@ -324,9 +322,9 @@ def _retrieve_suit(out_dir, resolution, space='MNI'): resolution = _closest_resolution(resolution, _valid_resolutions) # define file names - atlas_fname = out_dir / 'SUIT' / ( + atlas_fname = atlas_path / 'SUIT' / ( f'SUIT_{space}Space_{resolution}mm.nii') - atlas_lname = out_dir / 'SUIT' / ( + atlas_lname = atlas_path / 'SUIT' / ( f'SUIT_{space}Space_{resolution}mm.tsv') # check existance of atlas diff --git a/junifer/markers/parcel.py b/junifer/markers/parcel.py new file mode 100644 index 000000000..a60bc79ba --- /dev/null +++ b/junifer/markers/parcel.py @@ -0,0 +1,39 @@ +import numpy as np + +from nilearn.maskers import NiftiLabelsMasker + +from .base import PipelineStepMixin +from ..stats import get_aggfunc_by_name +from ..data import load_atlas + +class ParcelAggregation(PipelineStepMixin): + def __init__(self, atlas, method, method_params=None): + self.atlas = atlas + self.method = method + self.method_params = {} if method_params is None else method_params + self._valid_inputs = ['T1w', 'BOLD', 'VBM_GM', 'VBM_WM'] + + def validate_input(self, input): + return any(x in input for x in self._valid_inputs) + + def get_output_kind(self, input): + return None + + def fit_transform(self, input, storage=None): + out = [] + agg_func = get_aggfunc_by_name(self.method, **self.method_params) + for kind in self._valid_inputs: + if kind in input.keys(): + t_input = input[kind].data + + # Get the min of the voxels sizes and use it as the resolution + resolution = np.min(t_input.header.get_zooms()[:3]) + t_atlas, t_labels = load_atlas( + self.atlas, resolution=resolution) + masker = NiftiLabelsMasker( + labels_img=t_atlas, + labels=t_labels + , + ) + + return input \ No newline at end of file diff --git a/junifer/stats.py b/junifer/stats.py new file mode 100644 index 000000000..7d72aeece --- /dev/null +++ b/junifer/stats.py @@ -0,0 +1,80 @@ +from functools import partial +import numpy as np +from scipy.stats.mstats import winsorize +from scipy.stats import trim_mean +from .utils import logger, raise_error + + +def get_aggfunc_by_name(name, func_params): + """ + Helper function to get an aggregation function by its name. + + Parameters + ---------- + name : str + Name to identify the function. Currently supported names and + corresponding functions are: + 'winsorized_mean' -> scipy.stats.mstats.winsorize + 'mean' -> np.mean + 'std' -> np.std + 'trim_mean' -> scipy.stats.trim_mean + + func_params : dict + Parameters to pass to the function. + E.g. for 'winsorized_mean': func_params = {'limits': [0.1, 0.1]} + + Returns + ------- + func : function + Respective function with `func_params` parameter set. + """ + + # check validity of names + _valid_func_names = {'winsorized_mean', 'mean', 'std', 'trim_mean'} + + # apply functions + if name == 'winsorized_mean': + # check validity of func_params + limits = func_params.get('limits') + if all((lim >= 0.0 and lim <= 1) for lim in limits): + logger.info(f'Limits for winsorized mean are set to {limits}.') + else: + raise_error( + 'Limits for the winsorized mean must be between 0 and 1.') + # partially interpret func_params + func = partial(winsorized_mean, **func_params) + elif name == 'mean': + func = np.mean + elif name == 'std': + func = np.std + elif name == 'trim_mean': + func = partial(trim_mean, **func_params) + else: + raise_error(f'Function {name} unknown. Please provide any of ' + f'{_valid_func_names}') + return func + + +def winsorized_mean(data, axis=None, **win_params): + """ + Compute a winsorized mean by chaining winsorization and mean. + + Parameters + ---------- + data : array + Data to calculate winsorized mean on. + win_params : dict + Dictionary containing the keyword arguments for the winsorize function. + E.g. {'limits': [0.1, 0.1]} + + Returns + ------- + win_mean : np.ndarray + Winsorized mean of the inputted data with the winsorize settings + applied as specified in win_params. + """ + + win_dat = winsorize(data, axis=axis, **win_params) + win_mean = win_dat.mean(axis=axis) + + return win_mean diff --git a/scratch/test_niftimasker.py b/scratch/test_niftimasker.py new file mode 100644 index 000000000..6308cbcbd --- /dev/null +++ b/scratch/test_niftimasker.py @@ -0,0 +1,38 @@ +import numpy as np +from nilearn import datasets +import nibabel as nib + +from nilearn.image import resample_to_img, math_img +from nilearn.maskers import NiftiMasker, NiftiLabelsMasker +oasis_dataset = datasets.fetch_oasis_vbm(n_subjects=1) + +vbm = oasis_dataset.gray_matter_maps[0] +atlas = datasets.fetch_atlas_schaefer_2018(n_rois=100) + +nifti_masker = NiftiLabelsMasker(labels_img=atlas.maps) +auto = nifti_masker.fit_transform(vbm) + + +img = nib.load(vbm) + +atlas_img_res = resample_to_img( + atlas.maps, + img, + interpolation='nearest', +) +atlas_bin = math_img( + 'img != 0', + img=atlas_img_res, +) + + +masker = NiftiMasker(atlas_bin, target_affine=img.affine) + +data = masker.fit_transform(img) +atlas_values = masker.transform(atlas_img_res) +atlas_values = np.squeeze(atlas_values) + +manual = [] +for t_v in sorted(np.unique(atlas_values)): + t_values = np.mean(data[:, atlas_values == t_v]) + manual.append(t_values) \ No newline at end of file -- 2.52.0 From 47bd92adf243f919092070d93c4a8f8da557615f Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 25 Feb 2022 13:08:11 +0100 Subject: [PATCH 002/287] New example + Parcel Aggregation marker --- examples/run_compute_parcel_mean.py | 59 ++++++++++++++++++++++++ junifer/markers/base.py | 47 +++++++++++++++++++ junifer/markers/parcel.py | 70 ++++++++++++++++++----------- scratch/test_niftimasker.py | 42 ++++++++++++++++- 4 files changed, 191 insertions(+), 27 deletions(-) create mode 100644 examples/run_compute_parcel_mean.py diff --git a/examples/run_compute_parcel_mean.py b/examples/run_compute_parcel_mean.py new file mode 100644 index 000000000..4c842129d --- /dev/null +++ b/examples/run_compute_parcel_mean.py @@ -0,0 +1,59 @@ +""" +Computer Parcel Aggregation. +============================ + +This example uses a ParcelAggregation marker to compute the mean of each parcel +using the Schaefer atlas (100 rois, 7 Yeo networks) for both a 3D and 4D nifti + +Authors: Federico Raimondo + +License: BSD 3 clause +""" + +import nilearn + +from junifer.utils import configure_logging +from junifer.markers.parcel import ParcelAggregation + +############################################################################### +# Set the logging level to info to see extra information +configure_logging(level='INFO') + + +############################################################################### +# Load the VBM GM data (3d): +# - Fetch the Oasis dataset +oasis_dataset = nilearn.datasets.fetch_oasis_vbm(n_subjects=1) +vbm_fname = oasis_dataset.gray_matter_maps[0] +vbm_img = nilearn.image.load_img(vbm_fname) + +############################################################################### +# Load the functional data (4d): +# - Fetch the SPM auditory dataset +# - Concatenate the functional data into one 4D image +s_func_data = nilearn.datasets.fetch_spm_auditory() +fmri_img = nilearn.image.concat_imgs(s_func_data.func) + +############################################################################### +# Define the marker +marker = ParcelAggregation(atlas='Schaefer100x7', method='mean') + +############################################################################### +# Prepare the input +input = { + 'BOLD': {'data': fmri_img}, + 'VBM_GM': {'data': vbm_img} +} + +############################################################################### +# Fit transform the data +out = marker.fit_transform(input) + +############################################################################### +# Check the results + +print(out.keys()) +print(out['VBM_GM']['data'].shape) # Shape is (1 x parcels) + +print(out.keys()) +print(out['BOLD']['data'].shape) # Shape is (timepoints x parcels) \ No newline at end of file diff --git a/junifer/markers/base.py b/junifer/markers/base.py index 1aeb3b284..6531ba039 100644 --- a/junifer/markers/base.py +++ b/junifer/markers/base.py @@ -1,5 +1,7 @@ # Authors: Federico Raimondo # License: AGPL +from ..utils import logger + class PipelineStepMixin(): @@ -60,3 +62,48 @@ class PipelineStepMixin(): def fit_transform(self, input): raise NotImplementedError('fit_transform not implemented') + + +class BaseMarker(PipelineStepMixin): + """Base class for all markers.""" + + def __init__(self, on): + if not isinstance(on, list): + on = [on] + self._valid_inputs = on + + def get_meta(self): + t_meta = {} + t_meta['class'] = self.__class__.__name__ + for k, v in vars(self).items(): + if not k.startswith('_'): + t_meta[k] = v + return t_meta + + def validate_input(self, input): + return any(x in input for x in self._valid_inputs) + + def get_output_kind(self, input): + return None + + def compute(self, input): + raise NotImplementedError('compute not implemented') + + def fit_transform(self, input, storage=None): + out = {} + meta = input.get('meta', {}) + for kind in self._valid_inputs: + if kind in input.keys(): + logger.info(f'Computing {kind}') + t_input = input[kind] + t_meta = meta.copy() + t_meta.update(t_input.get('meta', {})) + t_meta.update(self.get_meta()) + t_out = self.compute(t_input) + t_out.update(meta=t_meta) + if storage is not None: + storage.store_2d(t_out) + else: + out[kind] = t_out + + return out diff --git a/junifer/markers/parcel.py b/junifer/markers/parcel.py index a60bc79ba..c370d644c 100644 --- a/junifer/markers/parcel.py +++ b/junifer/markers/parcel.py @@ -1,39 +1,59 @@ import numpy as np -from nilearn.maskers import NiftiLabelsMasker +from nilearn.maskers import NiftiMasker +from nilearn.image import resample_to_img, math_img -from .base import PipelineStepMixin +from .base import BaseMarker from ..stats import get_aggfunc_by_name from ..data import load_atlas -class ParcelAggregation(PipelineStepMixin): - def __init__(self, atlas, method, method_params=None): + +class ParcelAggregation(BaseMarker): + def __init__(self, atlas, method, method_params=None, on=None): + if on is None: + on = ['T1w', 'BOLD', 'VBM_GM', 'VBM_WM'] + super().__init__(on=on) self.atlas = atlas self.method = method self.method_params = {} if method_params is None else method_params - self._valid_inputs = ['T1w', 'BOLD', 'VBM_GM', 'VBM_WM'] - def validate_input(self, input): - return any(x in input for x in self._valid_inputs) + def compute(self, input): + t_input = input['data'] + agg_func = get_aggfunc_by_name( + self.method, func_params=self.method_params) + # Get the min of the voxels sizes and use it as the resolution + resolution = np.min(t_input.header.get_zooms()[:3]) + t_atlas, t_labels, _ = load_atlas( + self.atlas, resolution=resolution) + atlas_img_res = resample_to_img( + t_atlas, + t_input, + interpolation='nearest', + ) + atlas_bin = math_img( + 'img != 0', + img=atlas_img_res, + ) - def get_output_kind(self, input): - return None + masker = NiftiMasker(atlas_bin, target_affine=t_input.affine) - def fit_transform(self, input, storage=None): - out = [] - agg_func = get_aggfunc_by_name(self.method, **self.method_params) - for kind in self._valid_inputs: - if kind in input.keys(): - t_input = input[kind].data + # Mask the input data and the atlas + data = masker.fit_transform(t_input) + atlas_values = masker.transform(atlas_img_res) + atlas_values = np.squeeze(atlas_values).astype(int) - # Get the min of the voxels sizes and use it as the resolution - resolution = np.min(t_input.header.get_zooms()[:3]) - t_atlas, t_labels = load_atlas( - self.atlas, resolution=resolution) - masker = NiftiLabelsMasker( - labels_img=t_atlas, - labels=t_labels - , - ) + # Get the values for each parcel and apply agg function + atlas_roi_vals = sorted(np.unique(atlas_values)) + out_labels = [] + out_values = [] + # Iterate over the parcels (existing) + for t_v in atlas_roi_vals: + t_values = agg_func(data[:, atlas_values == t_v], axis=-1) + out_values.append(t_values) + # Update the labels just in case a parcel has no voxels + # in it + out_labels.append(t_labels[t_v - 1]) - return input \ No newline at end of file + out_values = np.array(out_values).transpose() + out = dict(data=out_values, labels=out_labels) + return out diff --git a/scratch/test_niftimasker.py b/scratch/test_niftimasker.py index 6308cbcbd..db7fcd0a2 100644 --- a/scratch/test_niftimasker.py +++ b/scratch/test_niftimasker.py @@ -1,9 +1,14 @@ +from distutils.command.config import config import numpy as np from nilearn import datasets import nibabel as nib +from junifer.utils import configure_logging from nilearn.image import resample_to_img, math_img from nilearn.maskers import NiftiMasker, NiftiLabelsMasker + +configure_logging(level='INFO') + oasis_dataset = datasets.fetch_oasis_vbm(n_subjects=1) vbm = oasis_dataset.gray_matter_maps[0] @@ -30,9 +35,42 @@ masker = NiftiMasker(atlas_bin, target_affine=img.affine) data = masker.fit_transform(img) atlas_values = masker.transform(atlas_img_res) -atlas_values = np.squeeze(atlas_values) +atlas_values = np.squeeze(atlas_values).astype(int) manual = [] for t_v in sorted(np.unique(atlas_values)): t_values = np.mean(data[:, atlas_values == t_v]) - manual.append(t_values) \ No newline at end of file + manual.append(t_values) + + +from junifer.markers.parcel import ParcelAggregation + +marker = ParcelAggregation(atlas='Schaefer100x7', method='mean') +input = dict(VBM_GM=dict(data=img)) +jun_values3d = marker.fit_transform(input)['VBM_GM']['data'] + + +# Now do the 4D case +from nilearn.datasets import fetch_spm_auditory +from nilearn.image import concat_imgs +subject_data = fetch_spm_auditory() +fmri_img = concat_imgs(subject_data.func) + +nifti_masker = NiftiLabelsMasker(labels_img=atlas.maps) +auto4d = nifti_masker.fit_transform(fmri_img) + +marker = ParcelAggregation(atlas='Schaefer100x7', method='mean') +input = dict(BOLD=dict(data=fmri_img)) +jun_values4d = marker.fit_transform(input)['BOLD']['data'] + + +print(auto.shape) +print(jun_values3d.shape) +print(auto4d.shape) +print(jun_values4d.shape) + + +# Now both +marker = ParcelAggregation(atlas='Schaefer100x7', method='mean') +input = dict(BOLD=dict(data=fmri_img), VBM_GM=dict(data=img)) +jun_both = marker.fit_transform(input) \ No newline at end of file -- 2.52.0 From b25656bc7ac8c4f819933f9c708baa6905ccf8be Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 4 Mar 2022 11:53:35 +0100 Subject: [PATCH 003/287] More tests + SQLITE --- examples/run_compute_parcel_mean.py | 2 +- junifer/configs/juseless.py | 1 + junifer/data/atlases.py | 14 +- junifer/data/tests/test_atlas.py | 104 ++++++++++++++ junifer/datagrabber/base.py | 29 +++- ...{test_base.py => test_base_datagrabber.py} | 21 ++- junifer/datareader/default.py | 3 + .../datareader/tests/test_default_reader.py | 16 +++ junifer/markers/base.py | 30 ++-- junifer/markers/parcel.py | 6 +- junifer/markers/tests/test_base_marker.py | 44 ++++++ junifer/storage/__init__.py | 1 - junifer/storage/base.py | 67 +++++++++ junifer/storage/sqlite.py | 130 ++++++++++++++++++ requirements.txt | 3 +- 15 files changed, 446 insertions(+), 25 deletions(-) create mode 100644 junifer/data/tests/test_atlas.py rename junifer/datagrabber/tests/{test_base.py => test_base_datagrabber.py} (80%) create mode 100644 junifer/markers/tests/test_base_marker.py create mode 100644 junifer/storage/base.py create mode 100644 junifer/storage/sqlite.py diff --git a/examples/run_compute_parcel_mean.py b/examples/run_compute_parcel_mean.py index 4c842129d..4383b96c7 100644 --- a/examples/run_compute_parcel_mean.py +++ b/examples/run_compute_parcel_mean.py @@ -56,4 +56,4 @@ print(out.keys()) print(out['VBM_GM']['data'].shape) # Shape is (1 x parcels) print(out.keys()) -print(out['BOLD']['data'].shape) # Shape is (timepoints x parcels) \ No newline at end of file +print(out['BOLD']['data'].shape) # Shape is (timepoints x parcels) diff --git a/junifer/configs/juseless.py b/junifer/configs/juseless.py index d54e1ee68..7bebc030b 100644 --- a/junifer/configs/juseless.py +++ b/junifer/configs/juseless.py @@ -62,4 +62,5 @@ class JuselessUKBVBM(DataladDataGrabber): out['VBM_GM'] = self.datadir / f'm0wp1{sub}_{ses}_T1w.nii.gz' self._dataset_get(out) + out['meta']['element'] = dict(subject=sub, session=ses) return out diff --git a/junifer/data/atlases.py b/junifer/data/atlases.py index 395517f50..ee10cdd39 100644 --- a/junifer/data/atlases.py +++ b/junifer/data/atlases.py @@ -26,7 +26,7 @@ Optional keys: _available_atlases = { 'SUITxSUIT': { 'family': 'SUIT', - 'sace': 'SUIT' + 'space': 'SUIT' }, 'SUITxMNI': { 'family': 'SUIT', @@ -76,6 +76,8 @@ def register_atlas(name, atlas_path, atl_labels, overwrite=False): raise_error( f'Atlas {name} already registered. Set `overwrite=True` to ' 'update its value.') + if not isinstance(atlas_path, Path): + atlas_path = Path(atlas_path) _available_atlases[name] = { 'path': atlas_path, 'labels': atl_labels, 'family': 'CustomUserAtlas'} @@ -161,7 +163,8 @@ def load_atlas(name, atlas_dir=None, resolution=None, path_only=False): else: # retrieve atlases by passing arguments on to _retrieve_atlas() atlas_fname, atlas_labels = _retrieve_atlas( - t_family, atlas_dir=atlas_dir, **atlas_definition) + t_family, resolution=resolution, atlas_dir=atlas_dir, + **atlas_definition) logger.info( f'Loading atlas {atlas_fname.as_posix()}') # type: ignore @@ -184,7 +187,7 @@ def _retrieve_atlas(family, atlas_dir=None, resolution=None, **kwargs): --------------------- family : str Specify by name of atlas family, e.g. 'Schaefer'. - atlas_dir: path + atlas_dir: str or Path Path to where to store the retrieved atlas file. Defaults to: $HOME/junifer/data/atlas resolution : int @@ -218,6 +221,8 @@ def _retrieve_atlas(family, atlas_dir=None, resolution=None, **kwargs): if atlas_dir is None: atlas_dir = Path().home() / 'junifer' / 'data' / 'atlas' atlas_dir.mkdir(exist_ok=True, parents=True) + elif not isinstance(atlas_dir, Path): + atlas_dir = Path(atlas_dir) logger.info(f"Fetching one of {family} atlas.") @@ -329,12 +334,13 @@ def _retrieve_suit(atlas_path, resolution, space='MNI'): # check existance 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.') url_basis = ( - 'https://github.com/DiedrichsenLab/cerebellar_atlases/blob' + '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' diff --git a/junifer/data/tests/test_atlas.py b/junifer/data/tests/test_atlas.py new file mode 100644 index 000000000..308739642 --- /dev/null +++ b/junifer/data/tests/test_atlas.py @@ -0,0 +1,104 @@ +import tempfile +import pytest +from numpy.testing import assert_array_equal + +from junifer.data.atlases import register_atlas, list_atlases, load_atlas + + +def test_register_atlas(): + """Test register_atlas""" + + 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) + + _, lbl, fname = load_atlas('testatlas', path_only=True) + + assert lbl == ['1', '2', '6'] + assert fname.name == 'testatlas_2.nii.gz' # type: ignore + + +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_1mm.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]) + + 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]) + + +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]) diff --git a/junifer/datagrabber/base.py b/junifer/datagrabber/base.py index 8c828c875..4577aa443 100644 --- a/junifer/datagrabber/base.py +++ b/junifer/datagrabber/base.py @@ -75,6 +75,14 @@ class BaseDataGrabber(ABC): self._datadir = datadir self.types = types + def get_meta(self): + t_meta = {} + t_meta['class'] = self.__class__.__name__ + for k, v in vars(self).items(): + if not k.startswith('_'): + t_meta[k] = v + return t_meta + @property def datadir(self): """ @@ -185,7 +193,9 @@ class BIDSDataGrabber(BaseDataGrabber): t_replace = t_pattern.replace('{subject}', element) t_out = self.datadir / element / t_replace out[t_type] = dict(path=t_out) - + # Meta here is element and types + out['meta'] = dict(datagrabber=self.get_meta()) + out['meta']['element'] = element return out @@ -266,14 +276,25 @@ class DataladDataGrabber(BaseDataGrabber): def _dataset_get(self, out): for _, v in out.items(): - self.dataset.get(v['path']) + if 'path' in v: + self.dataset.get(v['path']) + + # append the version of the dataset + out['meta']['datagrabber']['dataset_commit_id'] = \ + self.dataset.repo.get_hexsha( + self.dataset.repo.get_corresponding_branch()) + return out def __getitem__(self, element): """Index one element in the Datalad database. It will first obtain the paths from the parent class and then `datalad get` each of the - files.""" + files. + + This method only works with multiple inheritance. + + """ out = super().__getitem__(element) - self._dataset_get(out) + out = self._dataset_get(out) return out diff --git a/junifer/datagrabber/tests/test_base.py b/junifer/datagrabber/tests/test_base_datagrabber.py similarity index 80% rename from junifer/datagrabber/tests/test_base.py rename to junifer/datagrabber/tests/test_base_datagrabber.py index 9d80054e1..3f0b0daa0 100644 --- a/junifer/datagrabber/tests/test_base.py +++ b/junifer/datagrabber/tests/test_base_datagrabber.py @@ -6,6 +6,14 @@ from pathlib import Path from junifer.datagrabber.base import BIDSDataGrabber, BIDSDataladDataGrabber +_testing_dataset = { + 'example_bids': { + 'uri': 'https://gin.g-node.org/juaml/datalad-example-bids', + 'id': 'e2ce149bd723088769a86c72e57eded009258c6b' + } +} + + def test_BIDSDataGrabber(): """Test BIDSDataGrabber""" with pytest.raises(TypeError, match=r"types must be a list"): @@ -54,8 +62,9 @@ def test_BIDSDataladDataGrabber(): with pytest.raises(ValueError, match=r"uri must be provided"): BIDSDataladDataGrabber(datadir=None, types=types, patterns=patterns) - repo_uri = 'https://gin.g-node.org/juaml/datalad-example-bids' + repo_uri = _testing_dataset['example_bids']['uri'] rootdir = 'example_bids' + repo_commit = _testing_dataset['example_bids']['id'] with BIDSDataladDataGrabber(rootdir=rootdir, uri=repo_uri, types=types, patterns=patterns) as dg: @@ -72,5 +81,15 @@ def test_BIDSDataladDataGrabber(): assert t_sub['bold']['path'] == \ (dg.datadir / f'{elem}/func/{elem}_task-rest_bold.nii.gz') + assert 'meta' in t_sub + assert 'datagrabber' in t_sub['meta'] + dg_meta = t_sub['meta']['datagrabber'] + assert 'class' in dg_meta + assert dg_meta['class'] == 'BIDSDataladDataGrabber' + assert 'uri' in dg_meta + assert dg_meta['uri'] == repo_uri + assert 'dataset_commit_id' in dg_meta + assert dg_meta['dataset_commit_id'] == repo_commit + with open(t_sub['T1w']['path'], 'r') as f: assert f.readlines()[0] == 'placeholder' diff --git a/junifer/datareader/default.py b/junifer/datareader/default.py index d06494d63..80cc11e15 100644 --- a/junifer/datareader/default.py +++ b/junifer/datareader/default.py @@ -63,4 +63,7 @@ class DefaultDataReader(PipelineStepMixin): logger.info( f'Unknown file type {t_path.as_posix()}, skipping reading') out[kind]['data'] = fread + if 'meta' not in out: + out['meta'] = {} + out['meta']['datareader'] = self.get_meta() return out diff --git a/junifer/datareader/tests/test_default_reader.py b/junifer/datareader/tests/test_default_reader.py index e1781d643..10ca74fdd 100644 --- a/junifer/datareader/tests/test_default_reader.py +++ b/junifer/datareader/tests/test_default_reader.py @@ -29,6 +29,22 @@ def test_validation(): assert reader.validate(t_kind) == t_kind +def test_meta(): + """Test reader metadata""" + reader = DefaultDataReader() + t_meta = reader.get_meta() + assert t_meta['class'] == 'DefaultDataReader' + + nib_data_path = Path(nib_testing.data_path) + t_path = nib_data_path / 'example4d.nii.gz' + input = {'bold': t_path} + output = reader.fit_transform(input) + assert 'meta' in output + assert 'datareader' in output['meta'] + assert 'class' in output['meta']['datareader'] + assert output['meta']['datareader']['class'] == 'DefaultDataReader' + + def test_read_nifti(): """Test reading NIFTI files""" reader = DefaultDataReader() diff --git a/junifer/markers/base.py b/junifer/markers/base.py index 6531ba039..e06cd302b 100644 --- a/junifer/markers/base.py +++ b/junifer/markers/base.py @@ -5,9 +5,13 @@ from ..utils import logger class PipelineStepMixin(): - @property - def name(self): - return self.__class__.__name__ + def get_meta(self): + t_meta = {} + t_meta['class'] = self.__class__.__name__ + for k, v in vars(self).items(): + if not k.startswith('_'): + t_meta[k] = v + return t_meta def validate_input(self, input): """Validate the input to the pipeline step. @@ -67,21 +71,23 @@ class PipelineStepMixin(): class BaseMarker(PipelineStepMixin): """Base class for all markers.""" - def __init__(self, on): + def __init__(self, on, name=None): if not isinstance(on, list): on = [on] self._valid_inputs = on + self.name = self.__class__.__name__ if name is None else name def get_meta(self): - t_meta = {} - t_meta['class'] = self.__class__.__name__ - for k, v in vars(self).items(): - if not k.startswith('_'): - t_meta[k] = v - return t_meta + s_meta = super().get_meta() + s_meta['name'] = self.name + return dict(marker=s_meta) def validate_input(self, input): - return any(x in input for x in self._valid_inputs) + if not any(x in input for x in self._valid_inputs): + raise ValueError( + 'Input does not have the required data.' + f'\t Input: {input}' + f'\t Required (any of): {self._valid_inputs}') def get_output_kind(self, input): return None @@ -102,7 +108,7 @@ class BaseMarker(PipelineStepMixin): t_out = self.compute(t_input) t_out.update(meta=t_meta) if storage is not None: - storage.store_2d(t_out) + storage.store_2d(**t_out) else: out[kind] = t_out diff --git a/junifer/markers/parcel.py b/junifer/markers/parcel.py index c370d644c..df46c1a21 100644 --- a/junifer/markers/parcel.py +++ b/junifer/markers/parcel.py @@ -1,3 +1,5 @@ +# Authors: Federico Raimondo +# License: AGPL import numpy as np from nilearn.maskers import NiftiMasker @@ -55,5 +57,7 @@ class ParcelAggregation(BaseMarker): out_labels.append(t_labels[t_v - 1]) out_values = np.array(out_values).transpose() - out = dict(data=out_values, labels=out_labels) + out = dict(data=out_values, columns=out_labels) + if out_values.shape[0] > 1: + out['row_names'] = 'scan' # type: ignore return out diff --git a/junifer/markers/tests/test_base_marker.py b/junifer/markers/tests/test_base_marker.py new file mode 100644 index 000000000..2b0e6fd5e --- /dev/null +++ b/junifer/markers/tests/test_base_marker.py @@ -0,0 +1,44 @@ +import pytest +from junifer.markers.base import BaseMarker, PipelineStepMixin + + +def test_meta(): + """Test metadata""" + pipemixin = PipelineStepMixin() + t_meta = pipemixin.get_meta() + assert t_meta['class'] == 'PipelineStepMixin' + + base = BaseMarker(on=['bold', 'dwi']) + + t_meta = base.get_meta() + assert t_meta['marker']['class'] == 'BaseMarker' + assert t_meta['marker']['name'] == 'BaseMarker' + + base = BaseMarker(on=['bold', 'dwi'], name='mymarker') + + t_meta = base.get_meta() + assert t_meta['marker']['name'] == 'mymarker' + + +def test_base(): + """Test base class""" + base = BaseMarker(on=['bold', 'dwi'], name='mymarker') + input = {'bold': {'path': 'test'}, 't2': {'path': 'test'}} + base.validate_input(input) + + wrong_input = {'t2': {'path': 'test'}} + with pytest.raises(ValueError): + base.validate_input(wrong_input) + + output = base.get_output_kind(input) + assert output is None + + with pytest.raises(NotImplementedError): + base.fit_transform(input) + + base.compute = lambda x: dict(data=1) # type: ignore + + out = base.fit_transform(input) + assert out['bold']['data'] == 1 + assert out['bold']['meta']['marker']['name'] == 'mymarker' + assert out['bold']['meta']['marker']['class'] == 'BaseMarker' diff --git a/junifer/storage/__init__.py b/junifer/storage/__init__.py index 4e235bc8b..0f49b9785 100644 --- a/junifer/storage/__init__.py +++ b/junifer/storage/__init__.py @@ -1,3 +1,2 @@ # Authors: Federico Raimondo -# Leonard Sasse # License: AGPL \ No newline at end of file diff --git a/junifer/storage/base.py b/junifer/storage/base.py new file mode 100644 index 000000000..96d8dd8ce --- /dev/null +++ b/junifer/storage/base.py @@ -0,0 +1,67 @@ +# Authors: Federico Raimondo +# License: AGPL +import numpy as np +import pandas as pd + +from .. import __version__ + + +def _element_to_index(meta, n_rows=1, row_names=None): + """Convert the element meta to index + + Parameters + ---------- + meta: dict + The metadata. Must contain the key 'element' + n_rows: int + Number of rows to create + row_names: list + The row names to use ins case n_rows > 1 + + Returns + ------- + index: pd.MultiIndex + The index of the dataframe to store + """ + element = meta['element'] + if not isinstance(element, dict): + element = dict(element=element) + if n_rows > 1: + elem_idx = { + k: v * n_rows for k, v in element.items() + } + elem_idx[row_names] = np.arange(n_rows) + else: + elem_idx = element + index = pd.MultiIndex.from_frame( + pd.DataFrame(elem_idx, index=range(n_rows))) + + return index + + +class BaseFeatureStorage(): + """ + Base class for feature storage. + """ + + def __init__(self, uri): + self.uri = uri + + def get_meta(self): + meta = {} + meta['versions'] = { + 'junifer': __version__, + } + return meta + + def store_matrix2d(self, data, col_names=None, row_names=None, meta=None): + raise NotImplementedError('store_matrix2d not implemented') + + def store_table(self, data, columns=None, row_names=None, meta=None): + raise NotImplementedError('store_table not implemented') + + def store_df(self, df): + raise NotImplementedError('store_df not implemented') + + def store_timeseries(self, data): + raise NotImplementedError('store_timeseries not implemented') diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py new file mode 100644 index 000000000..954b34243 --- /dev/null +++ b/junifer/storage/sqlite.py @@ -0,0 +1,130 @@ +# Authors: Federico Raimondo +# License: AGPL +import pandas as pd +from pandas.core.base import NoNewAttributesMixin +from pandas.io.sql import pandasSQL_builder +from sqlalchemy import create_engine, inspect +import hashlib +import json +from .base import BaseFeatureStorage, _element_to_index + + +class SQLiteFeatureStorage(BaseFeatureStorage): + """ + SQLite feature storage. + """ + + def __init__(self, uri): + super().__init__(uri) + self._engine = create_engine(uri, echo=False) + + def store_metadata(self, meta): + meta_md5 = self._meta_hash(meta) + if meta_md5 not in inspect(self._engine).get_table_names(): + meta_df = self._meta_row(meta, meta_md5) + self._save_upsert(meta_df, 'meta') + return meta_md5 + + def store_matrix2d(self, data, col_names=None, row_names=None, meta=None): + # Same as store_2d, but order is important + raise NotImplementedError('store_matrix2d not implemented') + + def store_2d(self, data, meta, columns=None, row_names=None): + idx = _element_to_index( + meta, n_rows=data.shape[0], row_names=row_names) + data_df = pd.DataFrame(data, columns=columns, index=idx) + self.store_df(data_df, meta) + + def store_df(self, df, meta): + table_name = self.store_metadata(meta) + self._save_upsert(df, table_name) + + def store_timeseries(self, data): + raise NotImplementedError('store_timeseries not implemented') + + def _meta_hash(self, meta): + t_meta = meta.copy() + t_meta.update(self.get_meta()) + meta_md5 = hashlib.md5( + json.dumps(meta, sort_keys=True).encode('utf-8')).hexdigest() + return meta_md5 + + def _meta_row(self, meta, meta_md5): + data_df = {} + for k, v in meta: + data_df[k] = json.dumps(v, sort_keys=True) + data_df['name'] = meta['marker']['name'] + df = pd.DataFrame(data_df, index=[meta_md5]) + return df + + def _save_upsert(self, df, name, upsert='ignore', if_exist='append'): + if upsert not in ['delete', 'ignore']: + raise ValueError('upsert must be either "delete" or "ignore"') + + index_col = df.index.names + with self._engine.begin() as con: + if if_exist == 'replace': + # Case 1: replace all the existing elements + df.to_sql(name, con=con, if_exists='replace') + elif not inspect(self._engine).has_table(name): + # Case 2: new table, so no big issue + df.to_sql(name, con=con, if_exists='append') + else: + # Case 3: existing table, so we need to check if the index + # is present or not. + + # Step 1: split incoming data into existing and new data + pk_indb = _get_existing_pk( + con, table_name=name, index_col=index_col) + existing, new = _split_incoming_data(df, pk_indb, index_col) + + # Step 2: upsert existing data + pandas_sql = pandasSQL_builder(con) + pandas_sql.meta.reflect(only=[name]) + table = pandas_sql.get_table(name) + update_stmts = NoNewAttributesMixin + if upsert == 'delete': + update_stmts = _generate_update_statements( + table, index_col, existing) + for stmt in update_stmts: + con.execute(stmt) + + # Step 3: insert new data + new.to_sql(name, con=con, if_exists='append') + + +def _get_existing_pk(con, table_name, index_col): + pk_cols = ','.join(index_col) + query = f'SELECT {pk_cols} FROM {table_name};' + pk_indb = pd.read_sql(query, con=con) + return pk_indb + + +def _split_incoming_data(df, pk_indb, index_col): + incoming_pk = df.reset_index()[index_col] + exists_mask = ( + incoming_pk[index_col] + .apply(tuple, axis=1) + .isin(pk_indb[index_col].apply(tuple, axis=1)) + ) + existing, new = df.loc[exists_mask.values], df.loc[~exists_mask.values] + return existing, new + + +def _generate_update_statements(table, index_col, rows_to_update): + from sqlalchemy import and_ + + new_records = rows_to_update.to_dict(orient="records") + pk_indb = rows_to_update.reset_index()[index_col] + pk_cols = [table.c[key] for key in index_col] + + stmts = [] + for i, (_, keys) in enumerate(pk_indb.iterrows()): + stmt = ( + table.update() + .where(and_(col == keys[j] + for j, col in enumerate(pk_cols))) # type: ignore + .values(new_records[i]) + ) + stmts.append(stmt) + return stmts diff --git a/requirements.txt b/requirements.txt index ea87f8453..eca50d3f0 100644 --- a/requirements.txt +++ b/requirements.txt @@ -2,4 +2,5 @@ numpy>=1.20, <1.22 datalad>=0.15.4, <0.16 pandas>=0.18.0, <1.5 nibabel>=3.2.0, <4.0 -nilearn>=0.9.0, <1.0 \ No newline at end of file +nilearn>=0.9.0, <1.0 +sqlalchemy>=1.4.27, <= 1.5.0 \ No newline at end of file -- 2.52.0 From 836308d2e467f758dbc23cc1af978a65e1e30c32 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 4 Mar 2022 11:58:11 +0100 Subject: [PATCH 004/287] Update CI --- .github/workflows/ci.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 344ac188a..ab498cdc7 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -8,7 +8,7 @@ jobs: runs-on: ubuntu-latest strategy: matrix: - python-version: [3.6, 3.7, 3.8] + python-version: [3.7, 3.8, 3.9] steps: - uses: actions/checkout@v2 -- 2.52.0 From a62a86a524f4ce82877e22970034f9fc6fd79491 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 4 Mar 2022 12:04:19 +0100 Subject: [PATCH 005/287] Fixes --- .flake8 | 2 +- examples/run_compute_parcel_mean.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/.flake8 b/.flake8 index d0ed0eec7..e57304ad3 100644 --- a/.flake8 +++ b/.flake8 @@ -1,5 +1,5 @@ [flake8] -exclude = __init__.py,*externals*,constants.py,fixes.py,resources.py,nilearn_cache,venv,docs/auto_examples,docs/_build/,.eggs/ +exclude = __init__.py,*externals*,constants.py,fixes.py,resources.py,nilearn_cache,venv,docs/auto_examples,docs/_build/,.eggs/,scratch/ ignore = W503,W504,I100,I101,I201,N806,E201,E202,E221,E222,E241,F541 # We add A for the array-spacing plugin, and ignore the E ones it covers above select = A,E,F,W,C \ No newline at end of file diff --git a/examples/run_compute_parcel_mean.py b/examples/run_compute_parcel_mean.py index 4383b96c7..b5c7ddc41 100644 --- a/examples/run_compute_parcel_mean.py +++ b/examples/run_compute_parcel_mean.py @@ -41,7 +41,7 @@ marker = ParcelAggregation(atlas='Schaefer100x7', method='mean') ############################################################################### # Prepare the input input = { - 'BOLD': {'data': fmri_img}, + 'BOLD': {'data': fmri_img}, 'VBM_GM': {'data': vbm_img} } -- 2.52.0 From 93321e68a6945f149d8cd1a9bc206b02054103ee Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 4 Mar 2022 12:16:09 +0100 Subject: [PATCH 006/287] Tyos --- junifer/data/atlases.py | 4 ++-- junifer/datareader/tests/test_default_reader.py | 2 +- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/junifer/data/atlases.py b/junifer/data/atlases.py index ee10cdd39..6cc04070e 100644 --- a/junifer/data/atlases.py +++ b/junifer/data/atlases.py @@ -285,7 +285,7 @@ def _retrieve_schaefer(atlas_dir, resolution, n_rois=None, yeo_networks=7): atlas_lname = atlas_dir / 'schaefer_2018' / ( f'Schaefer2018_{n_rois}Parcels_{yeo_networks}Networks_order.txt') - # check existance of atlas + # check existence of atlas if not (atlas_fname.exists() and atlas_lname.exists()): logger.info( 'At least one of the atlas files is missing. ' @@ -332,7 +332,7 @@ def _retrieve_suit(atlas_path, resolution, space='MNI'): atlas_lname = atlas_path / 'SUIT' / ( f'SUIT_{space}Space_{resolution}mm.tsv') - # check existance of atlas + # check existence of atlas if not (atlas_fname.exists() and atlas_lname.exists()): atlas_fname.parent.mkdir(exist_ok=True, parents=True) logger.info( diff --git a/junifer/datareader/tests/test_default_reader.py b/junifer/datareader/tests/test_default_reader.py index 10ca74fdd..d3661e9a7 100644 --- a/junifer/datareader/tests/test_default_reader.py +++ b/junifer/datareader/tests/test_default_reader.py @@ -75,7 +75,7 @@ def test_read_unknown(): nib_data_path = Path(nib_testing.data_path) anat_path = nib_data_path / 'reoriented_anat_moved.nii' - whatever_path = nib_data_path / 'unexistant.unkwnownextension' + whatever_path = nib_data_path / 'unexistent.unkwnownextension' input = {'anat': anat_path, 'whatever': whatever_path} output = reader.fit_transform(input) -- 2.52.0 From 10059b4a67cebe60ab2a27e0413aafdb36eaab54 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 4 Mar 2022 12:21:50 +0100 Subject: [PATCH 007/287] Fix test --- junifer/data/tests/test_atlas.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/junifer/data/tests/test_atlas.py b/junifer/data/tests/test_atlas.py index 308739642..046084b84 100644 --- a/junifer/data/tests/test_atlas.py +++ b/junifer/data/tests/test_atlas.py @@ -49,7 +49,7 @@ def test_schaefer_atlas(): with tempfile.TemporaryDirectory() as tmpdir: fname1 = 'Schaefer2018_100Parcels_7Networks_order_FSLMNI152_1mm.nii.gz' - fname2 = '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) -- 2.52.0 From 50e8f8dc1a24d81c6229960fa083aa30197b82cf Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Mon, 7 Mar 2022 09:06:39 +0100 Subject: [PATCH 008/287] A bit more of testing + doc --- docs/data.rst | 5 ++-- docs/index.rst | 1 + junifer/datagrabber/base.py | 4 ++-- .../tests/test_base_datagrabber.py | 23 ++++++++++++++++++- 4 files changed, 27 insertions(+), 6 deletions(-) diff --git a/docs/data.rst b/docs/data.rst index 080c5e4d2..c62a7dfc3 100644 --- a/docs/data.rst +++ b/docs/data.rst @@ -1,9 +1,8 @@ .. include:: links.inc -Data Object -=========== - +The Data Object +=============== Introduction ^^^^^^^^^^^^ diff --git a/docs/index.rst b/docs/index.rst index 1f7f0d2aa..eee62d926 100644 --- a/docs/index.rst +++ b/docs/index.rst @@ -9,6 +9,7 @@ Welcome to the documentation! :caption: Contents: installation + data api auto_examples/index.rst maintaining diff --git a/junifer/datagrabber/base.py b/junifer/datagrabber/base.py index 4577aa443..05be8d539 100644 --- a/junifer/datagrabber/base.py +++ b/junifer/datagrabber/base.py @@ -221,8 +221,8 @@ class DataladDataGrabber(BaseDataGrabber): Remove the datalad dataset from the datadir. This method is called automatically when the datagrabber is used within a `with` statement. - Note - ---- + Notes + ----- By itself, this class is still abstract as the `__getitem__` method relies on the parent class `BaseDataGrabber.__getitem__` which is not yet implemented. This class is intended to be used as a superclass of a class diff --git a/junifer/datagrabber/tests/test_base_datagrabber.py b/junifer/datagrabber/tests/test_base_datagrabber.py index 3f0b0daa0..3a2a6139a 100644 --- a/junifer/datagrabber/tests/test_base_datagrabber.py +++ b/junifer/datagrabber/tests/test_base_datagrabber.py @@ -1,9 +1,10 @@ # Authors: Federico Raimondo # Leonard Sasse # License: AGPL +from typing import Type import pytest from pathlib import Path -from junifer.datagrabber.base import BIDSDataGrabber, BIDSDataladDataGrabber +from junifer.datagrabber.base import BIDSDataGrabber, BIDSDataladDataGrabber, BaseDataGrabber _testing_dataset = { @@ -14,6 +15,26 @@ _testing_dataset = { } +def test_BaseDataGrabber(): + """Test BaseDataGrabber""" + with pytest.raises(TypeError, match=r"abstract"): + BaseDataGrabber(datadir='/tmp', types=['func']) # type: ignore + + class MyDataGrabber(BaseDataGrabber): + def __getitem__(self, element): + return super().__getitem__(element) + + def get_elements(self): + return super().get_elements() + + dg = MyDataGrabber(datadir='/tmp', types=['func']) + with pytest.raises(NotImplementedError): + dg['elem'] + + with pytest.raises(NotImplementedError): + dg.get_elements() + + def test_BIDSDataGrabber(): """Test BIDSDataGrabber""" with pytest.raises(TypeError, match=r"types must be a list"): -- 2.52.0 From 4c762fbabc2c21669bc374c09a1dff7a7f30c280 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Mon, 7 Mar 2022 09:33:06 +0100 Subject: [PATCH 009/287] Damn! --- junifer/datagrabber/tests/test_base_datagrabber.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/junifer/datagrabber/tests/test_base_datagrabber.py b/junifer/datagrabber/tests/test_base_datagrabber.py index 3a2a6139a..8f8452519 100644 --- a/junifer/datagrabber/tests/test_base_datagrabber.py +++ b/junifer/datagrabber/tests/test_base_datagrabber.py @@ -1,10 +1,10 @@ # Authors: Federico Raimondo # Leonard Sasse # License: AGPL -from typing import Type import pytest from pathlib import Path -from junifer.datagrabber.base import BIDSDataGrabber, BIDSDataladDataGrabber, BaseDataGrabber +from junifer.datagrabber.base import BIDSDataGrabber, BIDSDataladDataGrabber, \ + BaseDataGrabber _testing_dataset = { -- 2.52.0 From 151bc34980c3ceafd7367c8196741972d7bc2157 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Thu, 10 Mar 2022 11:58:32 +0100 Subject: [PATCH 010/287] ParcelAggregation marker --- .coveragerc | 1 + junifer/data/atlases.py | 22 +-- junifer/data/tests/test_atlas.py | 50 +++++- .../tests/test_base_datagrabber.py | 11 ++ .../datareader/tests/test_default_reader.py | 4 + junifer/markers/parcel.py | 4 +- junifer/markers/tests/test_base_marker.py | 17 +++ junifer/markers/tests/test_parcel.py | 143 ++++++++++++++++++ 8 files changed, 240 insertions(+), 12 deletions(-) create mode 100644 junifer/markers/tests/test_parcel.py diff --git a/.coveragerc b/.coveragerc index f7d2883cb..23019aac4 100644 --- a/.coveragerc +++ b/.coveragerc @@ -5,6 +5,7 @@ include = */junifer/* omit = */setup.py */tests/* + junifer/configs/juseless.py [report] exclude_lines = diff --git a/junifer/data/atlases.py b/junifer/data/atlases.py index 6cc04070e..dc603121d 100644 --- a/junifer/data/atlases.py +++ b/junifer/data/atlases.py @@ -95,12 +95,12 @@ def list_atlases(): return sorted(_available_atlases.keys()) -def _check_resolution(resolution, valid_resolution): - if resolution is None: - return None - if resolution not in valid_resolution: - raise ValueError(f'Invalid resolution: {resolution}') - return resolution +# def _check_resolution(resolution, valid_resolution): +# if resolution is None: +# return None +# if resolution not in valid_resolution: +# raise ValueError(f'Invalid resolution: {resolution}') +# return resolution def load_atlas(name, atlas_dir=None, resolution=None, path_only=False): @@ -153,7 +153,9 @@ def load_atlas(name, atlas_dir=None, resolution=None, path_only=False): atlas_fname : 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') @@ -296,7 +298,8 @@ def _retrieve_schaefer(atlas_dir, resolution, n_rois=None, yeo_networks=7): resolution_mm=resolution, data_dir=atlas_dir.as_posix()) - if not (atlas_fname.exists() and atlas_lname.exists()): + if not (atlas_fname.exists() and + atlas_lname.exists()): # pragma: no cover raise_error('There was a problem fetching the atlases.') # Load labels @@ -363,7 +366,8 @@ def _retrieve_suit(atlas_path, resolution, space='MNI'): sep='\t', usecols=['name']) labels.to_csv(atlas_lname, sep='\t', index=False) - if not (atlas_fname.exists() and atlas_lname.exists()): + 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( diff --git a/junifer/data/tests/test_atlas.py b/junifer/data/tests/test_atlas.py index 046084b84..ed973f0cc 100644 --- a/junifer/data/tests/test_atlas.py +++ b/junifer/data/tests/test_atlas.py @@ -1,8 +1,11 @@ import tempfile import pytest +from pathlib import Path from numpy.testing import assert_array_equal -from junifer.data.atlases import register_atlas, list_atlases, load_atlas +from junifer.data.atlases import (register_atlas, list_atlases, load_atlas, + _retrieve_schaefer, _retrieve_suit, + _retrieve_atlas) def test_register_atlas(): @@ -31,12 +34,24 @@ def test_register_atlas(): 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"rovided atlas name"): + _retrieve_atlas('wrongatlas') + + def test_schaefer_atlas(): """Test Schaefer atlas""" @@ -59,6 +74,9 @@ def test_schaefer_atlas(): 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 @@ -80,6 +98,26 @@ def test_schaefer_atlas(): 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""" @@ -102,3 +140,13 @@ def test_suit(): 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') diff --git a/junifer/datagrabber/tests/test_base_datagrabber.py b/junifer/datagrabber/tests/test_base_datagrabber.py index 8f8452519..a26dfc87e 100644 --- a/junifer/datagrabber/tests/test_base_datagrabber.py +++ b/junifer/datagrabber/tests/test_base_datagrabber.py @@ -1,6 +1,7 @@ # Authors: Federico Raimondo # Leonard Sasse # License: AGPL +import tempfile import pytest from pathlib import Path from junifer.datagrabber.base import BIDSDataGrabber, BIDSDataladDataGrabber, \ @@ -34,6 +35,10 @@ def test_BaseDataGrabber(): with pytest.raises(NotImplementedError): dg.get_elements() + with dg: + assert dg.datadir == Path('/tmp') + assert dg.types == ['func'] + def test_BIDSDataGrabber(): """Test BIDSDataGrabber""" @@ -114,3 +119,9 @@ def test_BIDSDataladDataGrabber(): with open(t_sub['T1w']['path'], 'r') as f: assert f.readlines()[0] == 'placeholder' + + with tempfile.TemporaryDirectory() as tmpdir: + with BIDSDataladDataGrabber(rootdir=rootdir, uri=repo_uri, + types=types, patterns=patterns, + datadir=tmpdir) as dg: + assert dg.datadir == Path(tmpdir) / rootdir diff --git a/junifer/datareader/tests/test_default_reader.py b/junifer/datareader/tests/test_default_reader.py index d3661e9a7..b1a110885 100644 --- a/junifer/datareader/tests/test_default_reader.py +++ b/junifer/datareader/tests/test_default_reader.py @@ -68,6 +68,10 @@ def test_read_nifti(): t_read_img = nib.load(t_path) assert_array_equal(read_img.get_fdata(), t_read_img.get_fdata()) + input = {'bold': t_path.as_posix()} + output2 = reader.fit_transform(input) + assert output['bold']['path'] == output2['bold']['path'] + def test_read_unknown(): """Test (not) reading unknown files""" diff --git a/junifer/markers/parcel.py b/junifer/markers/parcel.py index df46c1a21..8f60cf15d 100644 --- a/junifer/markers/parcel.py +++ b/junifer/markers/parcel.py @@ -11,10 +11,10 @@ from ..data import load_atlas class ParcelAggregation(BaseMarker): - def __init__(self, atlas, method, method_params=None, on=None): + def __init__(self, atlas, method, method_params=None, on=None, name=None): if on is None: on = ['T1w', 'BOLD', 'VBM_GM', 'VBM_WM'] - super().__init__(on=on) + super().__init__(on=on, name=name) self.atlas = atlas self.method = method self.method_params = {} if method_params is None else method_params diff --git a/junifer/markers/tests/test_base_marker.py b/junifer/markers/tests/test_base_marker.py index 2b0e6fd5e..5b8cc5805 100644 --- a/junifer/markers/tests/test_base_marker.py +++ b/junifer/markers/tests/test_base_marker.py @@ -2,6 +2,16 @@ import pytest from junifer.markers.base import BaseMarker, PipelineStepMixin +def test_pipelinestepmixing(): + mixin = PipelineStepMixin() + with pytest.raises(NotImplementedError): + mixin.validate_input(None) + with pytest.raises(NotImplementedError): + mixin.get_output_kind(None) + with pytest.raises(NotImplementedError): + mixin.fit_transform(None) + + def test_meta(): """Test metadata""" pipemixin = PipelineStepMixin() @@ -42,3 +52,10 @@ def test_base(): assert out['bold']['data'] == 1 assert out['bold']['meta']['marker']['name'] == 'mymarker' assert out['bold']['meta']['marker']['class'] == 'BaseMarker' + + base2 = BaseMarker(on='bold', name='mymarker') + base2.compute = lambda x: dict(data=1) # type: ignore + out2 = base2.fit_transform(input) + assert out2['bold']['data'] == 1 + assert out2['bold']['meta']['marker']['name'] == 'mymarker' + assert out2['bold']['meta']['marker']['class'] == 'BaseMarker' diff --git a/junifer/markers/tests/test_parcel.py b/junifer/markers/tests/test_parcel.py new file mode 100644 index 000000000..dcce038c2 --- /dev/null +++ b/junifer/markers/tests/test_parcel.py @@ -0,0 +1,143 @@ +import numpy as np +from numpy.testing import assert_array_equal +from scipy.stats import trim_mean +import nibabel as nib + +from nilearn import datasets +from nilearn.image import resample_to_img, math_img, concat_imgs +from nilearn.maskers import NiftiMasker, NiftiLabelsMasker + +from junifer.markers.parcel import ParcelAggregation + + +def test_parcelagregation_3D(): + """Test ParcelAggregation object on 3D images""" + + # Get the testing atlas (for nilearn) + atlas = datasets.fetch_atlas_schaefer_2018(n_rois=100) + + # Get the oasis VBM data: + oasis_dataset = datasets.fetch_oasis_vbm(n_subjects=1) + vbm = oasis_dataset.gray_matter_maps[0] + img = nib.load(vbm) + + # Mask atlas manually + atlas_img_res = resample_to_img( + atlas.maps, + img, + interpolation='nearest', + ) + atlas_bin = math_img( + 'img != 0', + img=atlas_img_res, + ) + + masker = NiftiMasker(atlas_bin, target_affine=img.affine) + + data = masker.fit_transform(img) + atlas_values = masker.transform(atlas_img_res) + atlas_values = np.squeeze(atlas_values).astype(int) + + # Compute the mean manually + manual = [] + for t_v in sorted(np.unique(atlas_values)): + t_values = np.mean(data[:, atlas_values == t_v]) + manual.append(t_values) + manual = np.array(manual)[np.newaxis, :] + + # Use the ParcelAggregation object + marker = ParcelAggregation( + atlas='Schaefer100x7', method='mean', name='gmd_schaefer100x7_mean') + input = dict(VBM_GM=dict(data=img)) + jun_values3d_mean = marker.fit_transform(input)['VBM_GM']['data'] + + assert jun_values3d_mean.ndim == 2 + assert jun_values3d_mean.shape[0] == 1 + assert_array_equal(manual, jun_values3d_mean) + + meta = marker.get_meta()['marker'] + assert meta['method'] == 'mean' + assert meta['atlas'] == 'Schaefer100x7' + assert meta['name'] == 'gmd_schaefer100x7_mean' + assert meta['class'] == 'ParcelAggregation' + assert meta['method_params'] == {} + + # Test using another function (std) + manual = [] + for t_v in sorted(np.unique(atlas_values)): + t_values = np.std(data[:, atlas_values == t_v]) + manual.append(t_values) + manual = np.array(manual)[np.newaxis, :] + + # Use the ParcelAggregation object + marker = ParcelAggregation(atlas='Schaefer100x7', method='std') + input = dict(VBM_GM=dict(data=img)) + jun_values3d_std = marker.fit_transform(input)['VBM_GM']['data'] + + assert jun_values3d_std.ndim == 2 + assert jun_values3d_std.shape[0] == 1 + assert_array_equal(manual, jun_values3d_std) + + meta = marker.get_meta()['marker'] + assert meta['method'] == 'std' + assert meta['atlas'] == 'Schaefer100x7' + assert meta['name'] == 'ParcelAggregation' + assert meta['class'] == 'ParcelAggregation' + assert meta['method_params'] == {} + + # Test using another function with parameters + manual = [] + for t_v in sorted(np.unique(atlas_values)): + t_values = trim_mean( + data[:, atlas_values == t_v], proportiontocut=0.1, axis=None) + manual.append(t_values) + manual = np.array(manual)[np.newaxis, :] + + # Use the ParcelAggregation object + marker = ParcelAggregation( + atlas='Schaefer100x7', method='trim_mean', + method_params={'proportiontocut': 0.1}) + input = dict(VBM_GM=dict(data=img)) + jun_values3d_tm = marker.fit_transform(input)['VBM_GM']['data'] + + assert jun_values3d_tm.ndim == 2 + assert jun_values3d_tm.shape[0] == 1 + assert_array_equal(manual, jun_values3d_tm) + + meta = marker.get_meta()['marker'] + assert meta['method'] == 'trim_mean' + assert meta['atlas'] == 'Schaefer100x7' + assert meta['name'] == 'ParcelAggregation' + assert meta['class'] == 'ParcelAggregation' + assert meta['method_params'] == {'proportiontocut': 0.1} + + +def test_parcelagregation_4D(): + """Test ParcelAggregation object on 3D images""" + + # Get the testing atlas (for nilearn) + atlas = datasets.fetch_atlas_schaefer_2018(n_rois=100) + + # Get the SPM auditory data: + subject_data = datasets.fetch_spm_auditory() + fmri_img = concat_imgs(subject_data.func) # type: ignore + + nifti_masker = NiftiLabelsMasker(labels_img=atlas.maps) + auto4d = nifti_masker.fit_transform(fmri_img) + + marker = ParcelAggregation(atlas='Schaefer100x7', method='mean') + input = dict(BOLD=dict(data=fmri_img)) + jun_values4d = marker.fit_transform(input)['BOLD']['data'] + + # TODO: check https://github.com/nilearn/nilearn/issues/2686 + # assert_array_equal(auto4d, jun_values4d) + + assert jun_values4d.ndim == 2 + assert_array_equal(auto4d.shape, jun_values4d.shape) + + meta = marker.get_meta()['marker'] + assert meta['method'] == 'mean' + assert meta['atlas'] == 'Schaefer100x7' + assert meta['name'] == 'ParcelAggregation' + assert meta['class'] == 'ParcelAggregation' + assert meta['method_params'] == {} -- 2.52.0 From 9a15e9e681df449504f7630eba0420fa951500c6 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Thu, 10 Mar 2022 12:29:22 +0100 Subject: [PATCH 011/287] Marker Collection WIP --- junifer/datagrabber/base.py | 3 ++ junifer/markers/__init__.py | 4 +- junifer/markers/collection.py | 14 +++++- junifer/markers/tests/test_collection.py | 61 ++++++++++++++++++++++++ junifer/testing.py | 28 +++++++++++ 5 files changed, 107 insertions(+), 3 deletions(-) create mode 100644 junifer/markers/tests/test_collection.py create mode 100644 junifer/testing.py diff --git a/junifer/datagrabber/base.py b/junifer/datagrabber/base.py index 05be8d539..8b8b15fa3 100644 --- a/junifer/datagrabber/base.py +++ b/junifer/datagrabber/base.py @@ -75,6 +75,9 @@ class BaseDataGrabber(ABC): self._datadir = datadir self.types = types + def get_types(self): + return self.types.copy() + def get_meta(self): t_meta = {} t_meta['class'] = self.__class__.__name__ diff --git a/junifer/markers/__init__.py b/junifer/markers/__init__.py index 4e235bc8b..5e67e4920 100644 --- a/junifer/markers/__init__.py +++ b/junifer/markers/__init__.py @@ -1,3 +1,5 @@ # Authors: Federico Raimondo # Leonard Sasse -# License: AGPL \ No newline at end of file +# License: AGPL +from .collection import MarkerCollection +from .parcel import ParcelAggregation \ No newline at end of file diff --git a/junifer/markers/collection.py b/junifer/markers/collection.py index 9d06aa0a5..6bb58c450 100644 --- a/junifer/markers/collection.py +++ b/junifer/markers/collection.py @@ -4,6 +4,8 @@ from ..utils import logger from ..datareader import DefaultDataReader +from collections import Counter + class MarkerCollection(): def __init__(self, markers, datareader=None, preprocessing=None, @@ -15,6 +17,14 @@ class MarkerCollection(): self._markers = markers self._storage = storage + # Check that the markers have different names + marker_names = [m.name for m in self._markers] + if len(set(marker_names)) != len(marker_names): + counts = Counter(marker_names) + raise ValueError( + 'Markers must have different names. ' + f'Current names are: {counts}') + def fit(self, input): """Fit the pipeline. @@ -26,7 +36,7 @@ class MarkerCollection(): Returns ------- - output : dict[str -> object] + output : dict[str -> object] | None The output of the pipeline. Each key represents a marker name and the values are the computer marker values. If the pipeline has a storage configured, then the output will be None. @@ -54,7 +64,7 @@ class MarkerCollection(): check that the storage can handle the markers output. """ logger.info('Validating Marker Collection') - t_data = datagrabber.get_output_kind() + t_data = datagrabber.get_types() logger.info(f'DataGrabber output type: {t_data}') logger.info(f'Validating Data Reader:') diff --git a/junifer/markers/tests/test_collection.py b/junifer/markers/tests/test_collection.py new file mode 100644 index 000000000..2600ff80c --- /dev/null +++ b/junifer/markers/tests/test_collection.py @@ -0,0 +1,61 @@ +# Authors: Federico Raimondo +# License: AGPL + +import pytest +from junifer.datareader.default import DefaultDataReader +from junifer.markers import MarkerCollection, ParcelAggregation +from junifer.testing import OasisVBMTestingDatagrabber + + +def test_markercollection(): + """Test MarkerCollection""" + wrong_markers = [ + ParcelAggregation( + atlas='Schaefer100x7', method='mean', + name='gmd_schaefer100x7_mean'), + ParcelAggregation( + atlas='Schaefer100x7', method='mean', + name='gmd_schaefer100x7_mean'), + ] + + with pytest.raises(ValueError, match=r"must have different names"): + MarkerCollection(wrong_markers) + + markers = [ + ParcelAggregation( + atlas='Schaefer100x7', method='mean', + name='gmd_schaefer100x7_mean'), + ParcelAggregation( + atlas='Schaefer100x7', method='std', + name='gmd_schaefer100x7_std'), + ParcelAggregation( + atlas='Schaefer100x7', method='trim_mean', + method_params={'proportiontocut': 0.1}, + name='gmd_schaefer100x7_trim_mean90') + ] + mc = MarkerCollection(markers=markers) + assert mc._markers == markers + assert mc._preprocessing is None + assert mc._storage is None + assert isinstance(mc._datareader, DefaultDataReader) + + dg = OasisVBMTestingDatagrabber() + mc.validate(dg) + + with dg: + input = dg[1] + out = mc.fit(input) + assert out is not None + assert isinstance(out, dict) + assert len(out) == 3 + assert 'gmd_schaefer100x7_mean' in out + assert 'gmd_schaefer100x7_std' in out + assert 'gmd_schaefer100x7_trim_mean90' in out + + for t_marker in markers: + t_name = t_marker.name + assert 'VBM_GM' in out[t_name] + t_vbm = out[t_name]['VBM_GM'] + assert 'data' in t_vbm + assert 'columns' in t_vbm + assert 'meta' in t_vbm diff --git a/junifer/testing.py b/junifer/testing.py new file mode 100644 index 000000000..60a298ec4 --- /dev/null +++ b/junifer/testing.py @@ -0,0 +1,28 @@ +# Authors: Federico Raimondo +# License: AGPL +import tempfile +from nilearn import datasets + +from .datagrabber.base import BaseDataGrabber + + +class OasisVBMTestingDatagrabber(BaseDataGrabber): + """ + DataGrabber for Oasis VBM testing data. + """ + def __init__(self): + datadir = tempfile.mkdtemp() + types = ['VBM_GM'] + super().__init__(types=types, datadir=datadir) + + def get_elements(self): + return list(range(1, 11)) + + def __getitem__(self, element): + out = {} + out['VBM_GM'] = self._dataset.gray_matter_maps[element] + return out + + def __enter__(self): + self._dataset = datasets.fetch_oasis_vbm(n_subjects=10) + return self -- 2.52.0 From c9c262d136e79becbb636bd5e21cbd8c5df5e83d Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Thu, 10 Mar 2022 12:31:29 +0100 Subject: [PATCH 012/287] fix test --- junifer/data/tests/test_atlas.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/junifer/data/tests/test_atlas.py b/junifer/data/tests/test_atlas.py index ed973f0cc..c02b1b42b 100644 --- a/junifer/data/tests/test_atlas.py +++ b/junifer/data/tests/test_atlas.py @@ -48,7 +48,7 @@ def test_wrong_atlas(): with pytest.raises(ValueError, match=r"not found"): load_atlas('wrongatlas') - with pytest.raises(ValueError, match=r"rovided atlas name"): + with pytest.raises(ValueError, match=r"provided atlas name"): _retrieve_atlas('wrongatlas') -- 2.52.0 From bb108b868c7340532e33204cc6ab42a0528b1b32 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Thu, 10 Mar 2022 14:08:48 +0100 Subject: [PATCH 013/287] Increase coverage --- .coveragerc | 1 + junifer/data/atlases.py | 2 +- junifer/markers/tests/test_collection.py | 19 +++++++++++++++++++ junifer/markers/tests/test_parcel.py | 3 ++- 4 files changed, 23 insertions(+), 2 deletions(-) diff --git a/.coveragerc b/.coveragerc index 23019aac4..6cdbc7a3c 100644 --- a/.coveragerc +++ b/.coveragerc @@ -6,6 +6,7 @@ omit = */setup.py */tests/* junifer/configs/juseless.py + junifer/testing.py [report] exclude_lines = diff --git a/junifer/data/atlases.py b/junifer/data/atlases.py index dc603121d..33cd4747f 100644 --- a/junifer/data/atlases.py +++ b/junifer/data/atlases.py @@ -354,7 +354,7 @@ def _retrieve_suit(atlas_path, resolution, space='MNI'): atlas_download = requests.get(url_MNI) with open(atlas_fname, 'wb') as f: f.write(atlas_download.content) - elif space == 'SUIT': + else: # if not MNI, then SUIT logger.info(f'Downloading {url_SUIT}') atlas_download = requests.get(url_SUIT) with open(atlas_fname, 'wb') as f: diff --git a/junifer/markers/tests/test_collection.py b/junifer/markers/tests/test_collection.py index 2600ff80c..28316cff5 100644 --- a/junifer/markers/tests/test_collection.py +++ b/junifer/markers/tests/test_collection.py @@ -2,8 +2,10 @@ # License: AGPL import pytest +from numpy.testing import assert_array_equal from junifer.datareader.default import DefaultDataReader from junifer.markers import MarkerCollection, ParcelAggregation +from junifer.markers.base import PipelineStepMixin from junifer.testing import OasisVBMTestingDatagrabber @@ -59,3 +61,20 @@ def test_markercollection(): assert 'data' in t_vbm assert 'columns' in t_vbm assert 'meta' in t_vbm + + # Test preprocessing + class BypassPreprocessing(PipelineStepMixin): + def fit_transform(self, input): + return input + + mc2 = MarkerCollection( + markers=markers, preprocessing=BypassPreprocessing(), + datareader=DefaultDataReader()) + assert isinstance(mc2._datareader, DefaultDataReader) + with dg: + input = dg[1] + out2 = mc2.fit(input) + for t_marker in markers: + t_name = t_marker.name + assert_array_equal(out[t_name]['VBM_GM']['data'], + out2[t_name]['VBM_GM']['data']) # type: ignore diff --git a/junifer/markers/tests/test_parcel.py b/junifer/markers/tests/test_parcel.py index dcce038c2..cf58f5107 100644 --- a/junifer/markers/tests/test_parcel.py +++ b/junifer/markers/tests/test_parcel.py @@ -47,7 +47,8 @@ def test_parcelagregation_3D(): # Use the ParcelAggregation object marker = ParcelAggregation( - atlas='Schaefer100x7', method='mean', name='gmd_schaefer100x7_mean') + atlas='Schaefer100x7', method='mean', name='gmd_schaefer100x7_mean', + on='VBM_GM') # Test passing "on" as a keyword argument input = dict(VBM_GM=dict(data=img)) jun_values3d_mean = marker.fit_transform(input)['VBM_GM']['data'] -- 2.52.0 From 60cdcf69614d28174e632416fe0c3a12477de2ee Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 11 Mar 2022 10:51:53 +0100 Subject: [PATCH 014/287] Fix Parcel test --- junifer/markers/parcel.py | 2 +- junifer/markers/tests/test_parcel.py | 16 ++++++++++------ 2 files changed, 11 insertions(+), 7 deletions(-) diff --git a/junifer/markers/parcel.py b/junifer/markers/parcel.py index 8f60cf15d..c277bcacc 100644 --- a/junifer/markers/parcel.py +++ b/junifer/markers/parcel.py @@ -56,7 +56,7 @@ class ParcelAggregation(BaseMarker): # in it out_labels.append(t_labels[t_v - 1]) - out_values = np.array(out_values).transpose() + out_values = np.array(out_values).T out = dict(data=out_values, columns=out_labels) if out_values.shape[0] > 1: out['row_names'] = 'scan' # type: ignore diff --git a/junifer/markers/tests/test_parcel.py b/junifer/markers/tests/test_parcel.py index cf58f5107..40974168f 100644 --- a/junifer/markers/tests/test_parcel.py +++ b/junifer/markers/tests/test_parcel.py @@ -1,5 +1,5 @@ import numpy as np -from numpy.testing import assert_array_equal +from numpy.testing import assert_array_equal, assert_array_almost_equal from scipy.stats import trim_mean import nibabel as nib @@ -45,6 +45,11 @@ def test_parcelagregation_3D(): manual.append(t_values) manual = np.array(manual)[np.newaxis, :] + nifti_masker = NiftiLabelsMasker(labels_img=atlas.maps) + auto = nifti_masker.fit_transform(img) + + assert_array_almost_equal(auto, manual) + # Use the ParcelAggregation object marker = ParcelAggregation( atlas='Schaefer100x7', method='mean', name='gmd_schaefer100x7_mean', @@ -114,10 +119,11 @@ def test_parcelagregation_3D(): def test_parcelagregation_4D(): - """Test ParcelAggregation object on 3D images""" + """Test ParcelAggregation object on 4D images""" # Get the testing atlas (for nilearn) - atlas = datasets.fetch_atlas_schaefer_2018(n_rois=100) + atlas = datasets.fetch_atlas_schaefer_2018( + n_rois=100, yeo_networks=7, resolution_mm=2) # Get the SPM auditory data: subject_data = datasets.fetch_spm_auditory() @@ -130,11 +136,9 @@ def test_parcelagregation_4D(): input = dict(BOLD=dict(data=fmri_img)) jun_values4d = marker.fit_transform(input)['BOLD']['data'] - # TODO: check https://github.com/nilearn/nilearn/issues/2686 - # assert_array_equal(auto4d, jun_values4d) - assert jun_values4d.ndim == 2 assert_array_equal(auto4d.shape, jun_values4d.shape) + assert_array_equal(auto4d, jun_values4d) meta = marker.get_meta()['marker'] assert meta['method'] == 'mean' -- 2.52.0 From e030b01493b7d31f6c3a7b6191f7ce6742dec533 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 11 Mar 2022 11:08:45 +0100 Subject: [PATCH 015/287] Fix test --- junifer/datagrabber/tests/test_base_datagrabber.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/junifer/datagrabber/tests/test_base_datagrabber.py b/junifer/datagrabber/tests/test_base_datagrabber.py index a26dfc87e..a31fa8677 100644 --- a/junifer/datagrabber/tests/test_base_datagrabber.py +++ b/junifer/datagrabber/tests/test_base_datagrabber.py @@ -121,7 +121,8 @@ def test_BIDSDataladDataGrabber(): assert f.readlines()[0] == 'placeholder' with tempfile.TemporaryDirectory() as tmpdir: + datadir = Path(tmpdir) / 'dataset' # Need this for testing with BIDSDataladDataGrabber(rootdir=rootdir, uri=repo_uri, types=types, patterns=patterns, - datadir=tmpdir) as dg: - assert dg.datadir == Path(tmpdir) / rootdir + datadir=datadir) as dg: + assert dg.datadir == datadir / rootdir -- 2.52.0 From 385da81eb40a709f58783617d829afe265cc625a Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 11 Mar 2022 11:26:47 +0100 Subject: [PATCH 016/287] add one more python version + rename --- .github/workflows/ci.yml | 2 +- junifer/markers/tests/test_base_marker.py | 4 ++-- junifer/markers/tests/test_collection.py | 2 +- junifer/markers/tests/test_parcel.py | 4 ++-- 4 files changed, 6 insertions(+), 6 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index ab498cdc7..0569d8fdd 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -8,7 +8,7 @@ jobs: runs-on: ubuntu-latest strategy: matrix: - python-version: [3.7, 3.8, 3.9] + python-version: [3.7, 3.8, 3.9, 3.10] steps: - uses: actions/checkout@v2 diff --git a/junifer/markers/tests/test_base_marker.py b/junifer/markers/tests/test_base_marker.py index 5b8cc5805..cf9ed5298 100644 --- a/junifer/markers/tests/test_base_marker.py +++ b/junifer/markers/tests/test_base_marker.py @@ -2,7 +2,7 @@ import pytest from junifer.markers.base import BaseMarker, PipelineStepMixin -def test_pipelinestepmixing(): +def test_PipelineStepMixin(): mixin = PipelineStepMixin() with pytest.raises(NotImplementedError): mixin.validate_input(None) @@ -30,7 +30,7 @@ def test_meta(): assert t_meta['marker']['name'] == 'mymarker' -def test_base(): +def test_BaseMarker(): """Test base class""" base = BaseMarker(on=['bold', 'dwi'], name='mymarker') input = {'bold': {'path': 'test'}, 't2': {'path': 'test'}} diff --git a/junifer/markers/tests/test_collection.py b/junifer/markers/tests/test_collection.py index 28316cff5..2cf8adf64 100644 --- a/junifer/markers/tests/test_collection.py +++ b/junifer/markers/tests/test_collection.py @@ -9,7 +9,7 @@ from junifer.markers.base import PipelineStepMixin from junifer.testing import OasisVBMTestingDatagrabber -def test_markercollection(): +def test_MarkerCollection(): """Test MarkerCollection""" wrong_markers = [ ParcelAggregation( diff --git a/junifer/markers/tests/test_parcel.py b/junifer/markers/tests/test_parcel.py index 40974168f..1de773bfc 100644 --- a/junifer/markers/tests/test_parcel.py +++ b/junifer/markers/tests/test_parcel.py @@ -10,7 +10,7 @@ from nilearn.maskers import NiftiMasker, NiftiLabelsMasker from junifer.markers.parcel import ParcelAggregation -def test_parcelagregation_3D(): +def test_ParcelAggregation_3D(): """Test ParcelAggregation object on 3D images""" # Get the testing atlas (for nilearn) @@ -118,7 +118,7 @@ def test_parcelagregation_3D(): assert meta['method_params'] == {'proportiontocut': 0.1} -def test_parcelagregation_4D(): +def test_ParcelAggregation_4D(): """Test ParcelAggregation object on 4D images""" # Get the testing atlas (for nilearn) -- 2.52.0 From 733818d465bcc4b3de54b6e4d559c9d3ebcc9f56 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 11 Mar 2022 11:28:23 +0100 Subject: [PATCH 017/287] Fix --- .github/workflows/ci.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 0569d8fdd..03f9e4a24 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -8,7 +8,7 @@ jobs: runs-on: ubuntu-latest strategy: matrix: - python-version: [3.7, 3.8, 3.9, 3.10] + python-version: ['3.7', '3.8', '3.9', '3.10'] steps: - uses: actions/checkout@v2 -- 2.52.0 From a2755be299270955151d54a087232059754297b2 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 11 Mar 2022 11:32:03 +0100 Subject: [PATCH 018/287] Run all tests if one python version fails --- .github/workflows/ci.yml | 1 + 1 file changed, 1 insertion(+) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 03f9e4a24..091cebeb9 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -7,6 +7,7 @@ jobs: runs-on: ubuntu-latest strategy: + fail-fast: false matrix: python-version: ['3.7', '3.8', '3.9', '3.10'] -- 2.52.0 From 49a1c3b01938bdfb2844503e65a33b8a921a3b89 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Mon, 14 Mar 2022 12:49:31 +0100 Subject: [PATCH 019/287] Mores tests --- junifer/storage/base.py | 59 ++++++++++-- junifer/storage/sqlite.py | 27 ++---- junifer/storage/tests/test_base.py | 145 +++++++++++++++++++++++++++++ 3 files changed, 205 insertions(+), 26 deletions(-) create mode 100644 junifer/storage/tests/test_base.py diff --git a/junifer/storage/base.py b/junifer/storage/base.py index 96d8dd8ce..1d14338b0 100644 --- a/junifer/storage/base.py +++ b/junifer/storage/base.py @@ -2,11 +2,35 @@ # License: AGPL import numpy as np import pandas as pd +import json +import hashlib +from abc import ABC, abstractmethod from .. import __version__ -def _element_to_index(meta, n_rows=1, row_names=None): +def meta_hash(meta): + """Compute the md5 hash of the meta + + Parameters + ---------- + meta: dict + The metadata. Must contain the key 'element' + + Returns + ------- + md5: str + The md5 hash of the meta + """ + if meta is None: + raise ValueError('Meta must be a dict (currently is None)') + t_meta = meta.copy() + meta_md5 = hashlib.md5( + json.dumps(t_meta, sort_keys=True).encode('utf-8')).hexdigest() + return meta_md5 + + +def element_to_index(meta, n_rows=1, rows_col_name=None): """Convert the element meta to index Parameters @@ -15,8 +39,9 @@ def _element_to_index(meta, n_rows=1, row_names=None): The metadata. Must contain the key 'element' n_rows: int Number of rows to create - row_names: list - The row names to use ins case n_rows > 1 + rows_col_name: str + The column name to use in case n_rows > 1. If None (default) and + n_rows > 1, the name will be 'index'. Returns ------- @@ -27,10 +52,12 @@ def _element_to_index(meta, n_rows=1, row_names=None): if not isinstance(element, dict): element = dict(element=element) if n_rows > 1: + if rows_col_name is None: + rows_col_name = 'index' elem_idx = { - k: v * n_rows for k, v in element.items() + k: [v] * n_rows for k, v in element.items() } - elem_idx[row_names] = np.arange(n_rows) + elem_idx[rows_col_name] = np.arange(n_rows) else: elem_idx = element index = pd.MultiIndex.from_frame( @@ -39,7 +66,7 @@ def _element_to_index(meta, n_rows=1, row_names=None): return index -class BaseFeatureStorage(): +class BaseFeatureStorage(ABC): """ Base class for feature storage. """ @@ -54,14 +81,34 @@ class BaseFeatureStorage(): } return meta + @abstractmethod + def store_metadata(self, meta): + raise NotImplementedError('store_metadata not implemented') + + @abstractmethod def store_matrix2d(self, data, col_names=None, row_names=None, meta=None): raise NotImplementedError('store_matrix2d not implemented') + @abstractmethod def store_table(self, data, columns=None, row_names=None, meta=None): raise NotImplementedError('store_table not implemented') + @abstractmethod def store_df(self, df): raise NotImplementedError('store_df not implemented') + @abstractmethod def store_timeseries(self, data): raise NotImplementedError('store_timeseries not implemented') + + +class PandasFeatureStoreage(BaseFeatureStorage): + + def _meta_row(self, meta, meta_md5): + """Converta the meta to a dataframe row""" + data_df = {} + for k, v in meta: + data_df[k] = json.dumps(v, sort_keys=True) + data_df['name'] = meta['marker']['name'] + df = pd.DataFrame(data_df, index=[meta_md5]) + return df diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py index 954b34243..5db52baee 100644 --- a/junifer/storage/sqlite.py +++ b/junifer/storage/sqlite.py @@ -6,10 +6,10 @@ from pandas.io.sql import pandasSQL_builder from sqlalchemy import create_engine, inspect import hashlib import json -from .base import BaseFeatureStorage, _element_to_index +from .base import PandasFeatureStoreage, element_to_index, meta_hash -class SQLiteFeatureStorage(BaseFeatureStorage): +class SQLiteFeatureStorage(PandasFeatureStoreage): """ SQLite feature storage. """ @@ -19,9 +19,11 @@ class SQLiteFeatureStorage(BaseFeatureStorage): self._engine = create_engine(uri, echo=False) def store_metadata(self, meta): - meta_md5 = self._meta_hash(meta) + t_meta = meta.copy() + t_meta.update(self.get_meta()) + meta_md5 = meta_hash(t_meta) if meta_md5 not in inspect(self._engine).get_table_names(): - meta_df = self._meta_row(meta, meta_md5) + meta_df = self._meta_row(t_meta, meta_md5) self._save_upsert(meta_df, 'meta') return meta_md5 @@ -30,7 +32,7 @@ class SQLiteFeatureStorage(BaseFeatureStorage): raise NotImplementedError('store_matrix2d not implemented') def store_2d(self, data, meta, columns=None, row_names=None): - idx = _element_to_index( + idx = element_to_index( meta, n_rows=data.shape[0], row_names=row_names) data_df = pd.DataFrame(data, columns=columns, index=idx) self.store_df(data_df, meta) @@ -42,21 +44,6 @@ class SQLiteFeatureStorage(BaseFeatureStorage): def store_timeseries(self, data): raise NotImplementedError('store_timeseries not implemented') - def _meta_hash(self, meta): - t_meta = meta.copy() - t_meta.update(self.get_meta()) - meta_md5 = hashlib.md5( - json.dumps(meta, sort_keys=True).encode('utf-8')).hexdigest() - return meta_md5 - - def _meta_row(self, meta, meta_md5): - data_df = {} - for k, v in meta: - data_df[k] = json.dumps(v, sort_keys=True) - data_df['name'] = meta['marker']['name'] - df = pd.DataFrame(data_df, index=[meta_md5]) - return df - def _save_upsert(self, df, name, upsert='ignore', if_exist='append'): if upsert not in ['delete', 'ignore']: raise ValueError('upsert must be either "delete" or "ignore"') diff --git a/junifer/storage/tests/test_base.py b/junifer/storage/tests/test_base.py new file mode 100644 index 000000000..56779d212 --- /dev/null +++ b/junifer/storage/tests/test_base.py @@ -0,0 +1,145 @@ +import pytest + +from junifer.storage.base import (meta_hash, element_to_index, + BaseFeatureStorage) + + +def test_meta_hash(): + """Test meta_hash""" + + meta = {} + hash = meta_hash(meta) + + assert hash == '99914b932bd37a50b983c5e7c90ae93b' # empty dict + + meta = None + with pytest.raises(ValueError, match=r"Meta must be a dict"): + meta_hash(meta) + + meta = {'element': 'foo', 'A': 1, 'B': [2, 3, 4, 5, 6]} + hash1 = meta_hash(meta) + + meta = {'element': 'foo', 'B': [2, 3, 4, 5, 6], 'A': 1} + hash2 = meta_hash(meta) + + assert hash1 == hash2 + + meta = {'element': 'foo', 'A': 1, 'B': [2, 3, 1, 5, 6]} + hash3 = meta_hash(meta) + assert hash1 != hash3 + + meta1 = { + 'element': 'foo', + 'B': { + 'B2': [2, 3, 4, 5, 6], + 'B1': [9.22, 3.14, 1.41, 5.67, 6.28], + 'B3': (1, 'car'), + }, + 'A': 1} + + meta2 = { + 'A': 1, + 'B': { + 'B3': (1, 'car'), + 'B1': [9.22, 3.14, 1.41, 5.67, 6.28], + 'B2': [2, 3, 4, 5, 6], + }, + 'element': 'foo' + } + + hash4 = meta_hash(meta1) + hash5 = meta_hash(meta2) + assert hash4 == hash5 + + +def test_element_to_index(): + """Test element_to_index""" + + meta = {'element': 'foo', 'A': 1, 'B': [2, 3, 4, 5, 6]} + index = element_to_index(meta) + assert index.names == ['element'] + assert index.levels[0].name == 'element' + assert index.levels[0].values[0] == 'foo' + + index = element_to_index(meta, n_rows=10) + assert index.names == ['element', 'index'] + assert index.levels[0].name == 'element' + assert all(x == 'foo' for x in index.levels[0].values) + assert index.levels[0].values.shape == (1,) + + assert index.levels[1].name == 'index' + assert all(x == i for i, x in enumerate(index.levels[1].values)) + assert index.levels[1].values.shape == (10,) + + index = element_to_index(meta, n_rows=1, rows_col_name='scan') + assert index.names == ['element'] + assert index.levels[0].name == 'element' + assert all(x == 'foo' for x in index.levels[0].values) + assert index.levels[0].values.shape == (1,) + + index = element_to_index(meta, n_rows=7, rows_col_name='scan') + assert index.names == ['element', 'scan'] + assert index.levels[0].name == 'element' + assert all(x == 'foo' for x in index.levels[0].values) + assert index.levels[0].values.shape == (1,) + + assert index.levels[1].name == 'scan' + assert all(x == i for i, x in enumerate(index.levels[1].values)) + assert index.levels[1].values.shape == (7,) + + meta = { + 'element': {'subject': 'sub-01', 'session': 'ses-01'}, + 'A': 1, 'B': [2, 3, 4, 5, 6]} + index = element_to_index(meta, n_rows=10) + + assert index.levels[0].name == 'subject' + assert all(x == 'sub-01' for x in index.levels[0].values) + assert index.levels[0].values.shape == (1,) + + assert index.levels[1].name == 'session' + assert all(x == 'ses-01' for x in index.levels[1].values) + assert index.levels[1].values.shape == (1,) + + assert index.levels[2].name == 'index' + assert all(x == i for i, x in enumerate(index.levels[2].values)) + assert index.levels[2].values.shape == (10,) + + +def test_BaseFeatureStorage(): + """Test BaseFeatureStorage""" + with pytest.raises(TypeError, match=r"abstract"): + BaseFeatureStorage(uri='/tmp') # type: ignore + + class MyFeatureStorage(BaseFeatureStorage): + def store_metadata(self, metadata): + super().store_metadata(metadata) + + def store_matrix2d(self, matrix): + super().store_matrix2d(matrix) + + def store_table(self, table): + super().store_table(table) + + def store_df(self, df): + super().store_df(df) + + def store_timeseries(self, timeseries): + super().store_timeseries(timeseries) + + st = MyFeatureStorage(uri='/tmp') + with pytest.raises(NotImplementedError): + st.store_metadata(None) + + with pytest.raises(NotImplementedError): + st.store_matrix2d(None) + + with pytest.raises(NotImplementedError): + st.store_table(None) + + with pytest.raises(NotImplementedError): + st.store_df(None) + + with pytest.raises(NotImplementedError): + st.store_timeseries(None) + + assert st.uri == '/tmp' -- 2.52.0 From 188a8a87a289af9bf03cf887700ac3bef5816fb0 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Mon, 14 Mar 2022 12:53:44 +0100 Subject: [PATCH 020/287] flake! --- junifer/storage/sqlite.py | 10 +++++----- junifer/storage/tests/test_base.py | 2 +- 2 files changed, 6 insertions(+), 6 deletions(-) diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py index 5db52baee..1e62f9ca8 100644 --- a/junifer/storage/sqlite.py +++ b/junifer/storage/sqlite.py @@ -4,8 +4,7 @@ import pandas as pd from pandas.core.base import NoNewAttributesMixin from pandas.io.sql import pandasSQL_builder from sqlalchemy import create_engine, inspect -import hashlib -import json + from .base import PandasFeatureStoreage, element_to_index, meta_hash @@ -27,13 +26,14 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): self._save_upsert(meta_df, 'meta') return meta_md5 - def store_matrix2d(self, data, col_names=None, row_names=None, meta=None): + def store_matrix2d( + self, data, col_names=None, rows_col_name=None, meta=None): # Same as store_2d, but order is important raise NotImplementedError('store_matrix2d not implemented') - def store_2d(self, data, meta, columns=None, row_names=None): + def store_2d(self, data, meta, columns=None, rows_col_name=None): idx = element_to_index( - meta, n_rows=data.shape[0], row_names=row_names) + meta, n_rows=data.shape[0], rows_col_name=rows_col_name) data_df = pd.DataFrame(data, columns=columns, index=idx) self.store_df(data_df, meta) diff --git a/junifer/storage/tests/test_base.py b/junifer/storage/tests/test_base.py index 56779d212..37136ccab 100644 --- a/junifer/storage/tests/test_base.py +++ b/junifer/storage/tests/test_base.py @@ -31,7 +31,7 @@ def test_meta_hash(): meta1 = { 'element': 'foo', 'B': { - 'B2': [2, 3, 4, 5, 6], + 'B2': [2, 3, 4, 5, 6], 'B1': [9.22, 3.14, 1.41, 5.67, 6.28], 'B3': (1, 'car'), }, -- 2.52.0 From 06466b717ee1d685d1af54f20f6790a5a311074c Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Mon, 14 Mar 2022 15:25:12 +0100 Subject: [PATCH 021/287] More test! --- junifer/datagrabber/base.py | 7 +- junifer/storage/base.py | 18 ++-- junifer/storage/sqlite.py | 72 +++++++++++++-- junifer/storage/tests/test_base.py | 24 ++--- junifer/storage/tests/test_sqlite.py | 133 +++++++++++++++++++++++++++ 5 files changed, 221 insertions(+), 33 deletions(-) create mode 100644 junifer/storage/tests/test_sqlite.py diff --git a/junifer/datagrabber/base.py b/junifer/datagrabber/base.py index 8b8b15fa3..c9e73c2f5 100644 --- a/junifer/datagrabber/base.py +++ b/junifer/datagrabber/base.py @@ -16,9 +16,10 @@ def _validate_types(types): Validate the types """ if not isinstance(types, list): - raise_error("types must be a list", TypeError) + raise_error("types must be a list", TypeError) # type: ignore if any(not isinstance(x, str) for x in types): - raise_error("types must be a list of strings", TypeError) + raise_error( + "types must be a list of strings", TypeError) # type: ignore def _validate_patterns(types, patterns): @@ -27,7 +28,7 @@ def _validate_patterns(types, patterns): """ _validate_types(types) if not isinstance(patterns, dict): - raise_error("patterns must be a dict", TypeError) + raise_error("patterns must be a dict", TypeError) # type: ignore if len(types) != len(patterns): raise_error("types and patterns must have the same length", ValueError) diff --git a/junifer/storage/base.py b/junifer/storage/base.py index 1d14338b0..2e915f3d4 100644 --- a/junifer/storage/base.py +++ b/junifer/storage/base.py @@ -57,7 +57,7 @@ def element_to_index(meta, n_rows=1, rows_col_name=None): elem_idx = { k: [v] * n_rows for k, v in element.items() } - elem_idx[rows_col_name] = np.arange(n_rows) + elem_idx[rows_col_name] = np.arange(n_rows) # type: ignore else: elem_idx = element index = pd.MultiIndex.from_frame( @@ -86,29 +86,31 @@ class BaseFeatureStorage(ABC): raise NotImplementedError('store_metadata not implemented') @abstractmethod - def store_matrix2d(self, data, col_names=None, row_names=None, meta=None): + def store_matrix2d(self, data, meta, col_names=None, row_names=None): raise NotImplementedError('store_matrix2d not implemented') @abstractmethod - def store_table(self, data, columns=None, row_names=None, meta=None): + def store_table(self, data, meta, columns=None, rows_col_name=None): raise NotImplementedError('store_table not implemented') @abstractmethod - def store_df(self, df): + def store_df(self, df, meta): raise NotImplementedError('store_df not implemented') @abstractmethod - def store_timeseries(self, data): + def store_timeseries(self, data, meta): raise NotImplementedError('store_timeseries not implemented') class PandasFeatureStoreage(BaseFeatureStorage): def _meta_row(self, meta, meta_md5): - """Converta the meta to a dataframe row""" + """Convert the meta to a dataframe row""" data_df = {} - for k, v in meta: + for k, v in meta.items(): data_df[k] = json.dumps(v, sort_keys=True) - data_df['name'] = meta['marker']['name'] + if 'marker' in meta: + data_df['name'] = meta['marker']['name'] df = pd.DataFrame(data_df, index=[meta_md5]) + df.index.name = 'meta_md5' return df diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py index 1e62f9ca8..928733656 100644 --- a/junifer/storage/sqlite.py +++ b/junifer/storage/sqlite.py @@ -6,6 +6,7 @@ from pandas.io.sql import pandasSQL_builder from sqlalchemy import create_engine, inspect from .base import PandasFeatureStoreage, element_to_index, meta_hash +from ..utils.logging import warn class SQLiteFeatureStorage(PandasFeatureStoreage): @@ -13,9 +14,30 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): SQLite feature storage. """ - def __init__(self, uri): + def __init__(self, uri, upsert='update'): + """Initialise an SQLite feature storage + + Parameters + ---------- + uri : str + The connection URI. + Easy options: + 'sqlite://' for an in memory sqlite database + 'sqlite:///' to save in a file + + Check https://docs.sqlalchemy.org/en/14/core/engines.html for more + options + upsert : str + Upsert mode. Options are 'ignore' and 'update' (default). If + 'ignore', the existing elements are ignored. If update, the + existing elements are updated. + + """ + if upsert not in ['update', 'ignore']: + raise ValueError('upsert must be either "update" or "ignore"') super().__init__(uri) self._engine = create_engine(uri, echo=False) + self._upsert = upsert def store_metadata(self, meta): t_meta = meta.copy() @@ -24,16 +46,20 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): if meta_md5 not in inspect(self._engine).get_table_names(): meta_df = self._meta_row(t_meta, meta_md5) self._save_upsert(meta_df, 'meta') - return meta_md5 + return f'meta_{meta_md5}' def store_matrix2d( - self, data, col_names=None, rows_col_name=None, meta=None): + self, data, meta, col_names=None, rows_col_name=None): # Same as store_2d, but order is important raise NotImplementedError('store_matrix2d not implemented') + def store_table(self, data, meta, columns=None, rows_col_name=None): + self.store_2d(data, meta, columns, rows_col_name) + def store_2d(self, data, meta, columns=None, rows_col_name=None): + n_rows = len(data) idx = element_to_index( - meta, n_rows=data.shape[0], rows_col_name=rows_col_name) + meta, n_rows=n_rows, rows_col_name=rows_col_name) data_df = pd.DataFrame(data, columns=columns, index=idx) self.store_df(data_df, meta) @@ -41,13 +67,29 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): table_name = self.store_metadata(meta) self._save_upsert(df, table_name) - def store_timeseries(self, data): + def store_timeseries(self, data, meta): raise NotImplementedError('store_timeseries not implemented') - def _save_upsert(self, df, name, upsert='ignore', if_exist='append'): - if upsert not in ['delete', 'ignore']: - raise ValueError('upsert must be either "delete" or "ignore"') + def _save_upsert(self, df, name, if_exist='append'): + """ Implemention of UPSERT functionality. + Parameters + ---------- + df : pandas.DataFrame + DataFrame to save + name : str + Name of the table to save + if_exist : str + If the table exists, the behavior is controlled by this parameter. + Options are 'append' (default) and 'fail'. If 'fail' and th table + exists, it will raise an error. If 'append', the data will be + appended to the existing table (following the upsert mode). + + Raises + ______ + ValueError + If the table exists and if_exist is 'fail' + """ index_col = df.index.names with self._engine.begin() as con: if if_exist == 'replace': @@ -59,6 +101,8 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): else: # Case 3: existing table, so we need to check if the index # is present or not. + if if_exist == 'fail': + raise ValueError(f"Table ({name}) already exists") # Step 1: split incoming data into existing and new data pk_indb = _get_existing_pk( @@ -67,10 +111,18 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): # Step 2: upsert existing data pandas_sql = pandasSQL_builder(con) - pandas_sql.meta.reflect(only=[name]) + pandas_sql.meta.reflect(bind=con, only=[name]) table = pandas_sql.get_table(name) update_stmts = NoNewAttributesMixin - if upsert == 'delete': + if len(existing) > 0 and len(new) > 0: + warn( + f"Some rows (n={len(existing)}) are already present " + "in the database. The storage is configured to " + f"{self._upsert} the existing elements. The new rows " + f"(n={len(new)}) will be appended. This warning " + "is shown because normally all of the elements should " + "be updated") + if self._upsert == 'update': update_stmts = _generate_update_statements( table, index_col, existing) for stmt in update_stmts: diff --git a/junifer/storage/tests/test_base.py b/junifer/storage/tests/test_base.py index 37136ccab..5deceb110 100644 --- a/junifer/storage/tests/test_base.py +++ b/junifer/storage/tests/test_base.py @@ -114,32 +114,32 @@ def test_BaseFeatureStorage(): def store_metadata(self, metadata): super().store_metadata(metadata) - def store_matrix2d(self, matrix): - super().store_matrix2d(matrix) + def store_matrix2d(self, matrix, meta): + super().store_matrix2d(matrix, meta) - def store_table(self, table): - super().store_table(table) + def store_table(self, table, meta): + super().store_table(table, meta) - def store_df(self, df): - super().store_df(df) + def store_df(self, df, meta): + super().store_df(df, meta) - def store_timeseries(self, timeseries): - super().store_timeseries(timeseries) + def store_timeseries(self, timeseries, meta): + super().store_timeseries(timeseries, meta) st = MyFeatureStorage(uri='/tmp') with pytest.raises(NotImplementedError): st.store_metadata(None) with pytest.raises(NotImplementedError): - st.store_matrix2d(None) + st.store_matrix2d(None, None) with pytest.raises(NotImplementedError): - st.store_table(None) + st.store_table(None, None) with pytest.raises(NotImplementedError): - st.store_df(None) + st.store_df(None, None) with pytest.raises(NotImplementedError): - st.store_timeseries(None) + st.store_timeseries(None, None) assert st.uri == '/tmp' diff --git a/junifer/storage/tests/test_sqlite.py b/junifer/storage/tests/test_sqlite.py new file mode 100644 index 000000000..ee7fb8e3b --- /dev/null +++ b/junifer/storage/tests/test_sqlite.py @@ -0,0 +1,133 @@ +import pandas as pd +from pandas.testing import assert_frame_equal +import tempfile +from sqlalchemy import create_engine +import pytest + +from junifer.storage.sqlite import SQLiteFeatureStorage +from junifer.storage.base import element_to_index + + +df1 = pd.DataFrame({ + 'pk1': [1, 2, 3, 4, 5], + 'pk2': ['a', 'b', 'c', 'd', 'e'], + 'col1': [11, 22, 33, 44, 55], + 'col2': [111, 222, 333, 444, 555] +}).set_index(['pk1', 'pk2']) + +df2 = pd.DataFrame({ + 'pk1': [2, 5, 6], + 'pk2': ['b', 'e', 'f'], + 'col1': [2222, 5555, 66], + 'col2': [22222, 55555, 666] +}).set_index(['pk1', 'pk2']) + +df_update = pd.DataFrame({ + 'pk1': [1, 2, 3, 4, 5, 6], + 'pk2': ['a', 'b', 'c', 'd', 'e', 'f'], + 'col1': [11, 2222, 33, 44, 5555, 66], + 'col2': [111, 22222, 333, 444, 55555, 666] +}).set_index(['pk1', 'pk2']) + +df_ignore = pd.DataFrame({ + 'pk1': [1, 2, 3, 4, 5, 6], + 'pk2': ['a', 'b', 'c', 'd', 'e', 'f'], + 'col1': [11, 22, 33, 44, 55, 66], + 'col2': [111, 222, 333, 444, 555, 666] +}).set_index(['pk1', 'pk2']) + + +def _read_sql(table_name, uri, index_col): + engine = create_engine(uri, echo=False) + df = pd.read_sql(table_name, con=engine, index_col=index_col) + return df + + +def test_upsert_ignore(): + """Test store_df (upsert=ignore)""" + with tempfile.TemporaryDirectory() as _tmpdir: + uri = f'sqlite:///{_tmpdir}/test.db' + storage = SQLiteFeatureStorage(uri=uri, upsert='ignore') + meta = {'element': 'test', 'version': '0.0.1'} + + # Save to SQL + storage.store_df(df1, meta) + + # Test the internals + table_name = storage.store_metadata(meta) + + c_df1 = _read_sql(table_name, uri=uri, index_col=['pk1', 'pk2']) + assert_frame_equal(df1, c_df1) + + storage.store_df(df2, meta) + + c_dfignore = _read_sql(table_name, uri=uri, index_col=['pk1', 'pk2']) + assert_frame_equal(c_dfignore, df_ignore) + + with pytest.raises(ValueError, match=r"already exists"): + storage._save_upsert(df2, table_name, if_exist='fail') + + +def test_upsert_update(): + """Test store_df (upsert=delete)""" + meta = {'element': 'test', 'version': '0.0.1'} + with tempfile.TemporaryDirectory() as _tmpdir: + uri = f'sqlite:///{_tmpdir}/test.db' + storage = SQLiteFeatureStorage(uri=uri) + + # Save to SQL + storage.store_df(df1, meta) + + # Test the internals + table_name = storage.store_metadata(meta) + + c_df1 = _read_sql(table_name, uri=uri, index_col=['pk1', 'pk2']) + assert_frame_equal(df1, c_df1) + + storage.store_df(df2, meta) + + c_dfupdate = _read_sql(table_name, uri=uri, index_col=['pk1', 'pk2']) + assert_frame_equal(c_dfupdate, df_update) + + +def test_store_table(): + """Test store_df""" + meta = {'element': 'test', 'version': '0.0.1'} + with tempfile.TemporaryDirectory() as _tmpdir: + uri = f'sqlite:///{_tmpdir}/test.db' + storage = SQLiteFeatureStorage(uri=uri) + data = [ + [1, 10], + [2, 20], + [3, 30], + [4, 40], + [5, 50], + ] + idx = element_to_index(meta, n_rows=5, rows_col_name='scan') + df1 = pd.DataFrame(data, columns=['f1', 'f2'], index=idx) + + storage.store_table( + data, meta, columns=['f1', 'f2'], rows_col_name='scan') + + table_name = storage.store_metadata(meta) + c_df1 = _read_sql(table_name, uri=uri, index_col=['element', 'scan']) + assert_frame_equal(df1, c_df1) + + data2 = [ + [1, 10], + [2, 20], + [3, 300], + [4, 40], + [5, 50], + [6, 600] + ] + + idx = element_to_index(meta, n_rows=6, rows_col_name='scan') + df2 = pd.DataFrame(data2, columns=['f1', 'f2'], index=idx) + + with pytest.warns(RuntimeWarning, match=r"Some rows"): + storage.store_table( + data2, meta, columns=['f1', 'f2'], rows_col_name='scan') + + c_df2 = _read_sql(table_name, uri=uri, index_col=['element', 'scan']) + assert_frame_equal(df2, c_df2) -- 2.52.0 From c1fd143e0d226084f65758bb9f096869712bbf02 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Mon, 14 Mar 2022 15:27:51 +0100 Subject: [PATCH 022/287] typo --- junifer/storage/sqlite.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py index 928733656..1321ab2d3 100644 --- a/junifer/storage/sqlite.py +++ b/junifer/storage/sqlite.py @@ -71,7 +71,7 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): raise NotImplementedError('store_timeseries not implemented') def _save_upsert(self, df, name, if_exist='append'): - """ Implemention of UPSERT functionality. + """ Implementation of UPSERT functionality. Parameters ---------- -- 2.52.0 From 4ab08e08a0bf6dff987e047df536aaa85c9b9be0 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Mon, 14 Mar 2022 15:37:38 +0100 Subject: [PATCH 023/287] More tests!! --- junifer/storage/tests/test_sqlite.py | 27 ++++++++++++++++++++++++++- 1 file changed, 26 insertions(+), 1 deletion(-) diff --git a/junifer/storage/tests/test_sqlite.py b/junifer/storage/tests/test_sqlite.py index ee7fb8e3b..41e93d9c1 100644 --- a/junifer/storage/tests/test_sqlite.py +++ b/junifer/storage/tests/test_sqlite.py @@ -43,10 +43,35 @@ def _read_sql(table_name, uri, index_col): return df +def test_upsert_replace(): + """Test store_df (if_exist=replace)""" + with tempfile.TemporaryDirectory() as _tmpdir: + uri = f'sqlite:///{_tmpdir}/test.db' + storage = SQLiteFeatureStorage(uri=uri, upsert='ignore') + meta = {'element': 'test', 'version': '0.0.1'} + + # Save to SQL + storage.store_df(df1, meta) + + # Test the internals + table_name = storage.store_metadata(meta) + + c_df1 = _read_sql(table_name, uri=uri, index_col=['pk1', 'pk2']) + assert_frame_equal(df1, c_df1) + + storage._save_upsert(df2, table_name, if_exist='replace') + + c_df2 = _read_sql(table_name, uri=uri, index_col=['pk1', 'pk2']) + assert_frame_equal(df2, c_df2) + + def test_upsert_ignore(): """Test store_df (upsert=ignore)""" with tempfile.TemporaryDirectory() as _tmpdir: uri = f'sqlite:///{_tmpdir}/test.db' + with pytest.raises(ValueError): + SQLiteFeatureStorage(uri=uri, upsert='wrong') + storage = SQLiteFeatureStorage(uri=uri, upsert='ignore') meta = {'element': 'test', 'version': '0.0.1'} @@ -92,7 +117,7 @@ def test_upsert_update(): def test_store_table(): """Test store_df""" - meta = {'element': 'test', 'version': '0.0.1'} + meta = {'element': 'test', 'version': '0.0.1', 'marker': 'fc'} with tempfile.TemporaryDirectory() as _tmpdir: uri = f'sqlite:///{_tmpdir}/test.db' storage = SQLiteFeatureStorage(uri=uri) -- 2.52.0 From c37372ff932beb6c5bfa458db66e0c64a0729e20 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Mon, 14 Mar 2022 16:12:33 +0100 Subject: [PATCH 024/287] Fix one more test --- junifer/storage/tests/test_sqlite.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/junifer/storage/tests/test_sqlite.py b/junifer/storage/tests/test_sqlite.py index 41e93d9c1..73b7009ed 100644 --- a/junifer/storage/tests/test_sqlite.py +++ b/junifer/storage/tests/test_sqlite.py @@ -117,7 +117,7 @@ def test_upsert_update(): def test_store_table(): """Test store_df""" - meta = {'element': 'test', 'version': '0.0.1', 'marker': 'fc'} + meta = {'element': 'test', 'version': '0.0.1', 'marker': {'name': 'fc'}} with tempfile.TemporaryDirectory() as _tmpdir: uri = f'sqlite:///{_tmpdir}/test.db' storage = SQLiteFeatureStorage(uri=uri) -- 2.52.0 From 8d5254e4f4a5fd2de736ccb9dca46fdc8c0ad0e7 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Tue, 15 Mar 2022 08:10:19 +0100 Subject: [PATCH 025/287] WIP --- junifer/api/__init__.py | 2 +- junifer/api/decorators.py | 42 ++++++++- junifer/api/{pipeline.py => run.py} | 2 +- junifer/datagrabber/base.py | 22 +++-- junifer/datareader/default.py | 3 + junifer/markers/base.py | 10 ++- junifer/markers/parcel.py | 12 +++ junifer/markers/tests/test_collection.py | 46 ++++++++++ junifer/markers/tests/test_parcel.py | 12 ++- junifer/storage/__init__.py | 3 +- junifer/storage/base.py | 108 +++++++++++++++++++++-- junifer/storage/sqlite.py | 61 ++++++++++++- junifer/storage/tests/test_base.py | 84 +++++++++++++----- junifer/storage/tests/test_sqlite.py | 49 +++++++++- junifer/testing.py | 1 + 15 files changed, 405 insertions(+), 52 deletions(-) rename junifer/api/{pipeline.py => run.py} (99%) diff --git a/junifer/api/__init__.py b/junifer/api/__init__.py index 51642c08a..4bc42cc57 100644 --- a/junifer/api/__init__.py +++ b/junifer/api/__init__.py @@ -1 +1 @@ -from . pipeline import run_pipeline \ No newline at end of file +from . run import run \ No newline at end of file diff --git a/junifer/api/decorators.py b/junifer/api/decorators.py index fa8ac0257..4808fe1ec 100644 --- a/junifer/api/decorators.py +++ b/junifer/api/decorators.py @@ -1,13 +1,13 @@ # Authors: Federico Raimondo # Leonard Sasse # License: AGPL -from . pipeline import register +from .run import register def register_datagrabber(klass): """Datagrabber decorator. - Registers the datagrabber so it can be used by name in the pipeline. + Registers the datagrabber so it can be used by name. Parameters ---------- @@ -21,3 +21,41 @@ def register_datagrabber(klass): """ register('datagrabber', klass.__name__, klass) return klass + + +def register_marker(klass): + """marker decorator. + + Registers the marker so it can be used by name. + + Parameters + ---------- + klass: class + The class of the marker to register. + + Returns + ------- + klass: class + The unmodified input class + """ + register('marker', klass.__name__, klass) + return klass + + +def register_storage(klass): + """Storage decorator. + + Registers the storage so it can be used by name. + + Parameters + ---------- + klass: class + The class of the storage to register. + + Returns + ------- + klass: class + The unmodified input class + """ + register('storage', klass.__name__, klass) + return klass diff --git a/junifer/api/pipeline.py b/junifer/api/run.py similarity index 99% rename from junifer/api/pipeline.py rename to junifer/api/run.py index 0bf3a5c23..35a4c4038 100644 --- a/junifer/api/pipeline.py +++ b/junifer/api/run.py @@ -27,7 +27,7 @@ def register(step, name, klass): _registry[step][name] = klass -def run_pipeline( +def run( workdir, datagrabber, element, markers, storage, source_params=None, storage_params=None): """Run the pipeline on the selected element diff --git a/junifer/datagrabber/base.py b/junifer/datagrabber/base.py index c9e73c2f5..7df89e8d2 100644 --- a/junifer/datagrabber/base.py +++ b/junifer/datagrabber/base.py @@ -87,6 +87,9 @@ class BaseDataGrabber(ABC): t_meta[k] = v return t_meta + def get_element_keys(self): + return 'element' + @property def datadir(self): """ @@ -183,8 +186,10 @@ class BIDSDataGrabber(BaseDataGrabber): Parameters ---------- - element : str - The element to be indexed. + element : str or tuple + The element to be indexed. If one string is provided, it is + assumed to be a subject. If a tuple is provided, it is assumed to + be a (subject, session) pair. Returns ------- out : dict[str -> Path] @@ -192,14 +197,21 @@ class BIDSDataGrabber(BaseDataGrabber): specified element. """ out = {} + if not isinstance(element, tuple): + element = (element,) for t_type in self.types: t_pattern = self.patterns[t_type] # type: ignore - t_replace = t_pattern.replace('{subject}', element) - t_out = self.datadir / element / t_replace + t_replace = t_pattern.replace('{subject}', element[0]) + if len(element) > 1: + t_replace = t_replace.replace( + '{session}', element[1]) # type: ignore + t_out = self.datadir / element[0] / t_replace out[t_type] = dict(path=t_out) # Meta here is element and types out['meta'] = dict(datagrabber=self.get_meta()) - out['meta']['element'] = element + out['meta']['element'] = {'subject': element[0]} + if len(element) > 1: + out['meta']['element']['session'] = element[1] # type: ignore return out diff --git a/junifer/datareader/default.py b/junifer/datareader/default.py index 80cc11e15..aa71c7e63 100644 --- a/junifer/datareader/default.py +++ b/junifer/datareader/default.py @@ -41,6 +41,9 @@ class DefaultDataReader(PipelineStepMixin): if params is None: params = {} for kind in input.keys(): + if kind == 'meta': + out['meta'] = input['meta'] + continue t_path = input[kind] t_params = params.get(kind, {}) if not isinstance(t_path, Path): diff --git a/junifer/markers/base.py b/junifer/markers/base.py index e06cd302b..9d736aa15 100644 --- a/junifer/markers/base.py +++ b/junifer/markers/base.py @@ -77,9 +77,10 @@ class BaseMarker(PipelineStepMixin): self._valid_inputs = on self.name = self.__class__.__name__ if name is None else name - def get_meta(self): + def get_meta(self, kind): s_meta = super().get_meta() s_meta['name'] = self.name + s_meta['kind'] = kind # same marker can be fit into different kinds return dict(marker=s_meta) def validate_input(self, input): @@ -95,6 +96,9 @@ class BaseMarker(PipelineStepMixin): def compute(self, input): raise NotImplementedError('compute not implemented') + def store(self, input, out, storage): + raise NotImplementedError('store not implemented') + def fit_transform(self, input, storage=None): out = {} meta = input.get('meta', {}) @@ -104,11 +108,11 @@ class BaseMarker(PipelineStepMixin): t_input = input[kind] t_meta = meta.copy() t_meta.update(t_input.get('meta', {})) - t_meta.update(self.get_meta()) + t_meta.update(self.get_meta(kind)) t_out = self.compute(t_input) t_out.update(meta=t_meta) if storage is not None: - storage.store_2d(**t_out) + self.store(kind, t_out, storage) else: out[kind] = t_out diff --git a/junifer/markers/parcel.py b/junifer/markers/parcel.py index c277bcacc..0f1beb09c 100644 --- a/junifer/markers/parcel.py +++ b/junifer/markers/parcel.py @@ -19,6 +19,18 @@ class ParcelAggregation(BaseMarker): self.method = method self.method_params = {} if method_params is None else method_params + def get_output_kind(self, input): + if input in ['GMD_GM', 'GMD_WM']: + return 'table' + if input in ['BOLD']: + return 'timeseries' + + def store(self, kind, out, storage): + if kind in ['GMD_GM', 'GMD_WM']: + storage.store_table(**out) + if kind in ['BOLD']: + storage.store_timeseries(**out) + def compute(self, input): t_input = input['data'] agg_func = get_aggfunc_by_name( diff --git a/junifer/markers/tests/test_collection.py b/junifer/markers/tests/test_collection.py index 2cf8adf64..b68d2a7fd 100644 --- a/junifer/markers/tests/test_collection.py +++ b/junifer/markers/tests/test_collection.py @@ -2,11 +2,13 @@ # License: AGPL import pytest +import tempfile from numpy.testing import assert_array_equal from junifer.datareader.default import DefaultDataReader from junifer.markers import MarkerCollection, ParcelAggregation from junifer.markers.base import PipelineStepMixin from junifer.testing import OasisVBMTestingDatagrabber +from junifer.storage import SQLiteFeatureStorage def test_MarkerCollection(): @@ -78,3 +80,47 @@ def test_MarkerCollection(): t_name = t_marker.name assert_array_equal(out[t_name]['VBM_GM']['data'], out2[t_name]['VBM_GM']['data']) # type: ignore + + +def test_MarkerCollection_storage(): + """Test marker collection with storage""" + markers = [ + ParcelAggregation( + atlas='Schaefer100x7', method='mean', + name='gmd_schaefer100x7_mean'), + ParcelAggregation( + atlas='Schaefer100x7', method='std', + name='gmd_schaefer100x7_std'), + ParcelAggregation( + atlas='Schaefer100x7', method='trim_mean', + method_params={'proportiontocut': 0.1}, + name='gmd_schaefer100x7_trim_mean90') + ] + # Test storage + dg = OasisVBMTestingDatagrabber() + with tempfile.TemporaryDirectory() as tmpdir: + uri = f'sqlite:///{tmpdir}/test.db' + storage = SQLiteFeatureStorage(uri=uri) + mc = MarkerCollection( + markers=markers, storage=storage, datareader=DefaultDataReader()) + mc.validate(dg) + assert mc._storage.uri == storage.uri + with dg: + input = dg[1] + out = mc.fit(input) + assert out is None + + mc2 = MarkerCollection( + markers=markers, datareader=DefaultDataReader()) + mc2.validate(dg) + assert mc2._storage is None + + with dg: + input = dg[1] + out = mc2.fit(input) + + features = storage.list_features() + assert len(features) == 1 + feature_md5 = features.keys()[0] + t_feature = storage.read_df(feature_md5=feature_md5) + diff --git a/junifer/markers/tests/test_parcel.py b/junifer/markers/tests/test_parcel.py index 1de773bfc..7ab06d591 100644 --- a/junifer/markers/tests/test_parcel.py +++ b/junifer/markers/tests/test_parcel.py @@ -61,11 +61,12 @@ def test_ParcelAggregation_3D(): assert jun_values3d_mean.shape[0] == 1 assert_array_equal(manual, jun_values3d_mean) - meta = marker.get_meta()['marker'] + meta = marker.get_meta('VBM_GM')['marker'] assert meta['method'] == 'mean' assert meta['atlas'] == 'Schaefer100x7' assert meta['name'] == 'gmd_schaefer100x7_mean' assert meta['class'] == 'ParcelAggregation' + assert meta['kind'] == 'VBM_GM' assert meta['method_params'] == {} # Test using another function (std) @@ -84,11 +85,12 @@ def test_ParcelAggregation_3D(): assert jun_values3d_std.shape[0] == 1 assert_array_equal(manual, jun_values3d_std) - meta = marker.get_meta()['marker'] + meta = marker.get_meta('VBM_GM')['marker'] assert meta['method'] == 'std' assert meta['atlas'] == 'Schaefer100x7' assert meta['name'] == 'ParcelAggregation' assert meta['class'] == 'ParcelAggregation' + assert meta['kind'] == 'VBM_GM' assert meta['method_params'] == {} # Test using another function with parameters @@ -110,11 +112,12 @@ def test_ParcelAggregation_3D(): assert jun_values3d_tm.shape[0] == 1 assert_array_equal(manual, jun_values3d_tm) - meta = marker.get_meta()['marker'] + meta = marker.get_meta('VBM_GM')['marker'] assert meta['method'] == 'trim_mean' assert meta['atlas'] == 'Schaefer100x7' assert meta['name'] == 'ParcelAggregation' assert meta['class'] == 'ParcelAggregation' + assert meta['kind'] == 'VBM_GM' assert meta['method_params'] == {'proportiontocut': 0.1} @@ -140,9 +143,10 @@ def test_ParcelAggregation_4D(): assert_array_equal(auto4d.shape, jun_values4d.shape) assert_array_equal(auto4d, jun_values4d) - meta = marker.get_meta()['marker'] + meta = marker.get_meta('BOLD')['marker'] assert meta['method'] == 'mean' assert meta['atlas'] == 'Schaefer100x7' assert meta['name'] == 'ParcelAggregation' assert meta['class'] == 'ParcelAggregation' + assert meta['kind'] == 'BOLD' assert meta['method_params'] == {} diff --git a/junifer/storage/__init__.py b/junifer/storage/__init__.py index 0f49b9785..81239bdf4 100644 --- a/junifer/storage/__init__.py +++ b/junifer/storage/__init__.py @@ -1,2 +1,3 @@ # Authors: Federico Raimondo -# License: AGPL \ No newline at end of file +# License: AGPL +from .sqlite import SQLiteFeatureStorage \ No newline at end of file diff --git a/junifer/storage/base.py b/junifer/storage/base.py index 2e915f3d4..5ff8a9993 100644 --- a/junifer/storage/base.py +++ b/junifer/storage/base.py @@ -9,7 +9,58 @@ from abc import ABC, abstractmethod from .. import __version__ -def meta_hash(meta): +def process_meta(meta, return_idx=False, n_rows=1, rows_col_name=None): + """Process the metadata for storage. It removes the "element" key + and adds the "_element_keys" with the keys used to index the element. + + Parameters + ---------- + meta: dict + The metadata. Must contain the key 'element' + return_idx: bool + If true, return the pandas index to be stored. Defaults to false + n_rows: int + Number of rows to create (if return_idx is true) + rows_col_name: str + The column name to use in case n_rows > 1. If None (default) and + n_rows > 1, the name will be 'index'. + + Returns + ------- + md5_hash: str + The md5 hash of the meta + meta : dict + The metadata processed for storage + idx : pd.MultiIndex + The pandas index (if return_idx is True) + """ + if meta is None: + raise ValueError('Meta must be a dict (currently is None)') + t_meta = meta.copy() + idx = None + if return_idx is True: + idx = _element_to_index( + meta, n_rows=n_rows, rows_col_name=rows_col_name) + element = t_meta.pop('element', None) + if element is None: + if '_element_keys' not in t_meta: + raise ValueError( + 'Meta must contain the key "element" or "_element_keys"') + else: + if isinstance(element, dict): + t_meta['_element_keys'] = list(element.keys()) + else: + t_meta['_element_keys'] = ['element'] + + md5_hash = _meta_hash(meta) + if return_idx is True: + out = md5_hash, t_meta, idx + else: + out = md5_hash, t_meta + return out + + +def _meta_hash(meta): """Compute the md5 hash of the meta Parameters @@ -22,15 +73,12 @@ def meta_hash(meta): md5: str The md5 hash of the meta """ - if meta is None: - raise ValueError('Meta must be a dict (currently is None)') - t_meta = meta.copy() meta_md5 = hashlib.md5( - json.dumps(t_meta, sort_keys=True).encode('utf-8')).hexdigest() + json.dumps(meta, sort_keys=True).encode('utf-8')).hexdigest() return meta_md5 -def element_to_index(meta, n_rows=1, rows_col_name=None): +def _element_to_index(meta, n_rows=1, rows_col_name=None): """Convert the element meta to index Parameters @@ -47,7 +95,15 @@ def element_to_index(meta, n_rows=1, rows_col_name=None): ------- index: pd.MultiIndex The index of the dataframe to store + + Raises + ------ + ValueError + If the meta does not contain the key 'element' """ + if 'element' not in meta: + raise ValueError( + 'To create and index, meta must contain the key "element"') element = meta['element'] if not isinstance(element, dict): element = dict(element=element) @@ -81,6 +137,46 @@ class BaseFeatureStorage(ABC): } return meta + @abstractmethod + def validate(self, input): + """Validate the input to the pipeline step. + + Parameters + ---------- + input : Junifer Data dictionary + The input to the pipeline step. + + Raises + ------ + ValueError: + If the input does not have the required data. + """ + raise NotImplementedError('validate_input not implemented') + + @abstractmethod + def list_features(self): + """List the features in the storage + + Returns + ------- + features: dict(str, dict) + List of features in the storage. The keys are the feature names + to be used in read_features. The values are the metadata of each + feature + """ + raise NotImplementedError('list_features not implemented') + + @abstractmethod + def read_df(self, feature_name=None, feature_md5=None): + """Read the features from the storage + + Returns + ------- + out: pd.DataFrame + The features as a dataframe + """ + raise NotImplementedError('read_df not implemented') + @abstractmethod def store_metadata(self, meta): raise NotImplementedError('store_metadata not implemented') diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py index 1321ab2d3..26c072e7e 100644 --- a/junifer/storage/sqlite.py +++ b/junifer/storage/sqlite.py @@ -5,10 +5,12 @@ from pandas.core.base import NoNewAttributesMixin from pandas.io.sql import pandasSQL_builder from sqlalchemy import create_engine, inspect -from .base import PandasFeatureStoreage, element_to_index, meta_hash +from ..api.decorators import register_storage +from .base import PandasFeatureStoreage, process_meta from ..utils.logging import warn +@register_storage class SQLiteFeatureStorage(PandasFeatureStoreage): """ SQLite feature storage. @@ -38,11 +40,61 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): super().__init__(uri) self._engine = create_engine(uri, echo=False) self._upsert = upsert + self._valid_inputs = ['table', 'timeseries'] + + def validate(self, input): + if not isinstance(input, list): + input = [input] + return all(x in self._valid_inputs for x in input) + + def list_features(self): + meta_df = pd.read_sql('meta', con=self._engine, index_col='meta_md5') + return meta_df.to_dict(orient='index') + + def read_df(self, feature_name=None, feature_md5=None): + """Read features from the storage. + + Parameters + ---------- + feature_name : str + Name of the feature to read. At least one of feature_name or + feature_md5 must be specified. + feature_md5 : str + MD5 of the feature to read. At least one of feature_name or + feature_md5 must be specified. + + Returns + ------- + pandas.DataFrame + The features. + """ + if feature_md5 is not None and feature_name is not None: + raise ValueError('Only one of feature_name or feature_md5 can be ' + 'specified') + elif feature_md5 is None and feature_name is None: + raise ValueError('At least one of feature_name or feature_md5 ' + 'must be specified') + elif feature_md5 is not None: + feature_md5 = f'meta_{feature_md5}' + else: + meta_df = pd.read_sql( + 'meta', con=self._engine, index_col='meta_md5') + t_df = meta_df.query(f"name == '{feature_name}'") + if len(t_df) == 0: + raise ValueError(f'Feature {feature_name} not found') + elif len(t_df) > 1: + raise ValueError( + f'More than one feature with name {feature_name} found', + 'This file is invalid. You can bypass this issue by ' + 'specifying a feature_md5') + feature_md5 = t_df.index[0] + return pd.read_sql(feature_md5, con=self._engine) def store_metadata(self, meta): t_meta = meta.copy() t_meta.update(self.get_meta()) - meta_md5 = meta_hash(t_meta) + meta_md5, t_meta = process_meta( # type: ignore + t_meta, return_idx=False) if meta_md5 not in inspect(self._engine).get_table_names(): meta_df = self._meta_row(t_meta, meta_md5) self._save_upsert(meta_df, 'meta') @@ -58,13 +110,14 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): def store_2d(self, data, meta, columns=None, rows_col_name=None): n_rows = len(data) - idx = element_to_index( - meta, n_rows=n_rows, rows_col_name=rows_col_name) + _, _, idx = process_meta( # type: ignore + meta, return_idx=True, n_rows=n_rows, rows_col_name=rows_col_name) data_df = pd.DataFrame(data, columns=columns, index=idx) self.store_df(data_df, meta) def store_df(self, df, meta): table_name = self.store_metadata(meta) + # TODO: Check that the element is in the index of the dataframe self._save_upsert(df, table_name) def store_timeseries(self, data, meta): diff --git a/junifer/storage/tests/test_base.py b/junifer/storage/tests/test_base.py index 5deceb110..e051b27aa 100644 --- a/junifer/storage/tests/test_base.py +++ b/junifer/storage/tests/test_base.py @@ -1,31 +1,25 @@ import pytest -from junifer.storage.base import (meta_hash, element_to_index, +from junifer.storage.base import (process_meta, _element_to_index, BaseFeatureStorage) -def test_meta_hash(): +def test_process_meta_hash(): """Test meta_hash""" - meta = {} - hash = meta_hash(meta) - - assert hash == '99914b932bd37a50b983c5e7c90ae93b' # empty dict - meta = None with pytest.raises(ValueError, match=r"Meta must be a dict"): - meta_hash(meta) + process_meta(meta) meta = {'element': 'foo', 'A': 1, 'B': [2, 3, 4, 5, 6]} - hash1 = meta_hash(meta) + hash1, _ = process_meta(meta, return_idx=False) # type: ignore meta = {'element': 'foo', 'B': [2, 3, 4, 5, 6], 'A': 1} - hash2 = meta_hash(meta) - + hash2, _ = process_meta(meta, return_idx=False) # type: ignore assert hash1 == hash2 meta = {'element': 'foo', 'A': 1, 'B': [2, 3, 1, 5, 6]} - hash3 = meta_hash(meta) + hash3, _ = process_meta(meta, return_idx=False) # type: ignore assert hash1 != hash3 meta1 = { @@ -47,21 +41,49 @@ def test_meta_hash(): 'element': 'foo' } - hash4 = meta_hash(meta1) - hash5 = meta_hash(meta2) + hash4, _ = process_meta(meta1, return_idx=False) # type: ignore + hash5, _ = process_meta(meta2, return_idx=False) # type: ignore assert hash4 == hash5 -def test_element_to_index(): - """Test element_to_index""" +def test_process_meta_element(): + """Test meta element""" + + meta = {} + with pytest.raises(ValueError, match=r"_element_keys"): + process_meta(meta, return_idx=False) meta = {'element': 'foo', 'A': 1, 'B': [2, 3, 4, 5, 6]} - index = element_to_index(meta) + _, new_meta = process_meta(meta, return_idx=False) # type: ignore + assert '_element_keys' in new_meta + assert new_meta['_element_keys'] == ['element'] + assert 'A' in new_meta + assert 'B' in new_meta + + meta = { + 'element': {'subject': 'foo', 'session': 'bar'}, + 'B': [2, 3, 4, 5, 6], 'A': 1} + _, new_meta = process_meta(meta, return_idx=False) # type: ignore + assert '_element_keys' in new_meta + assert new_meta['_element_keys'] == ['subject', 'session'] + assert 'A' in new_meta + assert 'B' in new_meta + + +def test_process_meta_index(): + """Test _element_to_index""" + + meta = {'noelement': 'foo'} + with pytest.raises(ValueError, match=r'meta must contain the key'): + _element_to_index(meta) + + meta = {'element': 'foo', 'A': 1, 'B': [2, 3, 4, 5, 6]} + index = _element_to_index(meta) assert index.names == ['element'] assert index.levels[0].name == 'element' assert index.levels[0].values[0] == 'foo' - index = element_to_index(meta, n_rows=10) + index = _element_to_index(meta, n_rows=10) assert index.names == ['element', 'index'] assert index.levels[0].name == 'element' assert all(x == 'foo' for x in index.levels[0].values) @@ -71,13 +93,13 @@ def test_element_to_index(): assert all(x == i for i, x in enumerate(index.levels[1].values)) assert index.levels[1].values.shape == (10,) - index = element_to_index(meta, n_rows=1, rows_col_name='scan') + index = _element_to_index(meta, n_rows=1, rows_col_name='scan') assert index.names == ['element'] assert index.levels[0].name == 'element' assert all(x == 'foo' for x in index.levels[0].values) assert index.levels[0].values.shape == (1,) - index = element_to_index(meta, n_rows=7, rows_col_name='scan') + index = _element_to_index(meta, n_rows=7, rows_col_name='scan') assert index.names == ['element', 'scan'] assert index.levels[0].name == 'element' assert all(x == 'foo' for x in index.levels[0].values) @@ -90,7 +112,7 @@ def test_element_to_index(): meta = { 'element': {'subject': 'sub-01', 'session': 'ses-01'}, 'A': 1, 'B': [2, 3, 4, 5, 6]} - index = element_to_index(meta, n_rows=10) + index = _element_to_index(meta, n_rows=10) assert index.levels[0].name == 'subject' assert all(x == 'sub-01' for x in index.levels[0].values) @@ -111,6 +133,16 @@ def test_BaseFeatureStorage(): BaseFeatureStorage(uri='/tmp') # type: ignore class MyFeatureStorage(BaseFeatureStorage): + def validate(self, input): + super().validate(input) + + def list_features(self): + super().list_features() + + def read_df(self, feature_name=None, feature_md5=None): + super().read_df( + feature_name=feature_name, feature_md5=feature_md5) + def store_metadata(self, metadata): super().store_metadata(metadata) @@ -127,6 +159,16 @@ def test_BaseFeatureStorage(): super().store_timeseries(timeseries, meta) st = MyFeatureStorage(uri='/tmp') + + with pytest.raises(NotImplementedError): + st.validate(None) + + with pytest.raises(NotImplementedError): + st.list_features() + + with pytest.raises(NotImplementedError): + st.read_df(None) + with pytest.raises(NotImplementedError): st.store_metadata(None) diff --git a/junifer/storage/tests/test_sqlite.py b/junifer/storage/tests/test_sqlite.py index 73b7009ed..0c4d57083 100644 --- a/junifer/storage/tests/test_sqlite.py +++ b/junifer/storage/tests/test_sqlite.py @@ -5,7 +5,7 @@ from sqlalchemy import create_engine import pytest from junifer.storage.sqlite import SQLiteFeatureStorage -from junifer.storage.base import element_to_index +from junifer.storage.base import process_meta df1 = pd.DataFrame({ @@ -80,6 +80,7 @@ def test_upsert_ignore(): # Test the internals table_name = storage.store_metadata(meta) + c_df1 = _read_sql(table_name, uri=uri, index_col=['pk1', 'pk2']) assert_frame_equal(df1, c_df1) @@ -115,8 +116,46 @@ def test_upsert_update(): assert_frame_equal(c_dfupdate, df_update) -def test_store_table(): +def test_store_read_df(): """Test store_df""" + with tempfile.TemporaryDirectory() as _tmpdir: + uri = f'sqlite:///{_tmpdir}/test.db' + storage = SQLiteFeatureStorage(uri=uri, upsert='ignore') + meta = { + 'element': 'test', 'version': '0.0.1', + 'marker': {'name': 'fcname'}} + + # Save to SQL + storage.store_df(df1, meta) + + # Test the internals + table_name = storage.store_metadata(meta) + + features = storage.list_features() + assert len(features) == 1 + assert table_name.replace('meta_', '') in features + + with pytest.raises(ValueError, match='not found'): + storage.read_df('wrong_md5') + + with pytest.raises(ValueError, match='least one'): + storage.read_df() + + with pytest.raises(ValueError, match='Only one'): + storage.read_df('wrong_md5', 'wrong_name') + + feature_md5 = list(features.keys())[0] + assert 'fcname' == features[feature_md5]['name'] + read_df1 = storage.read_df(feature_md5=feature_md5) + read_df2 = storage.read_df(feature_name='fcname') + assert_frame_equal(read_df1, read_df2) + assert_frame_equal(read_df1, df1) + + + + +def test_store_table(): + """Test store_table""" meta = {'element': 'test', 'version': '0.0.1', 'marker': {'name': 'fc'}} with tempfile.TemporaryDirectory() as _tmpdir: uri = f'sqlite:///{_tmpdir}/test.db' @@ -128,7 +167,8 @@ def test_store_table(): [4, 40], [5, 50], ] - idx = element_to_index(meta, n_rows=5, rows_col_name='scan') + _, _, idx = process_meta( # type: ignore + meta, return_idx=True, n_rows=5, rows_col_name='scan') df1 = pd.DataFrame(data, columns=['f1', 'f2'], index=idx) storage.store_table( @@ -147,7 +187,8 @@ def test_store_table(): [6, 600] ] - idx = element_to_index(meta, n_rows=6, rows_col_name='scan') + _, _, idx = process_meta( # type: ignore + meta, return_idx=True, n_rows=6, rows_col_name='scan') df2 = pd.DataFrame(data2, columns=['f1', 'f2'], index=idx) with pytest.warns(RuntimeWarning, match=r"Some rows"): diff --git a/junifer/testing.py b/junifer/testing.py index 60a298ec4..46dca9f5e 100644 --- a/junifer/testing.py +++ b/junifer/testing.py @@ -21,6 +21,7 @@ class OasisVBMTestingDatagrabber(BaseDataGrabber): def __getitem__(self, element): out = {} out['VBM_GM'] = self._dataset.gray_matter_maps[element] + out['meta'] = {'subject': element} return out def __enter__(self): -- 2.52.0 From d3a1dca14a688d3791a0e33333565332f8eb39c3 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Tue, 15 Mar 2022 08:26:42 +0100 Subject: [PATCH 026/287] Test store and read_df --- junifer/storage/sqlite.py | 13 +++++++++++-- junifer/storage/tests/test_sqlite.py | 16 +++++++++++----- 2 files changed, 22 insertions(+), 7 deletions(-) diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py index 26c072e7e..c794d9a99 100644 --- a/junifer/storage/sqlite.py +++ b/junifer/storage/sqlite.py @@ -87,7 +87,7 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): f'More than one feature with name {feature_name} found', 'This file is invalid. You can bypass this issue by ' 'specifying a feature_md5') - feature_md5 = t_df.index[0] + feature_md5 = f'meta_{t_df.index[0]}' return pd.read_sql(feature_md5, con=self._engine) def store_metadata(self, meta): @@ -117,7 +117,16 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): def store_df(self, df, meta): table_name = self.store_metadata(meta) - # TODO: Check that the element is in the index of the dataframe + # Check that the index generated by meta is in the dataframe + n_rows = len(df) + _, _, idx = process_meta( # type: ignore + meta, return_idx=True, n_rows=n_rows) + if any(x not in df.index.names for x in idx.names): # type: ignore + raise ValueError( + 'The index of the dataframe does not match the ' + 'index of the meta data. This happens when the element ' + 'is not part of the index.') + # Save self._save_upsert(df, table_name) def store_timeseries(self, data, meta): diff --git a/junifer/storage/tests/test_sqlite.py b/junifer/storage/tests/test_sqlite.py index 0c4d57083..47593cc1a 100644 --- a/junifer/storage/tests/test_sqlite.py +++ b/junifer/storage/tests/test_sqlite.py @@ -80,7 +80,6 @@ def test_upsert_ignore(): # Test the internals table_name = storage.store_metadata(meta) - c_df1 = _read_sql(table_name, uri=uri, index_col=['pk1', 'pk2']) assert_frame_equal(df1, c_df1) @@ -125,8 +124,17 @@ def test_store_read_df(): 'element': 'test', 'version': '0.0.1', 'marker': {'name': 'fcname'}} + to_store = df1[['col1', 'col2']] + # Save to SQL - storage.store_df(df1, meta) + with pytest.raises(ValueError, match=r"index of the dataframe"): + storage.store_df(to_store, meta) + + _, _, idx = process_meta( # type: ignore + meta, return_idx=True, n_rows=len(to_store)) + to_store = to_store.set_index(idx) + + storage.store_df(to_store, meta) # Test the internals table_name = storage.store_metadata(meta) @@ -149,9 +157,7 @@ def test_store_read_df(): read_df1 = storage.read_df(feature_md5=feature_md5) read_df2 = storage.read_df(feature_name='fcname') assert_frame_equal(read_df1, read_df2) - assert_frame_equal(read_df1, df1) - - + assert_frame_equal(read_df1, to_store.reset_index()) def test_store_table(): -- 2.52.0 From 69bff67b1d776f27629167c7b2a493f79e5363b6 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Tue, 15 Mar 2022 11:59:48 +0100 Subject: [PATCH 027/287] SQLITE Element based storage by default --- examples/norun_ukbvm_gmd.py | 6 +- junifer/api/decorators.py | 2 +- junifer/api/registry.py | 106 +++++++++++++++++++++++++++ junifer/api/run.py | 59 ++++++++------- junifer/storage/base.py | 19 ++++- junifer/storage/sqlite.py | 45 ++++++++---- junifer/storage/tests/test_base.py | 7 ++ junifer/storage/tests/test_sqlite.py | 89 ++++++++++++++-------- 8 files changed, 260 insertions(+), 73 deletions(-) create mode 100644 junifer/api/registry.py diff --git a/examples/norun_ukbvm_gmd.py b/examples/norun_ukbvm_gmd.py index ef1da3544..2cde565c0 100644 --- a/examples/norun_ukbvm_gmd.py +++ b/examples/norun_ukbvm_gmd.py @@ -8,7 +8,7 @@ License: BSD 3 clause """ -from junifer.api import run_pipeline +from junifer.api import run markers = [ {'name': 'Schaefer1000x7_TrimMean80', @@ -26,10 +26,10 @@ markers = [ 'method': 'std'} ] -run_pipeline( +run( workdir='/tmp', datagrabber='JuselessUKBVBM', - element=('sub-1627474', 'ses-2'), + elements=('sub-1627474', 'ses-2'), markers=markers, storage='SQLDataFrameStorage', storage_params={'outpath': '/data/project/juniferexample'}, diff --git a/junifer/api/decorators.py b/junifer/api/decorators.py index 4808fe1ec..91c40af9e 100644 --- a/junifer/api/decorators.py +++ b/junifer/api/decorators.py @@ -1,7 +1,7 @@ # Authors: Federico Raimondo # Leonard Sasse # License: AGPL -from .run import register +from .registry import register def register_datagrabber(klass): diff --git a/junifer/api/registry.py b/junifer/api/registry.py new file mode 100644 index 000000000..3de096073 --- /dev/null +++ b/junifer/api/registry.py @@ -0,0 +1,106 @@ +# Authors: Federico Raimondo +# Leonard Sasse +# License: AGPL +from ..utils.logging import raise_error, logger + +_valid_steps = [ + 'datagrabber', 'datareader', 'preprocessing', 'marker', 'storage'] + +_registry = {x: {} for x in _valid_steps} + + +def register(step, name, klass): + """Register a function to be used in a pipeline step + + Parameters + ---------- + step : str + Name of the step + name : str + Name of the function + klass : class + Class to be registered + """ + if step not in _valid_steps: + raise_error(f'Invalid step: {step}', ValueError) + logger.info(f'Registering {name} in {step}') + _registry[step][name] = klass + + +def get_step_names(step): + """Get the names of the registered functions for a given step + + Parameters + ---------- + step : str + Name of the step + + Returns + ------- + list + List of registered function names + """ + if step not in _valid_steps: + raise_error(f'Invalid step: {step}', ValueError) + return list(_registry[step].keys()) + + +def get(step, name): + """Get the class of the registered function for a given step + + Parameters + ---------- + step : str + Name of the step + name : str + Name of the function + + Returns + ------- + class + Registered function class + """ + if step not in _valid_steps: + raise_error(f'Invalid step: {step}', ValueError) + if name not in _registry[step]: + raise_error(f'Invalid name: {name}', ValueError) + return _registry[step][name] + + +def build(step, name, baseclass, init_params=None): + """Ensure that the given object is an instance of the given class + + Parameters + ---------- + step : str + Name of the step + name : str + Name of the function. + baseclass : class + Class to be checked against + init_parms : dict + Parameters to pass to the class constructor + + Returns + ------- + object + Object if it is an instance of the given class, otherwise a + ValueError is raised + + Raises + ------ + ValueError + If the name is not a string or the object is not an instance of the + baseclass parameter. + """ + if not isinstance(name, str): + raise_error(f'Invalid name: {name}', ValueError) + klass = get(step, name) + if init_params is None: + init_params = {} + object = klass(**init_params) + if not isinstance(object, baseclass): + raise_error( + f'Invalid {step} ({object.__class__.name}). ' + f'Must inherit from {baseclass.name}', ValueError) + return object diff --git a/junifer/api/run.py b/junifer/api/run.py index 35a4c4038..6358ea9cd 100644 --- a/junifer/api/run.py +++ b/junifer/api/run.py @@ -1,34 +1,18 @@ # Authors: Federico Raimondo # Leonard Sasse # License: AGPL -from ..utils.logging import raise_error, logger -_valid_steps = [ - 'datagrabber', 'datareader', 'preprocessing', 'marker', 'storage'] +from pathlib import Path -_registry = {x: {} for x in _valid_steps} - - -def register(step, name, klass): - """Register a function to be used in a pipeline step - - Parameters - ---------- - step : str - Name of the step - name : str - Name of the function - klass : class - Class to be registered - """ - if step not in _valid_steps: - raise_error(f'Invalid step: {step}', ValueError) - logger.info(f'Registering {name} in {step}') - _registry[step][name] = klass +from .registry import build +from ..datagrabber.base import BaseDataGrabber +from ..markers.base import BaseMarker +from ..storage.base import BaseFeatureStorage +from ..markers.collection import MarkerCollection def run( - workdir, datagrabber, element, markers, storage, source_params=None, + workdir, datagrabber, elements, markers, storage, source_params=None, storage_params=None): """Run the pipeline on the selected element @@ -38,8 +22,8 @@ def run( Directory where the pipeline will be executed datagrabber : str Name of the datagrabber to use - element : str - Name of the element to process. Will be used to index the datagrabber. + elements : str, tuple or list[str or tuple] + Element(s) to process. Will be used to index the datagrabber. markers : list of dict List of markers to extract. Each marker is a dict with at least two keys: 'name' and 'kind'. The 'name' key is used to name the output @@ -58,3 +42,28 @@ def run( if storage_params is None: storage_params = {} + + if isinstance(workdir, str): + workdir = Path(workdir) + + datagrabber = build( + 'datagrabber', datagrabber, BaseDataGrabber, init_params=source_params) + + built_markers = [] + for t_marker in markers: + kind = t_marker.pop('kind') + t_m = build('marker', kind, BaseMarker, init_params=t_marker) + built_markers.append(t_m) + + storage = build( + 'storage', storage, BaseFeatureStorage, init_params=storage_params) + + mc = MarkerCollection(markers, storage=storage) + + with datagrabber: + if elements is not None: + for t_element in elements: + mc.fit(datagrabber[t_element]) + else: + for t_element in datagrabber: + mc.fit(datagrabber[t_element]) diff --git a/junifer/storage/base.py b/junifer/storage/base.py index 5ff8a9993..28778afdc 100644 --- a/junifer/storage/base.py +++ b/junifer/storage/base.py @@ -31,7 +31,7 @@ def process_meta(meta, return_idx=False, n_rows=1, rows_col_name=None): The md5 hash of the meta meta : dict The metadata processed for storage - idx : pd.MultiIndex + idx : pd.MultiIndex The pandas index (if return_idx is True) """ if meta is None: @@ -122,13 +122,28 @@ def _element_to_index(meta, n_rows=1, rows_col_name=None): return index +def element_to_prefix(element): + prefix = 'element' + if isinstance(element, str): + prefix = f'{prefix}_{element}' + elif isinstance(element, tuple): + prefix = f"{prefix}_{'_'.join(element)}" + elif isinstance(element, dict): + prefix = f"{prefix}_{'_'.join(element.values())}" + else: + raise ValueError(f'Cannot convert element {element} to prefix. ' + 'Must be a str, tuple or dict') + return f'{prefix}_' + + class BaseFeatureStorage(ABC): """ Base class for feature storage. """ - def __init__(self, uri): + def __init__(self, uri, single_output=False): self.uri = uri + self.single_output = single_output def get_meta(self): meta = {} diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py index c794d9a99..305729966 100644 --- a/junifer/storage/sqlite.py +++ b/junifer/storage/sqlite.py @@ -1,12 +1,13 @@ # Authors: Federico Raimondo # License: AGPL +from pathlib import Path import pandas as pd from pandas.core.base import NoNewAttributesMixin from pandas.io.sql import pandasSQL_builder from sqlalchemy import create_engine, inspect from ..api.decorators import register_storage -from .base import PandasFeatureStoreage, process_meta +from .base import PandasFeatureStoreage, process_meta, element_to_prefix from ..utils.logging import warn @@ -16,7 +17,7 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): SQLite feature storage. """ - def __init__(self, uri, upsert='update'): + def __init__(self, uri, single_output=False, upsert='update'): """Initialise an SQLite feature storage Parameters @@ -37,8 +38,9 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): """ if upsert not in ['update', 'ignore']: raise ValueError('upsert must be either "update" or "ignore"') - super().__init__(uri) - self._engine = create_engine(uri, echo=False) + if not isinstance(uri, Path): + uri = Path(uri) + super().__init__(uri, single_output=single_output) self._upsert = upsert self._valid_inputs = ['table', 'timeseries'] @@ -47,8 +49,23 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): input = [input] return all(x in self._valid_inputs for x in input) + def get_engine(self, meta=None): + if meta is None: + meta = {} + element = meta.get('element', None) + if self.single_output is False and element is None: + raise ValueError( + 'element must be specified when single_output is False') + prefix = '' + if self.single_output is False: + prefix = element_to_prefix(element) + + uri = f'sqlite:///{self.uri.parent}/{prefix}{self.uri.name}' + return create_engine(uri, echo=False) + def list_features(self): - meta_df = pd.read_sql('meta', con=self._engine, index_col='meta_md5') + meta_df = pd.read_sql( + 'meta', con=self.get_engine(), index_col='meta_md5') return meta_df.to_dict(orient='index') def read_df(self, feature_name=None, feature_md5=None): @@ -68,6 +85,7 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): pandas.DataFrame The features. """ + engine = self.get_engine() if feature_md5 is not None and feature_name is not None: raise ValueError('Only one of feature_name or feature_md5 can be ' 'specified') @@ -78,7 +96,7 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): feature_md5 = f'meta_{feature_md5}' else: meta_df = pd.read_sql( - 'meta', con=self._engine, index_col='meta_md5') + 'meta', con=engine, index_col='meta_md5') t_df = meta_df.query(f"name == '{feature_name}'") if len(t_df) == 0: raise ValueError(f'Feature {feature_name} not found') @@ -88,14 +106,14 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): 'This file is invalid. You can bypass this issue by ' 'specifying a feature_md5') feature_md5 = f'meta_{t_df.index[0]}' - return pd.read_sql(feature_md5, con=self._engine) + return pd.read_sql(feature_md5, con=engine) def store_metadata(self, meta): t_meta = meta.copy() t_meta.update(self.get_meta()) meta_md5, t_meta = process_meta( # type: ignore t_meta, return_idx=False) - if meta_md5 not in inspect(self._engine).get_table_names(): + if meta_md5 not in inspect(self.get_engine(meta)).get_table_names(): meta_df = self._meta_row(t_meta, meta_md5) self._save_upsert(meta_df, 'meta') return f'meta_{meta_md5}' @@ -117,10 +135,10 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): def store_df(self, df, meta): table_name = self.store_metadata(meta) - # Check that the index generated by meta is in the dataframe - n_rows = len(df) + # Check that the index generated by meta is in the dataframe, at least + # for one row _, _, idx = process_meta( # type: ignore - meta, return_idx=True, n_rows=n_rows) + meta, return_idx=True, n_rows=1) if any(x not in df.index.names for x in idx.names): # type: ignore raise ValueError( 'The index of the dataframe does not match the ' @@ -153,11 +171,12 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): If the table exists and if_exist is 'fail' """ index_col = df.index.names - with self._engine.begin() as con: + engine = self.get_engine() + with engine.begin() as con: if if_exist == 'replace': # Case 1: replace all the existing elements df.to_sql(name, con=con, if_exists='replace') - elif not inspect(self._engine).has_table(name): + elif not inspect(engine).has_table(name): # Case 2: new table, so no big issue df.to_sql(name, con=con, if_exists='append') else: diff --git a/junifer/storage/tests/test_base.py b/junifer/storage/tests/test_base.py index e051b27aa..01f2156f8 100644 --- a/junifer/storage/tests/test_base.py +++ b/junifer/storage/tests/test_base.py @@ -133,6 +133,9 @@ def test_BaseFeatureStorage(): BaseFeatureStorage(uri='/tmp') # type: ignore class MyFeatureStorage(BaseFeatureStorage): + def __init__(self, uri, single_output=False): + super().__init__(uri, single_output=single_output) + def validate(self, input): super().validate(input) @@ -159,6 +162,10 @@ def test_BaseFeatureStorage(): super().store_timeseries(timeseries, meta) st = MyFeatureStorage(uri='/tmp') + assert st.single_output is False + + st = MyFeatureStorage(uri='/tmp', single_output=True) + assert st.single_output is True with pytest.raises(NotImplementedError): st.validate(None) diff --git a/junifer/storage/tests/test_sqlite.py b/junifer/storage/tests/test_sqlite.py index 47593cc1a..ab4c378e1 100644 --- a/junifer/storage/tests/test_sqlite.py +++ b/junifer/storage/tests/test_sqlite.py @@ -9,45 +9,73 @@ from junifer.storage.base import process_meta df1 = pd.DataFrame({ - 'pk1': [1, 2, 3, 4, 5], + 'element': [1, 2, 3, 4, 5], 'pk2': ['a', 'b', 'c', 'd', 'e'], 'col1': [11, 22, 33, 44, 55], 'col2': [111, 222, 333, 444, 555] -}).set_index(['pk1', 'pk2']) +}).set_index(['element', 'pk2']) df2 = pd.DataFrame({ - 'pk1': [2, 5, 6], + 'element': [2, 5, 6], 'pk2': ['b', 'e', 'f'], 'col1': [2222, 5555, 66], 'col2': [22222, 55555, 666] -}).set_index(['pk1', 'pk2']) +}).set_index(['element', 'pk2']) df_update = pd.DataFrame({ - 'pk1': [1, 2, 3, 4, 5, 6], + 'element': [1, 2, 3, 4, 5, 6], 'pk2': ['a', 'b', 'c', 'd', 'e', 'f'], 'col1': [11, 2222, 33, 44, 5555, 66], 'col2': [111, 22222, 333, 444, 55555, 666] -}).set_index(['pk1', 'pk2']) +}).set_index(['element', 'pk2']) df_ignore = pd.DataFrame({ - 'pk1': [1, 2, 3, 4, 5, 6], + 'element': [1, 2, 3, 4, 5, 6], 'pk2': ['a', 'b', 'c', 'd', 'e', 'f'], 'col1': [11, 22, 33, 44, 55, 66], 'col2': [111, 222, 333, 444, 555, 666] -}).set_index(['pk1', 'pk2']) +}).set_index(['element', 'pk2']) def _read_sql(table_name, uri, index_col): - engine = create_engine(uri, echo=False) + engine = create_engine(f'sqlite:///{uri}', echo=False) df = pd.read_sql(table_name, con=engine, index_col=index_col) return df -def test_upsert_replace(): - """Test store_df (if_exist=replace)""" +def test_get_engine(): + """Test get_engine""" with tempfile.TemporaryDirectory() as _tmpdir: - uri = f'sqlite:///{_tmpdir}/test.db' - storage = SQLiteFeatureStorage(uri=uri, upsert='ignore') + uri = f'{_tmpdir}/test.db' + # Single storage, must be the uri + storage = SQLiteFeatureStorage( + uri=uri, single_output=True, upsert='ignore') + assert storage.single_output is True + engine = storage.get_engine() + assert engine.url.drivername == 'sqlite' + assert f'{engine.url.database}' == uri + + +def test_store_metadata(): + """Test store_metadata""" + with tempfile.TemporaryDirectory() as _tmpdir: + uri = f'{_tmpdir}/test.db' + # Single storage, must be the uri + storage = SQLiteFeatureStorage( + uri=uri, single_output=True, upsert='ignore') + meta = {'element': 'test', 'version': '0.0.1'} + table_name = storage.store_metadata(meta) + assert table_name.startswith('meta_') + + + +def test_upsert_replace(): + """Test store_df (upsert=replace)""" + with tempfile.TemporaryDirectory() as _tmpdir: + uri = f'{_tmpdir}/test.db' + # Single storage, must be the uri + storage = SQLiteFeatureStorage( + uri=uri, single_output=True, upsert='ignore') meta = {'element': 'test', 'version': '0.0.1'} # Save to SQL @@ -56,23 +84,24 @@ def test_upsert_replace(): # Test the internals table_name = storage.store_metadata(meta) - c_df1 = _read_sql(table_name, uri=uri, index_col=['pk1', 'pk2']) + c_df1 = _read_sql(table_name, uri=uri, index_col=['element', 'pk2']) assert_frame_equal(df1, c_df1) storage._save_upsert(df2, table_name, if_exist='replace') - c_df2 = _read_sql(table_name, uri=uri, index_col=['pk1', 'pk2']) + c_df2 = _read_sql(table_name, uri=uri, index_col=['element', 'pk2']) assert_frame_equal(df2, c_df2) def test_upsert_ignore(): """Test store_df (upsert=ignore)""" with tempfile.TemporaryDirectory() as _tmpdir: - uri = f'sqlite:///{_tmpdir}/test.db' + uri = f'{_tmpdir}/test.db' with pytest.raises(ValueError): - SQLiteFeatureStorage(uri=uri, upsert='wrong') + SQLiteFeatureStorage(uri=uri, single_output=True, upsert='wrong') - storage = SQLiteFeatureStorage(uri=uri, upsert='ignore') + storage = SQLiteFeatureStorage( + uri=uri, single_output=True, upsert='ignore') meta = {'element': 'test', 'version': '0.0.1'} # Save to SQL @@ -81,12 +110,13 @@ def test_upsert_ignore(): # Test the internals table_name = storage.store_metadata(meta) - c_df1 = _read_sql(table_name, uri=uri, index_col=['pk1', 'pk2']) + c_df1 = _read_sql(table_name, uri=uri, index_col=['element', 'pk2']) assert_frame_equal(df1, c_df1) storage.store_df(df2, meta) - c_dfignore = _read_sql(table_name, uri=uri, index_col=['pk1', 'pk2']) + c_dfignore = _read_sql( + table_name, uri=uri, index_col=['element', 'pk2']) assert_frame_equal(c_dfignore, df_ignore) with pytest.raises(ValueError, match=r"already exists"): @@ -97,8 +127,8 @@ def test_upsert_update(): """Test store_df (upsert=delete)""" meta = {'element': 'test', 'version': '0.0.1'} with tempfile.TemporaryDirectory() as _tmpdir: - uri = f'sqlite:///{_tmpdir}/test.db' - storage = SQLiteFeatureStorage(uri=uri) + uri = f'{_tmpdir}/test.db' + storage = SQLiteFeatureStorage(uri=uri, single_output=True) # Save to SQL storage.store_df(df1, meta) @@ -106,20 +136,21 @@ def test_upsert_update(): # Test the internals table_name = storage.store_metadata(meta) - c_df1 = _read_sql(table_name, uri=uri, index_col=['pk1', 'pk2']) + c_df1 = _read_sql(table_name, uri=uri, index_col=['element', 'pk2']) assert_frame_equal(df1, c_df1) storage.store_df(df2, meta) - c_dfupdate = _read_sql(table_name, uri=uri, index_col=['pk1', 'pk2']) + c_dfupdate = _read_sql(table_name, uri=uri, index_col=['element', 'pk2']) assert_frame_equal(c_dfupdate, df_update) def test_store_read_df(): """Test store_df""" with tempfile.TemporaryDirectory() as _tmpdir: - uri = f'sqlite:///{_tmpdir}/test.db' - storage = SQLiteFeatureStorage(uri=uri, upsert='ignore') + uri = f'{_tmpdir}/test.db' + storage = SQLiteFeatureStorage( + uri=uri, single_output=True, upsert='ignore') meta = { 'element': 'test', 'version': '0.0.1', 'marker': {'name': 'fcname'}} @@ -128,7 +159,7 @@ def test_store_read_df(): # Save to SQL with pytest.raises(ValueError, match=r"index of the dataframe"): - storage.store_df(to_store, meta) + storage.store_df(to_store.set_index('col1'), meta) _, _, idx = process_meta( # type: ignore meta, return_idx=True, n_rows=len(to_store)) @@ -164,8 +195,8 @@ def test_store_table(): """Test store_table""" meta = {'element': 'test', 'version': '0.0.1', 'marker': {'name': 'fc'}} with tempfile.TemporaryDirectory() as _tmpdir: - uri = f'sqlite:///{_tmpdir}/test.db' - storage = SQLiteFeatureStorage(uri=uri) + uri = f'{_tmpdir}/test.db' + storage = SQLiteFeatureStorage(uri=uri, single_output=True) data = [ [1, 10], [2, 20], -- 2.52.0 From 8c3302b54d64dec8aa698d715a1bd0344a033a8e Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Tue, 15 Mar 2022 15:34:00 +0100 Subject: [PATCH 028/287] Multiple output tested --- junifer/storage/base.py | 6 +- junifer/storage/sqlite.py | 17 +++--- junifer/storage/tests/test_sqlite.py | 82 +++++++++++++++++++++++++++- 3 files changed, 93 insertions(+), 12 deletions(-) diff --git a/junifer/storage/base.py b/junifer/storage/base.py index 28778afdc..c9a326d6e 100644 --- a/junifer/storage/base.py +++ b/junifer/storage/base.py @@ -6,6 +6,7 @@ import json import hashlib from abc import ABC, abstractmethod +from ..utils import logger from .. import __version__ @@ -51,8 +52,9 @@ def process_meta(meta, return_idx=False, n_rows=1, rows_col_name=None): t_meta['_element_keys'] = list(element.keys()) else: t_meta['_element_keys'] = ['element'] - - md5_hash = _meta_hash(meta) + logger.debug(f'Hasing meta {t_meta}') + md5_hash = _meta_hash(t_meta) + logger.debug(f'Hash computed: {md5_hash}') if return_idx is True: out = md5_hash, t_meta, idx else: diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py index 305729966..d872d9978 100644 --- a/junifer/storage/sqlite.py +++ b/junifer/storage/sqlite.py @@ -111,11 +111,12 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): def store_metadata(self, meta): t_meta = meta.copy() t_meta.update(self.get_meta()) - meta_md5, t_meta = process_meta( # type: ignore + meta_md5, t_meta_row = process_meta( # type: ignore t_meta, return_idx=False) - if meta_md5 not in inspect(self.get_engine(meta)).get_table_names(): - meta_df = self._meta_row(t_meta, meta_md5) - self._save_upsert(meta_df, 'meta') + engine = self.get_engine(t_meta) + if meta_md5 not in inspect(engine).get_table_names(): + meta_df = self._meta_row(t_meta_row, meta_md5) + self._save_upsert(meta_df, 'meta', engine) return f'meta_{meta_md5}' def store_matrix2d( @@ -145,12 +146,13 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): 'index of the meta data. This happens when the element ' 'is not part of the index.') # Save - self._save_upsert(df, table_name) + engine = self.get_engine(meta) + self._save_upsert(df, table_name, engine) def store_timeseries(self, data, meta): raise NotImplementedError('store_timeseries not implemented') - def _save_upsert(self, df, name, if_exist='append'): + def _save_upsert(self, df, name, engine=None, if_exist='append'): """ Implementation of UPSERT functionality. Parameters @@ -171,7 +173,8 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): If the table exists and if_exist is 'fail' """ index_col = df.index.names - engine = self.get_engine() + if engine is None: + engine = self.get_engine() with engine.begin() as con: if if_exist == 'replace': # Case 1: replace all the existing elements diff --git a/junifer/storage/tests/test_sqlite.py b/junifer/storage/tests/test_sqlite.py index ab4c378e1..31abd542d 100644 --- a/junifer/storage/tests/test_sqlite.py +++ b/junifer/storage/tests/test_sqlite.py @@ -1,3 +1,5 @@ +from pathlib import Path +import numpy as np import pandas as pd from pandas.testing import assert_frame_equal import tempfile @@ -5,7 +7,7 @@ from sqlalchemy import create_engine import pytest from junifer.storage.sqlite import SQLiteFeatureStorage -from junifer.storage.base import process_meta +from junifer.storage.base import process_meta, element_to_prefix df1 = pd.DataFrame({ @@ -68,7 +70,6 @@ def test_store_metadata(): assert table_name.startswith('meta_') - def test_upsert_replace(): """Test store_df (upsert=replace)""" with tempfile.TemporaryDirectory() as _tmpdir: @@ -141,7 +142,8 @@ def test_upsert_update(): storage.store_df(df2, meta) - c_dfupdate = _read_sql(table_name, uri=uri, index_col=['element', 'pk2']) + c_dfupdate = _read_sql( + table_name, uri=uri, index_col=['element', 'pk2']) assert_frame_equal(c_dfupdate, df_update) @@ -234,3 +236,77 @@ def test_store_table(): c_df2 = _read_sql(table_name, uri=uri, index_col=['element', 'scan']) assert_frame_equal(df2, c_df2) + + +def test_store_multiple_output(): + """Test storing using single_output=False""" + + meta1 = {'element': {'subject': 'test-01', 'session': 'ses-01'}, + 'version': '0.0.1', 'marker': {'name': 'fc'}} + meta2 = {'element': {'subject': 'test-02', 'session': 'ses-01'}, + 'version': '0.0.1', 'marker': {'name': 'fc'}} + meta3 = {'element': {'subject': 'test-01', 'session': 'ses-02'}, + 'version': '0.0.1', 'marker': {'name': 'fc'}} + with tempfile.TemporaryDirectory() as _tmpdir: + uri = Path(f'{_tmpdir}/test.db') + storage = SQLiteFeatureStorage(uri=uri, single_output=False) + data1 = np.array([ + [1, 10], + [2, 20], + [3, 30], + [4, 40], + [5, 50], + ]) + + data2 = data1 * 10 + data3 = data1 * 20 + + hash1, _, idx1 = process_meta( # type: ignore + meta1, return_idx=True, n_rows=5, rows_col_name='scan') + df1 = pd.DataFrame(data1, columns=['f1', 'f2'], index=idx1) + + hash2, _, idx2 = process_meta( # type: ignore + meta2, return_idx=True, n_rows=5, rows_col_name='scan') + df2 = pd.DataFrame(data2, columns=['f1', 'f2'], index=idx2) + + hash3, _, idx3 = process_meta( # type: ignore + meta3, return_idx=True, n_rows=5, rows_col_name='scan') + df3 = pd.DataFrame(data3, columns=['f1', 'f2'], index=idx3) + + assert hash1 == hash2 + assert hash2 == hash3 + + storage.store_table( + data1, meta1, columns=['f1', 'f2'], rows_col_name='scan') + + storage.store_table( + data2, meta2, columns=['f1', 'f2'], rows_col_name='scan') + + storage.store_table( + data3, meta3, columns=['f1', 'f2'], rows_col_name='scan') + + assert not uri.exists() + + prefix1 = element_to_prefix(meta1['element']) + prefix2 = element_to_prefix(meta2['element']) + prefix3 = element_to_prefix(meta3['element']) + + uri1 = uri.parent / f'{prefix1}{uri.name}' + uri2 = uri.parent / f'{prefix2}{uri.name}' + uri3 = uri.parent / f'{prefix3}{uri.name}' + + assert uri1.exists() + assert uri2.exists() + assert uri3.exists() + + table_name = storage.store_metadata(meta1) + + cols = ['subject', 'session', 'scan'] + + cdf1 = _read_sql(table_name, uri1, index_col=cols) + cdf2 = _read_sql(table_name, uri2, index_col=cols) + cdf3 = _read_sql(table_name, uri3, index_col=cols) + + assert_frame_equal(df1, cdf1) + assert_frame_equal(df2, cdf2) + assert_frame_equal(df3, cdf3) -- 2.52.0 From aa855d99f469ccca014e8db06e971f44333d9bef Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Tue, 15 Mar 2022 15:49:22 +0100 Subject: [PATCH 029/287] The new example "runs", still don't know if it works --- examples/run_run_gmd_mean.py | 41 ++++++++++++++++++++++++++++++++++++ junifer/api/run.py | 6 +++--- junifer/markers/parcel.py | 2 ++ junifer/testing.py | 10 ++++++++- 4 files changed, 55 insertions(+), 4 deletions(-) create mode 100644 examples/run_run_gmd_mean.py diff --git a/examples/run_run_gmd_mean.py b/examples/run_run_gmd_mean.py new file mode 100644 index 000000000..db461246d --- /dev/null +++ b/examples/run_run_gmd_mean.py @@ -0,0 +1,41 @@ +""" +UKB VBM GMD Extraction +====================== + +Authors: Federico Raimondo + +License: BSD 3 clause +""" +import tempfile + +from junifer.api import run +from junifer.testing import register_testing + + +register_testing() + +markers = [ + {'name': 'Schaefer1000x7_TrimMean80', + 'kind': 'ParcelAggregation', + 'atlas': 'Schaefer1000x7', + 'method': 'trim_mean', + 'method_params': {'proportiontocut': 0.2}}, + {'name': 'Schaefer1000x7_Mean', + 'kind': 'ParcelAggregation', + 'atlas': 'Schaefer1000x7', + 'method': 'mean'}, + {'name': 'Schaefer1000x7_Std', + 'kind': 'ParcelAggregation', + 'atlas': 'Schaefer1000x7', + 'method': 'std'} +] + +with tempfile.TemporaryDirectory() as tmpdir: + uri = f'{tmpdir}/test.db' + run( + workdir='/tmp', + datagrabber='OasisVBMTestingDatagrabber', + markers=markers, + storage='SQLiteFeatureStorage', + storage_params={'uri': uri}, + ) diff --git a/junifer/api/run.py b/junifer/api/run.py index 6358ea9cd..f33737a33 100644 --- a/junifer/api/run.py +++ b/junifer/api/run.py @@ -12,8 +12,8 @@ from ..markers.collection import MarkerCollection def run( - workdir, datagrabber, elements, markers, storage, source_params=None, - storage_params=None): + workdir, datagrabber, markers, storage, source_params=None, + storage_params=None, elements=None): """Run the pipeline on the selected element Parameters @@ -58,7 +58,7 @@ def run( storage = build( 'storage', storage, BaseFeatureStorage, init_params=storage_params) - mc = MarkerCollection(markers, storage=storage) + mc = MarkerCollection(built_markers, storage=storage) with datagrabber: if elements is not None: diff --git a/junifer/markers/parcel.py b/junifer/markers/parcel.py index 0f1beb09c..79e23d405 100644 --- a/junifer/markers/parcel.py +++ b/junifer/markers/parcel.py @@ -8,8 +8,10 @@ from nilearn.image import resample_to_img, math_img from .base import BaseMarker from ..stats import get_aggfunc_by_name from ..data import load_atlas +from ..api.decorators import register_marker +@register_marker class ParcelAggregation(BaseMarker): def __init__(self, atlas, method, method_params=None, on=None, name=None): if on is None: diff --git a/junifer/testing.py b/junifer/testing.py index 46dca9f5e..184754786 100644 --- a/junifer/testing.py +++ b/junifer/testing.py @@ -4,6 +4,14 @@ import tempfile from nilearn import datasets from .datagrabber.base import BaseDataGrabber +from .api.registry import register + + +def register_testing(): + """Register testing datagrabber""" + register( + 'datagrabber', 'OasisVBMTestingDatagrabber', + OasisVBMTestingDatagrabber) class OasisVBMTestingDatagrabber(BaseDataGrabber): @@ -20,7 +28,7 @@ class OasisVBMTestingDatagrabber(BaseDataGrabber): def __getitem__(self, element): out = {} - out['VBM_GM'] = self._dataset.gray_matter_maps[element] + out['VBM_GM'] = self._dataset.gray_matter_maps[element - 1] out['meta'] = {'subject': element} return out -- 2.52.0 From 33b94da4b1ecb3d6063526730d107b21d93e0753 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Tue, 15 Mar 2022 15:49:36 +0100 Subject: [PATCH 030/287] flake --- junifer/testing.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/junifer/testing.py b/junifer/testing.py index 184754786..1a20c8695 100644 --- a/junifer/testing.py +++ b/junifer/testing.py @@ -10,7 +10,7 @@ from .api.registry import register def register_testing(): """Register testing datagrabber""" register( - 'datagrabber', 'OasisVBMTestingDatagrabber', + 'datagrabber', 'OasisVBMTestingDatagrabber', OasisVBMTestingDatagrabber) -- 2.52.0 From b4b53e7822081c487a2cc654144048e05ce92f12 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Wed, 16 Mar 2022 14:39:27 +0100 Subject: [PATCH 031/287] Fixed tests, maybe it's ready to run? --- junifer/data/atlases.py | 2 +- junifer/markers/parcel.py | 4 ++-- junifer/markers/tests/test_base_marker.py | 4 ++-- junifer/markers/tests/test_collection.py | 27 ++++++++++++++++++----- junifer/testing.py | 2 +- 5 files changed, 28 insertions(+), 11 deletions(-) diff --git a/junifer/data/atlases.py b/junifer/data/atlases.py index 33cd4747f..519cbced5 100644 --- a/junifer/data/atlases.py +++ b/junifer/data/atlases.py @@ -293,7 +293,7 @@ def _retrieve_schaefer(atlas_dir, resolution, n_rois=None, yeo_networks=7): 'At least one of the atlas files is missing. ' 'Fetching using nilearn.') datasets.fetch_atlas_schaefer_2018( - n_rois=n_rois, + n_rois=n_rois, # type: ignore yeo_networks=yeo_networks, resolution_mm=resolution, data_dir=atlas_dir.as_posix()) diff --git a/junifer/markers/parcel.py b/junifer/markers/parcel.py index 79e23d405..6ede62d7e 100644 --- a/junifer/markers/parcel.py +++ b/junifer/markers/parcel.py @@ -22,13 +22,13 @@ class ParcelAggregation(BaseMarker): self.method_params = {} if method_params is None else method_params def get_output_kind(self, input): - if input in ['GMD_GM', 'GMD_WM']: + if input in ['VBM_GM', 'VBM_WM']: return 'table' if input in ['BOLD']: return 'timeseries' def store(self, kind, out, storage): - if kind in ['GMD_GM', 'GMD_WM']: + if kind in ['VBM_GM', 'VBM_WM']: storage.store_table(**out) if kind in ['BOLD']: storage.store_timeseries(**out) diff --git a/junifer/markers/tests/test_base_marker.py b/junifer/markers/tests/test_base_marker.py index cf9ed5298..f661af6df 100644 --- a/junifer/markers/tests/test_base_marker.py +++ b/junifer/markers/tests/test_base_marker.py @@ -20,13 +20,13 @@ def test_meta(): base = BaseMarker(on=['bold', 'dwi']) - t_meta = base.get_meta() + t_meta = base.get_meta('bold') assert t_meta['marker']['class'] == 'BaseMarker' assert t_meta['marker']['name'] == 'BaseMarker' base = BaseMarker(on=['bold', 'dwi'], name='mymarker') - t_meta = base.get_meta() + t_meta = base.get_meta('dwi') assert t_meta['marker']['name'] == 'mymarker' diff --git a/junifer/markers/tests/test_collection.py b/junifer/markers/tests/test_collection.py index b68d2a7fd..969f356e1 100644 --- a/junifer/markers/tests/test_collection.py +++ b/junifer/markers/tests/test_collection.py @@ -99,8 +99,8 @@ def test_MarkerCollection_storage(): # Test storage dg = OasisVBMTestingDatagrabber() with tempfile.TemporaryDirectory() as tmpdir: - uri = f'sqlite:///{tmpdir}/test.db' - storage = SQLiteFeatureStorage(uri=uri) + uri = f'{tmpdir}/test.db' + storage = SQLiteFeatureStorage(uri=uri, single_output=True) mc = MarkerCollection( markers=markers, storage=storage, datareader=DefaultDataReader()) mc.validate(dg) @@ -120,7 +120,24 @@ def test_MarkerCollection_storage(): out = mc2.fit(input) features = storage.list_features() - assert len(features) == 1 - feature_md5 = features.keys()[0] + assert len(features) == 3 + feature_md5 = list(features.keys())[0] t_feature = storage.read_df(feature_md5=feature_md5) - + fname = 'gmd_schaefer100x7_mean' + t_data = out[fname]['VBM_GM']['data'] # type: ignore + cols = out[fname]['VBM_GM']['columns'] # type: ignore + assert_array_equal(t_feature[cols].values, t_data) # type: ignore + + feature_md5 = list(features.keys())[1] + t_feature = storage.read_df(feature_md5=feature_md5) + fname = 'gmd_schaefer100x7_std' + t_data = out[fname]['VBM_GM']['data'] # type: ignore + cols = out[fname]['VBM_GM']['columns'] # type: ignore + assert_array_equal(t_feature[cols].values, t_data) # type: ignore + + feature_md5 = list(features.keys())[2] + t_feature = storage.read_df(feature_md5=feature_md5) + fname = 'gmd_schaefer100x7_trim_mean90' + t_data = out[fname]['VBM_GM']['data'] # type: ignore + cols = out[fname]['VBM_GM']['columns'] # type: ignore + assert_array_equal(t_feature[cols].values, t_data) # type: ignore diff --git a/junifer/testing.py b/junifer/testing.py index 1a20c8695..999f58f69 100644 --- a/junifer/testing.py +++ b/junifer/testing.py @@ -29,7 +29,7 @@ class OasisVBMTestingDatagrabber(BaseDataGrabber): def __getitem__(self, element): out = {} out['VBM_GM'] = self._dataset.gray_matter_maps[element - 1] - out['meta'] = {'subject': element} + out['meta'] = {'element': {'subject': element}} return out def __enter__(self): -- 2.52.0 From 0d66f7015e4ef468431a9421b27de816d5175c3a Mon Sep 17 00:00:00 2001 From: Vera Komeyer Date: Wed, 16 Mar 2022 17:02:26 +0100 Subject: [PATCH 032/287] Add retrieval of Tian atlases --- junifer/data/atlases.py | 160 +++++++++++++++++++++++++++++++++++++++- 1 file changed, 159 insertions(+), 1 deletion(-) diff --git a/junifer/data/atlases.py b/junifer/data/atlases.py index 33cd4747f..a45fa9943 100644 --- a/junifer/data/atlases.py +++ b/junifer/data/atlases.py @@ -3,7 +3,11 @@ # License: AGPL from pathlib import Path import io +import os import requests +import wget +import shutil +import zipfile import numpy as np import pandas as pd @@ -44,6 +48,36 @@ for n_rois in range(100, 1001, 100): 'yeo_networks': t_net, } +for scale in range(1, 5): + for field in ['3T', '7T']: + if field == '7T': + space = 'MNI6thgeneration' + t_name = f'Tian{scale}x{field}x{space}' + _available_atlases[t_name] = { + 'family': 'Tian', + 'scale': scale, + 'magneticfield': field, + 'valid_resolutions': [1.6] + } + else: + for space in ['MNI6thgeneration', 'MNInonlinear2009cAsym']: + if space == 'MNI6thgeneration': + t_name = f'Tian{scale}x{field}x{space}' + _available_atlases[t_name] = { + 'family': 'Tian', + 'scale': scale, + 'magneticfield': field, + 'valid_resolutions': [1, 2] + } + elif space == 'MNInonlinear2009cAsym': + t_name = f'Tian{scale}x{field}x{space}' + _available_atlases[t_name] = { + 'family': 'Tian', + 'scale': scale, + 'magneticfield': field, + 'valid_resolutions': [2] + } + def register_atlas(name, atlas_path, atl_labels, overwrite=False): """Register a custom user atlas. @@ -136,7 +170,15 @@ def load_atlas(name, atlas_dir=None, resolution=None, path_only=False): Number of yeo networks to use. Valid values: 7, 17. Defaults to 7. Tian : - # TODO add + 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 @@ -207,6 +249,16 @@ def _retrieve_atlas(family, atlas_dir=None, resolution=None, **kwargs): (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 @@ -235,6 +287,9 @@ def _retrieve_atlas(family, atlas_dir=None, resolution=None, **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) else: raise_error( f"The provided atlas name {family} cannot be retrieved. ") @@ -312,6 +367,109 @@ def _retrieve_schaefer(atlas_dir, resolution, n_rois=None, yeo_networks=7): return atlas_fname, labels +def _retrieve_tian( + atlas_dir, resolution, scale=None, space='MNI6thgeneration', + magneticfield='3T'): + + # check validity of atlas parameters + _valid_scales = [1, 2, 3, 4] + _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}') + if field not in _valid_fields: + raise_error( + f'The parameter `magneticfield` ({field}) needs to be one of ' + f'the following: {_valid_fields}') + + if magneticfield == '3T': + _valid_spaces = ['MNI6thgeneration', 'MNInonlinear2009cAsym'] + if space == 'MNI6thgeneration': + _valid_resolutions = [1, 2] + elif space == 'MNInonlinear2009cAsym': + _valid_resolutions = [2] + if 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}') + + resolution = _closest_resolution(resolution, _valid_resolutions) + + # 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}') + + # define file names + if magneticfield == '3T': + atlas_fname_base_3T = ( + 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': + atlas_fname = atlas_fname_base_3T / ( + f'Tian_Subcortex_S{scale}_{magneticfield}.nii.gz') + if resolution == 1: + atlas_fname = atlas_fname_base_3T / ( + f'Tian_Subcortex_S{scale}_{magneticfield}_{resolution}' + 'mm.nii.gz') + elif space == 'MNInonlinear2009cAsym': + space = '2009cAsym' + atlas_fname = atlas_fname_base_3T / ( + f'Tian_Subcortex_S{scale}_{magneticfield}_{space}.nii.gz') + elif 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') + # 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)] + atlas_lname = atlas_fname_base_7T / ( + f'Tian_Subcortex_S{scale}_7T_labelnumbering.txt') + with open(atlas_lname, 'w') as filehandle: + for listitem in labels: + 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.') + + # check existance of atlas + if not (atlas_fname.exists() and atlas_lname.exists()): + 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') + + logger.info(f'Downloading {url_basis}') + atlas_download_dir = wget.download(url_basis, atlas_dir.as_posix()) + with zipfile.ZipFile(atlas_download_dir, 'r') as zip_ref: + zip_ref.extractall(atlas_dir.as_posix()) + # clean after unzipping + if os.path.exists(atlas_download_dir): + os.remove(atlas_download_dir) + if os.path.exists((atlas_dir / '__MACOSX').as_posix()): + 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.') + + 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}') -- 2.52.0 From b49c701588fe7a41ab2d5f9927563199119fb032 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Wed, 23 Mar 2022 10:50:29 +0100 Subject: [PATCH 033/287] CLI RUN + COLLECT: missing tests on collect --- .gitignore | 3 +- examples/norun_hcpfc_pearson.py | 26 ++++--- examples/run_run_gmd_mean.py | 17 +++-- examples/yamls/gmd_mean.yaml | 25 ++++++ junifer/api/__init__.py | 3 +- junifer/api/cli.py | 49 ++++++++++++ junifer/api/parser.py | 23 ++++++ junifer/api/run.py | 44 ++++++----- junifer/markers/tests/test_collection.py | 2 +- junifer/storage/base.py | 53 ++++++------- junifer/storage/sqlite.py | 85 ++++++++++++++++----- junifer/storage/tests/test_base.py | 96 ++++++++++++++++++------ junifer/storage/tests/test_sqlite.py | 29 ++++--- junifer/testing/__init__.py | 3 + junifer/testing/datagrabbers.py | 29 +++++++ junifer/testing/registry.py | 6 ++ requirements.txt | 3 +- setup.py | 8 +- 18 files changed, 380 insertions(+), 124 deletions(-) create mode 100644 examples/yamls/gmd_mean.yaml create mode 100644 junifer/api/cli.py create mode 100644 junifer/api/parser.py create mode 100644 junifer/testing/__init__.py create mode 100644 junifer/testing/datagrabbers.py create mode 100644 junifer/testing/registry.py diff --git a/.gitignore b/.gitignore index 09d12056a..35e6aacdc 100644 --- a/.gitignore +++ b/.gitignore @@ -132,4 +132,5 @@ cython_debug/ # OS Stuff .DS_store -junifer/_version.py \ No newline at end of file +junifer/_version.py +scratch/ \ No newline at end of file diff --git a/examples/norun_hcpfc_pearson.py b/examples/norun_hcpfc_pearson.py index d2d17640d..fe2632e92 100644 --- a/examples/norun_hcpfc_pearson.py +++ b/examples/norun_hcpfc_pearson.py @@ -6,7 +6,14 @@ License: BSD 3 clause """ -from junifer.api import run_pipeline +from junifer.api import run + +datagrabber = { + 'kind': 'HCPOpenAccess', + 'modality': 'fMRI', + 'preprocessed': 'ICA+FIX', + 'space': 'volumetric', +} custom_confound_strategy = { 'filter': 'butterworth', @@ -48,18 +55,15 @@ markers = [ 'confound_strategy': custom_confound_strategy} ] -dg_params = { - 'modality': 'fMRI', - 'preprocessed': 'ICA+FIX', - 'space': 'volumetric', +storage = { + 'kind': 'SQLiteFeatureStorage', + 'uri': '/data/project/juniferexample' } -run_pipeline( +run( workdir='/tmp', - datagrabber='HCPOpenAccess', - datagrabber_params=dg_params, - element=('100408', 'REST1', "LR"), + datagrabber=datagrabber, + elements=[('100408', 'REST1', "LR")], markers=markers, - storage='SQLDataFrameStorage', - storage_params={'outpath': '/data/project/juniferexample'}, + storage=storage, ) diff --git a/examples/run_run_gmd_mean.py b/examples/run_run_gmd_mean.py index db461246d..6259530b7 100644 --- a/examples/run_run_gmd_mean.py +++ b/examples/run_run_gmd_mean.py @@ -9,10 +9,11 @@ License: BSD 3 clause import tempfile from junifer.api import run -from junifer.testing import register_testing +import junifer.testing.registry # noqa: F401 - -register_testing() +datagrabber = { + 'kind': 'OasisVBMTestingDatagrabber', +} markers = [ {'name': 'Schaefer1000x7_TrimMean80', @@ -30,12 +31,16 @@ markers = [ 'method': 'std'} ] +storage = { + 'kind': 'SQLiteFeatureStorage', +} + with tempfile.TemporaryDirectory() as tmpdir: uri = f'{tmpdir}/test.db' + storage['uri'] = uri run( workdir='/tmp', - datagrabber='OasisVBMTestingDatagrabber', + datagrabber=datagrabber, markers=markers, - storage='SQLiteFeatureStorage', - storage_params={'uri': uri}, + storage=storage, ) diff --git a/examples/yamls/gmd_mean.yaml b/examples/yamls/gmd_mean.yaml new file mode 100644 index 000000000..ce50d5d5c --- /dev/null +++ b/examples/yamls/gmd_mean.yaml @@ -0,0 +1,25 @@ +with: junifer.testing.registry +workdir: /tmp + +datagrabber: + kind: OasisVBMTestingDatagrabber +elements: +markers: + - name: Schaefer1000x7_TrimMean80 + kind: ParcelAggregation + atlas: Schaefer1000x7 + method: trim_mean + method_params: + proportiontocut: 0.2 + - name: Schaefer1000x7_Mean + kind: ParcelAggregation + atlas: Schaefer1000x7 + method: mean + - name: Schaefer1000x7_Std + kind: ParcelAggregation + atlas: Schaefer1000x7 + method: std +storage: + kind: SQLiteFeatureStorage + uri: /Users/fraimondo/dev/tbox/junifer/scratch/db/test.db + diff --git a/junifer/api/__init__.py b/junifer/api/__init__.py index 4bc42cc57..9ab1c35b8 100644 --- a/junifer/api/__init__.py +++ b/junifer/api/__init__.py @@ -1 +1,2 @@ -from . run import run \ No newline at end of file +from . run import run +from . import cli \ No newline at end of file diff --git a/junifer/api/cli.py b/junifer/api/cli.py new file mode 100644 index 000000000..fcc3da991 --- /dev/null +++ b/junifer/api/cli.py @@ -0,0 +1,49 @@ +import click + +from .parser import parse_yaml +from .run import run as api_run +from .run import collect as api_collect +from ..utils.logging import configure_logging + + +@click.group() +def cli(): + pass + + +@cli.command() +@click.argument('filepath') +@click.option('-v', '--verbose', + type=click.Choice(['warning', 'info', 'debug'], + case_sensitive=False), + default='warning') +def run(filepath, verbose): + configure_logging(level=verbose.upper()) + contents = parse_yaml(filepath) + workdir = contents['workdir'] + datagrabber = contents['datagrabber'] + markers = contents['markers'] + storage = contents['storage'] + elements = contents.get('elements', None) + api_run( + workdir=workdir, datagrabber=datagrabber, markers=markers, + storage=storage, elements=elements) + + +@cli.command() +@click.argument('filepath') +@click.option('-v', '--verbose', + type=click.Choice(['warning', 'info', 'debug'], + case_sensitive=False), + default='warning') +def collect(filepath, verbose): + configure_logging(level=verbose.upper()) + contents = parse_yaml(filepath) + storage = contents['storage'] + api_collect(storage) + + + +@cli.command() +def queue(): + click.echo('queue') diff --git a/junifer/api/parser.py b/junifer/api/parser.py new file mode 100644 index 000000000..03b041362 --- /dev/null +++ b/junifer/api/parser.py @@ -0,0 +1,23 @@ +import yaml +import importlib +from pathlib import Path + +from ..utils.logging import raise_error + + +def parse_yaml(filepath): + if not isinstance(filepath, Path): + filepath = Path(filepath) + if not filepath.exists(): + raise_error(f'File does not exist: {filepath.as_posix()}') + with open(filepath, 'r') as f: + contents = yaml.safe_load(f) + + if 'with' in contents: + to_load = contents.pop('with') + if not isinstance(to_load, list): + to_load = [to_load] + for t_module in to_load: + importlib.import_module(t_module) + + return contents diff --git a/junifer/api/run.py b/junifer/api/run.py index f33737a33..f26957272 100644 --- a/junifer/api/run.py +++ b/junifer/api/run.py @@ -12,16 +12,17 @@ from ..markers.collection import MarkerCollection def run( - workdir, datagrabber, markers, storage, source_params=None, - storage_params=None, elements=None): + workdir, datagrabber, markers, storage, elements=None): """Run the pipeline on the selected element Parameters ---------- workdir : str or path-like object Directory where the pipeline will be executed - datagrabber : str - Name of the datagrabber to use + datagrabber : dict + Datagrabber to use. Must have a key 'kind' with the kind of + datagrabber to use. All other keys are passed to the datagrabber + init function. elements : str, tuple or list[str or tuple] Element(s) to process. Will be used to index the datagrabber. markers : list of dict @@ -30,24 +31,22 @@ def run( marker. The 'kind' key is used to specify the kind of marker to extract. The rest of the keys are used to pass parameters to the marker calculation. - storage: str - Name of the storage to use. - source_params : dict - Parameters to pass to the datagrabber. - storage_params: dict - Parameters to pass to the storage. + storage : dict + Storage to use. Must have a key 'kind' with the kind of + storage to use. All other keys are passed to the storage + init function. """ - if source_params is None: - source_params = {} - - if storage_params is None: - storage_params = {} + datagrabber_params = datagrabber.copy() + datagrabber_kind = datagrabber_params.pop('kind') + storage_params = storage.copy() + storage_kind = storage_params.pop('kind') if isinstance(workdir, str): workdir = Path(workdir) datagrabber = build( - 'datagrabber', datagrabber, BaseDataGrabber, init_params=source_params) + 'datagrabber', datagrabber_kind, BaseDataGrabber, + init_params=datagrabber_params) built_markers = [] for t_marker in markers: @@ -56,7 +55,8 @@ def run( built_markers.append(t_m) storage = build( - 'storage', storage, BaseFeatureStorage, init_params=storage_params) + 'storage', storage_kind, BaseFeatureStorage, + init_params=storage_params) mc = MarkerCollection(built_markers, storage=storage) @@ -67,3 +67,13 @@ def run( else: for t_element in datagrabber: mc.fit(datagrabber[t_element]) + + +def collect(storage): + storage_params = storage.copy() + storage_kind = storage_params.pop('kind') + + storage = build( + 'storage', storage_kind, BaseFeatureStorage, + init_params=storage_params) + storage.collect() diff --git a/junifer/markers/tests/test_collection.py b/junifer/markers/tests/test_collection.py index 969f356e1..07f601992 100644 --- a/junifer/markers/tests/test_collection.py +++ b/junifer/markers/tests/test_collection.py @@ -7,7 +7,7 @@ from numpy.testing import assert_array_equal from junifer.datareader.default import DefaultDataReader from junifer.markers import MarkerCollection, ParcelAggregation from junifer.markers.base import PipelineStepMixin -from junifer.testing import OasisVBMTestingDatagrabber +from junifer.testing.datagrabbers import OasisVBMTestingDatagrabber from junifer.storage import SQLiteFeatureStorage diff --git a/junifer/storage/base.py b/junifer/storage/base.py index c9a326d6e..f995bd13f 100644 --- a/junifer/storage/base.py +++ b/junifer/storage/base.py @@ -10,7 +10,7 @@ from ..utils import logger from .. import __version__ -def process_meta(meta, return_idx=False, n_rows=1, rows_col_name=None): +def process_meta(meta): """Process the metadata for storage. It removes the "element" key and adds the "_element_keys" with the keys used to index the element. @@ -38,10 +38,6 @@ def process_meta(meta, return_idx=False, n_rows=1, rows_col_name=None): if meta is None: raise ValueError('Meta must be a dict (currently is None)') t_meta = meta.copy() - idx = None - if return_idx is True: - idx = _element_to_index( - meta, n_rows=n_rows, rows_col_name=rows_col_name) element = t_meta.pop('element', None) if element is None: if '_element_keys' not in t_meta: @@ -52,14 +48,8 @@ def process_meta(meta, return_idx=False, n_rows=1, rows_col_name=None): t_meta['_element_keys'] = list(element.keys()) else: t_meta['_element_keys'] = ['element'] - logger.debug(f'Hasing meta {t_meta}') md5_hash = _meta_hash(t_meta) - logger.debug(f'Hash computed: {md5_hash}') - if return_idx is True: - out = md5_hash, t_meta, idx - else: - out = md5_hash, t_meta - return out + return md5_hash, t_meta def _meta_hash(meta): @@ -75,12 +65,15 @@ def _meta_hash(meta): md5: str The md5 hash of the meta """ + logger.debug(f'Hashing meta {meta}') meta_md5 = hashlib.md5( json.dumps(meta, sort_keys=True).encode('utf-8')).hexdigest() + logger.debug(f'Hash computed: {meta_md5}') return meta_md5 -def _element_to_index(meta, n_rows=1, rows_col_name=None): +# TODO: Test this new functionality +def element_to_index(meta, n_rows=1, rows_col_name=None): """Convert the element meta to index Parameters @@ -88,7 +81,7 @@ def _element_to_index(meta, n_rows=1, rows_col_name=None): meta: dict The metadata. Must contain the key 'element' n_rows: int - Number of rows to create + Number of rows to create. Defaults to 1. rows_col_name: str The column name to use in case n_rows > 1. If None (default) and n_rows > 1, the name will be 'index'. @@ -109,32 +102,30 @@ def _element_to_index(meta, n_rows=1, rows_col_name=None): element = meta['element'] if not isinstance(element, dict): element = dict(element=element) - if n_rows > 1: - if rows_col_name is None: - rows_col_name = 'index' - elem_idx = { - k: [v] * n_rows for k, v in element.items() - } - elem_idx[rows_col_name] = np.arange(n_rows) # type: ignore - else: - elem_idx = element + if rows_col_name is None: + rows_col_name = 'idx' + elem_idx = { + k: [v] * n_rows for k, v in element.items() + } + elem_idx[rows_col_name] = np.arange(n_rows) # type: ignore index = pd.MultiIndex.from_frame( pd.DataFrame(elem_idx, index=range(n_rows))) - return index def element_to_prefix(element): + logger.debug(f'Converting element {element} to prefix') prefix = 'element' - if isinstance(element, str): - prefix = f'{prefix}_{element}' - elif isinstance(element, tuple): - prefix = f"{prefix}_{'_'.join(element)}" + if isinstance(element, tuple): + prefix = f"{prefix}_{'_'.join([f'{x}' for x in element])}" elif isinstance(element, dict): - prefix = f"{prefix}_{'_'.join(element.values())}" + prefix = f"{prefix}_{'_'.join([f'{x}' for x in element.values()])}" + elif isinstance(element, (str, int)): + prefix = f'{prefix}_{element}' else: raise ValueError(f'Cannot convert element {element} to prefix. ' 'Must be a str, tuple or dict') + logger.debug(f'Converted prefix {prefix}') return f'{prefix}_' @@ -214,6 +205,10 @@ class BaseFeatureStorage(ABC): def store_timeseries(self, data, meta): raise NotImplementedError('store_timeseries not implemented') + @abstractmethod + def collect(self): + raise NotImplementedError('collect not implemented') + class PandasFeatureStoreage(BaseFeatureStorage): diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py index d872d9978..389141d24 100644 --- a/junifer/storage/sqlite.py +++ b/junifer/storage/sqlite.py @@ -7,8 +7,9 @@ from pandas.io.sql import pandasSQL_builder from sqlalchemy import create_engine, inspect from ..api.decorators import register_storage -from .base import PandasFeatureStoreage, process_meta, element_to_prefix -from ..utils.logging import warn +from .base import (PandasFeatureStoreage, process_meta, element_to_prefix, + element_to_index) +from ..utils.logging import warn, logger @register_storage @@ -93,7 +94,7 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): raise ValueError('At least one of feature_name or feature_md5 ' 'must be specified') elif feature_md5 is not None: - feature_md5 = f'meta_{feature_md5}' + table_name = f'meta_{feature_md5}' else: meta_df = pd.read_sql( 'meta', con=engine, index_col='meta_md5') @@ -105,14 +106,22 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): f'More than one feature with name {feature_name} found', 'This file is invalid. You can bypass this issue by ' 'specifying a feature_md5') - feature_md5 = f'meta_{t_df.index[0]}' - return pd.read_sql(feature_md5, con=engine) + table_name = f'meta_{t_df.index[0]}' + df = pd.read_sql(table_name, con=engine) + # Read the index: + query = ("SELECT ii.name FROM sqlite_schema AS m, " + "pragma_index_list(m.name) AS il, " + "pragma_index_info(il.name) AS ii " + f"WHERE tbl_name='{table_name}' " + "ORDER BY cid;") + index_names = pd.read_sql(query, con=engine).values.squeeze().tolist() + df = df.set_index(index_names) + return df def store_metadata(self, meta): t_meta = meta.copy() t_meta.update(self.get_meta()) - meta_md5, t_meta_row = process_meta( # type: ignore - t_meta, return_idx=False) + meta_md5, t_meta_row = process_meta(t_meta) engine = self.get_engine(t_meta) if meta_md5 not in inspect(engine).get_table_names(): meta_df = self._meta_row(t_meta_row, meta_md5) @@ -129,22 +138,35 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): def store_2d(self, data, meta, columns=None, rows_col_name=None): n_rows = len(data) - _, _, idx = process_meta( # type: ignore - meta, return_idx=True, n_rows=n_rows, rows_col_name=rows_col_name) + idx = element_to_index( + meta, n_rows=n_rows, rows_col_name=rows_col_name) data_df = pd.DataFrame(data, columns=columns, index=idx) self.store_df(data_df, meta) def store_df(self, df, meta): - table_name = self.store_metadata(meta) - # Check that the index generated by meta is in the dataframe, at least - # for one row - _, _, idx = process_meta( # type: ignore - meta, return_idx=True, n_rows=1) - if any(x not in df.index.names for x in idx.names): # type: ignore + # TODO: Test this function + # Check that the index generated by meta matches the one in + # the dataframe. + idx = element_to_index(meta) + # Given the meta, we might not know if there is an extra column added + # when storing a timeseries or 2d elements. We need to check if the + # extra element is only one. + extra = [x for x in df.index.names if x not in idx.names] + if len(extra) > 1: raise ValueError( - 'The index of the dataframe does not match the ' - 'index of the meta data. This happens when the element ' - 'is not part of the index.') + 'The index of the dataframe has extra items that are not ' + 'in the index generated from the meta data.') + elif len(extra) == 1: + # The df has one extra index item, this should be the new name + # of the missing element in the index + idx = element_to_index(meta, rows_col_name=extra[0]) + + if any(x not in df.index.names for x in idx.names): + raise ValueError( + 'The index of the dataframe is missing index items that are ' + 'generated from the meta data.') + + table_name = self.store_metadata(meta) # Save engine = self.get_engine(meta) self._save_upsert(df, table_name, engine) @@ -152,6 +174,31 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): def store_timeseries(self, data, meta): raise NotImplementedError('store_timeseries not implemented') + def collect(self): + if self.single_output is True: + raise ValueError('collect is not implemented for single output') + logger.info( + f'Collecting data from {self.uri.parent}/*{self.uri.name}') + + out_storage = SQLiteFeatureStorage( + uri=self.uri, single_output=True, upsert='ignore') + + for elem in self.uri.parent.glob(f'*{self.uri.name}'): + logger.debug(f'Reading from {elem.as_posix()}') + in_storage = SQLiteFeatureStorage(uri=elem, single_output=True) + in_engine = in_storage.get_engine() + # Open "meta" table + t_meta_df = pd.read_sql( + 'meta', con=in_engine, index_col='meta_md5') + out_storage._save_upsert(t_meta_df, 'meta') + for meta_md5 in t_meta_df.index: + logger.debug(f'Collecting feature {meta_md5}') + # TODO: Fix this, needs that read_feature sets the index + # properly + table_name = f'meta_{meta_md5}' + t_df = in_storage.read_df(feature_md5=meta_md5) + out_storage._save_upsert(t_df, table_name) + def _save_upsert(self, df, name, engine=None, if_exist='append'): """ Implementation of UPSERT functionality. @@ -217,7 +264,7 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): def _get_existing_pk(con, table_name, index_col): - pk_cols = ','.join(index_col) + pk_cols = ', '.join(index_col) query = f'SELECT {pk_cols} FROM {table_name};' pk_indb = pd.read_sql(query, con=con) return pk_indb diff --git a/junifer/storage/tests/test_base.py b/junifer/storage/tests/test_base.py index 01f2156f8..532907250 100644 --- a/junifer/storage/tests/test_base.py +++ b/junifer/storage/tests/test_base.py @@ -1,7 +1,8 @@ import pytest -from junifer.storage.base import (process_meta, _element_to_index, - BaseFeatureStorage) +from junifer.storage.base import (process_meta, element_to_index, + BaseFeatureStorage, + element_to_prefix) def test_process_meta_hash(): @@ -12,14 +13,14 @@ def test_process_meta_hash(): process_meta(meta) meta = {'element': 'foo', 'A': 1, 'B': [2, 3, 4, 5, 6]} - hash1, _ = process_meta(meta, return_idx=False) # type: ignore + hash1, _ = process_meta(meta) meta = {'element': 'foo', 'B': [2, 3, 4, 5, 6], 'A': 1} - hash2, _ = process_meta(meta, return_idx=False) # type: ignore + hash2, _ = process_meta(meta) assert hash1 == hash2 meta = {'element': 'foo', 'A': 1, 'B': [2, 3, 1, 5, 6]} - hash3, _ = process_meta(meta, return_idx=False) # type: ignore + hash3, _ = process_meta(meta) assert hash1 != hash3 meta1 = { @@ -41,8 +42,8 @@ def test_process_meta_hash(): 'element': 'foo' } - hash4, _ = process_meta(meta1, return_idx=False) # type: ignore - hash5, _ = process_meta(meta2, return_idx=False) # type: ignore + hash4, _ = process_meta(meta1) + hash5, _ = process_meta(meta2) assert hash4 == hash5 @@ -51,10 +52,10 @@ def test_process_meta_element(): meta = {} with pytest.raises(ValueError, match=r"_element_keys"): - process_meta(meta, return_idx=False) + process_meta(meta) meta = {'element': 'foo', 'A': 1, 'B': [2, 3, 4, 5, 6]} - _, new_meta = process_meta(meta, return_idx=False) # type: ignore + _, new_meta = process_meta(meta) assert '_element_keys' in new_meta assert new_meta['_element_keys'] == ['element'] assert 'A' in new_meta @@ -63,7 +64,7 @@ def test_process_meta_element(): meta = { 'element': {'subject': 'foo', 'session': 'bar'}, 'B': [2, 3, 4, 5, 6], 'A': 1} - _, new_meta = process_meta(meta, return_idx=False) # type: ignore + _, new_meta = process_meta(meta) assert '_element_keys' in new_meta assert new_meta['_element_keys'] == ['subject', 'session'] assert 'A' in new_meta @@ -71,35 +72,44 @@ def test_process_meta_element(): def test_process_meta_index(): - """Test _element_to_index""" + """Test element_to_index""" meta = {'noelement': 'foo'} with pytest.raises(ValueError, match=r'meta must contain the key'): - _element_to_index(meta) + element_to_index(meta) meta = {'element': 'foo', 'A': 1, 'B': [2, 3, 4, 5, 6]} - index = _element_to_index(meta) - assert index.names == ['element'] + index = element_to_index(meta) + assert index.names == ['element', 'idx'] assert index.levels[0].name == 'element' assert index.levels[0].values[0] == 'foo' + assert all(x == 'foo' for x in index.levels[0].values) + assert index.levels[0].values.shape == (1,) + assert index.levels[1].name == 'idx' + assert all(x == i for i, x in enumerate(index.levels[1].values)) + assert index.levels[1].values.shape == (1,) - index = _element_to_index(meta, n_rows=10) - assert index.names == ['element', 'index'] + index = element_to_index(meta, n_rows=10) + assert index.names == ['element', 'idx'] assert index.levels[0].name == 'element' assert all(x == 'foo' for x in index.levels[0].values) assert index.levels[0].values.shape == (1,) - assert index.levels[1].name == 'index' + assert index.levels[1].name == 'idx' assert all(x == i for i, x in enumerate(index.levels[1].values)) assert index.levels[1].values.shape == (10,) - index = _element_to_index(meta, n_rows=1, rows_col_name='scan') - assert index.names == ['element'] + index = element_to_index(meta, n_rows=1, rows_col_name='scan') + assert index.names == ['element', 'scan'] assert index.levels[0].name == 'element' + assert index.levels[0].values[0] == 'foo' assert all(x == 'foo' for x in index.levels[0].values) assert index.levels[0].values.shape == (1,) + assert index.levels[1].name == 'scan' + assert all(x == i for i, x in enumerate(index.levels[1].values)) + assert index.levels[1].values.shape == (1,) - index = _element_to_index(meta, n_rows=7, rows_col_name='scan') + index = element_to_index(meta, n_rows=7, rows_col_name='scan') assert index.names == ['element', 'scan'] assert index.levels[0].name == 'element' assert all(x == 'foo' for x in index.levels[0].values) @@ -112,7 +122,7 @@ def test_process_meta_index(): meta = { 'element': {'subject': 'sub-01', 'session': 'ses-01'}, 'A': 1, 'B': [2, 3, 4, 5, 6]} - index = _element_to_index(meta, n_rows=10) + index = element_to_index(meta, n_rows=10) assert index.levels[0].name == 'subject' assert all(x == 'sub-01' for x in index.levels[0].values) @@ -122,7 +132,7 @@ def test_process_meta_index(): assert all(x == 'ses-01' for x in index.levels[1].values) assert index.levels[1].values.shape == (1,) - assert index.levels[2].name == 'index' + assert index.levels[2].name == 'idx' assert all(x == i for i, x in enumerate(index.levels[2].values)) assert index.levels[2].values.shape == (10,) @@ -161,6 +171,9 @@ def test_BaseFeatureStorage(): def store_timeseries(self, timeseries, meta): super().store_timeseries(timeseries, meta) + def collect(self): + return super().collect() + st = MyFeatureStorage(uri='/tmp') assert st.single_output is False @@ -191,4 +204,43 @@ def test_BaseFeatureStorage(): with pytest.raises(NotImplementedError): st.store_timeseries(None, None) + with pytest.raises(NotImplementedError): + st.collect() + assert st.uri == '/tmp' + + +def test_element_to_prefix(): + """Test converting element to prefix (for file naming)""" + + element = 'sub-01' + prefix = element_to_prefix(element) + assert prefix == 'element_sub-01_' + + element = 1 + prefix = element_to_prefix(element) + assert prefix == 'element_1_' + + element = {'subject': 'sub-01'} + prefix = element_to_prefix(element) + assert prefix == 'element_sub-01_' + + element = {'subject': 1} + prefix = element_to_prefix(element) + assert prefix == 'element_1_' + + element = {'subject': 'sub-01', 'session': 'ses-02'} + prefix = element_to_prefix(element) + assert prefix == 'element_sub-01_ses-02_' + + element = {'subject': 1, 'session': 2} + prefix = element_to_prefix(element) + assert prefix == 'element_1_2_' + + element = ('sub-01', 'ses-02') + prefix = element_to_prefix(element) + assert prefix == 'element_sub-01_ses-02_' + + element = (1, 2) + prefix = element_to_prefix(element) + assert prefix == 'element_1_2_' diff --git a/junifer/storage/tests/test_sqlite.py b/junifer/storage/tests/test_sqlite.py index 31abd542d..6d82c5993 100644 --- a/junifer/storage/tests/test_sqlite.py +++ b/junifer/storage/tests/test_sqlite.py @@ -7,7 +7,8 @@ from sqlalchemy import create_engine import pytest from junifer.storage.sqlite import SQLiteFeatureStorage -from junifer.storage.base import process_meta, element_to_prefix +from junifer.storage.base import (process_meta, element_to_prefix, + element_to_index) df1 = pd.DataFrame({ @@ -160,11 +161,10 @@ def test_store_read_df(): to_store = df1[['col1', 'col2']] # Save to SQL - with pytest.raises(ValueError, match=r"index of the dataframe"): + with pytest.raises(ValueError, match=r"missing index items"): storage.store_df(to_store.set_index('col1'), meta) - _, _, idx = process_meta( # type: ignore - meta, return_idx=True, n_rows=len(to_store)) + idx = element_to_index(meta, n_rows=len(to_store)) to_store = to_store.set_index(idx) storage.store_df(to_store, meta) @@ -190,7 +190,7 @@ def test_store_read_df(): read_df1 = storage.read_df(feature_md5=feature_md5) read_df2 = storage.read_df(feature_name='fcname') assert_frame_equal(read_df1, read_df2) - assert_frame_equal(read_df1, to_store.reset_index()) + assert_frame_equal(read_df1, to_store) def test_store_table(): @@ -206,8 +206,8 @@ def test_store_table(): [4, 40], [5, 50], ] - _, _, idx = process_meta( # type: ignore - meta, return_idx=True, n_rows=5, rows_col_name='scan') + + idx = element_to_index(meta, n_rows=5, rows_col_name='scan') df1 = pd.DataFrame(data, columns=['f1', 'f2'], index=idx) storage.store_table( @@ -226,8 +226,7 @@ def test_store_table(): [6, 600] ] - _, _, idx = process_meta( # type: ignore - meta, return_idx=True, n_rows=6, rows_col_name='scan') + idx = element_to_index(meta, n_rows=6, rows_col_name='scan') df2 = pd.DataFrame(data2, columns=['f1', 'f2'], index=idx) with pytest.warns(RuntimeWarning, match=r"Some rows"): @@ -261,16 +260,16 @@ def test_store_multiple_output(): data2 = data1 * 10 data3 = data1 * 20 - hash1, _, idx1 = process_meta( # type: ignore - meta1, return_idx=True, n_rows=5, rows_col_name='scan') + hash1, _ = process_meta(meta1) + idx1 = element_to_index(meta1, n_rows=5, rows_col_name='scan') df1 = pd.DataFrame(data1, columns=['f1', 'f2'], index=idx1) - hash2, _, idx2 = process_meta( # type: ignore - meta2, return_idx=True, n_rows=5, rows_col_name='scan') + hash2, _ = process_meta(meta2) + idx2 = element_to_index(meta2, n_rows=5, rows_col_name='scan') df2 = pd.DataFrame(data2, columns=['f1', 'f2'], index=idx2) - hash3, _, idx3 = process_meta( # type: ignore - meta3, return_idx=True, n_rows=5, rows_col_name='scan') + hash3, _ = process_meta(meta3) + idx3 = element_to_index(meta3, n_rows=5, rows_col_name='scan') df3 = pd.DataFrame(data3, columns=['f1', 'f2'], index=idx3) assert hash1 == hash2 diff --git a/junifer/testing/__init__.py b/junifer/testing/__init__.py new file mode 100644 index 000000000..4b2ba56f9 --- /dev/null +++ b/junifer/testing/__init__.py @@ -0,0 +1,3 @@ +# Authors: Federico Raimondo +# License: AGPL +from . import datagrabbers \ No newline at end of file diff --git a/junifer/testing/datagrabbers.py b/junifer/testing/datagrabbers.py new file mode 100644 index 000000000..fba285064 --- /dev/null +++ b/junifer/testing/datagrabbers.py @@ -0,0 +1,29 @@ +# Authors: Federico Raimondo +# License: AGPL +import tempfile +from nilearn import datasets + +from ..datagrabber.base import BaseDataGrabber + + +class OasisVBMTestingDatagrabber(BaseDataGrabber): + """ + DataGrabber for Oasis VBM testing data. + """ + def __init__(self): + datadir = tempfile.mkdtemp() + types = ['VBM_GM'] + super().__init__(types=types, datadir=datadir) + + def get_elements(self): + return list(range(1, 11)) + + def __getitem__(self, element): + out = {} + out['VBM_GM'] = self._dataset.gray_matter_maps[element - 1] + out['meta'] = {'element': {'subject': element}} + return out + + def __enter__(self): + self._dataset = datasets.fetch_oasis_vbm(n_subjects=10) + return self diff --git a/junifer/testing/registry.py b/junifer/testing/registry.py new file mode 100644 index 000000000..556005b53 --- /dev/null +++ b/junifer/testing/registry.py @@ -0,0 +1,6 @@ +from .datagrabbers import OasisVBMTestingDatagrabber +from ..api.registry import register + +register( + 'datagrabber', 'OasisVBMTestingDatagrabber', + OasisVBMTestingDatagrabber) diff --git a/requirements.txt b/requirements.txt index eca50d3f0..2eb1a7062 100644 --- a/requirements.txt +++ b/requirements.txt @@ -3,4 +3,5 @@ datalad>=0.15.4, <0.16 pandas>=0.18.0, <1.5 nibabel>=3.2.0, <4.0 nilearn>=0.9.0, <1.0 -sqlalchemy>=1.4.27, <= 1.5.0 \ No newline at end of file +sqlalchemy>=1.4.27, <= 1.5.0 +pyyaml>=5.1.2, <7.0 \ No newline at end of file diff --git a/setup.py b/setup.py index 846a52937..1c3f673e2 100644 --- a/setup.py +++ b/setup.py @@ -52,7 +52,13 @@ setuptools.setup( 'Source': DOWNLOAD_URL, 'Tracker': f'{DOWNLOAD_URL}issues/', }, - install_requires=[], # TODO: Complete + install_requires=['Click'], # TODO: Complete + py_modules=['junifer.api.cli'], + entry_points={ + 'console_scripts': [ + 'junifer=junifer.api.cli:cli', + ] + }, python_requires='>=3.6', use_scm_version=_getversion, setup_requires=['setuptools_scm'], -- 2.52.0 From 4e1d118dd61613b2ffd22245f44cd172a493b344 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Wed, 23 Mar 2022 11:48:57 +0100 Subject: [PATCH 034/287] Flake --- junifer/api/cli.py | 1 - 1 file changed, 1 deletion(-) diff --git a/junifer/api/cli.py b/junifer/api/cli.py index fcc3da991..6a14bfd2a 100644 --- a/junifer/api/cli.py +++ b/junifer/api/cli.py @@ -43,7 +43,6 @@ def collect(filepath, verbose): api_collect(storage) - @cli.command() def queue(): click.echo('queue') -- 2.52.0 From dd0077f7515e4313d4a03fe8d98253fcf51925b7 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Wed, 23 Mar 2022 12:02:34 +0100 Subject: [PATCH 035/287] Update reqs + use sqlite_master --- dev-requirements.txt | 3 ++- junifer/storage/sqlite.py | 2 +- junifer/storage/tests/test_sqlite.py | 3 ++- 3 files changed, 5 insertions(+), 3 deletions(-) diff --git a/dev-requirements.txt b/dev-requirements.txt index aad120b7c..2e4e347ff 100644 --- a/dev-requirements.txt +++ b/dev-requirements.txt @@ -1,2 +1,3 @@ flake8 -pytest \ No newline at end of file +pytest +click \ No newline at end of file diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py index 389141d24..9616ff7f6 100644 --- a/junifer/storage/sqlite.py +++ b/junifer/storage/sqlite.py @@ -109,7 +109,7 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): table_name = f'meta_{t_df.index[0]}' df = pd.read_sql(table_name, con=engine) # Read the index: - query = ("SELECT ii.name FROM sqlite_schema AS m, " + query = ("SELECT ii.name FROM sqlite_master AS m, " "pragma_index_list(m.name) AS il, " "pragma_index_info(il.name) AS ii " f"WHERE tbl_name='{table_name}' " diff --git a/junifer/storage/tests/test_sqlite.py b/junifer/storage/tests/test_sqlite.py index 6d82c5993..de185954f 100644 --- a/junifer/storage/tests/test_sqlite.py +++ b/junifer/storage/tests/test_sqlite.py @@ -115,7 +115,8 @@ def test_upsert_ignore(): c_df1 = _read_sql(table_name, uri=uri, index_col=['element', 'pk2']) assert_frame_equal(df1, c_df1) - storage.store_df(df2, meta) + with pytest.warns(RuntimeWarning, match='are already present'): + storage.store_df(df2, meta) c_dfignore = _read_sql( table_name, uri=uri, index_col=['element', 'pk2']) -- 2.52.0 From 5e1dfdbff45d4e5c27291d3347b22a42006948e9 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Wed, 23 Mar 2022 12:19:15 +0100 Subject: [PATCH 036/287] Test collect --- .coveragerc | 2 +- junifer/storage/tests/test_sqlite.py | 72 ++++++++++++++++++++++++++++ 2 files changed, 73 insertions(+), 1 deletion(-) diff --git a/.coveragerc b/.coveragerc index 6cdbc7a3c..2cd7255d0 100644 --- a/.coveragerc +++ b/.coveragerc @@ -6,7 +6,7 @@ omit = */setup.py */tests/* junifer/configs/juseless.py - junifer/testing.py + junifer/testing/* [report] exclude_lines = diff --git a/junifer/storage/tests/test_sqlite.py b/junifer/storage/tests/test_sqlite.py index de185954f..cde4e35ea 100644 --- a/junifer/storage/tests/test_sqlite.py +++ b/junifer/storage/tests/test_sqlite.py @@ -58,6 +58,11 @@ def test_get_engine(): assert engine.url.drivername == 'sqlite' assert f'{engine.url.database}' == uri + storage = SQLiteFeatureStorage( + uri=uri, single_output=False, upsert='ignore') + with pytest.raises(ValueError, match='element must be specified'): + storage.get_engine() + def test_store_metadata(): """Test store_metadata""" @@ -165,6 +170,10 @@ def test_store_read_df(): with pytest.raises(ValueError, match=r"missing index items"): storage.store_df(to_store.set_index('col1'), meta) + to_store = df1.reset_index.set_index(['element', 'pk2', 'col1']) + with pytest.raises(ValueError, match=r"extra items"): + storage.store_df(to_store.set_index('col1'), meta) + idx = element_to_index(meta, n_rows=len(to_store)) to_store = to_store.set_index(idx) @@ -310,3 +319,66 @@ def test_store_multiple_output(): assert_frame_equal(df1, cdf1) assert_frame_equal(df2, cdf2) assert_frame_equal(df3, cdf3) + + +def test_collect(): + meta1 = {'element': {'subject': 'test-01', 'session': 'ses-01'}, + 'version': '0.0.1', 'marker': {'name': 'fc'}} + meta2 = {'element': {'subject': 'test-02', 'session': 'ses-01'}, + 'version': '0.0.1', 'marker': {'name': 'fc'}} + meta3 = {'element': {'subject': 'test-01', 'session': 'ses-02'}, + 'version': '0.0.1', 'marker': {'name': 'fc'}} + with tempfile.TemporaryDirectory() as _tmpdir: + uri = Path(f'{_tmpdir}/test.db') + storage = SQLiteFeatureStorage(uri=uri) + + data1 = np.array([ + [1, 10], + [2, 20], + [3, 30], + [4, 40], + [5, 50], + ]) + + data2 = data1 * 10 + data3 = data1 * 20 + storage.store_table( + data1, meta1, columns=['f1', 'f2'], rows_col_name='scan') + + storage.store_table( + data2, meta2, columns=['f1', 'f2'], rows_col_name='scan') + + storage.store_table( + data3, meta3, columns=['f1', 'f2'], rows_col_name='scan') + + prefix1 = element_to_prefix(meta1['element']) + prefix2 = element_to_prefix(meta2['element']) + prefix3 = element_to_prefix(meta3['element']) + + uri1 = uri.parent / f'{prefix1}{uri.name}' + uri2 = uri.parent / f'{prefix2}{uri.name}' + uri3 = uri.parent / f'{prefix3}{uri.name}' + + assert uri1.exists() + assert uri2.exists() + assert uri3.exists() + + assert not uri.exists() + + storage.collect() + + assert uri.exists() + + cols = ['subject', 'session', 'scan'] + table_name = storage.store_metadata(meta1) + all_df = _read_sql(table_name, uri, index_col=cols) + + cdf1 = _read_sql(table_name, uri1, index_col=cols) + cdf2 = _read_sql(table_name, uri2, index_col=cols) + cdf3 = _read_sql(table_name, uri3, index_col=cols) + + all_cdf = pd.concat([cdf1, cdf2, cdf3]) + all_df.sort_index(level=cols, inplace=True) + all_cdf.sort_index(level=cols, inplace=True) + + assert_frame_equal(all_df, all_cdf) -- 2.52.0 From a39cff7d1cd9d1505926adae0874734bf8863aa6 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Wed, 23 Mar 2022 12:24:04 +0100 Subject: [PATCH 037/287] typo --- junifer/storage/tests/test_sqlite.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/junifer/storage/tests/test_sqlite.py b/junifer/storage/tests/test_sqlite.py index cde4e35ea..558cff610 100644 --- a/junifer/storage/tests/test_sqlite.py +++ b/junifer/storage/tests/test_sqlite.py @@ -170,7 +170,7 @@ def test_store_read_df(): with pytest.raises(ValueError, match=r"missing index items"): storage.store_df(to_store.set_index('col1'), meta) - to_store = df1.reset_index.set_index(['element', 'pk2', 'col1']) + to_store = df1.reset_index().set_index(['element', 'pk2', 'col1']) with pytest.raises(ValueError, match=r"extra items"): storage.store_df(to_store.set_index('col1'), meta) -- 2.52.0 From 56382b543b3a354b5364d2da3a62c4d4d35e7e5b Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Wed, 23 Mar 2022 12:44:29 +0100 Subject: [PATCH 038/287] getting tired of testing sqlite --- junifer/storage/tests/test_sqlite.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/junifer/storage/tests/test_sqlite.py b/junifer/storage/tests/test_sqlite.py index 558cff610..55e3670e7 100644 --- a/junifer/storage/tests/test_sqlite.py +++ b/junifer/storage/tests/test_sqlite.py @@ -172,7 +172,7 @@ def test_store_read_df(): to_store = df1.reset_index().set_index(['element', 'pk2', 'col1']) with pytest.raises(ValueError, match=r"extra items"): - storage.store_df(to_store.set_index('col1'), meta) + storage.store_df(to_store, meta) idx = element_to_index(meta, n_rows=len(to_store)) to_store = to_store.set_index(idx) -- 2.52.0 From bdf4c764a131533c894614a5cb94bb2d0868ab71 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Wed, 23 Mar 2022 14:18:04 +0100 Subject: [PATCH 039/287] More tests --- junifer/api/__init__.py | 4 +- junifer/api/cli.py | 4 +- junifer/api/{run.py => functions.py} | 0 junifer/api/registry.py | 6 +-- junifer/api/tests/test_functions.py | 0 junifer/api/tests/test_parser.py | 39 +++++++++++++++++++ junifer/api/tests/test_registry.py | 56 ++++++++++++++++++++++++++++ 7 files changed, 101 insertions(+), 8 deletions(-) rename junifer/api/{run.py => functions.py} (100%) create mode 100644 junifer/api/tests/test_functions.py create mode 100644 junifer/api/tests/test_parser.py create mode 100644 junifer/api/tests/test_registry.py diff --git a/junifer/api/__init__.py b/junifer/api/__init__.py index 9ab1c35b8..ecd3d2c69 100644 --- a/junifer/api/__init__.py +++ b/junifer/api/__init__.py @@ -1,2 +1,2 @@ -from . run import run -from . import cli \ No newline at end of file +from .functions import run, collect +from .import cli \ No newline at end of file diff --git a/junifer/api/cli.py b/junifer/api/cli.py index 6a14bfd2a..c32b21484 100644 --- a/junifer/api/cli.py +++ b/junifer/api/cli.py @@ -1,8 +1,8 @@ import click from .parser import parse_yaml -from .run import run as api_run -from .run import collect as api_collect +from .functions import run as api_run +from .functions import collect as api_collect from ..utils.logging import configure_logging diff --git a/junifer/api/run.py b/junifer/api/functions.py similarity index 100% rename from junifer/api/run.py rename to junifer/api/functions.py diff --git a/junifer/api/registry.py b/junifer/api/registry.py index 3de096073..922433f7c 100644 --- a/junifer/api/registry.py +++ b/junifer/api/registry.py @@ -93,14 +93,12 @@ def build(step, name, baseclass, init_params=None): If the name is not a string or the object is not an instance of the baseclass parameter. """ - if not isinstance(name, str): - raise_error(f'Invalid name: {name}', ValueError) klass = get(step, name) if init_params is None: init_params = {} object = klass(**init_params) if not isinstance(object, baseclass): raise_error( - f'Invalid {step} ({object.__class__.name}). ' - f'Must inherit from {baseclass.name}', ValueError) + f'Invalid {step} ({object.__class__.__name__}). ' + f'Must inherit from {baseclass.__name__}', ValueError) return object diff --git a/junifer/api/tests/test_functions.py b/junifer/api/tests/test_functions.py new file mode 100644 index 000000000..e69de29bb diff --git a/junifer/api/tests/test_parser.py b/junifer/api/tests/test_parser.py new file mode 100644 index 000000000..32447506f --- /dev/null +++ b/junifer/api/tests/test_parser.py @@ -0,0 +1,39 @@ +import sys +from pathlib import Path +import tempfile +import pytest + +from junifer.api.parser import parse_yaml + + +def test_parse_yaml(): + """Test parse yaml""" + with pytest.raises(ValueError, match='does not exist'): + parse_yaml('foo.yaml') + + with tempfile.TemporaryDirectory() as _tmpdir: + fname = Path(_tmpdir) / 'test.yaml' + with open(fname, 'w') as f: + f.write('foo: bar\n') + f.write('with: junifer.configs.juseless\n') + + assert 'junifer.configs.juseless' not in sys.modules + contents = parse_yaml(fname) + assert 'foo' in contents + assert 'bar' == contents['foo'] + assert 'with' not in contents + assert 'junifer.configs.juseless' in sys.modules + + assert 'junifer.testing.registry' not in sys.modules + + with open(fname, 'w') as f: + f.write('foo: bar\n') + f.write('with:\n') + f.write(' - junifer.configs.juseless\n') + f.write(' - junifer.testing.registry\n') + + contents = parse_yaml(fname.as_posix()) + assert 'foo' in contents + assert 'bar' == contents['foo'] + assert 'with' not in contents + assert 'junifer.testing.registry' in sys.modules diff --git a/junifer/api/tests/test_registry.py b/junifer/api/tests/test_registry.py new file mode 100644 index 000000000..d8716f541 --- /dev/null +++ b/junifer/api/tests/test_registry.py @@ -0,0 +1,56 @@ +import pytest +from abc import ABC + +from junifer.api.registry import register, get_step_names, get, build + + +def test_register_error(): + """Test register error""" + with pytest.raises(ValueError, match='Invalid ste'): + register('foo', 'bar', 'baz') + + +def test_gets(): + """Test get""" + with pytest.raises(ValueError, match='Invalid ste'): + get_step_names('foo') + + datagrabbers = get_step_names('datagrabber') + assert 'bar' not in datagrabbers + register('datagrabber', 'bar', 'baz') + datagrabbers = get_step_names('datagrabber') + assert 'bar' in datagrabbers + + with pytest.raises(ValueError, match='Invalid ste'): + get('foo', 'bar') + + with pytest.raises(ValueError, match='Invalid name'): + get('datagrabber', 'foo') + + obj = get('datagrabber', 'bar') + assert obj == 'baz' + + +def test_build(): + """Test building objects from names""" + import numpy as np + + class SuperClass(ABC): + pass + + class ConcreteClass(SuperClass): + def __init__(self, value=1): + self.value = value + + register('datagrabber', 'concrete', ConcreteClass) + + obj = build('datagrabber', 'concrete', SuperClass) + assert isinstance(obj, ConcreteClass) + assert obj.value == 1 + + obj = build('datagrabber', 'concrete', SuperClass, {'value': 2}) + assert isinstance(obj, ConcreteClass) + assert obj.value == 2 + + with pytest.raises(ValueError, match='Must inherit'): + build('datagrabber', 'concrete', np.ndarray) -- 2.52.0 From 5928bc910bcf14c178121373640772e962af0721 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Wed, 23 Mar 2022 14:58:09 +0100 Subject: [PATCH 040/287] Test run and collect --- junifer/api/functions.py | 5 +- junifer/api/tests/test_functions.py | 91 +++++++++++++++++++++++++++++ 2 files changed, 94 insertions(+), 2 deletions(-) diff --git a/junifer/api/functions.py b/junifer/api/functions.py index f26957272..b98273051 100644 --- a/junifer/api/functions.py +++ b/junifer/api/functions.py @@ -47,9 +47,10 @@ def run( datagrabber = build( 'datagrabber', datagrabber_kind, BaseDataGrabber, init_params=datagrabber_params) - + # Copy to avoid changing the original dict + _markers = [x.copy() for x in markers] built_markers = [] - for t_marker in markers: + for t_marker in _markers: kind = t_marker.pop('kind') t_m = build('marker', kind, BaseMarker, init_params=t_marker) built_markers.append(t_m) diff --git a/junifer/api/tests/test_functions.py b/junifer/api/tests/test_functions.py index e69de29bb..5fc65510d 100644 --- a/junifer/api/tests/test_functions.py +++ b/junifer/api/tests/test_functions.py @@ -0,0 +1,91 @@ +import tempfile +from pathlib import Path + +from junifer.api.registry import build +from junifer.api.functions import run, collect +from junifer.datagrabber.base import BaseDataGrabber +import junifer.testing.registry # noqa: F401 + +datagrabber = { + 'kind': 'OasisVBMTestingDatagrabber', +} + +markers = [ + {'name': 'Schaefer1000x7_Mean', + 'kind': 'ParcelAggregation', + 'atlas': 'Schaefer1000x7', + 'method': 'mean'}, + {'name': 'Schaefer1000x7_Std', + 'kind': 'ParcelAggregation', + 'atlas': 'Schaefer1000x7', + 'method': 'std'} +] + +storage = { + 'kind': 'SQLiteFeatureStorage', +} + + +def test_run(): + """Test run function""" + with tempfile.TemporaryDirectory() as tmpdir: + tmp_path = Path(tmpdir) + workdir = tmp_path / 'workdir' + workdir.mkdir() + outdir = tmp_path / 'out' + outdir.mkdir() + uri = outdir / 'test.db' + storage['uri'] = uri # type: ignore + run( + workdir=workdir, + datagrabber=datagrabber, + markers=markers, + storage=storage, + elements=[1] + ) + + files = list(outdir.glob('*.db')) + assert len(files) == 1 + + run( + workdir=workdir.as_posix(), + datagrabber=datagrabber, + markers=markers, + storage=storage, + elements=[1, 3] + ) + + files = list(outdir.glob('*.db')) + assert len(files) == 2 + + +def test_collect(): + """Test run and collect functions""" + + with tempfile.TemporaryDirectory() as tmpdir: + tmp_path = Path(tmpdir) + workdir = tmp_path / 'workdir' + workdir.mkdir() + outdir = tmp_path / 'out' + outdir.mkdir() + uri = outdir / 'test.db' + storage['uri'] = uri # type: ignore + run( + workdir=workdir, + datagrabber=datagrabber, + markers=markers, + storage=storage, + ) + dg = build('datagrabber', datagrabber['kind'], BaseDataGrabber) + elements = dg.get_elements() + + # This should create 10 files + files = list(outdir.glob('*.db')) + assert len(files) == len(elements) + + # But the test.db file should not exist + assert not uri.exists() + collect(storage) + + # Now the file exists + assert uri.exists() -- 2.52.0 From db43d379943ad4c2750bead8798ec864044f3b1b Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Thu, 24 Mar 2022 12:56:37 +0100 Subject: [PATCH 041/287] Fix parsing test --- junifer/api/run.py | 79 ++++++++++++++++++++++++++++ junifer/api/tests/data/gmd_mean.yaml | 15 ++++++ junifer/api/tests/test_parser.py | 17 +++--- junifer/storage/base.py | 1 - 4 files changed, 100 insertions(+), 12 deletions(-) create mode 100644 junifer/api/run.py create mode 100644 junifer/api/tests/data/gmd_mean.yaml diff --git a/junifer/api/run.py b/junifer/api/run.py new file mode 100644 index 000000000..f26957272 --- /dev/null +++ b/junifer/api/run.py @@ -0,0 +1,79 @@ +# Authors: Federico Raimondo +# Leonard Sasse +# License: AGPL + +from pathlib import Path + +from .registry import build +from ..datagrabber.base import BaseDataGrabber +from ..markers.base import BaseMarker +from ..storage.base import BaseFeatureStorage +from ..markers.collection import MarkerCollection + + +def run( + workdir, datagrabber, markers, storage, elements=None): + """Run the pipeline on the selected element + + Parameters + ---------- + workdir : str or path-like object + Directory where the pipeline will be executed + datagrabber : dict + Datagrabber to use. Must have a key 'kind' with the kind of + datagrabber to use. All other keys are passed to the datagrabber + init function. + elements : str, tuple or list[str or tuple] + Element(s) to process. Will be used to index the datagrabber. + markers : list of dict + List of markers to extract. Each marker is a dict with at least two + keys: 'name' and 'kind'. The 'name' key is used to name the output + marker. The 'kind' key is used to specify the kind of marker to + extract. The rest of the keys are used to pass parameters to the + marker calculation. + storage : dict + Storage to use. Must have a key 'kind' with the kind of + storage to use. All other keys are passed to the storage + init function. + """ + datagrabber_params = datagrabber.copy() + datagrabber_kind = datagrabber_params.pop('kind') + storage_params = storage.copy() + storage_kind = storage_params.pop('kind') + + if isinstance(workdir, str): + workdir = Path(workdir) + + datagrabber = build( + 'datagrabber', datagrabber_kind, BaseDataGrabber, + init_params=datagrabber_params) + + built_markers = [] + for t_marker in markers: + kind = t_marker.pop('kind') + t_m = build('marker', kind, BaseMarker, init_params=t_marker) + built_markers.append(t_m) + + storage = build( + 'storage', storage_kind, BaseFeatureStorage, + init_params=storage_params) + + mc = MarkerCollection(built_markers, storage=storage) + + with datagrabber: + if elements is not None: + for t_element in elements: + mc.fit(datagrabber[t_element]) + else: + for t_element in datagrabber: + mc.fit(datagrabber[t_element]) + + +def collect(storage): + storage_params = storage.copy() + storage_kind = storage_params.pop('kind') + + storage = build( + 'storage', storage_kind, BaseFeatureStorage, + init_params=storage_params) + storage.collect() diff --git a/junifer/api/tests/data/gmd_mean.yaml b/junifer/api/tests/data/gmd_mean.yaml new file mode 100644 index 000000000..af4df1798 --- /dev/null +++ b/junifer/api/tests/data/gmd_mean.yaml @@ -0,0 +1,15 @@ +with: junifer.testing.registry +workdir: /tmp + +datagrabber: + kind: OasisVBMTestingDatagrabber +elements: +markers: + - name: Schaefer1000x7_Mean + kind: ParcelAggregation + atlas: Schaefer1000x7 + method: mean +storage: + kind: SQLiteFeatureStorage + uri: /Users/fraimondo/dev/tbox/junifer/scratch/db/test.db + diff --git a/junifer/api/tests/test_parser.py b/junifer/api/tests/test_parser.py index 32447506f..a91604d03 100644 --- a/junifer/api/tests/test_parser.py +++ b/junifer/api/tests/test_parser.py @@ -15,25 +15,20 @@ def test_parse_yaml(): fname = Path(_tmpdir) / 'test.yaml' with open(fname, 'w') as f: f.write('foo: bar\n') - f.write('with: junifer.configs.juseless\n') + f.write('with: numpy\n') - assert 'junifer.configs.juseless' not in sys.modules contents = parse_yaml(fname) assert 'foo' in contents assert 'bar' == contents['foo'] assert 'with' not in contents - assert 'junifer.configs.juseless' in sys.modules - assert 'junifer.testing.registry' not in sys.modules + assert 'junifer.configs.wrong_config' not in sys.modules with open(fname, 'w') as f: f.write('foo: bar\n') f.write('with:\n') - f.write(' - junifer.configs.juseless\n') - f.write(' - junifer.testing.registry\n') + f.write(' - numpy\n') + f.write(' - junifer.testing.wrong_config\n') - contents = parse_yaml(fname.as_posix()) - assert 'foo' in contents - assert 'bar' == contents['foo'] - assert 'with' not in contents - assert 'junifer.testing.registry' in sys.modules + with pytest.raises(ImportError, match='wrong_config'): + contents = parse_yaml(fname.as_posix()) diff --git a/junifer/storage/base.py b/junifer/storage/base.py index f995bd13f..da1f26aa5 100644 --- a/junifer/storage/base.py +++ b/junifer/storage/base.py @@ -72,7 +72,6 @@ def _meta_hash(meta): return meta_md5 -# TODO: Test this new functionality def element_to_index(meta, n_rows=1, rows_col_name=None): """Convert the element meta to index -- 2.52.0 From 4b6c0d93f23fb6297d0531b2541c2d51cd8ec7a5 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Thu, 24 Mar 2022 14:44:23 +0100 Subject: [PATCH 042/287] Test cli --- Makefile | 2 +- junifer/api/tests/test_cli.py | 38 ++++++++++++++++++++++++++++++++ junifer/api/tests/test_parser.py | 7 ++++++ 3 files changed, 46 insertions(+), 1 deletion(-) create mode 100644 junifer/api/tests/test_cli.py diff --git a/Makefile b/Makefile index f34d0536d..d1b4e3a6c 100644 --- a/Makefile +++ b/Makefile @@ -12,4 +12,4 @@ spellcheck: codespell junifer/ docs/ examples/ test: - pytest -v \ No newline at end of file + pytest -vv --cov=junifer --cov-report html --cov-report term \ No newline at end of file diff --git a/junifer/api/tests/test_cli.py b/junifer/api/tests/test_cli.py new file mode 100644 index 000000000..486393b0c --- /dev/null +++ b/junifer/api/tests/test_cli.py @@ -0,0 +1,38 @@ +from pathlib import Path +import tempfile +import yaml +from junifer.api.cli import run, collect + +from click.testing import CliRunner + + +runner = CliRunner() + + +def _modify_path(tmpdir, in_file): + """Modify the path to use the temporary directory""" + if not isinstance(tmpdir, Path): + tmpdir = Path(tmpdir) + with open(in_file, 'r') as f: + contents = yaml.safe_load(f) + outfile = tmpdir / 'in.yaml' + outdir = tmpdir / 'out' + workdir = tmpdir / 'work' + contents['storage']['uri'] = outdir.as_posix() + contents['workdir'] = workdir.as_posix() + with open(outfile, 'w') as f: + yaml.dump(contents, f) + return outfile + + +def test_run_collect(): + """Test run and collect""" + infile = Path(__file__).parent / 'data' / 'gmd_mean.yaml' + with tempfile.TemporaryDirectory() as _tmpdir: + runfile = _modify_path(_tmpdir, infile) + args = [runfile.as_posix(), '--verbose', 'debug'] + response = runner.invoke(run, args) + assert response.exit_code == 0 + + response = runner.invoke(collect, args) + assert response.exit_code == 0 diff --git a/junifer/api/tests/test_parser.py b/junifer/api/tests/test_parser.py index a91604d03..a5dcee7a0 100644 --- a/junifer/api/tests/test_parser.py +++ b/junifer/api/tests/test_parser.py @@ -13,6 +13,13 @@ def test_parse_yaml(): with tempfile.TemporaryDirectory() as _tmpdir: fname = Path(_tmpdir) / 'test.yaml' + + with open(fname, 'w') as f: + f.write('foo: bar\n') + contents = parse_yaml(fname) + assert 'foo' in contents + assert 'bar' == contents['foo'] + with open(fname, 'w') as f: f.write('foo: bar\n') f.write('with: numpy\n') -- 2.52.0 From 585dd0fd76392699f3f5d868a98d0c4e72dbe058 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Thu, 24 Mar 2022 14:56:10 +0100 Subject: [PATCH 043/287] Renamed run, so delete this file --- junifer/api/run.py | 79 ---------------------------------------------- 1 file changed, 79 deletions(-) delete mode 100644 junifer/api/run.py diff --git a/junifer/api/run.py b/junifer/api/run.py deleted file mode 100644 index f26957272..000000000 --- a/junifer/api/run.py +++ /dev/null @@ -1,79 +0,0 @@ -# Authors: Federico Raimondo -# Leonard Sasse -# License: AGPL - -from pathlib import Path - -from .registry import build -from ..datagrabber.base import BaseDataGrabber -from ..markers.base import BaseMarker -from ..storage.base import BaseFeatureStorage -from ..markers.collection import MarkerCollection - - -def run( - workdir, datagrabber, markers, storage, elements=None): - """Run the pipeline on the selected element - - Parameters - ---------- - workdir : str or path-like object - Directory where the pipeline will be executed - datagrabber : dict - Datagrabber to use. Must have a key 'kind' with the kind of - datagrabber to use. All other keys are passed to the datagrabber - init function. - elements : str, tuple or list[str or tuple] - Element(s) to process. Will be used to index the datagrabber. - markers : list of dict - List of markers to extract. Each marker is a dict with at least two - keys: 'name' and 'kind'. The 'name' key is used to name the output - marker. The 'kind' key is used to specify the kind of marker to - extract. The rest of the keys are used to pass parameters to the - marker calculation. - storage : dict - Storage to use. Must have a key 'kind' with the kind of - storage to use. All other keys are passed to the storage - init function. - """ - datagrabber_params = datagrabber.copy() - datagrabber_kind = datagrabber_params.pop('kind') - storage_params = storage.copy() - storage_kind = storage_params.pop('kind') - - if isinstance(workdir, str): - workdir = Path(workdir) - - datagrabber = build( - 'datagrabber', datagrabber_kind, BaseDataGrabber, - init_params=datagrabber_params) - - built_markers = [] - for t_marker in markers: - kind = t_marker.pop('kind') - t_m = build('marker', kind, BaseMarker, init_params=t_marker) - built_markers.append(t_m) - - storage = build( - 'storage', storage_kind, BaseFeatureStorage, - init_params=storage_params) - - mc = MarkerCollection(built_markers, storage=storage) - - with datagrabber: - if elements is not None: - for t_element in elements: - mc.fit(datagrabber[t_element]) - else: - for t_element in datagrabber: - mc.fit(datagrabber[t_element]) - - -def collect(storage): - storage_params = storage.copy() - storage_kind = storage_params.pop('kind') - - storage = build( - 'storage', storage_kind, BaseFeatureStorage, - init_params=storage_params) - storage.collect() -- 2.52.0 From 02bd4a556e0238e653935b669992d004e2edc870 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Thu, 24 Mar 2022 16:23:25 +0100 Subject: [PATCH 044/287] Prototype of junifer queue --- examples/yamls/gmd_mean_htcondor.yaml | 20 +++ junifer/api/cli.py | 69 +++++++-- junifer/api/functions.py | 146 +++++++++++++++++- junifer/api/res/run_conda.sh | 17 ++ junifer/api/tests/data/gmd_mean.yaml | 2 +- junifer/api/tests/data/gmd_mean_htcondor.yaml | 20 +++ junifer_jobs/TestHTCondorQueue/condor.dag | 30 ++++ junifer_jobs/TestHTCondorQueue/condor.submit | 24 +++ junifer_jobs/TestHTCondorQueue/config.yaml | 11 ++ junifer_jobs/TestHTCondorQueue/run_conda.sh | 17 ++ 10 files changed, 333 insertions(+), 23 deletions(-) create mode 100644 examples/yamls/gmd_mean_htcondor.yaml create mode 100644 junifer/api/res/run_conda.sh create mode 100644 junifer/api/tests/data/gmd_mean_htcondor.yaml create mode 100644 junifer_jobs/TestHTCondorQueue/condor.dag create mode 100644 junifer_jobs/TestHTCondorQueue/condor.submit create mode 100644 junifer_jobs/TestHTCondorQueue/config.yaml create mode 100644 junifer_jobs/TestHTCondorQueue/run_conda.sh diff --git a/examples/yamls/gmd_mean_htcondor.yaml b/examples/yamls/gmd_mean_htcondor.yaml new file mode 100644 index 000000000..fddf7aa9b --- /dev/null +++ b/examples/yamls/gmd_mean_htcondor.yaml @@ -0,0 +1,20 @@ +with: junifer.testing.registry +workdir: /tmp + +datagrabber: + kind: OasisVBMTestingDatagrabber +markers: + - name: Schaefer1000x7_Mean + kind: ParcelAggregation + atlas: Schaefer1000x7 + method: mean +storage: + kind: SQLiteFeatureStorage + uri: /Users/fraimondo/dev/tbox/junifer/scratch/db/test.db +queue: + jobname: TestHTCondorQueue + kind: HTCondor + env: + kind: conda + name: junifer + mem: 8G \ No newline at end of file diff --git a/junifer/api/cli.py b/junifer/api/cli.py index c32b21484..5fdc6e276 100644 --- a/junifer/api/cli.py +++ b/junifer/api/cli.py @@ -3,7 +3,27 @@ import click from .parser import parse_yaml from .functions import run as api_run from .functions import collect as api_collect -from ..utils.logging import configure_logging +from .functions import queue as api_queue +from ..utils.logging import configure_logging, logger, warn + + +def _parse_elements(element, config): + logger.debug(f'Parsing elements: {element}') + if len(element) == 0: + return None + # TODO: If len == 1, check if its a file, then parse elements from file + elements = [x.split(',') if ',' in x else x for x in element] + logger.debug(f'Parsed elements: {elements}') + if elements is not None and 'elements' in config: + warn('One or more elements have been specified in both the command ' + 'line and in the config file. The command line has precedence ' + 'over the configuration file. That is, the elements specified ' + 'in the command line will be used. The elements specified in ' + 'the configuration file will be ignored. To remove this warning, ' + 'please remove the "elements" item from the configuration file.') + elif elements is None: + elements = config.get('elements', None) + return elements @click.group() @@ -12,37 +32,54 @@ def cli(): @cli.command() -@click.argument('filepath') +@click.argument('filepath', type=click.File('r')) @click.option('-v', '--verbose', type=click.Choice(['warning', 'info', 'debug'], case_sensitive=False), - default='warning') -def run(filepath, verbose): + default='info') +@click.option('--element', type=str, multiple=True) +def run(filepath, element, verbose): configure_logging(level=verbose.upper()) - contents = parse_yaml(filepath) - workdir = contents['workdir'] - datagrabber = contents['datagrabber'] - markers = contents['markers'] - storage = contents['storage'] - elements = contents.get('elements', None) + config = parse_yaml(filepath) + workdir = config['workdir'] + datagrabber = config['datagrabber'] + markers = config['markers'] + storage = config['storage'] + elements = _parse_elements(element, config) api_run( workdir=workdir, datagrabber=datagrabber, markers=markers, storage=storage, elements=elements) @cli.command() -@click.argument('filepath') +@click.argument('filepath', type=click.File('r')) @click.option('-v', '--verbose', type=click.Choice(['warning', 'info', 'debug'], case_sensitive=False), - default='warning') + default='info') def collect(filepath, verbose): configure_logging(level=verbose.upper()) - contents = parse_yaml(filepath) - storage = contents['storage'] + config = parse_yaml(filepath) + storage = config['storage'] api_collect(storage) @cli.command() -def queue(): - click.echo('queue') +@click.argument( + 'filepath', + type=click.Path(exists=True, readable=True, dir_okay=False)) +@click.option('-v', '--verbose', + type=click.Choice(['warning', 'info', 'debug'], + case_sensitive=False), + default='info') +@click.option('--overwrite', is_flag=True) +@click.option('--submit', is_flag=True) +@click.option('--element', type=str, multiple=True) +def queue(filepath, element, overwrite, submit, verbose): + configure_logging(level=verbose.upper()) + config = parse_yaml(filepath) + elements = _parse_elements(element, config) + queue_config = config.pop('queue') + kind = queue_config.pop('kind') + api_queue(config, kind=kind, overwrite=overwrite, submit=submit, + elements=elements, **queue_config) diff --git a/junifer/api/functions.py b/junifer/api/functions.py index b98273051..5311a2a04 100644 --- a/junifer/api/functions.py +++ b/junifer/api/functions.py @@ -3,14 +3,27 @@ # License: AGPL from pathlib import Path +import shutil +import yaml +import subprocess from .registry import build +from ..utils import logger, raise_error from ..datagrabber.base import BaseDataGrabber from ..markers.base import BaseMarker from ..storage.base import BaseFeatureStorage from ..markers.collection import MarkerCollection +def _get_datagrabber(datagrabber_config): + datagrabber_params = datagrabber_config.copy() + datagrabber_kind = datagrabber_params.pop('kind') + datagrabber = build( + 'datagrabber', datagrabber_kind, BaseDataGrabber, + init_params=datagrabber_params) + return datagrabber + + def run( workdir, datagrabber, markers, storage, elements=None): """Run the pipeline on the selected element @@ -36,17 +49,13 @@ def run( storage to use. All other keys are passed to the storage init function. """ - datagrabber_params = datagrabber.copy() - datagrabber_kind = datagrabber_params.pop('kind') storage_params = storage.copy() storage_kind = storage_params.pop('kind') if isinstance(workdir, str): workdir = Path(workdir) - datagrabber = build( - 'datagrabber', datagrabber_kind, BaseDataGrabber, - init_params=datagrabber_params) + datagrabber = _get_datagrabber(datagrabber) # Copy to avoid changing the original dict _markers = [x.copy() for x in markers] built_markers = [] @@ -73,8 +82,133 @@ def run( def collect(storage): storage_params = storage.copy() storage_kind = storage_params.pop('kind') - + logger.info(f'Collecting data using {storage_kind}') + logger.debug(f'\tStorage params: {storage_params}') storage = build( 'storage', storage_kind, BaseFeatureStorage, init_params=storage_params) + logger.debug('Running storage.collect()') storage.collect() + logger.info('Collect done') + + +def queue(config, kind, jobname='junifer_job', overwrite=False, elements=None, + **kwargs): + """Queue a job to be executed later + + Parameters + ---------- + kind : str + The kind of job to queue. + **kwargs : dict + The parameters to pass to the job. + """ + # Create a folder within the CWD to store the job files / config + cwd = Path.cwd() + job_dir = cwd / 'junifer_jobs' / jobname + logger.info(f'Creating job in {job_dir.as_posix()}') + if job_dir.exists(): + if overwrite is not True: + raise_error(f'Job folder for {jobname} already exists. ' + 'This error is raise to prevent overwriting files ' + 'of jobs that might be scheduled but yet not ' + 'executed. Either delete the directory ' + f'{job_dir.as_posix()} or set overwrite to True.') + job_dir.mkdir(exist_ok=True, parents=True) + + yaml_config = job_dir / 'config.yaml' + logger.info(f'Writing YAML config to {yaml_config}') + with open(yaml_config, 'w') as f: + f.write(yaml.dump(config)) + + # Get list of elements + if elements is None: + if 'elements' in config: + elements = config['elements'] + else: + # If no elements are specified, use all elements from the + # datagrabber + datagrabber = _get_datagrabber(config['datagrabber']) + with datagrabber as dg: + elements = dg.get_elements() + if kind == 'HTCondor': + _queue_condor(job_dir, yaml_config, elements, **kwargs) + elif kind == 'SLURM': + _queue_slurm(job_dir, yaml_config, elements, **kwargs) + else: + raise ValueError(f'Unknown queue kind: {kind}') + + logger.info('Queue done') + + +def _queue_condor(job_dir, yaml_config, elements, env, mem='8G', cpus=1, + disk='1G', extra_preamble='', verbose='info', submit=False): + logger.debug('Creating HTCondor job') + junifer_args = (f'run {yaml_config.as_posix()} ' + f'--verbose {verbose} --element $(element)') + if env is None: + env = {'kind': 'local'} + if env['kind'] == 'conda': + env_name = env['name'] + executable = 'run_conda.sh' + exec_string = f'{env_name} junifer {junifer_args}' + # TODO: Copy run_conda.sh to job_dir + shutil.copy(Path(__file__).parent / 'res' / executable, + job_dir / executable) + elif env['kind'] == 'venv': + env_name = env['name'] + executable = 'run_venv.sh' + exec_string = f'{env_name} junifer {junifer_args}' + # TODO: Copy run_venv.sh to job_dir + elif env['kind'] == 'local': + executable = 'junifer' + exec_string = junifer_args + else: + raise ValueError(f'Unknown env kind: {env["kind"]}') + log_dir = job_dir / 'logs' + log_dir.mkdir(exist_ok=True, parents=True) + preamble = f""" + # The environment + universe = vanilla + getenv = True + + # Resources + request_cpus = {cpus} + request_memory = {mem} + request_disk = {disk} + + # Executable + initial_dir = {job_dir.as_posix()} + executable = $(initial_dir)/{executable} + transfer_executable = False + + arguments = {exec_string} + + {extra_preamble} + + # Logs + log = {log_dir.as_posix()}/junifer_run_$(element).log + output = {log_dir.as_posix()}/junifer_run_$(element).out + error = {log_dir.as_posix()}/junifer_run_$(element).err + """ + + submit_fname = job_dir / 'condor.submit' + dag_fname = job_dir / 'condor.dag' + with open(submit_fname, 'w') as submit_file: + submit_file.write(preamble) + submit_file.write('queue\n') + + with open(dag_fname, 'w') as dag_file: + # Get all subject and session names from file list + for i_job, t_elem in enumerate(elements): + dag_file.write(f'JOB job{i_job} {submit_fname}\n') + dag_file.write(f'VARS job{i_job} element={t_elem}\n\n') + + if submit: + logger.info('Submitting HTCondor job') + subprocess.run(['condor_submit_dag', dag_fname]) + logger.info('HTCondor job submitted') + + +def _queue_slurm(job_dir, yaml_config, elements): + pass diff --git a/junifer/api/res/run_conda.sh b/junifer/api/res/run_conda.sh new file mode 100644 index 000000000..6f859d050 --- /dev/null +++ b/junifer/api/res/run_conda.sh @@ -0,0 +1,17 @@ +#!/bin/bash + +if [ $# -lt 2 ]; then + echo "This script is ment to run a command within a python environment" + echo "It needs at least 2 parameters." + echo "The first one must be the environment name." + echo "The rest will be the command" + exit -1 +fi + +eval "$(conda shell.bash hook)" +env_name=$1 +echo "Activating ${env_name}" +conda activate $1 +shift 1 +echo "Running ${@} in virtual environment" +$@ \ No newline at end of file diff --git a/junifer/api/tests/data/gmd_mean.yaml b/junifer/api/tests/data/gmd_mean.yaml index af4df1798..c492cc559 100644 --- a/junifer/api/tests/data/gmd_mean.yaml +++ b/junifer/api/tests/data/gmd_mean.yaml @@ -3,7 +3,7 @@ workdir: /tmp datagrabber: kind: OasisVBMTestingDatagrabber -elements: +elements: [1, 2] markers: - name: Schaefer1000x7_Mean kind: ParcelAggregation diff --git a/junifer/api/tests/data/gmd_mean_htcondor.yaml b/junifer/api/tests/data/gmd_mean_htcondor.yaml new file mode 100644 index 000000000..fddf7aa9b --- /dev/null +++ b/junifer/api/tests/data/gmd_mean_htcondor.yaml @@ -0,0 +1,20 @@ +with: junifer.testing.registry +workdir: /tmp + +datagrabber: + kind: OasisVBMTestingDatagrabber +markers: + - name: Schaefer1000x7_Mean + kind: ParcelAggregation + atlas: Schaefer1000x7 + method: mean +storage: + kind: SQLiteFeatureStorage + uri: /Users/fraimondo/dev/tbox/junifer/scratch/db/test.db +queue: + jobname: TestHTCondorQueue + kind: HTCondor + env: + kind: conda + name: junifer + mem: 8G \ No newline at end of file diff --git a/junifer_jobs/TestHTCondorQueue/condor.dag b/junifer_jobs/TestHTCondorQueue/condor.dag new file mode 100644 index 000000000..663280160 --- /dev/null +++ b/junifer_jobs/TestHTCondorQueue/condor.dag @@ -0,0 +1,30 @@ +JOB job0 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/condor.submit +VARS job0 element=1 + +JOB job1 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/condor.submit +VARS job1 element=2 + +JOB job2 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/condor.submit +VARS job2 element=3 + +JOB job3 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/condor.submit +VARS job3 element=4 + +JOB job4 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/condor.submit +VARS job4 element=5 + +JOB job5 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/condor.submit +VARS job5 element=6 + +JOB job6 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/condor.submit +VARS job6 element=7 + +JOB job7 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/condor.submit +VARS job7 element=8 + +JOB job8 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/condor.submit +VARS job8 element=9 + +JOB job9 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/condor.submit +VARS job9 element=10 + diff --git a/junifer_jobs/TestHTCondorQueue/condor.submit b/junifer_jobs/TestHTCondorQueue/condor.submit new file mode 100644 index 000000000..48194f711 --- /dev/null +++ b/junifer_jobs/TestHTCondorQueue/condor.submit @@ -0,0 +1,24 @@ + + # The environment + universe = vanilla + getenv = True + + # Resources + request_cpus = 1 + request_memory = 8G + request_disk = 1G + + # Executable + initial_dir = /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue + executable = $(initial_dir)/run_conda.sh + transfer_executable = False + + arguments = junifer junifer run /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/config.yaml --verbose info --element $(element) + + + + # Logs + log = /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/logs/junifer_run_$(element).log + output = /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/logs/junifer_run_$(element).out + error = /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/logs/junifer_run_$(element).err + queue diff --git a/junifer_jobs/TestHTCondorQueue/config.yaml b/junifer_jobs/TestHTCondorQueue/config.yaml new file mode 100644 index 000000000..750062fb0 --- /dev/null +++ b/junifer_jobs/TestHTCondorQueue/config.yaml @@ -0,0 +1,11 @@ +datagrabber: + kind: OasisVBMTestingDatagrabber +markers: +- atlas: Schaefer1000x7 + kind: ParcelAggregation + method: mean + name: Schaefer1000x7_Mean +storage: + kind: SQLiteFeatureStorage + uri: /Users/fraimondo/dev/tbox/junifer/scratch/db/test.db +workdir: /tmp diff --git a/junifer_jobs/TestHTCondorQueue/run_conda.sh b/junifer_jobs/TestHTCondorQueue/run_conda.sh new file mode 100644 index 000000000..6f859d050 --- /dev/null +++ b/junifer_jobs/TestHTCondorQueue/run_conda.sh @@ -0,0 +1,17 @@ +#!/bin/bash + +if [ $# -lt 2 ]; then + echo "This script is ment to run a command within a python environment" + echo "It needs at least 2 parameters." + echo "The first one must be the environment name." + echo "The rest will be the command" + exit -1 +fi + +eval "$(conda shell.bash hook)" +env_name=$1 +echo "Activating ${env_name}" +conda activate $1 +shift 1 +echo "Running ${@} in virtual environment" +$@ \ No newline at end of file -- 2.52.0 From d79fbf0aeda9181382855b21088cc34303057657 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Thu, 24 Mar 2022 16:24:06 +0100 Subject: [PATCH 045/287] damn flake --- junifer/api/functions.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/junifer/api/functions.py b/junifer/api/functions.py index 5311a2a04..a01a939d1 100644 --- a/junifer/api/functions.py +++ b/junifer/api/functions.py @@ -153,7 +153,7 @@ def _queue_condor(job_dir, yaml_config, elements, env, mem='8G', cpus=1, executable = 'run_conda.sh' exec_string = f'{env_name} junifer {junifer_args}' # TODO: Copy run_conda.sh to job_dir - shutil.copy(Path(__file__).parent / 'res' / executable, + shutil.copy(Path(__file__).parent / 'res' / executable, job_dir / executable) elif env['kind'] == 'venv': env_name = env['name'] -- 2.52.0 From 8f411c2cc3247be5962782f9c104aee5995a2924 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 25 Mar 2022 09:33:04 +0100 Subject: [PATCH 046/287] Ready to test HTCondor --- junifer/api/functions.py | 79 +++++++++++++++---- junifer/testing.py | 37 --------- junifer/utils/fs.py | 7 ++ junifer_jobs/TestHTCondorQueue/collect.submit | 24 ++++++ junifer_jobs/TestHTCondorQueue/condor.dag | 43 +++++----- junifer_jobs/TestHTCondorQueue/run.submit | 24 ++++++ junifer_jobs/TestHTCondorQueue/run_conda.sh | 0 7 files changed, 141 insertions(+), 73 deletions(-) delete mode 100644 junifer/testing.py create mode 100644 junifer/utils/fs.py create mode 100644 junifer_jobs/TestHTCondorQueue/collect.submit create mode 100644 junifer_jobs/TestHTCondorQueue/run.submit mode change 100644 => 100755 junifer_jobs/TestHTCondorQueue/run_conda.sh diff --git a/junifer/api/functions.py b/junifer/api/functions.py index a01a939d1..5682f55b3 100644 --- a/junifer/api/functions.py +++ b/junifer/api/functions.py @@ -9,6 +9,7 @@ import subprocess from .registry import build from ..utils import logger, raise_error +from ..utils.fs import make_executable from ..datagrabber.base import BaseDataGrabber from ..markers.base import BaseMarker from ..storage.base import BaseFeatureStorage @@ -142,32 +143,39 @@ def queue(config, kind, jobname='junifer_job', overwrite=False, elements=None, def _queue_condor(job_dir, yaml_config, elements, env, mem='8G', cpus=1, - disk='1G', extra_preamble='', verbose='info', submit=False): + disk='1G', extra_preamble='', verbose='info', collect=True, + submit=False): logger.debug('Creating HTCondor job') - junifer_args = (f'run {yaml_config.as_posix()} ' - f'--verbose {verbose} --element $(element)') + run_junifer_args = (f'run {yaml_config.as_posix()} ' + f'--verbose {verbose} --element $(element)') + collect_junifer_args = \ + f'collect {yaml_config.as_posix()} --verbose {verbose} ' + + # Set up the env_name, executable and arguments according to the + # environment type if env is None: env = {'kind': 'local'} if env['kind'] == 'conda': env_name = env['name'] executable = 'run_conda.sh' - exec_string = f'{env_name} junifer {junifer_args}' + arguments = f'{env_name} junifer' # TODO: Copy run_conda.sh to job_dir - shutil.copy(Path(__file__).parent / 'res' / executable, - job_dir / executable) + exec_path = job_dir / executable + shutil.copy(Path(__file__).parent / 'res' / executable, exec_path) + make_executable(exec_path) elif env['kind'] == 'venv': env_name = env['name'] executable = 'run_venv.sh' - exec_string = f'{env_name} junifer {junifer_args}' + arguments = f'{env_name} junifer' # TODO: Copy run_venv.sh to job_dir elif env['kind'] == 'local': executable = 'junifer' - exec_string = junifer_args + arguments = '' else: raise ValueError(f'Unknown env kind: {env["kind"]}') log_dir = job_dir / 'logs' log_dir.mkdir(exist_ok=True, parents=True) - preamble = f""" + run_preamble = f""" # The environment universe = vanilla getenv = True @@ -182,7 +190,7 @@ def _queue_condor(job_dir, yaml_config, elements, env, mem='8G', cpus=1, executable = $(initial_dir)/{executable} transfer_executable = False - arguments = {exec_string} + arguments = {arguments} {run_junifer_args} {extra_preamble} @@ -192,19 +200,58 @@ def _queue_condor(job_dir, yaml_config, elements, env, mem='8G', cpus=1, error = {log_dir.as_posix()}/junifer_run_$(element).err """ - submit_fname = job_dir / 'condor.submit' + submit_run_fname = job_dir / 'run.submit' + submit_collect_fname = job_dir / 'collect.submit' dag_fname = job_dir / 'condor.dag' - with open(submit_fname, 'w') as submit_file: - submit_file.write(preamble) + + # Write to run submit files + with open(submit_run_fname, 'w') as submit_file: + submit_file.write(run_preamble) + submit_file.write('queue\n') + + collect_preamble = f""" + # The environment + universe = vanilla + getenv = True + + # Resources + request_cpus = {cpus} + request_memory = {mem} + request_disk = {disk} + + # Executable + initial_dir = {job_dir.as_posix()} + executable = $(initial_dir)/{executable} + transfer_executable = False + + arguments = {arguments} {collect_junifer_args} + + {extra_preamble} + + # Logs + log = {log_dir.as_posix()}/junifer_collect.log + output = {log_dir.as_posix()}/junifer_collect.out + error = {log_dir.as_posix()}/junifer_collect.err + """ + + # Now create the collect submit file + with open(submit_collect_fname, 'w') as submit_file: + submit_file.write(collect_preamble) # Eval preamble here submit_file.write('queue\n') with open(dag_fname, 'w') as dag_file: # Get all subject and session names from file list for i_job, t_elem in enumerate(elements): - dag_file.write(f'JOB job{i_job} {submit_fname}\n') - dag_file.write(f'VARS job{i_job} element={t_elem}\n\n') + dag_file.write(f'JOB run{i_job} {submit_run_fname}\n') + dag_file.write(f'VARS run{i_job} element={t_elem}\n\n') + if collect is True: + dag_file.write(f'JOB collect {submit_collect_fname}\n') + dag_file.write('PARENT ') + for i_job, t_elem in enumerate(elements): + dag_file.write(f'run{i_job} ') + dag_file.write(f'CHILD collect\n\n') - if submit: + if submit is True: logger.info('Submitting HTCondor job') subprocess.run(['condor_submit_dag', dag_fname]) logger.info('HTCondor job submitted') diff --git a/junifer/testing.py b/junifer/testing.py deleted file mode 100644 index 999f58f69..000000000 --- a/junifer/testing.py +++ /dev/null @@ -1,37 +0,0 @@ -# Authors: Federico Raimondo -# License: AGPL -import tempfile -from nilearn import datasets - -from .datagrabber.base import BaseDataGrabber -from .api.registry import register - - -def register_testing(): - """Register testing datagrabber""" - register( - 'datagrabber', 'OasisVBMTestingDatagrabber', - OasisVBMTestingDatagrabber) - - -class OasisVBMTestingDatagrabber(BaseDataGrabber): - """ - DataGrabber for Oasis VBM testing data. - """ - def __init__(self): - datadir = tempfile.mkdtemp() - types = ['VBM_GM'] - super().__init__(types=types, datadir=datadir) - - def get_elements(self): - return list(range(1, 11)) - - def __getitem__(self, element): - out = {} - out['VBM_GM'] = self._dataset.gray_matter_maps[element - 1] - out['meta'] = {'element': {'subject': element}} - return out - - def __enter__(self): - self._dataset = datasets.fetch_oasis_vbm(n_subjects=10) - return self diff --git a/junifer/utils/fs.py b/junifer/utils/fs.py new file mode 100644 index 000000000..667485bb6 --- /dev/null +++ b/junifer/utils/fs.py @@ -0,0 +1,7 @@ +import os +import stat + + +def make_executable(path): + st = os.stat(path) + os.chmod(path, st.st_mode | stat.S_IEXEC) diff --git a/junifer_jobs/TestHTCondorQueue/collect.submit b/junifer_jobs/TestHTCondorQueue/collect.submit new file mode 100644 index 000000000..abf768d65 --- /dev/null +++ b/junifer_jobs/TestHTCondorQueue/collect.submit @@ -0,0 +1,24 @@ + + # The environment + universe = vanilla + getenv = True + + # Resources + request_cpus = 1 + request_memory = 8G + request_disk = 1G + + # Executable + initial_dir = /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue + executable = $(initial_dir)/run_conda.sh + transfer_executable = False + + arguments = junifer junifer collect /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/config.yaml --verbose info + + + + # Logs + log = /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/logs/junifer_run_$(element).log + output = /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/logs/junifer_run_$(element).out + error = /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/logs/junifer_run_$(element).err + queue diff --git a/junifer_jobs/TestHTCondorQueue/condor.dag b/junifer_jobs/TestHTCondorQueue/condor.dag index 663280160..c99e2889f 100644 --- a/junifer_jobs/TestHTCondorQueue/condor.dag +++ b/junifer_jobs/TestHTCondorQueue/condor.dag @@ -1,30 +1,33 @@ -JOB job0 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/condor.submit -VARS job0 element=1 +JOB run0 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/run.submit +VARS run0 element=1 -JOB job1 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/condor.submit -VARS job1 element=2 +JOB run1 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/run.submit +VARS run1 element=2 -JOB job2 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/condor.submit -VARS job2 element=3 +JOB run2 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/run.submit +VARS run2 element=3 -JOB job3 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/condor.submit -VARS job3 element=4 +JOB run3 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/run.submit +VARS run3 element=4 -JOB job4 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/condor.submit -VARS job4 element=5 +JOB run4 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/run.submit +VARS run4 element=5 -JOB job5 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/condor.submit -VARS job5 element=6 +JOB run5 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/run.submit +VARS run5 element=6 -JOB job6 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/condor.submit -VARS job6 element=7 +JOB run6 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/run.submit +VARS run6 element=7 -JOB job7 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/condor.submit -VARS job7 element=8 +JOB run7 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/run.submit +VARS run7 element=8 -JOB job8 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/condor.submit -VARS job8 element=9 +JOB run8 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/run.submit +VARS run8 element=9 -JOB job9 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/condor.submit -VARS job9 element=10 +JOB run9 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/run.submit +VARS run9 element=10 + +JOB collect /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/collect.submit +PARENT run0 run1 run2 run3 run4 run5 run6 run7 run8 run9 CHILD collect diff --git a/junifer_jobs/TestHTCondorQueue/run.submit b/junifer_jobs/TestHTCondorQueue/run.submit new file mode 100644 index 000000000..48194f711 --- /dev/null +++ b/junifer_jobs/TestHTCondorQueue/run.submit @@ -0,0 +1,24 @@ + + # The environment + universe = vanilla + getenv = True + + # Resources + request_cpus = 1 + request_memory = 8G + request_disk = 1G + + # Executable + initial_dir = /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue + executable = $(initial_dir)/run_conda.sh + transfer_executable = False + + arguments = junifer junifer run /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/config.yaml --verbose info --element $(element) + + + + # Logs + log = /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/logs/junifer_run_$(element).log + output = /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/logs/junifer_run_$(element).out + error = /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/logs/junifer_run_$(element).err + queue diff --git a/junifer_jobs/TestHTCondorQueue/run_conda.sh b/junifer_jobs/TestHTCondorQueue/run_conda.sh old mode 100644 new mode 100755 -- 2.52.0 From 62a8e3423fb6072adfea15e184c1ecaeff0c3c27 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 25 Mar 2022 10:36:35 +0100 Subject: [PATCH 047/287] Fixes for CI + test in juseless --- junifer/api/cli.py | 8 +++++-- junifer/api/res/run_conda.sh | 10 ++++----- junifer/api/tests/test_cli.py | 28 +++++++++++++++++++++--- junifer/api/tests/test_functions.py | 4 ++-- junifer/markers/tests/test_collection.py | 8 +++---- junifer/testing/datagrabbers.py | 5 +++-- 6 files changed, 45 insertions(+), 18 deletions(-) mode change 100644 => 100755 junifer/api/res/run_conda.sh diff --git a/junifer/api/cli.py b/junifer/api/cli.py index 5fdc6e276..6991ae713 100644 --- a/junifer/api/cli.py +++ b/junifer/api/cli.py @@ -32,7 +32,9 @@ def cli(): @cli.command() -@click.argument('filepath', type=click.File('r')) +@click.argument( + 'filepath', + type=click.Path(exists=True, readable=True, dir_okay=False)) @click.option('-v', '--verbose', type=click.Choice(['warning', 'info', 'debug'], case_sensitive=False), @@ -52,7 +54,9 @@ def run(filepath, element, verbose): @cli.command() -@click.argument('filepath', type=click.File('r')) +@click.argument( + 'filepath', + type=click.Path(exists=True, readable=True, dir_okay=False)) @click.option('-v', '--verbose', type=click.Choice(['warning', 'info', 'debug'], case_sensitive=False), diff --git a/junifer/api/res/run_conda.sh b/junifer/api/res/run_conda.sh old mode 100644 new mode 100755 index 6f859d050..f07c86c6f --- a/junifer/api/res/run_conda.sh +++ b/junifer/api/res/run_conda.sh @@ -1,17 +1,17 @@ #!/bin/bash if [ $# -lt 2 ]; then - echo "This script is ment to run a command within a python environment" + echo "This script is meant to run a command within a python environment" echo "It needs at least 2 parameters." echo "The first one must be the environment name." echo "The rest will be the command" - exit -1 + exit 255 fi eval "$(conda shell.bash hook)" env_name=$1 echo "Activating ${env_name}" -conda activate $1 +conda activate "$1" shift 1 -echo "Running ${@} in virtual environment" -$@ \ No newline at end of file +echo "Running ${*} in virtual environment" +"$@" \ No newline at end of file diff --git a/junifer/api/tests/test_cli.py b/junifer/api/tests/test_cli.py index 486393b0c..c29b7262e 100644 --- a/junifer/api/tests/test_cli.py +++ b/junifer/api/tests/test_cli.py @@ -30,9 +30,31 @@ def test_run_collect(): infile = Path(__file__).parent / 'data' / 'gmd_mean.yaml' with tempfile.TemporaryDirectory() as _tmpdir: runfile = _modify_path(_tmpdir, infile) - args = [runfile.as_posix(), '--verbose', 'debug'] - response = runner.invoke(run, args) + run_args = [runfile.as_posix(), '--verbose', 'debug', '--element', + 'sub-01', '--element', 'sub-02', '--element', 'sub-03'] + response = runner.invoke(run, run_args) assert response.exit_code == 0 - response = runner.invoke(collect, args) + # TODO: Check that there are 3 files in the output directory + + collect_args = [runfile.as_posix(), '--verbose', 'debug'] + response = runner.invoke(collect, collect_args) assert response.exit_code == 0 + + # TODO: Check that there are 4 files in the output directory + # TODO: Check that the collected file has the correct number of rows + + run_args = [runfile.as_posix(), '--verbose', 'debug', '--element', + 'sub-01', '--element', 'sub-02', '--element', 'sub-04'] + response = runner.invoke(run, run_args) + assert response.exit_code == 0 + + # TODO: Check that there are 5 files in the output directory + # TODO: Check that the collected file has 3 rows (sub-04 is not there) + + collect_args = [runfile.as_posix(), '--verbose', 'debug'] + response = runner.invoke(collect, collect_args) + assert response.exit_code == 0 + + # TODO: Check that there are 5 files in the output directory + # TODO: Check that the collected file has 4 rows (sub-04 is there) diff --git a/junifer/api/tests/test_functions.py b/junifer/api/tests/test_functions.py index 5fc65510d..ae44675fc 100644 --- a/junifer/api/tests/test_functions.py +++ b/junifer/api/tests/test_functions.py @@ -41,7 +41,7 @@ def test_run(): datagrabber=datagrabber, markers=markers, storage=storage, - elements=[1] + elements=['sub-01'] ) files = list(outdir.glob('*.db')) @@ -52,7 +52,7 @@ def test_run(): datagrabber=datagrabber, markers=markers, storage=storage, - elements=[1, 3] + elements=['sub-01', 'sub-03'] ) files = list(outdir.glob('*.db')) diff --git a/junifer/markers/tests/test_collection.py b/junifer/markers/tests/test_collection.py index 07f601992..bb0bdf154 100644 --- a/junifer/markers/tests/test_collection.py +++ b/junifer/markers/tests/test_collection.py @@ -47,7 +47,7 @@ def test_MarkerCollection(): mc.validate(dg) with dg: - input = dg[1] + input = dg['sub-01'] out = mc.fit(input) assert out is not None assert isinstance(out, dict) @@ -74,7 +74,7 @@ def test_MarkerCollection(): datareader=DefaultDataReader()) assert isinstance(mc2._datareader, DefaultDataReader) with dg: - input = dg[1] + input = dg['sub-01'] out2 = mc2.fit(input) for t_marker in markers: t_name = t_marker.name @@ -106,7 +106,7 @@ def test_MarkerCollection_storage(): mc.validate(dg) assert mc._storage.uri == storage.uri with dg: - input = dg[1] + input = dg['sub-01'] out = mc.fit(input) assert out is None @@ -116,7 +116,7 @@ def test_MarkerCollection_storage(): assert mc2._storage is None with dg: - input = dg[1] + input = dg['sub-01'] out = mc2.fit(input) features = storage.list_features() diff --git a/junifer/testing/datagrabbers.py b/junifer/testing/datagrabbers.py index fba285064..e39511163 100644 --- a/junifer/testing/datagrabbers.py +++ b/junifer/testing/datagrabbers.py @@ -16,11 +16,12 @@ class OasisVBMTestingDatagrabber(BaseDataGrabber): super().__init__(types=types, datadir=datadir) def get_elements(self): - return list(range(1, 11)) + return [f'sub-{x:02d}' for x in list(range(1, 11))] def __getitem__(self, element): out = {} - out['VBM_GM'] = self._dataset.gray_matter_maps[element - 1] + i_sub = int(element.split('-')[1]) - 1 + out['VBM_GM'] = self._dataset.gray_matter_maps[i_sub] out['meta'] = {'element': {'subject': element}} return out -- 2.52.0 From 2239256e517c926805f3b99d534c5773a909be72 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 25 Mar 2022 10:37:43 +0100 Subject: [PATCH 048/287] Remove files not belonging --- junifer_jobs/TestHTCondorQueue/collect.submit | 24 -------------- junifer_jobs/TestHTCondorQueue/condor.dag | 33 ------------------- junifer_jobs/TestHTCondorQueue/condor.submit | 24 -------------- junifer_jobs/TestHTCondorQueue/config.yaml | 11 ------- junifer_jobs/TestHTCondorQueue/run.submit | 24 -------------- junifer_jobs/TestHTCondorQueue/run_conda.sh | 17 ---------- 6 files changed, 133 deletions(-) delete mode 100644 junifer_jobs/TestHTCondorQueue/collect.submit delete mode 100644 junifer_jobs/TestHTCondorQueue/condor.dag delete mode 100644 junifer_jobs/TestHTCondorQueue/condor.submit delete mode 100644 junifer_jobs/TestHTCondorQueue/config.yaml delete mode 100644 junifer_jobs/TestHTCondorQueue/run.submit delete mode 100755 junifer_jobs/TestHTCondorQueue/run_conda.sh diff --git a/junifer_jobs/TestHTCondorQueue/collect.submit b/junifer_jobs/TestHTCondorQueue/collect.submit deleted file mode 100644 index abf768d65..000000000 --- a/junifer_jobs/TestHTCondorQueue/collect.submit +++ /dev/null @@ -1,24 +0,0 @@ - - # The environment - universe = vanilla - getenv = True - - # Resources - request_cpus = 1 - request_memory = 8G - request_disk = 1G - - # Executable - initial_dir = /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue - executable = $(initial_dir)/run_conda.sh - transfer_executable = False - - arguments = junifer junifer collect /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/config.yaml --verbose info - - - - # Logs - log = /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/logs/junifer_run_$(element).log - output = /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/logs/junifer_run_$(element).out - error = /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/logs/junifer_run_$(element).err - queue diff --git a/junifer_jobs/TestHTCondorQueue/condor.dag b/junifer_jobs/TestHTCondorQueue/condor.dag deleted file mode 100644 index c99e2889f..000000000 --- a/junifer_jobs/TestHTCondorQueue/condor.dag +++ /dev/null @@ -1,33 +0,0 @@ -JOB run0 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/run.submit -VARS run0 element=1 - -JOB run1 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/run.submit -VARS run1 element=2 - -JOB run2 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/run.submit -VARS run2 element=3 - -JOB run3 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/run.submit -VARS run3 element=4 - -JOB run4 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/run.submit -VARS run4 element=5 - -JOB run5 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/run.submit -VARS run5 element=6 - -JOB run6 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/run.submit -VARS run6 element=7 - -JOB run7 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/run.submit -VARS run7 element=8 - -JOB run8 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/run.submit -VARS run8 element=9 - -JOB run9 /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/run.submit -VARS run9 element=10 - -JOB collect /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/collect.submit -PARENT run0 run1 run2 run3 run4 run5 run6 run7 run8 run9 CHILD collect - diff --git a/junifer_jobs/TestHTCondorQueue/condor.submit b/junifer_jobs/TestHTCondorQueue/condor.submit deleted file mode 100644 index 48194f711..000000000 --- a/junifer_jobs/TestHTCondorQueue/condor.submit +++ /dev/null @@ -1,24 +0,0 @@ - - # The environment - universe = vanilla - getenv = True - - # Resources - request_cpus = 1 - request_memory = 8G - request_disk = 1G - - # Executable - initial_dir = /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue - executable = $(initial_dir)/run_conda.sh - transfer_executable = False - - arguments = junifer junifer run /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/config.yaml --verbose info --element $(element) - - - - # Logs - log = /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/logs/junifer_run_$(element).log - output = /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/logs/junifer_run_$(element).out - error = /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/logs/junifer_run_$(element).err - queue diff --git a/junifer_jobs/TestHTCondorQueue/config.yaml b/junifer_jobs/TestHTCondorQueue/config.yaml deleted file mode 100644 index 750062fb0..000000000 --- a/junifer_jobs/TestHTCondorQueue/config.yaml +++ /dev/null @@ -1,11 +0,0 @@ -datagrabber: - kind: OasisVBMTestingDatagrabber -markers: -- atlas: Schaefer1000x7 - kind: ParcelAggregation - method: mean - name: Schaefer1000x7_Mean -storage: - kind: SQLiteFeatureStorage - uri: /Users/fraimondo/dev/tbox/junifer/scratch/db/test.db -workdir: /tmp diff --git a/junifer_jobs/TestHTCondorQueue/run.submit b/junifer_jobs/TestHTCondorQueue/run.submit deleted file mode 100644 index 48194f711..000000000 --- a/junifer_jobs/TestHTCondorQueue/run.submit +++ /dev/null @@ -1,24 +0,0 @@ - - # The environment - universe = vanilla - getenv = True - - # Resources - request_cpus = 1 - request_memory = 8G - request_disk = 1G - - # Executable - initial_dir = /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue - executable = $(initial_dir)/run_conda.sh - transfer_executable = False - - arguments = junifer junifer run /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/config.yaml --verbose info --element $(element) - - - - # Logs - log = /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/logs/junifer_run_$(element).log - output = /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/logs/junifer_run_$(element).out - error = /Users/fraimondo/dev/tbox/junifer/junifer_jobs/TestHTCondorQueue/logs/junifer_run_$(element).err - queue diff --git a/junifer_jobs/TestHTCondorQueue/run_conda.sh b/junifer_jobs/TestHTCondorQueue/run_conda.sh deleted file mode 100755 index 6f859d050..000000000 --- a/junifer_jobs/TestHTCondorQueue/run_conda.sh +++ /dev/null @@ -1,17 +0,0 @@ -#!/bin/bash - -if [ $# -lt 2 ]; then - echo "This script is ment to run a command within a python environment" - echo "It needs at least 2 parameters." - echo "The first one must be the environment name." - echo "The rest will be the command" - exit -1 -fi - -eval "$(conda shell.bash hook)" -env_name=$1 -echo "Activating ${env_name}" -conda activate $1 -shift 1 -echo "Running ${@} in virtual environment" -$@ \ No newline at end of file -- 2.52.0 From abaa5d2917d5e7f1191250c472512596fbe877ca Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 25 Mar 2022 12:03:55 +0100 Subject: [PATCH 049/287] Fix example + Fix bug with testing datagrabber not using proper met --- examples/yamls/gmd_mean_htcondor.yaml | 2 +- junifer/api/functions.py | 52 +++++++++++++++------------ junifer/api/parser.py | 6 ++-- junifer/datagrabber/base.py | 8 ++--- junifer/storage/base.py | 16 +++++---- junifer/storage/sqlite.py | 28 +++++++++------ junifer/storage/tests/test_sqlite.py | 7 ++++ junifer/testing/datagrabbers.py | 5 +-- 8 files changed, 77 insertions(+), 47 deletions(-) diff --git a/examples/yamls/gmd_mean_htcondor.yaml b/examples/yamls/gmd_mean_htcondor.yaml index fddf7aa9b..fde41b733 100644 --- a/examples/yamls/gmd_mean_htcondor.yaml +++ b/examples/yamls/gmd_mean_htcondor.yaml @@ -10,7 +10,7 @@ markers: method: mean storage: kind: SQLiteFeatureStorage - uri: /Users/fraimondo/dev/tbox/junifer/scratch/db/test.db + uri: /data/group/appliedml/fraimondo/junifer_test/test.db queue: jobname: TestHTCondorQueue kind: HTCondor diff --git a/junifer/api/functions.py b/junifer/api/functions.py index 5682f55b3..e3fb8b80d 100644 --- a/junifer/api/functions.py +++ b/junifer/api/functions.py @@ -106,18 +106,22 @@ def queue(config, kind, jobname='junifer_job', overwrite=False, elements=None, """ # Create a folder within the CWD to store the job files / config cwd = Path.cwd() - job_dir = cwd / 'junifer_jobs' / jobname - logger.info(f'Creating job in {job_dir.as_posix()}') - if job_dir.exists(): + jobdir = cwd / 'junifer_jobs' / jobname + logger.info(f'Creating job in {jobdir.as_posix()}') + if jobdir.exists(): if overwrite is not True: raise_error(f'Job folder for {jobname} already exists. ' 'This error is raise to prevent overwriting files ' 'of jobs that might be scheduled but yet not ' 'executed. Either delete the directory ' - f'{job_dir.as_posix()} or set overwrite to True.') - job_dir.mkdir(exist_ok=True, parents=True) + f'{jobdir.as_posix()} or set overwrite to True.') + else: + logger.info( + f'Deleting previous job directory {jobdir.as_posix()}') + shutil.rmtree(jobdir) + jobdir.mkdir(exist_ok=True, parents=True) - yaml_config = job_dir / 'config.yaml' + yaml_config = jobdir / 'config.yaml' logger.info(f'Writing YAML config to {yaml_config}') with open(yaml_config, 'w') as f: f.write(yaml.dump(config)) @@ -133,18 +137,18 @@ def queue(config, kind, jobname='junifer_job', overwrite=False, elements=None, with datagrabber as dg: elements = dg.get_elements() if kind == 'HTCondor': - _queue_condor(job_dir, yaml_config, elements, **kwargs) + _queue_condor(jobname, jobdir, yaml_config, elements, **kwargs) elif kind == 'SLURM': - _queue_slurm(job_dir, yaml_config, elements, **kwargs) + _queue_slurm(jobname, jobdir, yaml_config, elements, **kwargs) else: raise ValueError(f'Unknown queue kind: {kind}') logger.info('Queue done') -def _queue_condor(job_dir, yaml_config, elements, env, mem='8G', cpus=1, - disk='1G', extra_preamble='', verbose='info', collect=True, - submit=False): +def _queue_condor(jobname, jobdir, yaml_config, elements, env, mem='8G', + cpus=1, disk='1G', extra_preamble='', verbose='info', + collect=True, submit=False): logger.debug('Creating HTCondor job') run_junifer_args = (f'run {yaml_config.as_posix()} ' f'--verbose {verbose} --element $(element)') @@ -159,21 +163,21 @@ def _queue_condor(job_dir, yaml_config, elements, env, mem='8G', cpus=1, env_name = env['name'] executable = 'run_conda.sh' arguments = f'{env_name} junifer' - # TODO: Copy run_conda.sh to job_dir - exec_path = job_dir / executable + # TODO: Copy run_conda.sh to jobdir + exec_path = jobdir / executable shutil.copy(Path(__file__).parent / 'res' / executable, exec_path) make_executable(exec_path) elif env['kind'] == 'venv': env_name = env['name'] executable = 'run_venv.sh' arguments = f'{env_name} junifer' - # TODO: Copy run_venv.sh to job_dir + # TODO: Copy run_venv.sh to jobdir elif env['kind'] == 'local': executable = 'junifer' arguments = '' else: raise ValueError(f'Unknown env kind: {env["kind"]}') - log_dir = job_dir / 'logs' + log_dir = jobdir / 'logs' log_dir.mkdir(exist_ok=True, parents=True) run_preamble = f""" # The environment @@ -186,7 +190,7 @@ def _queue_condor(job_dir, yaml_config, elements, env, mem='8G', cpus=1, request_disk = {disk} # Executable - initial_dir = {job_dir.as_posix()} + initial_dir = {jobdir.as_posix()} executable = $(initial_dir)/{executable} transfer_executable = False @@ -200,9 +204,9 @@ def _queue_condor(job_dir, yaml_config, elements, env, mem='8G', cpus=1, error = {log_dir.as_posix()}/junifer_run_$(element).err """ - submit_run_fname = job_dir / 'run.submit' - submit_collect_fname = job_dir / 'collect.submit' - dag_fname = job_dir / 'condor.dag' + submit_run_fname = jobdir / f'run_{jobname}.submit' + submit_collect_fname = jobdir / f'collect_{jobname}.submit' + dag_fname = jobdir / f'{jobname}.dag' # Write to run submit files with open(submit_run_fname, 'w') as submit_file: @@ -220,7 +224,7 @@ def _queue_condor(job_dir, yaml_config, elements, env, mem='8G', cpus=1, request_disk = {disk} # Executable - initial_dir = {job_dir.as_posix()} + initial_dir = {jobdir.as_posix()} executable = $(initial_dir)/{executable} transfer_executable = False @@ -243,7 +247,7 @@ def _queue_condor(job_dir, yaml_config, elements, env, mem='8G', cpus=1, # Get all subject and session names from file list for i_job, t_elem in enumerate(elements): dag_file.write(f'JOB run{i_job} {submit_run_fname}\n') - dag_file.write(f'VARS run{i_job} element={t_elem}\n\n') + dag_file.write(f'VARS run{i_job} element="{t_elem}"\n\n') if collect is True: dag_file.write(f'JOB collect {submit_collect_fname}\n') dag_file.write('PARENT ') @@ -255,7 +259,11 @@ def _queue_condor(job_dir, yaml_config, elements, env, mem='8G', cpus=1, logger.info('Submitting HTCondor job') subprocess.run(['condor_submit_dag', dag_fname]) logger.info('HTCondor job submitted') + else: + cmd = f'condor_submit_dag {dag_fname.as_posix()}' + logger.info('HTCondor job files created, to submit the job, ' + f'run "{cmd}"') -def _queue_slurm(job_dir, yaml_config, elements): +def _queue_slurm(jobname, jobdir, yaml_config, elements): pass diff --git a/junifer/api/parser.py b/junifer/api/parser.py index 03b041362..a5883ad0a 100644 --- a/junifer/api/parser.py +++ b/junifer/api/parser.py @@ -2,22 +2,24 @@ import yaml import importlib from pathlib import Path -from ..utils.logging import raise_error +from ..utils.logging import raise_error, logger def parse_yaml(filepath): if not isinstance(filepath, Path): filepath = Path(filepath) + logger.info(f'Parsing yaml file: {filepath.as_posix()}') if not filepath.exists(): raise_error(f'File does not exist: {filepath.as_posix()}') with open(filepath, 'r') as f: contents = yaml.safe_load(f) if 'with' in contents: - to_load = contents.pop('with') + to_load = contents['with'] if not isinstance(to_load, list): to_load = [to_load] for t_module in to_load: + logger.info(f'Importing module {t_module}') importlib.import_module(t_module) return contents diff --git a/junifer/datagrabber/base.py b/junifer/datagrabber/base.py index 7df89e8d2..732df21bf 100644 --- a/junifer/datagrabber/base.py +++ b/junifer/datagrabber/base.py @@ -111,9 +111,10 @@ class BaseDataGrabber(ABC): for elem in self.get_elements(): yield elem - @abstractmethod def __getitem__(self, element): - raise NotImplementedError('__getitem__ not implemented') + out = {} + out['meta'] = dict(datagrabber=self.get_meta()) + return out @abstractmethod def get_elements(self): @@ -196,7 +197,7 @@ class BIDSDataGrabber(BaseDataGrabber): Dictionary of paths for each type of data required for the specified element. """ - out = {} + out = super().__getitem__(element) if not isinstance(element, tuple): element = (element,) for t_type in self.types: @@ -208,7 +209,6 @@ class BIDSDataGrabber(BaseDataGrabber): t_out = self.datadir / element[0] / t_replace out[t_type] = dict(path=t_out) # Meta here is element and types - out['meta'] = dict(datagrabber=self.get_meta()) out['meta']['element'] = {'subject': element[0]} if len(element) > 1: out['meta']['element']['session'] = element[1] # type: ignore diff --git a/junifer/storage/base.py b/junifer/storage/base.py index da1f26aa5..dbdf0f204 100644 --- a/junifer/storage/base.py +++ b/junifer/storage/base.py @@ -161,15 +161,19 @@ class BaseFeatureStorage(ABC): raise NotImplementedError('validate_input not implemented') @abstractmethod - def list_features(self): + def list_features(self, return_df=False): """List the features in the storage - + Parameters + ---------- + return_df : bool + If True, return a dataframe. If False, (default) return a + dictionary Returns ------- - features: dict(str, dict) - List of features in the storage. The keys are the feature names - to be used in read_features. The values are the metadata of each - feature + features: dict(str, dict) | pd.DataFrame + List of features in the storage. If dictionarly, the keys are the + feature names to be used in read_features. The values are the + metadata of each feature. """ raise NotImplementedError('list_features not implemented') diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py index 9616ff7f6..d063d194c 100644 --- a/junifer/storage/sqlite.py +++ b/junifer/storage/sqlite.py @@ -23,14 +23,15 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): Parameters ---------- - uri : str - The connection URI. - Easy options: - 'sqlite://' for an in memory sqlite database - 'sqlite:///' to save in a file - - Check https://docs.sqlalchemy.org/en/14/core/engines.html for more - options + uri : str or Path (must be a file) + The Path to the file to be used. + single_output : bool + If False (default), will create one file per element. The name + of the file will be prefixed with the respective element. + If True, will create only one file, the specified in the URI and + store all the elements in the same file. This behaviour is only + suitable for non-parallel executions. SQLite does not support + concurrency. upsert : str Upsert mode. Options are 'ignore' and 'update' (default). If 'ignore', the existing elements are ignored. If update, the @@ -41,6 +42,10 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): raise ValueError('upsert must be either "update" or "ignore"') if not isinstance(uri, Path): uri = Path(uri) + if not uri.parent.exists(): + logger.info(f'Output directory ({uri.parent.as_posix()}) ' + 'does not exist, creating') + uri.parent.mkdir(parents=True, exist_ok=True) super().__init__(uri, single_output=single_output) self._upsert = upsert self._valid_inputs = ['table', 'timeseries'] @@ -64,10 +69,13 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): uri = f'sqlite:///{self.uri.parent}/{prefix}{self.uri.name}' return create_engine(uri, echo=False) - def list_features(self): + def list_features(self, return_df=False): meta_df = pd.read_sql( 'meta', con=self.get_engine(), index_col='meta_md5') - return meta_df.to_dict(orient='index') + out = meta_df + if return_df is False: + out = meta_df.to_dict(orient='index') + return out def read_df(self, feature_name=None, feature_md5=None): """Read features from the storage. diff --git a/junifer/storage/tests/test_sqlite.py b/junifer/storage/tests/test_sqlite.py index 55e3670e7..a4688497d 100644 --- a/junifer/storage/tests/test_sqlite.py +++ b/junifer/storage/tests/test_sqlite.py @@ -63,6 +63,13 @@ def test_get_engine(): with pytest.raises(ValueError, match='element must be specified'): storage.get_engine() + tocreate = Path(_tmpdir) / 'tocreate' + assert not tocreate.exists() + uri = f'{tocreate.as_posix()}/test.db' + storage = SQLiteFeatureStorage( + uri=uri, single_output=True, upsert='ignore') + assert tocreate.exists() + def test_store_metadata(): """Test store_metadata""" diff --git a/junifer/testing/datagrabbers.py b/junifer/testing/datagrabbers.py index e39511163..db8003a6f 100644 --- a/junifer/testing/datagrabbers.py +++ b/junifer/testing/datagrabbers.py @@ -19,10 +19,11 @@ class OasisVBMTestingDatagrabber(BaseDataGrabber): return [f'sub-{x:02d}' for x in list(range(1, 11))] def __getitem__(self, element): - out = {} + out = super().__getitem__(element) i_sub = int(element.split('-')[1]) - 1 out['VBM_GM'] = self._dataset.gray_matter_maps[i_sub] - out['meta'] = {'element': {'subject': element}} + # Set the element accordingly + out['meta']['element'] = {'subject': element} return out def __enter__(self): -- 2.52.0 From 98e26780b6069af4ef4c7793241e61f961d7a1fd Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 25 Mar 2022 12:04:26 +0100 Subject: [PATCH 050/287] Add one more TODO --- setup.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/setup.py b/setup.py index 1c3f673e2..6c76b7b05 100644 --- a/setup.py +++ b/setup.py @@ -25,6 +25,8 @@ def _getversion(): DOWNLOAD_URL = 'https://github.com/juaml/junifer' URL = 'https://juaml.github.io/junifer' +# TODO: Read requirementes from requirements.txt and use them + setuptools.setup( name='junifer', author='Fede Raimondo', -- 2.52.0 From 5746d9ed50d6185b42617acc4e95cc563ac07402 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 25 Mar 2022 12:08:58 +0100 Subject: [PATCH 051/287] Fixed test, ready for demo in Software meeting --- junifer/api/tests/test_parser.py | 3 ++- junifer/datagrabber/tests/test_base_datagrabber.py | 7 +++++-- 2 files changed, 7 insertions(+), 3 deletions(-) diff --git a/junifer/api/tests/test_parser.py b/junifer/api/tests/test_parser.py index a5dcee7a0..392a90d7f 100644 --- a/junifer/api/tests/test_parser.py +++ b/junifer/api/tests/test_parser.py @@ -27,7 +27,8 @@ def test_parse_yaml(): contents = parse_yaml(fname) assert 'foo' in contents assert 'bar' == contents['foo'] - assert 'with' not in contents + assert 'with' in contents + assert 'numpy' in contents['with'] assert 'junifer.configs.wrong_config' not in sys.modules diff --git a/junifer/datagrabber/tests/test_base_datagrabber.py b/junifer/datagrabber/tests/test_base_datagrabber.py index a31fa8677..29924744e 100644 --- a/junifer/datagrabber/tests/test_base_datagrabber.py +++ b/junifer/datagrabber/tests/test_base_datagrabber.py @@ -29,8 +29,11 @@ def test_BaseDataGrabber(): return super().get_elements() dg = MyDataGrabber(datadir='/tmp', types=['func']) - with pytest.raises(NotImplementedError): - dg['elem'] + elem = dg['elem'] + assert 'meta' in elem + assert 'datagrabber' in elem['meta'] + assert 'class' in elem['meta']['datagrabber'] + assert MyDataGrabber.__name__ in elem['meta']['datagrabber']['class'] with pytest.raises(NotImplementedError): dg.get_elements() -- 2.52.0 From 6977b0b780cfcaca93356f9deca95103b46279f4 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 25 Mar 2022 14:32:42 +0100 Subject: [PATCH 052/287] One more YAML example + typos + llinter --- examples/yamls/ukb_gmd_mean.yaml | 25 +++++++++++++++++++++++++ junifer/markers/tests/test_parcel.py | 3 ++- junifer/storage/sqlite.py | 2 +- 3 files changed, 28 insertions(+), 2 deletions(-) create mode 100644 examples/yamls/ukb_gmd_mean.yaml diff --git a/examples/yamls/ukb_gmd_mean.yaml b/examples/yamls/ukb_gmd_mean.yaml new file mode 100644 index 000000000..7d2303d7e --- /dev/null +++ b/examples/yamls/ukb_gmd_mean.yaml @@ -0,0 +1,25 @@ +with: junifer.configs.juseless +workdir: /tmp + +datagrabber: + kind: JuselessUKBVBM +elements: +markers: + - name: Schaefer1000x7_TrimMean80 + kind: ParcelAggregation + atlas: Schaefer1000x7 + method: trim_mean + method_params: + proportiontocut: 0.2 + - name: Schaefer1000x7_Mean + kind: ParcelAggregation + atlas: Schaefer1000x7 + method: mean + - name: Schaefer1000x7_Std + kind: ParcelAggregation + atlas: Schaefer1000x7 + method: std +storage: + kind: SQLiteFeatureStorage + uri: /data/project/ukb_motor/junifer_test/test.db + diff --git a/junifer/markers/tests/test_parcel.py b/junifer/markers/tests/test_parcel.py index 7ab06d591..fae7313ab 100644 --- a/junifer/markers/tests/test_parcel.py +++ b/junifer/markers/tests/test_parcel.py @@ -97,7 +97,8 @@ def test_ParcelAggregation_3D(): manual = [] for t_v in sorted(np.unique(atlas_values)): t_values = trim_mean( - data[:, atlas_values == t_v], proportiontocut=0.1, axis=None) + data[:, atlas_values == t_v], proportiontocut=0.1, + axis=None) # type: ignore manual.append(t_values) manual = np.array(manual)[np.newaxis, :] diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py index d063d194c..82d5b0590 100644 --- a/junifer/storage/sqlite.py +++ b/junifer/storage/sqlite.py @@ -218,7 +218,7 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): Name of the table to save if_exist : str If the table exists, the behavior is controlled by this parameter. - Options are 'append' (default) and 'fail'. If 'fail' and th table + Options are 'append' (default) and 'fail'. If 'fail' and the table exists, it will raise an error. If 'append', the data will be appended to the existing table (following the upsert mode). -- 2.52.0 From f93240da379c4ba19cd011bbeb3c5375e9a2f65e Mon Sep 17 00:00:00 2001 From: LeSasse Date: Mon, 28 Mar 2022 17:26:31 +0200 Subject: [PATCH 053/287] initial not-quite working version of HCP datagrabber --- junifer/datagrabber/hcp.py | 112 +++++++++++++++++++++++++++++++++++++ 1 file changed, 112 insertions(+) create mode 100644 junifer/datagrabber/hcp.py diff --git a/junifer/datagrabber/hcp.py b/junifer/datagrabber/hcp.py new file mode 100644 index 000000000..0c2a8e940 --- /dev/null +++ b/junifer/datagrabber/hcp.py @@ -0,0 +1,112 @@ +import os +from itertools import product +from ..datagrabber import DataladDataGrabber +from ..api.decorators import register_datagrabber + + +@register_datagrabber +class HCP1200(DataladDataGrabber): + """ Human Connectome Project Datalad DataGrabber class + + Implements a DataGrabber to access the Human Connectome Project + + """ + def __init__(self, datadir=None): + """Initialize a HCP object. + + Parameters + ---------- + datadir : str or Path + That directory where the datalad dataset will be cloned. If None, + (default), the datalad dataset will be cloned into a temporary + directory. + + """ + uri = ( + 'https://github.com/datalad-datasets/' + 'human-connectome-project-openaccess.git' + ) + rootdir = 'HCP1200' + types = ['BOLD', 'T1w'] + super().__init__( + types=types, datadir=datadir, uri=uri, rootdir=rootdir + ) + + def get_elements(self, subjects=None, tasks=None, phase_encodings=None): + """Get the list of subjects in the dataset. + + Returns + ------- + elements : list[str] + The list of subjects in the dataset. + """ + elems = [] + + if isinstance(subjects, str): + subjects = [subjects] + if isinstance(tasks, str): + tasks = [tasks] + if isinstance(phase_encodings, str): + phase_encodings = [phase_encodings] + + if subjects is None: + subjects = os.listdir(self.datadir) + if tasks is None: + tasks = [ + "REST1", + "REST2", + "SOCIAL", + "WM", + "RELATIONAL", + "EMOTION", + "LANGUAGE", + "GAMBLING", + "MOTOR", + ] + if phase_encodings is None: + phase_encodings = ["LR", "RL"] + + for subject, task, phase_encoding in product( + subjects, tasks, phase_encodings + ): + elems.append((subject, task, phase_encoding)) + + return elems + + def __getitem__(self, element): + """Index one element in the dataset. + + Parameters + ---------- + element : tuple[str, str] + The element to be indexed. First element in the tuple is the + subject, second element is the task, third element is the + phase encoding direction. + + Returns + ------- + out : dict[str -> Path] + Dictionary of paths for each type of data required for the + specified element. + """ + sub, task, phase_encoding = element + out = {} + + if "REST" in task: + task_name = f"rfMRI_{task}" + else: + task_name = f"tfMRI_{task}" + + out["BOLD"] = dict( + path=self.datadir / sub / "MNINonLinear" / "Results" / + f"{task_name}_{phase_encoding}" / + f"{task_name}_{phase_encoding}_hp2000_clean.nii.gz" + ) + + self._dataset_get(out) + + out['meta']['element'] = dict( + subject=sub, task=task, phase_encoding=phase_encoding + ) + + return out \ No newline at end of file -- 2.52.0 From 43b290ee86d2604a82dbba12bcadf3ca2be65ce7 Mon Sep 17 00:00:00 2001 From: LeSasse Date: Tue, 29 Mar 2022 09:16:25 +0200 Subject: [PATCH 054/287] moved HCP1200 datagrabber to configs/juseless, because confound files are so far only available on juseless. Updated meta dictionary for JuselessUKBVBM and HCP1200 and updated constructor for HCP1200 --- junifer/configs/juseless.py | 142 ++++++++++++++++++++++++++++++++++++ junifer/datagrabber/hcp.py | 112 ---------------------------- 2 files changed, 142 insertions(+), 112 deletions(-) delete mode 100644 junifer/datagrabber/hcp.py diff --git a/junifer/configs/juseless.py b/junifer/configs/juseless.py index 7bebc030b..ab4f88d17 100644 --- a/junifer/configs/juseless.py +++ b/junifer/configs/juseless.py @@ -1,3 +1,6 @@ +import os +from itertools import product +from pathlib import Path from ..datagrabber import DataladDataGrabber from ..api.decorators import register_datagrabber @@ -61,6 +64,145 @@ class JuselessUKBVBM(DataladDataGrabber): out = {} out['VBM_GM'] = self.datadir / f'm0wp1{sub}_{ses}_T1w.nii.gz' + out['meta'] = dict(datagrabber=self.get_meta()) self._dataset_get(out) out['meta']['element'] = dict(subject=sub, session=ses) return out + + +@register_datagrabber +class HCP1200(DataladDataGrabber): + """ Human Connectome Project Datalad DataGrabber class + + Implements a DataGrabber to access the Human Connectome Project + + """ + + def __init__( + self, datadir=None, subjects=None, tasks=None, phase_encodings=None + ): + """Initialize a HCP object. + + Parameters + ---------- + datadir : str or Path + That directory where the datalad dataset will be cloned. If None, + (default), the datalad dataset will be cloned into a temporary + directory. + subjects : str or list of strings + HCP subject ID's. If 'None' (default), all available subjects are + selected + tasks : str or list of strings + HCP task sessions. If 'None' (default), all available task + sessions are selected. Can be 'REST1', 'REST2', 'SOCIAL', 'WM', + 'RELATIONAL', 'EMOTION', 'LANGUAGE', 'GAMBLING', 'MOTOR', or a + list consisting of these names. + phase_encoding : str or list of strings + HCP phase encoding directions. Can be 'LR' or 'RL'. If 'None' + (default) both will be used. + + """ + uri = ( + 'https://github.com/datalad-datasets/' + 'human-connectome-project-openaccess.git' + ) + rootdir = 'HCP1200' + types = ['BOLD'] + super().__init__( + types=types, datadir=datadir, uri=uri, rootdir=rootdir + ) + + self.subjects = subjects + self.tasks = tasks + self.phase_encodings = phase_encodings + + if isinstance(self.subjects, str): + self.subjects = [self.subjects] + if isinstance(self.tasks, str): + self.tasks = [self.tasks] + if isinstance(self.phase_encodings, str): + self.phase_encodings = [self.phase_encodings] + + if self.tasks is None: + self.tasks = [ + 'REST1', + 'REST2', + 'SOCIAL', + 'WM', + 'RELATIONAL', + 'EMOTION', + 'LANGUAGE', + 'GAMBLING', + 'MOTOR', + ] + + if self.phase_encodings is None: + self.phase_encodings = ["LR", "RL"] + + def get_elements(self): + """Get the list of subjects in the dataset. + + Returns + ------- + elements : list[str] + The list of subjects in the dataset. + """ + elems = [] + + if self.subjects is None: + self.subjects = os.listdir(self.datadir) + + for subject, task, phase_encoding in product( + self.subjects, self.tasks, self.phase_encodings + ): + elems.append((subject, task, phase_encoding)) + + return elems + + def __getitem__(self, element): + """Index one element in the dataset. + + Parameters + ---------- + element : tuple[str, str] + The element to be indexed. First element in the tuple is the + subject, second element is the task, third element is the + phase encoding direction. + + Returns + ------- + out : dict[str -> Path] + Dictionary of paths for each type of data required for the + specified element. + """ + sub, task, phase_encoding = element + out = {} + + if 'REST' in task: + task_name = f'rfMRI_{task}' + else: + task_name = f'tfMRI_{task}' + + out['BOLD'] = dict( + path=self.datadir / sub / 'MNINonLinear' / 'Results' / + f'{task_name}_{phase_encoding}' / + f'{task_name}_{phase_encoding}_hp2000_clean.nii.gz' + ) + + conf_dir = ( + Path('/data') / 'group' / 'appliedml' / + 'data' / 'HCP1200_Confounds_tsv' + ) + out['BOLD']['confounds'] = ( + conf_dir / sub / 'MNINonLinear' / 'Results' / + f'{task_name}_{phase_encoding}' / f'Confounds_{sub}.tsv' + ) + + out['meta'] = dict(datagrabber=self.get_meta()) + self._dataset_get(out) + + out['meta']['element'] = dict( + subject=sub, task=task, phase_encoding=phase_encoding + ) + + return out diff --git a/junifer/datagrabber/hcp.py b/junifer/datagrabber/hcp.py deleted file mode 100644 index 0c2a8e940..000000000 --- a/junifer/datagrabber/hcp.py +++ /dev/null @@ -1,112 +0,0 @@ -import os -from itertools import product -from ..datagrabber import DataladDataGrabber -from ..api.decorators import register_datagrabber - - -@register_datagrabber -class HCP1200(DataladDataGrabber): - """ Human Connectome Project Datalad DataGrabber class - - Implements a DataGrabber to access the Human Connectome Project - - """ - def __init__(self, datadir=None): - """Initialize a HCP object. - - Parameters - ---------- - datadir : str or Path - That directory where the datalad dataset will be cloned. If None, - (default), the datalad dataset will be cloned into a temporary - directory. - - """ - uri = ( - 'https://github.com/datalad-datasets/' - 'human-connectome-project-openaccess.git' - ) - rootdir = 'HCP1200' - types = ['BOLD', 'T1w'] - super().__init__( - types=types, datadir=datadir, uri=uri, rootdir=rootdir - ) - - def get_elements(self, subjects=None, tasks=None, phase_encodings=None): - """Get the list of subjects in the dataset. - - Returns - ------- - elements : list[str] - The list of subjects in the dataset. - """ - elems = [] - - if isinstance(subjects, str): - subjects = [subjects] - if isinstance(tasks, str): - tasks = [tasks] - if isinstance(phase_encodings, str): - phase_encodings = [phase_encodings] - - if subjects is None: - subjects = os.listdir(self.datadir) - if tasks is None: - tasks = [ - "REST1", - "REST2", - "SOCIAL", - "WM", - "RELATIONAL", - "EMOTION", - "LANGUAGE", - "GAMBLING", - "MOTOR", - ] - if phase_encodings is None: - phase_encodings = ["LR", "RL"] - - for subject, task, phase_encoding in product( - subjects, tasks, phase_encodings - ): - elems.append((subject, task, phase_encoding)) - - return elems - - def __getitem__(self, element): - """Index one element in the dataset. - - Parameters - ---------- - element : tuple[str, str] - The element to be indexed. First element in the tuple is the - subject, second element is the task, third element is the - phase encoding direction. - - Returns - ------- - out : dict[str -> Path] - Dictionary of paths for each type of data required for the - specified element. - """ - sub, task, phase_encoding = element - out = {} - - if "REST" in task: - task_name = f"rfMRI_{task}" - else: - task_name = f"tfMRI_{task}" - - out["BOLD"] = dict( - path=self.datadir / sub / "MNINonLinear" / "Results" / - f"{task_name}_{phase_encoding}" / - f"{task_name}_{phase_encoding}_hp2000_clean.nii.gz" - ) - - self._dataset_get(out) - - out['meta']['element'] = dict( - subject=sub, task=task, phase_encoding=phase_encoding - ) - - return out \ No newline at end of file -- 2.52.0 From 7622c2c6786987371577586fbd364e25fa74b48a Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Tue, 29 Mar 2022 12:08:38 +0200 Subject: [PATCH 055/287] HCP Datagrabber WIP --- .gitignore | 3 +- junifer/configs/juseless.py | 268 ++++++++---------- junifer/configs/tests/test_juseless.py | 9 +- junifer/datagrabber/__init__.py | 3 +- junifer/datagrabber/base.py | 120 +++++--- junifer/datagrabber/hcp.py | 102 +++++++ .../tests/test_base_datagrabber.py | 108 ++++--- junifer/utils/logging.py | 3 +- pyproject.toml | 6 +- 9 files changed, 401 insertions(+), 221 deletions(-) create mode 100644 junifer/datagrabber/hcp.py diff --git a/.gitignore b/.gitignore index 35e6aacdc..94ac529c9 100644 --- a/.gitignore +++ b/.gitignore @@ -133,4 +133,5 @@ cython_debug/ .DS_store junifer/_version.py -scratch/ \ No newline at end of file +scratch/ +junifer_jobs/ \ No newline at end of file diff --git a/junifer/configs/juseless.py b/junifer/configs/juseless.py index ab4f88d17..1918c5d82 100644 --- a/junifer/configs/juseless.py +++ b/junifer/configs/juseless.py @@ -1,13 +1,11 @@ -import os -from itertools import product -from pathlib import Path -from ..datagrabber import DataladDataGrabber +from ..datagrabber import PatternDataladDataGrabber from ..api.decorators import register_datagrabber +from ..utils import logger @register_datagrabber -class JuselessUKBVBM(DataladDataGrabber): - """Juseless UKB VMG DataGrabber class. +class JuselessDataladUKBVBM(PatternDataladDataGrabber): + """Juseless UKB VBM DataGrabber class. Implements a DataGrabber to access the UKB VBM data in Juseless. @@ -26,8 +24,13 @@ class JuselessUKBVBM(DataladDataGrabber): uri = 'ria+http://ukb.ds.inm7.de#~cat_m0wp1' rootdir = 'm0wp1' types = ['VBM_GM'] + replacements = ['subject', 'session'] + patterns = { + 'VBM_GM': 'm0wp1{subject}_{session}_T1w.nii.gz' + } super().__init__( - types=types, datadir=datadir, uri=uri, rootdir=rootdir) + types=types, datadir=datadir, uri=uri, rootdir=rootdir, + replacements=replacements, patterns=patterns) def get_elements(self): """Get the list of subjects in the dataset. @@ -37,6 +40,7 @@ class JuselessUKBVBM(DataladDataGrabber): elements : list[str] The list of subjects in the dataset. """ + logger.debug('Getting the list of subjects in the dataset') elems = [] for x in self.datadir.glob('*._T1w.nii.gz'): sub, ses = x.name.split('_') @@ -45,164 +49,140 @@ class JuselessUKBVBM(DataladDataGrabber): elems.append((sub, ses)) return elems - def __getitem__(self, element): - """Index one element in the dataset. - Parameters - ---------- - element : tuple[str, str] - The element to be indexed. First element in the tuple is the - subject, second element is the session. +# @register_datagrabber +# class HCP1200(PatternDataGrabber): +# """ Human Connectome Project Datalad DataGrabber class - Returns - ------- - out : dict[str -> Path] - Dictionary of paths for each type of data required for the - specified element. - """ - sub, ses = element - out = {} +# Implements a DataGrabber to access the Human Connectome Project - out['VBM_GM'] = self.datadir / f'm0wp1{sub}_{ses}_T1w.nii.gz' - out['meta'] = dict(datagrabber=self.get_meta()) - self._dataset_get(out) - out['meta']['element'] = dict(subject=sub, session=ses) - return out +# """ +# def __init__( +# self, datadir=None, subjects=None, tasks=None, phase_encodings=None +# ): +# """Initialize a HCP object. -@register_datagrabber -class HCP1200(DataladDataGrabber): - """ Human Connectome Project Datalad DataGrabber class +# Parameters +# ---------- +# datadir : str or Path +# That directory where the datalad dataset will be cloned. If None, +# (default), the datalad dataset will be cloned into a temporary +# directory. +# subjects : str or list of strings +# HCP subject ID's. If 'None' (default), all available subjects are +# selected +# tasks : str or list of strings +# HCP task sessions. If 'None' (default), all available task +# sessions are selected. Can be 'REST1', 'REST2', 'SOCIAL', 'WM', +# 'RELATIONAL', 'EMOTION', 'LANGUAGE', 'GAMBLING', 'MOTOR', or a +# list consisting of these names. +# phase_encoding : str or list of strings +# HCP phase encoding directions. Can be 'LR' or 'RL'. If 'None' +# (default) both will be used. - Implements a DataGrabber to access the Human Connectome Project +# """ +# uri = ( +# 'https://github.com/datalad-datasets/' +# 'human-connectome-project-openaccess.git' +# ) +# rootdir = 'HCP1200' +# types = ['BOLD'] +# super().__init__( +# types=types, datadir=datadir, uri=uri, rootdir=rootdir +# ) - """ +# self.subjects = subjects +# self.tasks = tasks +# self.phase_encodings = phase_encodings - def __init__( - self, datadir=None, subjects=None, tasks=None, phase_encodings=None - ): - """Initialize a HCP object. +# if isinstance(self.subjects, str): +# self.subjects = [self.subjects] +# if isinstance(self.tasks, str): +# self.tasks = [self.tasks] +# if isinstance(self.phase_encodings, str): +# self.phase_encodings = [self.phase_encodings] - Parameters - ---------- - datadir : str or Path - That directory where the datalad dataset will be cloned. If None, - (default), the datalad dataset will be cloned into a temporary - directory. - subjects : str or list of strings - HCP subject ID's. If 'None' (default), all available subjects are - selected - tasks : str or list of strings - HCP task sessions. If 'None' (default), all available task - sessions are selected. Can be 'REST1', 'REST2', 'SOCIAL', 'WM', - 'RELATIONAL', 'EMOTION', 'LANGUAGE', 'GAMBLING', 'MOTOR', or a - list consisting of these names. - phase_encoding : str or list of strings - HCP phase encoding directions. Can be 'LR' or 'RL'. If 'None' - (default) both will be used. +# if self.tasks is None: +# self.tasks = [ +# 'REST1', +# 'REST2', +# 'SOCIAL', +# 'WM', +# 'RELATIONAL', +# 'EMOTION', +# 'LANGUAGE', +# 'GAMBLING', +# 'MOTOR', +# ] - """ - uri = ( - 'https://github.com/datalad-datasets/' - 'human-connectome-project-openaccess.git' - ) - rootdir = 'HCP1200' - types = ['BOLD'] - super().__init__( - types=types, datadir=datadir, uri=uri, rootdir=rootdir - ) +# if self.phase_encodings is None: +# self.phase_encodings = ["LR", "RL"] - self.subjects = subjects - self.tasks = tasks - self.phase_encodings = phase_encodings +# def get_elements(self): +# """Get the list of subjects in the dataset. - if isinstance(self.subjects, str): - self.subjects = [self.subjects] - if isinstance(self.tasks, str): - self.tasks = [self.tasks] - if isinstance(self.phase_encodings, str): - self.phase_encodings = [self.phase_encodings] +# Returns +# ------- +# elements : list[str] +# The list of subjects in the dataset. +# """ +# elems = [] - if self.tasks is None: - self.tasks = [ - 'REST1', - 'REST2', - 'SOCIAL', - 'WM', - 'RELATIONAL', - 'EMOTION', - 'LANGUAGE', - 'GAMBLING', - 'MOTOR', - ] +# if self.subjects is None: +# self.subjects = os.listdir(self.datadir) - if self.phase_encodings is None: - self.phase_encodings = ["LR", "RL"] +# for subject, task, phase_encoding in product( +# self.subjects, self.tasks, self.phase_encodings +# ): +# elems.append((subject, task, phase_encoding)) - def get_elements(self): - """Get the list of subjects in the dataset. +# return elems - Returns - ------- - elements : list[str] - The list of subjects in the dataset. - """ - elems = [] +# def __getitem__(self, element): +# """Index one element in the dataset. - if self.subjects is None: - self.subjects = os.listdir(self.datadir) +# Parameters +# ---------- +# element : tuple[str, str] +# The element to be indexed. First element in the tuple is the +# subject, second element is the task, third element is the +# phase encoding direction. - for subject, task, phase_encoding in product( - self.subjects, self.tasks, self.phase_encodings - ): - elems.append((subject, task, phase_encoding)) +# Returns +# ------- +# out : dict[str -> Path] +# Dictionary of paths for each type of data required for the +# specified element. +# """ +# sub, task, phase_encoding = element +# out = {} - return elems +# if 'REST' in task: +# task_name = f'rfMRI_{task}' +# else: +# task_name = f'tfMRI_{task}' - def __getitem__(self, element): - """Index one element in the dataset. +# out['BOLD'] = dict( +# path=self.datadir / sub / 'MNINonLinear' / 'Results' / +# f'{task_name}_{phase_encoding}' / +# f'{task_name}_{phase_encoding}_hp2000_clean.nii.gz' +# ) - Parameters - ---------- - element : tuple[str, str] - The element to be indexed. First element in the tuple is the - subject, second element is the task, third element is the - phase encoding direction. +# conf_dir = ( +# Path('/data') / 'group' / 'appliedml' / +# 'data' / 'HCP1200_Confounds_tsv' +# ) +# out['BOLD']['confounds'] = ( +# conf_dir / sub / 'MNINonLinear' / 'Results' / +# f'{task_name}_{phase_encoding}' / f'Confounds_{sub}.tsv' +# ) - Returns - ------- - out : dict[str -> Path] - Dictionary of paths for each type of data required for the - specified element. - """ - sub, task, phase_encoding = element - out = {} +# out['meta'] = dict(datagrabber=self.get_meta()) +# self._dataset_get(out) - if 'REST' in task: - task_name = f'rfMRI_{task}' - else: - task_name = f'tfMRI_{task}' +# out['meta']['element'] = dict( +# subject=sub, task=task, phase_encoding=phase_encoding +# ) - out['BOLD'] = dict( - path=self.datadir / sub / 'MNINonLinear' / 'Results' / - f'{task_name}_{phase_encoding}' / - f'{task_name}_{phase_encoding}_hp2000_clean.nii.gz' - ) - - conf_dir = ( - Path('/data') / 'group' / 'appliedml' / - 'data' / 'HCP1200_Confounds_tsv' - ) - out['BOLD']['confounds'] = ( - conf_dir / sub / 'MNINonLinear' / 'Results' / - f'{task_name}_{phase_encoding}' / f'Confounds_{sub}.tsv' - ) - - out['meta'] = dict(datagrabber=self.get_meta()) - self._dataset_get(out) - - out['meta']['element'] = dict( - subject=sub, task=task, phase_encoding=phase_encoding - ) - - return out +# return out diff --git a/junifer/configs/tests/test_juseless.py b/junifer/configs/tests/test_juseless.py index 7c5131327..03a6bb8fc 100644 --- a/junifer/configs/tests/test_juseless.py +++ b/junifer/configs/tests/test_juseless.py @@ -1,14 +1,17 @@ import socket import pytest -from junifer.configs.juseless import JuselessUKBVBM +from junifer.configs.juseless import JuselessDataladUKBVBM +from junifer.utils.logging import configure_logging if socket.gethostname() != 'juseless': pytest.skip('This tests are only for juseless', allow_module_level=True) +configure_logging(level='DEBUG') -def test_juselessukbvbm_datagrabber(): - with JuselessUKBVBM() as dg: + +def test_juselessdataladukbvbm_datagrabber(): + with JuselessDataladUKBVBM() as dg: out = dg[('sub-2670511', 'ses-2')] assert 'VBM_GM' in out assert out['VBM_GM'].name == 'm0wp1sub-2670511_ses-2_T1w.nii.gz' diff --git a/junifer/datagrabber/__init__.py b/junifer/datagrabber/__init__.py index e9866a1be..ce54e9f35 100644 --- a/junifer/datagrabber/__init__.py +++ b/junifer/datagrabber/__init__.py @@ -1,4 +1,5 @@ # Authors: Federico Raimondo # Leonard Sasse # License: AGPL -from .base import BIDSDataladDataGrabber, DataladDataGrabber, BIDSDataGrabber \ No newline at end of file +from .base import (DataladDataGrabber, PatternDataGrabber, + PatternDataladDataGrabber) \ No newline at end of file diff --git a/junifer/datagrabber/base.py b/junifer/datagrabber/base.py index 732df21bf..d99ec3821 100644 --- a/junifer/datagrabber/base.py +++ b/junifer/datagrabber/base.py @@ -8,7 +8,7 @@ import datalad.api as dl from abc import ABC, abstractmethod from ..api.decorators import register_datagrabber -from ..utils.logging import logger, raise_error +from ..utils.logging import logger, raise_error, warn def _validate_types(types): @@ -22,6 +22,22 @@ def _validate_types(types): "types must be a list of strings", TypeError) # type: ignore +def _validate_replacements(replacements, patterns): + """ + Validate the replacements + """ + if not isinstance(replacements, list): + raise_error("replacements must be a list", TypeError) # type: ignore + if any(not isinstance(x, str) for x in replacements): + raise_error( + "replacements must be a list of strings", + TypeError) # type: ignore + + for x in replacements: + if all(f'{{x}}' not in y for y in patterns.values()): + warn(f"Replacement {x} is not part of any pattern") + + def _validate_patterns(types, patterns): """ Validate the patterns. @@ -73,6 +89,9 @@ class BaseDataGrabber(ABC): _validate_types(types) if not isinstance(datadir, Path): datadir = Path(datadir) + logger.debug('Initializing BaseDataGrabber') + logger.debug(f'\t_datadir = {datadir}') + logger.debug(f'\ttypes = {types}') self._datadir = datadir self.types = types @@ -112,13 +131,16 @@ class BaseDataGrabber(ABC): yield elem def __getitem__(self, element): + logger.info(f'Getting element {element}') out = {} out['meta'] = dict(datagrabber=self.get_meta()) return out @abstractmethod def get_elements(self): - raise NotImplementedError('get_elements not implemented') + raise_error( + 'get_elements not implemented', + NotImplementedError) # type: ignore def __enter__(self): return self @@ -128,34 +150,37 @@ class BaseDataGrabber(ABC): @register_datagrabber -class BIDSDataGrabber(BaseDataGrabber): - """BIDS DataGrabber class. Implements a DataGrabber that understands BIDS - database format. +class PatternDataGrabber(BaseDataGrabber): + """Patternd DataGrabber class (abstract). Implements a DataGrabber that + understands patterns to grab data. Attributes ---------- - datadir - types : list + datadir: Path + Directory where the data is stored + types : list[str] List of data types to be grabbed. patterns : dict[str -> str] Patterns for each type of data. + replacements: list[str] + Replacements in the patterns for each item in the `element` tuple. Methods ------- get_elements: list[str] Returns a list of elements that can be grabbed. Each element is a subject in the BIDS database. - __getitem__(str): dict[str -> Path] + __getitem__(str): dict[str -> dict] Returns a dictionary of paths for each type of data required for the specified element. Each occurrence of the string `{subject}` is replaced by the indexed element """ - def __init__(self, types=None, patterns=None, **kwargs): + def __init__(self, types=None, patterns=None, replacements=None, **kwargs): """Initialize a BaseDataGrabber object. Parameters ---------- - types : list of str + types : list[str] The types of data to be grabbed. patterns : dict[str -> str] Patterns for each type of data. The keys are the types and the @@ -165,32 +190,52 @@ class BIDSDataGrabber(BaseDataGrabber): That directory where the data is/will be stored. """ _validate_patterns(types, patterns) + if not isinstance(replacements, list): + replacements = [replacements] + _validate_replacements(replacements, patterns) super().__init__(types=types, **kwargs) + logger.debug('Initializing PatternDataGrabber') + logger.debug(f'\tpatterns = {patterns}') + logger.debug(f'\treplacements = {replacements}') self.patterns = patterns + self.replacements = replacements - def get_elements(self): - """Get all the elements in the BIDS database + def _replace_patterns(self, element, pattern): + """Replace the patterns in the pattern with the element. + + Parameters + ---------- + element : tuple + The element to be used in the replacement. + pattern : str + The pattern to be replaced. Returns ------- - elems : list[str] - List of all the elements in the database root directory + str + The pattern with the element replaced. """ - elems = [x.name for x in self.datadir.iterdir() if x.is_dir()] - return elems + if len(element) != len(self.replacements): + raise_error( + f'The element lenght must be {len(self.replacements)}, ' + f'indicating {self.replacements}') + to_replace = dict(zip(self.replacements, element)) + return pattern.format(**to_replace) def __getitem__(self, element): - """Index one element in the BIDS database. + """Index one element in the database. - Each occurrence of the string `{subject}` is replaced by the indexed - element. + Each occurrence of the strings in `replacements` is replaced by the + corresponding item in the element tuple. Parameters ---------- element : str or tuple The element to be indexed. If one string is provided, it is - assumed to be a subject. If a tuple is provided, it is assumed to - be a (subject, session) pair. + assumed to be a tuple with only one item. If a tuple is provided, + each item in the tuple is the value for the replacement string + specified in `replacements`. + Returns ------- out : dict[str -> Path] @@ -202,24 +247,19 @@ class BIDSDataGrabber(BaseDataGrabber): element = (element,) for t_type in self.types: t_pattern = self.patterns[t_type] # type: ignore - t_replace = t_pattern.replace('{subject}', element[0]) - if len(element) > 1: - t_replace = t_replace.replace( - '{session}', element[1]) # type: ignore + t_replace = self._replace_patterns(element, t_pattern) t_out = self.datadir / element[0] / t_replace out[t_type] = dict(path=t_out) # Meta here is element and types - out['meta']['element'] = {'subject': element[0]} - if len(element) > 1: - out['meta']['element']['session'] = element[1] # type: ignore + out['meta']['element'] = dict(zip(self.replacements, element)) return out @register_datagrabber class DataladDataGrabber(BaseDataGrabber): """ - Datalad DataGrabber class. Implements a DataGrabber that gets data from - a datalad sibling. + Datalad DataGrabber class (abstract). Implements a DataGrabber that gets + data from a datalad sibling. Attributes ---------- @@ -267,6 +307,9 @@ class DataladDataGrabber(BaseDataGrabber): datadir = tempfile.mkdtemp() logger.info(f'datadir set to {datadir}') super().__init__(datadir=datadir, **kwargs) + logger.debug('Initializing DataladDataGrabber') + logger.debug(f'\turi = {uri}') + logger.debug(f'\t_rootdir = {rootdir}') self.uri = uri self._rootdir = rootdir @@ -280,11 +323,15 @@ class DataladDataGrabber(BaseDataGrabber): def install(self): """Install the datalad dataset into the datadir.""" + logger.debug(f'Installing dataset {self.uri} to {self._datadir}') self.dataset = dl.install( # type: ignore self._datadir, source=self.uri) + logger.debug('Dataset installed') def __exit__(self, exc_type, exc_value, exc_traceback): + logger.debug('Removing dataset') self.remove() + logger.debug('Dataset removed') def remove(self): """Remove the datalad dataset from the datadir.""" @@ -293,7 +340,9 @@ class DataladDataGrabber(BaseDataGrabber): def _dataset_get(self, out): for _, v in out.items(): if 'path' in v: + logger.debug(f'Getting {v["path"]}') self.dataset.get(v['path']) + logger.debug(f'Get done') # append the version of the dataset out['meta']['datagrabber']['dataset_commit_id'] = \ @@ -314,15 +363,16 @@ class DataladDataGrabber(BaseDataGrabber): return out -class BIDSDataladDataGrabber(DataladDataGrabber, BIDSDataGrabber): - """BIDS Datalad DataGrabber class. - Implements a DataGrabber that gets data from a datalad sibling which - follows a BIDS format. +class PatternDataladDataGrabber(DataladDataGrabber, PatternDataGrabber): + """Pattern-based Datalad DataGrabber class (abstract). + Implements a DataGrabber that gets data from a datalad sibling, + interpreting patterns. + See Also -------- DataladDataGrabber - BIDSDataGrabber + PatternDataGrabber """ def __init__(self, types=None, patterns=None, **kwargs): diff --git a/junifer/datagrabber/hcp.py b/junifer/datagrabber/hcp.py new file mode 100644 index 000000000..a063fbd16 --- /dev/null +++ b/junifer/datagrabber/hcp.py @@ -0,0 +1,102 @@ +from itertools import product + +from junifer.datagrabber.base import DataladDataGrabber + +from ..datagrabber import PatternDataGrabber +from ..api.decorators import register_datagrabber + + +@register_datagrabber +class HCP1200(PatternDataGrabber): + + def __init__( + self, datadir=None, tasks=None, phase_encodings=None + ): + """Initialize a HCP object. + + Parameters + ---------- + datadir : str or Path + That directory where the datalad dataset will be cloned. If None, + (default), the datalad dataset will be cloned into a temporary + directory. + tasks : str or list of strings + HCP task sessions. If 'None' (default), all available task + sessions are selected. Can be 'REST1', 'REST2', 'SOCIAL', 'WM', + 'RELATIONAL', 'EMOTION', 'LANGUAGE', 'GAMBLING', 'MOTOR', or a + list consisting of these names. + phase_encoding : str or list of strings + HCP phase encoding directions. Can be 'LR' or 'RL'. If 'None' + (default) both will be used. + + """ + + types = ['BOLD'] + # TODO: Validate tasks + # TODO: Validate phase_encodings + + replacements = ['subject', 'task', 'phase_encoding'] + patterns = { + 'BOLD': ('{subject}/MNINonLinear/Results/' + '*fMRI_{task}_{phase_encoding}/' + '*fMRI_{task}_{phase_encoding}_hp2000_clean.nii.gz') + } + super().__init__( + types=types, datadir=datadir, patterns=patterns, + replacements=replacements + ) + + self.tasks = tasks + self.phase_encodings = phase_encodings + + if isinstance(self.tasks, str): + self.tasks = [self.tasks] + if isinstance(self.phase_encodings, str): + self.phase_encodings = [self.phase_encodings] + + if self.tasks is None: + self.tasks = [ + 'REST1', + 'REST2', + 'SOCIAL', + 'WM', + 'RELATIONAL', + 'EMOTION', + 'LANGUAGE', + 'GAMBLING', + 'MOTOR', + ] + + if self.phase_encodings is None: + self.phase_encodings = ["LR", "RL"] + + def get_elements(self): + """Get the list of subjects in the dataset. + + Returns + ------- + elements : list[str] + The list of subjects in the dataset. + """ + + subjects = [x.name for x in self.datadir.iterdir() if x.is_dir()] + elems = [] + for subject, task, phase_encoding in product( + subjects, self.tasks, self.phase_encodings + ): + elems.append((subject, task, phase_encoding)) + + return elems + + +@register_datagrabber +class DataladHCP1200(HCP1200, DataladDataGrabber): + def __init__(self, datadir=None, tasks=None, phase_encodings=None): + uri = ( + 'https://github.com/datalad-datasets/' + 'human-connectome-project-openaccess.git' + ) + rootdir = 'HCP1200' + super().__init__(datadir=datadir, tasks=tasks, + phase_encodings=phase_encodings, + uri=uri, rootdir=rootdir) # type: ignore diff --git a/junifer/datagrabber/tests/test_base_datagrabber.py b/junifer/datagrabber/tests/test_base_datagrabber.py index 29924744e..3b5f2f8f1 100644 --- a/junifer/datagrabber/tests/test_base_datagrabber.py +++ b/junifer/datagrabber/tests/test_base_datagrabber.py @@ -4,8 +4,8 @@ import tempfile import pytest from pathlib import Path -from junifer.datagrabber.base import BIDSDataGrabber, BIDSDataladDataGrabber, \ - BaseDataGrabber +from junifer.datagrabber.base import (PatternDataGrabber, BaseDataGrabber, + PatternDataladDataGrabber) _testing_dataset = { @@ -43,60 +43,98 @@ def test_BaseDataGrabber(): assert dg.types == ['func'] -def test_BIDSDataGrabber(): - """Test BIDSDataGrabber""" +def test_PatternDataGrabber(): + class MyDataGrabber(PatternDataGrabber): + def get_elements(self): + return super().get_elements() + + """Test test_PatternDataGrabber""" with pytest.raises(TypeError, match=r"types must be a list"): - BIDSDataGrabber(datadir='/tmp', types='wrong', - patterns=dict(wrong='pattern')) + MyDataGrabber(datadir='/tmp', types='wrong', + patterns=dict(wrong='pattern'), + replacements='subject') with pytest.raises(TypeError, match=r"must be a list of strings"): - BIDSDataGrabber(datadir='/tmp', types=[1, 2, 3], - patterns={'1': 'pattern', '2': 'pattern', - '3': 'pattern'}) + MyDataGrabber(datadir='/tmp', types=[1, 2, 3], + patterns={'1': 'pattern', '2': 'pattern', + '3': 'pattern'}, + replacements='subject') - datagrabber = BIDSDataGrabber( - datadir='/tmp/data', types=['func', 'anat'], - patterns=dict(func='pattern1', anat='pattern2')) - assert datagrabber.datadir == Path('/tmp/data') - assert datagrabber.types == ['func', 'anat'] - - datagrabber = BIDSDataGrabber( - datadir=Path('/tmp/data'), types=['func', 'anat'], - patterns=dict(func='pattern1', anat='pattern2')) - assert datagrabber.datadir == Path('/tmp/data') - assert datagrabber.types == ['func', 'anat'] + with pytest.raises(ValueError, match=r"must have the same length"): + MyDataGrabber(datadir='/tmp', types=['func', 'anat'], + patterns={'1': 'pattern', '2': 'pattern', + '3': 'pattern'}, + replacements=1) with pytest.raises(TypeError, match=r"patterns must be a dict"): - BIDSDataGrabber(datadir='/tmp', types=['func', 'anat'], - patterns='wrong') + MyDataGrabber(datadir='/tmp', types=['func', 'anat'], + patterns='wrong', replacements='subject') with pytest.raises(ValueError, match=r"patterns must have the same length"): - BIDSDataGrabber(datadir='/tmp', types=['func', 'anat'], - patterns={'wrong': 'pattern'}) + MyDataGrabber(datadir='/tmp', types=['func', 'anat'], + patterns={'wrong': 'pattern'}, replacements='subject') with pytest.raises(ValueError, match=r"patterns must contain all types"): - BIDSDataGrabber(datadir='/tmp', types=['func', 'anat'], - patterns={'wrong': 'pattern', 'func': 'pattern'}) + MyDataGrabber(datadir='/tmp', types=['func', 'anat'], + patterns={'wrong': 'pattern', 'func': 'pattern'}, + replacements='subject') + + with pytest.raises(TypeError, match=r"must be a list of strings"): + MyDataGrabber(datadir='/tmp', types=['func', 'anat'], + patterns={'func': 'func/test', 'anat': 'anat/test'}, + replacements=1) + + with pytest.warns(RuntimeWarning, match=r"not part of any pattern"): + MyDataGrabber(datadir='/tmp', types=['func', 'anat'], + patterns={'func': 'func/{subject}.nii', + 'anat': 'anat/{subject}.nii'}, + replacements=['subject', 'wrong']) + + datagrabber = MyDataGrabber( + datadir='/tmp/data', types=['func', 'anat'], + patterns={'func': 'func/{subject}.nii', + 'anat': 'anat/{subject}.nii'}, + replacements='subject') + assert datagrabber.datadir == Path('/tmp/data') + assert datagrabber.types == ['func', 'anat'] + assert datagrabber.replacements == ['subject'] + + datagrabber = MyDataGrabber( + datadir=Path('/tmp/data'), types=['func', 'anat'], + patterns={'func': 'func/{subject}.nii', + 'anat': 'anat/{subject}_{session}.nii'}, + replacements=['subject', 'session']) + assert datagrabber.datadir == Path('/tmp/data') + assert datagrabber.types == ['func', 'anat'] + assert datagrabber.replacements == ['subject', 'session'] -def test_BIDSDataladDataGrabber(): - """Test BIDSDataladDataGrabber""" +def test_bids_datalad_PatternDataGrabber(): + """Test a subject-based BIDS datalad datagrabber""" types = ['T1w', 'bold'] patterns = { 'T1w': 'anat/{subject}_T1w.nii.gz', 'bold': 'func/{subject}_task-rest_bold.nii.gz' } + replacements = ['subject'] + + class MyDataGrabber(PatternDataladDataGrabber): + def get_elements(self): + elems = [x.name for x in self.datadir.iterdir() if x.is_dir()] + return elems with pytest.raises(ValueError, match=r"uri must be provided"): - BIDSDataladDataGrabber(datadir=None, types=types, patterns=patterns) + MyDataGrabber(datadir=None, types=types, patterns=patterns, + replacements=replacements) repo_uri = _testing_dataset['example_bids']['uri'] rootdir = 'example_bids' repo_commit = _testing_dataset['example_bids']['id'] - with BIDSDataladDataGrabber(rootdir=rootdir, uri=repo_uri, - types=types, patterns=patterns) as dg: + with MyDataGrabber(rootdir=rootdir, uri=repo_uri, + types=types, patterns=patterns, + replacements=replacements) as dg: subs = [x for x in dg] expected_subs = [f'sub-{i:02d}' for i in range(1, 10)] assert set(subs) == set(expected_subs) @@ -114,7 +152,7 @@ def test_BIDSDataladDataGrabber(): assert 'datagrabber' in t_sub['meta'] dg_meta = t_sub['meta']['datagrabber'] assert 'class' in dg_meta - assert dg_meta['class'] == 'BIDSDataladDataGrabber' + assert dg_meta['class'] == 'MyDataGrabber' assert 'uri' in dg_meta assert dg_meta['uri'] == repo_uri assert 'dataset_commit_id' in dg_meta @@ -125,7 +163,7 @@ def test_BIDSDataladDataGrabber(): with tempfile.TemporaryDirectory() as tmpdir: datadir = Path(tmpdir) / 'dataset' # Need this for testing - with BIDSDataladDataGrabber(rootdir=rootdir, uri=repo_uri, - types=types, patterns=patterns, - datadir=datadir) as dg: + with MyDataGrabber(rootdir=rootdir, uri=repo_uri, + types=types, patterns=patterns, + datadir=datadir, replacements=replacements) as dg: assert dg.datadir == datadir / rootdir diff --git a/junifer/utils/logging.py b/junifer/utils/logging.py index a5558aad8..343ae4954 100644 --- a/junifer/utils/logging.py +++ b/junifer/utils/logging.py @@ -126,7 +126,8 @@ def configure_logging(level='WARNING', fname=None, overwrite=None, """ _close_handlers(logger) if output_format is None: - output_format = '%(asctime)s - %(name)s - %(levelname)s - %(message)s' + output_format = ('%(asctime)s [%(levelname)8s] %(message)s ' + '(%(filename)s:%(lineno)s)') formatter = logging.Formatter(output_format) if fname is not None: diff --git a/pyproject.toml b/pyproject.toml index 9cae268d7..22429039a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -6,4 +6,8 @@ build-backend = "setuptools.build_meta" version_scheme = "python-simplified-semver" local_scheme = "no-local-version" write_to = "junifer/_version.py" -write_to_template = "__version__ = '{version}'\n" \ No newline at end of file +write_to_template = "__version__ = '{version}'\n" + +[tool.pytest.ini_options] +log_cli = true +log_cli_level = "WARNING" \ No newline at end of file -- 2.52.0 From ce73351bb8a3185968be4ca9145c67e7eacc8689 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Tue, 29 Mar 2022 13:17:44 +0200 Subject: [PATCH 056/287] Fix pattern data grabber: add option to use '*' --- junifer/datagrabber/base.py | 12 +++++++++++- .../datagrabber/tests/test_base_datagrabber.py | 16 ++++++++++++++-- 2 files changed, 25 insertions(+), 3 deletions(-) diff --git a/junifer/datagrabber/base.py b/junifer/datagrabber/base.py index d99ec3821..a7c97084f 100644 --- a/junifer/datagrabber/base.py +++ b/junifer/datagrabber/base.py @@ -248,7 +248,17 @@ class PatternDataGrabber(BaseDataGrabber): for t_type in self.types: t_pattern = self.patterns[t_type] # type: ignore t_replace = self._replace_patterns(element, t_pattern) - t_out = self.datadir / element[0] / t_replace + if '*' in t_replace: + t_matches = list(self.datadir.glob(t_replace)) + if len(t_matches) > 1: + raise_error( + f'More than one file matches for {element} / {t_type}: ' + f'{t_matches}') + elif len(t_matches) == 0: + raise_error(f'No file matches for {element} / {t_type}') + t_out = t_matches[0] + else: + t_out = self.datadir / t_replace out[t_type] = dict(path=t_out) # Meta here is element and types out['meta']['element'] = dict(zip(self.replacements, element)) diff --git a/junifer/datagrabber/tests/test_base_datagrabber.py b/junifer/datagrabber/tests/test_base_datagrabber.py index 3b5f2f8f1..2a4df25d7 100644 --- a/junifer/datagrabber/tests/test_base_datagrabber.py +++ b/junifer/datagrabber/tests/test_base_datagrabber.py @@ -114,8 +114,8 @@ def test_bids_datalad_PatternDataGrabber(): """Test a subject-based BIDS datalad datagrabber""" types = ['T1w', 'bold'] patterns = { - 'T1w': 'anat/{subject}_T1w.nii.gz', - 'bold': 'func/{subject}_task-rest_bold.nii.gz' + 'T1w': '{subject}/anat/{subject}_T1w.nii.gz', + 'bold': '{subject}/func/{subject}_task-rest_bold.nii.gz' } replacements = ['subject'] @@ -163,7 +163,19 @@ def test_bids_datalad_PatternDataGrabber(): with tempfile.TemporaryDirectory() as tmpdir: datadir = Path(tmpdir) / 'dataset' # Need this for testing + patterns = { + 'T1w': '{subject}/anat/{subject}_T*w.nii.gz', + 'bold': '{subject}/func/{subject}_task-rest_*.nii.gz' + } with MyDataGrabber(rootdir=rootdir, uri=repo_uri, types=types, patterns=patterns, datadir=datadir, replacements=replacements) as dg: assert dg.datadir == datadir / rootdir + for elem in dg: + t_sub = dg[elem] + assert 'path' in t_sub['T1w'] + assert t_sub['T1w']['path'] == \ + (dg.datadir / f'{elem}/anat/{elem}_T1w.nii.gz') + assert 'path' in t_sub['bold'] + assert t_sub['bold']['path'] == \ + (dg.datadir / f'{elem}/func/{elem}_task-rest_bold.nii.gz') -- 2.52.0 From 7106f3d02cd958ed2995d77cf19a677d8fe06189 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Tue, 29 Mar 2022 13:18:25 +0200 Subject: [PATCH 057/287] flake --- junifer/datagrabber/base.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/junifer/datagrabber/base.py b/junifer/datagrabber/base.py index a7c97084f..1e8ac9c8c 100644 --- a/junifer/datagrabber/base.py +++ b/junifer/datagrabber/base.py @@ -252,8 +252,8 @@ class PatternDataGrabber(BaseDataGrabber): t_matches = list(self.datadir.glob(t_replace)) if len(t_matches) > 1: raise_error( - f'More than one file matches for {element} / {t_type}: ' - f'{t_matches}') + f'More than one file matches for {element} / {t_type}:' + f' {t_matches}') elif len(t_matches) == 0: raise_error(f'No file matches for {element} / {t_type}') t_out = t_matches[0] -- 2.52.0 From f3c1bd1e3e9488e7443b85554e5d2f73994c81b6 Mon Sep 17 00:00:00 2001 From: LeSasse Date: Tue, 29 Mar 2022 16:25:05 +0200 Subject: [PATCH 058/287] updated HCP1200 pattern and __getitem__() method and added test for juseless --- junifer/configs/tests/test_juseless.py | 13 ++++++ junifer/datagrabber/hcp.py | 56 ++++++++++++++++++++++++-- 2 files changed, 65 insertions(+), 4 deletions(-) diff --git a/junifer/configs/tests/test_juseless.py b/junifer/configs/tests/test_juseless.py index 03a6bb8fc..894eaf793 100644 --- a/junifer/configs/tests/test_juseless.py +++ b/junifer/configs/tests/test_juseless.py @@ -1,8 +1,10 @@ import socket import pytest +import os from junifer.configs.juseless import JuselessDataladUKBVBM from junifer.utils.logging import configure_logging +from junifer.datagrabber.hcp import DataladHCP1200 if socket.gethostname() != 'juseless': pytest.skip('This tests are only for juseless', allow_module_level=True) @@ -16,3 +18,14 @@ def test_juselessdataladukbvbm_datagrabber(): assert 'VBM_GM' in out assert out['VBM_GM'].name == 'm0wp1sub-2670511_ses-2_T1w.nii.gz' assert out['VBM_GM'].exists() + + +def test_juselessdataladhcp_datagrabber(): + with DataladHCP1200() as dg: + all_elements = dg.get_elements() + test_element = all_elements[0] + + out = dg[test_element] + + assert out['BOLD'].exists() + assert os.path.isfile(out['BOLD']['path']) diff --git a/junifer/datagrabber/hcp.py b/junifer/datagrabber/hcp.py index a063fbd16..5336f1ab9 100644 --- a/junifer/datagrabber/hcp.py +++ b/junifer/datagrabber/hcp.py @@ -38,8 +38,8 @@ class HCP1200(PatternDataGrabber): replacements = ['subject', 'task', 'phase_encoding'] patterns = { 'BOLD': ('{subject}/MNINonLinear/Results/' - '*fMRI_{task}_{phase_encoding}/' - '*fMRI_{task}_{phase_encoding}_hp2000_clean.nii.gz') + '{task}_{phase_encoding}/' + '{task}_{phase_encoding}_hp2000_clean.nii.gz') } super().__init__( types=types, datadir=datadir, patterns=patterns, @@ -68,7 +68,16 @@ class HCP1200(PatternDataGrabber): ] if self.phase_encodings is None: - self.phase_encodings = ["LR", "RL"] + self.phase_encodings = ['LR', 'RL'] + + repl_tasks = [] + for t in self.tasks: + if 'REST' in t: + repl_tasks.append(f'rfMRI_{t}') + else: + repl_tasks.append(f'tfMRI_{t}') + + self.tasks = repl_tasks def get_elements(self): """Get the list of subjects in the dataset. @@ -88,9 +97,48 @@ class HCP1200(PatternDataGrabber): return elems + def __getitem__(self, element): + """Index one element in the dataset. + + Parameters + ---------- + element : tuple[str, str] + The element to be indexed. First element in the tuple is the + subject, second element is the task, third element is the + phase encoding direction. + + Returns + ------- + out : dict[str -> Path] + Dictionary of paths for each type of data required for the + specified element. + """ + + sub, task, phase_encoding = element + + out = super().__getitem__(element) + + self.tasks = [x.split("_")[1] for x in self.tasks] + + out['meta'] = dict(datagrabber=self.get_meta()) + out['meta']['element'] = dict( + subject=sub, task=task, phase_encoding=phase_encoding + ) + + repl_tasks = [] + for t in self.tasks: + if 'REST' in t: + repl_tasks.append(f'rfMRI_{t}') + else: + repl_tasks.append(f'tfMRI_{t}') + + self.tasks = repl_tasks + + return out + @register_datagrabber -class DataladHCP1200(HCP1200, DataladDataGrabber): +class DataladHCP1200(DataladDataGrabber, HCP1200,): def __init__(self, datadir=None, tasks=None, phase_encodings=None): uri = ( 'https://github.com/datalad-datasets/' -- 2.52.0 From 7026ba71927485e44b49ea56d552ce7631e027ef Mon Sep 17 00:00:00 2001 From: LeSasse Date: Tue, 29 Mar 2022 16:38:11 +0200 Subject: [PATCH 059/287] improved __getitem__() --- junifer/datagrabber/hcp.py | 27 +++++---------------------- 1 file changed, 5 insertions(+), 22 deletions(-) diff --git a/junifer/datagrabber/hcp.py b/junifer/datagrabber/hcp.py index 5336f1ab9..decf94873 100644 --- a/junifer/datagrabber/hcp.py +++ b/junifer/datagrabber/hcp.py @@ -70,15 +70,6 @@ class HCP1200(PatternDataGrabber): if self.phase_encodings is None: self.phase_encodings = ['LR', 'RL'] - repl_tasks = [] - for t in self.tasks: - if 'REST' in t: - repl_tasks.append(f'rfMRI_{t}') - else: - repl_tasks.append(f'tfMRI_{t}') - - self.tasks = repl_tasks - def get_elements(self): """Get the list of subjects in the dataset. @@ -116,24 +107,16 @@ class HCP1200(PatternDataGrabber): sub, task, phase_encoding = element - out = super().__getitem__(element) + if "REST" in task: + new_task = f'rfMRI_{task}' + else: + new_task = f'tfMRI_{task}' - self.tasks = [x.split("_")[1] for x in self.tasks] - - out['meta'] = dict(datagrabber=self.get_meta()) + out = super().__getitem__((sub, new_task, phase_encoding)) out['meta']['element'] = dict( subject=sub, task=task, phase_encoding=phase_encoding ) - repl_tasks = [] - for t in self.tasks: - if 'REST' in t: - repl_tasks.append(f'rfMRI_{t}') - else: - repl_tasks.append(f'tfMRI_{t}') - - self.tasks = repl_tasks - return out -- 2.52.0 From 81c2e1bf1785b70b62d86aa5b7258d513b054743 Mon Sep 17 00:00:00 2001 From: LeSasse Date: Wed, 30 Mar 2022 12:11:48 +0200 Subject: [PATCH 060/287] validating HCP datagrabber tasks and phase encodings, also added initial confound removal module --- junifer/datagrabber/hcp.py | 43 +++++--- junifer/preprocess/confounds.py | 169 ++++++++++++++++++++++++++++++++ 2 files changed, 198 insertions(+), 14 deletions(-) create mode 100644 junifer/preprocess/confounds.py diff --git a/junifer/datagrabber/hcp.py b/junifer/datagrabber/hcp.py index decf94873..f3ef35215 100644 --- a/junifer/datagrabber/hcp.py +++ b/junifer/datagrabber/hcp.py @@ -32,8 +32,6 @@ class HCP1200(PatternDataGrabber): """ types = ['BOLD'] - # TODO: Validate tasks - # TODO: Validate phase_encodings replacements = ['subject', 'task', 'phase_encoding'] patterns = { @@ -54,22 +52,39 @@ class HCP1200(PatternDataGrabber): if isinstance(self.phase_encodings, str): self.phase_encodings = [self.phase_encodings] + all_tasks = [ + 'REST1', + 'REST2', + 'SOCIAL', + 'WM', + 'RELATIONAL', + 'EMOTION', + 'LANGUAGE', + 'GAMBLING', + 'MOTOR', + ] + if self.tasks is None: - self.tasks = [ - 'REST1', - 'REST2', - 'SOCIAL', - 'WM', - 'RELATIONAL', - 'EMOTION', - 'LANGUAGE', - 'GAMBLING', - 'MOTOR', - ] + self.tasks = all_tasks if self.phase_encodings is None: self.phase_encodings = ['LR', 'RL'] + for task in self.tasks: + if task not in all_tasks: + raise ValueError( + f'{task} not a valid HCP-YA fMRI task input! \n' + f'task can be any of {all_tasks}' + ) + + for pe in self.phase_encodings: + if pe not in ['LR', 'RL']: + raise ValueError( + f'{pe} not a valid HCP-YA phase encoding. \n' + 'phase_encoding can be LR or RL (or both)!' + ) + + def get_elements(self): """Get the list of subjects in the dataset. @@ -121,7 +136,7 @@ class HCP1200(PatternDataGrabber): @register_datagrabber -class DataladHCP1200(DataladDataGrabber, HCP1200,): +class DataladHCP1200(DataladDataGrabber, HCP1200): def __init__(self, datadir=None, tasks=None, phase_encodings=None): uri = ( 'https://github.com/datalad-datasets/' diff --git a/junifer/preprocess/confounds.py b/junifer/preprocess/confounds.py new file mode 100644 index 000000000..5504f5cdf --- /dev/null +++ b/junifer/preprocess/confounds.py @@ -0,0 +1,169 @@ +import numpy as np + +from junifer.markers.base import PipelineStepMixin + +class BaseConfoundRemover(PipelineStepMixin): + """ A base class to read confound files and select columns according to + a pre-defined strategy + + """ + + # TODO: Properly Validate Input + # TODO: implement get_output_kind + # TODO: implement fit_transform + + # lower priority + # TODO: implement more strategies from + # nilearn.interfaces.fmriprep.load_confounds for Felix's confound files, + # in particular scrubbing + # TODO: Implement read_confounds for fmriprep data + + def __init__(self, strategy): + """ Initialise a BaseConfoundReader object + + Parameters + ----------- + strategy : dict + keys of dictionary should correspond to names of noise components + to include: + - 'motion' + - 'wm_csf' + - 'global_signal' + values of dictionary should correspond to types of confounds + extracted from each signal: + - 'basic' - only the confounding time series + - 'power2' - signal + quadratic term + - 'derivatives' - signal + derivatives + - 'full' - signal + deriv. + quadratic terms + power2 deriv. + """ + self.strategy = strategy + + + def validate_input(self, input): + """Validate the input to the pipeline step. + + Parameters + ---------- + input : Junifer Data dictionary + The input to the pipeline step. + + Raises + ------ + ValueError: + If the input does not have the required data. + """ + for k, v in self.strategy.items(): + if k not in ['motion','wm_csf','global_signal']: + raise ValueError( + f'{k} not a valid noise component to include!\n' + f'If {k} is a valid parameter in ' + 'nilearn.interfaces.fmriprep.load_confounds we may ' + 'include it in the future' + ) + if v not in ['basic','power2','derivatives','full']: + raise ValueError( + f'{v} not a valid type of confound to extract from input ' + 'signal!' + ) + + + def get_output_kind(self, input): + """Get the kind of the pipeline step. + + Parameters + ---------- + input : Junifer Data dictionary + The input to the pipeline step. + + Returns + ------- + output : Junifer Data dictionary + The output of the pipeline step. + """ + raise NotImplementedError('get_output_kind not implemented') + + def read_confounds(self): + raise NotImplementedError('read_confounds not implemented') + + def remove_confounds(self): + pass + + def fit_transform(self, input): + out = {} + meta = input.get('meta', {}) + bold_img = input['BOLD']['data'] + confounds_df = input['confounds'] + selected_confounds = self.read_confounds(confounds_df) + clean_img = self.remove_confounds(bold_img, confounds_df) + + + return out + +class FelixConfoundRemover(BaseConfoundRemover): + + def read_confounds(self, confound_dataframe): + + confounds_to_select = [] + # for some confounds we need to manually calculate derivatives using + # numpy and add them to the output + derivatives_to_calculate = [] + + for comp, param in self.strategy.items(): + if comp == 'motion': + + # there should be six rigid body parameters + for i in range(1,7): + # select basic + confounds_to_select.append(f'RP.{i}') + + # select squares + if param in ['power2', 'full']: + confounds_to_select.append(f'RP^2.{i}') + + # select derivatives + if param in ['derivatives', 'full']: + confounds_to_select.append(f'DRP.{i}') + + # if 'full' we should not forget the derivative + # of the squares + if param in ['full']: + confounds_to_select.append(f'DRP^2.{i}') + + elif comp == 'wm_csf': + confounds = ['WM', 'CSF'] + elif comp == 'global_signal': + confounds = ['GS'] + + for conf in confounds: + + confounds_to_select.append(conf) + + # select squares + if param in ['power2', 'full']: + confounds_to_select.append(f'{conf}^2') + + # we have to calculate derivatives (not included in felix' + # confound files) + if param in ['derivatives', 'full']: + derivatives_to_calculate.append(conf) + + if param in ['full']: + derivatives_to_calculate.append(f'{conf}^2') + + confounds_to_remove = confound_dataframe[confounds_to_select] + + for conf in derivatives_to_calculate: + confounds_to_remove[f'D{conf}'] = np.append( + np.diff(confound_dataframe[conf]), 0 + ) + + return confounds_to_remove + +class FmriprepConfoundRemover(BaseConfoundRemover): + """ A ConfoundRemover class for fmriprep output utilising + nilearn's nilearn.interfaces.fmriprep.load_confounds + """ + + def read_confounds(self): + raise NotImplementedError('read_confounds not implemented') + \ No newline at end of file -- 2.52.0 From 8bcc54abc6ee8ba5b1e67120138fcd223e228c71 Mon Sep 17 00:00:00 2001 From: LeSasse Date: Thu, 31 Mar 2022 09:10:22 +0200 Subject: [PATCH 061/287] completed first version of FelixConfoundRemover and added tests --- junifer/preprocess/confounds.py | 232 ++++++++++++++++----- junifer/preprocess/tests/test_confounds.py | 140 +++++++++++++ 2 files changed, 324 insertions(+), 48 deletions(-) create mode 100644 junifer/preprocess/tests/test_confounds.py diff --git a/junifer/preprocess/confounds.py b/junifer/preprocess/confounds.py index 5504f5cdf..774dc36a1 100644 --- a/junifer/preprocess/confounds.py +++ b/junifer/preprocess/confounds.py @@ -1,26 +1,49 @@ +# Authors: Federico Raimondo +# Leonard Sasse +# License: AGPL + import numpy as np +import pandas as pd + +from nilearn._utils.niimg_conversions import check_niimg_4d +from nilearn.masking import compute_brain_mask +from nilearn.image import clean_img from junifer.markers.base import PipelineStepMixin +from ..utils.logging import logger, raise_error + class BaseConfoundRemover(PipelineStepMixin): - """ A base class to read confound files and select columns according to + """ A base class to read confound files and select columns according to a pre-defined strategy - + """ - # TODO: Properly Validate Input - # TODO: implement get_output_kind - # TODO: implement fit_transform - # lower priority - # TODO: implement more strategies from + # TODO: implement more strategies from # nilearn.interfaces.fmriprep.load_confounds for Felix's confound files, # in particular scrubbing # TODO: Implement read_confounds for fmriprep data - def __init__(self, strategy): + def __init__( + self, + strategy={ + 'motion': 'full', + 'wm_csf': 'full', + 'global_signal': 'full' + }, + spike=None, + detrend=True, + standardize=True, + low_pass=None, + high_pass=None, + t_r=None, + mask_img=None + ): """ Initialise a BaseConfoundReader object - + + Confound removal is based on nilearn.image.clean_img + Parameters ----------- strategy : dict @@ -29,15 +52,58 @@ class BaseConfoundRemover(PipelineStepMixin): - 'motion' - 'wm_csf' - 'global_signal' - values of dictionary should correspond to types of confounds + values of dictionary should correspond to types of confounds extracted from each signal: - 'basic' - only the confounding time series - 'power2' - signal + quadratic term - 'derivatives' - signal + derivatives - 'full' - signal + deriv. + quadratic terms + power2 deriv. - """ - self.strategy = strategy + spike : None (default) | float + If None, no spike regressor is added. If spike is a float, it will + add a spike regressor for every point at which FD exceeds the + specified float. + detrend : bool + If detrending should be applied on timeseries + (before confound removal). Default=True. + standardize : bool + If True, returned signals are set to unit variance. Default=True. + low_pass : float + Low cutoff frequencies, in Hertz. + high_pass : float + High cutoff frequencies, in Hertz. + t_r : float + Repetition time, in second (sampling period). + If set to None it will use t_r from nifti header. + mask_img: Niimg-like object + If provided, signal is only cleaned from voxels inside the mask. + If mask is provided, it should have same shape and affine as imgs. + If not provided, a mask is computed using + nilearn.masking.compute_brain_mask + """ + + self.strategy = strategy + self.spike = spike + self.detrend = detrend + self.standardize = standardize + self.low_pass = low_pass + self.high_pass = high_pass + self.t_r = t_r + self.mask_img = mask_img + + for k, v in self.strategy.items(): + if k not in ['motion', 'wm_csf', 'global_signal']: + raise_error( + f'{k} not a valid noise component to include!\n' + f'If {k} is a valid parameter in ' + 'nilearn.interfaces.fmriprep.load_confounds we may ' + 'include it in the future', ValueError + ) + if v not in ['basic', 'power2', 'derivatives', 'full']: + raise_error( + f'{v} not a valid type of confound to extract from {k} ' + 'signal!', ValueError + ) def validate_input(self, input): """Validate the input to the pipeline step. @@ -52,20 +118,26 @@ class BaseConfoundRemover(PipelineStepMixin): ValueError: If the input does not have the required data. """ - for k, v in self.strategy.items(): - if k not in ['motion','wm_csf','global_signal']: - raise ValueError( - f'{k} not a valid noise component to include!\n' - f'If {k} is a valid parameter in ' - 'nilearn.interfaces.fmriprep.load_confounds we may ' - 'include it in the future' - ) - if v not in ['basic','power2','derivatives','full']: - raise ValueError( - f'{v} not a valid type of confound to extract from input ' - 'signal!' - ) + if 'BOLD' not in input: + raise_error( + 'Input does not have the required data. \n' + f'Input: {input} \n' + f'Required: BOLD \n', ValueError + ) + if 'data' not in input['BOLD']: + raise_error( + 'Input does not have the required data. \n' + f'Input: {input} \n' + f'Required: [\'BOLD\'][\'data\'] \n', ValueError + ) + # check that keys actually return a valid 4d image + # and that confounds are a dataframe + check_niimg_4d(input['BOLD']['data']) + if not isinstance(input['confounds'], pd.DataFrame): + raise_error( + 'Input does not have the required data.', ValueError + ) def get_output_kind(self, input): """Get the kind of the pipeline step. @@ -80,29 +152,85 @@ class BaseConfoundRemover(PipelineStepMixin): output : Junifer Data dictionary The output of the pipeline step. """ - raise NotImplementedError('get_output_kind not implemented') + return 'BOLD' - def read_confounds(self): - raise NotImplementedError('read_confounds not implemented') + def read_confounds(self, confound_dataframe): + """ Select relevant confounds from the specified file """ - def remove_confounds(self): - pass + logger.warning( + 'BaseConfoundRemover removes all confounds from the file without' + ' applying any selection strategy!' + ) + return confound_dataframe + + def remove_confounds(self, bold_img, confounds_df): + """ Remove confounds from the BOLD image + + bold_img : Niimg-like object + 4D image. The signals in the last dimension are filtered + (see http://nilearn.github.io/manipulating_images/input_output.html + for a detailed description of the valid input types). + confounds_df : pd.DataFrame + Dataframe containing confounds to remove. Number of rows should + correspond to number of volumes in the BOLD image. + + returns + -------- + clean_bold : Niimg-like object + input image, cleaned. + + """ + + confounds_array = confounds_df.values + + assert bold_img.get_fdata().shape[3] == confounds_array.shape[0], ( + 'Image time series and confounds have different length!' + ) + + if self.t_r is None: + logger.warning('No t_r specified, using t_r from nifti header!') + zooms = bold_img.header.get_zooms() + self.t_r = zooms[3] + + if self.mask_img is None: + self.mask_img = compute_brain_mask(bold_img) + + clean_bold = clean_img( + imgs=bold_img, + detrend=self.detrend, + standardize=self.standardize, + confounds=confounds_array, + low_pass=self.low_pass, + high_pass=self.high_pass, + t_r=self.t_r, + mask_img=self.mask_img + ) + + return clean_bold def fit_transform(self, input): + self.validate_input(input) out = {} meta = input.get('meta', {}) bold_img = input['BOLD']['data'] - confounds_df = input['confounds'] - selected_confounds = self.read_confounds(confounds_df) - clean_img = self.remove_confounds(bold_img, confounds_df) - + confounds_df = self.read_confounds(input['confounds']) + out['BOLD'] = {} + out['BOLD']['data'] = self.remove_confounds(bold_img, confounds_df) + out['confounds'] = confounds_df + out['meta'] = meta return out + class FelixConfoundRemover(BaseConfoundRemover): + """ A class to read confounds from confound files generated by Felix's + pipeline using CAT and some other custom scripts. It is meant to emulate + the new nilearn.interfaces.fmriprep.load_confounds as closely as possible. + + """ def read_confounds(self, confound_dataframe): - + confounds_to_select = [] # for some confounds we need to manually calculate derivatives using # numpy and add them to the output @@ -110,21 +238,21 @@ class FelixConfoundRemover(BaseConfoundRemover): for comp, param in self.strategy.items(): if comp == 'motion': - - # there should be six rigid body parameters - for i in range(1,7): + confounds = [] + # there should be six rigid body parameters + for i in range(1, 7): # select basic confounds_to_select.append(f'RP.{i}') - - # select squares + + # select squares if param in ['power2', 'full']: confounds_to_select.append(f'RP^2.{i}') - + # select derivatives if param in ['derivatives', 'full']: confounds_to_select.append(f'DRP.{i}') - - # if 'full' we should not forget the derivative + + # if 'full' we should not forget the derivative # of the squares if param in ['full']: confounds_to_select.append(f'DRP^2.{i}') @@ -133,11 +261,11 @@ class FelixConfoundRemover(BaseConfoundRemover): confounds = ['WM', 'CSF'] elif comp == 'global_signal': confounds = ['GS'] - + for conf in confounds: confounds_to_select.append(conf) - + # select squares if param in ['power2', 'full']: confounds_to_select.append(f'{conf}^2') @@ -146,19 +274,28 @@ class FelixConfoundRemover(BaseConfoundRemover): # confound files) if param in ['derivatives', 'full']: derivatives_to_calculate.append(conf) - + if param in ['full']: derivatives_to_calculate.append(f'{conf}^2') confounds_to_remove = confound_dataframe[confounds_to_select] - + + # calc additional derivatives for conf in derivatives_to_calculate: confounds_to_remove[f'D{conf}'] = np.append( np.diff(confound_dataframe[conf]), 0 ) + # add binary spike regressor if needed at given threshold + if self.spike is not None: + fd = confound_dataframe["FD"].copy() + fd.loc[fd > self.spike] = 1 + fd.loc[fd != 1] = 0 + confounds_to_remove['spike'] = fd + return confounds_to_remove + class FmriprepConfoundRemover(BaseConfoundRemover): """ A ConfoundRemover class for fmriprep output utilising nilearn's nilearn.interfaces.fmriprep.load_confounds @@ -166,4 +303,3 @@ class FmriprepConfoundRemover(BaseConfoundRemover): def read_confounds(self): raise NotImplementedError('read_confounds not implemented') - \ No newline at end of file diff --git a/junifer/preprocess/tests/test_confounds.py b/junifer/preprocess/tests/test_confounds.py new file mode 100644 index 000000000..baf635e3b --- /dev/null +++ b/junifer/preprocess/tests/test_confounds.py @@ -0,0 +1,140 @@ +# Authors: Federico Raimondo +# Leonard Sasse +# License: AGPL + +import numpy as np +import pandas as pd +from nibabel import Nifti1Image +from nilearn._utils.niimg_conversions import check_niimg_4d +import random +import string +from junifer.preprocess.confounds import FelixConfoundRemover + +np.random.seed(1234567) + + +def generate_conf_name(size=6, chars=string.ascii_uppercase + string.digits): + return ''.join(random.choice(chars) for _ in range(size)) + + +def _simu_img(): + + # Random 4D volume with 100 time points + vol = 100 + 10 * np.random.randn(5, 5, 2, 100) + img = Nifti1Image(vol, np.eye(4)) + # Create an nifti image with the data, and corresponding mask + mask = Nifti1Image(np.ones([5, 5, 2]), np.eye(4)) + return img, mask + + +def test_felixconfoundremover(): + # Generate a simulated BOLD img + siimg, simsk = _simu_img() + + # generate random confound dataframe with Felix's column naming + confound_column_names = [] + + confound_column_names.append('FD') + + for i in range(1, 7): + confound_column_names.append(f'RP.{i}') + confound_column_names.append(f'RP^2.{i}') + confound_column_names.append(f'DRP.{i}') + confound_column_names.append(f'DRP^2.{i}') + + for conf in ['WM', 'CSF', 'GS']: + confound_column_names.append(conf) + confound_column_names.append(f'{conf}^2') + + # add some random irrelevant confounds + for i in range(10): + confound_column_names.append(generate_conf_name()) + + np.random.shuffle(confound_column_names) + n_cols = len(confound_column_names) + confounds_df = pd.DataFrame( + np.random.randint(0, 100, size=(100, n_cols)), + columns=confound_column_names + ) + + # generate confound removal strategies with varying numbers of parameters + list_of_strategy_tuples = [] + + # 36 params + strat1 = { + 'motion': 'full', + 'wm_csf': 'full', + 'global_signal': 'full' + } + list_of_strategy_tuples.append((strat1, 36)) + + # 24 params + strat2 = { + 'motion': 'full', + } + list_of_strategy_tuples.append((strat2, 24)) + + # 9 params + strat3 = { + 'motion': 'basic', + 'wm_csf': 'basic', + 'global_signal': 'basic' + } + list_of_strategy_tuples.append((strat3, 9)) + + # 6 params + strat4 = { + 'motion': 'basic', + } + list_of_strategy_tuples.append((strat4, 6)) + + # 2 params + strat5 = { + 'wm_csf': 'basic', + } + list_of_strategy_tuples.append((strat5, 2)) + + # generate a junifer pipeline data object dictionary + input_data_obj = {} + input_data_obj['meta'] = {} + input_data_obj['BOLD'] = {} + input_data_obj['BOLD']['data'] = siimg + input_data_obj['confounds'] = confounds_df + + # run confound removal + for spike in [None, 0.25]: + for strategy, n_params in list_of_strategy_tuples: + FCR = FelixConfoundRemover( + strategy=strategy, + spike=spike, + mask_img=simsk, + t_r=0.75 + ) + + # first test some methods manually + FCR.validate_input(input_data_obj) + out_type = FCR.get_output_kind(input_data_obj) + + assert out_type in ['BOLD'] + + selected_conf_frame = FCR.read_confounds(confounds_df) + timepoints, n_sel_confs = selected_conf_frame.shape + + assert timepoints == 100 + assert isinstance(selected_conf_frame, pd.DataFrame) + + if spike is None: + assert n_sel_confs == n_params + else: + assert n_sel_confs == n_params + 1 + + cl_bold = FCR.remove_confounds(siimg, selected_conf_frame) + check_niimg_4d(cl_bold) + + # finally test fit_transform + output_data_obj = FCR.fit_transform(input_data_obj) + + for key in ['BOLD', 'confounds', 'meta']: + assert key in output_data_obj + + check_niimg_4d(output_data_obj['BOLD']['data']) -- 2.52.0 From d799c7028696a01b5646cf8057982981ff6db25f Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Thu, 31 Mar 2022 15:36:52 +0200 Subject: [PATCH 062/287] In theory, this should work --- junifer/markers/base.py | 15 +- junifer/preprocess/confounds.py | 465 ++++++++++++++++++++------------ 2 files changed, 299 insertions(+), 181 deletions(-) diff --git a/junifer/markers/base.py b/junifer/markers/base.py index 9d736aa15..0471744d0 100644 --- a/junifer/markers/base.py +++ b/junifer/markers/base.py @@ -18,8 +18,9 @@ class PipelineStepMixin(): Parameters ---------- - input : Junifer Data dictionary - The input to the pipeline step. + input : list[str] + The input to the pipeline step. The list must contain the + available Junifer Data dictionary keys. Raises ------ @@ -33,13 +34,15 @@ class PipelineStepMixin(): Parameters ---------- - input : Junifer Data dictionary - The input to the pipeline step. + input : list[str] + The input to the pipeline step. The list must contain the + available Junifer Data dictionary keys. Returns ------- - output : Junifer Data dictionary - The output of the pipeline step. + output : list[str] + The updated list of available Junifer Data dictionary keys after + the pipeline step. """ raise NotImplementedError('get_output_kind not implemented') diff --git a/junifer/preprocess/confounds.py b/junifer/preprocess/confounds.py index 774dc36a1..632bdbd24 100644 --- a/junifer/preprocess/confounds.py +++ b/junifer/preprocess/confounds.py @@ -25,62 +25,56 @@ class BaseConfoundRemover(PipelineStepMixin): # in particular scrubbing # TODO: Implement read_confounds for fmriprep data - def __init__( - self, - strategy={ - 'motion': 'full', - 'wm_csf': 'full', - 'global_signal': 'full' - }, - spike=None, - detrend=True, - standardize=True, - low_pass=None, - high_pass=None, - t_r=None, - mask_img=None - ): + def __init__(self, strategy=None, spike=None, detrend=True, + standardize=True, low_pass=None, high_pass=None, t_r=None, + mask_img=None): """ Initialise a BaseConfoundReader object Confound removal is based on nilearn.image.clean_img Parameters ----------- - strategy : dict - keys of dictionary should correspond to names of noise components - to include: + strategy : dict[str -> str] + The keys of the dictionary should correspond to names of noise + components to include: - 'motion' - 'wm_csf' - 'global_signal' - values of dictionary should correspond to types of confounds + The values of dictionary should correspond to types of confounds extracted from each signal: - - 'basic' - only the confounding time series - - 'power2' - signal + quadratic term - - 'derivatives' - signal + derivatives - - 'full' - signal + deriv. + quadratic terms + power2 deriv. - spike : None (default) | float + - 'basic': only the confounding time series + - 'power2': signal + quadratic term + - 'derivatives': signal + derivatives + - 'full': signal + deriv. + quadratic terms + power2 deriv. + spike : float | None (default) If None, no spike regressor is added. If spike is a float, it will add a spike regressor for every point at which FD exceeds the specified float. detrend : bool - If detrending should be applied on timeseries - (before confound removal). Default=True. + If True (default), detrending will be applied on timeseries + (before confound removal). standardize : bool - If True, returned signals are set to unit variance. Default=True. - low_pass : float - Low cutoff frequencies, in Hertz. - high_pass : float - High cutoff frequencies, in Hertz. + If True (default), returned signals are set to unit variance. + low_pass : float | None (default) + Low cutoff frequencies, in Hertz. If None, no filtering is applied. + high_pass : float | None (default) + High cutoff frequencies, in Hertz. If None, no filtering is + applied. t_r : float Repetition time, in second (sampling period). - If set to None it will use t_r from nifti header. + If None (default) it will use t_r from nifti header. mask_img: Niimg-like object If provided, signal is only cleaned from voxels inside the mask. If mask is provided, it should have same shape and affine as imgs. If not provided, a mask is computed using nilearn.masking.compute_brain_mask - """ + if strategy is None: + strategy = { + 'motion': 'full', + 'wm_csf': 'full', + 'global_signal': 'full' + } self.strategy = strategy self.spike = spike @@ -91,52 +85,57 @@ class BaseConfoundRemover(PipelineStepMixin): self.t_r = t_r self.mask_img = mask_img - for k, v in self.strategy.items(): - if k not in ['motion', 'wm_csf', 'global_signal']: - raise_error( - f'{k} not a valid noise component to include!\n' - f'If {k} is a valid parameter in ' - 'nilearn.interfaces.fmriprep.load_confounds we may ' - 'include it in the future', ValueError - ) - if v not in ['basic', 'power2', 'derivatives', 'full']: - raise_error( - f'{v} not a valid type of confound to extract from {k} ' - 'signal!', ValueError - ) + self._valid_components = ['motion', 'wm_csf', 'global_signal'] + self._valid_confounds = ['basic', 'power2', 'derivatives', 'full'] + + if any(not isinstance(k, str) for k in strategy.keys()): + raise_error( + 'Strategy keys must be strings', ValueError + ) + + if any(not isinstance(v, str) for v in strategy.values()): + raise_error( + 'Strategy values must be strings', ValueError + ) + + if any(x not in self._valid_components for x in strategy.keys()): + raise_error( + f'Invalid component names {list(strategy.keys())}. ' + f'Valid components are {self._valid_components}.\n' + f'If any of them is a valid parameter in ' + 'nilearn.interfaces.fmriprep.load_confounds we may ' + 'include it in the future', ValueError + ) + + if any(x not in self._valid_confounds for x in strategy.values()): + raise_error( + f'Invalid component names {list(strategy.values())}. ' + f'Valid confound types are {self._valid_confounds}.\n' + f'If any of them is a valid parameter in ' + 'nilearn.interfaces.fmriprep.load_confounds we may ' + 'include it in the future', ValueError + ) def validate_input(self, input): """Validate the input to the pipeline step. Parameters ---------- - input : Junifer Data dictionary - The input to the pipeline step. + input : list[str] + The input to the pipeline step. The list must contain the + available Junifer Data dictionary keys. Raises ------ ValueError: If the input does not have the required data. """ - if 'BOLD' not in input: + _required_inputs = ['BOLD', 'confounds'] + if any(x not in input for x in _required_inputs): raise_error( 'Input does not have the required data. \n' f'Input: {input} \n' - f'Required: BOLD \n', ValueError - ) - if 'data' not in input['BOLD']: - raise_error( - 'Input does not have the required data. \n' - f'Input: {input} \n' - f'Required: [\'BOLD\'][\'data\'] \n', ValueError - ) - - # check that keys actually return a valid 4d image - # and that confounds are a dataframe - check_niimg_4d(input['BOLD']['data']) - if not isinstance(input['confounds'], pd.DataFrame): - raise_error( - 'Input does not have the required data.', ValueError + f'Required (all off): {_required_inputs} \n', ValueError ) def get_output_kind(self, input): @@ -144,26 +143,50 @@ class BaseConfoundRemover(PipelineStepMixin): Parameters ---------- - input : Junifer Data dictionary - The input to the pipeline step. + input : list[str] + The input to the pipeline step. The list must contain the + available Junifer Data dictionary keys. Returns ------- - output : Junifer Data dictionary - The output of the pipeline step. + output : list[str] + The updated list of available Junifer Data dictionary keys after + the pipeline step. """ - return 'BOLD' + # Does not add any new keys + return input - def read_confounds(self, confound_dataframe): + def _pick_confounds(self, input): """ Select relevant confounds from the specified file """ + to_select = [] + confounds_df = input['data'] + confounds_spec = input['names']['spec'] + derivatives_to_compute = input['names']['derivatives'] + spike_name = input['names']['spike'] - logger.warning( - 'BaseConfoundRemover removes all confounds from the file without' - ' applying any selection strategy!' - ) - return confound_dataframe + # Get all the column names according to the strategy + for comp, param in self.strategy.items(): + to_select.append(confounds_spec[comp][param]) - def remove_confounds(self, bold_img, confounds_df): + # Add derivatives if needed + to_compute = [x in derivatives_to_compute.keys() for x in to_select] + out_df = confounds_df.copy() + if any(to_compute): + for t_dst in to_compute: + t_src = derivatives_to_compute[t_dst] + out_df[t_dst] = np.append( # type: ignore + np.diff(out_df[t_src]), 0) # type: ignore + + # add binary spike regressor if needed at given threshold + if self.spike is not None: + fd = confounds_df[spike_name].copy() + fd.loc[fd > self.spike] = 1 + fd.loc[fd != 1] = 0 + out_df['spike'] = fd + + return out_df + + def _remove_confounds(self, bold_img, confounds_df): """ Remove confounds from the BOLD image bold_img : Niimg-like object @@ -183,17 +206,17 @@ class BaseConfoundRemover(PipelineStepMixin): confounds_array = confounds_df.values - assert bold_img.get_fdata().shape[3] == confounds_array.shape[0], ( - 'Image time series and confounds have different length!' - ) - - if self.t_r is None: - logger.warning('No t_r specified, using t_r from nifti header!') + t_r = self.t_r + if t_r is None: + logger.info('No t_r specified, using t_r from nifti header!') zooms = bold_img.header.get_zooms() - self.t_r = zooms[3] + t_r = zooms[3] + logger.info(f'Read t_r from nifti header: {t_r}', ) - if self.mask_img is None: - self.mask_img = compute_brain_mask(bold_img) + mask_img = self.mask_img + if mask_img is None: + logger.info('Computing brain mask from image') + mask_img = compute_brain_mask(bold_img) clean_bold = clean_img( imgs=bold_img, @@ -202,104 +225,196 @@ class BaseConfoundRemover(PipelineStepMixin): confounds=confounds_array, low_pass=self.low_pass, high_pass=self.high_pass, - t_r=self.t_r, - mask_img=self.mask_img + t_r=t_r, + mask_img=mask_img ) return clean_bold - def fit_transform(self, input): - self.validate_input(input) - out = {} - meta = input.get('meta', {}) - bold_img = input['BOLD']['data'] - confounds_df = self.read_confounds(input['confounds']) - out['BOLD'] = {} - out['BOLD']['data'] = self.remove_confounds(bold_img, confounds_df) - out['confounds'] = confounds_df - out['meta'] = meta + def _validate_data(self, input): + # Bold must be 4D niimg + check_niimg_4d(input['BOLD']['data']) - return out - - -class FelixConfoundRemover(BaseConfoundRemover): - """ A class to read confounds from confound files generated by Felix's - pipeline using CAT and some other custom scripts. It is meant to emulate - the new nilearn.interfaces.fmriprep.load_confounds as closely as possible. - - """ - - def read_confounds(self, confound_dataframe): - - confounds_to_select = [] - # for some confounds we need to manually calculate derivatives using - # numpy and add them to the output - derivatives_to_calculate = [] - - for comp, param in self.strategy.items(): - if comp == 'motion': - confounds = [] - # there should be six rigid body parameters - for i in range(1, 7): - # select basic - confounds_to_select.append(f'RP.{i}') - - # select squares - if param in ['power2', 'full']: - confounds_to_select.append(f'RP^2.{i}') - - # select derivatives - if param in ['derivatives', 'full']: - confounds_to_select.append(f'DRP.{i}') - - # if 'full' we should not forget the derivative - # of the squares - if param in ['full']: - confounds_to_select.append(f'DRP^2.{i}') - - elif comp == 'wm_csf': - confounds = ['WM', 'CSF'] - elif comp == 'global_signal': - confounds = ['GS'] - - for conf in confounds: - - confounds_to_select.append(conf) - - # select squares - if param in ['power2', 'full']: - confounds_to_select.append(f'{conf}^2') - - # we have to calculate derivatives (not included in felix' - # confound files) - if param in ['derivatives', 'full']: - derivatives_to_calculate.append(conf) - - if param in ['full']: - derivatives_to_calculate.append(f'{conf}^2') - - confounds_to_remove = confound_dataframe[confounds_to_select] - - # calc additional derivatives - for conf in derivatives_to_calculate: - confounds_to_remove[f'D{conf}'] = np.append( - np.diff(confound_dataframe[conf]), 0 + # Confounds must be a dataframe + if not isinstance(input['confounds']['data'], pd.DataFrame): + raise_error( + 'confounds data must be a pandas dataframe', ValueError ) - # add binary spike regressor if needed at given threshold - if self.spike is not None: - fd = confound_dataframe["FD"].copy() - fd.loc[fd > self.spike] = 1 - fd.loc[fd != 1] = 0 - confounds_to_remove['spike'] = fd + confound_df = input['confounds']['data'] + bold_img = input['BOLD']['data'] + if bold_img.get_fdata().shape[3] != len(confound_df): + raise_error( + 'Image time series and confounds have different length!\n' + f'\tImage time series: { bold_img.get_fdata().shape[3]}\n' + f'\tConfounds: {len(confound_df)}') - return confounds_to_remove + # Check the column names of the dataframe and the spec + # spec must be a dictionary: + # { + # 'motion': { + # 'basic': [(list of columns)] + # 'power2', [(list of columns)] + # 'derivatives', [(list of columns)] + # 'full', [(list of columns)]}, + # 'wm_csf': { + # 'basic': [(list of columns)] + # 'power2', [(list of columns)] + # 'derivatives', [(list of columns)] + # 'full', [(list of columns)]} + # 'global_signal': { + # 'basic': [(list of columns)] + # 'power2', [(list of columns)] + # 'derivatives', [(list of columns)] + # 'full', [(list of columns)]} + # } + + # Check the columns in the dataframe + conf_spec = input['confounds']['names']['spec'] + if any(x not in conf_spec.keys() for x in self._valid_components): + raise_error( + 'All of the component types must be in the confounds data ' + 'object `spec`. Please check your datagrabber.', ValueError) + + if any(x not in v.keys() for x in self._valid_confounds + for v in conf_spec.values()): + raise_error( + 'All of the confound types must be in the confounds data ' + 'object `spec`. Please check your datagrabber.', ValueError) + + spike_name = input['confounds']['names']['spike'] + + derivatives_to_compute = input['confounds']['names']['derivatives'] + if not(isinstance(derivatives_to_compute, dict)): + raise_error( + 'input["confounds"]["names"]["derivatives"] ' + 'must be a dictionary. Please check your datagrabber', + ValueError) + + if any(not (isinstance(k, str) or isinstance(v, str)) + for k, v in derivatives_to_compute.items()): + raise_error( + 'input["confounds"]["names"]["derivatives"] ' + 'must be a dictionary with string keys and values. ' + 'Please check your datagrabber', + ValueError) + + missing_derivatives = [ + x for x in derivatives_to_compute.values() + if x not in confound_df.columns] + if len(missing_derivatives) > 0: + raise_error( + 'Some of the derivatives to calculate are not in the confounds' + f' dataframe: {missing_derivatives}.' + 'Please check your data ' + f'({input["confounds"]["path"].as_posix()}) ' + 'and the datagrabber.', ValueError) + + column_names = set([x for y in conf_spec.values() for x in y]) + column_names.add(spike_name) + + missing_columns = [ + x for x in column_names + if x not in confound_df.columns and + x not in derivatives_to_compute.keys()] + + if len(missing_columns) > 0: + raise_error( + 'Some of the columns in the confound spec are not in the ' + f'confounds dataframe: {missing_columns}. ' + 'Please check your data ' + f'({input["confounds"]["path"].as_posix()}) ' + 'and the datagrabber.', ValueError) + + def fit_transform(self, input): + self._validate_data(input) + bold_img = input['BOLD']['data'] + confounds_df = self._pick_confounds(input['confounds']) + input['BOLD']['data'] = self._remove_confounds(bold_img, confounds_df) + + # TODO: Update meta + return input -class FmriprepConfoundRemover(BaseConfoundRemover): - """ A ConfoundRemover class for fmriprep output utilising - nilearn's nilearn.interfaces.fmriprep.load_confounds - """ +# class FelixConfoundRemover(BaseConfoundRemover): +# """ A class to read confounds from confound files generated by Felix's +# pipeline using CAT and some other custom scripts. It is meant to emulate +# the new nilearn.interfaces.fmriprep.load_confounds as closely as +# possible. - def read_confounds(self): - raise NotImplementedError('read_confounds not implemented') +# """ + +# def read_confounds(self, confound_dataframe): + +# confounds_to_select = [] +# # for some confounds we need to manually calculate derivatives using +# # numpy and add them to the output +# derivatives_to_compute = [] + +# for comp, param in self.strategy.items(): +# if comp == 'motion': +# confounds = [] +# # there should be six rigid body parameters +# for i in range(1, 7): +# # select basic +# confounds_to_select.append(f'RP.{i}') + +# # select squares +# if param in ['power2', 'full']: +# confounds_to_select.append(f'RP^2.{i}') + +# # select derivatives +# if param in ['derivatives', 'full']: +# confounds_to_select.append(f'DRP.{i}') + +# # if 'full' we should not forget the derivative +# # of the squares +# if param in ['full']: +# confounds_to_select.append(f'DRP^2.{i}') + +# elif comp == 'wm_csf': +# confounds = ['WM', 'CSF'] +# elif comp == 'global_signal': +# confounds = ['GS'] + +# for conf in confounds: + +# confounds_to_select.append(conf) + +# # select squares +# if param in ['power2', 'full']: +# confounds_to_select.append(f'{conf}^2') + +# # we have to calculate derivatives (not included in felix' +# # confound files) +# if param in ['derivatives', 'full']: +# derivatives_to_compute.append(conf) + +# if param in ['full']: +# derivatives_to_compute.append(f'{conf}^2') + +# confounds_to_remove = confound_dataframe[confounds_to_select] + +# # calc additional derivatives +# for conf in derivatives_to_compute: +# confounds_to_remove[f'D{conf}'] = np.append( +# np.diff(confound_dataframe[conf]), 0 +# ) + +# # add binary spike regressor if needed at given threshold +# if self.spike is not None: +# fd = confound_dataframe["FD"].copy() +# fd.loc[fd > self.spike] = 1 +# fd.loc[fd != 1] = 0 +# confounds_to_remove['spike'] = fd + +# return confounds_to_remove + + +# class FmriprepConfoundRemover(BaseConfoundRemover): +# """ A ConfoundRemover class for fmriprep output utilising +# nilearn's nilearn.interfaces.fmriprep.load_confounds +# """ + +# def read_confounds(self): +# raise NotImplementedError('read_confounds not implemented') -- 2.52.0 From 98e344145ea6410924fd25dfb264ba6d35a5455a Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Tue, 12 Apr 2022 13:10:48 +0300 Subject: [PATCH 063/287] WIP: Multiple datagrabber + test confound removal --- junifer/datagrabber/meta.py | 32 ++ junifer/preprocess/confounds.py | 21 +- junifer/preprocess/tests/test_confounds.py | 350 ++++++++++++++++----- 3 files changed, 319 insertions(+), 84 deletions(-) create mode 100644 junifer/datagrabber/meta.py diff --git a/junifer/datagrabber/meta.py b/junifer/datagrabber/meta.py new file mode 100644 index 000000000..1b20949d2 --- /dev/null +++ b/junifer/datagrabber/meta.py @@ -0,0 +1,32 @@ +from .base import BaseDataGrabber + + +class MultipleDataGrabber(BaseDataGrabber): + def __init__(self, datagrabbers): + # TODO: Check datagrabbers consistency + # - same element keys + # - no overlapping types + self._datagrabbers = datagrabbers + + def get_types(self): + types = [x for dg in self._datagrabbers for x in dg.get_types()] + return types + + def get_meta(self): + t_meta = {} + t_meta['class'] = self.__class__.__name__ + t_meta['datagrabbers'] = [dg.get_meta() for dg in self._datagrabbers] + + def __enter__(self): + for dg in self._datagrabbers: + dg.__enter__() + + def __exit__(self, exc_type, exc_value, exc_traceback): + for dg in self._datagrabbers: + dg.__exit__(exc_type, exc_value, exc_traceback) + + def __getitem__(self, element): + out = {} + for dg in self._datagrabbers: + t_out = dg[element] + out.update(t_out) diff --git a/junifer/preprocess/confounds.py b/junifer/preprocess/confounds.py index 632bdbd24..0d56f818a 100644 --- a/junifer/preprocess/confounds.py +++ b/junifer/preprocess/confounds.py @@ -161,22 +161,21 @@ class BaseConfoundRemover(PipelineStepMixin): to_select = [] confounds_df = input['data'] confounds_spec = input['names']['spec'] - derivatives_to_compute = input['names']['derivatives'] + derivatives_to_compute = input['names'].get('derivatives', {}) spike_name = input['names']['spike'] # Get all the column names according to the strategy for comp, param in self.strategy.items(): - to_select.append(confounds_spec[comp][param]) + to_select.extend(confounds_spec[comp][param]) # Add derivatives if needed to_compute = [x in derivatives_to_compute.keys() for x in to_select] out_df = confounds_df.copy() if any(to_compute): - for t_dst in to_compute: - t_src = derivatives_to_compute[t_dst] - out_df[t_dst] = np.append( # type: ignore - np.diff(out_df[t_src]), 0) # type: ignore - + for t_dst, t_src in derivatives_to_compute.items(): + out_df[t_dst] = np.append( # type: ignore + np.diff(out_df[t_src]), 0) # type: ignore + out_df = out_df[to_select] # add binary spike regressor if needed at given threshold if self.spike is not None: fd = confounds_df[spike_name].copy() @@ -284,7 +283,8 @@ class BaseConfoundRemover(PipelineStepMixin): spike_name = input['confounds']['names']['spike'] - derivatives_to_compute = input['confounds']['names']['derivatives'] + derivatives_to_compute = input['confounds']['names'].get( + 'derivatives', {}) if not(isinstance(derivatives_to_compute, dict)): raise_error( 'input["confounds"]["names"]["derivatives"] ' @@ -310,7 +310,10 @@ class BaseConfoundRemover(PipelineStepMixin): f'({input["confounds"]["path"].as_posix()}) ' 'and the datagrabber.', ValueError) - column_names = set([x for y in conf_spec.values() for x in y]) + t_conf_spec = {k: input['confounds']['names']['spec'][k][v] + for k, v in self.strategy.items()} + + column_names = set([x for y in t_conf_spec.values() for x in y]) column_names.add(spike_name) missing_columns = [ diff --git a/junifer/preprocess/tests/test_confounds.py b/junifer/preprocess/tests/test_confounds.py index baf635e3b..ede77c861 100644 --- a/junifer/preprocess/tests/test_confounds.py +++ b/junifer/preprocess/tests/test_confounds.py @@ -2,13 +2,14 @@ # Leonard Sasse # License: AGPL +from pathlib import Path import numpy as np import pandas as pd from nibabel import Nifti1Image from nilearn._utils.niimg_conversions import check_niimg_4d import random import string -from junifer.preprocess.confounds import FelixConfoundRemover +from junifer.preprocess.confounds import BaseConfoundRemover np.random.seed(1234567) @@ -27,24 +28,45 @@ def _simu_img(): return img, mask -def test_felixconfoundremover(): +def test_baseconfoundremover(): # Generate a simulated BOLD img siimg, simsk = _simu_img() # generate random confound dataframe with Felix's column naming + + motion_basic = [f'RP.{i}' for i in range(1, 7)] + motion_power2 = [f'RP^2.{i}' for i in range(1, 7)] + motion_derivatives = [f'DRP.{i}' for i in range(1, 7)] + motion_full = [f'DRP^2.{i}' for i in range(1, 7)] + + wm_csf_basic = ['WM', 'CSF'] + wm_csf_power2 = ['WM^2', 'CSF^2'] + wm_csf_derivatives = ['DWM', 'DCSF'] + wm_csf_full = ['DWM^2', 'DCSF^2'] + + gs_basic = ['GS'] + gs_power2 = ['GS^2'] + gs_derivatives = ['DGS'] + gs_full = ['DGS^2'] + confound_column_names = [] - confound_column_names.append('FD') + confound_column_names.append('FD') # spike - for i in range(1, 7): - confound_column_names.append(f'RP.{i}') - confound_column_names.append(f'RP^2.{i}') - confound_column_names.append(f'DRP.{i}') - confound_column_names.append(f'DRP^2.{i}') + confound_column_names.extend(motion_basic) + confound_column_names.extend(motion_power2) + confound_column_names.extend(motion_derivatives) + confound_column_names.extend(motion_full) - for conf in ['WM', 'CSF', 'GS']: - confound_column_names.append(conf) - confound_column_names.append(f'{conf}^2') + confound_column_names.extend(wm_csf_basic) + confound_column_names.extend(wm_csf_power2) + confound_column_names.extend(wm_csf_derivatives) + confound_column_names.extend(wm_csf_full) + + confound_column_names.extend(gs_basic) + confound_column_names.extend(gs_power2) + confound_column_names.extend(gs_derivatives) + confound_column_names.extend(gs_full) # add some random irrelevant confounds for i in range(10): @@ -57,84 +79,262 @@ def test_felixconfoundremover(): columns=confound_column_names ) - # generate confound removal strategies with varying numbers of parameters - list_of_strategy_tuples = [] - - # 36 params - strat1 = { - 'motion': 'full', - 'wm_csf': 'full', - 'global_signal': 'full' + # Generate spec from Felix's column naming + spec = { + 'motion': { + 'basic': motion_basic, + 'power2': motion_basic + motion_power2, + 'derivatives': motion_basic + motion_derivatives, + 'full': + motion_basic + motion_derivatives + motion_power2 + motion_full + }, + 'wm_csf': { + 'basic': wm_csf_basic, + 'power2': wm_csf_basic + wm_csf_power2, + 'derivatives': wm_csf_basic + wm_csf_derivatives, + 'full': + wm_csf_basic + wm_csf_derivatives + wm_csf_power2 + wm_csf_full + }, + 'global_signal': { + 'basic': gs_basic, + 'power2': gs_basic + gs_power2, + 'derivatives': gs_basic + gs_derivatives, + 'full': gs_basic + gs_derivatives + gs_power2 + gs_full + } } - list_of_strategy_tuples.append((strat1, 36)) - - # 24 params - strat2 = { - 'motion': 'full', - } - list_of_strategy_tuples.append((strat2, 24)) - - # 9 params - strat3 = { - 'motion': 'basic', - 'wm_csf': 'basic', - 'global_signal': 'basic' - } - list_of_strategy_tuples.append((strat3, 9)) - - # 6 params - strat4 = { - 'motion': 'basic', - } - list_of_strategy_tuples.append((strat4, 6)) - - # 2 params - strat5 = { - 'wm_csf': 'basic', - } - list_of_strategy_tuples.append((strat5, 2)) # generate a junifer pipeline data object dictionary input_data_obj = {} input_data_obj['meta'] = {} input_data_obj['BOLD'] = {} input_data_obj['BOLD']['data'] = siimg - input_data_obj['confounds'] = confounds_df + input_data_obj['confounds'] = {} + input_data_obj['confounds']['path'] = Path('/test.df') + input_data_obj['confounds']['data'] = confounds_df + input_data_obj['confounds']['names'] = {} + input_data_obj['confounds']['names']['spec'] = spec + input_data_obj['confounds']['names']['spike'] = 'FD' - # run confound removal - for spike in [None, 0.25]: - for strategy, n_params in list_of_strategy_tuples: - FCR = FelixConfoundRemover( - strategy=strategy, - spike=spike, - mask_img=simsk, - t_r=0.75 - ) + # generate confound removal strategies with varying numbers of parameters - # first test some methods manually - FCR.validate_input(input_data_obj) - out_type = FCR.get_output_kind(input_data_obj) + # Test #1: 36 params, no derivatives to compute, no spike + # 36 params + strat1 = { + 'motion': 'full', + 'wm_csf': 'full', + 'global_signal': 'full' + } - assert out_type in ['BOLD'] + cr = BaseConfoundRemover( + strategy=strat1, spike=None, mask_img=simsk, t_r=0.75 + ) + cr.validate_input(input_data_obj.keys()) + out_type = cr.get_output_kind(input_data_obj.keys()) - selected_conf_frame = FCR.read_confounds(confounds_df) - timepoints, n_sel_confs = selected_conf_frame.shape + assert 'BOLD' in out_type - assert timepoints == 100 - assert isinstance(selected_conf_frame, pd.DataFrame) + # Check if the input data is valid + cr._validate_data(input_data_obj) - if spike is None: - assert n_sel_confs == n_params - else: - assert n_sel_confs == n_params + 1 + # Check that the confounds are picked correctly: + t_df = cr._pick_confounds(input_data_obj['confounds']) + assert len(t_df.columns) == 36 + assert all(x in t_df.columns for x in motion_basic) + assert all(x in t_df.columns for x in motion_power2) + assert all(x in t_df.columns for x in motion_derivatives) + assert all(x in t_df.columns for x in motion_full) + assert all(x in t_df.columns for x in wm_csf_basic) + assert all(x in t_df.columns for x in wm_csf_power2) + assert all(x in t_df.columns for x in wm_csf_derivatives) + assert all(x in t_df.columns for x in wm_csf_full) + assert all(x in t_df.columns for x in gs_basic) + assert all(x in t_df.columns for x in gs_power2) + assert all(x in t_df.columns for x in gs_derivatives) + assert all(x in t_df.columns for x in gs_full) + assert 'FD' not in t_df.columns + assert 'spike' not in t_df.columns - cl_bold = FCR.remove_confounds(siimg, selected_conf_frame) - check_niimg_4d(cl_bold) + # Test #2: 24 params, no derivatives to compute, no spike + # 24 params + strat2 = { + 'motion': 'full', + } - # finally test fit_transform - output_data_obj = FCR.fit_transform(input_data_obj) + cr = BaseConfoundRemover( + strategy=strat2, spike=None, mask_img=simsk, t_r=0.75 + ) + cr.validate_input(input_data_obj.keys()) + out_type = cr.get_output_kind(input_data_obj.keys()) - for key in ['BOLD', 'confounds', 'meta']: - assert key in output_data_obj + assert 'BOLD' in out_type - check_niimg_4d(output_data_obj['BOLD']['data']) + # Check if the input data is valid + cr._validate_data(input_data_obj) + + # Check that the confounds are picked correctly: + t_df = cr._pick_confounds(input_data_obj['confounds']) + assert len(t_df.columns) == 24 + assert all(x in t_df.columns for x in motion_basic) + assert all(x in t_df.columns for x in motion_power2) + assert all(x in t_df.columns for x in motion_derivatives) + assert all(x in t_df.columns for x in motion_full) + assert all(x not in t_df.columns for x in wm_csf_basic) + assert all(x not in t_df.columns for x in wm_csf_power2) + assert all(x not in t_df.columns for x in wm_csf_derivatives) + assert all(x not in t_df.columns for x in wm_csf_full) + assert all(x not in t_df.columns for x in gs_basic) + assert all(x not in t_df.columns for x in gs_power2) + assert all(x not in t_df.columns for x in gs_derivatives) + assert all(x not in t_df.columns for x in gs_full) + assert 'FD' not in t_df.columns + assert 'spike' not in t_df.columns + + # Test #3: 9 params, no derivatives to compute, no spike + strat3 = { + 'motion': 'basic', + 'wm_csf': 'basic', + 'global_signal': 'basic' + } + + cr = BaseConfoundRemover( + strategy=strat3, spike=None, mask_img=simsk, t_r=0.75 + ) + cr.validate_input(input_data_obj.keys()) + out_type = cr.get_output_kind(input_data_obj.keys()) + + assert 'BOLD' in out_type + + # Check if the input data is valid + cr._validate_data(input_data_obj) + + # Check that the confounds are picked correctly: + t_df = cr._pick_confounds(input_data_obj['confounds']) + assert len(t_df.columns) == 9 + assert all(x in t_df.columns for x in motion_basic) + assert all(x not in t_df.columns for x in motion_power2) + assert all(x not in t_df.columns for x in motion_derivatives) + assert all(x not in t_df.columns for x in motion_full) + assert all(x in t_df.columns for x in wm_csf_basic) + assert all(x not in t_df.columns for x in wm_csf_power2) + assert all(x not in t_df.columns for x in wm_csf_derivatives) + assert all(x not in t_df.columns for x in wm_csf_full) + assert all(x in t_df.columns for x in gs_basic) + assert all(x not in t_df.columns for x in gs_power2) + assert all(x not in t_df.columns for x in gs_derivatives) + assert all(x not in t_df.columns for x in gs_full) + assert 'FD' not in t_df.columns + assert 'spike' not in t_df.columns + + # Test #4: 6 params, no derivatives to compute, no spike + strat4 = { + 'motion': 'basic', + } + cr = BaseConfoundRemover( + strategy=strat4, spike=None, mask_img=simsk, t_r=0.75 + ) + cr.validate_input(input_data_obj.keys()) + out_type = cr.get_output_kind(input_data_obj.keys()) + + assert 'BOLD' in out_type + + # Check if the input data is valid + cr._validate_data(input_data_obj) + + # Check that the confounds are picked correctly: + t_df = cr._pick_confounds(input_data_obj['confounds']) + assert len(t_df.columns) == 6 + assert all(x in t_df.columns for x in motion_basic) + assert all(x not in t_df.columns for x in motion_power2) + assert all(x not in t_df.columns for x in motion_derivatives) + assert all(x not in t_df.columns for x in motion_full) + assert all(x not in t_df.columns for x in wm_csf_basic) + assert all(x not in t_df.columns for x in wm_csf_power2) + assert all(x not in t_df.columns for x in wm_csf_derivatives) + assert all(x not in t_df.columns for x in wm_csf_full) + assert all(x not in t_df.columns for x in gs_basic) + assert all(x not in t_df.columns for x in gs_power2) + assert all(x not in t_df.columns for x in gs_derivatives) + assert all(x not in t_df.columns for x in gs_full) + assert 'FD' not in t_df.columns + assert 'spike' not in t_df.columns + + # Test #5: 2 params, no derivatives to compute, no spike + strat5 = { + 'wm_csf': 'basic', + } + cr = BaseConfoundRemover( + strategy=strat5, spike=None, mask_img=simsk, t_r=0.75 + ) + cr.validate_input(input_data_obj.keys()) + out_type = cr.get_output_kind(input_data_obj.keys()) + + assert 'BOLD' in out_type + + # Check if the input data is valid + cr._validate_data(input_data_obj) + + # Check that the confounds are picked correctly: + t_df = cr._pick_confounds(input_data_obj['confounds']) + assert len(t_df.columns) == 2 + assert all(x not in t_df.columns for x in motion_basic) + assert all(x not in t_df.columns for x in motion_power2) + assert all(x not in t_df.columns for x in motion_derivatives) + assert all(x not in t_df.columns for x in motion_full) + assert all(x in t_df.columns for x in wm_csf_basic) + assert all(x not in t_df.columns for x in wm_csf_power2) + assert all(x not in t_df.columns for x in wm_csf_derivatives) + assert all(x not in t_df.columns for x in wm_csf_full) + assert all(x not in t_df.columns for x in gs_basic) + assert all(x not in t_df.columns for x in gs_power2) + assert all(x not in t_df.columns for x in gs_derivatives) + assert all(x not in t_df.columns for x in gs_full) + assert 'FD' not in t_df.columns + assert 'spike' not in t_df.columns + + out = cr.fit_transform(input_data_obj) + check_niimg_4d(out['BOLD']['data']) + # TODO: check meta + + # Test #6: 12 params, derivatives to compute, spike + to_select = [x for x in confounds_df.columns + if x not in motion_derivatives] + no_d_df = confounds_df[to_select] + input_data_obj['confounds']['data'] = no_d_df + + derivatives = { + f'D{x}': x for x in motion_basic + } + + input_data_obj['confounds']['names']['derivatives'] = derivatives + + strat6 = { + 'motion': 'derivatives', + } + cr = BaseConfoundRemover( + strategy=strat6, spike=0.75, mask_img=simsk, t_r=0.75 + ) + cr.validate_input(input_data_obj.keys()) + out_type = cr.get_output_kind(input_data_obj.keys()) + + assert 'BOLD' in out_type + + # Check if the input data is valid + cr._validate_data(input_data_obj) + + # Check that the confounds are picked correctly: + t_df = cr._pick_confounds(input_data_obj['confounds']) + assert len(t_df.columns) == 13 + assert all(x in t_df.columns for x in motion_basic) + assert all(x not in t_df.columns for x in motion_power2) + assert all(x in t_df.columns for x in motion_derivatives) + assert all(x not in t_df.columns for x in motion_full) + assert all(x not in t_df.columns for x in wm_csf_basic) + assert all(x not in t_df.columns for x in wm_csf_power2) + assert all(x not in t_df.columns for x in wm_csf_derivatives) + assert all(x not in t_df.columns for x in wm_csf_full) + assert all(x not in t_df.columns for x in gs_basic) + assert all(x not in t_df.columns for x in gs_power2) + assert all(x not in t_df.columns for x in gs_derivatives) + assert all(x not in t_df.columns for x in gs_full) + assert 'FD' not in t_df.columns + assert 'spike' in t_df.columns -- 2.52.0 From ffa75f0aba3c2d108edd2764394856221581100f Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Wed, 13 Apr 2022 13:03:45 +0300 Subject: [PATCH 064/287] Updated pattern datagrabber --- junifer/configs/juseless.py | 158 +----------------- junifer/configs/tests/test_juseless.py | 2 +- junifer/datagrabber/base.py | 72 +++++++- .../tests/test_base_datagrabber.py | 68 ++++++-- tools/create_bids_example_dataset_sessions.py | 41 +++++ 5 files changed, 166 insertions(+), 175 deletions(-) create mode 100644 tools/create_bids_example_dataset_sessions.py diff --git a/junifer/configs/juseless.py b/junifer/configs/juseless.py index 1918c5d82..be3b427f5 100644 --- a/junifer/configs/juseless.py +++ b/junifer/configs/juseless.py @@ -1,6 +1,5 @@ from ..datagrabber import PatternDataladDataGrabber from ..api.decorators import register_datagrabber -from ..utils import logger @register_datagrabber @@ -26,163 +25,8 @@ class JuselessDataladUKBVBM(PatternDataladDataGrabber): types = ['VBM_GM'] replacements = ['subject', 'session'] patterns = { - 'VBM_GM': 'm0wp1{subject}_{session}_T1w.nii.gz' + 'VBM_GM': 'm0wp1_sub-{subject}_ses-{session}_T1w.nii.gz' } super().__init__( types=types, datadir=datadir, uri=uri, rootdir=rootdir, replacements=replacements, patterns=patterns) - - def get_elements(self): - """Get the list of subjects in the dataset. - - Returns - ------- - elements : list[str] - The list of subjects in the dataset. - """ - logger.debug('Getting the list of subjects in the dataset') - elems = [] - for x in self.datadir.glob('*._T1w.nii.gz'): - sub, ses = x.name.split('_') - sub = sub.replace('m0wp1', '') - ses = ses[:5] - elems.append((sub, ses)) - return elems - - -# @register_datagrabber -# class HCP1200(PatternDataGrabber): -# """ Human Connectome Project Datalad DataGrabber class - -# Implements a DataGrabber to access the Human Connectome Project - -# """ - -# def __init__( -# self, datadir=None, subjects=None, tasks=None, phase_encodings=None -# ): -# """Initialize a HCP object. - -# Parameters -# ---------- -# datadir : str or Path -# That directory where the datalad dataset will be cloned. If None, -# (default), the datalad dataset will be cloned into a temporary -# directory. -# subjects : str or list of strings -# HCP subject ID's. If 'None' (default), all available subjects are -# selected -# tasks : str or list of strings -# HCP task sessions. If 'None' (default), all available task -# sessions are selected. Can be 'REST1', 'REST2', 'SOCIAL', 'WM', -# 'RELATIONAL', 'EMOTION', 'LANGUAGE', 'GAMBLING', 'MOTOR', or a -# list consisting of these names. -# phase_encoding : str or list of strings -# HCP phase encoding directions. Can be 'LR' or 'RL'. If 'None' -# (default) both will be used. - -# """ -# uri = ( -# 'https://github.com/datalad-datasets/' -# 'human-connectome-project-openaccess.git' -# ) -# rootdir = 'HCP1200' -# types = ['BOLD'] -# super().__init__( -# types=types, datadir=datadir, uri=uri, rootdir=rootdir -# ) - -# self.subjects = subjects -# self.tasks = tasks -# self.phase_encodings = phase_encodings - -# if isinstance(self.subjects, str): -# self.subjects = [self.subjects] -# if isinstance(self.tasks, str): -# self.tasks = [self.tasks] -# if isinstance(self.phase_encodings, str): -# self.phase_encodings = [self.phase_encodings] - -# if self.tasks is None: -# self.tasks = [ -# 'REST1', -# 'REST2', -# 'SOCIAL', -# 'WM', -# 'RELATIONAL', -# 'EMOTION', -# 'LANGUAGE', -# 'GAMBLING', -# 'MOTOR', -# ] - -# if self.phase_encodings is None: -# self.phase_encodings = ["LR", "RL"] - -# def get_elements(self): -# """Get the list of subjects in the dataset. - -# Returns -# ------- -# elements : list[str] -# The list of subjects in the dataset. -# """ -# elems = [] - -# if self.subjects is None: -# self.subjects = os.listdir(self.datadir) - -# for subject, task, phase_encoding in product( -# self.subjects, self.tasks, self.phase_encodings -# ): -# elems.append((subject, task, phase_encoding)) - -# return elems - -# def __getitem__(self, element): -# """Index one element in the dataset. - -# Parameters -# ---------- -# element : tuple[str, str] -# The element to be indexed. First element in the tuple is the -# subject, second element is the task, third element is the -# phase encoding direction. - -# Returns -# ------- -# out : dict[str -> Path] -# Dictionary of paths for each type of data required for the -# specified element. -# """ -# sub, task, phase_encoding = element -# out = {} - -# if 'REST' in task: -# task_name = f'rfMRI_{task}' -# else: -# task_name = f'tfMRI_{task}' - -# out['BOLD'] = dict( -# path=self.datadir / sub / 'MNINonLinear' / 'Results' / -# f'{task_name}_{phase_encoding}' / -# f'{task_name}_{phase_encoding}_hp2000_clean.nii.gz' -# ) - -# conf_dir = ( -# Path('/data') / 'group' / 'appliedml' / -# 'data' / 'HCP1200_Confounds_tsv' -# ) -# out['BOLD']['confounds'] = ( -# conf_dir / sub / 'MNINonLinear' / 'Results' / -# f'{task_name}_{phase_encoding}' / f'Confounds_{sub}.tsv' -# ) - -# out['meta'] = dict(datagrabber=self.get_meta()) -# self._dataset_get(out) - -# out['meta']['element'] = dict( -# subject=sub, task=task, phase_encoding=phase_encoding -# ) - -# return out diff --git a/junifer/configs/tests/test_juseless.py b/junifer/configs/tests/test_juseless.py index 894eaf793..87e1e863f 100644 --- a/junifer/configs/tests/test_juseless.py +++ b/junifer/configs/tests/test_juseless.py @@ -14,7 +14,7 @@ configure_logging(level='DEBUG') def test_juselessdataladukbvbm_datagrabber(): with JuselessDataladUKBVBM() as dg: - out = dg[('sub-2670511', 'ses-2')] + out = dg[('2670511', '2')] assert 'VBM_GM' in out assert out['VBM_GM'].name == 'm0wp1sub-2670511_ses-2_T1w.nii.gz' assert out['VBM_GM'].exists() diff --git a/junifer/datagrabber/base.py b/junifer/datagrabber/base.py index 1e8ac9c8c..431f7ca4d 100644 --- a/junifer/datagrabber/base.py +++ b/junifer/datagrabber/base.py @@ -2,6 +2,7 @@ # Leonard Sasse # License: AGPL from pathlib import Path +import re import tempfile import datalad.api as dl @@ -200,8 +201,73 @@ class PatternDataGrabber(BaseDataGrabber): self.patterns = patterns self.replacements = replacements - def _replace_patterns(self, element, pattern): - """Replace the patterns in the pattern with the element. + def _replace_patterns_regex(self, pattern): + """Replace the patterns in the pattern with the named groups so the + elements can be obtained from the filesystem. + + Parameters + ---------- + pattern : str + The pattern to be replaced. + + Returns + ------- + re_pattern : str + The regular expression with the named groups. + glob_pattern : str + The search pattern to be used with glob + + """ + re_pattern = pattern + glob_pattern = pattern + for t_r in self.replacements: + # Replace the first of each with a named group definition + re_pattern = re_pattern.replace( + f'{{{t_r}}}', f'(?P<{t_r}>.*)', 1) + + for t_r in self.replacements: + # Replace the second appearance of each with the named group + # back reference + re_pattern = re_pattern.replace(f'{{{t_r}}}', f'(?P={t_r})') + + for t_r in self.replacements: + glob_pattern = glob_pattern.replace(f'{{{t_r}}}', f'*') + return re_pattern, glob_pattern + + def get_elements(self): + """Get the list of elements in the dataset. It will use regex + to search for `replacements` in the `patterns` and return the + intersection of the results for each type. That is, build a list + of elements that have all the required types. + + Returns + ------- + elements : list + The list of elements in the dataset. + """ + elements = None + for t_type in self.types: + types_element = set() + t_pattern = self.patterns[t_type] # get the pattern + re_pattern, glob_pattern = self._replace_patterns_regex(t_pattern) + for fname in self.datadir.glob(glob_pattern): + suffix = fname.relative_to(self.datadir).as_posix() + m = re.match(re_pattern, suffix) + if m is not None: + t_element = tuple(m.group(k) for k in self.replacements) + if len(self.replacements) == 1: + t_element = t_element[0] + types_element.add(t_element) + if elements is None: + elements = types_element + else: + elements = elements.intersection(types_element) + + return list(elements) + + def _replace_patterns_glob(self, element, pattern): + """Replace the patterns in the pattern with the element so it can + be globbed. Parameters ---------- @@ -247,7 +313,7 @@ class PatternDataGrabber(BaseDataGrabber): element = (element,) for t_type in self.types: t_pattern = self.patterns[t_type] # type: ignore - t_replace = self._replace_patterns(element, t_pattern) + t_replace = self._replace_patterns_glob(element, t_pattern) if '*' in t_replace: t_matches = list(self.datadir.glob(t_replace)) if len(t_matches) > 1: diff --git a/junifer/datagrabber/tests/test_base_datagrabber.py b/junifer/datagrabber/tests/test_base_datagrabber.py index 2a4df25d7..689dd24a6 100644 --- a/junifer/datagrabber/tests/test_base_datagrabber.py +++ b/junifer/datagrabber/tests/test_base_datagrabber.py @@ -12,6 +12,10 @@ _testing_dataset = { 'example_bids': { 'uri': 'https://gin.g-node.org/juaml/datalad-example-bids', 'id': 'e2ce149bd723088769a86c72e57eded009258c6b' + }, + 'example_bids_ses': { + 'uri': 'https://gin.g-node.org/juaml/datalad-example-bids-ses', + 'id': '3d08d55d1faad4f12ab64ac9497544a0d924d47a' } } @@ -119,22 +123,17 @@ def test_bids_datalad_PatternDataGrabber(): } replacements = ['subject'] - class MyDataGrabber(PatternDataladDataGrabber): - def get_elements(self): - elems = [x.name for x in self.datadir.iterdir() if x.is_dir()] - return elems - with pytest.raises(ValueError, match=r"uri must be provided"): - MyDataGrabber(datadir=None, types=types, patterns=patterns, - replacements=replacements) + PatternDataladDataGrabber(datadir=None, types=types, patterns=patterns, + replacements=replacements) repo_uri = _testing_dataset['example_bids']['uri'] rootdir = 'example_bids' repo_commit = _testing_dataset['example_bids']['id'] - with MyDataGrabber(rootdir=rootdir, uri=repo_uri, - types=types, patterns=patterns, - replacements=replacements) as dg: + with PatternDataladDataGrabber( + rootdir=rootdir, uri=repo_uri, types=types, patterns=patterns, + replacements=replacements) as dg: subs = [x for x in dg] expected_subs = [f'sub-{i:02d}' for i in range(1, 10)] assert set(subs) == set(expected_subs) @@ -152,7 +151,7 @@ def test_bids_datalad_PatternDataGrabber(): assert 'datagrabber' in t_sub['meta'] dg_meta = t_sub['meta']['datagrabber'] assert 'class' in dg_meta - assert dg_meta['class'] == 'MyDataGrabber' + assert dg_meta['class'] == 'PatternDataladDataGrabber' assert 'uri' in dg_meta assert dg_meta['uri'] == repo_uri assert 'dataset_commit_id' in dg_meta @@ -167,9 +166,9 @@ def test_bids_datalad_PatternDataGrabber(): 'T1w': '{subject}/anat/{subject}_T*w.nii.gz', 'bold': '{subject}/func/{subject}_task-rest_*.nii.gz' } - with MyDataGrabber(rootdir=rootdir, uri=repo_uri, - types=types, patterns=patterns, - datadir=datadir, replacements=replacements) as dg: + with PatternDataladDataGrabber( + rootdir=rootdir, uri=repo_uri, types=types, patterns=patterns, + datadir=datadir, replacements=replacements) as dg: assert dg.datadir == datadir / rootdir for elem in dg: t_sub = dg[elem] @@ -179,3 +178,44 @@ def test_bids_datalad_PatternDataGrabber(): assert 'path' in t_sub['bold'] assert t_sub['bold']['path'] == \ (dg.datadir / f'{elem}/func/{elem}_task-rest_bold.nii.gz') + + +def test_bids_datalad_PatternDataGrabber_session(): + """Test a subject and session-based BIDS datalad datagrabber""" + types = ['T1w', 'bold'] + patterns = { + 'T1w': '{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz', + 'bold': '{subject}/{session}/func/' + '{subject}_{session}_task-rest_bold.nii.gz' + } + replacements = ['subject', 'session'] + + with pytest.raises(ValueError, match=r"uri must be provided"): + PatternDataladDataGrabber(datadir=None, types=types, patterns=patterns, + replacements=replacements) + + repo_uri = _testing_dataset['example_bids_ses']['uri'] + rootdir = 'example_bids_ses' + # repo_commit = _testing_dataset['example_bids_ses']['id'] + + # With T1W and bold, only 2 sessions are available + with PatternDataladDataGrabber( + rootdir=rootdir, uri=repo_uri, types=types, patterns=patterns, + replacements=replacements) as dg: + subs = [x for x in dg] + expected_subs = [(f'sub-{i:02d}', f'ses-{j:02d}') for j in range(1, 3) + for i in range(1, 10)] + assert set(subs) == set(expected_subs) + + # Test with a different T1w only, it should have 3 sessions + types = ['T1w'] + patterns = { + 'T1w': '{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz', + } + with PatternDataladDataGrabber( + rootdir=rootdir, uri=repo_uri, types=types, patterns=patterns, + replacements=replacements) as dg: + subs = [x for x in dg] + expected_subs = [(f'sub-{i:02d}', f'ses-{j:02d}') for j in range(1, 4) + for i in range(1, 10)] + assert set(subs) == set(expected_subs) diff --git a/tools/create_bids_example_dataset_sessions.py b/tools/create_bids_example_dataset_sessions.py new file mode 100644 index 000000000..3301c6f8a --- /dev/null +++ b/tools/create_bids_example_dataset_sessions.py @@ -0,0 +1,41 @@ +# Authors: Federico Raimondo +# License: AGPL +from tempfile import TemporaryDirectory +from pathlib import Path + +import datalad.api as dl + +dst = 'git@gin.g-node.org:/juaml/datalad-example-bids-ses.git' + +with TemporaryDirectory() as tmpdir_name: + tmpdir = Path(tmpdir_name) + ds = dl.create(tmpdir) # type: ignore + + base_dir = tmpdir / 'example_bids_ses' + base_dir.mkdir() + + for i_sub in range(1, 10): + t_sub = f'sub-{i_sub:02d}' + sub_dir = base_dir / t_sub + sub_dir.mkdir() + + for i_ses in range(1, 4): + t_ses = f'ses-{i_ses:02d}' + ses_dir = sub_dir / t_ses + ses_dir.mkdir() + + for dname in ['anat', 'func']: + (ses_dir / dname).mkdir() + + fnames = [f'anat/{t_sub}_{t_ses}_T1w.nii.gz'] + if i_ses != 3: # Session 3 does not have functional data + fnames.extend([ + f'func/{t_sub}_{t_ses}_task-rest_bold.nii.gz', + f'func/{t_sub}_{t_ses}_task-rest_bold.json']) + for fname in fnames: + with open(ses_dir / fname, 'w') as f: + f.write('placeholder') + + ds.save(recursive=True) + ds.siblings('add', name='gin', url=dst) + ds.push(to='gin', force='all') -- 2.52.0 From 9bc310af4ca380f0605991129ff6913c631dee2b Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Wed, 13 Apr 2022 13:06:25 +0300 Subject: [PATCH 065/287] Better juseless test --- junifer/configs/tests/test_juseless.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/junifer/configs/tests/test_juseless.py b/junifer/configs/tests/test_juseless.py index 87e1e863f..48cdcb275 100644 --- a/junifer/configs/tests/test_juseless.py +++ b/junifer/configs/tests/test_juseless.py @@ -14,9 +14,12 @@ configure_logging(level='DEBUG') def test_juselessdataladukbvbm_datagrabber(): with JuselessDataladUKBVBM() as dg: - out = dg[('2670511', '2')] + all_elements = dg.get_elements() + test_element = all_elements[0] + out = dg[test_element] assert 'VBM_GM' in out - assert out['VBM_GM'].name == 'm0wp1sub-2670511_ses-2_T1w.nii.gz' + assert out['VBM_GM'].name == \ + f'm0wp1sub-{test_element[0]}_ses-{test_element[1]}_T1w.nii.gz' assert out['VBM_GM'].exists() -- 2.52.0 From fc7cab2a306e5dd4680813176f91f4e6888f3c13 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Wed, 13 Apr 2022 12:17:56 +0200 Subject: [PATCH 066/287] JuselessUKBVBM data grabber tested --- junifer/configs/juseless.py | 2 +- junifer/configs/tests/test_juseless.py | 8 ++++---- 2 files changed, 5 insertions(+), 5 deletions(-) diff --git a/junifer/configs/juseless.py b/junifer/configs/juseless.py index be3b427f5..db88315b7 100644 --- a/junifer/configs/juseless.py +++ b/junifer/configs/juseless.py @@ -25,7 +25,7 @@ class JuselessDataladUKBVBM(PatternDataladDataGrabber): types = ['VBM_GM'] replacements = ['subject', 'session'] patterns = { - 'VBM_GM': 'm0wp1_sub-{subject}_ses-{session}_T1w.nii.gz' + 'VBM_GM': 'm0wp1sub-{subject}_ses-{session}_T1w.nii.gz' } super().__init__( types=types, datadir=datadir, uri=uri, rootdir=rootdir, diff --git a/junifer/configs/tests/test_juseless.py b/junifer/configs/tests/test_juseless.py index 48cdcb275..005173806 100644 --- a/junifer/configs/tests/test_juseless.py +++ b/junifer/configs/tests/test_juseless.py @@ -18,9 +18,9 @@ def test_juselessdataladukbvbm_datagrabber(): test_element = all_elements[0] out = dg[test_element] assert 'VBM_GM' in out - assert out['VBM_GM'].name == \ + assert out['VBM_GM']['path'].name == \ f'm0wp1sub-{test_element[0]}_ses-{test_element[1]}_T1w.nii.gz' - assert out['VBM_GM'].exists() + assert out['VBM_GM']['path'].exists() def test_juselessdataladhcp_datagrabber(): @@ -30,5 +30,5 @@ def test_juselessdataladhcp_datagrabber(): out = dg[test_element] - assert out['BOLD'].exists() - assert os.path.isfile(out['BOLD']['path']) + assert out['BOLD']['path'].exists() + assert out['BOLD']['path'].isfile() -- 2.52.0 From dad7ce69fd31d528fcf24deb9589d0bacee2a841 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Wed, 13 Apr 2022 12:19:19 +0200 Subject: [PATCH 067/287] Flake --- junifer/configs/tests/test_juseless.py | 1 + 1 file changed, 1 insertion(+) diff --git a/junifer/configs/tests/test_juseless.py b/junifer/configs/tests/test_juseless.py index 005173806..4fda2f849 100644 --- a/junifer/configs/tests/test_juseless.py +++ b/junifer/configs/tests/test_juseless.py @@ -23,6 +23,7 @@ def test_juselessdataladukbvbm_datagrabber(): assert out['VBM_GM']['path'].exists() + def test_juselessdataladhcp_datagrabber(): with DataladHCP1200() as dg: all_elements = dg.get_elements() -- 2.52.0 From dd0f8192d98cdc94767a7f3e156dddb6bbd2f30f Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Wed, 13 Apr 2022 12:19:50 +0200 Subject: [PATCH 068/287] flake again --- junifer/configs/tests/test_juseless.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/junifer/configs/tests/test_juseless.py b/junifer/configs/tests/test_juseless.py index 4fda2f849..9a266c19c 100644 --- a/junifer/configs/tests/test_juseless.py +++ b/junifer/configs/tests/test_juseless.py @@ -1,6 +1,5 @@ import socket import pytest -import os from junifer.configs.juseless import JuselessDataladUKBVBM from junifer.utils.logging import configure_logging @@ -23,7 +22,6 @@ def test_juselessdataladukbvbm_datagrabber(): assert out['VBM_GM']['path'].exists() - def test_juselessdataladhcp_datagrabber(): with DataladHCP1200() as dg: all_elements = dg.get_elements() -- 2.52.0 From cfd8715bab199cadbb4fb047268d7c0d399cb1fc Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Thu, 14 Apr 2022 09:38:35 +0300 Subject: [PATCH 069/287] Use kind in name for storage --- junifer/markers/base.py | 6 ++++-- junifer/markers/tests/test_base_marker.py | 8 ++++---- junifer/markers/tests/test_parcel.py | 8 ++++---- 3 files changed, 12 insertions(+), 10 deletions(-) diff --git a/junifer/markers/base.py b/junifer/markers/base.py index 0471744d0..0cfec4fc6 100644 --- a/junifer/markers/base.py +++ b/junifer/markers/base.py @@ -82,8 +82,10 @@ class BaseMarker(PipelineStepMixin): def get_meta(self, kind): s_meta = super().get_meta() - s_meta['name'] = self.name - s_meta['kind'] = kind # same marker can be fit into different kinds + # same marker can be fit into different kinds, so the name + # is created from the kind and the name of the marker + s_meta['name'] = f'{kind}_{self.name}' + s_meta['kind'] = kind return dict(marker=s_meta) def validate_input(self, input): diff --git a/junifer/markers/tests/test_base_marker.py b/junifer/markers/tests/test_base_marker.py index f661af6df..a0d43c81c 100644 --- a/junifer/markers/tests/test_base_marker.py +++ b/junifer/markers/tests/test_base_marker.py @@ -22,12 +22,12 @@ def test_meta(): t_meta = base.get_meta('bold') assert t_meta['marker']['class'] == 'BaseMarker' - assert t_meta['marker']['name'] == 'BaseMarker' + assert t_meta['marker']['name'] == 'bold_BaseMarker' base = BaseMarker(on=['bold', 'dwi'], name='mymarker') t_meta = base.get_meta('dwi') - assert t_meta['marker']['name'] == 'mymarker' + assert t_meta['marker']['name'] == 'dwi_mymarker' def test_BaseMarker(): @@ -50,12 +50,12 @@ def test_BaseMarker(): out = base.fit_transform(input) assert out['bold']['data'] == 1 - assert out['bold']['meta']['marker']['name'] == 'mymarker' + assert out['bold']['meta']['marker']['name'] == 'bold_mymarker' assert out['bold']['meta']['marker']['class'] == 'BaseMarker' base2 = BaseMarker(on='bold', name='mymarker') base2.compute = lambda x: dict(data=1) # type: ignore out2 = base2.fit_transform(input) assert out2['bold']['data'] == 1 - assert out2['bold']['meta']['marker']['name'] == 'mymarker' + assert out2['bold']['meta']['marker']['name'] == 'bold_mymarker' assert out2['bold']['meta']['marker']['class'] == 'BaseMarker' diff --git a/junifer/markers/tests/test_parcel.py b/junifer/markers/tests/test_parcel.py index fae7313ab..5098f9b01 100644 --- a/junifer/markers/tests/test_parcel.py +++ b/junifer/markers/tests/test_parcel.py @@ -64,7 +64,7 @@ def test_ParcelAggregation_3D(): meta = marker.get_meta('VBM_GM')['marker'] assert meta['method'] == 'mean' assert meta['atlas'] == 'Schaefer100x7' - assert meta['name'] == 'gmd_schaefer100x7_mean' + assert meta['name'] == 'VBM_GM_gmd_schaefer100x7_mean' assert meta['class'] == 'ParcelAggregation' assert meta['kind'] == 'VBM_GM' assert meta['method_params'] == {} @@ -88,7 +88,7 @@ def test_ParcelAggregation_3D(): meta = marker.get_meta('VBM_GM')['marker'] assert meta['method'] == 'std' assert meta['atlas'] == 'Schaefer100x7' - assert meta['name'] == 'ParcelAggregation' + assert meta['name'] == 'VBM_GM_ParcelAggregation' assert meta['class'] == 'ParcelAggregation' assert meta['kind'] == 'VBM_GM' assert meta['method_params'] == {} @@ -116,7 +116,7 @@ def test_ParcelAggregation_3D(): meta = marker.get_meta('VBM_GM')['marker'] assert meta['method'] == 'trim_mean' assert meta['atlas'] == 'Schaefer100x7' - assert meta['name'] == 'ParcelAggregation' + assert meta['name'] == 'VBM_GM_ParcelAggregation' assert meta['class'] == 'ParcelAggregation' assert meta['kind'] == 'VBM_GM' assert meta['method_params'] == {'proportiontocut': 0.1} @@ -147,7 +147,7 @@ def test_ParcelAggregation_4D(): meta = marker.get_meta('BOLD')['marker'] assert meta['method'] == 'mean' assert meta['atlas'] == 'Schaefer100x7' - assert meta['name'] == 'ParcelAggregation' + assert meta['name'] == 'BOLD_ParcelAggregation' assert meta['class'] == 'ParcelAggregation' assert meta['kind'] == 'BOLD' assert meta['method_params'] == {} -- 2.52.0 From 3cec900f80a147baf56f7738947303bdccb18dca Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Thu, 14 Apr 2022 11:28:43 +0300 Subject: [PATCH 070/287] Add tests + simplify logic in retrieve_tian --- junifer/data/atlases.py | 113 ++++++++++++++----------------- junifer/data/tests/test_atlas.py | 68 ++++++++++++++++++- 2 files changed, 118 insertions(+), 63 deletions(-) diff --git a/junifer/data/atlases.py b/junifer/data/atlases.py index a45fa9943..5d92f2463 100644 --- a/junifer/data/atlases.py +++ b/junifer/data/atlases.py @@ -3,9 +3,8 @@ # License: AGPL from pathlib import Path import io -import os +import tempfile import requests -import wget import shutil import zipfile import numpy as np @@ -49,34 +48,27 @@ for n_rois in range(100, 1001, 100): } for scale in range(1, 5): - for field in ['3T', '7T']: - if field == '7T': - space = 'MNI6thgeneration' - t_name = f'Tian{scale}x{field}x{space}' - _available_atlases[t_name] = { - 'family': 'Tian', - 'scale': scale, - 'magneticfield': field, - 'valid_resolutions': [1.6] - } - else: - for space in ['MNI6thgeneration', 'MNInonlinear2009cAsym']: - if space == 'MNI6thgeneration': - t_name = f'Tian{scale}x{field}x{space}' - _available_atlases[t_name] = { - 'family': 'Tian', - 'scale': scale, - 'magneticfield': field, - 'valid_resolutions': [1, 2] - } - elif space == 'MNInonlinear2009cAsym': - t_name = f'Tian{scale}x{field}x{space}' - _available_atlases[t_name] = { - 'family': 'Tian', - 'scale': scale, - 'magneticfield': field, - 'valid_resolutions': [2] - } + t_name = f'TianxS{scale}x7TxMNI6thgeneration' + _available_atlases[t_name] = { + 'family': 'Tian', + 'scale': scale, + 'magneticfield': '7T', + 'space': 'MNI6thgeneration' + } + t_name = f'TianxS{scale}x3TxMNI6thgeneration' + _available_atlases[t_name] = { + 'family': 'Tian', + 'scale': scale, + 'magneticfield': '3T', + 'space': 'MNI6thgeneration' + } + t_name = f'TianxS{scale}x3TxMNInonlinear2009cAsym' + _available_atlases[t_name] = { + 'family': 'Tian', + 'scale': scale, + 'magneticfield': '3T', + 'space': 'MNInonlinear2009cAsym' + } def register_atlas(name, atlas_path, atl_labels, overwrite=False): @@ -255,7 +247,7 @@ def _retrieve_atlas(family, atlas_dir=None, resolution=None, **kwargs): space (optional) : str Space of atlas can be either 'MNI6thgeneration' or 'MNInonlinear2009cAsym' (for some cases). - Defaults to 'MNI6thgeneration'. (For more information see + Defaults to 'MNI6thgeneration'. (For more information see https://github.com/yetianmed/subcortex) magneticfield (optional) : str Options are 3T and 7T, defaults to 3T. @@ -370,7 +362,12 @@ def _retrieve_schaefer(atlas_dir, resolution, n_rois=None, yeo_networks=7): def _retrieve_tian( atlas_dir, resolution, scale=None, space='MNI6thgeneration', magneticfield='3T'): - + # 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}') # check validity of atlas parameters _valid_scales = [1, 2, 3, 4] _valid_fields = ['3T', '7T'] @@ -378,18 +375,18 @@ def _retrieve_tian( raise_error( f'The parameter `scale` ({scale}) needs to be one of the ' f'following: {_valid_scales}') - if field not in _valid_fields: + if magneticfield not in _valid_fields: raise_error( - f'The parameter `magneticfield` ({field}) needs to be one of ' - f'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': _valid_resolutions = [1, 2] - elif space == 'MNInonlinear2009cAsym': + else: # space == 'MNInonlinear2009cAsym': _valid_resolutions = [2] - if magneticfield == '7T': + else: # magneticfield == '7T': _valid_spaces = ['MNI6thgeneration'] _valid_resolutions = [1.6] if space not in _valid_spaces: @@ -399,13 +396,6 @@ def _retrieve_tian( resolution = _closest_resolution(resolution, _valid_resolutions) - # 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}') - # define file names if magneticfield == '3T': atlas_fname_base_3T = ( @@ -416,14 +406,13 @@ def _retrieve_tian( atlas_fname = atlas_fname_base_3T / ( f'Tian_Subcortex_S{scale}_{magneticfield}.nii.gz') if resolution == 1: - atlas_fname = atlas_fname_base_3T / ( - f'Tian_Subcortex_S{scale}_{magneticfield}_{resolution}' - 'mm.nii.gz') - elif space == 'MNInonlinear2009cAsym': + 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}_{space}.nii.gz') - elif magneticfield == '7T': + else: # magneticfield == '7T': atlas_fname_base_7T = ( atlas_dir / 'Tian2020MSA_v1.1' / '7T') atlas_fname = atlas_dir / 'Tian2020MSA_v1.1' / f'{magneticfield}' / ( @@ -431,7 +420,7 @@ def _retrieve_tian( # 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: @@ -447,18 +436,20 @@ def _retrieve_tian( '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 {url_basis}') - atlas_download_dir = wget.download(url_basis, atlas_dir.as_posix()) - with zipfile.ZipFile(atlas_download_dir, 'r') as zip_ref: - zip_ref.extractall(atlas_dir.as_posix()) - # clean after unzipping - if os.path.exists(atlas_download_dir): - os.remove(atlas_download_dir) - if os.path.exists((atlas_dir / '__MACOSX').as_posix()): - shutil.rmtree((atlas_dir / '__MACOSX').as_posix()) + 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: + f.write(atlas_download.content) + 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()) labels = pd.read_csv(atlas_lname, sep=" ", header=None)[0].to_list() diff --git a/junifer/data/tests/test_atlas.py b/junifer/data/tests/test_atlas.py index c02b1b42b..e42c02228 100644 --- a/junifer/data/tests/test_atlas.py +++ b/junifer/data/tests/test_atlas.py @@ -1,11 +1,11 @@ import tempfile import pytest from pathlib import Path -from numpy.testing import assert_array_equal +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_atlas, _retrieve_tian) def test_register_atlas(): @@ -150,3 +150,67 @@ def test_suit(): 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') -- 2.52.0 From a6f927bfbece96ea2fed0e59b414939ef1c6f866 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 15 Apr 2022 09:49:49 +0200 Subject: [PATCH 071/287] Add support for fALFF, GCOR and LCOR in ParcelAggregation --- junifer/markers/parcel.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/junifer/markers/parcel.py b/junifer/markers/parcel.py index 6ede62d7e..9615337a0 100644 --- a/junifer/markers/parcel.py +++ b/junifer/markers/parcel.py @@ -15,7 +15,7 @@ from ..api.decorators import register_marker class ParcelAggregation(BaseMarker): def __init__(self, atlas, method, method_params=None, on=None, name=None): if on is None: - on = ['T1w', 'BOLD', 'VBM_GM', 'VBM_WM'] + on = ['T1w', 'BOLD', 'VBM_GM', 'VBM_WM', 'fALFF', 'GCOR', 'LCOR'] super().__init__(on=on, name=name) self.atlas = atlas self.method = method -- 2.52.0 From 88692ca71ffce3c222bb6fc73c6da87e636a0010 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 15 Apr 2022 11:42:29 +0200 Subject: [PATCH 072/287] Update requirements + debug + various fixes --- junifer/datagrabber/base.py | 11 ++++++----- junifer/datareader/default.py | 20 +++++++++++++++----- junifer/markers/base.py | 2 ++ junifer/markers/collection.py | 2 +- junifer/markers/parcel.py | 10 +++++++--- junifer/storage/base.py | 5 +++++ requirements.txt | 2 +- 7 files changed, 37 insertions(+), 15 deletions(-) diff --git a/junifer/datagrabber/base.py b/junifer/datagrabber/base.py index 431f7ca4d..fb86fa658 100644 --- a/junifer/datagrabber/base.py +++ b/junifer/datagrabber/base.py @@ -400,7 +400,7 @@ class DataladDataGrabber(BaseDataGrabber): def install(self): """Install the datalad dataset into the datadir.""" logger.debug(f'Installing dataset {self.uri} to {self._datadir}') - self.dataset = dl.install( # type: ignore + self._dataset = dl.install( # type: ignore self._datadir, source=self.uri) logger.debug('Dataset installed') @@ -411,19 +411,19 @@ class DataladDataGrabber(BaseDataGrabber): def remove(self): """Remove the datalad dataset from the datadir.""" - self.dataset.remove(recursive=True) + self._dataset.remove(recursive=True) def _dataset_get(self, out): for _, v in out.items(): if 'path' in v: logger.debug(f'Getting {v["path"]}') - self.dataset.get(v['path']) + self._dataset.get(v['path']) logger.debug(f'Get done') # append the version of the dataset out['meta']['datagrabber']['dataset_commit_id'] = \ - self.dataset.repo.get_hexsha( - self.dataset.repo.get_corresponding_branch()) + self._dataset.repo.get_hexsha( + self._dataset.repo.get_corresponding_branch()) return out def __getitem__(self, element): @@ -439,6 +439,7 @@ class DataladDataGrabber(BaseDataGrabber): return out +@register_datagrabber class PatternDataladDataGrabber(DataladDataGrabber, PatternDataGrabber): """Pattern-based Datalad DataGrabber class (abstract). Implements a DataGrabber that gets data from a datalad sibling, diff --git a/junifer/datareader/default.py b/junifer/datareader/default.py index aa71c7e63..b1a4acf83 100644 --- a/junifer/datareader/default.py +++ b/junifer/datareader/default.py @@ -6,7 +6,7 @@ from pathlib import Path import nibabel as nib import pandas as pd -from ..utils.logging import logger +from ..utils.logging import logger, warn from ..markers.base import PipelineStepMixin # Map each filenanm end to a kind @@ -37,24 +37,34 @@ class DefaultDataReader(PipelineStepMixin): def fit_transform(self, input, params=None): # For each kind of data, try to read it - out = {} + + # out is the same, but with the 'data' key set in + # each kind dictionary, except for meta + out = input if params is None: params = {} for kind in input.keys(): if kind == 'meta': out['meta'] = input['meta'] continue - t_path = input[kind] + if 'path' not in input[kind]: + warn( + f'Input kind {kind} does not provide a path. Skipping.') + continue + t_path = input[kind]['path'] t_params = params.get(kind, {}) + + # Convert to Path if datareader is not well done if not isinstance(t_path, Path): t_path = Path(t_path) - out[kind] = {'path': t_path} + out[kind]['path'] = t_path + logger.info(f'Reading {kind} from {t_path.as_posix()}') fread = None fname = t_path.name.lower() for ext, ftype in _extensions.items(): if fname.endswith(ext): - logger.info(f'Reading {ftype} file {t_path.as_posix()}') + logger.info(f'{kind} is type {ftype}') reader_func = _readers[ftype]['func'] reader_params = _readers[ftype]['params'] if reader_params is not None: diff --git a/junifer/markers/base.py b/junifer/markers/base.py index 0cfec4fc6..bcec81d9f 100644 --- a/junifer/markers/base.py +++ b/junifer/markers/base.py @@ -117,8 +117,10 @@ class BaseMarker(PipelineStepMixin): t_out = self.compute(t_input) t_out.update(meta=t_meta) if storage is not None: + logger.info(f'Storing in {storage}') self.store(kind, t_out, storage) else: + logger.info('No storage specified, returning dictionary') out[kind] = t_out return out diff --git a/junifer/markers/collection.py b/junifer/markers/collection.py index 6bb58c450..fefd2d6ec 100644 --- a/junifer/markers/collection.py +++ b/junifer/markers/collection.py @@ -52,7 +52,7 @@ class MarkerCollection(): m_value = marker.fit_transform(data, storage=self._storage) if self._storage is None: out[marker.name] = m_value - + logger.info(f'Marker collection fitting done') return None if self._storage else out def validate(self, datagrabber): diff --git a/junifer/markers/parcel.py b/junifer/markers/parcel.py index 9615337a0..defc9bb3d 100644 --- a/junifer/markers/parcel.py +++ b/junifer/markers/parcel.py @@ -9,6 +9,7 @@ from .base import BaseMarker from ..stats import get_aggfunc_by_name from ..data import load_atlas from ..api.decorators import register_marker +from ..utils import logger @register_marker @@ -22,19 +23,21 @@ class ParcelAggregation(BaseMarker): self.method_params = {} if method_params is None else method_params def get_output_kind(self, input): - if input in ['VBM_GM', 'VBM_WM']: + if input in ['VBM_GM', 'VBM_WM', 'fALFF', 'GCOR', 'LCOR']: return 'table' if input in ['BOLD']: return 'timeseries' def store(self, kind, out, storage): - if kind in ['VBM_GM', 'VBM_WM']: + logger.debug(f'Storing {kind} in {storage}') + if kind in ['VBM_GM', 'VBM_WM', 'fALFF', 'GCOR', 'LCOR']: storage.store_table(**out) if kind in ['BOLD']: storage.store_timeseries(**out) def compute(self, input): t_input = input['data'] + logger.debug(f'Parcel aggregation using {self.method}') agg_func = get_aggfunc_by_name( self.method, func_params=self.method_params) # Get the min of the voxels sizes and use it as the resolution @@ -50,7 +53,7 @@ class ParcelAggregation(BaseMarker): 'img != 0', img=atlas_img_res, ) - + logger.debug('Masking') masker = NiftiMasker(atlas_bin, target_affine=t_input.affine) # Mask the input data and the atlas @@ -59,6 +62,7 @@ class ParcelAggregation(BaseMarker): atlas_values = np.squeeze(atlas_values).astype(int) # Get the values for each parcel and apply agg function + logger.debug('Computing ROI means') atlas_roi_vals = sorted(np.unique(atlas_values)) out_labels = [] out_values = [] diff --git a/junifer/storage/base.py b/junifer/storage/base.py index dbdf0f204..68a83244e 100644 --- a/junifer/storage/base.py +++ b/junifer/storage/base.py @@ -212,6 +212,11 @@ class BaseFeatureStorage(ABC): def collect(self): raise NotImplementedError('collect not implemented') + def __str__(self): + single = '(single output)' \ + if self.single_output is True else '(multiple output)' + return f'<{self.__class__.__name__} @ {self.uri} {single}>' + class PandasFeatureStoreage(BaseFeatureStorage): diff --git a/requirements.txt b/requirements.txt index 2eb1a7062..a91a82997 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,6 +1,6 @@ numpy>=1.20, <1.22 datalad>=0.15.4, <0.16 -pandas>=0.18.0, <1.5 +pandas>=1.4.0, <1.5 nibabel>=3.2.0, <4.0 nilearn>=0.9.0, <1.0 sqlalchemy>=1.4.27, <= 1.5.0 -- 2.52.0 From 3faa45e961bb3a3e532d81f380447b1fd6e64f04 Mon Sep 17 00:00:00 2001 From: LeSasse Date: Fri, 15 Apr 2022 13:49:39 +0200 Subject: [PATCH 073/287] adding code to compute squares if needed for _pick_confounds --- junifer/preprocess/confounds.py | 10 ++++++++++ 1 file changed, 10 insertions(+) diff --git a/junifer/preprocess/confounds.py b/junifer/preprocess/confounds.py index 0d56f818a..7026cab75 100644 --- a/junifer/preprocess/confounds.py +++ b/junifer/preprocess/confounds.py @@ -161,7 +161,10 @@ class BaseConfoundRemover(PipelineStepMixin): to_select = [] confounds_df = input['data'] confounds_spec = input['names']['spec'] + # for every confound there is a derivative + # and for every confound + derivative there should be squares derivatives_to_compute = input['names'].get('derivatives', {}) + squares_to_compute = input['names'].get('squares', {}) spike_name = input['names']['spike'] # Get all the column names according to the strategy @@ -175,7 +178,14 @@ class BaseConfoundRemover(PipelineStepMixin): for t_dst, t_src in derivatives_to_compute.items(): out_df[t_dst] = np.append( # type: ignore np.diff(out_df[t_src]), 0) # type: ignore + + # Add squares (of base confounds and derivatives) if needed + to_compute = [x in squares_to_compute.keys() for x in to_select] + if any(to_compute): + for t_dst, t_src in derivatives_to_compute.items(): + out_df[t_dst] = out_df[t_src] ** 2 out_df = out_df[to_select] + # add binary spike regressor if needed at given threshold if self.spike is not None: fd = confounds_df[spike_name].copy() -- 2.52.0 From 3cd0b7bbffc69432411ed766e338c26505cdaeb2 Mon Sep 17 00:00:00 2001 From: LeSasse Date: Fri, 15 Apr 2022 13:59:49 +0200 Subject: [PATCH 074/287] tricky typo --- junifer/preprocess/confounds.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/junifer/preprocess/confounds.py b/junifer/preprocess/confounds.py index 7026cab75..f1e3f20aa 100644 --- a/junifer/preprocess/confounds.py +++ b/junifer/preprocess/confounds.py @@ -182,7 +182,7 @@ class BaseConfoundRemover(PipelineStepMixin): # Add squares (of base confounds and derivatives) if needed to_compute = [x in squares_to_compute.keys() for x in to_select] if any(to_compute): - for t_dst, t_src in derivatives_to_compute.items(): + for t_dst, t_src in squares_to_compute.items(): out_df[t_dst] = out_df[t_src] ** 2 out_df = out_df[to_select] -- 2.52.0 From 174de449eb2baf0426756e631cc7eccd5ae19b91 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Wed, 27 Apr 2022 12:43:45 +0200 Subject: [PATCH 075/287] Update logging #8 + #6 WIP --- junifer/storage/sqlite.py | 14 ++++++++++---- junifer/utils/logging.py | 2 +- 2 files changed, 11 insertions(+), 5 deletions(-) diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py index 82d5b0590..f398f895e 100644 --- a/junifer/storage/sqlite.py +++ b/junifer/storage/sqlite.py @@ -5,6 +5,7 @@ import pandas as pd from pandas.core.base import NoNewAttributesMixin from pandas.io.sql import pandasSQL_builder from sqlalchemy import create_engine, inspect +import tqdm from ..api.decorators import register_storage from .base import (PandasFeatureStoreage, process_meta, element_to_prefix, @@ -191,7 +192,8 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): out_storage = SQLiteFeatureStorage( uri=self.uri, single_output=True, upsert='ignore') - for elem in self.uri.parent.glob(f'*{self.uri.name}'): + files = self.uri.parent.glob(f'*{self.uri.name}') + for elem in tqdm.tqdm(files, desc='file'): logger.debug(f'Reading from {elem.as_posix()}') in_storage = SQLiteFeatureStorage(uri=elem, single_output=True) in_engine = in_storage.get_engine() @@ -199,13 +201,13 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): t_meta_df = pd.read_sql( 'meta', con=in_engine, index_col='meta_md5') out_storage._save_upsert(t_meta_df, 'meta') - for meta_md5 in t_meta_df.index: + for meta_md5 in tqdm.tqdm(t_meta_df.index, desc='feature'): logger.debug(f'Collecting feature {meta_md5}') # TODO: Fix this, needs that read_feature sets the index # properly table_name = f'meta_{meta_md5}' t_df = in_storage.read_df(feature_md5=meta_md5) - out_storage._save_upsert(t_df, table_name) + out_storage._save_upsert(t_df, table_name, if_exist='nocheck') def _save_upsert(self, df, name, engine=None, if_exist='append'): """ Implementation of UPSERT functionality. @@ -237,8 +239,12 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): elif not inspect(engine).has_table(name): # Case 2: new table, so no big issue df.to_sql(name, con=con, if_exists='append') + elif if_exist == 'nocheck': + # Case 3: existing table, but we will not check for + # existing + df.to_sql(name, con=con, if_exists='append') else: - # Case 3: existing table, so we need to check if the index + # Case 4: existing table, so we need to check if the index # is present or not. if if_exist == 'fail': raise ValueError(f"Table ({name}) already exists") diff --git a/junifer/utils/logging.py b/junifer/utils/logging.py index 343ae4954..b9ef4471b 100644 --- a/junifer/utils/logging.py +++ b/junifer/utils/logging.py @@ -8,7 +8,7 @@ from pathlib import Path import warnings -logging.basicConfig(stream=sys.stdout, level=logging.WARN) +# logging.basicConfig(stream=sys.stdout, level=logging.WARN) logger = logging.getLogger('JUNIFER') -- 2.52.0 From ad8d44fe5feb5058f4d94b83027b1e4fe7bc9230 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 20 Jul 2022 16:20:12 +0200 Subject: [PATCH 076/287] update: add initial tox support for flake8 --- tox.ini | 31 +++++++++++++++++++++++++++++++ 1 file changed, 31 insertions(+) create mode 100644 tox.ini diff --git a/tox.ini b/tox.ini new file mode 100644 index 000000000..a39a7eb43 --- /dev/null +++ b/tox.ini @@ -0,0 +1,31 @@ +[tox] +envlist = flake8 +isolated_build = true + +[testenv:flake8] +skip_install = true +deps = + flake8 + flake8-docstrings + flake8-bugbear +commands = + flake8 {toxinidir}/junifer {toxinidir}/examples {toxinidir}/scratch {toxinidir}/tools {toxinidir}/setup.py + +[flake8] +exclude = + __init__.py +max-line-length = 79 +ignore = + D202 + E201 # whitespace after ‘(’ + E202 # whitespace before ‘)’ + E203 # whitespace before ‘,’, ‘;’, or ‘:’ + E221 # multiple spaces before operator + E222 # multiple spaces after operator + E241 # multiple spaces after ‘,’ + I100 + I101 + I201 + N806 + W503 # line break before binary operator + W504 # line break after binary operator -- 2.52.0 From d570f480c165a9893e65d862b6d079f9826e40c2 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 20 Jul 2022 17:21:33 +0200 Subject: [PATCH 077/287] update: add pytest support in tox --- tox.ini | 23 ++++++++++++++++++++++- 1 file changed, 22 insertions(+), 1 deletion(-) diff --git a/tox.ini b/tox.ini index a39a7eb43..ef3266cab 100644 --- a/tox.ini +++ b/tox.ini @@ -1,5 +1,5 @@ [tox] -envlist = flake8 +envlist = flake8, test isolated_build = true [testenv:flake8] @@ -11,6 +11,14 @@ deps = commands = flake8 {toxinidir}/junifer {toxinidir}/examples {toxinidir}/scratch {toxinidir}/tools {toxinidir}/setup.py +[testenv:test] +skip_install = false +deps = + pytest +commands = + pytest + + [flake8] exclude = __init__.py @@ -29,3 +37,16 @@ ignore = N806 W503 # line break before binary operator W504 # line break after binary operator + +[pytest] +testpaths = + junifer/api/tests + junifer/configs/tests + junifer/data/tests + junifer/datagrabber/tests + junifer/datareader/tests + junifer/markers/tests + junifer/preprocess/tests + junifer/storage/tests + junifer/testing + junifer/utils/tests -- 2.52.0 From 85ec07fae4d2b91900b0ce1eb2694c4e471d6a43 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 20 Jul 2022 18:44:48 +0200 Subject: [PATCH 078/287] update: add coverage support in tox --- tox.ini | 35 ++++++++++++++++++++++++++++++++++- 1 file changed, 34 insertions(+), 1 deletion(-) diff --git a/tox.ini b/tox.ini index ef3266cab..bfd3bea38 100644 --- a/tox.ini +++ b/tox.ini @@ -1,5 +1,5 @@ [tox] -envlist = flake8, test +envlist = flake8, test, coverage isolated_build = true [testenv:flake8] @@ -18,6 +18,17 @@ deps = commands = pytest +[testenv:coverage] +skip_install = false +deps = + pytest + pytest-cov +commands = + pytest --cov=junifer --cov-report=xml + +################ +# Tool configs # +################ [flake8] exclude = @@ -50,3 +61,25 @@ testpaths = junifer/storage/tests junifer/testing junifer/utils/tests + +[coverage:paths] +source = + junifer + */site-packages/junifer + +[coverage:run] +branch = true +omit = + */setup.py + */tests/* + junifer/configs/juseless.py + junifer/testing/* +parallel = false + +[coverage:report] +exclude_lines = + # Have to re-enable the standard pragma + pragma: no cover + # Don't complain if non-runnable code isn't run: + if __name__ == .__main__.: +precision = 2 \ No newline at end of file -- 2.52.0 From 5e79651dea5aee59e1bf3d41b3665d004b550f21 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 20 Jul 2022 19:00:14 +0200 Subject: [PATCH 079/287] update: add support for codespell in tox --- tox.ini | 19 +++++++++++++++++-- 1 file changed, 17 insertions(+), 2 deletions(-) diff --git a/tox.ini b/tox.ini index bfd3bea38..3e200f0f9 100644 --- a/tox.ini +++ b/tox.ini @@ -1,5 +1,5 @@ [tox] -envlist = flake8, test, coverage +envlist = flake8, test, coverage, codespell isolated_build = true [testenv:flake8] @@ -26,6 +26,13 @@ deps = commands = pytest --cov=junifer --cov-report=xml +[testenv:codespell] +skip_install = true +deps = + codespell +commands = + codespell --config tox.ini examples/ junifer/ scratch/ tools/ + ################ # Tool configs # ################ @@ -82,4 +89,12 @@ exclude_lines = pragma: no cover # Don't complain if non-runnable code isn't run: if __name__ == .__main__.: -precision = 2 \ No newline at end of file +precision = 2 + +[codespell] +skip = docs/auto_*,*.html,.git/,*.pyc,docs/_build +count = +quiet-level = 3 +ignore-words = ignore_words.txt +interactive = 0 +builtin = clear,rare,informal,names,usage -- 2.52.0 From b5757834932cb5d2627d4d5bdcf73cb2d9236a81 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 20 Jul 2022 19:02:17 +0200 Subject: [PATCH 080/287] chore: remove unnecessary tool config files and makefile --- .codespellrc | 7 ------- .coveragerc | 14 -------------- .flake8 | 5 ----- Makefile | 15 --------------- 4 files changed, 41 deletions(-) delete mode 100644 .codespellrc delete mode 100644 .coveragerc delete mode 100644 .flake8 delete mode 100644 Makefile diff --git a/.codespellrc b/.codespellrc deleted file mode 100644 index ddfbbf2f8..000000000 --- a/.codespellrc +++ /dev/null @@ -1,7 +0,0 @@ -[codespell] -skip = docs/auto_*,*.html,.git/,*.pyc,docs/_build -count = -quiet-level = 3 -ignore-words = ignore_words.txt -interactive = 0 -builtin = clear,rare,informal,names,usage \ No newline at end of file diff --git a/.coveragerc b/.coveragerc deleted file mode 100644 index 2cd7255d0..000000000 --- a/.coveragerc +++ /dev/null @@ -1,14 +0,0 @@ -[run] -branch = True -source = junifer -include = */junifer/* -omit = - */setup.py - */tests/* - junifer/configs/juseless.py - junifer/testing/* - -[report] -exclude_lines = - pragma: no cover - if __name__ == .__main__.: \ No newline at end of file diff --git a/.flake8 b/.flake8 deleted file mode 100644 index e57304ad3..000000000 --- a/.flake8 +++ /dev/null @@ -1,5 +0,0 @@ -[flake8] -exclude = __init__.py,*externals*,constants.py,fixes.py,resources.py,nilearn_cache,venv,docs/auto_examples,docs/_build/,.eggs/,scratch/ -ignore = W503,W504,I100,I101,I201,N806,E201,E202,E221,E222,E241,F541 -# We add A for the array-spacing plugin, and ignore the E ones it covers above -select = A,E,F,W,C \ No newline at end of file diff --git a/Makefile b/Makefile deleted file mode 100644 index d1b4e3a6c..000000000 --- a/Makefile +++ /dev/null @@ -1,15 +0,0 @@ -# Makefile before PR -# - -.PHONY: checks - -checks: flake spellcheck - -flake: - flake8 - -spellcheck: - codespell junifer/ docs/ examples/ - -test: - pytest -vv --cov=junifer --cov-report html --cov-report term \ No newline at end of file -- 2.52.0 From 8c09f0f26810b54e532257af73d04d4bad5b9c93 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 21 Jul 2022 13:22:29 +0200 Subject: [PATCH 081/287] chore: remove dev-requirements.txt --- dev-requirements.txt | 3 --- 1 file changed, 3 deletions(-) delete mode 100644 dev-requirements.txt diff --git a/dev-requirements.txt b/dev-requirements.txt deleted file mode 100644 index 2e4e347ff..000000000 --- a/dev-requirements.txt +++ /dev/null @@ -1,3 +0,0 @@ -flake8 -pytest -click \ No newline at end of file -- 2.52.0 From a56f8bdb7d87c202fbd885d0a81021d1e2bea851 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 21 Jul 2022 13:22:46 +0200 Subject: [PATCH 082/287] chore: remove test-requirements.txt --- test-requirements.txt | 5 ----- 1 file changed, 5 deletions(-) delete mode 100644 test-requirements.txt diff --git a/test-requirements.txt b/test-requirements.txt deleted file mode 100644 index 17b606f42..000000000 --- a/test-requirements.txt +++ /dev/null @@ -1,5 +0,0 @@ -flake8 -pytest -pytest-cov -codecov -https://github.com/codespell-project/codespell/archive/master.zip \ No newline at end of file -- 2.52.0 From 95046b1806da65820b1b7c36224eb23f3283c33a Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 21 Jul 2022 14:29:07 +0200 Subject: [PATCH 083/287] update: remove metadata and improve code style for setup.py --- setup.py | 51 ++++++++------------------------------------------- 1 file changed, 8 insertions(+), 43 deletions(-) diff --git a/setup.py b/setup.py index 6c76b7b05..b0601efcc 100644 --- a/setup.py +++ b/setup.py @@ -1,15 +1,18 @@ +"""Set up junifer package.""" + # Authors: Federico Raimondo # Sami Hamdan +# Synchon Mandal # License: AGPL -import setuptools -with open('README.md', 'r') as fh: - long_description = fh.read() +from setuptools import setup def _getversion(): - from setuptools_scm.version import get_local_node_and_date, \ - simplified_semver_version + from setuptools_scm.version import ( + get_local_node_and_date, + simplified_semver_version, + ) def clean_scheme(version): print(version) @@ -22,46 +25,8 @@ def _getversion(): 'write_to_template': "__version__ = '{version}'\n"} -DOWNLOAD_URL = 'https://github.com/juaml/junifer' -URL = 'https://juaml.github.io/junifer' - -# TODO: Read requirementes from requirements.txt and use them setuptools.setup( - name='junifer', - author='Fede Raimondo', - author_email='f.raimondo@fz-juelich.de', - description='JUelich NeuroImaging FEature extractoR', - long_description=long_description, - long_description_content_type='text/markdown', - url=URL, - download_url=DOWNLOAD_URL, - packages=setuptools.find_packages(), - zip_safe=False, - classifiers=['Intended Audience :: Science/Research', - 'Intended Audience :: Developers', - 'License :: OSI Approved', - 'Programming Language :: Python', - 'Topic :: Software Development', - 'Topic :: Scientific/Engineering', - 'Operating System :: Microsoft :: Windows', - 'Operating System :: POSIX', - 'Operating System :: Unix', - 'Operating System :: MacOS', - 'Programming Language :: Python :: 3'], - project_urls={ - 'Documentation': URL, - 'Source': DOWNLOAD_URL, - 'Tracker': f'{DOWNLOAD_URL}issues/', - }, - install_requires=['Click'], # TODO: Complete - py_modules=['junifer.api.cli'], - entry_points={ - 'console_scripts': [ - 'junifer=junifer.api.cli:cli', - ] - }, - python_requires='>=3.6', use_scm_version=_getversion, setup_requires=['setuptools_scm'], ) -- 2.52.0 From 90dcb8aa4fe5d392005e3a31b66f49ff6bc94ad1 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 21 Jul 2022 14:30:31 +0200 Subject: [PATCH 084/287] update: add metadata and improve code style for pyproject.toml --- pyproject.toml | 74 ++++++++++++++++++++++++++++++++++++++++++++++---- 1 file changed, 68 insertions(+), 6 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 22429039a..8816aea11 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,13 +1,75 @@ [build-system] -requires = ["setuptools>=45", "wheel", "setuptools_scm[toml]>=6.2"] +requires = [ + "setuptools >= 61.0.0", + "wheel", + "setuptools_scm[toml] >= 6.2" +] build-backend = "setuptools.build_meta" +[project] +name = "junifer" +description = "JUelich NeuroImaging FEature extractoR" +readme = "README.md" +requires-python = ">=3.7" +license = {file = "LICENSE.md"} +authors = [ + {email = "f.raimondo@fz-juelich.de"}, + {name = "Fede Raimondo"} +] +maintainers = [ + {email = "s.mandal@fz-juelich.de"}, + {name = "Synchon Mandal"} +] +keywords = [ + "neuroimaging", +] +classifiers = [ + "Development Status :: 4 - Beta", + "Intended Audience :: Science/Research", + "Intended Audience :: Developers", + "License :: OSI Approved", + "Natural Language :: English", + "Topic :: Software Development", + "Topic :: Scientific/Engineering", + "Topic :: Scientific/Engineering :: Bio-Informatics", + "Operating System :: OS Independent", + "Programming Language :: Python :: 3.7", + "Programming Language :: Python :: 3.8", + "Programming Language :: Python :: 3.9", + "Programming Language :: Python :: 3.10", +] +dependencies = [ + "click>=8.1.3,<8.2", + "numpy>=1.20,<1.22", + "datalad>=0.15.4,<0.16", + "pandas>=1.4.0,<1.5", + "nibabel>=3.2.0,<4.0", + "nilearn>=0.9.0,<1.0", + "sqlalchemy>=1.4.27,<= 1.5.0", + "pyyaml>=5.1.2,<7.0", +] +dynamic = ["version"] + +[project.urls] +homepage = "https://juaml.github.io/junifer" +documentation = "https://juaml.github.io/junifer" +repository = "https://github.com/juaml/junifer" + +[project.scripts] +junifer = "junifer.api.cli:cli" + +[project.optional-dependencies] +dev = ["tox"] + +################ +# Tool configs # +################ + +[tool.setuptools] +packages = ["junifer"] + [tool.setuptools_scm] version_scheme = "python-simplified-semver" local_scheme = "no-local-version" write_to = "junifer/_version.py" -write_to_template = "__version__ = '{version}'\n" - -[tool.pytest.ini_options] -log_cli = true -log_cli_level = "WARNING" \ No newline at end of file +write_to_template = "__version__ = '{version}'\n" \ No newline at end of file -- 2.52.0 From 18bbff7d2804c77bc66fb90983080d2a46bfe82f Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 21 Jul 2022 14:31:11 +0200 Subject: [PATCH 085/287] chore: remove requirements.txt --- requirements.txt | 7 ------- 1 file changed, 7 deletions(-) delete mode 100644 requirements.txt diff --git a/requirements.txt b/requirements.txt deleted file mode 100644 index a91a82997..000000000 --- a/requirements.txt +++ /dev/null @@ -1,7 +0,0 @@ -numpy>=1.20, <1.22 -datalad>=0.15.4, <0.16 -pandas>=1.4.0, <1.5 -nibabel>=3.2.0, <4.0 -nilearn>=0.9.0, <1.0 -sqlalchemy>=1.4.27, <= 1.5.0 -pyyaml>=5.1.2, <7.0 \ No newline at end of file -- 2.52.0 From 6436d3818a839991401dbe62874135eb43557f68 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 21 Jul 2022 14:34:15 +0200 Subject: [PATCH 086/287] fix: add missing guard block for setup.py --- setup.py | 9 ++++----- 1 file changed, 4 insertions(+), 5 deletions(-) diff --git a/setup.py b/setup.py index b0601efcc..8c99b5acd 100644 --- a/setup.py +++ b/setup.py @@ -25,8 +25,7 @@ def _getversion(): 'write_to_template': "__version__ = '{version}'\n"} - -setuptools.setup( - use_scm_version=_getversion, - setup_requires=['setuptools_scm'], -) +if __name__ == "__main__": + setup( + use_scm_version=_getversion, + ) -- 2.52.0 From a4e0eeb7bf2bb79595ed6f44fd41776b1fe004e8 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 21 Jul 2022 16:37:18 +0200 Subject: [PATCH 087/287] update: add support for python envs in tox --- tox.ini | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/tox.ini b/tox.ini index 3e200f0f9..80caf2596 100644 --- a/tox.ini +++ b/tox.ini @@ -1,7 +1,14 @@ [tox] -envlist = flake8, test, coverage, codespell +envlist = flake8, test, coverage, codespell, py3{7,8,9,10} isolated_build = true +[gh-actions] +python = + 3.7: py37 + 3.8: py38 + 3.9: py39 + 3.10: py310 + [testenv:flake8] skip_install = true deps = -- 2.52.0 From 53ce28a82d2b3ce34c57c88a35e2550e845ae21b Mon Sep 17 00:00:00 2001 From: Fede Date: Fri, 22 Jul 2022 12:01:25 +0300 Subject: [PATCH 088/287] Fix tests --- junifer/datareader/default.py | 2 +- junifer/datareader/tests/test_default_reader.py | 14 +++++++------- junifer/testing/datagrabbers.py | 2 +- pyproject.toml | 6 +++--- 4 files changed, 12 insertions(+), 12 deletions(-) diff --git a/junifer/datareader/default.py b/junifer/datareader/default.py index b1a4acf83..226c4d608 100644 --- a/junifer/datareader/default.py +++ b/junifer/datareader/default.py @@ -40,7 +40,7 @@ class DefaultDataReader(PipelineStepMixin): # out is the same, but with the 'data' key set in # each kind dictionary, except for meta - out = input + out = input.copy() if params is None: params = {} for kind in input.keys(): diff --git a/junifer/datareader/tests/test_default_reader.py b/junifer/datareader/tests/test_default_reader.py index b1a110885..5ad9ccb39 100644 --- a/junifer/datareader/tests/test_default_reader.py +++ b/junifer/datareader/tests/test_default_reader.py @@ -37,7 +37,7 @@ def test_meta(): nib_data_path = Path(nib_testing.data_path) t_path = nib_data_path / 'example4d.nii.gz' - input = {'bold': t_path} + input = {'bold': {'path': t_path}} output = reader.fit_transform(input) assert 'meta' in output assert 'datareader' in output['meta'] @@ -54,7 +54,7 @@ def test_read_nifti(): 'reoriented_anat_moved.nii']: t_path = nib_data_path / fname - input = {'bold': t_path} + input = {'bold': {'path': t_path}} output = reader.fit_transform(input) assert isinstance(output, dict) @@ -68,7 +68,7 @@ def test_read_nifti(): t_read_img = nib.load(t_path) assert_array_equal(read_img.get_fdata(), t_read_img.get_fdata()) - input = {'bold': t_path.as_posix()} + input = {'bold': {'path': t_path.as_posix()}} output2 = reader.fit_transform(input) assert output['bold']['path'] == output2['bold']['path'] @@ -81,7 +81,7 @@ def test_read_unknown(): anat_path = nib_data_path / 'reoriented_anat_moved.nii' whatever_path = nib_data_path / 'unexistent.unkwnownextension' - input = {'anat': anat_path, 'whatever': whatever_path} + input = {'anat': {'path': anat_path}, 'whatever': {'path': whatever_path}} output = reader.fit_transform(input) assert isinstance(output, dict) @@ -108,7 +108,7 @@ def test_read_csv(): df.to_csv(tmpdir / 'test.csv') reader = DefaultDataReader() - input = {'csv': tmpdir / 'test.csv'} + input = {'csv': {'path': tmpdir / 'test.csv'}} output = reader.fit_transform(input) assert isinstance(output, dict) @@ -121,7 +121,7 @@ def test_read_csv(): assert_frame_equal(df, read_df) df.to_csv(tmpdir / 'test.csv', sep=';') - input = {'csv': tmpdir / 'test.csv'} + input = {'csv': {'path': tmpdir / 'test.csv'}} params = {'csv': {'sep': ';'}} output = reader.fit_transform(input, params) @@ -135,7 +135,7 @@ def test_read_csv(): assert_frame_equal(df, read_df) df.to_csv(tmpdir / 'test.tsv', sep='\t') - input = {'csv': tmpdir / 'test.tsv'} + input = {'csv': {'path': tmpdir / 'test.tsv'}} output = reader.fit_transform(input) assert isinstance(output, dict) diff --git a/junifer/testing/datagrabbers.py b/junifer/testing/datagrabbers.py index db8003a6f..99c0d0642 100644 --- a/junifer/testing/datagrabbers.py +++ b/junifer/testing/datagrabbers.py @@ -21,7 +21,7 @@ class OasisVBMTestingDatagrabber(BaseDataGrabber): def __getitem__(self, element): out = super().__getitem__(element) i_sub = int(element.split('-')[1]) - 1 - out['VBM_GM'] = self._dataset.gray_matter_maps[i_sub] + out['VBM_GM'] = {'path': self._dataset.gray_matter_maps[i_sub]} # Set the element accordingly out['meta']['element'] = {'subject': element} return out diff --git a/pyproject.toml b/pyproject.toml index 8816aea11..0c41cc18f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -40,10 +40,10 @@ classifiers = [ ] dependencies = [ "click>=8.1.3,<8.2", - "numpy>=1.20,<1.22", - "datalad>=0.15.4,<0.16", + "numpy>=1.22,<1.23", + "datalad>=0.15.4,<0.18", "pandas>=1.4.0,<1.5", - "nibabel>=3.2.0,<4.0", + "nibabel>=3.2.0,<4.1", "nilearn>=0.9.0,<1.0", "sqlalchemy>=1.4.27,<= 1.5.0", "pyyaml>=5.1.2,<7.0", -- 2.52.0 From 932d5da2a14042205acef26876bcd70fbc1b44e7 Mon Sep 17 00:00:00 2001 From: Fede Date: Fri, 22 Jul 2022 12:27:22 +0300 Subject: [PATCH 089/287] Doc in data --- docs/data.rst | 17 +++++++++++++++++ 1 file changed, 17 insertions(+) diff --git a/docs/data.rst b/docs/data.rst index c62a7dfc3..29a52a591 100644 --- a/docs/data.rst +++ b/docs/data.rst @@ -29,3 +29,20 @@ Data types * - `VBM_WM` - VBM White Matter segmentation (3D) - CAT output (`m0wp2` images) + + +The Data Object +^^^^^^^^^^^^^^^ + +This is the _object_ that traverses the steps of the pipeline. It is indeed a +dictionary of dictionaries. The first level of keys are the _data types_ and a +special key named 'meta' that contains all the information on the data object +including source and previous transformation steps. + +The second level of keys are the actual data. So far, there are two keys used: +- `path`: path to the file containing the data. +- `data`: the data loaded in memory. + +The _DataGrabber_ step will only fill the `path` value. The `data` value will +be filled by the _DataReader_ step, if it is one of the possible file types +that the datareader can read. -- 2.52.0 From 6b541fd3cd27d0d7798752788cf43e6593396af0 Mon Sep 17 00:00:00 2001 From: Fede Date: Fri, 22 Jul 2022 12:42:02 +0300 Subject: [PATCH 090/287] Fix examples + doc --- docs/api.rst | 4 ++-- docs/data.rst | 8 ++++---- examples/run_datagrabber_bids_datalad.py | 23 +++++++++++++---------- 3 files changed, 19 insertions(+), 16 deletions(-) diff --git a/docs/api.rst b/docs/api.rst index 5f50340cb..a2789d71d 100644 --- a/docs/api.rst +++ b/docs/api.rst @@ -9,9 +9,9 @@ Data Grabbers .. autoclass:: junifer.datagrabber.base.BaseDataGrabber :members: -.. autoclass:: junifer.datagrabber.base.BIDSDataGrabber +.. autoclass:: junifer.datagrabber.base.PatternDataGrabber :members: .. autoclass:: junifer.datagrabber.base.DataladDataGrabber :members: -.. autoclass:: junifer.datagrabber.base.BIDSDataladDataGrabber +.. autoclass:: junifer.datagrabber.base.PatternDataladDataGrabber :members: \ No newline at end of file diff --git a/docs/data.rst b/docs/data.rst index 29a52a591..00e387178 100644 --- a/docs/data.rst +++ b/docs/data.rst @@ -34,8 +34,8 @@ Data types The Data Object ^^^^^^^^^^^^^^^ -This is the _object_ that traverses the steps of the pipeline. It is indeed a -dictionary of dictionaries. The first level of keys are the _data types_ and a +This is the *object* that traverses the steps of the pipeline. It is indeed a +dictionary of dictionaries. The first level of keys are the *data types* and a special key named 'meta' that contains all the information on the data object including source and previous transformation steps. @@ -43,6 +43,6 @@ The second level of keys are the actual data. So far, there are two keys used: - `path`: path to the file containing the data. - `data`: the data loaded in memory. -The _DataGrabber_ step will only fill the `path` value. The `data` value will -be filled by the _DataReader_ step, if it is one of the possible file types +The *DataGrabber* step will only fill the `path` value. The `data` value will +be filled by the *DataReader* step, if it is one of the possible file types that the datareader can read. diff --git a/examples/run_datagrabber_bids_datalad.py b/examples/run_datagrabber_bids_datalad.py index 221b4f74c..28a1eeb77 100644 --- a/examples/run_datagrabber_bids_datalad.py +++ b/examples/run_datagrabber_bids_datalad.py @@ -10,7 +10,7 @@ Authors: Federico Raimondo License: BSD 3 clause """ -from junifer.datagrabber.base import BIDSDataladDataGrabber +from junifer.datagrabber.base import PatternDataladDataGrabber from junifer.utils import configure_logging ############################################################################### @@ -19,14 +19,15 @@ configure_logging(level='INFO') ############################################################################### -# The BIDS datagrabber requires two parameters: the types of data we want, -# and the specific pattern that matches each type. +# The BIDS datagrabber requires three parameters: the types of data we want, +# the specific pattern that matches each type, and the variables that will be +# replaced int he patterns. types = ['T1w', 'bold'] patterns = { - 'T1w': 'anat/{subject}_T1w.nii.gz', - 'bold': 'func/{subject}_task-rest_bold.nii.gz' + 'T1w': '{subject}/anat/{subject}_T1w.nii.gz', + 'bold': '{subject}/func/{subject}_task-rest_bold.nii.gz' } - +replacements = ['subject'] ############################################################################### # Additionally, a datalad datagrabber requires the URI of the remote sibling # and the location of the dataset within the remote sibling. @@ -37,8 +38,9 @@ rootdir = 'example_bids' # Now we can use the datagrabber within a `with` context # One thing we can do with any datagrabber is iterate over the elements. # In this case, each element of the datagrabber is one session. -with BIDSDataladDataGrabber(rootdir=rootdir, types=types, - patterns=patterns, uri=repo_uri) as dg: +with PatternDataladDataGrabber(rootdir=rootdir, types=types, + patterns=patterns, uri=repo_uri, + replacements=replacements) as dg: for elem in dg: print(elem) @@ -46,7 +48,8 @@ with BIDSDataladDataGrabber(rootdir=rootdir, types=types, # Another feature of the datagrabber is the ability to get a specific # element by its name. In this case, we index `sub-01` and we get the file # paths for the two types of data we want (T1w and bold). -with BIDSDataladDataGrabber(rootdir=rootdir, types=types, - patterns=patterns, uri=repo_uri) as dg: +with PatternDataladDataGrabber(rootdir=rootdir, types=types, + patterns=patterns, uri=repo_uri, + replacements=replacements) as dg: sub01 = dg['sub-01'] print(sub01) -- 2.52.0 From 2fe1a9afe74fbbc773f55875ffdbb4adda73d2f7 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 22 Jul 2022 12:57:36 +0200 Subject: [PATCH 091/287] update: fix ci.yml to work with gh-actions --- .github/workflows/ci.yml | 26 ++++++-------------------- 1 file changed, 6 insertions(+), 20 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 091cebeb9..c0d6a824d 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -4,7 +4,6 @@ on: [push, pull_request] jobs: build: - runs-on: ubuntu-latest strategy: fail-fast: false @@ -27,30 +26,17 @@ jobs: $SUDO bash -c "$(curl -fsSL http://neuro.debian.net/_files/neurodebian-travis.sh)" $SUDO apt-get update -qq $SUDO apt-get install git-annex-standalone - python -m pip install --upgrade pip - pip install -r test-requirements.txt - pip install -r requirements.txt + python -m pip install --upgrade pip setuptools wheel + python -m pip install tox tox-gh-actions - name: Configure git for datalad run: | git config --global user.email "runner@github.com" git config --global user.name "GITHUB CI Runner" - - name: Install junifer - shell: bash -el {0} + - name: Test with tox run: | - python setup.py build - python setup.py install - - name: Lint with flake8 - run: | - # stop the build if there are Python syntax errors or undefined names - flake8 . --count --show-source --statistics - - name: Spell check - run: | - codespell junifer/ docs/ examples/ - - name: Test with pytest - run: | - PYTHONPATH="." pytest --cov=junifer --cov-report xml -vv junifer/ - - name: 'Upload coverage to CodeCov' + tox + - name: Upload coverage to Codecov uses: codecov/codecov-action@master with: token: ${{ secrets.CODECOV_TOKEN }} - if: success() && matrix.python-version == 3.8 + if: success() && matrix.python-version == 3.9 -- 2.52.0 From 70e35682821ce3bf0fd170eae25ab99c896cbf1e Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 22 Jul 2022 13:19:18 +0200 Subject: [PATCH 092/287] update: fix docs.yml to work with gh-actions --- .github/workflows/docs.yml | 110 ++++++++++++++++++------------------- 1 file changed, 52 insertions(+), 58 deletions(-) diff --git a/.github/workflows/docs.yml b/.github/workflows/docs.yml index 58b7314ef..64a033b2c 100644 --- a/.github/workflows/docs.yml +++ b/.github/workflows/docs.yml @@ -7,64 +7,58 @@ jobs: runs-on: ubuntu-latest strategy: fail-fast: false - steps: - - name: Checkout Source - uses: actions/checkout@v2 - with: - # require all of history to see all tagged versions' docs - fetch-depth: 0 - - name: Set up Python ${{ matrix.python-version }} - uses: actions/setup-python@v2 - with: - python-version: 3.8 - - name: Check for sudo - shell: bash - run: | - if type sudo >/dev/null 2>&1; then SUDO="sudo"; else SUDO=""; fi - echo "SUDO=$SUDO" >> $GITHUB_ENV - - name: Install Dependencies - run: | - $SUDO bash -c "$(curl -fsSL http://neuro.debian.net/_files/neurodebian-travis.sh)" - $SUDO apt-get update -qq - $SUDO apt-get install git-annex-standalone - python -m pip install --upgrade pip - pip install -r requirements.txt - pip install -r docs-requirements.txt - python setup.py build - python setup.py install - - name: Configure git for datalad - run: | - git config --global user.email "runner@github.com" - git config --global user.name "GITHUB CI Runner" - - name: Checkout gh-pages - # As we already did a deploy of gh-pages above, it is guaranteed to be there - # so check it out so we can selectively build docs below - uses: actions/checkout@v2 - with: + steps: + - name: Checkout source + uses: actions/checkout@v2 + with: + # require all of history to see all tagged versions' docs + fetch-depth: 0 + - name: Set up Python 3.9 + uses: actions/setup-python@v2 + with: + python-version: 3.9 + - name: Check for sudo + shell: bash + run: | + if type sudo >/dev/null 2>&1; then SUDO="sudo"; else SUDO=""; fi + echo "SUDO=$SUDO" >> $GITHUB_ENV + - name: Install dependencies + run: | + $SUDO bash -c "$(curl -fsSL http://neuro.debian.net/_files/neurodebian-travis.sh)" + $SUDO apt-get update -qq + $SUDO apt-get install git-annex-standalone + python -m pip install --upgrade pip setuptools wheel + python -m pip install -e .[docs] + - name: Configure git for datalad + run: | + git config --global user.email "runner@github.com" + git config --global user.name "GITHUB CI Runner" + - name: Checkout gh-pages + # As we already did a deploy of gh-pages above, it is guaranteed to be there + # so check it out so we can selectively build docs below + uses: actions/checkout@v2 + with: ref: gh-pages path: docs/_build - - - name: Test Build Docs - if: github.ref != 'refs/heads/main' && ! startsWith(github.ref, 'refs/tags') - run: | - BUILDDIR=_build/main make -C docs/ local - - - name: Build Docs - if: github.ref == 'refs/heads/main' || startsWith(github.ref, 'refs/tags') - # Use the args we normally pass to sphinx-build, but run sphinx-multiversion - run: | - make -C docs/ html - touch docs/_build/.nojekyll - cp docs/redirect.html docs/_build/index.html - - - name: Publish Docs to gh-pages - # Only once from main or a tag - if: github.ref == 'refs/heads/main' || startsWith(github.ref, 'refs/tags') - # We pin to the SHA, not the tag, for security reasons. - # https://docs.github.com/en/free-pro-team@latest/actions/learn-github-actions/security-hardening-for-github-actions#using-third-party-actions - uses: peaceiris/actions-gh-pages@bbdfb200618d235585ad98e965f4aafc39b4c501 # v3.7.3 - with: - github_token: ${{ secrets.GITHUB_TOKEN }} - publish_dir: docs/_build - keep_files: true + - name: Test build docs + if: github.ref != 'refs/heads/main' && ! startsWith(github.ref, 'refs/tags') + run: | + BUILDDIR=_build/main make -C docs/ local + - name: Build docs + if: github.ref == 'refs/heads/main' || startsWith(github.ref, 'refs/tags') + # Use the args we normally pass to sphinx-build, but run sphinx-multiversion + run: | + make -C docs/ html + touch docs/_build/.nojekyll + cp docs/redirect.html docs/_build/index.html + - name: Publish docs to gh-pages + # Only once from main or a tag + if: github.ref == 'refs/heads/main' || startsWith(github.ref, 'refs/tags') + # We pin to the SHA, not the tag, for security reasons. + # https://docs.github.com/en/free-pro-team@latest/actions/learn-github-actions/security-hardening-for-github-actions#using-third-party-actions + uses: peaceiris/actions-gh-pages@bbdfb200618d235585ad98e965f4aafc39b4c501 # v3.7.3 + with: + github_token: ${{ secrets.GITHUB_TOKEN }} + publish_dir: docs/_build + keep_files: true -- 2.52.0 From f113c6772bdf422ae7a7fd1dcea3a9f0fb754fa5 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 22 Jul 2022 15:04:54 +0200 Subject: [PATCH 093/287] update: add docs requirements in pyproject.toml --- pyproject.toml | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/pyproject.toml b/pyproject.toml index 0c41cc18f..ddd3e2bc5 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -60,6 +60,14 @@ junifer = "junifer.api.cli:cli" [project.optional-dependencies] dev = ["tox"] +docs = [ + "seaborn>=0.11.2,<0.12", + "Sphinx>=5.0.2,<5.1", + "sphinx-gallery>=0.10.1,<0.11", + "sphinx-rtd-theme>=1.0.0,<1.1", + "sphinx-multiversion>=0.2.4,<0.3", + "numpydoc>=1.4.0,<1.5", +] ################ # Tool configs # -- 2.52.0 From 789b9376bdc3893821dbe1d8ee32f13628151c5d Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 22 Jul 2022 15:05:30 +0200 Subject: [PATCH 094/287] chore: remove docs-requirements.txt --- docs-requirements.txt | 6 ------ 1 file changed, 6 deletions(-) delete mode 100644 docs-requirements.txt diff --git a/docs-requirements.txt b/docs-requirements.txt deleted file mode 100644 index 08da25ce4..000000000 --- a/docs-requirements.txt +++ /dev/null @@ -1,6 +0,0 @@ -seaborn -sphinx -sphinx-gallery -sphinx_rtd_theme -git+https://github.com/dls-controls/sphinx-multiversion.git@only-arg -numpydoc \ No newline at end of file -- 2.52.0 From 196dcd7947a9cc287d23f76e391c4efe0eaf6589 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 22 Jul 2022 15:08:10 +0200 Subject: [PATCH 095/287] update: drop Python 3.7 support --- .github/workflows/ci.yml | 2 +- pyproject.toml | 3 +-- tox.ini | 3 +-- 3 files changed, 3 insertions(+), 5 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index c0d6a824d..8704f9451 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -8,7 +8,7 @@ jobs: strategy: fail-fast: false matrix: - python-version: ['3.7', '3.8', '3.9', '3.10'] + python-version: ['3.8', '3.9', '3.10'] steps: - uses: actions/checkout@v2 diff --git a/pyproject.toml b/pyproject.toml index ddd3e2bc5..4191774b5 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -10,7 +10,7 @@ build-backend = "setuptools.build_meta" name = "junifer" description = "JUelich NeuroImaging FEature extractoR" readme = "README.md" -requires-python = ">=3.7" +requires-python = ">=3.8" license = {file = "LICENSE.md"} authors = [ {email = "f.raimondo@fz-juelich.de"}, @@ -33,7 +33,6 @@ classifiers = [ "Topic :: Scientific/Engineering", "Topic :: Scientific/Engineering :: Bio-Informatics", "Operating System :: OS Independent", - "Programming Language :: Python :: 3.7", "Programming Language :: Python :: 3.8", "Programming Language :: Python :: 3.9", "Programming Language :: Python :: 3.10", diff --git a/tox.ini b/tox.ini index 80caf2596..3f10074f5 100644 --- a/tox.ini +++ b/tox.ini @@ -1,10 +1,9 @@ [tox] -envlist = flake8, test, coverage, codespell, py3{7,8,9,10} +envlist = flake8, test, coverage, codespell, py3{8,9,10} isolated_build = true [gh-actions] python = - 3.7: py37 3.8: py38 3.9: py39 3.10: py310 -- 2.52.0 From ba4336c92404419874bbc338cb6df828b3fd29bd Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 22 Jul 2022 16:03:23 +0200 Subject: [PATCH 096/287] update: add coverage to 3.9 for gh-actions in tox.ini --- tox.ini | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tox.ini b/tox.ini index 3f10074f5..6b870308d 100644 --- a/tox.ini +++ b/tox.ini @@ -5,7 +5,7 @@ isolated_build = true [gh-actions] python = 3.8: py38 - 3.9: py39 + 3.9: py39, coverage 3.10: py310 [testenv:flake8] -- 2.52.0 From d7616e9c9dec82c835f037b793cd7da3ce938823 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 25 Jul 2022 11:30:19 +0200 Subject: [PATCH 097/287] update: add pytest for python envs in tox.ini --- tox.ini | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/tox.ini b/tox.ini index 6b870308d..e836386f5 100644 --- a/tox.ini +++ b/tox.ini @@ -8,6 +8,13 @@ python = 3.9: py39, coverage 3.10: py310 +[testenv] +skip_install = false +deps = + pytest +commands = + pytest + [testenv:flake8] skip_install = true deps = -- 2.52.0 From be707ad183f0d98a5c3993a8cbb4341be36fba5f Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 25 Jul 2022 11:53:52 +0200 Subject: [PATCH 098/287] fix: use absolute path for setting git config in ci.yml --- .github/workflows/ci.yml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 8704f9451..f0e47d36a 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -30,8 +30,8 @@ jobs: python -m pip install tox tox-gh-actions - name: Configure git for datalad run: | - git config --global user.email "runner@github.com" - git config --global user.name "GITHUB CI Runner" + /usr/bin/git config --global user.email "runner@github.com" + /usr/bin/git config --global user.name "GITHUB CI Runner" - name: Test with tox run: | tox -- 2.52.0 From 7fcd09a0cbcf2c55f424890cb4111295f3a7197f Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 25 Jul 2022 13:55:11 +0200 Subject: [PATCH 099/287] chore: update contributors list --- AUTHORS.rst | 3 ++- docs/changes/contributors.inc | 3 ++- docs/index.rst | 1 - 3 files changed, 4 insertions(+), 3 deletions(-) diff --git a/AUTHORS.rst b/AUTHORS.rst index ac9115336..36fbf49de 100644 --- a/AUTHORS.rst +++ b/AUTHORS.rst @@ -1,4 +1,5 @@ Original Authors ================ * Federico Raimondo -* Leonard Sasse \ No newline at end of file +* Leonard Sasse +* Synchon Mandal diff --git a/docs/changes/contributors.inc b/docs/changes/contributors.inc index 50ed21edf..ff1c379f9 100644 --- a/docs/changes/contributors.inc +++ b/docs/changes/contributors.inc @@ -1,4 +1,5 @@ .. _Fede Raimondo: https://fraimondo.github.io .. _Kaustubh Patil: https://github.com/kaurao .. _Leonard Sasse: https://github.com/LeSasse -.. _Amir Omidvarnia: https://github.com/omidvarnia \ No newline at end of file +.. _Amir Omidvarnia: https://github.com/omidvarnia +.. _Synchon Mandal: https://github.com/synchon \ No newline at end of file diff --git a/docs/index.rst b/docs/index.rst index eee62d926..3deeb2c34 100644 --- a/docs/index.rst +++ b/docs/index.rst @@ -22,4 +22,3 @@ Indices and tables * :ref:`genindex` * :ref:`modindex` * :ref:`search` - -- 2.52.0 From 20cb20712da761f1838c298c851e54969b461110 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 25 Jul 2022 13:56:15 +0200 Subject: [PATCH 100/287] fix: update ci.yml to make git config work properly --- .github/workflows/ci.yml | 18 ++++++++++-------- 1 file changed, 10 insertions(+), 8 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index f0e47d36a..3822a54e5 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -11,27 +11,29 @@ jobs: python-version: ['3.8', '3.9', '3.10'] steps: - - uses: actions/checkout@v2 - - name: Set up Python ${{ matrix.python-version }} - uses: actions/setup-python@v2 - with: - python-version: ${{ matrix.python-version }} + - uses: actions/checkout@v3 - name: Check for sudo shell: bash run: | if type sudo >/dev/null 2>&1; then SUDO="sudo"; else SUDO=""; fi echo "SUDO=$SUDO" >> $GITHUB_ENV - - name: Install dependencies + - name: Set up system run: | $SUDO bash -c "$(curl -fsSL http://neuro.debian.net/_files/neurodebian-travis.sh)" $SUDO apt-get update -qq $SUDO apt-get install git-annex-standalone - python -m pip install --upgrade pip setuptools wheel - python -m pip install tox tox-gh-actions - name: Configure git for datalad run: | /usr/bin/git config --global user.email "runner@github.com" /usr/bin/git config --global user.name "GITHUB CI Runner" + - name: Set up Python ${{ matrix.python-version }} + uses: actions/setup-python@v4 + with: + python-version: ${{ matrix.python-version }} + - name: Install dependencies + run: | + python -m pip install --upgrade pip setuptools wheel + python -m pip install tox tox-gh-actions - name: Test with tox run: | tox -- 2.52.0 From 0461d9d016bca2b51b0302fc2630119f9c8b7cef Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 25 Jul 2022 14:18:40 +0200 Subject: [PATCH 101/287] fix: reorder checkout action in ci.yml --- .github/workflows/ci.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 3822a54e5..057467eac 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -11,7 +11,6 @@ jobs: python-version: ['3.8', '3.9', '3.10'] steps: - - uses: actions/checkout@v3 - name: Check for sudo shell: bash run: | @@ -26,6 +25,7 @@ jobs: run: | /usr/bin/git config --global user.email "runner@github.com" /usr/bin/git config --global user.name "GITHUB CI Runner" + - uses: actions/checkout@v3 - name: Set up Python ${{ matrix.python-version }} uses: actions/setup-python@v4 with: -- 2.52.0 From 1f79edf45b44ba1a76449fea696c370702a35c7b Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 25 Jul 2022 14:37:11 +0200 Subject: [PATCH 102/287] fix: improve ci.yml --- .github/workflows/ci.yml | 15 +++++---------- 1 file changed, 5 insertions(+), 10 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 057467eac..4b5ba70fc 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -11,20 +11,15 @@ jobs: python-version: ['3.8', '3.9', '3.10'] steps: - - name: Check for sudo - shell: bash - run: | - if type sudo >/dev/null 2>&1; then SUDO="sudo"; else SUDO=""; fi - echo "SUDO=$SUDO" >> $GITHUB_ENV - name: Set up system run: | - $SUDO bash -c "$(curl -fsSL http://neuro.debian.net/_files/neurodebian-travis.sh)" - $SUDO apt-get update -qq - $SUDO apt-get install git-annex-standalone + bash -c "$(curl -fsSL http://neuro.debian.net/_files/neurodebian-travis.sh)" + sudo apt-get update -qq + sudo apt-get install git-annex-standalone - name: Configure git for datalad run: | - /usr/bin/git config --global user.email "runner@github.com" - /usr/bin/git config --global user.name "GITHUB CI Runner" + git config --global user.email "runner@github.com" + git config --global user.name "GitHub CI Runner" - uses: actions/checkout@v3 - name: Set up Python ${{ matrix.python-version }} uses: actions/setup-python@v4 -- 2.52.0 From 90b88d68a1677e56c51161edf227a90004ad3cf1 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 25 Jul 2022 14:49:28 +0200 Subject: [PATCH 103/287] update: add git config listing in tox.ini --- tox.ini | 1 + 1 file changed, 1 insertion(+) diff --git a/tox.ini b/tox.ini index e836386f5..73c1eaaca 100644 --- a/tox.ini +++ b/tox.ini @@ -29,6 +29,7 @@ skip_install = false deps = pytest commands = + git config --global -l | cat pytest [testenv:coverage] -- 2.52.0 From 20cdfc6dfe9917f78c4c64142cec03bab0be9c15 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 25 Jul 2022 15:11:27 +0200 Subject: [PATCH 104/287] fix: move git config listing to ci.yml --- .github/workflows/ci.yml | 1 + tox.ini | 1 - 2 files changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 4b5ba70fc..d9ffb6be1 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -31,6 +31,7 @@ jobs: python -m pip install tox tox-gh-actions - name: Test with tox run: | + git config --global -l | cat tox - name: Upload coverage to Codecov uses: codecov/codecov-action@master diff --git a/tox.ini b/tox.ini index 73c1eaaca..e836386f5 100644 --- a/tox.ini +++ b/tox.ini @@ -29,7 +29,6 @@ skip_install = false deps = pytest commands = - git config --global -l | cat pytest [testenv:coverage] -- 2.52.0 From 8bd6e336567674685fc48eb4ddf78a3a9b5eb0c7 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 25 Jul 2022 17:07:14 +0200 Subject: [PATCH 105/287] fix: generate .gitconfig for tox pytest tests to make git-annex work --- tox.ini | 10 ++++++++++ 1 file changed, 10 insertions(+) diff --git a/tox.ini b/tox.ini index e836386f5..b3e413916 100644 --- a/tox.ini +++ b/tox.ini @@ -10,9 +10,14 @@ python = [testenv] skip_install = false +# Required for git-annex +setenv = + HOME = {env:HOME:toxworkdir} deps = pytest commands = + # Required for git-annex + python -c "with open('{toxworkdir}/.gitconfig', 'w') as f: f.write('[user]\n name = GitHub Runner\n email = runner@github.com')" pytest [testenv:flake8] @@ -26,9 +31,14 @@ commands = [testenv:test] skip_install = false +# Required for git-annex +setenv = + HOME = {env:HOME:toxworkdir} deps = pytest commands = + # Required for git-annex + python -c "with open('{toxworkdir}/.gitconfig', 'w') as f: f.write('[user]\n name = GitHub Runner\n email = runner@github.com')" pytest [testenv:coverage] -- 2.52.0 From 73bf601ff8a45a7c12bc2b788fddca5efc9c22ea Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 25 Jul 2022 17:08:15 +0200 Subject: [PATCH 106/287] fix: remove git config generation from ci.yml --- .github/workflows/ci.yml | 5 ----- 1 file changed, 5 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index d9ffb6be1..3761cec7a 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -16,10 +16,6 @@ jobs: bash -c "$(curl -fsSL http://neuro.debian.net/_files/neurodebian-travis.sh)" sudo apt-get update -qq sudo apt-get install git-annex-standalone - - name: Configure git for datalad - run: | - git config --global user.email "runner@github.com" - git config --global user.name "GitHub CI Runner" - uses: actions/checkout@v3 - name: Set up Python ${{ matrix.python-version }} uses: actions/setup-python@v4 @@ -31,7 +27,6 @@ jobs: python -m pip install tox tox-gh-actions - name: Test with tox run: | - git config --global -l | cat tox - name: Upload coverage to Codecov uses: codecov/codecov-action@master -- 2.52.0 From 64c19f03638da62432709bf14201d4073842b156 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 25 Jul 2022 17:31:29 +0200 Subject: [PATCH 107/287] fix: move git config generation to ci.yml and add passenv in tox pytest tests --- .github/workflows/ci.yml | 4 ++++ tox.ini | 13 ++++--------- 2 files changed, 8 insertions(+), 9 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 3761cec7a..77a059095 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -16,6 +16,10 @@ jobs: bash -c "$(curl -fsSL http://neuro.debian.net/_files/neurodebian-travis.sh)" sudo apt-get update -qq sudo apt-get install git-annex-standalone + - name: Configure git for datalad + run: | + git config --global user.email "runner@github.com" + git config --global user.name "GitHub Runner" - uses: actions/checkout@v3 - name: Set up Python ${{ matrix.python-version }} uses: actions/setup-python@v4 diff --git a/tox.ini b/tox.ini index b3e413916..178ea2439 100644 --- a/tox.ini +++ b/tox.ini @@ -11,13 +11,11 @@ python = [testenv] skip_install = false # Required for git-annex -setenv = - HOME = {env:HOME:toxworkdir} +passenv = + HOME deps = pytest commands = - # Required for git-annex - python -c "with open('{toxworkdir}/.gitconfig', 'w') as f: f.write('[user]\n name = GitHub Runner\n email = runner@github.com')" pytest [testenv:flake8] @@ -31,14 +29,11 @@ commands = [testenv:test] skip_install = false -# Required for git-annex -setenv = - HOME = {env:HOME:toxworkdir} +passenv = + HOME deps = pytest commands = - # Required for git-annex - python -c "with open('{toxworkdir}/.gitconfig', 'w') as f: f.write('[user]\n name = GitHub Runner\n email = runner@github.com')" pytest [testenv:coverage] -- 2.52.0 From 3e9256715e34e355df1679e13ab290e954881c6e Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 25 Jul 2022 18:38:34 +0200 Subject: [PATCH 108/287] fix: update config for coverage:run in tox.ini --- tox.ini | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/tox.ini b/tox.ini index 178ea2439..f69322cfd 100644 --- a/tox.ini +++ b/tox.ini @@ -94,11 +94,12 @@ source = [coverage:run] branch = true +source = junifer +include = + */junifer/* omit = */setup.py */tests/* - junifer/configs/juseless.py - junifer/testing/* parallel = false [coverage:report] -- 2.52.0 From 1067272c4dfe99fb7c414fd8f55810c0bb257bb6 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 25 Jul 2022 18:42:15 +0200 Subject: [PATCH 109/287] update: increase verbosity for pytest in tox.ini --- tox.ini | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tox.ini b/tox.ini index f69322cfd..dded29853 100644 --- a/tox.ini +++ b/tox.ini @@ -34,7 +34,7 @@ passenv = deps = pytest commands = - pytest + pytest -vv [testenv:coverage] skip_install = false @@ -42,7 +42,7 @@ deps = pytest pytest-cov commands = - pytest --cov=junifer --cov-report=xml + pytest --cov=junifer --cov-report=xml -vv [testenv:codespell] skip_install = true -- 2.52.0 From a5c0e898838db82120d2b16909f48441ba3ac2f1 Mon Sep 17 00:00:00 2001 From: Fede Date: Tue, 26 Jul 2022 11:27:34 +0300 Subject: [PATCH 110/287] Fix tox coverage path --- tox.ini | 5 +---- 1 file changed, 1 insertion(+), 4 deletions(-) diff --git a/tox.ini b/tox.ini index dded29853..5dd6f3e13 100644 --- a/tox.ini +++ b/tox.ini @@ -42,7 +42,7 @@ deps = pytest pytest-cov commands = - pytest --cov=junifer --cov-report=xml -vv + pytest --cov={envsitepackagesdir}/junifer --cov-report=xml -vv --cov-report=term [testenv:codespell] skip_install = true @@ -94,9 +94,6 @@ source = [coverage:run] branch = true -source = junifer -include = - */junifer/* omit = */setup.py */tests/* -- 2.52.0 From 0fd08c6d52aed4295fc0e9a80af64091ea1b1b1c Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 11:53:40 +0200 Subject: [PATCH 111/287] update: remove py39 for Python 3.9 tox env in gh-actions --- tox.ini | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tox.ini b/tox.ini index 5dd6f3e13..80cae3360 100644 --- a/tox.ini +++ b/tox.ini @@ -5,7 +5,7 @@ isolated_build = true [gh-actions] python = 3.8: py38 - 3.9: py39, coverage + 3.9: coverage 3.10: py310 [testenv] -- 2.52.0 From 89a52fb6d21dbddfbf156ba1ed8228229728ac0b Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 13:59:05 +0200 Subject: [PATCH 112/287] update: add lint workflow for gh-actions --- .github/workflows/lint.yml | 31 +++++++++++++++++++++++++++++++ 1 file changed, 31 insertions(+) create mode 100644 .github/workflows/lint.yml diff --git a/.github/workflows/lint.yml b/.github/workflows/lint.yml new file mode 100644 index 000000000..edf01ba2e --- /dev/null +++ b/.github/workflows/lint.yml @@ -0,0 +1,31 @@ +name: Lint + +on: + - push + - pull_request + +jobs: + lint: + runs-on: ${{ matrix.os }} + strategy: + fail-fast: false + matrix: + os: [ubuntu-latest] + python-version: ['3.10'] + + steps: + - uses: actions/checkout@v3 + - name: Set up Python ${{ matrix.python-version }} + uses: actions/setup-python@v4 + with: + python-version: ${{ matrix.python-version }} + - name: Install dependencies + run: | + python -m pip install --upgrade pip setuptools wheel + python -m pip install tox tox-gh-actions + - name: Check with flake8 + run: | + tox -e flake8 + - name: Check with codespell + run: | + tox -e codespell -- 2.52.0 From 0e2f69c8ce5bd95bf0645c51439fb3d44303c90d Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 14:18:07 +0200 Subject: [PATCH 113/287] update: add support for macOS and Windows testing in ci.yml --- .github/workflows/ci.yml | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 77a059095..54efe2900 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -4,10 +4,11 @@ on: [push, pull_request] jobs: build: - runs-on: ubuntu-latest + runs-on: ${{ matrix.os }} strategy: fail-fast: false matrix: + os: [ubuntu-latest, macos-latest, windows-latest] python-version: ['3.8', '3.9', '3.10'] steps: -- 2.52.0 From d2f3af33fc4cf5308073b1f11040d9bfeb0caa6c Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 14:52:35 +0200 Subject: [PATCH 114/287] fix: revert ci.yml to Ubuntu to pass tests --- .github/workflows/ci.yml | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 54efe2900..77a059095 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -4,11 +4,10 @@ on: [push, pull_request] jobs: build: - runs-on: ${{ matrix.os }} + runs-on: ubuntu-latest strategy: fail-fast: false matrix: - os: [ubuntu-latest, macos-latest, windows-latest] python-version: ['3.8', '3.9', '3.10'] steps: -- 2.52.0 From 4bbf4544b63a6b4660cdc695769a8195038ceb2e Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 14:53:14 +0200 Subject: [PATCH 115/287] chore: fix codecov-action to v3 in ci.yml --- .github/workflows/ci.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 77a059095..b7cbe6d0d 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -33,7 +33,7 @@ jobs: run: | tox - name: Upload coverage to Codecov - uses: codecov/codecov-action@master + uses: codecov/codecov-action@v3 with: token: ${{ secrets.CODECOV_TOKEN }} if: success() && matrix.python-version == 3.9 -- 2.52.0 From 069f0eb005fbbcdf4ac040f5a31a446e066ef42c Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 14:53:59 +0200 Subject: [PATCH 116/287] chore: omit _version.py from coverage measurement and reporting --- tox.ini | 1 + 1 file changed, 1 insertion(+) diff --git a/tox.ini b/tox.ini index 80cae3360..c5b8633d6 100644 --- a/tox.ini +++ b/tox.ini @@ -96,6 +96,7 @@ source = branch = true omit = */setup.py + */_version.py */tests/* parallel = false -- 2.52.0 From 485ce67882f05b510bb908b7b829208d38b217cd Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 15:05:04 +0200 Subject: [PATCH 117/287] fix: make codespell work --- junifer/data/atlases.py | 2 +- junifer/datagrabber/base.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/junifer/data/atlases.py b/junifer/data/atlases.py index 2ad8e96e6..b6bb7b3ef 100644 --- a/junifer/data/atlases.py +++ b/junifer/data/atlases.py @@ -430,7 +430,7 @@ def _retrieve_tian( 'Currently there are no labels provided for the 7T Tian atlas. A ' 'simple numbering scheme for distinction was therefore used.') - # check existance of atlas + # check existence of atlas if not (atlas_fname.exists() and atlas_lname.exists()): logger.info( 'At least one of the atlas files is missing. ' diff --git a/junifer/datagrabber/base.py b/junifer/datagrabber/base.py index fb86fa658..bf54f10ad 100644 --- a/junifer/datagrabber/base.py +++ b/junifer/datagrabber/base.py @@ -283,7 +283,7 @@ class PatternDataGrabber(BaseDataGrabber): """ if len(element) != len(self.replacements): raise_error( - f'The element lenght must be {len(self.replacements)}, ' + f'The element length must be {len(self.replacements)}, ' f'indicating {self.replacements}') to_replace = dict(zip(self.replacements, element)) return pattern.format(**to_replace) -- 2.52.0 From 6688e155b1267fd937afb3e30b347693e33a2086 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 16:22:08 +0200 Subject: [PATCH 118/287] fix: flake8 fix for api.cli --- junifer/api/cli.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/junifer/api/cli.py b/junifer/api/cli.py index 6991ae713..e338968e7 100644 --- a/junifer/api/cli.py +++ b/junifer/api/cli.py @@ -1,3 +1,5 @@ +"""Provide functions for cli.""" + import click from .parser import parse_yaml @@ -28,6 +30,7 @@ def _parse_elements(element, config): @click.group() def cli(): + """CLI wrapper.""" pass @@ -41,6 +44,7 @@ def cli(): default='info') @click.option('--element', type=str, multiple=True) def run(filepath, element, verbose): + """Run command for CLI.""" configure_logging(level=verbose.upper()) config = parse_yaml(filepath) workdir = config['workdir'] @@ -62,6 +66,7 @@ def run(filepath, element, verbose): case_sensitive=False), default='info') def collect(filepath, verbose): + """Collect command for CLI.""" configure_logging(level=verbose.upper()) config = parse_yaml(filepath) storage = config['storage'] @@ -80,6 +85,7 @@ def collect(filepath, verbose): @click.option('--submit', is_flag=True) @click.option('--element', type=str, multiple=True) def queue(filepath, element, overwrite, submit, verbose): + """Queue command for CLI.""" configure_logging(level=verbose.upper()) config = parse_yaml(filepath) elements = _parse_elements(element, config) -- 2.52.0 From f9198dcac63d2aa69e496d76418959c3523039e2 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 16:22:25 +0200 Subject: [PATCH 119/287] fix: flake8 fix for api.decorators --- junifer/api/decorators.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/junifer/api/decorators.py b/junifer/api/decorators.py index 91c40af9e..ed624a58d 100644 --- a/junifer/api/decorators.py +++ b/junifer/api/decorators.py @@ -1,6 +1,10 @@ +"""Provide decorators for api.""" + # Authors: Federico Raimondo # Leonard Sasse +# Synchon Mandal # License: AGPL + from .registry import register @@ -24,7 +28,7 @@ def register_datagrabber(klass): def register_marker(klass): - """marker decorator. + """Marker decorator. Registers the marker so it can be used by name. -- 2.52.0 From e05f705b475d46d1c10ee9a6f96546799cd1b3d9 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 16:22:37 +0200 Subject: [PATCH 120/287] fix: flake8 fix for api.functions --- junifer/api/functions.py | 24 ++++++++++++++++-------- 1 file changed, 16 insertions(+), 8 deletions(-) diff --git a/junifer/api/functions.py b/junifer/api/functions.py index e3fb8b80d..fd16a2b28 100644 --- a/junifer/api/functions.py +++ b/junifer/api/functions.py @@ -1,3 +1,5 @@ +"""Provide functions for cli.""" + # Authors: Federico Raimondo # Leonard Sasse # License: AGPL @@ -25,9 +27,8 @@ def _get_datagrabber(datagrabber_config): return datagrabber -def run( - workdir, datagrabber, markers, storage, elements=None): - """Run the pipeline on the selected element +def run(workdir, datagrabber, markers, storage, elements=None): + """Run the pipeline on the selected element. Parameters ---------- @@ -81,6 +82,7 @@ def run( def collect(storage): + """Collect data.""" storage_params = storage.copy() storage_kind = storage_params.pop('kind') logger.info(f'Collecting data using {storage_kind}') @@ -93,9 +95,15 @@ def collect(storage): logger.info('Collect done') -def queue(config, kind, jobname='junifer_job', overwrite=False, elements=None, - **kwargs): - """Queue a job to be executed later +def queue( + config, + kind, + jobname='junifer_job', + overwrite=False, + elements=None, + **kwargs +): + """Queue a job to be executed later. Parameters ---------- @@ -251,9 +259,9 @@ def _queue_condor(jobname, jobdir, yaml_config, elements, env, mem='8G', if collect is True: dag_file.write(f'JOB collect {submit_collect_fname}\n') dag_file.write('PARENT ') - for i_job, t_elem in enumerate(elements): + for i_job, _t_elem in enumerate(elements): dag_file.write(f'run{i_job} ') - dag_file.write(f'CHILD collect\n\n') + dag_file.write('CHILD collect\n\n') if submit is True: logger.info('Submitting HTCondor job') -- 2.52.0 From c3ccc3d0d79ee146ca9d94df21a4e319fd197e61 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 16:22:48 +0200 Subject: [PATCH 121/287] fix: flake8 fix for api.parser --- junifer/api/parser.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/junifer/api/parser.py b/junifer/api/parser.py index a5883ad0a..2cf7d1c33 100644 --- a/junifer/api/parser.py +++ b/junifer/api/parser.py @@ -1,3 +1,5 @@ +"""Provide functions for parser.""" + import yaml import importlib from pathlib import Path @@ -6,6 +8,7 @@ from ..utils.logging import raise_error, logger def parse_yaml(filepath): + """Parse YAML.""" if not isinstance(filepath, Path): filepath = Path(filepath) logger.info(f'Parsing yaml file: {filepath.as_posix()}') -- 2.52.0 From b3baa83f3db2fa89a398d213f456a0b2c8f73000 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 16:22:58 +0200 Subject: [PATCH 122/287] fix: flake8 fix for api.registry --- junifer/api/registry.py | 11 +++++++---- 1 file changed, 7 insertions(+), 4 deletions(-) diff --git a/junifer/api/registry.py b/junifer/api/registry.py index 922433f7c..e7c8f2412 100644 --- a/junifer/api/registry.py +++ b/junifer/api/registry.py @@ -1,6 +1,9 @@ +"""Provide functions for registry.""" + # Authors: Federico Raimondo # Leonard Sasse # License: AGPL + from ..utils.logging import raise_error, logger _valid_steps = [ @@ -10,7 +13,7 @@ _registry = {x: {} for x in _valid_steps} def register(step, name, klass): - """Register a function to be used in a pipeline step + """Register a function to be used in a pipeline step. Parameters ---------- @@ -28,7 +31,7 @@ def register(step, name, klass): def get_step_names(step): - """Get the names of the registered functions for a given step + """Get the names of the registered functions for a given step. Parameters ---------- @@ -46,7 +49,7 @@ def get_step_names(step): def get(step, name): - """Get the class of the registered function for a given step + """Get the class of the registered function for a given step. Parameters ---------- @@ -68,7 +71,7 @@ def get(step, name): def build(step, name, baseclass, init_params=None): - """Ensure that the given object is an instance of the given class + """Ensure that the given object is an instance of the given class. Parameters ---------- -- 2.52.0 From fbd53a6c360f9841dc0b1c35120f97cc60a4c9e9 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 16:23:12 +0200 Subject: [PATCH 123/287] fix: flake8 fix for api.tests.test_cli --- junifer/api/tests/test_cli.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/junifer/api/tests/test_cli.py b/junifer/api/tests/test_cli.py index c29b7262e..db99ab9e1 100644 --- a/junifer/api/tests/test_cli.py +++ b/junifer/api/tests/test_cli.py @@ -1,3 +1,5 @@ +"""Provide tests for cli.""" + from pathlib import Path import tempfile import yaml @@ -10,7 +12,7 @@ runner = CliRunner() def _modify_path(tmpdir, in_file): - """Modify the path to use the temporary directory""" + """Modify the path to use the temporary directory.""" if not isinstance(tmpdir, Path): tmpdir = Path(tmpdir) with open(in_file, 'r') as f: @@ -26,7 +28,7 @@ def _modify_path(tmpdir, in_file): def test_run_collect(): - """Test run and collect""" + """Test run and collect.""" infile = Path(__file__).parent / 'data' / 'gmd_mean.yaml' with tempfile.TemporaryDirectory() as _tmpdir: runfile = _modify_path(_tmpdir, infile) -- 2.52.0 From 203c063192874813ebf77fa35004a8b7bee1f943 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 16:23:29 +0200 Subject: [PATCH 124/287] fix: flake8 fix for api.tests.test_functions --- junifer/api/tests/test_functions.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/junifer/api/tests/test_functions.py b/junifer/api/tests/test_functions.py index ae44675fc..9a00a5e8f 100644 --- a/junifer/api/tests/test_functions.py +++ b/junifer/api/tests/test_functions.py @@ -1,3 +1,5 @@ +"""Provide tests for functions.""" + import tempfile from pathlib import Path @@ -27,7 +29,7 @@ storage = { def test_run(): - """Test run function""" + """Test run function.""" with tempfile.TemporaryDirectory() as tmpdir: tmp_path = Path(tmpdir) workdir = tmp_path / 'workdir' @@ -60,7 +62,7 @@ def test_run(): def test_collect(): - """Test run and collect functions""" + """Test run and collect functions.""" with tempfile.TemporaryDirectory() as tmpdir: tmp_path = Path(tmpdir) -- 2.52.0 From a14c08dd53540bd3bb0c8b1a69dc1ed707227699 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 16:23:41 +0200 Subject: [PATCH 125/287] fix: flake8 fix for api.tests.test_parser --- junifer/api/tests/test_parser.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/junifer/api/tests/test_parser.py b/junifer/api/tests/test_parser.py index 392a90d7f..0d712340a 100644 --- a/junifer/api/tests/test_parser.py +++ b/junifer/api/tests/test_parser.py @@ -1,3 +1,5 @@ +"""Provide tests for parser.""" + import sys from pathlib import Path import tempfile @@ -7,7 +9,7 @@ from junifer.api.parser import parse_yaml def test_parse_yaml(): - """Test parse yaml""" + """Test parse yaml.""" with pytest.raises(ValueError, match='does not exist'): parse_yaml('foo.yaml') -- 2.52.0 From a680d0c9c250850c181e8d72b964bb239ce4ca87 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 16:24:00 +0200 Subject: [PATCH 126/287] fix: flake8 fix for api.tests.test_registry --- junifer/api/tests/test_registry.py | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/junifer/api/tests/test_registry.py b/junifer/api/tests/test_registry.py index d8716f541..afded33d0 100644 --- a/junifer/api/tests/test_registry.py +++ b/junifer/api/tests/test_registry.py @@ -1,3 +1,5 @@ +"""Provide tests for registry.""" + import pytest from abc import ABC @@ -5,13 +7,13 @@ from junifer.api.registry import register, get_step_names, get, build def test_register_error(): - """Test register error""" + """Test register error.""" with pytest.raises(ValueError, match='Invalid ste'): register('foo', 'bar', 'baz') def test_gets(): - """Test get""" + """Test get.""" with pytest.raises(ValueError, match='Invalid ste'): get_step_names('foo') @@ -32,7 +34,7 @@ def test_gets(): def test_build(): - """Test building objects from names""" + """Test building objects from names.""" import numpy as np class SuperClass(ABC): -- 2.52.0 From 65eb01a6717db788474a5501c97bb47abf818230 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 16:24:16 +0200 Subject: [PATCH 127/287] fix: flake8 fix for data.atlases --- junifer/data/atlases.py | 20 ++++++++++++-------- 1 file changed, 12 insertions(+), 8 deletions(-) diff --git a/junifer/data/atlases.py b/junifer/data/atlases.py index b6bb7b3ef..d429bbf82 100644 --- a/junifer/data/atlases.py +++ b/junifer/data/atlases.py @@ -1,6 +1,10 @@ +"""Provide functions for atlases.""" + # Authors: Federico Raimondo # Vera Komeyer +# Synchon Mandal # License: AGPL + from pathlib import Path import io import tempfile @@ -130,8 +134,8 @@ def list_atlases(): def load_atlas(name, atlas_dir=None, resolution=None, path_only=False): - """ - Loads a brain atlas (including a label file). + """Load a brain atlas (including a label file). + If it is built-in atlas and file is not present in the `atlas_dir` directory, it will be downloaded. @@ -213,9 +217,9 @@ def load_atlas(name, atlas_dir=None, resolution=None, path_only=False): def _retrieve_atlas(family, atlas_dir=None, resolution=None, **kwargs): - """ - Retrieves a brain atlas object either from nilearn or a specified online - source. Only returns one atlas per call. Call function multiple times for + """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. @@ -295,7 +299,7 @@ def _closest_resolution(resolution, valid_resolution): valid_resolution = np.array(valid_resolution) if resolution is None: logger.info('Resolution set to None, using highest resolution.') - closest = np.min(valid_resolution) + closest = np.min(valid_resolution) elif any(x <= resolution for x in valid_resolution): # Case 1: get the highest closest resolution closest = np.max(valid_resolution[valid_resolution <= resolution]) @@ -515,8 +519,8 @@ def _retrieve_suit(atlas_path, resolution, space='MNI'): 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 + 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( -- 2.52.0 From fb6854f1180c307152ba975c4e72cff26b640f3e Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 16:24:32 +0200 Subject: [PATCH 128/287] fix: flake8 fix for data.tests.test_atlas --- junifer/data/tests/test_atlas.py | 23 +++++++++++++++-------- 1 file changed, 15 insertions(+), 8 deletions(-) diff --git a/junifer/data/tests/test_atlas.py b/junifer/data/tests/test_atlas.py index e42c02228..0e4cd17e5 100644 --- a/junifer/data/tests/test_atlas.py +++ b/junifer/data/tests/test_atlas.py @@ -1,15 +1,22 @@ +"""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) +from junifer.data.atlases import ( + register_atlas, list_atlases, + load_atlas, + _retrieve_schaefer, + _retrieve_suit, + _retrieve_atlas, + _retrieve_tian, +) def test_register_atlas(): - """Test register_atlas""" + """Test atlas registration.""" atlases = list_atlases() assert 'testatlas' not in atlases @@ -44,7 +51,7 @@ def test_register_atlas(): def test_wrong_atlas(): - """Test invalid atlas""" + """Test invalid atlas.""" with pytest.raises(ValueError, match=r"not found"): load_atlas('wrongatlas') @@ -53,7 +60,7 @@ def test_wrong_atlas(): def test_schaefer_atlas(): - """Test Schaefer atlas""" + """Test Schaefer atlas.""" atlases = list_atlases() @@ -120,7 +127,7 @@ def test_schaefer_atlas(): def test_suit(): - """Test SUIT atlas""" + """Test SUIT atlas.""" atlases = list_atlases() assert 'SUITxSUIT' in atlases @@ -153,7 +160,7 @@ def test_suit(): def test_tian(): - """Test TIAN atlas""" + """Test TIAN atlas.""" atlases = list_atlases() assert 'TianxS1x3TxMNI6thgeneration' in atlases -- 2.52.0 From f4f8ad71345c42456a310740b97da06aa8a82280 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 16:27:14 +0200 Subject: [PATCH 129/287] fix: flake8 fix for stats --- junifer/stats.py | 10 ++++------ 1 file changed, 4 insertions(+), 6 deletions(-) diff --git a/junifer/stats.py b/junifer/stats.py index 7d72aeece..a4ea428c2 100644 --- a/junifer/stats.py +++ b/junifer/stats.py @@ -1,3 +1,5 @@ +"""Provide functions for statistics.""" + from functools import partial import numpy as np from scipy.stats.mstats import winsorize @@ -6,8 +8,7 @@ from .utils import logger, raise_error def get_aggfunc_by_name(name, func_params): - """ - Helper function to get an aggregation function by its name. + """Get an aggregation function by its name. Parameters ---------- @@ -28,7 +29,6 @@ def get_aggfunc_by_name(name, func_params): func : function Respective function with `func_params` parameter set. """ - # check validity of names _valid_func_names = {'winsorized_mean', 'mean', 'std', 'trim_mean'} @@ -56,8 +56,7 @@ def get_aggfunc_by_name(name, func_params): def winsorized_mean(data, axis=None, **win_params): - """ - Compute a winsorized mean by chaining winsorization and mean. + """Compute a winsorized mean by chaining winsorization and mean. Parameters ---------- @@ -73,7 +72,6 @@ def winsorized_mean(data, axis=None, **win_params): Winsorized mean of the inputted data with the winsorize settings applied as specified in win_params. """ - win_dat = winsorize(data, axis=axis, **win_params) win_mean = win_dat.mean(axis=axis) -- 2.52.0 From b0edc9198732abbdd486ec876969298d808982e4 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 16:30:22 +0200 Subject: [PATCH 130/287] fix: flake8 fix for testing.datagrabbers --- junifer/testing/datagrabbers.py | 13 ++++++++++--- 1 file changed, 10 insertions(+), 3 deletions(-) diff --git a/junifer/testing/datagrabbers.py b/junifer/testing/datagrabbers.py index 99c0d0642..549e84e7c 100644 --- a/junifer/testing/datagrabbers.py +++ b/junifer/testing/datagrabbers.py @@ -1,5 +1,9 @@ +"""Provide testing datagrabbers.""" + # Authors: Federico Raimondo +# Synchon Mandal # License: AGPL + import tempfile from nilearn import datasets @@ -7,18 +11,20 @@ from ..datagrabber.base import BaseDataGrabber class OasisVBMTestingDatagrabber(BaseDataGrabber): - """ - DataGrabber for Oasis VBM testing data. - """ + """DataGrabber for Oasis VBM testing data.""" + def __init__(self): + """Initialize class.""" datadir = tempfile.mkdtemp() types = ['VBM_GM'] super().__init__(types=types, datadir=datadir) def get_elements(self): + """Get elements.""" return [f'sub-{x:02d}' for x in list(range(1, 11))] def __getitem__(self, element): + """Get item implementation.""" out = super().__getitem__(element) i_sub = int(element.split('-')[1]) - 1 out['VBM_GM'] = {'path': self._dataset.gray_matter_maps[i_sub]} @@ -27,5 +33,6 @@ class OasisVBMTestingDatagrabber(BaseDataGrabber): return out def __enter__(self): + """Context enter implementation.""" self._dataset = datasets.fetch_oasis_vbm(n_subjects=10) return self -- 2.52.0 From ff85027ecd9a8a2341036f57c4fe560ec24453e2 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 16:31:30 +0200 Subject: [PATCH 131/287] fix: flake8 fix for testing.registry --- junifer/testing/registry.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/junifer/testing/registry.py b/junifer/testing/registry.py index 556005b53..0464c6c76 100644 --- a/junifer/testing/registry.py +++ b/junifer/testing/registry.py @@ -1,3 +1,5 @@ +"""Provide testing registry.""" + from .datagrabbers import OasisVBMTestingDatagrabber from ..api.registry import register -- 2.52.0 From d36b4b553bf13a21dde859a9738a00c031d07c2c Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 16:34:47 +0200 Subject: [PATCH 132/287] fix: flake8 fix for configs.juseless --- junifer/configs/juseless.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/junifer/configs/juseless.py b/junifer/configs/juseless.py index db88315b7..f7f81493a 100644 --- a/junifer/configs/juseless.py +++ b/junifer/configs/juseless.py @@ -1,3 +1,5 @@ +"""Provide class for juseless datagrabber.""" + from ..datagrabber import PatternDataladDataGrabber from ..api.decorators import register_datagrabber -- 2.52.0 From 52db045e4bcaac703e9402f100083573bc35f947 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 16:35:03 +0200 Subject: [PATCH 133/287] fix: flake8 fix for configs.tests.test_juseless --- junifer/configs/tests/test_juseless.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/junifer/configs/tests/test_juseless.py b/junifer/configs/tests/test_juseless.py index 9a266c19c..3234fb86a 100644 --- a/junifer/configs/tests/test_juseless.py +++ b/junifer/configs/tests/test_juseless.py @@ -1,3 +1,5 @@ +"""Provide tests for juseless datagrabber.""" + import socket import pytest @@ -12,6 +14,7 @@ configure_logging(level='DEBUG') def test_juselessdataladukbvbm_datagrabber(): + """Test datalad UKBVBM datagrabber.""" with JuselessDataladUKBVBM() as dg: all_elements = dg.get_elements() test_element = all_elements[0] @@ -23,6 +26,7 @@ def test_juselessdataladukbvbm_datagrabber(): def test_juselessdataladhcp_datagrabber(): + """Test datalad HCP datagrabber.""" with DataladHCP1200() as dg: all_elements = dg.get_elements() test_element = all_elements[0] -- 2.52.0 From a8fa3d60db91a7b098b989862408fed1e14937d8 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 16:52:59 +0200 Subject: [PATCH 134/287] fix: flake8 fix for storage.tests.test_sqlite --- junifer/storage/tests/test_sqlite.py | 19 +++++++++++-------- 1 file changed, 11 insertions(+), 8 deletions(-) diff --git a/junifer/storage/tests/test_sqlite.py b/junifer/storage/tests/test_sqlite.py index a4688497d..25cc75662 100644 --- a/junifer/storage/tests/test_sqlite.py +++ b/junifer/storage/tests/test_sqlite.py @@ -1,3 +1,5 @@ +"""Provide tests for sqlite.""" + from pathlib import Path import numpy as np import pandas as pd @@ -47,7 +49,7 @@ def _read_sql(table_name, uri, index_col): def test_get_engine(): - """Test get_engine""" + """Test engine retrieval.""" with tempfile.TemporaryDirectory() as _tmpdir: uri = f'{_tmpdir}/test.db' # Single storage, must be the uri @@ -72,7 +74,7 @@ def test_get_engine(): def test_store_metadata(): - """Test store_metadata""" + """Test metadata store.""" with tempfile.TemporaryDirectory() as _tmpdir: uri = f'{_tmpdir}/test.db' # Single storage, must be the uri @@ -84,7 +86,7 @@ def test_store_metadata(): def test_upsert_replace(): - """Test store_df (upsert=replace)""" + """Test dataframe store (upsert=replace).""" with tempfile.TemporaryDirectory() as _tmpdir: uri = f'{_tmpdir}/test.db' # Single storage, must be the uri @@ -108,7 +110,7 @@ def test_upsert_replace(): def test_upsert_ignore(): - """Test store_df (upsert=ignore)""" + """Test dataframe store (upsert=ignore).""" with tempfile.TemporaryDirectory() as _tmpdir: uri = f'{_tmpdir}/test.db' with pytest.raises(ValueError): @@ -139,7 +141,7 @@ def test_upsert_ignore(): def test_upsert_update(): - """Test store_df (upsert=delete)""" + """Test dataframe store (upsert=delete).""" meta = {'element': 'test', 'version': '0.0.1'} with tempfile.TemporaryDirectory() as _tmpdir: uri = f'{_tmpdir}/test.db' @@ -162,7 +164,7 @@ def test_upsert_update(): def test_store_read_df(): - """Test store_df""" + """Test dataframe store.""" with tempfile.TemporaryDirectory() as _tmpdir: uri = f'{_tmpdir}/test.db' storage = SQLiteFeatureStorage( @@ -211,7 +213,7 @@ def test_store_read_df(): def test_store_table(): - """Test store_table""" + """Test table store.""" meta = {'element': 'test', 'version': '0.0.1', 'marker': {'name': 'fc'}} with tempfile.TemporaryDirectory() as _tmpdir: uri = f'{_tmpdir}/test.db' @@ -255,7 +257,7 @@ def test_store_table(): def test_store_multiple_output(): - """Test storing using single_output=False""" + """Test storing using single_output=False.""" meta1 = {'element': {'subject': 'test-01', 'session': 'ses-01'}, 'version': '0.0.1', 'marker': {'name': 'fc'}} @@ -329,6 +331,7 @@ def test_store_multiple_output(): def test_collect(): + """Test collect.""" meta1 = {'element': {'subject': 'test-01', 'session': 'ses-01'}, 'version': '0.0.1', 'marker': {'name': 'fc'}} meta2 = {'element': {'subject': 'test-02', 'session': 'ses-01'}, -- 2.52.0 From 1055cb043f3d32ffe2a8d09d5f0ea5c4fe5a54eb Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 16:53:19 +0200 Subject: [PATCH 135/287] fix: flake8 fix for storage.tests.test_base --- junifer/storage/tests/test_base.py | 15 +++++++-------- 1 file changed, 7 insertions(+), 8 deletions(-) diff --git a/junifer/storage/tests/test_base.py b/junifer/storage/tests/test_base.py index 532907250..4c433735f 100644 --- a/junifer/storage/tests/test_base.py +++ b/junifer/storage/tests/test_base.py @@ -1,3 +1,5 @@ +"""Provide tests for base.""" + import pytest from junifer.storage.base import (process_meta, element_to_index, @@ -6,8 +8,7 @@ from junifer.storage.base import (process_meta, element_to_index, def test_process_meta_hash(): - """Test meta_hash""" - + """Test metadata hash.""" meta = None with pytest.raises(ValueError, match=r"Meta must be a dict"): process_meta(meta) @@ -48,8 +49,7 @@ def test_process_meta_hash(): def test_process_meta_element(): - """Test meta element""" - + """Test metadata element.""" meta = {} with pytest.raises(ValueError, match=r"_element_keys"): process_meta(meta) @@ -72,8 +72,7 @@ def test_process_meta_element(): def test_process_meta_index(): - """Test element_to_index""" - + """Test metadata element to index.""" meta = {'noelement': 'foo'} with pytest.raises(ValueError, match=r'meta must contain the key'): element_to_index(meta) @@ -138,7 +137,7 @@ def test_process_meta_index(): def test_BaseFeatureStorage(): - """Test BaseFeatureStorage""" + """Test BaseFeatureStorage.""" with pytest.raises(TypeError, match=r"abstract"): BaseFeatureStorage(uri='/tmp') # type: ignore @@ -211,7 +210,7 @@ def test_BaseFeatureStorage(): def test_element_to_prefix(): - """Test converting element to prefix (for file naming)""" + """Test converting element to prefix (for file naming).""" element = 'sub-01' prefix = element_to_prefix(element) -- 2.52.0 From 06d27ca735ac0754cf1532da255370ac68d5e0a2 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 16:53:33 +0200 Subject: [PATCH 136/287] fix: flake8 fix for storage.sqlite --- junifer/storage/sqlite.py | 27 ++++++++++++++++++++------- 1 file changed, 20 insertions(+), 7 deletions(-) diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py index f398f895e..f6f9c6d8e 100644 --- a/junifer/storage/sqlite.py +++ b/junifer/storage/sqlite.py @@ -1,11 +1,15 @@ +"""Provide class and functions for sqlite.""" + # Authors: Federico Raimondo +# Synchon Mandal # License: AGPL + from pathlib import Path import pandas as pd from pandas.core.base import NoNewAttributesMixin from pandas.io.sql import pandasSQL_builder from sqlalchemy import create_engine, inspect -import tqdm +import tqdm from ..api.decorators import register_storage from .base import (PandasFeatureStoreage, process_meta, element_to_prefix, @@ -15,12 +19,10 @@ from ..utils.logging import warn, logger @register_storage class SQLiteFeatureStorage(PandasFeatureStoreage): - """ - SQLite feature storage. - """ + """SQLite feature storage.""" def __init__(self, uri, single_output=False, upsert='update'): - """Initialise an SQLite feature storage + """Initialise an SQLite feature storage. Parameters ---------- @@ -52,11 +54,13 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): self._valid_inputs = ['table', 'timeseries'] def validate(self, input): + """Validate input.""" if not isinstance(input, list): input = [input] return all(x in self._valid_inputs for x in input) def get_engine(self, meta=None): + """Get engine.""" if meta is None: meta = {} element = meta.get('element', None) @@ -71,6 +75,7 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): return create_engine(uri, echo=False) def list_features(self, return_df=False): + """List features.""" meta_df = pd.read_sql( 'meta', con=self.get_engine(), index_col='meta_md5') out = meta_df @@ -128,6 +133,7 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): return df def store_metadata(self, meta): + """Store metadata.""" t_meta = meta.copy() t_meta.update(self.get_meta()) meta_md5, t_meta_row = process_meta(t_meta) @@ -138,14 +144,18 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): return f'meta_{meta_md5}' def store_matrix2d( - self, data, meta, col_names=None, rows_col_name=None): + self, data, meta, col_names=None, rows_col_name=None + ): + """Store 2D matrix.""" # Same as store_2d, but order is important raise NotImplementedError('store_matrix2d not implemented') def store_table(self, data, meta, columns=None, rows_col_name=None): + """Store table.""" self.store_2d(data, meta, columns, rows_col_name) def store_2d(self, data, meta, columns=None, rows_col_name=None): + """Store 2D dataframe.""" n_rows = len(data) idx = element_to_index( meta, n_rows=n_rows, rows_col_name=rows_col_name) @@ -153,6 +163,7 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): self.store_df(data_df, meta) def store_df(self, df, meta): + """Store dataframe.""" # TODO: Test this function # Check that the index generated by meta matches the one in # the dataframe. @@ -181,9 +192,11 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): self._save_upsert(df, table_name, engine) def store_timeseries(self, data, meta): + """Store timeseries.""" raise NotImplementedError('store_timeseries not implemented') def collect(self): + """Collect data.""" if self.single_output is True: raise ValueError('collect is not implemented for single output') logger.info( @@ -210,7 +223,7 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): out_storage._save_upsert(t_df, table_name, if_exist='nocheck') def _save_upsert(self, df, name, engine=None, if_exist='append'): - """ Implementation of UPSERT functionality. + """Implement of UPSERT functionality. Parameters ---------- -- 2.52.0 From fc00455c94b80b3a1efc9be3bca93c3f95b0c300 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 26 Jul 2022 16:54:23 +0200 Subject: [PATCH 137/287] fix: flake8 fix for storage.base --- junifer/storage/base.py | 34 ++++++++++++++++++++++++---------- 1 file changed, 24 insertions(+), 10 deletions(-) diff --git a/junifer/storage/base.py b/junifer/storage/base.py index 68a83244e..53d98703e 100644 --- a/junifer/storage/base.py +++ b/junifer/storage/base.py @@ -1,3 +1,5 @@ +"""Provide class and functions for storage.""" + # Authors: Federico Raimondo # License: AGPL import numpy as np @@ -11,8 +13,10 @@ from .. import __version__ def process_meta(meta): - """Process the metadata for storage. It removes the "element" key - and adds the "_element_keys" with the keys used to index the element. + """Process the metadata for storage. + + It removes the "element" key and adds the "_element_keys" with the keys + used to index the element. Parameters ---------- @@ -53,7 +57,7 @@ def process_meta(meta): def _meta_hash(meta): - """Compute the md5 hash of the meta + """Compute the md5 hash of the meta. Parameters ---------- @@ -73,7 +77,7 @@ def _meta_hash(meta): def element_to_index(meta, n_rows=1, rows_col_name=None): - """Convert the element meta to index + """Convert the element meta to index. Parameters ---------- @@ -113,6 +117,7 @@ def element_to_index(meta, n_rows=1, rows_col_name=None): def element_to_prefix(element): + """Convert the element meta to prefix.""" logger.debug(f'Converting element {element} to prefix') prefix = 'element' if isinstance(element, tuple): @@ -129,15 +134,15 @@ def element_to_prefix(element): class BaseFeatureStorage(ABC): - """ - Base class for feature storage. - """ + """Base class for feature storage.""" def __init__(self, uri, single_output=False): + """Initialize the class.""" self.uri = uri self.single_output = single_output def get_meta(self): + """Get metadata.""" meta = {} meta['versions'] = { 'junifer': __version__, @@ -162,7 +167,8 @@ class BaseFeatureStorage(ABC): @abstractmethod def list_features(self, return_df=False): - """List the features in the storage + """List the features in the storage. + Parameters ---------- return_df : bool @@ -179,7 +185,7 @@ class BaseFeatureStorage(ABC): @abstractmethod def read_df(self, feature_name=None, feature_md5=None): - """Read the features from the storage + """Read the features from the storage. Returns ------- @@ -190,38 +196,46 @@ class BaseFeatureStorage(ABC): @abstractmethod def store_metadata(self, meta): + """Store metadata.""" raise NotImplementedError('store_metadata not implemented') @abstractmethod def store_matrix2d(self, data, meta, col_names=None, row_names=None): + """Store 2D matrix.""" raise NotImplementedError('store_matrix2d not implemented') @abstractmethod def store_table(self, data, meta, columns=None, rows_col_name=None): + """Store table.""" raise NotImplementedError('store_table not implemented') @abstractmethod def store_df(self, df, meta): + """Store dataframe.""" raise NotImplementedError('store_df not implemented') @abstractmethod def store_timeseries(self, data, meta): + """Store timeseries.""" raise NotImplementedError('store_timeseries not implemented') @abstractmethod def collect(self): + """Collect data.""" raise NotImplementedError('collect not implemented') def __str__(self): + """Represent object as string.""" single = '(single output)' \ if self.single_output is True else '(multiple output)' return f'<{self.__class__.__name__} @ {self.uri} {single}>' class PandasFeatureStoreage(BaseFeatureStorage): + """Store features via pandas.""" def _meta_row(self, meta, meta_md5): - """Convert the meta to a dataframe row""" + """Convert the meta to a dataframe row.""" data_df = {} for k, v in meta.items(): data_df[k] = json.dumps(v, sort_keys=True) -- 2.52.0 From 99705b8a9b9b733b2e03a28055587a0aafe8e03e Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 27 Jul 2022 12:03:19 +0200 Subject: [PATCH 138/287] fix: flake fix for utils.fs --- junifer/utils/fs.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/junifer/utils/fs.py b/junifer/utils/fs.py index 667485bb6..77e0347dc 100644 --- a/junifer/utils/fs.py +++ b/junifer/utils/fs.py @@ -1,7 +1,10 @@ +"""Provide functions for filesystem.""" + import os import stat def make_executable(path): + """Make executable.""" st = os.stat(path) os.chmod(path, st.st_mode | stat.S_IEXEC) -- 2.52.0 From 2eae5edddb2db422479957283acf1a9f4cfc3d5d Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 27 Jul 2022 12:03:31 +0200 Subject: [PATCH 139/287] fix: flake8 fix for utils.logging --- junifer/utils/logging.py | 22 +++++++++++++++------- 1 file changed, 15 insertions(+), 7 deletions(-) diff --git a/junifer/utils/logging.py b/junifer/utils/logging.py index b9ef4471b..375bac230 100644 --- a/junifer/utils/logging.py +++ b/junifer/utils/logging.py @@ -1,5 +1,8 @@ +"""Provide class and functions for logging.""" + # Authors: Federico Raimondo # License: AGPL + import logging import subprocess import sys @@ -13,7 +16,7 @@ logger = logging.getLogger('JUNIFER') def _get_git_head(path): - """Aux function to read HEAD from git""" + """Aux function to read HEAD from git.""" if not path.exists(): raise ValueError('This path does not exist: {}'.format(path)) command = ('cd {gitpath}; ' @@ -27,7 +30,8 @@ def _get_git_head(path): def get_versions(sys): - """Import stuff and get versions if module + """Import stuff and get versions if module. + Parameters ---------- sys : module @@ -58,7 +62,7 @@ def get_versions(sys): def get_ext_versions(tbox_path): - """ Get versions of external tools used by JUNIFER.""" + """Get versions of external tools used by JUNIFER.""" versions = {} # spm_path = tbox_path / 'spm12' # if spm_path.exists(): @@ -74,6 +78,7 @@ def _safe_log(versions, name): def log_versions(tbox_path=None): + """Log versions of dependencies and junifer.""" versions = get_versions(sys) logger.info('===== Lib Versions =====') @@ -97,9 +102,10 @@ _logging_types = dict(DEBUG=logging.DEBUG, INFO=logging.INFO, WARNING=logging.WARNING, ERROR=logging.ERROR) -def configure_logging(level='WARNING', fname=None, overwrite=None, - output_format=None): - """Configure the logging functionality +def configure_logging( + level='WARNING', fname=None, overwrite=None, output_format=None +): + """Configure the logging functionality. Parameters ---------- @@ -161,12 +167,14 @@ def _close_handlers(logger): def raise_error(msg, klass=ValueError): + """Raise error.""" logger.error(msg) raise klass(msg) def warn(msg, category=RuntimeWarning): - """Warn, but first log it + """Warn, but first log it. + Parameters ---------- msg : str -- 2.52.0 From 8dc19912c7e088d24854840bec8f35ea9106ffa1 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 27 Jul 2022 12:03:41 +0200 Subject: [PATCH 140/287] fix: flake8 fix for utils.tests.test_logging --- junifer/utils/tests/test_logging.py | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/junifer/utils/tests/test_logging.py b/junifer/utils/tests/test_logging.py index 1281fb8f7..c5a9a41a7 100644 --- a/junifer/utils/tests/test_logging.py +++ b/junifer/utils/tests/test_logging.py @@ -1,6 +1,9 @@ +"""Provide tests for logging.""" + # Authors: Federico Raimondo # Sami Hamdan # License: AGPL + from junifer.utils import logger, configure_logging, raise_error, warn from junifer.utils.logging import _close_handlers import pytest @@ -9,7 +12,7 @@ from pathlib import Path def test_log_file(): - """Test logging to a file""" + """Test logging to a file.""" with tempfile.TemporaryDirectory() as tmp: tmpdir = Path(tmp) configure_logging(fname=tmpdir / 'test1.log') @@ -112,13 +115,13 @@ def test_log_file(): def test_log(): - """Simple log test""" + """Simple log test.""" configure_logging() logger.info('Testing') def test_lib_logging(): - """Test logging versions""" + """Test logging versions.""" import numpy as np # noqa import pandas # noqa -- 2.52.0 From 2fe9d7f37fb25b495488180447db7362b67e8953 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 27 Jul 2022 12:04:01 +0200 Subject: [PATCH 141/287] fix: flake8 fix for datagrabber.base --- junifer/datagrabber/base.py | 75 +++++++++++++++++++++++-------------- 1 file changed, 46 insertions(+), 29 deletions(-) diff --git a/junifer/datagrabber/base.py b/junifer/datagrabber/base.py index bf54f10ad..d65ebc34e 100644 --- a/junifer/datagrabber/base.py +++ b/junifer/datagrabber/base.py @@ -1,6 +1,9 @@ +"""Provide class and functions for base datagrabber.""" + # Authors: Federico Raimondo # Leonard Sasse # License: AGPL + from pathlib import Path import re import tempfile @@ -13,9 +16,7 @@ from ..utils.logging import logger, raise_error, warn def _validate_types(types): - """ - Validate the types - """ + """Validate the types.""" if not isinstance(types, list): raise_error("types must be a list", TypeError) # type: ignore if any(not isinstance(x, str) for x in types): @@ -24,9 +25,7 @@ def _validate_types(types): def _validate_replacements(replacements, patterns): - """ - Validate the replacements - """ + """Validate the replacements.""" if not isinstance(replacements, list): raise_error("replacements must be a list", TypeError) # type: ignore if any(not isinstance(x, str) for x in replacements): @@ -35,14 +34,12 @@ def _validate_replacements(replacements, patterns): TypeError) # type: ignore for x in replacements: - if all(f'{{x}}' not in y for y in patterns.values()): + if all(x not in y for y in patterns.values()): warn(f"Replacement {x} is not part of any pattern") def _validate_patterns(types, patterns): - """ - Validate the patterns. - """ + """Validate the patterns.""" _validate_types(types) if not isinstance(patterns, dict): raise_error("patterns must be a dict", TypeError) # type: ignore @@ -77,6 +74,7 @@ class BaseDataGrabber(ABC): Does nothing. Can be overridden by subclasses to clean up after `__enter__` """ + def __init__(self, types, datadir): """Initialize a BaseDataGrabber object. @@ -97,9 +95,11 @@ class BaseDataGrabber(ABC): self.types = types def get_types(self): + """Get types.""" return self.types.copy() def get_meta(self): + """Get metadata.""" t_meta = {} t_meta['class'] = self.__class__.__name__ for k, v in vars(self).items(): @@ -108,11 +108,13 @@ class BaseDataGrabber(ABC): return t_meta def get_element_keys(self): + """Get element keys.""" return 'element' @property def datadir(self): - """ + """Get data directory path. + Returns ------- Path to the data directory. Implemented as a property, can be @@ -132,6 +134,7 @@ class BaseDataGrabber(ABC): yield elem def __getitem__(self, element): + """Get item implementation.""" logger.info(f'Getting element {element}') out = {} out['meta'] = dict(datagrabber=self.get_meta()) @@ -139,21 +142,25 @@ class BaseDataGrabber(ABC): @abstractmethod def get_elements(self): + """Get elements.""" raise_error( 'get_elements not implemented', NotImplementedError) # type: ignore def __enter__(self): + """Context entry implementation.""" return self def __exit__(self, exc_type, exc_value, exc_traceback): + """Context exit implementation.""" pass @register_datagrabber class PatternDataGrabber(BaseDataGrabber): - """Patternd DataGrabber class (abstract). Implements a DataGrabber that - understands patterns to grab data. + """Pattern DataGrabber class (abstract). + + Implements a DataGrabber that understands patterns to grab data. Attributes ---------- @@ -176,6 +183,7 @@ class PatternDataGrabber(BaseDataGrabber): specified element. Each occurrence of the string `{subject}` is replaced by the indexed element """ + def __init__(self, types=None, patterns=None, replacements=None, **kwargs): """Initialize a BaseDataGrabber object. @@ -202,8 +210,9 @@ class PatternDataGrabber(BaseDataGrabber): self.replacements = replacements def _replace_patterns_regex(self, pattern): - """Replace the patterns in the pattern with the named groups so the - elements can be obtained from the filesystem. + """Replace the patterns in the pattern with the named groups. + + It allows elements to be obtained from the filesystem. Parameters ---------- @@ -231,14 +240,15 @@ class PatternDataGrabber(BaseDataGrabber): re_pattern = re_pattern.replace(f'{{{t_r}}}', f'(?P={t_r})') for t_r in self.replacements: - glob_pattern = glob_pattern.replace(f'{{{t_r}}}', f'*') + glob_pattern = glob_pattern.replace(f'{{{t_r}}}', '*') return re_pattern, glob_pattern def get_elements(self): - """Get the list of elements in the dataset. It will use regex - to search for `replacements` in the `patterns` and return the - intersection of the results for each type. That is, build a list - of elements that have all the required types. + """Get the list of elements in the dataset. + + It will use regex to search for `replacements` in the `patterns` and + return the intersection of the results for each type. That is, build a + list of elements that have all the required types. Returns ------- @@ -266,8 +276,7 @@ class PatternDataGrabber(BaseDataGrabber): return list(elements) def _replace_patterns_glob(self, element, pattern): - """Replace the patterns in the pattern with the element so it can - be globbed. + """Replace patterns with the element so it can be globbed. Parameters ---------- @@ -333,9 +342,9 @@ class PatternDataGrabber(BaseDataGrabber): @register_datagrabber class DataladDataGrabber(BaseDataGrabber): - """ - Datalad DataGrabber class (abstract). Implements a DataGrabber that gets - data from a datalad sibling. + """Datalad DataGrabber class (abstract). + + Implements a DataGrabber that gets data from a datalad sibling. Attributes ---------- @@ -362,6 +371,7 @@ class DataladDataGrabber(BaseDataGrabber): concrete class implementation. """ + def __init__(self, rootdir='.', datadir=None, uri=None, **kwargs): """Initialize a DataladDataGrabber object. @@ -390,11 +400,13 @@ class DataladDataGrabber(BaseDataGrabber): self._rootdir = rootdir def __enter__(self): + """Context entry implementation.""" self.install() return self @property def datadir(self): + """Get data directory path.""" return super().datadir / self._rootdir def install(self): @@ -405,6 +417,7 @@ class DataladDataGrabber(BaseDataGrabber): logger.debug('Dataset installed') def __exit__(self, exc_type, exc_value, exc_traceback): + """Context exit implementation.""" logger.debug('Removing dataset') self.remove() logger.debug('Dataset removed') @@ -418,7 +431,7 @@ class DataladDataGrabber(BaseDataGrabber): if 'path' in v: logger.debug(f'Getting {v["path"]}') self._dataset.get(v['path']) - logger.debug(f'Get done') + logger.debug('Get done') # append the version of the dataset out['meta']['datagrabber']['dataset_commit_id'] = \ @@ -427,9 +440,10 @@ class DataladDataGrabber(BaseDataGrabber): return out def __getitem__(self, element): - """Index one element in the Datalad database. It will first obtain - the paths from the parent class and then `datalad get` each of the - files. + """Index one element in the Datalad database. + + It will first obtain the paths from the parent class and then + `datalad get` each of the files. This method only works with multiple inheritance. @@ -442,6 +456,7 @@ class DataladDataGrabber(BaseDataGrabber): @register_datagrabber class PatternDataladDataGrabber(DataladDataGrabber, PatternDataGrabber): """Pattern-based Datalad DataGrabber class (abstract). + Implements a DataGrabber that gets data from a datalad sibling, interpreting patterns. @@ -452,7 +467,9 @@ class PatternDataladDataGrabber(DataladDataGrabber, PatternDataGrabber): PatternDataGrabber """ + def __init__(self, types=None, patterns=None, **kwargs): + """Initialize the class.""" _validate_patterns(types, patterns) super().__init__(types=types, patterns=patterns, **kwargs) self.patterns = patterns -- 2.52.0 From 04e42ba50dfbff6d07990acd3ffcb71f1dab4e37 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 27 Jul 2022 12:04:14 +0200 Subject: [PATCH 142/287] fix: flake8 fix for datagrabber.hcp --- junifer/datagrabber/hcp.py | 19 ++++++++++++------- 1 file changed, 12 insertions(+), 7 deletions(-) diff --git a/junifer/datagrabber/hcp.py b/junifer/datagrabber/hcp.py index f3ef35215..dd9b801f2 100644 --- a/junifer/datagrabber/hcp.py +++ b/junifer/datagrabber/hcp.py @@ -1,3 +1,5 @@ +"""Provide classes for HCP data access.""" + from itertools import product from junifer.datagrabber.base import DataladDataGrabber @@ -8,6 +10,7 @@ from ..api.decorators import register_datagrabber @register_datagrabber class HCP1200(PatternDataGrabber): + """PatternDataGrabber implementation for HCP1200.""" def __init__( self, datadir=None, tasks=None, phase_encodings=None @@ -30,7 +33,6 @@ class HCP1200(PatternDataGrabber): (default) both will be used. """ - types = ['BOLD'] replacements = ['subject', 'task', 'phase_encoding'] @@ -84,7 +86,6 @@ class HCP1200(PatternDataGrabber): 'phase_encoding can be LR or RL (or both)!' ) - def get_elements(self): """Get the list of subjects in the dataset. @@ -93,7 +94,6 @@ class HCP1200(PatternDataGrabber): elements : list[str] The list of subjects in the dataset. """ - subjects = [x.name for x in self.datadir.iterdir() if x.is_dir()] elems = [] for subject, task, phase_encoding in product( @@ -119,7 +119,6 @@ class HCP1200(PatternDataGrabber): Dictionary of paths for each type of data required for the specified element. """ - sub, task, phase_encoding = element if "REST" in task: @@ -137,12 +136,18 @@ class HCP1200(PatternDataGrabber): @register_datagrabber class DataladHCP1200(DataladDataGrabber, HCP1200): + """DataladDataGrabber implementation for HCP1200.""" + def __init__(self, datadir=None, tasks=None, phase_encodings=None): + """Initialize the class.""" uri = ( 'https://github.com/datalad-datasets/' 'human-connectome-project-openaccess.git' ) rootdir = 'HCP1200' - super().__init__(datadir=datadir, tasks=tasks, - phase_encodings=phase_encodings, - uri=uri, rootdir=rootdir) # type: ignore + super().__init__( + datadir=datadir, + tasks=tasks, + phase_encodings=phase_encodings, + uri=uri, rootdir=rootdir + ) # type: ignore -- 2.52.0 From b101686a0d8380c24d5e6bfd44ac7fc138e85925 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 27 Jul 2022 12:04:31 +0200 Subject: [PATCH 143/287] fix: flake8 fix for datagrabber.meta --- junifer/datagrabber/meta.py | 10 ++++++++++ 1 file changed, 10 insertions(+) diff --git a/junifer/datagrabber/meta.py b/junifer/datagrabber/meta.py index 1b20949d2..100475b65 100644 --- a/junifer/datagrabber/meta.py +++ b/junifer/datagrabber/meta.py @@ -1,31 +1,41 @@ +"""Provide class for metadata collection.""" + from .base import BaseDataGrabber class MultipleDataGrabber(BaseDataGrabber): + """MultipleDataGrabber implementation.""" + def __init__(self, datagrabbers): + """Initialize the class.""" # TODO: Check datagrabbers consistency # - same element keys # - no overlapping types self._datagrabbers = datagrabbers def get_types(self): + """Get types.""" types = [x for dg in self._datagrabbers for x in dg.get_types()] return types def get_meta(self): + """Get metadata.""" t_meta = {} t_meta['class'] = self.__class__.__name__ t_meta['datagrabbers'] = [dg.get_meta() for dg in self._datagrabbers] def __enter__(self): + """Context entry implementation.""" for dg in self._datagrabbers: dg.__enter__() def __exit__(self, exc_type, exc_value, exc_traceback): + """Context exit implementation.""" for dg in self._datagrabbers: dg.__exit__(exc_type, exc_value, exc_traceback) def __getitem__(self, element): + """Get item implementation.""" out = {} for dg in self._datagrabbers: t_out = dg[element] -- 2.52.0 From 438809753ac4fe82bf6bc1ef158f5ad0071ef1cb Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 27 Jul 2022 12:04:43 +0200 Subject: [PATCH 144/287] fix: flake8 fix for datagrabber.tests.test_base_datagrabber --- .../tests/test_base_datagrabber.py | 24 ++++++++++++------- 1 file changed, 15 insertions(+), 9 deletions(-) diff --git a/junifer/datagrabber/tests/test_base_datagrabber.py b/junifer/datagrabber/tests/test_base_datagrabber.py index 689dd24a6..330687e0f 100644 --- a/junifer/datagrabber/tests/test_base_datagrabber.py +++ b/junifer/datagrabber/tests/test_base_datagrabber.py @@ -1,11 +1,17 @@ +"""Provide tests for base datagrabber.""" + # Authors: Federico Raimondo # Leonard Sasse # License: AGPL + import tempfile import pytest from pathlib import Path -from junifer.datagrabber.base import (PatternDataGrabber, BaseDataGrabber, - PatternDataladDataGrabber) +from junifer.datagrabber.base import ( + PatternDataGrabber, + BaseDataGrabber, + PatternDataladDataGrabber, +) _testing_dataset = { @@ -21,7 +27,7 @@ _testing_dataset = { def test_BaseDataGrabber(): - """Test BaseDataGrabber""" + """Test BaseDataGrabber.""" with pytest.raises(TypeError, match=r"abstract"): BaseDataGrabber(datadir='/tmp', types=['func']) # type: ignore @@ -48,11 +54,11 @@ def test_BaseDataGrabber(): def test_PatternDataGrabber(): + """Test PatternDataGrabber.""" class MyDataGrabber(PatternDataGrabber): def get_elements(self): return super().get_elements() - """Test test_PatternDataGrabber""" with pytest.raises(TypeError, match=r"types must be a list"): MyDataGrabber(datadir='/tmp', types='wrong', patterns=dict(wrong='pattern'), @@ -115,7 +121,7 @@ def test_PatternDataGrabber(): def test_bids_datalad_PatternDataGrabber(): - """Test a subject-based BIDS datalad datagrabber""" + """Test a subject-based BIDS datalad datagrabber.""" types = ['T1w', 'bold'] patterns = { 'T1w': '{subject}/anat/{subject}_T1w.nii.gz', @@ -127,9 +133,9 @@ def test_bids_datalad_PatternDataGrabber(): PatternDataladDataGrabber(datadir=None, types=types, patterns=patterns, replacements=replacements) - repo_uri = _testing_dataset['example_bids']['uri'] + repo_uri = _testing_dataset['example_bids']['uri'] rootdir = 'example_bids' - repo_commit = _testing_dataset['example_bids']['id'] + repo_commit = _testing_dataset['example_bids']['id'] with PatternDataladDataGrabber( rootdir=rootdir, uri=repo_uri, types=types, patterns=patterns, @@ -181,7 +187,7 @@ def test_bids_datalad_PatternDataGrabber(): def test_bids_datalad_PatternDataGrabber_session(): - """Test a subject and session-based BIDS datalad datagrabber""" + """Test a subject and session-based BIDS datalad datagrabber.""" types = ['T1w', 'bold'] patterns = { 'T1w': '{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz', @@ -194,7 +200,7 @@ def test_bids_datalad_PatternDataGrabber_session(): PatternDataladDataGrabber(datadir=None, types=types, patterns=patterns, replacements=replacements) - repo_uri = _testing_dataset['example_bids_ses']['uri'] + repo_uri = _testing_dataset['example_bids_ses']['uri'] rootdir = 'example_bids_ses' # repo_commit = _testing_dataset['example_bids_ses']['id'] -- 2.52.0 From bfc89ea5d0db8d0dbcbcc76270b6dbb5b836acc7 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 27 Jul 2022 12:06:12 +0200 Subject: [PATCH 145/287] fix: flake8 fix for tests.test_main --- junifer/tests/test_main.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/junifer/tests/test_main.py b/junifer/tests/test_main.py index b5fa71ce7..ea7f2b76c 100644 --- a/junifer/tests/test_main.py +++ b/junifer/tests/test_main.py @@ -1,6 +1,11 @@ +"""Provide tests for package.""" + # Authors: Federico Raimondo # Leonard Sasse # License: AGPL + + def test_import(): + """Test junifer import.""" import junifer print(junifer.__version__) -- 2.52.0 From 57ff4aec5ade8eab1cc1aac1bedf0868e34bd24b Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 27 Jul 2022 12:22:47 +0200 Subject: [PATCH 146/287] fix: flake8 fix for preprocess.confounds --- junifer/preprocess/confounds.py | 34 ++++++++++++++++++++++----------- 1 file changed, 23 insertions(+), 11 deletions(-) diff --git a/junifer/preprocess/confounds.py b/junifer/preprocess/confounds.py index f1e3f20aa..ea53097c7 100644 --- a/junifer/preprocess/confounds.py +++ b/junifer/preprocess/confounds.py @@ -1,3 +1,5 @@ +"""Provide classes for confound removal.""" + # Authors: Federico Raimondo # Leonard Sasse # License: AGPL @@ -14,8 +16,10 @@ from ..utils.logging import logger, raise_error class BaseConfoundRemover(PipelineStepMixin): - """ A base class to read confound files and select columns according to - a pre-defined strategy + """Base class for cofound removal. + + Read confound files and select columns according to + a pre-defined strategy. """ @@ -25,10 +29,18 @@ class BaseConfoundRemover(PipelineStepMixin): # in particular scrubbing # TODO: Implement read_confounds for fmriprep data - def __init__(self, strategy=None, spike=None, detrend=True, - standardize=True, low_pass=None, high_pass=None, t_r=None, - mask_img=None): - """ Initialise a BaseConfoundReader object + def __init__( + self, + strategy=None, + spike=None, + detrend=True, + standardize=True, + low_pass=None, + high_pass=None, + t_r=None, + mask_img=None, + ): + """Initialise a BaseConfoundReader object. Confound removal is based on nilearn.image.clean_img @@ -157,7 +169,7 @@ class BaseConfoundRemover(PipelineStepMixin): return input def _pick_confounds(self, input): - """ Select relevant confounds from the specified file """ + """Select relevant confounds from the specified file.""" to_select = [] confounds_df = input['data'] confounds_spec = input['names']['spec'] @@ -179,13 +191,13 @@ class BaseConfoundRemover(PipelineStepMixin): out_df[t_dst] = np.append( # type: ignore np.diff(out_df[t_src]), 0) # type: ignore - # Add squares (of base confounds and derivatives) if needed + # Add squares (of base confounds and derivatives) if needed to_compute = [x in squares_to_compute.keys() for x in to_select] if any(to_compute): for t_dst, t_src in squares_to_compute.items(): out_df[t_dst] = out_df[t_src] ** 2 out_df = out_df[to_select] - + # add binary spike regressor if needed at given threshold if self.spike is not None: fd = confounds_df[spike_name].copy() @@ -196,7 +208,7 @@ class BaseConfoundRemover(PipelineStepMixin): return out_df def _remove_confounds(self, bold_img, confounds_df): - """ Remove confounds from the BOLD image + """Remove confounds from the BOLD image. bold_img : Niimg-like object 4D image. The signals in the last dimension are filtered @@ -212,7 +224,6 @@ class BaseConfoundRemover(PipelineStepMixin): input image, cleaned. """ - confounds_array = confounds_df.values t_r = self.t_r @@ -340,6 +351,7 @@ class BaseConfoundRemover(PipelineStepMixin): 'and the datagrabber.', ValueError) def fit_transform(self, input): + """Fit and transform.""" self._validate_data(input) bold_img = input['BOLD']['data'] confounds_df = self._pick_confounds(input['confounds']) -- 2.52.0 From 8e6f7aaf6d7d4e29eb319141fdb61f825b3c3aaa Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 27 Jul 2022 12:22:58 +0200 Subject: [PATCH 147/287] fix: flake8 fix for preprocess.tests.test_confounds --- junifer/preprocess/tests/test_confounds.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/junifer/preprocess/tests/test_confounds.py b/junifer/preprocess/tests/test_confounds.py index ede77c861..69afa9f76 100644 --- a/junifer/preprocess/tests/test_confounds.py +++ b/junifer/preprocess/tests/test_confounds.py @@ -1,3 +1,5 @@ +"""Provide tests for confound removal.""" + # Authors: Federico Raimondo # Leonard Sasse # License: AGPL @@ -15,11 +17,11 @@ np.random.seed(1234567) def generate_conf_name(size=6, chars=string.ascii_uppercase + string.digits): + """Generate configuration name.""" return ''.join(random.choice(chars) for _ in range(size)) def _simu_img(): - # Random 4D volume with 100 time points vol = 100 + 10 * np.random.randn(5, 5, 2, 100) img = Nifti1Image(vol, np.eye(4)) @@ -29,6 +31,7 @@ def _simu_img(): def test_baseconfoundremover(): + """Test BaseConfoundRemover.""" # Generate a simulated BOLD img siimg, simsk = _simu_img() @@ -69,7 +72,7 @@ def test_baseconfoundremover(): confound_column_names.extend(gs_full) # add some random irrelevant confounds - for i in range(10): + for _ in range(10): confound_column_names.append(generate_conf_name()) np.random.shuffle(confound_column_names) -- 2.52.0 From e5dfa942b4f17e288255efe115a277cedd8ac51d Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 27 Jul 2022 12:25:33 +0200 Subject: [PATCH 148/287] fix: flake8 fix for datareader.default --- junifer/datareader/default.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/junifer/datareader/default.py b/junifer/datareader/default.py index 226c4d608..d691f2849 100644 --- a/junifer/datareader/default.py +++ b/junifer/datareader/default.py @@ -1,3 +1,5 @@ +"""Provide class for default data reader.""" + # Authors: Federico Raimondo # License: AGPL @@ -26,16 +28,20 @@ _readers['TSV'] = dict(func=pd.read_csv, params={'sep': '\t'}) class DefaultDataReader(PipelineStepMixin): + """Mixin class for default data reader.""" def validate_input(self, input): + """Validate input.""" # Nothing to validate, any input is fine pass def get_output_kind(self, input): + """Get output kind.""" # It will output the same kind of data as the input return input def fit_transform(self, input, params=None): + """Fit and transform.""" # For each kind of data, try to read it # out is the same, but with the 'data' key set in -- 2.52.0 From 013cbefb5706446ce7cc53e7a24514eeb91784a1 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 27 Jul 2022 12:25:46 +0200 Subject: [PATCH 149/287] fix: flake8 fix for datareader.tests.test_default_reader --- junifer/datareader/tests/test_default_reader.py | 12 +++++++----- 1 file changed, 7 insertions(+), 5 deletions(-) diff --git a/junifer/datareader/tests/test_default_reader.py b/junifer/datareader/tests/test_default_reader.py index 5ad9ccb39..ab4867159 100644 --- a/junifer/datareader/tests/test_default_reader.py +++ b/junifer/datareader/tests/test_default_reader.py @@ -1,3 +1,5 @@ +"""Provide tests for default data reader.""" + # Authors: Federico Raimondo # License: AGPL @@ -13,7 +15,7 @@ from junifer.datareader import DefaultDataReader def test_validation(): - """Test validating input/output""" + """Test validating input/output.""" kinds = [ ['T1w', 'BOLD', 'T2', 'dwi'], [], @@ -30,7 +32,7 @@ def test_validation(): def test_meta(): - """Test reader metadata""" + """Test reader metadata.""" reader = DefaultDataReader() t_meta = reader.get_meta() assert t_meta['class'] == 'DefaultDataReader' @@ -46,7 +48,7 @@ def test_meta(): def test_read_nifti(): - """Test reading NIFTI files""" + """Test reading NIFTI files.""" reader = DefaultDataReader() nib_data_path = Path(nib_testing.data_path) @@ -74,7 +76,7 @@ def test_read_nifti(): def test_read_unknown(): - """Test (not) reading unknown files""" + """Test (not) reading unknown files.""" reader = DefaultDataReader() nib_data_path = Path(nib_testing.data_path) @@ -100,7 +102,7 @@ def test_read_unknown(): def test_read_csv(): - """Test reading CSV files""" + """Test reading CSV files.""" d = {'col1': [1, 2, 3, 4, 5], 'col2': [3, 4, 5, 6, 7]} df = pd.DataFrame(d) with tempfile.TemporaryDirectory() as tmpdir: -- 2.52.0 From c5100272ae0b9a52a7a6450b77c670f06614c38a Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 27 Jul 2022 12:33:06 +0200 Subject: [PATCH 150/287] fix: flake8 fix for markers.base --- junifer/markers/base.py | 13 +++++++++++++ 1 file changed, 13 insertions(+) diff --git a/junifer/markers/base.py b/junifer/markers/base.py index bcec81d9f..3357123af 100644 --- a/junifer/markers/base.py +++ b/junifer/markers/base.py @@ -1,11 +1,16 @@ +"""Provide base class and mixin class for markers.""" + # Authors: Federico Raimondo # License: AGPL + from ..utils import logger class PipelineStepMixin(): + """Mixin class for pipeline.""" def get_meta(self): + """Get metadata.""" t_meta = {} t_meta['class'] = self.__class__.__name__ for k, v in vars(self).items(): @@ -68,6 +73,7 @@ class PipelineStepMixin(): return self.get_output_kind(input) def fit_transform(self, input): + """Fit and transform.""" raise NotImplementedError('fit_transform not implemented') @@ -75,12 +81,14 @@ class BaseMarker(PipelineStepMixin): """Base class for all markers.""" def __init__(self, on, name=None): + """Initialize the class.""" if not isinstance(on, list): on = [on] self._valid_inputs = on self.name = self.__class__.__name__ if name is None else name def get_meta(self, kind): + """Get metadata.""" s_meta = super().get_meta() # same marker can be fit into different kinds, so the name # is created from the kind and the name of the marker @@ -89,6 +97,7 @@ class BaseMarker(PipelineStepMixin): return dict(marker=s_meta) def validate_input(self, input): + """Validate input.""" if not any(x in input for x in self._valid_inputs): raise ValueError( 'Input does not have the required data.' @@ -96,15 +105,19 @@ class BaseMarker(PipelineStepMixin): f'\t Required (any of): {self._valid_inputs}') def get_output_kind(self, input): + """Get output kind.""" return None def compute(self, input): + """Compute.""" raise NotImplementedError('compute not implemented') def store(self, input, out, storage): + """Store.""" raise NotImplementedError('store not implemented') def fit_transform(self, input, storage=None): + """Fit and transform.""" out = {} meta = input.get('meta', {}) for kind in self._valid_inputs: -- 2.52.0 From 9f98d866d0846c4a6bae7a64225a245a4818f466 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 27 Jul 2022 12:33:14 +0200 Subject: [PATCH 151/287] fix: flake8 fix for markers.collection --- junifer/markers/collection.py | 14 ++++++++++---- 1 file changed, 10 insertions(+), 4 deletions(-) diff --git a/junifer/markers/collection.py b/junifer/markers/collection.py index fefd2d6ec..84c812cfd 100644 --- a/junifer/markers/collection.py +++ b/junifer/markers/collection.py @@ -1,3 +1,5 @@ +"""Provide class for marker collection.""" + # Authors: Federico Raimondo # License: AGPL @@ -8,8 +10,12 @@ from collections import Counter class MarkerCollection(): - def __init__(self, markers, datareader=None, preprocessing=None, - storage=None): + """Class for marker collection.""" + + def __init__( + self, markers, datareader=None, preprocessing=None, storage=None + ): + """Initialize the class.""" if datareader is None: datareader = DefaultDataReader() self._datareader = datareader @@ -52,7 +58,7 @@ class MarkerCollection(): m_value = marker.fit_transform(data, storage=self._storage) if self._storage is None: out[marker.name] = m_value - logger.info(f'Marker collection fitting done') + logger.info('Marker collection fitting done') return None if self._storage else out def validate(self, datagrabber): @@ -67,7 +73,7 @@ class MarkerCollection(): t_data = datagrabber.get_types() logger.info(f'DataGrabber output type: {t_data}') - logger.info(f'Validating Data Reader:') + logger.info('Validating Data Reader:') t_data = self._datareader.validate(t_data) logger.info(f'Data Reader output type: {t_data}') -- 2.52.0 From 85c9bd0d0270bd6c6d10b42538c29264a93e8fc2 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 27 Jul 2022 12:33:30 +0200 Subject: [PATCH 152/287] fix: flake8 fix for markers.parcel --- junifer/markers/parcel.py | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/junifer/markers/parcel.py b/junifer/markers/parcel.py index defc9bb3d..8b86f3009 100644 --- a/junifer/markers/parcel.py +++ b/junifer/markers/parcel.py @@ -1,3 +1,5 @@ +"""Provide class for parcel aggregation.""" + # Authors: Federico Raimondo # License: AGPL import numpy as np @@ -14,7 +16,10 @@ from ..utils import logger @register_marker class ParcelAggregation(BaseMarker): + """Class for parcel aggregation.""" + def __init__(self, atlas, method, method_params=None, on=None, name=None): + """Initialize the class.""" if on is None: on = ['T1w', 'BOLD', 'VBM_GM', 'VBM_WM', 'fALFF', 'GCOR', 'LCOR'] super().__init__(on=on, name=name) @@ -23,12 +28,14 @@ class ParcelAggregation(BaseMarker): self.method_params = {} if method_params is None else method_params def get_output_kind(self, input): + """Get output kind.""" if input in ['VBM_GM', 'VBM_WM', 'fALFF', 'GCOR', 'LCOR']: return 'table' if input in ['BOLD']: return 'timeseries' def store(self, kind, out, storage): + """Store.""" logger.debug(f'Storing {kind} in {storage}') if kind in ['VBM_GM', 'VBM_WM', 'fALFF', 'GCOR', 'LCOR']: storage.store_table(**out) @@ -36,6 +43,7 @@ class ParcelAggregation(BaseMarker): storage.store_timeseries(**out) def compute(self, input): + """Compute.""" t_input = input['data'] logger.debug(f'Parcel aggregation using {self.method}') agg_func = get_aggfunc_by_name( -- 2.52.0 From 037b8f4b87b950be99e32f992fb3196e2772f9fd Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 27 Jul 2022 12:33:41 +0200 Subject: [PATCH 153/287] fix: flake8 fix for markers.tests.test_base_marker --- junifer/markers/tests/test_base_marker.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/junifer/markers/tests/test_base_marker.py b/junifer/markers/tests/test_base_marker.py index a0d43c81c..ccfdd039a 100644 --- a/junifer/markers/tests/test_base_marker.py +++ b/junifer/markers/tests/test_base_marker.py @@ -1,8 +1,11 @@ +"""Provide tests for base marker.""" + import pytest from junifer.markers.base import BaseMarker, PipelineStepMixin def test_PipelineStepMixin(): + """Test PipelineStepMixin.""" mixin = PipelineStepMixin() with pytest.raises(NotImplementedError): mixin.validate_input(None) @@ -13,7 +16,7 @@ def test_PipelineStepMixin(): def test_meta(): - """Test metadata""" + """Test metadata.""" pipemixin = PipelineStepMixin() t_meta = pipemixin.get_meta() assert t_meta['class'] == 'PipelineStepMixin' @@ -31,7 +34,7 @@ def test_meta(): def test_BaseMarker(): - """Test base class""" + """Test base class.""" base = BaseMarker(on=['bold', 'dwi'], name='mymarker') input = {'bold': {'path': 'test'}, 't2': {'path': 'test'}} base.validate_input(input) -- 2.52.0 From ebaac4c953bf7bef56b5a1f96168af84cf81839a Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 27 Jul 2022 12:35:32 +0200 Subject: [PATCH 154/287] fix: flake8 fix for markers.tests.test_parcel --- junifer/markers/tests/test_parcel.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/junifer/markers/tests/test_parcel.py b/junifer/markers/tests/test_parcel.py index 5098f9b01..a4e07a93f 100644 --- a/junifer/markers/tests/test_parcel.py +++ b/junifer/markers/tests/test_parcel.py @@ -1,3 +1,5 @@ +"""Provide test for parcel aggregation.""" + import numpy as np from numpy.testing import assert_array_equal, assert_array_almost_equal from scipy.stats import trim_mean @@ -11,7 +13,7 @@ from junifer.markers.parcel import ParcelAggregation def test_ParcelAggregation_3D(): - """Test ParcelAggregation object on 3D images""" + """Test ParcelAggregation object on 3D images.""" # Get the testing atlas (for nilearn) atlas = datasets.fetch_atlas_schaefer_2018(n_rois=100) @@ -123,7 +125,7 @@ def test_ParcelAggregation_3D(): def test_ParcelAggregation_4D(): - """Test ParcelAggregation object on 4D images""" + """Test ParcelAggregation object on 4D images.""" # Get the testing atlas (for nilearn) atlas = datasets.fetch_atlas_schaefer_2018( -- 2.52.0 From 25f78c90b8840add0caf82fdc81e6508abeeeb03 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 27 Jul 2022 12:35:45 +0200 Subject: [PATCH 155/287] fix: flake8 fix for markers.tests.test_collection --- junifer/markers/tests/test_collection.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/junifer/markers/tests/test_collection.py b/junifer/markers/tests/test_collection.py index bb0bdf154..f2b185b83 100644 --- a/junifer/markers/tests/test_collection.py +++ b/junifer/markers/tests/test_collection.py @@ -1,3 +1,5 @@ +"""Provide tests for marker collection.""" + # Authors: Federico Raimondo # License: AGPL @@ -12,7 +14,7 @@ from junifer.storage import SQLiteFeatureStorage def test_MarkerCollection(): - """Test MarkerCollection""" + """Test MarkerCollection.""" wrong_markers = [ ParcelAggregation( atlas='Schaefer100x7', method='mean', @@ -83,7 +85,7 @@ def test_MarkerCollection(): def test_MarkerCollection_storage(): - """Test marker collection with storage""" + """Test marker collection with storage.""" markers = [ ParcelAggregation( atlas='Schaefer100x7', method='mean', -- 2.52.0 From cc9a9f2b9a2897e61b14fbe442eef63bd4ca0fcb Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 27 Jul 2022 12:36:58 +0200 Subject: [PATCH 156/287] fix: remove unnecessary sources for flake8 check --- tox.ini | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tox.ini b/tox.ini index c5b8633d6..65094068c 100644 --- a/tox.ini +++ b/tox.ini @@ -25,7 +25,7 @@ deps = flake8-docstrings flake8-bugbear commands = - flake8 {toxinidir}/junifer {toxinidir}/examples {toxinidir}/scratch {toxinidir}/tools {toxinidir}/setup.py + flake8 {toxinidir}/junifer {toxinidir}/setup.py [testenv:test] skip_install = false -- 2.52.0 From c868f151e49cd66b2e12a001ed3f55d455e4206b Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 27 Jul 2022 14:13:23 +0200 Subject: [PATCH 157/287] fix: refactor get_elements docstring in datagrabber.base to build docs --- junifer/datagrabber/base.py | 15 ++++++++++----- 1 file changed, 10 insertions(+), 5 deletions(-) diff --git a/junifer/datagrabber/base.py b/junifer/datagrabber/base.py index d65ebc34e..b4731e902 100644 --- a/junifer/datagrabber/base.py +++ b/junifer/datagrabber/base.py @@ -61,10 +61,6 @@ class BaseDataGrabber(ABC): Methods ------- - get_elements : List - Returns a list of elements that can be grabbed. The elements can be - strings, tuples or any object that will be then used as a key to - index the datagrabber __getitem__(element) : dict[str -> Path] Returns a dictionary of paths for each type of data required for the specified element. Use the element as a key to index the datagrabber. @@ -142,7 +138,16 @@ class BaseDataGrabber(ABC): @abstractmethod def get_elements(self): - """Get elements.""" + """Get elements. + + Returns + ------- + list + List of elements that can be grabbed. The elements can be strings, + tuples or any object that will be then used as a key to index the + datagrabber. + + """ raise_error( 'get_elements not implemented', NotImplementedError) # type: ignore -- 2.52.0 From 9d6fb26b2c565b68458221318683f9a6bafdf557 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 10 Aug 2022 14:37:54 +0200 Subject: [PATCH 158/287] update: allow setuptools_scm to read config from pyproject.toml --- junifer/__init__.py | 4 ++-- pyproject.toml | 1 - setup.py | 22 +--------------------- 3 files changed, 3 insertions(+), 24 deletions(-) diff --git a/junifer/__init__.py b/junifer/__init__.py index ac53f2bcb..b6d36c58b 100644 --- a/junifer/__init__.py +++ b/junifer/__init__.py @@ -1,6 +1,6 @@ -from . _version import __version__ +from ._version import __version__ from . import api from . import utils from . import datagrabber from . import markers -from . import configs \ No newline at end of file +from . import configs diff --git a/pyproject.toml b/pyproject.toml index 4191774b5..49bd06177 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -79,4 +79,3 @@ packages = ["junifer"] version_scheme = "python-simplified-semver" local_scheme = "no-local-version" write_to = "junifer/_version.py" -write_to_template = "__version__ = '{version}'\n" \ No newline at end of file diff --git a/setup.py b/setup.py index 8c99b5acd..c6e4450ca 100644 --- a/setup.py +++ b/setup.py @@ -7,25 +7,5 @@ from setuptools import setup - -def _getversion(): - from setuptools_scm.version import ( - get_local_node_and_date, - simplified_semver_version, - ) - - def clean_scheme(version): - print(version) - return get_local_node_and_date(version) if version.dirty else "" - - return { - 'version_scheme': simplified_semver_version, - 'local_scheme': clean_scheme, - 'write_to': 'junifer/_version.py', - 'write_to_template': "__version__ = '{version}'\n"} - - if __name__ == "__main__": - setup( - use_scm_version=_getversion, - ) + setup() -- 2.52.0 From e765b09c1ef103a24fca09cf063aa1d67be97c93 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 10 Aug 2022 14:42:50 +0200 Subject: [PATCH 159/287] update: add support for isort in tox.ini --- tox.ini | 27 ++++++++++++++++++++++++++- 1 file changed, 26 insertions(+), 1 deletion(-) diff --git a/tox.ini b/tox.ini index 65094068c..4a7812cdd 100644 --- a/tox.ini +++ b/tox.ini @@ -1,5 +1,5 @@ [tox] -envlist = flake8, test, coverage, codespell, py3{8,9,10} +envlist = isort, flake8, test, coverage, codespell, py3{8,9,10} isolated_build = true [gh-actions] @@ -18,6 +18,13 @@ deps = commands = pytest +[testenv:isort] +skip_install = true +deps = + isort +commands = + isort --check-only --diff {toxinidir}/junifer {toxinidir}/setup.py + [testenv:flake8] skip_install = true deps = @@ -55,6 +62,24 @@ commands = # Tool configs # ################ +[isort] +skip = + __init__.py +profile = black +line_length = 79 +lines_after_imports = 2 +known_first_party = junifer +known_third_party = + click + numpy + datalad + pandas + nibabel + nilearn + sqlalchemy + yaml + pytest + [flake8] exclude = __init__.py -- 2.52.0 From 5d9c39d3dc8af8e8c2ef812e1aa51d73e9ca10e9 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 11 Aug 2022 15:26:26 +0200 Subject: [PATCH 160/287] refactor: add docstrings and type annotations to utils/fs.py --- junifer/utils/fs.py | 23 +++++++++++++++++------ 1 file changed, 17 insertions(+), 6 deletions(-) diff --git a/junifer/utils/fs.py b/junifer/utils/fs.py index 77e0347dc..b45e33345 100644 --- a/junifer/utils/fs.py +++ b/junifer/utils/fs.py @@ -1,10 +1,21 @@ -"""Provide functions for filesystem.""" +"""Provide functions for filesystem manipulation.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL -import os import stat +from pathlib import Path -def make_executable(path): - """Make executable.""" - st = os.stat(path) - os.chmod(path, st.st_mode | stat.S_IEXEC) +def make_executable(path: Path) -> None: + """Make `path` executable. + + Parameters + ---------- + path : pathlib.Path + The path to make executable. + + """ + st = path.stat() + path.chmod(mode=st.st_mode | stat.S_IEXEC) -- 2.52.0 From 0a619705d75b250348e8c0d5753b3cf873094603 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 11 Aug 2022 18:47:10 +0200 Subject: [PATCH 161/287] refactor: add docstrings and type annotations in utils/logging.py --- junifer/utils/logging.py | 364 ++++++++++++++++++++++++--------------- 1 file changed, 224 insertions(+), 140 deletions(-) diff --git a/junifer/utils/logging.py b/junifer/utils/logging.py index 375bac230..37056e3a8 100644 --- a/junifer/utils/logging.py +++ b/junifer/utils/logging.py @@ -1,164 +1,140 @@ """Provide class and functions for logging.""" # Authors: Federico Raimondo +# Synchon Mandal # License: AGPL import logging -import subprocess import sys from distutils.version import LooseVersion from pathlib import Path -import warnings +from subprocess import PIPE, Popen, TimeoutExpired +from typing import Dict, NoReturn, Optional, Union +from warnings import warn -# logging.basicConfig(stream=sys.stdout, level=logging.WARN) -logger = logging.getLogger('JUNIFER') +logger = logging.getLogger("JUNIFER") + +_logging_types = { + "DEBUG": logging.DEBUG, + "INFO": logging.INFO, + "WARNING": logging.WARNING, + "ERROR": logging.ERROR, +} -def _get_git_head(path): - """Aux function to read HEAD from git.""" - if not path.exists(): - raise ValueError('This path does not exist: {}'.format(path)) - command = ('cd {gitpath}; ' - 'git rev-parse --verify HEAD').format(gitpath=path) - process = subprocess.Popen(command, - stdout=subprocess.PIPE, - shell=True) - proc_stdout = process.communicate()[0].strip() - del process - return proc_stdout +class WrapStdOut: + """ + Dynamically wrap to sys.stdout. + + This makes packages that monkey-patch sys.stdout (e.g.doctest, + sphinx-gallery) work properly. + + """ + + def __getattr__(self, name: str) -> str: + """Implement attribute fetch.""" + # Even more ridiculous than this class, this must be sys.stdout (not + # just stdout) in order for this to work (tested on OSX and Linux) + if hasattr(sys.stdout, name): + return getattr(sys.stdout, name) + else: + raise AttributeError(f"'file' object has not attribute '{name}'") -def get_versions(sys): - """Import stuff and get versions if module. +def _get_git_head(path: Path) -> str: + """Aux function to read HEAD from git. Parameters ---------- - sys : module - The sys module object. + path : pathlib.Path + The path to read git HEAD from. + + Returns + ------- + str + Empty string if timeout expired for subprocess command execution else + git HEAD information. + + """ + if not path.exists(): + raise ValueError(f"This path does not exist: {path}") + command = f"cd {path}; git rev-parse --verify HEAD" + process = Popen( + args=command, + stdout=PIPE, + shell=True, + ) + try: + stdout, _ = process.communicate(timeout=10) + proc_stdout = stdout.strip() + except TimeoutExpired: + process.kill() + proc_stdout = "" + return proc_stdout + + +def get_versions() -> Dict: + """Import stuff and get versions if module. + Returns ------- module_versions : dict The module names and corresponding versions. + """ module_versions = {} for name, module in sys.modules.items(): - if '.' in name: + if "." in name: continue - if name in ['_curses']: + if name in ["_curses"]: continue - vstring = str(getattr(module, '__version__', None)) + vstring = str(getattr(module, "__version__", None)) module_version = LooseVersion(vstring) - module_version = getattr(module_version, 'vstring', None) + module_version = getattr(module_version, "vstring", None) if module_version is None: module_version = None - elif 'git' in module_version: + elif "git" in module_version: git_path = Path(module.__file__).resolve().parent head = _get_git_head(git_path) - module_version += '-HEAD:{}'.format(head) + module_version += f"-HEAD:{head}" module_versions[name] = module_version return module_versions -def get_ext_versions(tbox_path): - """Get versions of external tools used by JUNIFER.""" - versions = {} - # spm_path = tbox_path / 'spm12' - # if spm_path.exists(): - # head = _get_git_head(spm_path) - # module_version = 'SPM12-HEAD:{}'.format(head) - # versions['spm'] = module_version - return versions +# def get_ext_versions(tbox_path: Path) -> Dict: +# """Get versions of external tools used by junifer. + +# Parameters +# ---------- +# tbox_path : pathlib.Path +# The path to external toolboxes. + +# Returns +# ------- +# dict +# The dependency information. + +# """ +# versions = {} +# # spm_path = tbox_path / 'spm12' +# # if spm_path.exists(): +# # head = _get_git_head(spm_path) +# # module_version = 'SPM12-HEAD:{}'.format(head) +# # versions['spm'] = module_version +# return versions -def _safe_log(versions, name): - if name in versions: - logger.info(f'{name}: {versions[name]}') - - -def log_versions(tbox_path=None): - """Log versions of dependencies and junifer.""" - versions = get_versions(sys) - - logger.info('===== Lib Versions =====') - _safe_log(versions, 'numpy') - _safe_log(versions, 'scipy') - _safe_log(versions, 'pandas') - _safe_log(versions, 'nipype') - _safe_log(versions, 'nitime') - _safe_log(versions, 'nilearn') - _safe_log(versions, 'nibabel') - _safe_log(versions, 'junifer') - logger.info('========================') - - if tbox_path is not None: - # ext_versions = get_ext_versions(tbox_path) - # logger.info('spm: {}'.format(ext_versions['spm'])) - logger.info('========================') - - -_logging_types = dict(DEBUG=logging.DEBUG, INFO=logging.INFO, - WARNING=logging.WARNING, ERROR=logging.ERROR) - - -def configure_logging( - level='WARNING', fname=None, overwrite=None, output_format=None -): - """Configure the logging functionality. +def _close_handlers(logger: logging.Logger) -> None: + """Safely close relevant handlers for logger. Parameters ---------- - level : int or string - The level of the messages to print. If string, it will be interpreted - as elements of logging. - Options are: ['DEBUG', 'INFO', 'WARNING', 'ERROR']. Defaults to - 'WARNING'. - fname : str, Path or None - Filename of the log to print to. If None, stdout is used. - overwrite : bool | None - Overwrite the log file (if it exists). Otherwise, statements - will be appended to the log (default). None is the same as False, - but additionally raises a warning to notify the user that log - entries will be appended. - output_format : str - Format of the output messages. See the following for examples: + logger : logging.logger + The logger to close handlers for. - https://docs.python.org/dev/howto/logging.html - - e.g., "%(asctime)s - %(levelname)s - %(message)s". - - Defaults to "%(asctime)s - %(name)s - %(levelname)s - %(message)s" """ - _close_handlers(logger) - if output_format is None: - output_format = ('%(asctime)s [%(levelname)8s] %(message)s ' - '(%(filename)s:%(lineno)s)') - formatter = logging.Formatter(output_format) - - if fname is not None: - if not isinstance(fname, Path): - fname = Path(fname) - if fname.exists() and overwrite is None: - warnings.warn( - f'File ({fname.as_posix()}) exists. ' - 'Messages will be appended. Use overwrite=True to ' - 'overwrite or overwrite=False to avoid this message') - overwrite = False - mode = 'w' if overwrite else 'a' - lh = logging.FileHandler(fname, mode=mode) - else: - lh = logging.StreamHandler(WrapStdOut()) # type: ignore - - if isinstance(level, str): - level = _logging_types[level] - lh.setFormatter(formatter) - logger.setLevel(level) - logger.addHandler(lh) - log_versions() - - -def _close_handlers(logger): for handler in list(logger.handlers): if isinstance(handler, (logging.FileHandler, logging.StreamHandler)): if isinstance(handler, logging.FileHandler): @@ -166,37 +142,145 @@ def _close_handlers(logger): logger.removeHandler(handler) -def raise_error(msg, klass=ValueError): - """Raise error.""" +def _safe_log(versions: Dict, name: str) -> None: + """Log with safety. + + Parameters + ---------- + versions : dict + The dictionary with keys as dependency names and values as the + versions. + name : str + The dependency to look up in `versions`. + + """ + if name in versions: + logger.info(f"{name}: {versions[name]}") + + +def log_versions(tbox_path: Optional[Path] = None) -> None: + """Log versions of dependencies and junifer. + + If `tbox_path` is specified, can also log versions of external toolboxes. + + Parameters + ---------- + tbox_path : pathlib.Path, optional + The path to external toolboxes (default None). + + """ + # Get versions of all found packages + versions = get_versions() + + logger.info("===== Lib Versions =====") + _safe_log(versions, "numpy") + _safe_log(versions, "scipy") + _safe_log(versions, "pandas") + _safe_log(versions, "nipype") + _safe_log(versions, "nitime") + _safe_log(versions, "nilearn") + _safe_log(versions, "nibabel") + _safe_log(versions, "junifer") + logger.info("========================") + + if tbox_path is not None: + # ext_versions = get_ext_versions(tbox_path) + # logger.info('spm: {}'.format(ext_versions['spm'])) + # logger.info('========================') + pass + + +def configure_logging( + level: Union[int, str] = "WARNING", + fname: Optional[Union[str, Path]] = None, + overwrite: Optional[bool] = None, + output_format=None, +) -> None: + """Configure the logging functionality. + + Parameters + ---------- + level : int or {"DEBUG", "INFO", "WARNING", "ERROR"} + The level of the messages to print. If string, it will be interpreted + as elements of logging (default "WARNING"). + fname : str or pathlib.Path, optional + Filename of the log to print to. If None, stdout is used + (default None). + overwrite : bool, optional + Overwrite the log file (if it exists). Otherwise, statements + will be appended to the log (default). None is the same as False, + but additionally raises a warning to notify the user that log + entries will be appended (default None). + output_format : str, optional + Format of the output messages. See the following for examples: + https://docs.python.org/dev/howto/logging.html + e.g., "%(asctime)s - %(levelname)s - %(message)s". + If None, default string format is used + (default "%(asctime)s - %(name)s - %(levelname)s - %(message)s"). + + """ + _close_handlers(logger) # close relevant logger handlers + + # Set logging level + if isinstance(level, str): + level = _logging_types[level] + + # Set logging output handler + if fname is not None: + # Convert str to Path + if not isinstance(fname, Path): + fname = Path(fname) + if fname.exists() and overwrite is None: + warn( + f"File ({str(fname.absolute())}) exists. " + "Messages will be appended. Use overwrite=True to " + "overwrite or overwrite=False to avoid this message." + ) + overwrite = False + mode = "w" if overwrite else "a" + lh = logging.FileHandler(fname, mode=mode) + else: + lh = logging.StreamHandler(WrapStdOut()) + + # Set logging format + if output_format is None: + output_format = "%(asctime)s - %(name)s - %(levelname)s - %(message)s" + # ( + # "%(asctime)s [%(levelname)s] %(message)s (%(filename)s:%(lineno)s)" + # ) + formatter = logging.Formatter(fmt=output_format) + + lh.setFormatter(formatter) # set formatter + logger.setLevel(level) # set level + logger.addHandler(lh) # set handler + log_versions() # log versions of installed packages + + +def raise_error(msg: str, klass: Exception = ValueError) -> NoReturn: + """Raise error, but first log it. + + Parameters + ---------- + msg : str + The message for the exception. + klass : subclass of Exception, optional + The subclass of Exception to raise using (default ValueError). + + """ logger.error(msg) raise klass(msg) -def warn(msg, category=RuntimeWarning): +def warn_with_log(msg: str, category: Warning = RuntimeWarning) -> None: """Warn, but first log it. Parameters ---------- msg : str - Warning message - category : instance of Warning - The warning class. Defaults to ``RuntimeWarning``. + Warning message. + category : subclass of Warning, optional + The warning subclass (default RuntimeWarning). + """ logger.warning(msg) - warnings.warn(msg, category=category) - - -class WrapStdOut(object): - """Dynamically wrap to sys.stdout. - - This makes packages that monkey-patch sys.stdout (e.g.doctest, - sphinx-gallery) work properly. - """ - - def __getattr__(self, name): # noqa: D105 - # Even more ridiculous than this class, this must be sys.stdout (not - # just stdout) in order for this to work (tested on OSX and Linux) - if hasattr(sys.stdout, name): - return getattr(sys.stdout, name) - else: - raise AttributeError(f"'file' object has not attribute '{name}'") + warn(msg, category=category) -- 2.52.0 From d5d81aa3056ca6e94fc031cc21fb41d4b88f06f4 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 11 Aug 2022 18:48:08 +0200 Subject: [PATCH 162/287] refactor: prune import for utils sub-package --- junifer/utils/__init__.py | 10 ++++++++-- 1 file changed, 8 insertions(+), 2 deletions(-) diff --git a/junifer/utils/__init__.py b/junifer/utils/__init__.py index c8b961c0b..b7fde7ed9 100644 --- a/junifer/utils/__init__.py +++ b/junifer/utils/__init__.py @@ -1,2 +1,8 @@ -from . import logging -from .logging import configure_logging, logger, raise_error, warn \ No newline at end of file +"""Provide imports for utils sub-package.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from .fs import make_executable +from .logging import configure_logging, logger, raise_error, warn_with_log -- 2.52.0 From a65b3ef92b5ac436fad8194f82edf5b35f30577b Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 11 Aug 2022 18:49:08 +0200 Subject: [PATCH 163/287] refactor: prune test_logging.py and add docstrings and type annotations --- junifer/utils/tests/test_logging.py | 317 ++++++++++++++++++---------- 1 file changed, 201 insertions(+), 116 deletions(-) diff --git a/junifer/utils/tests/test_logging.py b/junifer/utils/tests/test_logging.py index c5a9a41a7..5673ccebf 100644 --- a/junifer/utils/tests/test_logging.py +++ b/junifer/utils/tests/test_logging.py @@ -2,135 +2,220 @@ # Authors: Federico Raimondo # Sami Hamdan +# Synchon Mandal # License: AGPL -from junifer.utils import logger, configure_logging, raise_error, warn -from junifer.utils.logging import _close_handlers -import pytest -import tempfile +import logging from pathlib import Path +import pytest -def test_log_file(): - """Test logging to a file.""" - with tempfile.TemporaryDirectory() as tmp: - tmpdir = Path(tmp) - configure_logging(fname=tmpdir / 'test1.log') - logger.debug('Debug message') - logger.info('Info message') - logger.warning('Warn message') - logger.error('Error message') - _close_handlers(logger) - with open(tmpdir / 'test1.log') as f: +from junifer.utils.logging import ( + logger, + _close_handlers, + configure_logging, + get_versions, + log_versions, + raise_error, + warn_with_log, +) + + +def test_get_versions() -> None: + """Test version info fetch for modules.""" + module_versions = get_versions() + assert "junifer" in module_versions.keys() + + +def test_log_versions(caplog: pytest.LogCaptureFixture) -> None: + """Test logging of dependency and junifer versions. + + Parameters + ---------- + caplog : pytest.LogCaptureFixture + A pytest fixture to capture logging. + + """ + # Set log capturing at INFO + caplog.set_level(logging.INFO) + # Log versions + log_versions() + # Check logging levels + for record in caplog.records: + assert record.levelname not in ("DEBUG", "WARNING", "ERROR") + assert record.levelname == "INFO" + # Check logging message + assert "junifer" in caplog.text + + +def test_log_file(tmp_path: Path) -> None: + """Test logging to a file. + + tmp_path : Path + The path to the test directory. + + """ + configure_logging(fname=tmp_path / "test1.log") + logger.debug("Debug message") + logger.info("Info message") + logger.warning("Warn message") + logger.error("Error message") + _close_handlers(logger) + with open(tmp_path / "test1.log") as f: + lines = f.readlines() + assert not any("Debug message" in line for line in lines) + assert not any("Info message" in line for line in lines) + assert any("Warn message" in line for line in lines) + assert any("Error message" in line for line in lines) + + configure_logging(fname=tmp_path / "test2.log", level="INFO") + logger.debug("Debug message") + logger.info("Info message") + logger.warning("Warn message") + logger.error("Error message") + _close_handlers(logger) + with open(tmp_path / "test2.log") as f: + lines = f.readlines() + assert not any("Debug message" in line for line in lines) + assert any("Info message" in line for line in lines) + assert any("Warn message" in line for line in lines) + assert any("Error message" in line for line in lines) + + configure_logging(fname=tmp_path / "test3.log", level="WARNING") + logger.debug("Debug message") + logger.info("Info message") + logger.warning("Warn message") + logger.error("Error message") + _close_handlers(logger) + with open(tmp_path / "test3.log") as f: + lines = f.readlines() + assert not any("Debug message" in line for line in lines) + assert not any("Info message" in line for line in lines) + assert any("Warn message" in line for line in lines) + assert any("Error message" in line for line in lines) + + configure_logging(fname=tmp_path / "test4.log", level="ERROR") + logger.debug("Debug message") + logger.info("Info message") + logger.warning("Warn message") + logger.error("Error message") + with open(tmp_path / "test4.log") as f: + lines = f.readlines() + assert not any("Debug message" in line for line in lines) + assert not any("Info message" in line for line in lines) + assert not any("Warn message" in line for line in lines) + assert any("Error message" in line for line in lines) + + with pytest.warns(UserWarning, match="to avoid this message"): + configure_logging(fname=tmp_path / "test4.log", level="WARNING") + logger.debug("Debug2 message") + logger.info("Info2 message") + logger.warning("Warn2 message") + logger.error("Error2 message") + with open(tmp_path / "test4.log") as f: lines = f.readlines() - assert not any('Debug message' in line for line in lines) - assert not any('Info message' in line for line in lines) - assert any('Warn message' in line for line in lines) - assert any('Error message' in line for line in lines) + assert not any("Debug message" in line for line in lines) + assert not any("Info message" in line for line in lines) + assert not any("Warn message" in line for line in lines) + assert any("Error message" in line for line in lines) + assert not any("Debug2 message" in line for line in lines) + assert not any("Info2 message" in line for line in lines) + assert any("Warn2 message" in line for line in lines) + assert any("Error2 message" in line for line in lines) - configure_logging(fname=tmpdir / 'test2.log', level='INFO') - logger.debug('Debug message') - logger.info('Info message') - logger.warning('Warn message') - logger.error('Error message') - _close_handlers(logger) - with open(tmpdir / 'test2.log') as f: - lines = f.readlines() - assert not any('Debug message' in line for line in lines) - assert any('Info message' in line for line in lines) - assert any('Warn message' in line for line in lines) - assert any('Error message' in line for line in lines) + configure_logging( + fname=tmp_path / "test4.log", level="WARNING", overwrite=True + ) + logger.debug("Debug3 message") + logger.info("Info3 message") + logger.warning("Warn3 message") + logger.error("Error3 message") + with open(tmp_path / "test4.log") as f: + lines = f.readlines() + assert not any("Debug message" in line for line in lines) + assert not any("Info message" in line for line in lines) + assert not any("Warn message" in line for line in lines) + assert not any("Error message" in line for line in lines) + assert not any("Debug2 message" in line for line in lines) + assert not any("Info2 message" in line for line in lines) + assert not any("Warn2 message" in line for line in lines) + assert not any("Error2 message" in line for line in lines) + assert not any("Debug3 message" in line for line in lines) + assert not any("Info3 message" in line for line in lines) + assert any("Warn3 message" in line for line in lines) + assert any("Error3 message" in line for line in lines) - configure_logging(fname=tmpdir / 'test3.log', level='WARNING') - logger.debug('Debug message') - logger.info('Info message') - logger.warning('Warn message') - logger.error('Error message') - _close_handlers(logger) - with open(tmpdir / 'test3.log') as f: - lines = f.readlines() - assert not any('Debug message' in line for line in lines) - assert not any('Info message' in line for line in lines) - assert any('Warn message' in line for line in lines) - assert any('Error message' in line for line in lines) - - configure_logging(fname=tmpdir / 'test4.log', level='ERROR') - logger.debug('Debug message') - logger.info('Info message') - logger.warning('Warn message') - logger.error('Error message') - with open(tmpdir / 'test4.log') as f: - lines = f.readlines() - assert not any('Debug message' in line for line in lines) - assert not any('Info message' in line for line in lines) - assert not any('Warn message' in line for line in lines) - assert any('Error message' in line for line in lines) - - with pytest.warns(UserWarning, match='to avoid this message'): - configure_logging(fname=tmpdir / 'test4.log', level='WARNING') - logger.debug('Debug2 message') - logger.info('Info2 message') - logger.warning('Warn2 message') - logger.error('Error2 message') - with open(tmpdir / 'test4.log') as f: - lines = f.readlines() - assert not any('Debug message' in line for line in lines) - assert not any('Info message' in line for line in lines) - assert not any('Warn message' in line for line in lines) - assert any('Error message' in line for line in lines) - assert not any('Debug2 message' in line for line in lines) - assert not any('Info2 message' in line for line in lines) - assert any('Warn2 message' in line for line in lines) - assert any('Error2 message' in line for line in lines) - - configure_logging(fname=tmpdir / 'test4.log', level='WARNING', - overwrite=True) - logger.debug('Debug3 message') - logger.info('Info3 message') - logger.warning('Warn3 message') - logger.error('Error3 message') - with open(tmpdir / 'test4.log') as f: - lines = f.readlines() - assert not any('Debug message' in line for line in lines) - assert not any('Info message' in line for line in lines) - assert not any('Warn message' in line for line in lines) - assert not any('Error message' in line for line in lines) - assert not any('Debug2 message' in line for line in lines) - assert not any('Info2 message' in line for line in lines) - assert not any('Warn2 message' in line for line in lines) - assert not any('Error2 message' in line for line in lines) - assert not any('Debug3 message' in line for line in lines) - assert not any('Info3 message' in line for line in lines) - assert any('Warn3 message' in line for line in lines) - assert any('Error3 message' in line for line in lines) - - with pytest.warns(RuntimeWarning, match=r"Warn raised"): - warn('Warn raised') - with pytest.raises(ValueError, match=r"Error raised"): - raise_error('Error raised') - with open(tmpdir / 'test4.log') as f: - lines = f.readlines() - assert any('Warn raised' in line for line in lines) - assert any('Error raised' in line for line in lines) + with pytest.warns(RuntimeWarning, match=r"Warn raised"): + warn_with_log("Warn raised") + with pytest.raises(ValueError, match=r"Error raised"): + raise_error("Error raised") + with open(tmp_path / "test4.log") as f: + lines = f.readlines() + assert any("Warn raised" in line for line in lines) + assert any("Error raised" in line for line in lines) -def test_log(): - """Simple log test.""" +def test_log_stdout(caplog: pytest.LogCaptureFixture) -> None: + """Test logging to stdout. + + Parameters + ---------- + caplog : pytest.LogCaptureFixture + A pytest fixture to capture logging. + + """ configure_logging() - logger.info('Testing') + logger.info("Testing") + for record in caplog.records: + assert record.levelname == "INFO" -def test_lib_logging(): - """Test logging versions.""" +def test_lib_logging(tmp_path: Path) -> None: + """Test logging versions. + tmp_path : Path + The path to the test directory. + + """ + # Import third-party packages import numpy as np # noqa import pandas # noqa - with tempfile.TemporaryDirectory() as tmp: - tmpdir = Path(tmp) - configure_logging(fname=tmpdir / 'test1.log', level='INFO') - logger.info('first message') - with open(tmpdir / 'test1.log') as f: - lines = f.readlines() - assert any('numpy' in line for line in lines) - assert any('pandas' in line for line in lines) - assert any('junifer' in line for line in lines) + + log_file_path = tmp_path / "test_lib_logging.log" + configure_logging(fname=log_file_path, level="INFO") + logger.info("first message") + with open(log_file_path) as f: + lines = f.readlines() + assert any("numpy" in line for line in lines) + assert any("pandas" in line for line in lines) + assert any("junifer" in line for line in lines) + + +def test_raise_error(caplog: pytest.LogCaptureFixture) -> None: + """Test logging and raising error. + + Parameters + ---------- + caplog : pytest.LogCaptureFixture + A pytest fixture to capture logging. + + """ + with pytest.raises(ValueError, match="test error"): + raise_error(msg="test error") + for record in caplog.records: + assert record.levelname == "ERROR" + + +def test_warn_with_log(caplog: pytest.LogCaptureFixture) -> None: + """Test logging and warning. + + Parameters + ---------- + caplog : pytest.LogCaptureFixture + A pytest fixture to capture logging. + + """ + with pytest.warns(RuntimeWarning, match="test warning"): + warn_with_log("test warning") + for record in caplog.records: + assert record.levelname == "WARNING" -- 2.52.0 From 165b8a19a7e81410fa43d7dcce21f0573b757192 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 11 Aug 2022 18:49:46 +0200 Subject: [PATCH 164/287] refactor: add separate module for unit tests of utils/fs.py --- junifer/utils/tests/test_fs.py | 30 ++++++++++++++++++++++++++++++ 1 file changed, 30 insertions(+) create mode 100644 junifer/utils/tests/test_fs.py diff --git a/junifer/utils/tests/test_fs.py b/junifer/utils/tests/test_fs.py new file mode 100644 index 000000000..17912a6a8 --- /dev/null +++ b/junifer/utils/tests/test_fs.py @@ -0,0 +1,30 @@ +"""Provide tests for filesystem manipulation.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from pathlib import Path + +from junifer.utils.fs import make_executable + + +def test_make_executable(tmp_path: Path) -> None: + """Test making path executable. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + test_file_path = tmp_path / "make_me_executable.txt" + test_file_path.write_bytes(b"umm") + test_file_stat_initial = test_file_path.stat() + # Check initial file mode + assert test_file_stat_initial.st_mode == 33188 + # Make the path executable + make_executable(test_file_path) + test_file_stat_final = test_file_path.stat() + # Check final file mode + assert test_file_stat_final.st_mode == 33252 -- 2.52.0 From 34113d2ef27ce6364b620a6f2a5ab5a476f7c407 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 15 Aug 2022 10:49:46 +0200 Subject: [PATCH 165/287] update: use context manager for pytest caplog in test_logging.py --- junifer/utils/tests/test_logging.py | 18 +++++++++--------- 1 file changed, 9 insertions(+), 9 deletions(-) diff --git a/junifer/utils/tests/test_logging.py b/junifer/utils/tests/test_logging.py index 5673ccebf..c4ab534bc 100644 --- a/junifer/utils/tests/test_logging.py +++ b/junifer/utils/tests/test_logging.py @@ -37,15 +37,15 @@ def test_log_versions(caplog: pytest.LogCaptureFixture) -> None: """ # Set log capturing at INFO - caplog.set_level(logging.INFO) - # Log versions - log_versions() - # Check logging levels - for record in caplog.records: - assert record.levelname not in ("DEBUG", "WARNING", "ERROR") - assert record.levelname == "INFO" - # Check logging message - assert "junifer" in caplog.text + with caplog.at_level(logging.INFO): + # Log versions + log_versions() + # Check logging levels + for record in caplog.records: + assert record.levelname not in ("DEBUG", "WARNING", "ERROR") + assert record.levelname == "INFO" + # Check logging message + assert "junifer" in caplog.text def test_log_file(tmp_path: Path) -> None: -- 2.52.0 From a4a1e66a09205a1c83d844e16e9a4e43d6125177 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 15 Aug 2022 15:29:04 +0200 Subject: [PATCH 166/287] chore: flake8 fix for logging.py --- junifer/utils/logging.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/junifer/utils/logging.py b/junifer/utils/logging.py index 37056e3a8..53431187b 100644 --- a/junifer/utils/logging.py +++ b/junifer/utils/logging.py @@ -246,7 +246,8 @@ def configure_logging( if output_format is None: output_format = "%(asctime)s - %(name)s - %(levelname)s - %(message)s" # ( - # "%(asctime)s [%(levelname)s] %(message)s (%(filename)s:%(lineno)s)" + # "%(asctime)s [%(levelname)s] %(message)s " + # "(%(filename)s:%(lineno)s)" # ) formatter = logging.Formatter(fmt=output_format) -- 2.52.0 From 0681dc7476aae7117388a8f1b0c5437d5e943502 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 10 Aug 2022 14:54:16 +0200 Subject: [PATCH 167/287] refactor: prune base.py and add docstrings and type annotations --- junifer/storage/base.py | 407 +++++++++++++++++++++------------------- 1 file changed, 212 insertions(+), 195 deletions(-) diff --git a/junifer/storage/base.py b/junifer/storage/base.py index 53d98703e..4b2bc94c7 100644 --- a/junifer/storage/base.py +++ b/junifer/storage/base.py @@ -1,246 +1,263 @@ -"""Provide class and functions for storage.""" +"""Provide abstract base class for feature storage.""" # Authors: Federico Raimondo +# Synchon Mandal # License: AGPL -import numpy as np -import pandas as pd -import json -import hashlib + from abc import ABC, abstractmethod +from pathlib import Path +from typing import Dict, Iterable, List, Optional, Union + +import pandas as pd -from ..utils import logger from .. import __version__ - - -def process_meta(meta): - """Process the metadata for storage. - - It removes the "element" key and adds the "_element_keys" with the keys - used to index the element. - - Parameters - ---------- - meta: dict - The metadata. Must contain the key 'element' - return_idx: bool - If true, return the pandas index to be stored. Defaults to false - n_rows: int - Number of rows to create (if return_idx is true) - rows_col_name: str - The column name to use in case n_rows > 1. If None (default) and - n_rows > 1, the name will be 'index'. - - Returns - ------- - md5_hash: str - The md5 hash of the meta - meta : dict - The metadata processed for storage - idx : pd.MultiIndex - The pandas index (if return_idx is True) - """ - if meta is None: - raise ValueError('Meta must be a dict (currently is None)') - t_meta = meta.copy() - element = t_meta.pop('element', None) - if element is None: - if '_element_keys' not in t_meta: - raise ValueError( - 'Meta must contain the key "element" or "_element_keys"') - else: - if isinstance(element, dict): - t_meta['_element_keys'] = list(element.keys()) - else: - t_meta['_element_keys'] = ['element'] - md5_hash = _meta_hash(t_meta) - return md5_hash, t_meta - - -def _meta_hash(meta): - """Compute the md5 hash of the meta. - - Parameters - ---------- - meta: dict - The metadata. Must contain the key 'element' - - Returns - ------- - md5: str - The md5 hash of the meta - """ - logger.debug(f'Hashing meta {meta}') - meta_md5 = hashlib.md5( - json.dumps(meta, sort_keys=True).encode('utf-8')).hexdigest() - logger.debug(f'Hash computed: {meta_md5}') - return meta_md5 - - -def element_to_index(meta, n_rows=1, rows_col_name=None): - """Convert the element meta to index. - - Parameters - ---------- - meta: dict - The metadata. Must contain the key 'element' - n_rows: int - Number of rows to create. Defaults to 1. - rows_col_name: str - The column name to use in case n_rows > 1. If None (default) and - n_rows > 1, the name will be 'index'. - - Returns - ------- - index: pd.MultiIndex - The index of the dataframe to store - - Raises - ------ - ValueError - If the meta does not contain the key 'element' - """ - if 'element' not in meta: - raise ValueError( - 'To create and index, meta must contain the key "element"') - element = meta['element'] - if not isinstance(element, dict): - element = dict(element=element) - if rows_col_name is None: - rows_col_name = 'idx' - elem_idx = { - k: [v] * n_rows for k, v in element.items() - } - elem_idx[rows_col_name] = np.arange(n_rows) # type: ignore - index = pd.MultiIndex.from_frame( - pd.DataFrame(elem_idx, index=range(n_rows))) - return index - - -def element_to_prefix(element): - """Convert the element meta to prefix.""" - logger.debug(f'Converting element {element} to prefix') - prefix = 'element' - if isinstance(element, tuple): - prefix = f"{prefix}_{'_'.join([f'{x}' for x in element])}" - elif isinstance(element, dict): - prefix = f"{prefix}_{'_'.join([f'{x}' for x in element.values()])}" - elif isinstance(element, (str, int)): - prefix = f'{prefix}_{element}' - else: - raise ValueError(f'Cannot convert element {element} to prefix. ' - 'Must be a str, tuple or dict') - logger.debug(f'Converted prefix {prefix}') - return f'{prefix}_' +from ..utils import raise_error class BaseFeatureStorage(ABC): - """Base class for feature storage.""" + """Abstract base class for feature storage. - def __init__(self, uri, single_output=False): + For every interface that is required, one needs to provide a concrete + implementation of this abstract class. + + Parameters + ---------- + uri : str or pathlib.Path + The path to the storage. + single_output : bool, optional + Whether to have single output (default False). + + """ + + def __init__( + self, uri: Union[str, Path], single_output: bool = False + ) -> None: """Initialize the class.""" self.uri = uri self.single_output = single_output - def get_meta(self): - """Get metadata.""" + def get_meta(self) -> Dict: + """Get metadata. + + Returns + ------- + meta : dict + The metadata as a dictionary. + + """ meta = {} - meta['versions'] = { - 'junifer': __version__, + meta["versions"] = { + "junifer": __version__, } return meta + # TODO: is raising ValueError required? @abstractmethod - def validate(self, input): + def validate(self, input_: List[str]) -> bool: """Validate the input to the pipeline step. Parameters ---------- - input : Junifer Data dictionary + input_ : list The input to the pipeline step. + Returns + ------- + bool + Whether the `input` is valid or not. + Raises ------ - ValueError: + ValueError If the input does not have the required data. + """ - raise NotImplementedError('validate_input not implemented') + raise_error( + msg="Concrete classes need to implement validate_input().", + klass=NotImplementedError, + ) @abstractmethod - def list_features(self, return_df=False): + def list_features( + self, return_df: bool = False + ) -> Union[Dict[str, Dict], pd.DataFrame]: """List the features in the storage. Parameters ---------- - return_df : bool - If True, return a dataframe. If False, (default) return a - dictionary - Returns - ------- - features: dict(str, dict) | pd.DataFrame - List of features in the storage. If dictionarly, the keys are the - feature names to be used in read_features. The values are the - metadata of each feature. - """ - raise NotImplementedError('list_features not implemented') - - @abstractmethod - def read_df(self, feature_name=None, feature_md5=None): - """Read the features from the storage. + return_df : bool, optional + If True, returns a pandas DataFrame. If False, returns a + dictionary (default False). Returns ------- - out: pd.DataFrame - The features as a dataframe + dict or pandas.DataFrame + List of features in the storage. If dictionary is returned, the + keys are the feature names to be used in read_features() and the + values are the metadata of each feature. + """ - raise NotImplementedError('read_df not implemented') + raise_error( + msg="Concrete classes need to implement list_features().", + klass=NotImplementedError, + ) @abstractmethod - def store_metadata(self, meta): - """Store metadata.""" - raise NotImplementedError('store_metadata not implemented') + def read_df( + self, + feature_name: Optional[str] = None, + feature_md5: Optional[bool] = None, + ) -> pd.DataFrame: + """Read feature from the storage. + + Parameters + ---------- + feature_name : str, optional + Name of the feature to read (default None). + feature_md5 : str, optional + MD5 hash of the feature to read (default None). + + Returns + ------- + pandas.DataFrame + The features as a dataframe. + + """ + raise_error( + msg="Concrete classes need to implement read_df().", + klass=NotImplementedError, + ) @abstractmethod - def store_matrix2d(self, data, meta, col_names=None, row_names=None): - """Store 2D matrix.""" - raise NotImplementedError('store_matrix2d not implemented') + def store_metadata(self, meta: Dict) -> str: + """Store metadata. + + Parameters + ---------- + meta : dict + The metadata as a dictionary. + + Returns + ------- + str + The metadata column + + """ + raise_error( + msg="Concrete classes need to implement store_metadata().", + klass=NotImplementedError, + ) + + # TODO: complete type annotations + @abstractmethod + def store_matrix2d( + self, + data, + meta: Dict, + col_names: Optional[Iterable[str]] = None, + row_names: Optional[Iterable[str]] = None, + ) -> None: + """Store 2D matrix. + + Parameters + ---------- + data + meta : dict + The metadata as a dictionary. + col_names : list or tuple of str, optional + The column names (default None). + row_names : list of tuple of str, optional + The row names (default None). + + """ + raise_error( + msg="Concrete classes need to implement store_matrix2d().", + klass=NotImplementedError, + ) + + # TODO: complete type annotations + @abstractmethod + def store_table( + self, + data, + meta: Dict, + columns: Optional[Iterable[str]] = None, + rows_col_name: str = None, + ) -> None: + """Store table. + + Parameters + ---------- + data + meta : dict + The metadata as a dictionary. + columns : list or tuple of str, optional + The columns (default None). + rows_col_name : str, optional + The column name ot use in case number of rows greater than 1. + If None and number of rows greater than 1, then the name will be + "index" (default None). + + """ + raise_error( + msg="Concrete classes need to implement store_table().", + klass=NotImplementedError, + ) @abstractmethod - def store_table(self, data, meta, columns=None, rows_col_name=None): - """Store table.""" - raise NotImplementedError('store_table not implemented') + def store_df(self, df: pd.DataFrame, meta: Dict) -> None: + """Store pandas DataFerame. + + Parameters + ---------- + df : pandas.DataFrame + The DataFrame to store. + meta : dict + The metadata as a dictionary. + + """ + raise_error( + msg="Concrete classes need to implement store_df().", + klass=NotImplementedError, + ) + + # TODO: complete type annotations + @abstractmethod + def store_timeseries(self, data, meta: Dict) -> None: + """Store timeseries. + + Parameters + ---------- + data + meta : dict + The metadata as a dictionary. + + """ + raise_error( + msg="Concrete classes need to implement store_timeseries().", + klass=NotImplementedError, + ) @abstractmethod - def store_df(self, df, meta): - """Store dataframe.""" - raise NotImplementedError('store_df not implemented') - - @abstractmethod - def store_timeseries(self, data, meta): - """Store timeseries.""" - raise NotImplementedError('store_timeseries not implemented') - - @abstractmethod - def collect(self): + def collect(self) -> None: """Collect data.""" - raise NotImplementedError('collect not implemented') + raise_error( + msg="Concrete classes need to implement collect().", + klass=NotImplementedError, + ) - def __str__(self): - """Represent object as string.""" - single = '(single output)' \ - if self.single_output is True else '(multiple output)' - return f'<{self.__class__.__name__} @ {self.uri} {single}>' + def __str__(self) -> str: + """Represent object as string. + Returns + ------- + str + The string representation. -class PandasFeatureStoreage(BaseFeatureStorage): - """Store features via pandas.""" - - def _meta_row(self, meta, meta_md5): - """Convert the meta to a dataframe row.""" - data_df = {} - for k, v in meta.items(): - data_df[k] = json.dumps(v, sort_keys=True) - if 'marker' in meta: - data_df['name'] = meta['marker']['name'] - df = pd.DataFrame(data_df, index=[meta_md5]) - df.index.name = 'meta_md5' - return df + """ + single = ( + "(single output)" + if self.single_output is True + else "(multiple output)" + ) + return f"<{self.__class__.__name__} @ {self.uri} {single}>" -- 2.52.0 From 338077e72f529d4d9b6d3c72bf04e1f4cbfa518d Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 10 Aug 2022 14:57:56 +0200 Subject: [PATCH 168/287] refactor: move PandasBaseFeatureStorage to separate module and add docstrings and type annotations --- junifer/storage/pandas_base.py | 53 ++++++++++++++++++++++++++++++++++ 1 file changed, 53 insertions(+) create mode 100644 junifer/storage/pandas_base.py diff --git a/junifer/storage/pandas_base.py b/junifer/storage/pandas_base.py new file mode 100644 index 000000000..a6a1b430a --- /dev/null +++ b/junifer/storage/pandas_base.py @@ -0,0 +1,53 @@ +"""Provide abstract base class for feature storage via pandas.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +import json +from typing import Dict + +import pandas as pd + +from .base import BaseFeatureStorage + + +class PandasBaseFeatureStorage(BaseFeatureStorage): + """Abstract base class for feature storage via pandas. + + For every interface that is required, one needs to provide a concrete + implementation of this abstract class. + + See Also + -------- + BaseFeatureStorage + + """ + + def __init__(**kwargs) -> None: + """Initialize the class.""" + super().__init__(**kwargs) + + def _meta_row(self, meta: Dict, meta_md5: str) -> pd.DataFrame: + """Convert the metadata to a pandas DataFrame. + + Parameters + ---------- + meta : dict + The metadata as a dictionary. + meta_md5 : str + The MD5 hash of the metadata. + + Returns + ------- + pandas.DataFrame + + """ + data_df = {} + for k, v in meta.items(): + data_df[k] = json.dumps(v, sort_keys=True) + if "marker" in meta: + data_df["name"] = meta["marker"]["name"] + df = pd.DataFrame(data_df, index=[meta_md5]) + df.index.name = "meta_md5" + return df -- 2.52.0 From b32d02fd0f8727109495d4d863220a6ecc6dd101 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 10 Aug 2022 15:02:21 +0200 Subject: [PATCH 169/287] refactor: move storage utility functions to separate module and add docstrings and type annotations --- junifer/storage/utils.py | 162 +++++++++++++++++++++++++++++++++++++++ 1 file changed, 162 insertions(+) create mode 100644 junifer/storage/utils.py diff --git a/junifer/storage/utils.py b/junifer/storage/utils.py new file mode 100644 index 000000000..d24b36c44 --- /dev/null +++ b/junifer/storage/utils.py @@ -0,0 +1,162 @@ +"""Provide utility functions for the storage sub-package.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +import hashlib +import json +from typing import Dict, Optional, Tuple, Union + +import numpy as np +import pandas as pd + +from ..utils.logging import logger, raise_error + + +def _meta_hash(meta: Dict) -> str: + """Compute the MD5 hash of the metadata. + + Parameters + ---------- + meta : dict + The metadata as a dictionary. Must contain the key "element". + + Returns + ------- + str + The MD5 hash of the metadata. + + """ + logger.debug(f"Hashing metadata: {meta}") + meta_md5 = hashlib.md5( + json.dumps(meta, sort_keys=True).encode("utf-8") + ).hexdigest() + logger.debug(f"Hash computed: {meta_md5}") + return meta_md5 + + +def process_meta(meta: Dict) -> Tuple[str, Dict]: + """Process the metadata for storage. + + It removes the key "element" and adds the "_element_keys" with the keys + used to index the element. + + Parameters + ---------- + meta : dict + The metadata as a dictionary. Must contain the key "element". + + Returns + ------- + str + The MD5 hash of the metadata. + dict + The processed metadata for storage. + + Raises + ------ + ValueError + If `meta` is None or if it does not contain the key "element" or + "_element_keys". + + """ + if meta is None: + raise_error(msg="`meta` must be a dict (currently is None)") + # Remove key "element" + element = meta.pop("element", None) + if element is None: + if "_element_keys" not in meta: + raise_error( + msg="`meta` must contain the key 'element' or '_element_keys'" + ) + else: + if isinstance(element, dict): + meta["_element_keys"] = list(element.keys()) + else: + meta["_element_keyes"] = ["element"] + # MD5 hash of the metadata + md5_hash = _meta_hash(meta) + return md5_hash, meta + + +def element_to_prefix(element: Union[Tuple, Dict, str, int]) -> str: + """Convert the element metadata to prefix. + + Parameters + ---------- + element : tuple, dict, str or int + The element to convert to prefix. + + Returns + ------- + str + The element converted to prefix. + + Raises + ------ + ValueError + If invalid type is passed for `element`. + + """ + logger.debug(f"Converting element {element} to prefix.") + prefix = "element" + if isinstance(element, tuple): + prefix = f"{prefix}_{'_'.join([f'{x}' for x in element])}" + elif isinstance(element, dict): + prefix = f"{prefix}_{'_'.join([f'{x}' for x in element.values()])}" + elif isinstance(element, (str, int)): + prefix = f"{prefix}_{element}" + else: + raise_error( + f"Cannot convert element of type {type(element)} to prefix. " + "Must be a str, int, tuple or dict." + ) + logger.debug(f"Converted prefix: {prefix}") + return f"{prefix}_" + + +def element_to_index( + meta: Dict, n_rows: int = 1, rows_col_name: Optional[str] = None +) -> pd.MultiIndex: + """Convert the element metadata to index. + + Parameters + ---------- + meta : dict + The metadata as a dictionary. Must contain the key "element"." + n_rows : int, optional + Number of rows to create (default 1). + rows_col_name: str, optional + The column name to use in case `n_rows` > 1. If None and + n_rows > 1, the name will be "index" (default None). + + Returns + ------- + pandas.MultiIndex + The index of the dataframe to store. + + Raises + ------ + ValueError + If `meta` does not contain the key "element". + + """ + if "element" not in meta: + raise_error( + msg="To create and index, metadata must contain the key 'element'." + ) + # Get element + element = meta["element"] + if not isinstance(element, dict): + element = {"element": element} + # Check rows_col_name + if rows_col_name is None: + rows_col_name = "index" + elem_idx = {k: [v] * n_rows for k, v in element.items()} + elem_idx[rows_col_name] = np.arange(n_rows) + # Create index + index = pd.MultiIndex.from_frame( + pd.DataFrame(elem_idx, index=range(n_rows)) + ) + return index -- 2.52.0 From 6addbf232d9a744d9535acd7f461a574e41ed49f Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 10 Aug 2022 15:12:53 +0200 Subject: [PATCH 170/287] refactor: prune sqlite.py and add docstring and type annotations --- junifer/storage/sqlite.py | 688 ++++++++++++++++++++++++++------------ 1 file changed, 479 insertions(+), 209 deletions(-) diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py index f6f9c6d8e..20326ecf2 100644 --- a/junifer/storage/sqlite.py +++ b/junifer/storage/sqlite.py @@ -1,169 +1,476 @@ -"""Provide class and functions for sqlite.""" +"""Provide concrete implementation for feature storage via SQLite.""" # Authors: Federico Raimondo # Synchon Mandal # License: AGPL from pathlib import Path +from typing import TYPE_CHECKING, Dict, Iterable, List, Optional, Union + import pandas as pd from pandas.core.base import NoNewAttributesMixin from pandas.io.sql import pandasSQL_builder from sqlalchemy import create_engine, inspect -import tqdm +from tqdm import tqdm from ..api.decorators import register_storage -from .base import (PandasFeatureStoreage, process_meta, element_to_prefix, - element_to_index) -from ..utils.logging import warn, logger +from ..utils import logger, raise_error, warn_with_log +from .pandas_base import PandasBaseFeatureStorage +from .utils import element_to_index, element_to_prefix, process_meta + + +if TYPE_CHECKING: + from sqlalchemy import Engine @register_storage -class SQLiteFeatureStorage(PandasFeatureStoreage): - """SQLite feature storage.""" +class SQLiteFeatureStorage(PandasBaseFeatureStorage): + """Concrete implementation for feature storage via SQLite. - def __init__(self, uri, single_output=False, upsert='update'): - """Initialise an SQLite feature storage. + Parameters + ---------- + uri : str or pathlib.Path + The path to the file to be used. + single_output : bool, optional + If False, will create one file per element. The name + of the file will be prefixed with the respective element. + If True, will create only one file as specified in the `uri` and + store all the elements in the same file. This behaviour is only + suitable for non-parallel executions. SQLite does not support + concurrency (default False). + upsert : {"ignore", "update"}, optional + Upsert mode. If "ignore" is used, the existing elements are ignored. + If "update", the existing elements are updated (default "update"). - Parameters - ---------- - uri : str or Path (must be a file) - The Path to the file to be used. - single_output : bool - If False (default), will create one file per element. The name - of the file will be prefixed with the respective element. - If True, will create only one file, the specified in the URI and - store all the elements in the same file. This behaviour is only - suitable for non-parallel executions. SQLite does not support - concurrency. - upsert : str - Upsert mode. Options are 'ignore' and 'update' (default). If - 'ignore', the existing elements are ignored. If update, the - existing elements are updated. + See Also + -------- + PandasBaseFeatureStorage + + """ + + def __init__( + self, + uri: Union[str, Path], + single_output: bool = False, + upsert: str = "update", + **kwargs: str, + ) -> None: + """Initialize the class. + + Extra Parameters + ---------------- + **kwargs : dict + The keyword arguments passed to the superclass. """ - if upsert not in ['update', 'ignore']: - raise ValueError('upsert must be either "update" or "ignore"') + if upsert not in ["update", "ignore"]: + raise_error( + msg=( + "Invalid choice for `upsert`. " + "Must be either 'update' or 'ignore'." + ) + ) + # Convert str to Path if not isinstance(uri, Path): uri = Path(uri) + # Create parent directories if not present if not uri.parent.exists(): - logger.info(f'Output directory ({uri.parent.as_posix()}) ' - 'does not exist, creating') + logger.info( + f"Output directory ({str(uri.parent.absolute())}) " + "does not exist, creating now." + ) uri.parent.mkdir(parents=True, exist_ok=True) - super().__init__(uri, single_output=single_output) + super().__init__(uri=uri, single_output=single_output, **kwargs) self._upsert = upsert - self._valid_inputs = ['table', 'timeseries'] + self._valid_inputs = ["table", "timeseries"] - def validate(self, input): - """Validate input.""" - if not isinstance(input, list): - input = [input] - return all(x in self._valid_inputs for x in input) - - def get_engine(self, meta=None): - """Get engine.""" - if meta is None: - meta = {} - element = meta.get('element', None) - if self.single_output is False and element is None: - raise ValueError( - 'element must be specified when single_output is False') - prefix = '' - if self.single_output is False: - prefix = element_to_prefix(element) - - uri = f'sqlite:///{self.uri.parent}/{prefix}{self.uri.name}' - return create_engine(uri, echo=False) - - def list_features(self, return_df=False): - """List features.""" - meta_df = pd.read_sql( - 'meta', con=self.get_engine(), index_col='meta_md5') - out = meta_df - if return_df is False: - out = meta_df.to_dict(orient='index') - return out - - def read_df(self, feature_name=None, feature_md5=None): - """Read features from the storage. + def get_engine(self, meta: Optional[Dict] = None) -> "Engine": + """Get engine. Parameters ---------- - feature_name : str - Name of the feature to read. At least one of feature_name or - feature_md5 must be specified. - feature_md5 : str - MD5 of the feature to read. At least one of feature_name or - feature_md5 must be specified. + meta : dict, optional + The metadata as dictionary (default None). + + Returns + ------- + sqlalchemy.Engine + The sqlalcemy engine. + + """ + # Set metadata as empty dictionary if None + if meta is None: + meta = {} + # Retrieve element key from metadata + element = meta.get("element", None) + # Functionality check + if self.single_output is False and element is None: + raise_error( + msg="element must be specified when single_output is False." + ) + # Prefixed elements + prefix = "" + if self.single_output is False: + prefix = element_to_prefix(element) + # Format URI for engine creation + uri = f"sqlite:///{self.uri.parent}/{prefix}{self.uri.name}" + return create_engine(uri, echo=False) + + def _save_upsert( + self, + df: pd.DataFrame, + name: str, + engine: Optional["Engine"] = None, + if_exists: str = "append", + ) -> None: + """Implement UPSERT functionality. + + Parameters + ---------- + df : pandas.DataFrame + DataFrame to save. + name : str + Name of the table to save. + engine : sqlalchemy.Engine, optional + The sqlalchemy engine to use (default None). + if_exists : {"replace", "nocheck", "append", "fail"}, optional + Action to take if the table exists. If "replace", existing table + will be dropped before inserting new values. If "nocheck", + existing table will be ignored. If "append", the data will be + appended to the existing table. If "fail", it will raise an error + (default "append"). + + Raises + ------ + ValueError + If the table exists and if_exists is "fail" or if invalid option is + passed to `if_exists`. + + """ + # Get index names + index_col = df.index.names + # Get sqlalchemy engine if None + if engine is None: + engine = self.get_engine() + # Write data + with engine.begin() as con: + # Check for table's existence + if not inspect(engine).has_table(name): + # New table, so no big issue + df.to_sql(name=name, con=con, if_exists="append") + else: + if if_exists == "replace": + # Replace all the existing elements + df.to_sql(name=name, con=con, if_exists="replace") + elif if_exists == "nocheck": + # Ignore check + df.to_sql(name, con=con, if_exists="append") + elif if_exists == "append": + # TODO: improve + # Step 1: split incoming data into existing and new data + pk_indb = _get_existing_pk( + con, table_name=name, index_col=index_col + ) + existing, new = _split_incoming_data( + df, pk_indb, index_col + ) + # Step 2: upsert existing data + pandas_sql = pandasSQL_builder(con) + pandas_sql.meta.reflect(bind=con, only=[name]) + table = pandas_sql.get_table(name) + update_stmts = NoNewAttributesMixin + if len(existing) > 0 and len(new) > 0: + warn_with_log( + f"Some rows (n={len(existing)}) are already " + "present in the database. The storage is " + f"configured to {self._upsert} the existing " + f"elements. The new rows (n={len(new)}) will be " + "appended. This warning is shown because normally " + "all of the elements should be updated." + ) + if self._upsert == "update": + update_stmts = _generate_update_statements( + table, index_col, existing + ) + for stmt in update_stmts: + con.execute(stmt) + # Step 3: insert new data + new.to_sql(name=name, con=con, if_exists="append") + elif if_exists == "fail": + # Case 4: existing table, so we need to check if the index + # is present or not. + raise_error(msg=f"Table ({name}) already exists.") + else: + raise_error( + msg=f"Invalid option {if_exists} for if_exists." + ) + + # TODO: complete type annotations + def store_2d( + self, + data, + meta: Dict, + columns: Optional[Iterable[str]] = None, + rows_col_name: str = None, + ) -> None: + """Store 2D dataframe. + + Parameters + ---------- + data + meta : dict + The metadata as a dictionary. + columns : list or tuple of str, optional + The columns (default None). + rows_col_name : str, optional + The column name ot use in case number of rows greater than 1. + If None and number of rows greater than 1, then the name will be + "index" (default None). + + """ + n_rows = len(data) + # Convert element metadata to index + idx = element_to_index( + meta=meta, n_rows=n_rows, rows_col_name=rows_col_name + ) + # Prepare new dataframe + data_df = pd.DataFrame(data, columns=columns, index=idx) + # Store dataframe + self.store_df(df=data_df, meta=meta) + + def validate(self, input_: List[str]) -> bool: + """Implement input validation. + + Parameters + ---------- + input_ : list of str + The input to the pipeline step. + + Returns + ------- + bool + Whether the `input` is valid or not. + + """ + # Convert input to list + if not isinstance(input_, list): + input_ = [input_] + + return all(x in self._valid_inputs for x in input_) + + def list_features( + self, return_df: bool = False + ) -> Union[Dict[str, Dict], pd.DataFrame]: + """Implement features listing from the storage. + + Parameters + ---------- + return_df : bool, optional + If True, returns a pandas DataFrame. If False, returns a + dictionary (default False). + + Returns + ------- + dict or pandas.DataFrame + List of features in the storage. If dictionary is returned, the + keys are the feature names to be used in read_features() and the + values are the metadata of each feature. + + """ + meta_df = pd.read_sql( + sql="meta", + con=self.get_engine(), + index_col="meta_md5", + ) + out = meta_df + # Return dictionary + if return_df is False: + out = meta_df.to_dict(orient="index") + return out + + def read_df( + self, + feature_name: Optional[str] = None, + feature_md5: Optional[str] = None, + ) -> pd.DataFrame: + """Implement feature reading from the storage. + + Either one of `feature_name` or `feature_md5` needs to be specified. + + Parameters + ---------- + feature_name : str, optional + Name of the feature to read (default None). + feature_md5 : str, optional + MD5 hash of the feature to read (default None). Returns ------- pandas.DataFrame - The features. + The features as a dataframe. + + Raises + ------ + ValueError + If parameter values are invalid or feature is not found or + mulitple features are found. + """ + # Get sqlalchemy engine engine = self.get_engine() + # Parameter value check if feature_md5 is not None and feature_name is not None: - raise ValueError('Only one of feature_name or feature_md5 can be ' - 'specified') + raise_error( + msg=( + "Only one of `feature_name` or `feature_md5` can be " + "specified." + ) + ) elif feature_md5 is None and feature_name is None: - raise ValueError('At least one of feature_name or feature_md5 ' - 'must be specified') + raise_error( + msg=( + "At least one of `feature_name` or `feature_md5` " + "must be specified." + ) + ) elif feature_md5 is not None: - table_name = f'meta_{feature_md5}' + table_name = f"meta_{feature_md5}" else: meta_df = pd.read_sql( - 'meta', con=engine, index_col='meta_md5') + sql="meta", + con=engine, + index_col="meta_md5", + ) t_df = meta_df.query(f"name == '{feature_name}'") if len(t_df) == 0: - raise ValueError(f'Feature {feature_name} not found') + raise_error(msg=f"Feature {feature_name} not found") elif len(t_df) > 1: - raise ValueError( - f'More than one feature with name {feature_name} found', - 'This file is invalid. You can bypass this issue by ' - 'specifying a feature_md5') - table_name = f'meta_{t_df.index[0]}' - df = pd.read_sql(table_name, con=engine) - # Read the index: - query = ("SELECT ii.name FROM sqlite_master AS m, " - "pragma_index_list(m.name) AS il, " - "pragma_index_info(il.name) AS ii " - f"WHERE tbl_name='{table_name}' " - "ORDER BY cid;") - index_names = pd.read_sql(query, con=engine).values.squeeze().tolist() + raise_error( + msg=( + f"More than one feature with name {feature_name} " + "found. This file is invalid. You can bypass this " + "issue by specifying a `feature_md5`." + ) + ) + table_name = f"meta_{t_df.index[0]}" + # Read metadata from table + df = pd.read_sql(sql=table_name, con=engine) + # Read the index + query = ( + "SELECT ii.name FROM sqlite_master AS m, " + "pragma_index_list(m.name) AS il, " + "pragma_index_info(il.name) AS ii " + f"WHERE tbl_name='{table_name}' " + "ORDER BY cid;" + ) + index_names = ( + pd.read_sql(sql=query, con=engine).values.squeeze().tolist() + ) + # Set index on dataframe df = df.set_index(index_names) return df - def store_metadata(self, meta): - """Store metadata.""" + def store_metadata(self, meta: Dict) -> str: + """Implement metadata storing in the storage. + + Parameters + ---------- + meta : dict + The metadata as a dictionary. + + Returns + ------- + str + The MD5 hash of the metadata prefixed with "meta_" . + + """ + # Copy metadata t_meta = meta.copy() + # Update metadata t_meta.update(self.get_meta()) + # Process metadata meta_md5, t_meta_row = process_meta(t_meta) + # Get sqlalchemy engine engine = self.get_engine(t_meta) if meta_md5 not in inspect(engine).get_table_names(): - meta_df = self._meta_row(t_meta_row, meta_md5) - self._save_upsert(meta_df, 'meta', engine) - return f'meta_{meta_md5}' + # Convert metadata to dataframe + meta_df = self._meta_row(meta=t_meta_row, meta_md5=meta_md5) + # Save dataframe + self._save_upsert(meta_df, "meta", engine) + return f"meta_{meta_md5}" + # TODO: complete type annotations def store_matrix2d( - self, data, meta, col_names=None, rows_col_name=None - ): - """Store 2D matrix.""" + self, + data, + meta: Dict, + col_names: Optional[Iterable[str]] = None, + rows_col_name: str = None, + ) -> None: + """Implement 2D matrix storing. + + Parameters + ---------- + data + meta : dict + The metadata as a dictionary. + col_names : list or tuple of str, optional + The column names (default None). + rows_col_name : str, optional + The column name ot use in case number of rows greater than 1. + If None and number of rows greater than 1, then the name will be + "index" (default None). + + """ # Same as store_2d, but order is important - raise NotImplementedError('store_matrix2d not implemented') + raise_error( + msg="store_matrix2d() not implemented", klass=NotImplementedError + ) - def store_table(self, data, meta, columns=None, rows_col_name=None): - """Store table.""" - self.store_2d(data, meta, columns, rows_col_name) + # TODO: complete type annotations + def store_table( + self, + data, + meta: Dict, + columns: Optional[Iterable[str]] = None, + rows_col_name: str = None, + ) -> None: + """Implement table storing. - def store_2d(self, data, meta, columns=None, rows_col_name=None): - """Store 2D dataframe.""" - n_rows = len(data) - idx = element_to_index( - meta, n_rows=n_rows, rows_col_name=rows_col_name) - data_df = pd.DataFrame(data, columns=columns, index=idx) - self.store_df(data_df, meta) + Parameters + ---------- + data + meta : dict + The metadata as a dictionary. + columns : list or tuple of str, optional + The columns (default None). + rows_col_name : str, optional + The column name ot use in case number of rows greater than 1. + If None and number of rows greater than 1, then the name will be + "index" (default None). - def store_df(self, df, meta): - """Store dataframe.""" + """ + self.store_2d( + data=data, meta=meta, columns=columns, rows_col_name=rows_col_name + ) + + def store_df(self, df: pd.DataFrame, meta: Dict) -> None: + """Implement dataframe storing. + + Parameters + ---------- + df : pandas.DataFrame + The DataFrame to store. + meta : dict + The metadata as a dictionary. + + Raises + ------ + ValueError + If the dataframe index has items that are not in the index + generated from the metadata. + + """ # TODO: Test this function # Check that the index generated by meta matches the one in # the dataframe. @@ -173,130 +480,91 @@ class SQLiteFeatureStorage(PandasFeatureStoreage): # extra element is only one. extra = [x for x in df.index.names if x not in idx.names] if len(extra) > 1: - raise ValueError( - 'The index of the dataframe has extra items that are not ' - 'in the index generated from the meta data.') + raise_error( + "The index of the dataframe has extra items that are not " + "in the index generated from the metadata." + ) elif len(extra) == 1: # The df has one extra index item, this should be the new name # of the missing element in the index idx = element_to_index(meta, rows_col_name=extra[0]) if any(x not in df.index.names for x in idx.names): - raise ValueError( - 'The index of the dataframe is missing index items that are ' - 'generated from the meta data.') - + raise_error( + "The index of the dataframe is missing index items that are " + "generated from the metadata." + ) + # Get table name table_name = self.store_metadata(meta) - # Save + # Get sqlalchemy engine engine = self.get_engine(meta) + # Save data self._save_upsert(df, table_name, engine) - def store_timeseries(self, data, meta): - """Store timeseries.""" - raise NotImplementedError('store_timeseries not implemented') + # TODO: complete type annotations + def store_timeseries(self, data, meta: Dict) -> None: + """Implement timeseries storing. - def collect(self): - """Collect data.""" + Parameters + ---------- + data + meta : dict + The metadata as a dictionary. + + """ + raise_error( + msg="store_timeseries() not implemented.", + klass=NotImplementedError, + ) + + def collect(self) -> None: + """Implement data collection. + + Raises + ------ + NotImplementedError + If `single_output` is True. + + """ if self.single_output is True: - raise ValueError('collect is not implemented for single output') - logger.info( - f'Collecting data from {self.uri.parent}/*{self.uri.name}') - + raise_error(msg="collect() is not implemented for single output.") + logger.info(f"Collecting data from {self.uri.parent}/*{self.uri.name}") + # Create new instance out_storage = SQLiteFeatureStorage( - uri=self.uri, single_output=True, upsert='ignore') - - files = self.uri.parent.glob(f'*{self.uri.name}') - for elem in tqdm.tqdm(files, desc='file'): - logger.debug(f'Reading from {elem.as_posix()}') + uri=self.uri, single_output=True, upsert="ignore" + ) + # Glob files + files = self.uri.parent.glob(f"*{self.uri.name}") + for elem in tqdm(files, desc="file"): + logger.debug(f"Reading from {str(elem.absolute())}") in_storage = SQLiteFeatureStorage(uri=elem, single_output=True) in_engine = in_storage.get_engine() # Open "meta" table t_meta_df = pd.read_sql( - 'meta', con=in_engine, index_col='meta_md5') - out_storage._save_upsert(t_meta_df, 'meta') - for meta_md5 in tqdm.tqdm(t_meta_df.index, desc='feature'): - logger.debug(f'Collecting feature {meta_md5}') + sql="meta", con=in_engine, index_col="meta_md5" + ) + # Save metadata + out_storage._save_upsert(t_meta_df, "meta") + # Save dataframes + for meta_md5 in tqdm(t_meta_df.index, desc="feature"): + logger.debug(f"Collecting feature {meta_md5}") # TODO: Fix this, needs that read_feature sets the index # properly - table_name = f'meta_{meta_md5}' + table_name = f"meta_{meta_md5}" t_df = in_storage.read_df(feature_md5=meta_md5) - out_storage._save_upsert(t_df, table_name, if_exist='nocheck') - - def _save_upsert(self, df, name, engine=None, if_exist='append'): - """Implement of UPSERT functionality. - - Parameters - ---------- - df : pandas.DataFrame - DataFrame to save - name : str - Name of the table to save - if_exist : str - If the table exists, the behavior is controlled by this parameter. - Options are 'append' (default) and 'fail'. If 'fail' and the table - exists, it will raise an error. If 'append', the data will be - appended to the existing table (following the upsert mode). - - Raises - ______ - ValueError - If the table exists and if_exist is 'fail' - """ - index_col = df.index.names - if engine is None: - engine = self.get_engine() - with engine.begin() as con: - if if_exist == 'replace': - # Case 1: replace all the existing elements - df.to_sql(name, con=con, if_exists='replace') - elif not inspect(engine).has_table(name): - # Case 2: new table, so no big issue - df.to_sql(name, con=con, if_exists='append') - elif if_exist == 'nocheck': - # Case 3: existing table, but we will not check for - # existing - df.to_sql(name, con=con, if_exists='append') - else: - # Case 4: existing table, so we need to check if the index - # is present or not. - if if_exist == 'fail': - raise ValueError(f"Table ({name}) already exists") - - # Step 1: split incoming data into existing and new data - pk_indb = _get_existing_pk( - con, table_name=name, index_col=index_col) - existing, new = _split_incoming_data(df, pk_indb, index_col) - - # Step 2: upsert existing data - pandas_sql = pandasSQL_builder(con) - pandas_sql.meta.reflect(bind=con, only=[name]) - table = pandas_sql.get_table(name) - update_stmts = NoNewAttributesMixin - if len(existing) > 0 and len(new) > 0: - warn( - f"Some rows (n={len(existing)}) are already present " - "in the database. The storage is configured to " - f"{self._upsert} the existing elements. The new rows " - f"(n={len(new)}) will be appended. This warning " - "is shown because normally all of the elements should " - "be updated") - if self._upsert == 'update': - update_stmts = _generate_update_statements( - table, index_col, existing) - for stmt in update_stmts: - con.execute(stmt) - - # Step 3: insert new data - new.to_sql(name, con=con, if_exists='append') + # Save data + out_storage._save_upsert(t_df, table_name, if_exist="nocheck") +# TODO: refactor def _get_existing_pk(con, table_name, index_col): - pk_cols = ', '.join(index_col) - query = f'SELECT {pk_cols} FROM {table_name};' + pk_cols = ", ".join(index_col) + query = f"SELECT {pk_cols} FROM {table_name};" pk_indb = pd.read_sql(query, con=con) return pk_indb +# TODO: refactor def _split_incoming_data(df, pk_indb, index_col): incoming_pk = df.reset_index()[index_col] exists_mask = ( @@ -308,6 +576,7 @@ def _split_incoming_data(df, pk_indb, index_col): return existing, new +# TODO: refactor def _generate_update_statements(table, index_col, rows_to_update): from sqlalchemy import and_ @@ -319,8 +588,9 @@ def _generate_update_statements(table, index_col, rows_to_update): for i, (_, keys) in enumerate(pk_indb.iterrows()): stmt = ( table.update() - .where(and_(col == keys[j] - for j, col in enumerate(pk_cols))) # type: ignore + .where( + and_(col == keys[j] for j, col in enumerate(pk_cols)) + ) # type: ignore .values(new_records[i]) ) stmts.append(stmt) -- 2.52.0 From 17ba60cd549db3ec7dfbb8e5369f4b19518f640c Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 10 Aug 2022 15:13:22 +0200 Subject: [PATCH 171/287] update: improve storage sub-package import --- junifer/storage/__init__.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/junifer/storage/__init__.py b/junifer/storage/__init__.py index 81239bdf4..b8c4f8942 100644 --- a/junifer/storage/__init__.py +++ b/junifer/storage/__init__.py @@ -1,3 +1,9 @@ +"""Provide imports for storage sub-package.""" + # Authors: Federico Raimondo +# Synchon Mandal # License: AGPL -from .sqlite import SQLiteFeatureStorage \ No newline at end of file + +from .base import BaseFeatureStorage +from .pandas_base import PandasBaseFeatureStorage +from .sqlite import SQLiteFeatureStorage -- 2.52.0 From a969d79c4555422c6cacf03794ec96c23495480b Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 10 Aug 2022 22:25:00 +0200 Subject: [PATCH 172/287] fix: update __init__ for PandasBaseFeatureStorage --- junifer/storage/pandas_base.py | 18 +++++++++++++++--- 1 file changed, 15 insertions(+), 3 deletions(-) diff --git a/junifer/storage/pandas_base.py b/junifer/storage/pandas_base.py index a6a1b430a..eaa2a0a32 100644 --- a/junifer/storage/pandas_base.py +++ b/junifer/storage/pandas_base.py @@ -5,7 +5,8 @@ # License: AGPL import json -from typing import Dict +from pathlib import Path +from typing import Dict, Union import pandas as pd @@ -18,15 +19,26 @@ class PandasBaseFeatureStorage(BaseFeatureStorage): For every interface that is required, one needs to provide a concrete implementation of this abstract class. + Parameters + ---------- + uri : str or pathlib.Path + The path to the storage. + single_output : bool, optional + Whether to have single output (default False). + **kwargs + Keyword arguments passed to superclass. + See Also -------- BaseFeatureStorage """ - def __init__(**kwargs) -> None: + def __init__( + self, uri: Union[str, Path], single_output: bool = False, **kwargs + ) -> None: """Initialize the class.""" - super().__init__(**kwargs) + super().__init__(uri=uri, single_output=single_output, **kwargs) def _meta_row(self, meta: Dict, meta_md5: str) -> pd.DataFrame: """Convert the metadata to a pandas DataFrame. -- 2.52.0 From 575039e6d5e94f8a8a4666b0741e369cd22c142e Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 10 Aug 2022 22:26:00 +0200 Subject: [PATCH 173/287] fix: correct key usage in process_meta() --- junifer/storage/utils.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/junifer/storage/utils.py b/junifer/storage/utils.py index d24b36c44..42a71f551 100644 --- a/junifer/storage/utils.py +++ b/junifer/storage/utils.py @@ -74,7 +74,7 @@ def process_meta(meta: Dict) -> Tuple[str, Dict]: if isinstance(element, dict): meta["_element_keys"] = list(element.keys()) else: - meta["_element_keyes"] = ["element"] + meta["_element_keys"] = ["element"] # MD5 hash of the metadata md5_hash = _meta_hash(meta) return md5_hash, meta -- 2.52.0 From d65c8ff39643643aa6b236d23e1e030984d224cb Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 10 Aug 2022 22:26:39 +0200 Subject: [PATCH 174/287] fix: correct default value for rows_col_name in element_to_index() --- junifer/storage/utils.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/junifer/storage/utils.py b/junifer/storage/utils.py index 42a71f551..bc0bd8faf 100644 --- a/junifer/storage/utils.py +++ b/junifer/storage/utils.py @@ -129,7 +129,7 @@ def element_to_index( Number of rows to create (default 1). rows_col_name: str, optional The column name to use in case `n_rows` > 1. If None and - n_rows > 1, the name will be "index" (default None). + n_rows > 1, the name will be "idx" (default None). Returns ------- @@ -152,7 +152,7 @@ def element_to_index( element = {"element": element} # Check rows_col_name if rows_col_name is None: - rows_col_name = "index" + rows_col_name = "idx" elem_idx = {k: [v] * n_rows for k, v in element.items()} elem_idx[rows_col_name] = np.arange(n_rows) # Create index -- 2.52.0 From a901f44acd0527e75c83a41879d8d0455495eee1 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 10 Aug 2022 22:28:06 +0200 Subject: [PATCH 175/287] refactor: separate module for unit tests of storage/utils.py --- junifer/storage/tests/test_base.py | 175 +----------------------- junifer/storage/tests/test_utils.py | 200 ++++++++++++++++++++++++++++ 2 files changed, 202 insertions(+), 173 deletions(-) create mode 100644 junifer/storage/tests/test_utils.py diff --git a/junifer/storage/tests/test_base.py b/junifer/storage/tests/test_base.py index 4c433735f..29eca2d85 100644 --- a/junifer/storage/tests/test_base.py +++ b/junifer/storage/tests/test_base.py @@ -2,142 +2,7 @@ import pytest -from junifer.storage.base import (process_meta, element_to_index, - BaseFeatureStorage, - element_to_prefix) - - -def test_process_meta_hash(): - """Test metadata hash.""" - meta = None - with pytest.raises(ValueError, match=r"Meta must be a dict"): - process_meta(meta) - - meta = {'element': 'foo', 'A': 1, 'B': [2, 3, 4, 5, 6]} - hash1, _ = process_meta(meta) - - meta = {'element': 'foo', 'B': [2, 3, 4, 5, 6], 'A': 1} - hash2, _ = process_meta(meta) - assert hash1 == hash2 - - meta = {'element': 'foo', 'A': 1, 'B': [2, 3, 1, 5, 6]} - hash3, _ = process_meta(meta) - assert hash1 != hash3 - - meta1 = { - 'element': 'foo', - 'B': { - 'B2': [2, 3, 4, 5, 6], - 'B1': [9.22, 3.14, 1.41, 5.67, 6.28], - 'B3': (1, 'car'), - }, - 'A': 1} - - meta2 = { - 'A': 1, - 'B': { - 'B3': (1, 'car'), - 'B1': [9.22, 3.14, 1.41, 5.67, 6.28], - 'B2': [2, 3, 4, 5, 6], - }, - 'element': 'foo' - } - - hash4, _ = process_meta(meta1) - hash5, _ = process_meta(meta2) - assert hash4 == hash5 - - -def test_process_meta_element(): - """Test metadata element.""" - meta = {} - with pytest.raises(ValueError, match=r"_element_keys"): - process_meta(meta) - - meta = {'element': 'foo', 'A': 1, 'B': [2, 3, 4, 5, 6]} - _, new_meta = process_meta(meta) - assert '_element_keys' in new_meta - assert new_meta['_element_keys'] == ['element'] - assert 'A' in new_meta - assert 'B' in new_meta - - meta = { - 'element': {'subject': 'foo', 'session': 'bar'}, - 'B': [2, 3, 4, 5, 6], 'A': 1} - _, new_meta = process_meta(meta) - assert '_element_keys' in new_meta - assert new_meta['_element_keys'] == ['subject', 'session'] - assert 'A' in new_meta - assert 'B' in new_meta - - -def test_process_meta_index(): - """Test metadata element to index.""" - meta = {'noelement': 'foo'} - with pytest.raises(ValueError, match=r'meta must contain the key'): - element_to_index(meta) - - meta = {'element': 'foo', 'A': 1, 'B': [2, 3, 4, 5, 6]} - index = element_to_index(meta) - assert index.names == ['element', 'idx'] - assert index.levels[0].name == 'element' - assert index.levels[0].values[0] == 'foo' - assert all(x == 'foo' for x in index.levels[0].values) - assert index.levels[0].values.shape == (1,) - assert index.levels[1].name == 'idx' - assert all(x == i for i, x in enumerate(index.levels[1].values)) - assert index.levels[1].values.shape == (1,) - - index = element_to_index(meta, n_rows=10) - assert index.names == ['element', 'idx'] - assert index.levels[0].name == 'element' - assert all(x == 'foo' for x in index.levels[0].values) - assert index.levels[0].values.shape == (1,) - - assert index.levels[1].name == 'idx' - assert all(x == i for i, x in enumerate(index.levels[1].values)) - assert index.levels[1].values.shape == (10,) - - index = element_to_index(meta, n_rows=1, rows_col_name='scan') - assert index.names == ['element', 'scan'] - assert index.levels[0].name == 'element' - assert index.levels[0].values[0] == 'foo' - assert all(x == 'foo' for x in index.levels[0].values) - assert index.levels[0].values.shape == (1,) - assert index.levels[1].name == 'scan' - assert all(x == i for i, x in enumerate(index.levels[1].values)) - assert index.levels[1].values.shape == (1,) - - index = element_to_index(meta, n_rows=7, rows_col_name='scan') - assert index.names == ['element', 'scan'] - assert index.levels[0].name == 'element' - assert all(x == 'foo' for x in index.levels[0].values) - assert index.levels[0].values.shape == (1,) - - assert index.levels[1].name == 'scan' - assert all(x == i for i, x in enumerate(index.levels[1].values)) - assert index.levels[1].values.shape == (7,) - - meta = { - 'element': {'subject': 'sub-01', 'session': 'ses-01'}, - 'A': 1, 'B': [2, 3, 4, 5, 6]} - index = element_to_index(meta, n_rows=10) - - assert index.levels[0].name == 'subject' - assert all(x == 'sub-01' for x in index.levels[0].values) - assert index.levels[0].values.shape == (1,) - - assert index.levels[1].name == 'session' - assert all(x == 'ses-01' for x in index.levels[1].values) - assert index.levels[1].values.shape == (1,) - - assert index.levels[2].name == 'idx' - assert all(x == i for i, x in enumerate(index.levels[2].values)) - assert index.levels[2].values.shape == (10,) - - -def test_BaseFeatureStorage(): - """Test BaseFeatureStorage.""" +from junifer.storage.base import BaseFeatureStorage with pytest.raises(TypeError, match=r"abstract"): BaseFeatureStorage(uri='/tmp') # type: ignore @@ -206,40 +71,4 @@ def test_BaseFeatureStorage(): with pytest.raises(NotImplementedError): st.collect() - assert st.uri == '/tmp' - - -def test_element_to_prefix(): - """Test converting element to prefix (for file naming).""" - - element = 'sub-01' - prefix = element_to_prefix(element) - assert prefix == 'element_sub-01_' - - element = 1 - prefix = element_to_prefix(element) - assert prefix == 'element_1_' - - element = {'subject': 'sub-01'} - prefix = element_to_prefix(element) - assert prefix == 'element_sub-01_' - - element = {'subject': 1} - prefix = element_to_prefix(element) - assert prefix == 'element_1_' - - element = {'subject': 'sub-01', 'session': 'ses-02'} - prefix = element_to_prefix(element) - assert prefix == 'element_sub-01_ses-02_' - - element = {'subject': 1, 'session': 2} - prefix = element_to_prefix(element) - assert prefix == 'element_1_2_' - - element = ('sub-01', 'ses-02') - prefix = element_to_prefix(element) - assert prefix == 'element_sub-01_ses-02_' - - element = (1, 2) - prefix = element_to_prefix(element) - assert prefix == 'element_1_2_' + assert st.uri == "/tmp" diff --git a/junifer/storage/tests/test_utils.py b/junifer/storage/tests/test_utils.py new file mode 100644 index 000000000..751410b85 --- /dev/null +++ b/junifer/storage/tests/test_utils.py @@ -0,0 +1,200 @@ +"""Provide tests for utils.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from typing import Dict, List, Tuple, Union + +import pytest + +from junifer.storage.utils import ( + element_to_index, + element_to_prefix, + process_meta, +) + + +def test_process_meta_invalid_metadata_type() -> None: + """Test invalid metadata type check for metadata hash processing.""" + meta = None + with pytest.raises(ValueError, match=r"`meta` must be a dict"): + process_meta(meta) + + +# TODO: parameterize +def test_process_meta_hash() -> None: + """Test metadata hash processing.""" + meta = {"element": "foo", "A": 1, "B": [2, 3, 4, 5, 6]} + hash1, _ = process_meta(meta) + + meta = {"element": "foo", "B": [2, 3, 4, 5, 6], "A": 1} + hash2, _ = process_meta(meta) + assert hash1 == hash2 + + meta = {"element": "foo", "A": 1, "B": [2, 3, 1, 5, 6]} + hash3, _ = process_meta(meta) + assert hash1 != hash3 + + meta1 = { + "element": "foo", + "B": { + "B2": [2, 3, 4, 5, 6], + "B1": [9.22, 3.14, 1.41, 5.67, 6.28], + "B3": (1, "car"), + }, + "A": 1, + } + + meta2 = { + "A": 1, + "B": { + "B3": (1, "car"), + "B1": [9.22, 3.14, 1.41, 5.67, 6.28], + "B2": [2, 3, 4, 5, 6], + }, + "element": "foo", + } + + hash4, _ = process_meta(meta1) + hash5, _ = process_meta(meta2) + assert hash4 == hash5 + + +def test_process_meta_invalid_metadata_key() -> None: + """Test invalid metadata key check for metadata hash processing.""" + meta = {} + with pytest.raises(ValueError, match=r"_element_keys"): + process_meta(meta) + + +@pytest.mark.parametrize( + "meta,elements", + [ + ({"element": "foo", "A": 1, "B": [2, 3, 4, 5, 6]}, ["element"]), + ( + { + "element": {"subject": "foo", "session": "bar"}, + "B": [2, 3, 4, 5, 6], + "A": 1, + }, + ["subject", "session"], + ), + ], +) +def test_process_meta_element(meta: Dict, elements: List[str]) -> None: + """Test metadata element after processing. + + Parameters + ---------- + meta : dict + The parametrized metadata dictionary. + elements : list of str + The parametrized elements to assert against. + + """ + _, processed_meta = process_meta(meta) + assert "_element_keys" in processed_meta + assert processed_meta["_element_keys"] == elements + assert "A" in processed_meta + assert "B" in processed_meta + + +@pytest.mark.parametrize( + "element,prefix", + [ + ("sub-01", "element_sub-01_"), + (1, "element_1_"), + ({"subject": "sub-01"}, "element_sub-01_"), + ({"subject": 1}, "element_1_"), + ({"subject": "sub-01", "session": "ses-02"}, "element_sub-01_ses-02_"), + ({"subject": 1, "session": 2}, "element_1_2_"), + (("sub-01", "ses-02"), "element_sub-01_ses-02_"), + ((1, 2), "element_1_2_"), + ], +) +def test_element_to_prefix( + element: Union[str, int, Dict, Tuple], prefix: str +) -> None: + """Test converting element to prefix (for file naming). + + Parameters + ---------- + element : str, int, dict or tuple + The parameterized element. + prefix : str + The parametrized prefix to assert against. + + """ + prefix_generated = element_to_prefix(element) + assert prefix_generated == prefix + + +def test_element_to_index_check_meta_invalid_key() -> None: + """Test element to index metadata key checking.""" + meta = {"noelement": "foo"} + with pytest.raises(ValueError, match=r"metadata must contain the key"): + element_to_index(meta) + + +def test_element_to_index() -> None: + """Test element to index.""" + meta = {"element": "foo", "A": 1, "B": [2, 3, 4, 5, 6]} + index = element_to_index(meta) + assert index.names == ["element", "idx"] + assert index.levels[0].name == "element" + assert index.levels[0].values[0] == "foo" + assert all(x == "foo" for x in index.levels[0].values) + assert index.levels[0].values.shape == (1,) + assert index.levels[1].name == "idx" + assert all(x == i for i, x in enumerate(index.levels[1].values)) + assert index.levels[1].values.shape == (1,) + + index = element_to_index(meta, n_rows=10) + assert index.names == ["element", "idx"] + assert index.levels[0].name == "element" + assert all(x == "foo" for x in index.levels[0].values) + assert index.levels[0].values.shape == (1,) + + assert index.levels[1].name == "idx" + assert all(x == i for i, x in enumerate(index.levels[1].values)) + assert index.levels[1].values.shape == (10,) + + index = element_to_index(meta, n_rows=1, rows_col_name="scan") + assert index.names == ["element", "scan"] + assert index.levels[0].name == "element" + assert index.levels[0].values[0] == "foo" + assert all(x == "foo" for x in index.levels[0].values) + assert index.levels[0].values.shape == (1,) + assert index.levels[1].name == "scan" + assert all(x == i for i, x in enumerate(index.levels[1].values)) + assert index.levels[1].values.shape == (1,) + + index = element_to_index(meta, n_rows=7, rows_col_name="scan") + assert index.names == ["element", "scan"] + assert index.levels[0].name == "element" + assert all(x == "foo" for x in index.levels[0].values) + assert index.levels[0].values.shape == (1,) + + assert index.levels[1].name == "scan" + assert all(x == i for i, x in enumerate(index.levels[1].values)) + assert index.levels[1].values.shape == (7,) + + meta = { + "element": {"subject": "sub-01", "session": "ses-01"}, + "A": 1, + "B": [2, 3, 4, 5, 6], + } + index = element_to_index(meta, n_rows=10) + + assert index.levels[0].name == "subject" + assert all(x == "sub-01" for x in index.levels[0].values) + assert index.levels[0].values.shape == (1,) + + assert index.levels[1].name == "session" + assert all(x == "ses-01" for x in index.levels[1].values) + assert index.levels[1].values.shape == (1,) + + assert index.levels[2].name == "idx" + assert all(x == i for i, x in enumerate(index.levels[2].values)) + assert index.levels[2].values.shape == (10,) -- 2.52.0 From 9b24684541fd35c36f9d98c7ffb76150226efcc4 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 10 Aug 2022 22:30:11 +0200 Subject: [PATCH 176/287] refactor: prune test_base.py and add type annotations --- junifer/storage/tests/test_base.py | 23 +++++++++++++++++------ 1 file changed, 17 insertions(+), 6 deletions(-) diff --git a/junifer/storage/tests/test_base.py b/junifer/storage/tests/test_base.py index 29eca2d85..6a69d9503 100644 --- a/junifer/storage/tests/test_base.py +++ b/junifer/storage/tests/test_base.py @@ -1,11 +1,23 @@ """Provide tests for base.""" +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + import pytest from junifer.storage.base import BaseFeatureStorage - with pytest.raises(TypeError, match=r"abstract"): - BaseFeatureStorage(uri='/tmp') # type: ignore + +def test_BaseFeatureStorage_abstractness() -> None: + """Test BaseFeatureStorage is abstract base class.""" + with pytest.raises(TypeError, match=r"abstract"): + BaseFeatureStorage(uri="/tmp") + + +def test_BaseFeatureStorage() -> None: + """Test BaseFeatureStorage.""" + # Create concrete class class MyFeatureStorage(BaseFeatureStorage): def __init__(self, uri, single_output=False): super().__init__(uri, single_output=single_output) @@ -17,8 +29,7 @@ from junifer.storage.base import BaseFeatureStorage super().list_features() def read_df(self, feature_name=None, feature_md5=None): - super().read_df( - feature_name=feature_name, feature_md5=feature_md5) + super().read_df(feature_name=feature_name, feature_md5=feature_md5) def store_metadata(self, metadata): super().store_metadata(metadata) @@ -38,10 +49,10 @@ from junifer.storage.base import BaseFeatureStorage def collect(self): return super().collect() - st = MyFeatureStorage(uri='/tmp') + st = MyFeatureStorage(uri="/tmp") assert st.single_output is False - st = MyFeatureStorage(uri='/tmp', single_output=True) + st = MyFeatureStorage(uri="/tmp", single_output=True) assert st.single_output is True with pytest.raises(NotImplementedError): -- 2.52.0 From 682c4dd1156684a617ca5657e27c8a28477773cd Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 10 Aug 2022 22:31:08 +0200 Subject: [PATCH 177/287] fix: correct spelling in sqlite.py --- junifer/storage/sqlite.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py index 20326ecf2..c5a5690d0 100644 --- a/junifer/storage/sqlite.py +++ b/junifer/storage/sqlite.py @@ -95,7 +95,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): Returns ------- sqlalchemy.Engine - The sqlalcemy engine. + The sqlalchemy engine. """ # Set metadata as empty dictionary if None -- 2.52.0 From 35a9afb13110afa034c30c7db1ef425577662099 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 11 Aug 2022 14:24:41 +0200 Subject: [PATCH 178/287] refactor: add docstrings and type annotations in test_sqlite.py --- junifer/storage/tests/test_sqlite.py | 868 ++++++++++++++++----------- 1 file changed, 534 insertions(+), 334 deletions(-) diff --git a/junifer/storage/tests/test_sqlite.py b/junifer/storage/tests/test_sqlite.py index 25cc75662..e5c251553 100644 --- a/junifer/storage/tests/test_sqlite.py +++ b/junifer/storage/tests/test_sqlite.py @@ -1,394 +1,594 @@ """Provide tests for sqlite.""" +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + from pathlib import Path + import numpy as np import pandas as pd -from pandas.testing import assert_frame_equal -import tempfile -from sqlalchemy import create_engine import pytest +from pandas.testing import assert_frame_equal +from sqlalchemy import create_engine from junifer.storage.sqlite import SQLiteFeatureStorage -from junifer.storage.base import (process_meta, element_to_prefix, - element_to_index) +from junifer.storage.utils import ( + element_to_index, + element_to_prefix, + process_meta, +) -df1 = pd.DataFrame({ - 'element': [1, 2, 3, 4, 5], - 'pk2': ['a', 'b', 'c', 'd', 'e'], - 'col1': [11, 22, 33, 44, 55], - 'col2': [111, 222, 333, 444, 555] -}).set_index(['element', 'pk2']) +df1 = pd.DataFrame( + { + "element": [1, 2, 3, 4, 5], + "pk2": ["a", "b", "c", "d", "e"], + "col1": [11, 22, 33, 44, 55], + "col2": [111, 222, 333, 444, 555], + } +).set_index(["element", "pk2"]) -df2 = pd.DataFrame({ - 'element': [2, 5, 6], - 'pk2': ['b', 'e', 'f'], - 'col1': [2222, 5555, 66], - 'col2': [22222, 55555, 666] -}).set_index(['element', 'pk2']) +df2 = pd.DataFrame( + { + "element": [2, 5, 6], + "pk2": ["b", "e", "f"], + "col1": [2222, 5555, 66], + "col2": [22222, 55555, 666], + } +).set_index(["element", "pk2"]) -df_update = pd.DataFrame({ - 'element': [1, 2, 3, 4, 5, 6], - 'pk2': ['a', 'b', 'c', 'd', 'e', 'f'], - 'col1': [11, 2222, 33, 44, 5555, 66], - 'col2': [111, 22222, 333, 444, 55555, 666] -}).set_index(['element', 'pk2']) +df_update = pd.DataFrame( + { + "element": [1, 2, 3, 4, 5, 6], + "pk2": ["a", "b", "c", "d", "e", "f"], + "col1": [11, 2222, 33, 44, 5555, 66], + "col2": [111, 22222, 333, 444, 55555, 666], + } +).set_index(["element", "pk2"]) -df_ignore = pd.DataFrame({ - 'element': [1, 2, 3, 4, 5, 6], - 'pk2': ['a', 'b', 'c', 'd', 'e', 'f'], - 'col1': [11, 22, 33, 44, 55, 66], - 'col2': [111, 222, 333, 444, 555, 666] -}).set_index(['element', 'pk2']) +df_ignore = pd.DataFrame( + { + "element": [1, 2, 3, 4, 5, 6], + "pk2": ["a", "b", "c", "d", "e", "f"], + "col1": [11, 22, 33, 44, 55, 66], + "col2": [111, 222, 333, 444, 555, 666], + } +).set_index(["element", "pk2"]) -def _read_sql(table_name, uri, index_col): - engine = create_engine(f'sqlite:///{uri}', echo=False) - df = pd.read_sql(table_name, con=engine, index_col=index_col) +def _read_sql(table_name: str, uri: str, index_col: str) -> pd.DataFrame: + """Read database table into a pandas DataFrame. + + Parameters + ---------- + table_name : str + The table name. + uri : str + The URI of the database. + index_col : str + The index column name. + + Returns + ------- + pandas.DataFrame + The contents of the table in a DataFrame. + + """ + engine = create_engine(f"sqlite:///{uri}", echo=False) + df = pd.read_sql(sql=table_name, con=engine, index_col=index_col) return df -def test_get_engine(): - """Test engine retrieval.""" - with tempfile.TemporaryDirectory() as _tmpdir: - uri = f'{_tmpdir}/test.db' - # Single storage, must be the uri - storage = SQLiteFeatureStorage( - uri=uri, single_output=True, upsert='ignore') - assert storage.single_output is True - engine = storage.get_engine() - assert engine.url.drivername == 'sqlite' - assert f'{engine.url.database}' == uri +def test_get_engine_single_output(tmp_path: Path) -> None: + """Test engine retrieval with single output. - storage = SQLiteFeatureStorage( - uri=uri, single_output=False, upsert='ignore') - with pytest.raises(ValueError, match='element must be specified'): - storage.get_engine() + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. - tocreate = Path(_tmpdir) / 'tocreate' - assert not tocreate.exists() - uri = f'{tocreate.as_posix()}/test.db' - storage = SQLiteFeatureStorage( - uri=uri, single_output=True, upsert='ignore') - assert tocreate.exists() + """ + uri = tmp_path / "test_single_output.db" + # Single storage, must be the uri + storage = SQLiteFeatureStorage( + uri=uri, single_output=True, upsert="ignore" + ) + assert storage.single_output is True + engine = storage.get_engine() + assert engine.url.drivername == "sqlite" + assert engine.url.database == str(uri.absolute()) -def test_store_metadata(): - """Test metadata store.""" - with tempfile.TemporaryDirectory() as _tmpdir: - uri = f'{_tmpdir}/test.db' - # Single storage, must be the uri - storage = SQLiteFeatureStorage( - uri=uri, single_output=True, upsert='ignore') - meta = {'element': 'test', 'version': '0.0.1'} - table_name = storage.store_metadata(meta) - assert table_name.startswith('meta_') +def test_get_engine_multi_output(tmp_path: Path) -> None: + """Test engine retrieval with multi output. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + uri = tmp_path / "test_multi_output.db" + storage = SQLiteFeatureStorage( + uri=uri, single_output=False, upsert="ignore" + ) + with pytest.raises(ValueError, match="element must be specified"): + storage.get_engine() -def test_upsert_replace(): - """Test dataframe store (upsert=replace).""" - with tempfile.TemporaryDirectory() as _tmpdir: - uri = f'{_tmpdir}/test.db' - # Single storage, must be the uri - storage = SQLiteFeatureStorage( - uri=uri, single_output=True, upsert='ignore') - meta = {'element': 'test', 'version': '0.0.1'} +def test_get_engine_single_output_creation(tmp_path: Path) -> None: + """Test engine retrieval with single output creation. - # Save to SQL - storage.store_df(df1, meta) + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. - # Test the internals - table_name = storage.store_metadata(meta) - - c_df1 = _read_sql(table_name, uri=uri, index_col=['element', 'pk2']) - assert_frame_equal(df1, c_df1) - - storage._save_upsert(df2, table_name, if_exist='replace') - - c_df2 = _read_sql(table_name, uri=uri, index_col=['element', 'pk2']) - assert_frame_equal(df2, c_df2) + """ + tocreate = tmp_path / "tocreate" + # Path does not exist yet + assert not tocreate.exists() + uri = tocreate.absolute() / "test_single_output.db" + _ = SQLiteFeatureStorage(uri=uri, single_output=True, upsert="ignore") + # Path exists now + assert tocreate.exists() -def test_upsert_ignore(): - """Test dataframe store (upsert=ignore).""" - with tempfile.TemporaryDirectory() as _tmpdir: - uri = f'{_tmpdir}/test.db' - with pytest.raises(ValueError): - SQLiteFeatureStorage(uri=uri, single_output=True, upsert='wrong') +def test_upsert_replace(tmp_path: Path) -> None: + """Test dataframe store with upsert=replace. - storage = SQLiteFeatureStorage( - uri=uri, single_output=True, upsert='ignore') - meta = {'element': 'test', 'version': '0.0.1'} + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. - # Save to SQL - storage.store_df(df1, meta) - - # Test the internals - table_name = storage.store_metadata(meta) - - c_df1 = _read_sql(table_name, uri=uri, index_col=['element', 'pk2']) - assert_frame_equal(df1, c_df1) - - with pytest.warns(RuntimeWarning, match='are already present'): - storage.store_df(df2, meta) - - c_dfignore = _read_sql( - table_name, uri=uri, index_col=['element', 'pk2']) - assert_frame_equal(c_dfignore, df_ignore) - - with pytest.raises(ValueError, match=r"already exists"): - storage._save_upsert(df2, table_name, if_exist='fail') + """ + uri = tmp_path / "test_upsert_replace.db" + # Single storage, must be the uri + storage = SQLiteFeatureStorage( + uri=uri, single_output=True, upsert="ignore" + ) + # Metadata to store + meta = {"element": "test", "version": "0.0.1"} + # Save to database + storage.store_df(df1, meta) + # Store metadata + table_name = storage.store_metadata(meta) + # Read stored table + c_df1 = _read_sql( + table_name=table_name, uri=uri, index_col=["element", "pk2"] + ) + # Check if dataframes are equal + assert_frame_equal(df1, c_df1) + # Upsert using replace + storage._save_upsert(df2, table_name, if_exist="replace") + # Read stored table + c_df2 = _read_sql( + table_name=table_name, uri=uri, index_col=["element", "pk2"] + ) + # Check if dataframes are equal + assert_frame_equal(df2, c_df2) -def test_upsert_update(): - """Test dataframe store (upsert=delete).""" - meta = {'element': 'test', 'version': '0.0.1'} - with tempfile.TemporaryDirectory() as _tmpdir: - uri = f'{_tmpdir}/test.db' - storage = SQLiteFeatureStorage(uri=uri, single_output=True) +def test_upsert_ignore(tmp_path: Path) -> None: + """Test dataframe store with upsert=ignore. - # Save to SQL - storage.store_df(df1, meta) - - # Test the internals - table_name = storage.store_metadata(meta) - - c_df1 = _read_sql(table_name, uri=uri, index_col=['element', 'pk2']) - assert_frame_equal(df1, c_df1) + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + """ + uri = tmp_path / "test_upsert_ignore.db" + # Single storage, must be the uri + storage = SQLiteFeatureStorage( + uri=uri, single_output=True, upsert="ignore" + ) + # Metadata to store + meta = {"element": "test", "version": "0.0.1"} + # Save to database + storage.store_df(df1, meta) + # Store metadata + table_name = storage.store_metadata(meta) + # Read stored table + c_df1 = _read_sql( + table_name=table_name, uri=uri, index_col=["element", "pk2"] + ) + # Check if dataframes are equal + assert_frame_equal(df1, c_df1) + # Check for warning + with pytest.warns(RuntimeWarning, match="are already present"): storage.store_df(df2, meta) - - c_dfupdate = _read_sql( - table_name, uri=uri, index_col=['element', 'pk2']) - assert_frame_equal(c_dfupdate, df_update) + # Read stored table + c_dfignore = _read_sql(table_name, uri=uri, index_col=["element", "pk2"]) + # Check if dataframes are equal + assert_frame_equal(c_dfignore, df_ignore) + # Check for error + with pytest.raises(ValueError, match=r"already exists"): + storage._save_upsert(df2, table_name, if_exist="fail") -def test_store_read_df(): - """Test dataframe store.""" - with tempfile.TemporaryDirectory() as _tmpdir: - uri = f'{_tmpdir}/test.db' - storage = SQLiteFeatureStorage( - uri=uri, single_output=True, upsert='ignore') - meta = { - 'element': 'test', 'version': '0.0.1', - 'marker': {'name': 'fcname'}} +def test_upsert_update(tmp_path: Path) -> None: + """Test dataframe store with upsert=delete. - to_store = df1[['col1', 'col2']] + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. - # Save to SQL - with pytest.raises(ValueError, match=r"missing index items"): - storage.store_df(to_store.set_index('col1'), meta) + """ + uri = tmp_path / "test_upsert_delete.db" + storage = SQLiteFeatureStorage(uri=uri, single_output=True) + # Metadata to store + meta = {"element": "test", "version": "0.0.1"} + # Save to database + storage.store_df(df1, meta) + # Store metadata + table_name = storage.store_metadata(meta) + # Read stored table + c_df1 = _read_sql( + table_name=table_name, uri=uri, index_col=["element", "pk2"] + ) + # Check if dataframes are equal + assert_frame_equal(df1, c_df1) + # Save to database + storage.store_df(df2, meta) + # Read stored table + c_dfupdate = _read_sql(table_name, uri=uri, index_col=["element", "pk2"]) + # Check if dataframes are equal + assert_frame_equal(c_dfupdate, df_update) - to_store = df1.reset_index().set_index(['element', 'pk2', 'col1']) - with pytest.raises(ValueError, match=r"extra items"): - storage.store_df(to_store, meta) - idx = element_to_index(meta, n_rows=len(to_store)) - to_store = to_store.set_index(idx) +def test_upsert_invalid_option(tmp_path: Path) -> None: + """Test dataframe store wtih invalid option for upsert. + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + uri = tmp_path / "test_upsert_invalid.db" + with pytest.raises(ValueError): + SQLiteFeatureStorage(uri=uri, single_output=True, upsert="wrong") + + +# TODO: can the tests be separated? +def test_store_df_and_read_df(tmp_path: Path) -> None: + """Test storing dataframe and reading of stored table into dataframe. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + uri = tmp_path / "test_store_df_and_read_df.db" + storage = SQLiteFeatureStorage( + uri=uri, single_output=True, upsert="ignore" + ) + # Metadata to store + meta = { + "element": "test", + "version": "0.0.1", + "marker": {"name": "fcname"}, + } + # Columns to store + to_store = df1[["col1", "col2"]] + # Check for error while storing + with pytest.raises(ValueError, match=r"missing index items"): + storage.store_df(to_store.set_index("col1"), meta) + # Set index + to_store = df1.reset_index().set_index(["element", "pk2", "col1"]) + # Check for error while storing + with pytest.raises(ValueError, match=r"extra items"): storage.store_df(to_store, meta) - - # Test the internals - table_name = storage.store_metadata(meta) - - features = storage.list_features() - assert len(features) == 1 - assert table_name.replace('meta_', '') in features - - with pytest.raises(ValueError, match='not found'): - storage.read_df('wrong_md5') - - with pytest.raises(ValueError, match='least one'): - storage.read_df() - - with pytest.raises(ValueError, match='Only one'): - storage.read_df('wrong_md5', 'wrong_name') - - feature_md5 = list(features.keys())[0] - assert 'fcname' == features[feature_md5]['name'] - read_df1 = storage.read_df(feature_md5=feature_md5) - read_df2 = storage.read_df(feature_name='fcname') - assert_frame_equal(read_df1, read_df2) - assert_frame_equal(read_df1, to_store) + # Convert element to index + idx = element_to_index(meta, n_rows=len(to_store)) + # Set index + to_store = to_store.set_index(idx) + # Store dataframe + storage.store_df(to_store, meta) + # Store metadata + table_name = storage.store_metadata(meta) + # List stored features + features = storage.list_features() + # Check correct usage + assert len(features) == 1 + assert table_name.replace("meta_", "") in features + # Check for missing feature + with pytest.raises(ValueError, match="not found"): + storage.read_df("wrong_md5") + # Check for missing feature to fetch + with pytest.raises(ValueError, match="least one"): + storage.read_df() + # Check for multiple features to fetch + with pytest.raises(ValueError, match="Only one"): + storage.read_df("wrong_md5", "wrong_name") + # Get MD5 hash of features + feature_md5 = list(features.keys())[0] + # Check for key + assert "fcname" == features[feature_md5]["name"] + # Read into dataframes + read_df1 = storage.read_df(feature_md5=feature_md5) + read_df2 = storage.read_df(feature_name="fcname") + # Check if dataframes are equal + assert_frame_equal(read_df1, read_df2) + assert_frame_equal(read_df1, to_store) -def test_store_table(): - """Test table store.""" - meta = {'element': 'test', 'version': '0.0.1', 'marker': {'name': 'fc'}} - with tempfile.TemporaryDirectory() as _tmpdir: - uri = f'{_tmpdir}/test.db' - storage = SQLiteFeatureStorage(uri=uri, single_output=True) - data = [ +def test_store_metadata(tmp_path: Path) -> None: + """Test metadata store. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + uri = tmp_path / "test_metadata_store.db" + # Single storage, must be the uri + storage = SQLiteFeatureStorage( + uri=uri, single_output=True, upsert="ignore" + ) + # Metadata to store + meta = {"element": "test", "version": "0.0.1"} + # Store metadata + table_name = storage.store_metadata(meta) + assert table_name.startswith("meta_") + + +def test_store_table(tmp_path: Path) -> None: + """Test table store. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + uri = tmp_path / "test_store_table.db" + storage = SQLiteFeatureStorage(uri=uri, single_output=True) + # Metadata to store + meta = {"element": "test", "version": "0.0.1", "marker": {"name": "fc"}} + # Data to store + data = [ + [1, 10], + [2, 20], + [3, 30], + [4, 40], + [5, 50], + ] + # Convert element to index + idx = element_to_index(meta, n_rows=5, rows_col_name="scan") + # Create dataframe + df = pd.DataFrame(data, columns=["f1", "f2"], index=idx) + # Store table + storage.store_table(data, meta, columns=["f1", "f2"], rows_col_name="scan") + # Store metadata + table_name = storage.store_metadata(meta) + # Read stored table + c_df = _read_sql( + table_name=table_name, uri=uri, index_col=["element", "scan"] + ) + # Check if dataframes are equal + assert_frame_equal(df, c_df) + + +def test_store_table_check_warning(tmp_path: Path) -> None: + """Test table store and check warning. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + uri = tmp_path / "test_store_table_check_warning.db" + storage = SQLiteFeatureStorage(uri=uri, single_output=True) + # Metadata to store + meta = {"element": "test", "version": "0.0.1", "marker": {"name": "fc"}} + # Data to store + data = [[1, 10], [2, 20], [3, 300], [4, 40], [5, 50], [6, 600]] + # Convert element to index + idx = element_to_index(meta, n_rows=6, rows_col_name="scan") + # Create dataframe + df = pd.DataFrame(data, columns=["f1", "f2"], index=idx) + # Check warning + with pytest.warns(RuntimeWarning, match=r"Some rows"): + storage.store_table( + data, meta, columns=["f1", "f2"], rows_col_name="scan" + ) + # Store metadata + table_name = storage.store_metadata(meta) + # Read stored table + c_df = _read_sql( + table_name=table_name, uri=uri, index_col=["element", "scan"] + ) + # Check if dataframes are equal + assert_frame_equal(df, c_df) + + +# TODO: can the test be parametrized? +def test_store_multiple_output(tmp_path: Path): + """Test storing using single_output=False. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + uri = tmp_path / "test_store_multiple_output.db" + storage = SQLiteFeatureStorage(uri=uri, single_output=False) + # Metadata to store + meta1 = { + "element": {"subject": "test-01", "session": "ses-01"}, + "version": "0.0.1", + "marker": {"name": "fc"}, + } + meta2 = { + "element": {"subject": "test-02", "session": "ses-01"}, + "version": "0.0.1", + "marker": {"name": "fc"}, + } + meta3 = { + "element": {"subject": "test-01", "session": "ses-02"}, + "version": "0.0.1", + "marker": {"name": "fc"}, + } + # Data to store + data1 = np.array( + [ [1, 10], [2, 20], [3, 30], [4, 40], [5, 50], ] + ) + data2 = data1 * 10 + data3 = data1 * 20 + # Process metadata for storage + hash1, _ = process_meta(meta1) + # Convert element to index + idx1 = element_to_index(meta1, n_rows=5, rows_col_name="scan") + # Create dataframe + df1 = pd.DataFrame(data1, columns=["f1", "f2"], index=idx1) + # Process metadata for storage + hash2, _ = process_meta(meta2) + # Convert element to index + idx2 = element_to_index(meta2, n_rows=5, rows_col_name="scan") + # Create dataframe + df2 = pd.DataFrame(data2, columns=["f1", "f2"], index=idx2) + # Process metadata for storage + hash3, _ = process_meta(meta3) + # Convert element to index + idx3 = element_to_index(meta3, n_rows=5, rows_col_name="scan") + # Create dataframe + df3 = pd.DataFrame(data3, columns=["f1", "f2"], index=idx3) + # Check hash equality + assert hash1 == hash2 + assert hash2 == hash3 + # Store tables + storage.store_table( + data1, meta1, columns=["f1", "f2"], rows_col_name="scan" + ) + storage.store_table( + data2, meta2, columns=["f1", "f2"], rows_col_name="scan" + ) + storage.store_table( + data3, meta3, columns=["f1", "f2"], rows_col_name="scan" + ) + # Check that URI does not exist yet + assert not uri.exists() + # Convert element to preifx + prefix1 = element_to_prefix(meta1["element"]) + prefix2 = element_to_prefix(meta2["element"]) + prefix3 = element_to_prefix(meta3["element"]) + # URIs for data storage + uri1 = uri.parent / f"{prefix1}{uri.name}" + uri2 = uri.parent / f"{prefix2}{uri.name}" + uri3 = uri.parent / f"{prefix3}{uri.name}" + # Check URIs for data storage exist + assert uri1.exists() + assert uri2.exists() + assert uri3.exists() + # Store metadata + table_name = storage.store_metadata(meta1) + # Set index columns + cols = ["subject", "session", "scan"] + # Read stored tables + cdf1 = _read_sql(table_name, uri1, index_col=cols) + cdf2 = _read_sql(table_name, uri2, index_col=cols) + cdf3 = _read_sql(table_name, uri3, index_col=cols) + # Check if dataframes are equal + assert_frame_equal(df1, cdf1) + assert_frame_equal(df2, cdf2) + assert_frame_equal(df3, cdf3) - idx = element_to_index(meta, n_rows=5, rows_col_name='scan') - df1 = pd.DataFrame(data, columns=['f1', 'f2'], index=idx) - storage.store_table( - data, meta, columns=['f1', 'f2'], rows_col_name='scan') +# TODO: can test be paramtrized? +def test_collect(tmp_path: Path) -> None: + """Test collect. - table_name = storage.store_metadata(meta) - c_df1 = _read_sql(table_name, uri=uri, index_col=['element', 'scan']) - assert_frame_equal(df1, c_df1) + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. - data2 = [ + """ + uri = tmp_path / "test_collect.db" + storage = SQLiteFeatureStorage(uri=uri) + # Metadata for storage + meta1 = { + "element": {"subject": "test-01", "session": "ses-01"}, + "version": "0.0.1", + "marker": {"name": "fc"}, + } + meta2 = { + "element": {"subject": "test-02", "session": "ses-01"}, + "version": "0.0.1", + "marker": {"name": "fc"}, + } + meta3 = { + "element": {"subject": "test-01", "session": "ses-02"}, + "version": "0.0.1", + "marker": {"name": "fc"}, + } + # Data for storage + data1 = np.array( + [ [1, 10], [2, 20], - [3, 300], + [3, 30], [4, 40], [5, 50], - [6, 600] ] - - idx = element_to_index(meta, n_rows=6, rows_col_name='scan') - df2 = pd.DataFrame(data2, columns=['f1', 'f2'], index=idx) - - with pytest.warns(RuntimeWarning, match=r"Some rows"): - storage.store_table( - data2, meta, columns=['f1', 'f2'], rows_col_name='scan') - - c_df2 = _read_sql(table_name, uri=uri, index_col=['element', 'scan']) - assert_frame_equal(df2, c_df2) - - -def test_store_multiple_output(): - """Test storing using single_output=False.""" - - meta1 = {'element': {'subject': 'test-01', 'session': 'ses-01'}, - 'version': '0.0.1', 'marker': {'name': 'fc'}} - meta2 = {'element': {'subject': 'test-02', 'session': 'ses-01'}, - 'version': '0.0.1', 'marker': {'name': 'fc'}} - meta3 = {'element': {'subject': 'test-01', 'session': 'ses-02'}, - 'version': '0.0.1', 'marker': {'name': 'fc'}} - with tempfile.TemporaryDirectory() as _tmpdir: - uri = Path(f'{_tmpdir}/test.db') - storage = SQLiteFeatureStorage(uri=uri, single_output=False) - data1 = np.array([ - [1, 10], - [2, 20], - [3, 30], - [4, 40], - [5, 50], - ]) - - data2 = data1 * 10 - data3 = data1 * 20 - - hash1, _ = process_meta(meta1) - idx1 = element_to_index(meta1, n_rows=5, rows_col_name='scan') - df1 = pd.DataFrame(data1, columns=['f1', 'f2'], index=idx1) - - hash2, _ = process_meta(meta2) - idx2 = element_to_index(meta2, n_rows=5, rows_col_name='scan') - df2 = pd.DataFrame(data2, columns=['f1', 'f2'], index=idx2) - - hash3, _ = process_meta(meta3) - idx3 = element_to_index(meta3, n_rows=5, rows_col_name='scan') - df3 = pd.DataFrame(data3, columns=['f1', 'f2'], index=idx3) - - assert hash1 == hash2 - assert hash2 == hash3 - - storage.store_table( - data1, meta1, columns=['f1', 'f2'], rows_col_name='scan') - - storage.store_table( - data2, meta2, columns=['f1', 'f2'], rows_col_name='scan') - - storage.store_table( - data3, meta3, columns=['f1', 'f2'], rows_col_name='scan') - - assert not uri.exists() - - prefix1 = element_to_prefix(meta1['element']) - prefix2 = element_to_prefix(meta2['element']) - prefix3 = element_to_prefix(meta3['element']) - - uri1 = uri.parent / f'{prefix1}{uri.name}' - uri2 = uri.parent / f'{prefix2}{uri.name}' - uri3 = uri.parent / f'{prefix3}{uri.name}' - - assert uri1.exists() - assert uri2.exists() - assert uri3.exists() - - table_name = storage.store_metadata(meta1) - - cols = ['subject', 'session', 'scan'] - - cdf1 = _read_sql(table_name, uri1, index_col=cols) - cdf2 = _read_sql(table_name, uri2, index_col=cols) - cdf3 = _read_sql(table_name, uri3, index_col=cols) - - assert_frame_equal(df1, cdf1) - assert_frame_equal(df2, cdf2) - assert_frame_equal(df3, cdf3) - - -def test_collect(): - """Test collect.""" - meta1 = {'element': {'subject': 'test-01', 'session': 'ses-01'}, - 'version': '0.0.1', 'marker': {'name': 'fc'}} - meta2 = {'element': {'subject': 'test-02', 'session': 'ses-01'}, - 'version': '0.0.1', 'marker': {'name': 'fc'}} - meta3 = {'element': {'subject': 'test-01', 'session': 'ses-02'}, - 'version': '0.0.1', 'marker': {'name': 'fc'}} - with tempfile.TemporaryDirectory() as _tmpdir: - uri = Path(f'{_tmpdir}/test.db') - storage = SQLiteFeatureStorage(uri=uri) - - data1 = np.array([ - [1, 10], - [2, 20], - [3, 30], - [4, 40], - [5, 50], - ]) - - data2 = data1 * 10 - data3 = data1 * 20 - storage.store_table( - data1, meta1, columns=['f1', 'f2'], rows_col_name='scan') - - storage.store_table( - data2, meta2, columns=['f1', 'f2'], rows_col_name='scan') - - storage.store_table( - data3, meta3, columns=['f1', 'f2'], rows_col_name='scan') - - prefix1 = element_to_prefix(meta1['element']) - prefix2 = element_to_prefix(meta2['element']) - prefix3 = element_to_prefix(meta3['element']) - - uri1 = uri.parent / f'{prefix1}{uri.name}' - uri2 = uri.parent / f'{prefix2}{uri.name}' - uri3 = uri.parent / f'{prefix3}{uri.name}' - - assert uri1.exists() - assert uri2.exists() - assert uri3.exists() - - assert not uri.exists() - - storage.collect() - - assert uri.exists() - - cols = ['subject', 'session', 'scan'] - table_name = storage.store_metadata(meta1) - all_df = _read_sql(table_name, uri, index_col=cols) - - cdf1 = _read_sql(table_name, uri1, index_col=cols) - cdf2 = _read_sql(table_name, uri2, index_col=cols) - cdf3 = _read_sql(table_name, uri3, index_col=cols) - - all_cdf = pd.concat([cdf1, cdf2, cdf3]) - all_df.sort_index(level=cols, inplace=True) - all_cdf.sort_index(level=cols, inplace=True) - - assert_frame_equal(all_df, all_cdf) + ) + data2 = data1 * 10 + data3 = data1 * 20 + # Store tables + storage.store_table( + data1, meta1, columns=["f1", "f2"], rows_col_name="scan" + ) + storage.store_table( + data2, meta2, columns=["f1", "f2"], rows_col_name="scan" + ) + storage.store_table( + data3, meta3, columns=["f1", "f2"], rows_col_name="scan" + ) + # Convert element to prefix + prefix1 = element_to_prefix(meta1["element"]) + prefix2 = element_to_prefix(meta2["element"]) + prefix3 = element_to_prefix(meta3["element"]) + # URIs for data storage + uri1 = uri.parent / f"{prefix1}{uri.name}" + uri2 = uri.parent / f"{prefix2}{uri.name}" + uri3 = uri.parent / f"{prefix3}{uri.name}" + # Check URIs for data storage exist + assert uri1.exists() + assert uri2.exists() + assert uri3.exists() + # Check that URI does not exist yet + assert not uri.exists() + # Collect data + storage.collect() + # Check that URI exists now + assert uri.exists() + # Set index columns + cols = ["subject", "session", "scan"] + # Store metadata + table_name = storage.store_metadata(meta1) + # Read stored tables + all_df = _read_sql(table_name, uri, index_col=cols) + cdf1 = _read_sql(table_name, uri1, index_col=cols) + cdf2 = _read_sql(table_name, uri2, index_col=cols) + cdf3 = _read_sql(table_name, uri3, index_col=cols) + # Operate on retrieved tables + all_cdf = pd.concat([cdf1, cdf2, cdf3]) + all_df.sort_index(level=cols, inplace=True) + all_cdf.sort_index(level=cols, inplace=True) + # Check if dataframes are equal + assert_frame_equal(all_df, all_cdf) -- 2.52.0 From 5df610e4fe6982ba7540014bc1ee8e08bfc97c4d Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 12 Aug 2022 17:20:35 +0200 Subject: [PATCH 179/287] fix: correct import in base.py --- junifer/storage/base.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/junifer/storage/base.py b/junifer/storage/base.py index 4b2bc94c7..39e43e1e9 100644 --- a/junifer/storage/base.py +++ b/junifer/storage/base.py @@ -10,7 +10,7 @@ from typing import Dict, Iterable, List, Optional, Union import pandas as pd -from .. import __version__ +from .._version import __version__ from ..utils import raise_error -- 2.52.0 From 17f3db451c6dce75a952fc04dd850006668f6b2c Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 15 Aug 2022 15:36:58 +0200 Subject: [PATCH 180/287] refactor: prune, add docstrings and type annotations in markers/base.py --- junifer/markers/base.py | 219 +++++++++++++++++++++------------------- 1 file changed, 117 insertions(+), 102 deletions(-) diff --git a/junifer/markers/base.py b/junifer/markers/base.py index 3357123af..1fef738bf 100644 --- a/junifer/markers/base.py +++ b/junifer/markers/base.py @@ -1,139 +1,154 @@ -"""Provide base class and mixin class for markers.""" +"""Provide base class for markers.""" # Authors: Federico Raimondo +# Synchon Mandal # License: AGPL -from ..utils import logger +from typing import Dict, List, Optional - -class PipelineStepMixin(): - """Mixin class for pipeline.""" - - def get_meta(self): - """Get metadata.""" - t_meta = {} - t_meta['class'] = self.__class__.__name__ - for k, v in vars(self).items(): - if not k.startswith('_'): - t_meta[k] = v - return t_meta - - def validate_input(self, input): - """Validate the input to the pipeline step. - - Parameters - ---------- - input : list[str] - The input to the pipeline step. The list must contain the - available Junifer Data dictionary keys. - - Raises - ------ - ValueError: - If the input does not have the required data. - """ - raise NotImplementedError('validate_input not implemented') - - def get_output_kind(self, input): - """Get the kind of the pipeline step. - - Parameters - ---------- - input : list[str] - The input to the pipeline step. The list must contain the - available Junifer Data dictionary keys. - - Returns - ------- - output : list[str] - The updated list of available Junifer Data dictionary keys after - the pipeline step. - """ - raise NotImplementedError('get_output_kind not implemented') - - def validate(self, input): - """Validate the the pipeline step. - - Parameters - ---------- - input : Junifer Data dictionary - The input to the pipeline step. - - Returns - ------- - output : Junifer Data dictionary - The output of the pipeline step. - - Raises - ------ - ValueError: - If the input does not have the required data. - """ - self.validate_input(input) - return self.get_output_kind(input) - - def fit_transform(self, input): - """Fit and transform.""" - raise NotImplementedError('fit_transform not implemented') +from ..utils import logger, raise_error +from .pipeline_mixin import PipelineStepMixin class BaseMarker(PipelineStepMixin): - """Base class for all markers.""" + """Base class for all markers. - def __init__(self, on, name=None): + Parameters + ---------- + on : list of str + The kind of data to work on. + name : str, optional + The name of the marker (default None). + + """ + + def __init__(self, on: List, name: Optional[str] = None) -> None: """Initialize the class.""" if not isinstance(on, list): on = [on] self._valid_inputs = on self.name = self.__class__.__name__ if name is None else name - def get_meta(self, kind): - """Get metadata.""" + def get_meta(self, kind: str) -> Dict: + """Get metadata. + + Parameters + ---------- + kind : str + The kind of pipeline step. + + Returns + ------- + dict + The metadata as a dictionary. + + """ s_meta = super().get_meta() - # same marker can be fit into different kinds, so the name + # same marker can be "fit"ted into different kinds, so the name # is created from the kind and the name of the marker - s_meta['name'] = f'{kind}_{self.name}' - s_meta['kind'] = kind - return dict(marker=s_meta) + s_meta["name"] = f"{kind}_{self.name}" + s_meta["kind"] = kind + return {"marker": s_meta} - def validate_input(self, input): - """Validate input.""" + def validate_input(self, input: List[str]) -> None: + """Validate input. + + Parameters + ---------- + input : list of str + The input to the pipeline step. The list must contain the + available Junifer Data dictionary keys. + + Raises + ------ + ValueError + If the input does not have the required data. + + """ if not any(x in input for x in self._valid_inputs): - raise ValueError( - 'Input does not have the required data.' - f'\t Input: {input}' - f'\t Required (any of): {self._valid_inputs}') + raise_error( + "Input does not have the required data." + f"\t Input: {input}" + f"\t Required (any of): {self._valid_inputs}" + ) - def get_output_kind(self, input): - """Get output kind.""" - return None + def get_output_kind(self, input: List[str]) -> List[str]: + """Get output kind. - def compute(self, input): - """Compute.""" - raise NotImplementedError('compute not implemented') + Parameters + ---------- + input : list of str + The input to the pipeline step. The list must contain the + available Junifer Data dictionary keys. - def store(self, input, out, storage): - """Store.""" - raise NotImplementedError('store not implemented') + Returns + ------- + list of str + The updated list of available Junifer Data dictionary keys after + the pipeline step. - def fit_transform(self, input, storage=None): - """Fit and transform.""" + """ + pass + + def compute(self, input: List[str]) -> Dict: + """Compute. + + Parameters + ---------- + input : list of str + The input to the pipeline step. The list must contain the + available Junifer Data dictionary keys. + + Returns + ------- + dict + The computed result as dictionary. + + """ + raise_error(msg="compute() not implemented", klass=NotImplementedError) + + # TODO: complete type annotations + def store(self, input: List[str], out: Dict, storage) -> None: + """Store. + + Parameters + ---------- + input + out + + """ + raise_error(msg="store() not implemented", klass=NotImplementedError) + + # TODO: complete type annotations + def fit_transform(self, input: List[str], storage=None) -> Dict: + """Fit and transform. + + Parameters + ---------- + input + storage + + Returns + ------- + dict + + """ out = {} - meta = input.get('meta', {}) + meta = input.get("meta", {}) for kind in self._valid_inputs: if kind in input.keys(): - logger.info(f'Computing {kind}') + logger.info(f"Computing {kind}") t_input = input[kind] t_meta = meta.copy() - t_meta.update(t_input.get('meta', {})) + t_meta.update(t_input.get("meta", {})) t_meta.update(self.get_meta(kind)) t_out = self.compute(t_input) t_out.update(meta=t_meta) if storage is not None: - logger.info(f'Storing in {storage}') + logger.info(f"Storing in {storage}") self.store(kind, t_out, storage) else: - logger.info('No storage specified, returning dictionary') + logger.info("No storage specified, returning dictionary") out[kind] = t_out return out -- 2.52.0 From 30c409610bd5f6a66fbf91cc80b9313c696906d2 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 15 Aug 2022 15:37:38 +0200 Subject: [PATCH 181/287] refactor: add new module for PipelineStepMixin --- junifer/markers/pipeline_mixin.py | 107 ++++++++++++++++++++++++++++++ 1 file changed, 107 insertions(+) create mode 100644 junifer/markers/pipeline_mixin.py diff --git a/junifer/markers/pipeline_mixin.py b/junifer/markers/pipeline_mixin.py new file mode 100644 index 000000000..22cd52da4 --- /dev/null +++ b/junifer/markers/pipeline_mixin.py @@ -0,0 +1,107 @@ +"""Provide mixin class for pipeline step.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from typing import Dict, List + +from ..utils import raise_error + + +class PipelineStepMixin: + """Mixin class for pipeline.""" + + def get_meta(self) -> Dict: + """Get metadata. + + Returns + ------- + dict + The metadata as a dictionary. + + """ + t_meta = {} + t_meta["class"] = self.__class__.__name__ + for k, v in vars(self).items(): + if not k.startswith("_"): + t_meta[k] = v + return t_meta + + def validate_input(self, input: List[str]) -> None: + """Validate the input to the pipeline step. + + Parameters + ---------- + input : list of str + The input to the pipeline step. The list must contain the + available Junifer Data dictionary keys. + + Raises + ------ + ValueError + If the input does not have the required data. + + """ + raise_error( + msg="Concrete classes need to implement validate_input().", + klass=NotImplementedError, + ) + + def get_output_kind(self, input: List[str]) -> List[str]: + """Get the kind of the pipeline step. + + Parameters + ---------- + input : list of str + The input to the pipeline step. The list must contain the + available Junifer Data dictionary keys. + + Returns + ------- + list of str + The updated list of available Junifer Data dictionary keys after + the pipeline step. + + """ + raise_error( + msg="Concrete classes need to implement get_output_kind().", + klass=NotImplementedError, + ) + + def validate(self, input: List[str]) -> List[str]: + """Validate the the pipeline step. + + Parameters + ---------- + input : list of str + The input to the pipeline step. + + Returns + ------- + list of str + The output of the pipeline step. + + Raises + ------ + ValueError + If the input does not have the required data. + + """ + self.validate_input(input=input) + return self.get_output_kind(input=input) + + def fit_transform(self, input: List[str]) -> None: + """Fit and transform. + + Parameters + ---------- + input : list of str + The input to the pipeline step. The list must contain the + available Junifer Data dictionary keys. + + """ + raise_error( + msg="Concrete classes need to implement fit_transform().", + klass=NotImplementedError, + ) -- 2.52.0 From 1518eb9b04eb2001f14043f499688cd88e9bd044 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 15 Aug 2022 15:38:06 +0200 Subject: [PATCH 182/287] refactor: add docstrings and type annotations for markers/collection.py --- junifer/markers/collection.py | 83 +++++++++++++++++++++-------------- 1 file changed, 51 insertions(+), 32 deletions(-) diff --git a/junifer/markers/collection.py b/junifer/markers/collection.py index 84c812cfd..f658e1b9f 100644 --- a/junifer/markers/collection.py +++ b/junifer/markers/collection.py @@ -1,86 +1,105 @@ """Provide class for marker collection.""" # Authors: Federico Raimondo +# Synchon Mandal # License: AGPL -from ..utils import logger -from ..datareader import DefaultDataReader - from collections import Counter +from typing import Dict, Optional + +from ..datareader import DefaultDataReader +from ..utils import logger -class MarkerCollection(): - """Class for marker collection.""" +class MarkerCollection: + """Class for marker collection. + + Parameters + ---------- + markers + datareader + preprocessing + storage + + """ def __init__( - self, markers, datareader=None, preprocessing=None, storage=None + self, markers, datareader=None, preprocessing=None, storage=None ): """Initialize the class.""" + # Check that the markers have different names + marker_names = [m.name for m in markers] + if len(set(marker_names)) != len(marker_names): + counts = Counter(marker_names) + raise ValueError( + "Markers must have different names. " + f"Current names are: {counts}" + ) + self._markers = markers if datareader is None: datareader = DefaultDataReader() self._datareader = datareader self._preprocessing = preprocessing - self._markers = markers self._storage = storage - # Check that the markers have different names - marker_names = [m.name for m in self._markers] - if len(set(marker_names)) != len(marker_names): - counts = Counter(marker_names) - raise ValueError( - 'Markers must have different names. ' - f'Current names are: {counts}') - - def fit(self, input): + # TODO: complete type annotations + def fit(self, input) -> Optional[Dict]: """Fit the pipeline. Parameters ---------- - input : Junifer Data dictionary (input) + input The input data to fit the pipeline on. Should be the output of indexing the DataGrabber with one element. Returns ------- - output : dict[str -> object] | None + output : dict or None The output of the pipeline. Each key represents a marker name and the values are the computer marker values. If the pipeline has a storage configured, then the output will be None. + """ - logger.info('Fitting pipeline') + logger.info("Fitting pipeline") data = self._datareader.fit_transform(input) if self._preprocessing is not None: - logger.info('Preprocessing data') + logger.info("Preprocessing data") data = self._preprocessing.fit_transform(data) out = {} for marker in self._markers: - logger.info(f'Fitting marker {marker.name}') + logger.info(f"Fitting marker {marker.name}") m_value = marker.fit_transform(data, storage=self._storage) if self._storage is None: out[marker.name] = m_value - logger.info('Marker collection fitting done') + logger.info("Marker collection fitting done") return None if self._storage else out - def validate(self, datagrabber): + # TODO: complete type annotations + def validate(self, datagrabber) -> None: """Validate the pipeline. Without doing any computation, check if the Marker Collection can be fit without problems. That is, the data required for each marker is present and streamed down the steps. Also, if a storage is configured, check that the storage can handle the markers output. - """ - logger.info('Validating Marker Collection') - t_data = datagrabber.get_types() - logger.info(f'DataGrabber output type: {t_data}') - logger.info('Validating Data Reader:') + Parameters + ---------- + datagrabber + + """ + logger.info("Validating Marker Collection") + t_data = datagrabber.get_types() + logger.info(f"DataGrabber output type: {t_data}") + + logger.info("Validating Data Reader:") t_data = self._datareader.validate(t_data) - logger.info(f'Data Reader output type: {t_data}') + logger.info(f"Data Reader output type: {t_data}") for marker in self._markers: - logger.info(f'Validating Marker: {marker.name}') + logger.info(f"Validating Marker: {marker.name}") m_data = marker.validate(t_data) - logger.info(f'Marker output type: {m_data}') + logger.info(f"Marker output type: {m_data}") if self._storage is not None: - logger.info(f'Validating storage for {marker.name}') + logger.info(f"Validating storage for {marker.name}") self._storage.validate(m_data) -- 2.52.0 From ed65194bd76997b505a44bd8da08af940328e67a Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 15 Aug 2022 15:38:37 +0200 Subject: [PATCH 183/287] refactor: add docstrings and type annotations in markers/parcel.py --- junifer/markers/parcel.py | 116 +++++++++++++++++++++++++++----------- 1 file changed, 82 insertions(+), 34 deletions(-) diff --git a/junifer/markers/parcel.py b/junifer/markers/parcel.py index 8b86f3009..6b217d2c5 100644 --- a/junifer/markers/parcel.py +++ b/junifer/markers/parcel.py @@ -1,67 +1,115 @@ """Provide class for parcel aggregation.""" # Authors: Federico Raimondo +# Synchon Mandal # License: AGPL + +from typing import Dict, List + import numpy as np - +from nilearn.image import math_img, resample_to_img from nilearn.maskers import NiftiMasker -from nilearn.image import resample_to_img, math_img -from .base import BaseMarker -from ..stats import get_aggfunc_by_name -from ..data import load_atlas from ..api.decorators import register_marker +from ..data import load_atlas +from ..stats import get_aggfunc_by_name from ..utils import logger +from .base import BaseMarker @register_marker class ParcelAggregation(BaseMarker): - """Class for parcel aggregation.""" + """Class for parcel aggregation. - def __init__(self, atlas, method, method_params=None, on=None, name=None): + Parameters + ---------- + atlas + method + method_params + on + name + + """ + + def __init__( + self, atlas, method, method_params=None, on=None, name=None + ) -> None: """Initialize the class.""" - if on is None: - on = ['T1w', 'BOLD', 'VBM_GM', 'VBM_WM', 'fALFF', 'GCOR', 'LCOR'] - super().__init__(on=on, name=name) self.atlas = atlas self.method = method self.method_params = {} if method_params is None else method_params + if on is None: + on = ["T1w", "BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"] + super().__init__(on=on, name=name) - def get_output_kind(self, input): - """Get output kind.""" - if input in ['VBM_GM', 'VBM_WM', 'fALFF', 'GCOR', 'LCOR']: - return 'table' - if input in ['BOLD']: - return 'timeseries' + def get_output_kind(self, input: List[str]) -> str: + """Get output kind. - def store(self, kind, out, storage): - """Store.""" - logger.debug(f'Storing {kind} in {storage}') - if kind in ['VBM_GM', 'VBM_WM', 'fALFF', 'GCOR', 'LCOR']: + Parameters + ---------- + input : list of str + The kind of data to work on. + + Returns + ------- + str + The kind of output. + + """ + if input in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]: + return "table" + if input in ["BOLD"]: + return "timeseries" + + # TODO: complete type annotations + def store(self, kind: List[str], out, storage) -> None: + """Store. + + Parameters + ---------- + kind + out + storage + + """ + logger.debug(f"Storing {kind} in {storage}") + if kind in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]: storage.store_table(**out) - if kind in ['BOLD']: + if kind in ["BOLD"]: storage.store_timeseries(**out) - def compute(self, input): - """Compute.""" - t_input = input['data'] - logger.debug(f'Parcel aggregation using {self.method}') + # TODO: complete type annotations + def compute(self, input) -> Dict: + """Compute. + + Parameters + ---------- + input + + Returns + ------- + dict + The computed result as dictionary. + + """ + t_input = input["data"] + logger.debug(f"Parcel aggregation using {self.method}") agg_func = get_aggfunc_by_name( - self.method, func_params=self.method_params) + self.method, func_params=self.method_params + ) # Get the min of the voxels sizes and use it as the resolution resolution = np.min(t_input.header.get_zooms()[:3]) - t_atlas, t_labels, _ = load_atlas( - self.atlas, resolution=resolution) + t_atlas, t_labels, _ = load_atlas(self.atlas, resolution=resolution) atlas_img_res = resample_to_img( t_atlas, t_input, - interpolation='nearest', + interpolation="nearest", ) atlas_bin = math_img( - 'img != 0', + "img != 0", img=atlas_img_res, ) - logger.debug('Masking') + logger.debug("Masking") masker = NiftiMasker(atlas_bin, target_affine=t_input.affine) # Mask the input data and the atlas @@ -70,7 +118,7 @@ class ParcelAggregation(BaseMarker): atlas_values = np.squeeze(atlas_values).astype(int) # Get the values for each parcel and apply agg function - logger.debug('Computing ROI means') + logger.debug("Computing ROI means") atlas_roi_vals = sorted(np.unique(atlas_values)) out_labels = [] out_values = [] @@ -83,7 +131,7 @@ class ParcelAggregation(BaseMarker): out_labels.append(t_labels[t_v - 1]) out_values = np.array(out_values).T - out = dict(data=out_values, columns=out_labels) + out = {"data": out_values, "columns": out_labels} if out_values.shape[0] > 1: - out['row_names'] = 'scan' # type: ignore + out["row_names"] = "scan" return out -- 2.52.0 From 05dceab4811865d59d3b28b1c120eeb7c9569b43 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 15 Aug 2022 15:39:28 +0200 Subject: [PATCH 184/287] refactor: rename, prune, add docstrings and type annotations in test_base.py --- junifer/markers/tests/test_base.py | 78 +++++++++++++++++++++++ junifer/markers/tests/test_base_marker.py | 64 ------------------- 2 files changed, 78 insertions(+), 64 deletions(-) create mode 100644 junifer/markers/tests/test_base.py delete mode 100644 junifer/markers/tests/test_base_marker.py diff --git a/junifer/markers/tests/test_base.py b/junifer/markers/tests/test_base.py new file mode 100644 index 000000000..f91ac3eb0 --- /dev/null +++ b/junifer/markers/tests/test_base.py @@ -0,0 +1,78 @@ +"""Provide tests for base marker.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from typing import List, Optional + +import pytest + +from junifer.markers.base import BaseMarker + + +@pytest.mark.parametrize( + "on, name, kind, expected_class, expected_name", + [ + (["bold", "dwi"], None, "bold", "BaseMarker", "bold_BaseMarker"), + (["bold", "dwi"], "mymarker", "dwi", "BaseMarker", "dwi_mymarker"), + ], +) +def test_base_marker_meta( + on: List[str], + name: Optional[str], + kind: str, + expected_class: str, + expected_name: str, +) -> None: + """Test metadata for BaseMarker. + + Parameters + ---------- + on : list of str + The parametrized kind of data to work on. + name : str or None + The parametrized name of the marker. + kind : str + The parametrized kind of data to get metadata for. + expected_class : str + The paramtrized expected class of the marker. + expected_name : str + The parametrized expected name of the marker. + + """ + base = BaseMarker(on=on, name=name) + t_meta = base.get_meta(kind=kind) + assert t_meta["marker"]["class"] == expected_class + assert t_meta["marker"]["name"] == expected_name + + +def test_BaseMarker() -> None: + """Test base class.""" + base = BaseMarker(on=["bold", "dwi"], name="mymarker") + input_ = {"bold": {"path": "test"}, "t2": {"path": "test"}} + base.validate_input(input_) + + wrong_input = {"t2": {"path": "test"}} + with pytest.raises(ValueError): + base.validate_input(wrong_input) + + output = base.get_output_kind(input_) + assert output is None + + with pytest.raises(NotImplementedError): + base.fit_transform(input_) + + base.compute = lambda x: {"data": 1} + + out = base.fit_transform(input_) + assert out["bold"]["data"] == 1 + assert out["bold"]["meta"]["marker"]["name"] == "bold_mymarker" + assert out["bold"]["meta"]["marker"]["class"] == "BaseMarker" + + base2 = BaseMarker(on="bold", name="mymarker") + base2.compute = lambda x: {"data": 1} + out2 = base2.fit_transform(input_) + assert out2["bold"]["data"] == 1 + assert out2["bold"]["meta"]["marker"]["name"] == "bold_mymarker" + assert out2["bold"]["meta"]["marker"]["class"] == "BaseMarker" diff --git a/junifer/markers/tests/test_base_marker.py b/junifer/markers/tests/test_base_marker.py deleted file mode 100644 index ccfdd039a..000000000 --- a/junifer/markers/tests/test_base_marker.py +++ /dev/null @@ -1,64 +0,0 @@ -"""Provide tests for base marker.""" - -import pytest -from junifer.markers.base import BaseMarker, PipelineStepMixin - - -def test_PipelineStepMixin(): - """Test PipelineStepMixin.""" - mixin = PipelineStepMixin() - with pytest.raises(NotImplementedError): - mixin.validate_input(None) - with pytest.raises(NotImplementedError): - mixin.get_output_kind(None) - with pytest.raises(NotImplementedError): - mixin.fit_transform(None) - - -def test_meta(): - """Test metadata.""" - pipemixin = PipelineStepMixin() - t_meta = pipemixin.get_meta() - assert t_meta['class'] == 'PipelineStepMixin' - - base = BaseMarker(on=['bold', 'dwi']) - - t_meta = base.get_meta('bold') - assert t_meta['marker']['class'] == 'BaseMarker' - assert t_meta['marker']['name'] == 'bold_BaseMarker' - - base = BaseMarker(on=['bold', 'dwi'], name='mymarker') - - t_meta = base.get_meta('dwi') - assert t_meta['marker']['name'] == 'dwi_mymarker' - - -def test_BaseMarker(): - """Test base class.""" - base = BaseMarker(on=['bold', 'dwi'], name='mymarker') - input = {'bold': {'path': 'test'}, 't2': {'path': 'test'}} - base.validate_input(input) - - wrong_input = {'t2': {'path': 'test'}} - with pytest.raises(ValueError): - base.validate_input(wrong_input) - - output = base.get_output_kind(input) - assert output is None - - with pytest.raises(NotImplementedError): - base.fit_transform(input) - - base.compute = lambda x: dict(data=1) # type: ignore - - out = base.fit_transform(input) - assert out['bold']['data'] == 1 - assert out['bold']['meta']['marker']['name'] == 'bold_mymarker' - assert out['bold']['meta']['marker']['class'] == 'BaseMarker' - - base2 = BaseMarker(on='bold', name='mymarker') - base2.compute = lambda x: dict(data=1) # type: ignore - out2 = base2.fit_transform(input) - assert out2['bold']['data'] == 1 - assert out2['bold']['meta']['marker']['name'] == 'bold_mymarker' - assert out2['bold']['meta']['marker']['class'] == 'BaseMarker' -- 2.52.0 From abfb2774b773678cf9f8feaf5c4fa29b83b7bc12 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 15 Aug 2022 15:40:08 +0200 Subject: [PATCH 185/287] refactor: create new module for PipelineStepMixin unit tests --- junifer/markers/tests/test_pipeline_mixin.py | 27 ++++++++++++++++++++ 1 file changed, 27 insertions(+) create mode 100644 junifer/markers/tests/test_pipeline_mixin.py diff --git a/junifer/markers/tests/test_pipeline_mixin.py b/junifer/markers/tests/test_pipeline_mixin.py new file mode 100644 index 000000000..ace380e9a --- /dev/null +++ b/junifer/markers/tests/test_pipeline_mixin.py @@ -0,0 +1,27 @@ +"""Provide tests for pipeline mixin.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +import pytest + +from junifer.markers.pipeline_mixin import PipelineStepMixin + + +def test_PipelineStepMixin() -> None: + """Test PipelineStepMixin.""" + mixin = PipelineStepMixin() + with pytest.raises(NotImplementedError): + mixin.validate_input(None) + with pytest.raises(NotImplementedError): + mixin.get_output_kind(None) + with pytest.raises(NotImplementedError): + mixin.fit_transform(None) + + +def test_pipeline_step_mixin_meta(): + """Test metadata for PipelineStepMixin.""" + pipemixin = PipelineStepMixin() + t_meta = pipemixin.get_meta() + assert t_meta["class"] == "PipelineStepMixin" -- 2.52.0 From 6ffb523495cc38b131a5113743904ac2cbfb6ea3 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 15 Aug 2022 15:40:54 +0200 Subject: [PATCH 186/287] refactor: add docstrings and type annotations in test_collection.py --- junifer/markers/tests/test_collection.py | 172 +++++++++++++---------- 1 file changed, 95 insertions(+), 77 deletions(-) diff --git a/junifer/markers/tests/test_collection.py b/junifer/markers/tests/test_collection.py index f2b185b83..4645ab7ae 100644 --- a/junifer/markers/tests/test_collection.py +++ b/junifer/markers/tests/test_collection.py @@ -1,43 +1,48 @@ """Provide tests for marker collection.""" # Authors: Federico Raimondo +# Synchon Mandal # License: AGPL import pytest -import tempfile from numpy.testing import assert_array_equal + from junifer.datareader.default import DefaultDataReader from junifer.markers import MarkerCollection, ParcelAggregation from junifer.markers.base import PipelineStepMixin -from junifer.testing.datagrabbers import OasisVBMTestingDatagrabber from junifer.storage import SQLiteFeatureStorage +from junifer.testing.datagrabbers import OasisVBMTestingDatagrabber -def test_MarkerCollection(): - """Test MarkerCollection.""" +def test_marker_collection_incorrect_markers() -> None: + """Test incorrect markers for MarkerCollection.""" wrong_markers = [ ParcelAggregation( - atlas='Schaefer100x7', method='mean', - name='gmd_schaefer100x7_mean'), + atlas="Schaefer100x7", method="mean", name="gmd_schaefer100x7_mean" + ), ParcelAggregation( - atlas='Schaefer100x7', method='mean', - name='gmd_schaefer100x7_mean'), + atlas="Schaefer100x7", method="mean", name="gmd_schaefer100x7_mean" + ), ] - with pytest.raises(ValueError, match=r"must have different names"): MarkerCollection(wrong_markers) + +def test_marker_collection(): + """Test MarkerCollection.""" markers = [ ParcelAggregation( - atlas='Schaefer100x7', method='mean', - name='gmd_schaefer100x7_mean'), + atlas="Schaefer100x7", method="mean", name="gmd_schaefer100x7_mean" + ), ParcelAggregation( - atlas='Schaefer100x7', method='std', - name='gmd_schaefer100x7_std'), + atlas="Schaefer100x7", method="std", name="gmd_schaefer100x7_std" + ), ParcelAggregation( - atlas='Schaefer100x7', method='trim_mean', - method_params={'proportiontocut': 0.1}, - name='gmd_schaefer100x7_trim_mean90') + atlas="Schaefer100x7", + method="trim_mean", + method_params={"proportiontocut": 0.1}, + name="gmd_schaefer100x7_trim_mean90", + ), ] mc = MarkerCollection(markers=markers) assert mc._markers == markers @@ -45,26 +50,27 @@ def test_MarkerCollection(): assert mc._storage is None assert isinstance(mc._datareader, DefaultDataReader) + # Create testing datagrabber dg = OasisVBMTestingDatagrabber() mc.validate(dg) with dg: - input = dg['sub-01'] + input = dg["sub-01"] out = mc.fit(input) assert out is not None assert isinstance(out, dict) assert len(out) == 3 - assert 'gmd_schaefer100x7_mean' in out - assert 'gmd_schaefer100x7_std' in out - assert 'gmd_schaefer100x7_trim_mean90' in out + assert "gmd_schaefer100x7_mean" in out + assert "gmd_schaefer100x7_std" in out + assert "gmd_schaefer100x7_trim_mean90" in out for t_marker in markers: t_name = t_marker.name - assert 'VBM_GM' in out[t_name] - t_vbm = out[t_name]['VBM_GM'] - assert 'data' in t_vbm - assert 'columns' in t_vbm - assert 'meta' in t_vbm + assert "VBM_GM" in out[t_name] + t_vbm = out[t_name]["VBM_GM"] + assert "data" in t_vbm + assert "columns" in t_vbm + assert "meta" in t_vbm # Test preprocessing class BypassPreprocessing(PipelineStepMixin): @@ -72,74 +78,86 @@ def test_MarkerCollection(): return input mc2 = MarkerCollection( - markers=markers, preprocessing=BypassPreprocessing(), - datareader=DefaultDataReader()) + markers=markers, + preprocessing=BypassPreprocessing(), + datareader=DefaultDataReader(), + ) assert isinstance(mc2._datareader, DefaultDataReader) with dg: - input = dg['sub-01'] + input = dg["sub-01"] out2 = mc2.fit(input) for t_marker in markers: t_name = t_marker.name - assert_array_equal(out[t_name]['VBM_GM']['data'], - out2[t_name]['VBM_GM']['data']) # type: ignore + assert_array_equal( + out[t_name]["VBM_GM"]["data"], out2[t_name]["VBM_GM"]["data"] + ) # type: ignore -def test_MarkerCollection_storage(): - """Test marker collection with storage.""" +def test_MarkerCollection_storage(tmp_path) -> None: + """Test marker collection with storage. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ markers = [ ParcelAggregation( - atlas='Schaefer100x7', method='mean', - name='gmd_schaefer100x7_mean'), + atlas="Schaefer100x7", method="mean", name="gmd_schaefer100x7_mean" + ), ParcelAggregation( - atlas='Schaefer100x7', method='std', - name='gmd_schaefer100x7_std'), + atlas="Schaefer100x7", method="std", name="gmd_schaefer100x7_std" + ), ParcelAggregation( - atlas='Schaefer100x7', method='trim_mean', - method_params={'proportiontocut': 0.1}, - name='gmd_schaefer100x7_trim_mean90') + atlas="Schaefer100x7", + method="trim_mean", + method_params={"proportiontocut": 0.1}, + name="gmd_schaefer100x7_trim_mean90", + ), ] # Test storage dg = OasisVBMTestingDatagrabber() - with tempfile.TemporaryDirectory() as tmpdir: - uri = f'{tmpdir}/test.db' - storage = SQLiteFeatureStorage(uri=uri, single_output=True) - mc = MarkerCollection( - markers=markers, storage=storage, datareader=DefaultDataReader()) - mc.validate(dg) - assert mc._storage.uri == storage.uri - with dg: - input = dg['sub-01'] - out = mc.fit(input) - assert out is None - mc2 = MarkerCollection( - markers=markers, datareader=DefaultDataReader()) - mc2.validate(dg) - assert mc2._storage is None + uri = tmp_path / "test_marker_collection_storage.db" + storage = SQLiteFeatureStorage(uri=uri, single_output=True) + mc = MarkerCollection( + markers=markers, storage=storage, datareader=DefaultDataReader() + ) + mc.validate(dg) + assert mc._storage.uri == storage.uri + with dg: + input = dg["sub-01"] + out = mc.fit(input) + assert out is None - with dg: - input = dg['sub-01'] - out = mc2.fit(input) + mc2 = MarkerCollection(markers=markers, datareader=DefaultDataReader()) + mc2.validate(dg) + assert mc2._storage is None - features = storage.list_features() - assert len(features) == 3 - feature_md5 = list(features.keys())[0] - t_feature = storage.read_df(feature_md5=feature_md5) - fname = 'gmd_schaefer100x7_mean' - t_data = out[fname]['VBM_GM']['data'] # type: ignore - cols = out[fname]['VBM_GM']['columns'] # type: ignore - assert_array_equal(t_feature[cols].values, t_data) # type: ignore + with dg: + input = dg["sub-01"] + out = mc2.fit(input) - feature_md5 = list(features.keys())[1] - t_feature = storage.read_df(feature_md5=feature_md5) - fname = 'gmd_schaefer100x7_std' - t_data = out[fname]['VBM_GM']['data'] # type: ignore - cols = out[fname]['VBM_GM']['columns'] # type: ignore - assert_array_equal(t_feature[cols].values, t_data) # type: ignore + features = storage.list_features() + assert len(features) == 3 + feature_md5 = list(features.keys())[0] + t_feature = storage.read_df(feature_md5=feature_md5) + fname = "gmd_schaefer100x7_mean" + t_data = out[fname]["VBM_GM"]["data"] # type: ignore + cols = out[fname]["VBM_GM"]["columns"] # type: ignore + assert_array_equal(t_feature[cols].values, t_data) # type: ignore - feature_md5 = list(features.keys())[2] - t_feature = storage.read_df(feature_md5=feature_md5) - fname = 'gmd_schaefer100x7_trim_mean90' - t_data = out[fname]['VBM_GM']['data'] # type: ignore - cols = out[fname]['VBM_GM']['columns'] # type: ignore - assert_array_equal(t_feature[cols].values, t_data) # type: ignore + feature_md5 = list(features.keys())[1] + t_feature = storage.read_df(feature_md5=feature_md5) + fname = "gmd_schaefer100x7_std" + t_data = out[fname]["VBM_GM"]["data"] # type: ignore + cols = out[fname]["VBM_GM"]["columns"] # type: ignore + assert_array_equal(t_feature[cols].values, t_data) # type: ignore + + feature_md5 = list(features.keys())[2] + t_feature = storage.read_df(feature_md5=feature_md5) + fname = "gmd_schaefer100x7_trim_mean90" + t_data = out[fname]["VBM_GM"]["data"] # type: ignore + cols = out[fname]["VBM_GM"]["columns"] # type: ignore + assert_array_equal(t_feature[cols].values, t_data) # type: ignore -- 2.52.0 From c3bc3c2bb2a93b8b387e48ac9d50b32aa1a2c925 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 15 Aug 2022 15:41:10 +0200 Subject: [PATCH 187/287] refactor: add docstrings and type annotations in test_parcel.py --- junifer/markers/tests/test_parcel.py | 119 +++++++++++++++------------ 1 file changed, 65 insertions(+), 54 deletions(-) diff --git a/junifer/markers/tests/test_parcel.py b/junifer/markers/tests/test_parcel.py index a4e07a93f..1bd78a50f 100644 --- a/junifer/markers/tests/test_parcel.py +++ b/junifer/markers/tests/test_parcel.py @@ -1,24 +1,26 @@ """Provide test for parcel aggregation.""" -import numpy as np -from numpy.testing import assert_array_equal, assert_array_almost_equal -from scipy.stats import trim_mean -import nibabel as nib +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL +import nibabel as nib +import numpy as np from nilearn import datasets -from nilearn.image import resample_to_img, math_img, concat_imgs -from nilearn.maskers import NiftiMasker, NiftiLabelsMasker +from nilearn.image import concat_imgs, math_img, resample_to_img +from nilearn.maskers import NiftiLabelsMasker, NiftiMasker +from numpy.testing import assert_array_almost_equal, assert_array_equal +from scipy.stats import trim_mean from junifer.markers.parcel import ParcelAggregation -def test_ParcelAggregation_3D(): +def test_ParcelAggregation_3D() -> None: """Test ParcelAggregation object on 3D images.""" - # Get the testing atlas (for nilearn) atlas = datasets.fetch_atlas_schaefer_2018(n_rois=100) - # Get the oasis VBM data: + # Get the oasis VBM data oasis_dataset = datasets.fetch_oasis_vbm(n_subjects=1) vbm = oasis_dataset.gray_matter_maps[0] img = nib.load(vbm) @@ -27,15 +29,15 @@ def test_ParcelAggregation_3D(): atlas_img_res = resample_to_img( atlas.maps, img, - interpolation='nearest', + interpolation="nearest", ) atlas_bin = math_img( - 'img != 0', + "img != 0", img=atlas_img_res, ) + # Create NiftiMasker masker = NiftiMasker(atlas_bin, target_affine=img.affine) - data = masker.fit_transform(img) atlas_values = masker.transform(atlas_img_res) atlas_values = np.squeeze(atlas_values).astype(int) @@ -47,29 +49,34 @@ def test_ParcelAggregation_3D(): manual.append(t_values) manual = np.array(manual)[np.newaxis, :] + # Create NiftiLabelsMasker nifti_masker = NiftiLabelsMasker(labels_img=atlas.maps) auto = nifti_masker.fit_transform(img) + # Check that arrays are almost equal assert_array_almost_equal(auto, manual) # Use the ParcelAggregation object marker = ParcelAggregation( - atlas='Schaefer100x7', method='mean', name='gmd_schaefer100x7_mean', - on='VBM_GM') # Test passing "on" as a keyword argument + atlas="Schaefer100x7", + method="mean", + name="gmd_schaefer100x7_mean", + on="VBM_GM", + ) # Test passing "on" as a keyword argument input = dict(VBM_GM=dict(data=img)) - jun_values3d_mean = marker.fit_transform(input)['VBM_GM']['data'] + jun_values3d_mean = marker.fit_transform(input)["VBM_GM"]["data"] assert jun_values3d_mean.ndim == 2 assert jun_values3d_mean.shape[0] == 1 assert_array_equal(manual, jun_values3d_mean) - meta = marker.get_meta('VBM_GM')['marker'] - assert meta['method'] == 'mean' - assert meta['atlas'] == 'Schaefer100x7' - assert meta['name'] == 'VBM_GM_gmd_schaefer100x7_mean' - assert meta['class'] == 'ParcelAggregation' - assert meta['kind'] == 'VBM_GM' - assert meta['method_params'] == {} + meta = marker.get_meta("VBM_GM")["marker"] + assert meta["method"] == "mean" + assert meta["atlas"] == "Schaefer100x7" + assert meta["name"] == "VBM_GM_gmd_schaefer100x7_mean" + assert meta["class"] == "ParcelAggregation" + assert meta["kind"] == "VBM_GM" + assert meta["method_params"] == {} # Test using another function (std) manual = [] @@ -79,77 +86,81 @@ def test_ParcelAggregation_3D(): manual = np.array(manual)[np.newaxis, :] # Use the ParcelAggregation object - marker = ParcelAggregation(atlas='Schaefer100x7', method='std') + marker = ParcelAggregation(atlas="Schaefer100x7", method="std") input = dict(VBM_GM=dict(data=img)) - jun_values3d_std = marker.fit_transform(input)['VBM_GM']['data'] + jun_values3d_std = marker.fit_transform(input)["VBM_GM"]["data"] assert jun_values3d_std.ndim == 2 assert jun_values3d_std.shape[0] == 1 assert_array_equal(manual, jun_values3d_std) - meta = marker.get_meta('VBM_GM')['marker'] - assert meta['method'] == 'std' - assert meta['atlas'] == 'Schaefer100x7' - assert meta['name'] == 'VBM_GM_ParcelAggregation' - assert meta['class'] == 'ParcelAggregation' - assert meta['kind'] == 'VBM_GM' - assert meta['method_params'] == {} + meta = marker.get_meta("VBM_GM")["marker"] + assert meta["method"] == "std" + assert meta["atlas"] == "Schaefer100x7" + assert meta["name"] == "VBM_GM_ParcelAggregation" + assert meta["class"] == "ParcelAggregation" + assert meta["kind"] == "VBM_GM" + assert meta["method_params"] == {} # Test using another function with parameters manual = [] for t_v in sorted(np.unique(atlas_values)): t_values = trim_mean( - data[:, atlas_values == t_v], proportiontocut=0.1, - axis=None) # type: ignore + data[:, atlas_values == t_v], proportiontocut=0.1, axis=None + ) # type: ignore manual.append(t_values) manual = np.array(manual)[np.newaxis, :] # Use the ParcelAggregation object marker = ParcelAggregation( - atlas='Schaefer100x7', method='trim_mean', - method_params={'proportiontocut': 0.1}) + atlas="Schaefer100x7", + method="trim_mean", + method_params={"proportiontocut": 0.1}, + ) input = dict(VBM_GM=dict(data=img)) - jun_values3d_tm = marker.fit_transform(input)['VBM_GM']['data'] + jun_values3d_tm = marker.fit_transform(input)["VBM_GM"]["data"] assert jun_values3d_tm.ndim == 2 assert jun_values3d_tm.shape[0] == 1 assert_array_equal(manual, jun_values3d_tm) - meta = marker.get_meta('VBM_GM')['marker'] - assert meta['method'] == 'trim_mean' - assert meta['atlas'] == 'Schaefer100x7' - assert meta['name'] == 'VBM_GM_ParcelAggregation' - assert meta['class'] == 'ParcelAggregation' - assert meta['kind'] == 'VBM_GM' - assert meta['method_params'] == {'proportiontocut': 0.1} + meta = marker.get_meta("VBM_GM")["marker"] + assert meta["method"] == "trim_mean" + assert meta["atlas"] == "Schaefer100x7" + assert meta["name"] == "VBM_GM_ParcelAggregation" + assert meta["class"] == "ParcelAggregation" + assert meta["kind"] == "VBM_GM" + assert meta["method_params"] == {"proportiontocut": 0.1} def test_ParcelAggregation_4D(): """Test ParcelAggregation object on 4D images.""" - # Get the testing atlas (for nilearn) atlas = datasets.fetch_atlas_schaefer_2018( - n_rois=100, yeo_networks=7, resolution_mm=2) + n_rois=100, yeo_networks=7, resolution_mm=2 + ) # Get the SPM auditory data: subject_data = datasets.fetch_spm_auditory() fmri_img = concat_imgs(subject_data.func) # type: ignore + # Create NiftiLabelsMasker nifti_masker = NiftiLabelsMasker(labels_img=atlas.maps) auto4d = nifti_masker.fit_transform(fmri_img) - marker = ParcelAggregation(atlas='Schaefer100x7', method='mean') + # Create ParcelAggregation object + marker = ParcelAggregation(atlas="Schaefer100x7", method="mean") input = dict(BOLD=dict(data=fmri_img)) - jun_values4d = marker.fit_transform(input)['BOLD']['data'] + jun_values4d = marker.fit_transform(input)["BOLD"]["data"] assert jun_values4d.ndim == 2 assert_array_equal(auto4d.shape, jun_values4d.shape) assert_array_equal(auto4d, jun_values4d) - meta = marker.get_meta('BOLD')['marker'] - assert meta['method'] == 'mean' - assert meta['atlas'] == 'Schaefer100x7' - assert meta['name'] == 'BOLD_ParcelAggregation' - assert meta['class'] == 'ParcelAggregation' - assert meta['kind'] == 'BOLD' - assert meta['method_params'] == {} + meta = marker.get_meta("BOLD")["marker"] + assert meta["method"] == "mean" + assert meta["atlas"] == "Schaefer100x7" + assert meta["name"] == "BOLD_ParcelAggregation" + assert meta["class"] == "ParcelAggregation" + assert meta["kind"] == "BOLD" + assert meta["method_params"] == {} -- 2.52.0 From 78105e625ddacb8e15166acf5a10bb452978a63b Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 15 Aug 2022 15:42:04 +0200 Subject: [PATCH 188/287] refactor: improve imports for markers sub-package --- junifer/markers/__init__.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/junifer/markers/__init__.py b/junifer/markers/__init__.py index 5e67e4920..823d27bcd 100644 --- a/junifer/markers/__init__.py +++ b/junifer/markers/__init__.py @@ -1,5 +1,10 @@ +"""Provide imports for markers sub-package.""" + # Authors: Federico Raimondo # Leonard Sasse # License: AGPL + +from .base import BaseMarker from .collection import MarkerCollection -from .parcel import ParcelAggregation \ No newline at end of file +from .parcel import ParcelAggregation +from .pipeline_mixin import PipelineStepMixin -- 2.52.0 From 518fd3e338f0f7c273ce1213235aa4967c3a10e2 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 16 Aug 2022 16:39:31 +0200 Subject: [PATCH 189/287] refactor: add docstrings and type annotations in preprocess/confounds.py --- junifer/preprocess/confounds.py | 337 ++++++++++++++++++-------------- 1 file changed, 189 insertions(+), 148 deletions(-) diff --git a/junifer/preprocess/confounds.py b/junifer/preprocess/confounds.py index ea53097c7..869097b08 100644 --- a/junifer/preprocess/confounds.py +++ b/junifer/preprocess/confounds.py @@ -1,26 +1,73 @@ -"""Provide classes for confound removal.""" +"""Provide base class for confound removal.""" # Authors: Federico Raimondo # Leonard Sasse +# Synchon Mandal # License: AGPL +from typing import TYPE_CHECKING, Dict, List, Optional + import numpy as np import pandas as pd - from nilearn._utils.niimg_conversions import check_niimg_4d -from nilearn.masking import compute_brain_mask from nilearn.image import clean_img +from nilearn.masking import compute_brain_mask -from junifer.markers.base import PipelineStepMixin -from ..utils.logging import logger, raise_error +from ..markers import PipelineStepMixin +from ..utils import logger, raise_error + + +if TYPE_CHECKING: + from nibabel import Nifti1Image class BaseConfoundRemover(PipelineStepMixin): - """Base class for cofound removal. + """Base class for confound removal. Read confound files and select columns according to a pre-defined strategy. + Confound removal is based on `nilearn.image.clean_img`. + + Parameters + ----------- + strategy : dict, optional + The keys of the dictionary should correspond to names of noise + components to include: + - 'motion' + - 'wm_csf' + - 'global_signal' + The values of dictionary should correspond to types of confounds + extracted from each signal: + - 'basic': only the confounding time series + - 'power2': signal + quadratic term + - 'derivatives': signal + derivatives + - 'full': signal + deriv. + quadratic terms + power2 deriv. + (default None). + spike : float, optional + If None, no spike regressor is added. If spike is a float, it will + add a spike regressor for every point at which FD exceeds the + specified float (default None). + detrend : bool, Optional + If True, detrending will be applied on timeseries + (before confound removal) (default True). + standardize : bool, optional + If True, returned signals are set to unit variance (default True). + low_pass : float, optional + Low cutoff frequencies, in Hertz. If None, no filtering is applied + (default None). + high_pass : float, optional + High cutoff frequencies, in Hertz. If None, no filtering is + applied (default None). + t_r : float, optional + Repetition time, in second (sampling period). + If None, it will use t_r from nifti header (default None). + mask_img: Niimg-like object, optional + If provided, signal is only cleaned from voxels inside the mask. + If mask is provided, it should have same shape and affine as imgs. + If not provided, a mask is computed using + `nilearn.masking.compute_brain_mask` (default None). + """ # lower priority @@ -30,64 +77,23 @@ class BaseConfoundRemover(PipelineStepMixin): # TODO: Implement read_confounds for fmriprep data def __init__( - self, - strategy=None, - spike=None, - detrend=True, - standardize=True, - low_pass=None, - high_pass=None, - t_r=None, - mask_img=None, - ): - """Initialise a BaseConfoundReader object. - - Confound removal is based on nilearn.image.clean_img - - Parameters - ----------- - strategy : dict[str -> str] - The keys of the dictionary should correspond to names of noise - components to include: - - 'motion' - - 'wm_csf' - - 'global_signal' - The values of dictionary should correspond to types of confounds - extracted from each signal: - - 'basic': only the confounding time series - - 'power2': signal + quadratic term - - 'derivatives': signal + derivatives - - 'full': signal + deriv. + quadratic terms + power2 deriv. - spike : float | None (default) - If None, no spike regressor is added. If spike is a float, it will - add a spike regressor for every point at which FD exceeds the - specified float. - detrend : bool - If True (default), detrending will be applied on timeseries - (before confound removal). - standardize : bool - If True (default), returned signals are set to unit variance. - low_pass : float | None (default) - Low cutoff frequencies, in Hertz. If None, no filtering is applied. - high_pass : float | None (default) - High cutoff frequencies, in Hertz. If None, no filtering is - applied. - t_r : float - Repetition time, in second (sampling period). - If None (default) it will use t_r from nifti header. - mask_img: Niimg-like object - If provided, signal is only cleaned from voxels inside the mask. - If mask is provided, it should have same shape and affine as imgs. - If not provided, a mask is computed using - nilearn.masking.compute_brain_mask - """ + self, + strategy: Optional[Dict[str, str]] = None, + spike: Optional[float] = None, + detrend: bool = True, + standardize: bool = True, + low_pass: Optional[float] = None, + high_pass: Optional[float] = None, + t_r: Optional[float] = None, + mask_img: Optional["Nifti1Image"] = None, + ) -> None: + """Initialise the class.""" if strategy is None: strategy = { - 'motion': 'full', - 'wm_csf': 'full', - 'global_signal': 'full' + "motion": "full", + "wm_csf": "full", + "global_signal": "full", } - self.strategy = strategy self.spike = spike self.detrend = detrend @@ -97,87 +103,89 @@ class BaseConfoundRemover(PipelineStepMixin): self.t_r = t_r self.mask_img = mask_img - self._valid_components = ['motion', 'wm_csf', 'global_signal'] - self._valid_confounds = ['basic', 'power2', 'derivatives', 'full'] + self._valid_components = ["motion", "wm_csf", "global_signal"] + self._valid_confounds = ["basic", "power2", "derivatives", "full"] if any(not isinstance(k, str) for k in strategy.keys()): - raise_error( - 'Strategy keys must be strings', ValueError - ) + raise_error("Strategy keys must be strings", ValueError) if any(not isinstance(v, str) for v in strategy.values()): - raise_error( - 'Strategy values must be strings', ValueError - ) + raise_error("Strategy values must be strings", ValueError) if any(x not in self._valid_components for x in strategy.keys()): raise_error( - f'Invalid component names {list(strategy.keys())}. ' - f'Valid components are {self._valid_components}.\n' - f'If any of them is a valid parameter in ' - 'nilearn.interfaces.fmriprep.load_confounds we may ' - 'include it in the future', ValueError + msg=f"Invalid component names {list(strategy.keys())}. " + f"Valid components are {self._valid_components}.\n" + f"If any of them is a valid parameter in " + "nilearn.interfaces.fmriprep.load_confounds we may " + "include it in the future", + klass=ValueError, ) if any(x not in self._valid_confounds for x in strategy.values()): raise_error( - f'Invalid component names {list(strategy.values())}. ' - f'Valid confound types are {self._valid_confounds}.\n' - f'If any of them is a valid parameter in ' - 'nilearn.interfaces.fmriprep.load_confounds we may ' - 'include it in the future', ValueError + msg=f"Invalid component names {list(strategy.values())}. " + f"Valid confound types are {self._valid_confounds}.\n" + f"If any of them is a valid parameter in " + "nilearn.interfaces.fmriprep.load_confounds we may " + "include it in the future", + klass=ValueError, ) - def validate_input(self, input): + def validate_input(self, input: List[str]) -> None: """Validate the input to the pipeline step. Parameters ---------- - input : list[str] + input : list of str The input to the pipeline step. The list must contain the available Junifer Data dictionary keys. Raises ------ - ValueError: + ValueError If the input does not have the required data. + """ - _required_inputs = ['BOLD', 'confounds'] + _required_inputs = ["BOLD", "confounds"] if any(x not in input for x in _required_inputs): raise_error( - 'Input does not have the required data. \n' - f'Input: {input} \n' - f'Required (all off): {_required_inputs} \n', ValueError + msg="Input does not have the required data. \n" + f"Input: {input} \n" + f"Required (all off): {_required_inputs} \n", + klass=ValueError, ) - def get_output_kind(self, input): + def get_output_kind(self, input: List[str]) -> List[str]: """Get the kind of the pipeline step. Parameters ---------- - input : list[str] + input : list of str The input to the pipeline step. The list must contain the available Junifer Data dictionary keys. Returns ------- - output : list[str] + list of str The updated list of available Junifer Data dictionary keys after the pipeline step. + """ # Does not add any new keys return input + # TODO: complete type annotations def _pick_confounds(self, input): """Select relevant confounds from the specified file.""" to_select = [] - confounds_df = input['data'] - confounds_spec = input['names']['spec'] + confounds_df = input["data"] + confounds_spec = input["names"]["spec"] # for every confound there is a derivative # and for every confound + derivative there should be squares - derivatives_to_compute = input['names'].get('derivatives', {}) - squares_to_compute = input['names'].get('squares', {}) - spike_name = input['names']['spike'] + derivatives_to_compute = input["names"].get("derivatives", {}) + squares_to_compute = input["names"].get("squares", {}) + spike_name = input["names"]["spike"] # Get all the column names according to the strategy for comp, param in self.strategy.items(): @@ -189,7 +197,8 @@ class BaseConfoundRemover(PipelineStepMixin): if any(to_compute): for t_dst, t_src in derivatives_to_compute.items(): out_df[t_dst] = np.append( # type: ignore - np.diff(out_df[t_src]), 0) # type: ignore + np.diff(out_df[t_src]), 0 + ) # type: ignore # Add squares (of base confounds and derivatives) if needed to_compute = [x in squares_to_compute.keys() for x in to_select] @@ -203,13 +212,17 @@ class BaseConfoundRemover(PipelineStepMixin): fd = confounds_df[spike_name].copy() fd.loc[fd > self.spike] = 1 fd.loc[fd != 1] = 0 - out_df['spike'] = fd + out_df["spike"] = fd return out_df - def _remove_confounds(self, bold_img, confounds_df): + def _remove_confounds( + self, bold_img: "Nifti1Image", confounds_df: pd.DataFrame + ) -> "Nifti1Image": """Remove confounds from the BOLD image. + Parameters + ---------- bold_img : Niimg-like object 4D image. The signals in the last dimension are filtered (see http://nilearn.github.io/manipulating_images/input_output.html @@ -218,24 +231,26 @@ class BaseConfoundRemover(PipelineStepMixin): Dataframe containing confounds to remove. Number of rows should correspond to number of volumes in the BOLD image. - returns + Returns -------- - clean_bold : Niimg-like object - input image, cleaned. + Niimg-like object + Input image with confounds removed. """ confounds_array = confounds_df.values t_r = self.t_r if t_r is None: - logger.info('No t_r specified, using t_r from nifti header!') + logger.info("No `t_r` specified, using t_r from nifti header") zooms = bold_img.header.get_zooms() t_r = zooms[3] - logger.info(f'Read t_r from nifti header: {t_r}', ) + logger.info( + f"Read t_r from nifti header: {t_r}", + ) mask_img = self.mask_img if mask_img is None: - logger.info('Computing brain mask from image') + logger.info("Computing brain mask from image") mask_img = compute_brain_mask(bold_img) clean_bold = clean_img( @@ -246,28 +261,31 @@ class BaseConfoundRemover(PipelineStepMixin): low_pass=self.low_pass, high_pass=self.high_pass, t_r=t_r, - mask_img=mask_img + mask_img=mask_img, ) return clean_bold + # TODO: complete type annotations def _validate_data(self, input): + """Validate input data.""" # Bold must be 4D niimg - check_niimg_4d(input['BOLD']['data']) + check_niimg_4d(input["BOLD"]["data"]) # Confounds must be a dataframe - if not isinstance(input['confounds']['data'], pd.DataFrame): + if not isinstance(input["confounds"]["data"], pd.DataFrame): raise_error( - 'confounds data must be a pandas dataframe', ValueError + "confounds data must be a pandas dataframe", ValueError ) - confound_df = input['confounds']['data'] - bold_img = input['BOLD']['data'] + confound_df = input["confounds"]["data"] + bold_img = input["BOLD"]["data"] if bold_img.get_fdata().shape[3] != len(confound_df): raise_error( - 'Image time series and confounds have different length!\n' - f'\tImage time series: { bold_img.get_fdata().shape[3]}\n' - f'\tConfounds: {len(confound_df)}') + "Image time series and confounds have different length!\n" + f"\tImage time series: { bold_img.get_fdata().shape[3]}\n" + f"\tConfounds: {len(confound_df)}" + ) # Check the column names of the dataframe and the spec # spec must be a dictionary: @@ -290,72 +308,95 @@ class BaseConfoundRemover(PipelineStepMixin): # } # Check the columns in the dataframe - conf_spec = input['confounds']['names']['spec'] + conf_spec = input["confounds"]["names"]["spec"] if any(x not in conf_spec.keys() for x in self._valid_components): raise_error( - 'All of the component types must be in the confounds data ' - 'object `spec`. Please check your datagrabber.', ValueError) + "All of the component types must be in the confounds data " + "object `spec`. Please check your datagrabber.", + ValueError, + ) - if any(x not in v.keys() for x in self._valid_confounds - for v in conf_spec.values()): + if any( + x not in v.keys() + for x in self._valid_confounds + for v in conf_spec.values() + ): raise_error( - 'All of the confound types must be in the confounds data ' - 'object `spec`. Please check your datagrabber.', ValueError) + "All of the confound types must be in the confounds data " + "object `spec`. Please check your datagrabber.", + ValueError, + ) - spike_name = input['confounds']['names']['spike'] + spike_name = input["confounds"]["names"]["spike"] - derivatives_to_compute = input['confounds']['names'].get( - 'derivatives', {}) - if not(isinstance(derivatives_to_compute, dict)): + derivatives_to_compute = input["confounds"]["names"].get( + "derivatives", {} + ) + if not (isinstance(derivatives_to_compute, dict)): raise_error( 'input["confounds"]["names"]["derivatives"] ' - 'must be a dictionary. Please check your datagrabber', - ValueError) + "must be a dictionary. Please check your datagrabber", + ValueError, + ) - if any(not (isinstance(k, str) or isinstance(v, str)) - for k, v in derivatives_to_compute.items()): + if any( + not (isinstance(k, str) or isinstance(v, str)) + for k, v in derivatives_to_compute.items() + ): raise_error( 'input["confounds"]["names"]["derivatives"] ' - 'must be a dictionary with string keys and values. ' - 'Please check your datagrabber', - ValueError) + "must be a dictionary with string keys and values. " + "Please check your datagrabber", + ValueError, + ) missing_derivatives = [ - x for x in derivatives_to_compute.values() - if x not in confound_df.columns] + x + for x in derivatives_to_compute.values() + if x not in confound_df.columns + ] if len(missing_derivatives) > 0: raise_error( - 'Some of the derivatives to calculate are not in the confounds' - f' dataframe: {missing_derivatives}.' - 'Please check your data ' + "Some of the derivatives to calculate are not in the confounds" + f" dataframe: {missing_derivatives}." + "Please check your data " f'({input["confounds"]["path"].as_posix()}) ' - 'and the datagrabber.', ValueError) + "and the datagrabber.", + ValueError, + ) - t_conf_spec = {k: input['confounds']['names']['spec'][k][v] - for k, v in self.strategy.items()} + t_conf_spec = { + k: input["confounds"]["names"]["spec"][k][v] + for k, v in self.strategy.items() + } column_names = set([x for y in t_conf_spec.values() for x in y]) column_names.add(spike_name) missing_columns = [ - x for x in column_names - if x not in confound_df.columns and - x not in derivatives_to_compute.keys()] + x + for x in column_names + if x not in confound_df.columns + and x not in derivatives_to_compute.keys() + ] if len(missing_columns) > 0: raise_error( - 'Some of the columns in the confound spec are not in the ' - f'confounds dataframe: {missing_columns}. ' - 'Please check your data ' + "Some of the columns in the confound spec are not in the " + f"confounds dataframe: {missing_columns}. " + "Please check your data " f'({input["confounds"]["path"].as_posix()}) ' - 'and the datagrabber.', ValueError) + "and the datagrabber.", + ValueError, + ) + # TODO: complete type annotations def fit_transform(self, input): """Fit and transform.""" self._validate_data(input) - bold_img = input['BOLD']['data'] - confounds_df = self._pick_confounds(input['confounds']) - input['BOLD']['data'] = self._remove_confounds(bold_img, confounds_df) + bold_img = input["BOLD"]["data"] + confounds_df = self._pick_confounds(input["confounds"]) + input["BOLD"]["data"] = self._remove_confounds(bold_img, confounds_df) # TODO: Update meta return input -- 2.52.0 From 97a6609d38eb2f399a59724de0428b1393d6b5fd Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 16 Aug 2022 16:40:39 +0200 Subject: [PATCH 190/287] refactor: add docstrings and type annotations in test_confounds.py --- junifer/preprocess/tests/test_confounds.py | 190 +++++++++++---------- 1 file changed, 97 insertions(+), 93 deletions(-) diff --git a/junifer/preprocess/tests/test_confounds.py b/junifer/preprocess/tests/test_confounds.py index 69afa9f76..01463cdd9 100644 --- a/junifer/preprocess/tests/test_confounds.py +++ b/junifer/preprocess/tests/test_confounds.py @@ -2,26 +2,34 @@ # Authors: Federico Raimondo # Leonard Sasse +# Synchon Mandal # License: AGPL +import random +import string from pathlib import Path +from typing import Tuple + import numpy as np import pandas as pd from nibabel import Nifti1Image from nilearn._utils.niimg_conversions import check_niimg_4d -import random -import string + from junifer.preprocess.confounds import BaseConfoundRemover + +# Set RNG seed for reproducibility np.random.seed(1234567) -def generate_conf_name(size=6, chars=string.ascii_uppercase + string.digits): +def generate_conf_name( + size: int = 6, chars: str = string.ascii_uppercase + string.digits +) -> str: """Generate configuration name.""" - return ''.join(random.choice(chars) for _ in range(size)) + return "".join(random.choice(chars) for _ in range(size)) -def _simu_img(): +def _simu_img() -> Tuple[Nifti1Image, Nifti1Image]: # Random 4D volume with 100 time points vol = 100 + 10 * np.random.randn(5, 5, 2, 100) img = Nifti1Image(vol, np.eye(4)) @@ -30,31 +38,32 @@ def _simu_img(): return img, mask -def test_baseconfoundremover(): +# TODO: split the tests +def test_baseconfoundremover() -> None: """Test BaseConfoundRemover.""" # Generate a simulated BOLD img siimg, simsk = _simu_img() # generate random confound dataframe with Felix's column naming - motion_basic = [f'RP.{i}' for i in range(1, 7)] - motion_power2 = [f'RP^2.{i}' for i in range(1, 7)] - motion_derivatives = [f'DRP.{i}' for i in range(1, 7)] - motion_full = [f'DRP^2.{i}' for i in range(1, 7)] + motion_basic = [f"RP.{i}" for i in range(1, 7)] + motion_power2 = [f"RP^2.{i}" for i in range(1, 7)] + motion_derivatives = [f"DRP.{i}" for i in range(1, 7)] + motion_full = [f"DRP^2.{i}" for i in range(1, 7)] - wm_csf_basic = ['WM', 'CSF'] - wm_csf_power2 = ['WM^2', 'CSF^2'] - wm_csf_derivatives = ['DWM', 'DCSF'] - wm_csf_full = ['DWM^2', 'DCSF^2'] + wm_csf_basic = ["WM", "CSF"] + wm_csf_power2 = ["WM^2", "CSF^2"] + wm_csf_derivatives = ["DWM", "DCSF"] + wm_csf_full = ["DWM^2", "DCSF^2"] - gs_basic = ['GS'] - gs_power2 = ['GS^2'] - gs_derivatives = ['DGS'] - gs_full = ['DGS^2'] + gs_basic = ["GS"] + gs_power2 = ["GS^2"] + gs_derivatives = ["DGS"] + gs_full = ["DGS^2"] confound_column_names = [] - confound_column_names.append('FD') # spike + confound_column_names.append("FD") # spike confound_column_names.extend(motion_basic) confound_column_names.extend(motion_power2) @@ -79,54 +88,54 @@ def test_baseconfoundremover(): n_cols = len(confound_column_names) confounds_df = pd.DataFrame( np.random.randint(0, 100, size=(100, n_cols)), - columns=confound_column_names + columns=confound_column_names, ) # Generate spec from Felix's column naming spec = { - 'motion': { - 'basic': motion_basic, - 'power2': motion_basic + motion_power2, - 'derivatives': motion_basic + motion_derivatives, - 'full': - motion_basic + motion_derivatives + motion_power2 + motion_full + "motion": { + "basic": motion_basic, + "power2": motion_basic + motion_power2, + "derivatives": motion_basic + motion_derivatives, + "full": motion_basic + + motion_derivatives + + motion_power2 + + motion_full, }, - 'wm_csf': { - 'basic': wm_csf_basic, - 'power2': wm_csf_basic + wm_csf_power2, - 'derivatives': wm_csf_basic + wm_csf_derivatives, - 'full': - wm_csf_basic + wm_csf_derivatives + wm_csf_power2 + wm_csf_full + "wm_csf": { + "basic": wm_csf_basic, + "power2": wm_csf_basic + wm_csf_power2, + "derivatives": wm_csf_basic + wm_csf_derivatives, + "full": wm_csf_basic + + wm_csf_derivatives + + wm_csf_power2 + + wm_csf_full, + }, + "global_signal": { + "basic": gs_basic, + "power2": gs_basic + gs_power2, + "derivatives": gs_basic + gs_derivatives, + "full": gs_basic + gs_derivatives + gs_power2 + gs_full, }, - 'global_signal': { - 'basic': gs_basic, - 'power2': gs_basic + gs_power2, - 'derivatives': gs_basic + gs_derivatives, - 'full': gs_basic + gs_derivatives + gs_power2 + gs_full - } } # generate a junifer pipeline data object dictionary input_data_obj = {} - input_data_obj['meta'] = {} - input_data_obj['BOLD'] = {} - input_data_obj['BOLD']['data'] = siimg - input_data_obj['confounds'] = {} - input_data_obj['confounds']['path'] = Path('/test.df') - input_data_obj['confounds']['data'] = confounds_df - input_data_obj['confounds']['names'] = {} - input_data_obj['confounds']['names']['spec'] = spec - input_data_obj['confounds']['names']['spike'] = 'FD' + input_data_obj["meta"] = {} + input_data_obj["BOLD"] = {} + input_data_obj["BOLD"]["data"] = siimg + input_data_obj["confounds"] = {} + input_data_obj["confounds"]["path"] = Path("/test.df") + input_data_obj["confounds"]["data"] = confounds_df + input_data_obj["confounds"]["names"] = {} + input_data_obj["confounds"]["names"]["spec"] = spec + input_data_obj["confounds"]["names"]["spike"] = "FD" # generate confound removal strategies with varying numbers of parameters # Test #1: 36 params, no derivatives to compute, no spike # 36 params - strat1 = { - 'motion': 'full', - 'wm_csf': 'full', - 'global_signal': 'full' - } + strat1 = {"motion": "full", "wm_csf": "full", "global_signal": "full"} cr = BaseConfoundRemover( strategy=strat1, spike=None, mask_img=simsk, t_r=0.75 @@ -134,13 +143,13 @@ def test_baseconfoundremover(): cr.validate_input(input_data_obj.keys()) out_type = cr.get_output_kind(input_data_obj.keys()) - assert 'BOLD' in out_type + assert "BOLD" in out_type # Check if the input data is valid cr._validate_data(input_data_obj) # Check that the confounds are picked correctly: - t_df = cr._pick_confounds(input_data_obj['confounds']) + t_df = cr._pick_confounds(input_data_obj["confounds"]) assert len(t_df.columns) == 36 assert all(x in t_df.columns for x in motion_basic) assert all(x in t_df.columns for x in motion_power2) @@ -154,13 +163,13 @@ def test_baseconfoundremover(): assert all(x in t_df.columns for x in gs_power2) assert all(x in t_df.columns for x in gs_derivatives) assert all(x in t_df.columns for x in gs_full) - assert 'FD' not in t_df.columns - assert 'spike' not in t_df.columns + assert "FD" not in t_df.columns + assert "spike" not in t_df.columns # Test #2: 24 params, no derivatives to compute, no spike # 24 params strat2 = { - 'motion': 'full', + "motion": "full", } cr = BaseConfoundRemover( @@ -169,13 +178,13 @@ def test_baseconfoundremover(): cr.validate_input(input_data_obj.keys()) out_type = cr.get_output_kind(input_data_obj.keys()) - assert 'BOLD' in out_type + assert "BOLD" in out_type # Check if the input data is valid cr._validate_data(input_data_obj) # Check that the confounds are picked correctly: - t_df = cr._pick_confounds(input_data_obj['confounds']) + t_df = cr._pick_confounds(input_data_obj["confounds"]) assert len(t_df.columns) == 24 assert all(x in t_df.columns for x in motion_basic) assert all(x in t_df.columns for x in motion_power2) @@ -189,15 +198,11 @@ def test_baseconfoundremover(): assert all(x not in t_df.columns for x in gs_power2) assert all(x not in t_df.columns for x in gs_derivatives) assert all(x not in t_df.columns for x in gs_full) - assert 'FD' not in t_df.columns - assert 'spike' not in t_df.columns + assert "FD" not in t_df.columns + assert "spike" not in t_df.columns # Test #3: 9 params, no derivatives to compute, no spike - strat3 = { - 'motion': 'basic', - 'wm_csf': 'basic', - 'global_signal': 'basic' - } + strat3 = {"motion": "basic", "wm_csf": "basic", "global_signal": "basic"} cr = BaseConfoundRemover( strategy=strat3, spike=None, mask_img=simsk, t_r=0.75 @@ -205,13 +210,13 @@ def test_baseconfoundremover(): cr.validate_input(input_data_obj.keys()) out_type = cr.get_output_kind(input_data_obj.keys()) - assert 'BOLD' in out_type + assert "BOLD" in out_type # Check if the input data is valid cr._validate_data(input_data_obj) # Check that the confounds are picked correctly: - t_df = cr._pick_confounds(input_data_obj['confounds']) + t_df = cr._pick_confounds(input_data_obj["confounds"]) assert len(t_df.columns) == 9 assert all(x in t_df.columns for x in motion_basic) assert all(x not in t_df.columns for x in motion_power2) @@ -225,12 +230,12 @@ def test_baseconfoundremover(): assert all(x not in t_df.columns for x in gs_power2) assert all(x not in t_df.columns for x in gs_derivatives) assert all(x not in t_df.columns for x in gs_full) - assert 'FD' not in t_df.columns - assert 'spike' not in t_df.columns + assert "FD" not in t_df.columns + assert "spike" not in t_df.columns # Test #4: 6 params, no derivatives to compute, no spike strat4 = { - 'motion': 'basic', + "motion": "basic", } cr = BaseConfoundRemover( strategy=strat4, spike=None, mask_img=simsk, t_r=0.75 @@ -238,13 +243,13 @@ def test_baseconfoundremover(): cr.validate_input(input_data_obj.keys()) out_type = cr.get_output_kind(input_data_obj.keys()) - assert 'BOLD' in out_type + assert "BOLD" in out_type # Check if the input data is valid cr._validate_data(input_data_obj) # Check that the confounds are picked correctly: - t_df = cr._pick_confounds(input_data_obj['confounds']) + t_df = cr._pick_confounds(input_data_obj["confounds"]) assert len(t_df.columns) == 6 assert all(x in t_df.columns for x in motion_basic) assert all(x not in t_df.columns for x in motion_power2) @@ -258,12 +263,12 @@ def test_baseconfoundremover(): assert all(x not in t_df.columns for x in gs_power2) assert all(x not in t_df.columns for x in gs_derivatives) assert all(x not in t_df.columns for x in gs_full) - assert 'FD' not in t_df.columns - assert 'spike' not in t_df.columns + assert "FD" not in t_df.columns + assert "spike" not in t_df.columns # Test #5: 2 params, no derivatives to compute, no spike strat5 = { - 'wm_csf': 'basic', + "wm_csf": "basic", } cr = BaseConfoundRemover( strategy=strat5, spike=None, mask_img=simsk, t_r=0.75 @@ -271,13 +276,13 @@ def test_baseconfoundremover(): cr.validate_input(input_data_obj.keys()) out_type = cr.get_output_kind(input_data_obj.keys()) - assert 'BOLD' in out_type + assert "BOLD" in out_type # Check if the input data is valid cr._validate_data(input_data_obj) # Check that the confounds are picked correctly: - t_df = cr._pick_confounds(input_data_obj['confounds']) + t_df = cr._pick_confounds(input_data_obj["confounds"]) assert len(t_df.columns) == 2 assert all(x not in t_df.columns for x in motion_basic) assert all(x not in t_df.columns for x in motion_power2) @@ -291,27 +296,26 @@ def test_baseconfoundremover(): assert all(x not in t_df.columns for x in gs_power2) assert all(x not in t_df.columns for x in gs_derivatives) assert all(x not in t_df.columns for x in gs_full) - assert 'FD' not in t_df.columns - assert 'spike' not in t_df.columns + assert "FD" not in t_df.columns + assert "spike" not in t_df.columns out = cr.fit_transform(input_data_obj) - check_niimg_4d(out['BOLD']['data']) + check_niimg_4d(out["BOLD"]["data"]) # TODO: check meta # Test #6: 12 params, derivatives to compute, spike - to_select = [x for x in confounds_df.columns - if x not in motion_derivatives] + to_select = [ + x for x in confounds_df.columns if x not in motion_derivatives + ] no_d_df = confounds_df[to_select] - input_data_obj['confounds']['data'] = no_d_df + input_data_obj["confounds"]["data"] = no_d_df - derivatives = { - f'D{x}': x for x in motion_basic - } + derivatives = {f"D{x}": x for x in motion_basic} - input_data_obj['confounds']['names']['derivatives'] = derivatives + input_data_obj["confounds"]["names"]["derivatives"] = derivatives strat6 = { - 'motion': 'derivatives', + "motion": "derivatives", } cr = BaseConfoundRemover( strategy=strat6, spike=0.75, mask_img=simsk, t_r=0.75 @@ -319,13 +323,13 @@ def test_baseconfoundremover(): cr.validate_input(input_data_obj.keys()) out_type = cr.get_output_kind(input_data_obj.keys()) - assert 'BOLD' in out_type + assert "BOLD" in out_type # Check if the input data is valid cr._validate_data(input_data_obj) # Check that the confounds are picked correctly: - t_df = cr._pick_confounds(input_data_obj['confounds']) + t_df = cr._pick_confounds(input_data_obj["confounds"]) assert len(t_df.columns) == 13 assert all(x in t_df.columns for x in motion_basic) assert all(x not in t_df.columns for x in motion_power2) @@ -339,5 +343,5 @@ def test_baseconfoundremover(): assert all(x not in t_df.columns for x in gs_power2) assert all(x not in t_df.columns for x in gs_derivatives) assert all(x not in t_df.columns for x in gs_full) - assert 'FD' not in t_df.columns - assert 'spike' in t_df.columns + assert "FD" not in t_df.columns + assert "spike" in t_df.columns -- 2.52.0 From 25313c84131b78f98f07a51aefd71bf9501c6b82 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 16 Aug 2022 16:41:03 +0200 Subject: [PATCH 191/287] refactor: improve imports for preprocess sub-package --- junifer/preprocess/__init__.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/junifer/preprocess/__init__.py b/junifer/preprocess/__init__.py index 4e235bc8b..b7016a59d 100644 --- a/junifer/preprocess/__init__.py +++ b/junifer/preprocess/__init__.py @@ -1,3 +1,7 @@ +"""Provide imports for preprocess sub-package.""" + # Authors: Federico Raimondo # Leonard Sasse -# License: AGPL \ No newline at end of file +# License: AGPL + +from .confounds import BaseConfoundRemover -- 2.52.0 From e0a138a66d3407d1efa4543787526089aeb80006 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 15 Aug 2022 16:16:05 +0200 Subject: [PATCH 192/287] refactor: add docstrings and type annotations for datareader/default.py --- junifer/datareader/default.py | 92 +++++++++++++++++++++++------------ 1 file changed, 61 insertions(+), 31 deletions(-) diff --git a/junifer/datareader/default.py b/junifer/datareader/default.py index d691f2849..34996f6de 100644 --- a/junifer/datareader/default.py +++ b/junifer/datareader/default.py @@ -1,47 +1,75 @@ """Provide class for default data reader.""" # Authors: Federico Raimondo +# Synchon Mandal # License: AGPL from pathlib import Path +from typing import Dict import nibabel as nib import pandas as pd -from ..utils.logging import logger, warn from ..markers.base import PipelineStepMixin +from ..utils.logging import logger, warn_with_log -# Map each filenanm end to a kind + +# Map each file extension to a kind _extensions = { - '.nii': 'NIFTI', - '.nii.gz': 'NIFTI', - '.csv': 'CSV', - '.tsv': 'TSV' - + ".nii": "NIFTI", + ".nii.gz": "NIFTI", + ".csv": "CSV", + ".tsv": "TSV", } # Map each kind to a function and arguments _readers = {} -_readers['NIFTI'] = dict(func=nib.load, params=None) -_readers['CSV'] = dict(func=pd.read_csv, params=None) -_readers['TSV'] = dict(func=pd.read_csv, params={'sep': '\t'}) +_readers["NIFTI"] = {"func": nib.load, "params": None} +_readers["CSV"] = {"func": pd.read_csv, "params": None} +_readers["TSV"] = {"func": pd.read_csv, "params": {"sep": "\t"}} class DefaultDataReader(PipelineStepMixin): """Mixin class for default data reader.""" + # TODO: complete type annotations def validate_input(self, input): - """Validate input.""" + """Validate input. + + Parameters + ---------- + input + + """ # Nothing to validate, any input is fine pass + # TODO: complete type annotations def get_output_kind(self, input): - """Get output kind.""" + """Get output kind. + + Parameters + ---------- + input + + """ # It will output the same kind of data as the input return input - def fit_transform(self, input, params=None): - """Fit and transform.""" + # TODO: complete type annotations + def fit_transform(self, input, params=None) -> Dict: + """Fit and transform. + + Parameters + ---------- + input + params + + Returns + ------- + dict + + """ # For each kind of data, try to read it # out is the same, but with the 'data' key set in @@ -50,39 +78,41 @@ class DefaultDataReader(PipelineStepMixin): if params is None: params = {} for kind in input.keys(): - if kind == 'meta': - out['meta'] = input['meta'] + if kind == "meta": + out["meta"] = input["meta"] continue - if 'path' not in input[kind]: - warn( - f'Input kind {kind} does not provide a path. Skipping.') + if "path" not in input[kind]: + warn_with_log( + f"Input kind {kind} does not provide a path. Skipping." + ) continue - t_path = input[kind]['path'] + t_path = input[kind]["path"] t_params = params.get(kind, {}) # Convert to Path if datareader is not well done if not isinstance(t_path, Path): t_path = Path(t_path) - out[kind]['path'] = t_path - logger.info(f'Reading {kind} from {t_path.as_posix()}') + out[kind]["path"] = t_path + logger.info(f"Reading {kind} from {t_path.as_posix()}") fread = None fname = t_path.name.lower() for ext, ftype in _extensions.items(): if fname.endswith(ext): - logger.info(f'{kind} is type {ftype}') - reader_func = _readers[ftype]['func'] - reader_params = _readers[ftype]['params'] + logger.info(f"{kind} is type {ftype}") + reader_func = _readers[ftype]["func"] + reader_params = _readers[ftype]["params"] if reader_params is not None: t_params.update(reader_params) - logger.debug(f'Calling {reader_func} with {t_params}') + logger.debug(f"Calling {reader_func} with {t_params}") fread = reader_func(t_path, **t_params) break if fread is None: logger.info( - f'Unknown file type {t_path.as_posix()}, skipping reading') - out[kind]['data'] = fread - if 'meta' not in out: - out['meta'] = {} - out['meta']['datareader'] = self.get_meta() + f"Unknown file type {t_path.as_posix()}, skipping reading" + ) + out[kind]["data"] = fread + if "meta" not in out: + out["meta"] = {} + out["meta"]["datareader"] = self.get_meta() return out -- 2.52.0 From 8fa1ada40738a11df58923d6d36d793d3a341b2a Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 15 Aug 2022 16:16:42 +0200 Subject: [PATCH 193/287] refactor: prune, add docstrings and type annotations in test_default_reader.py --- .../datareader/tests/test_default_reader.py | 206 ++++++++++-------- 1 file changed, 112 insertions(+), 94 deletions(-) diff --git a/junifer/datareader/tests/test_default_reader.py b/junifer/datareader/tests/test_default_reader.py index ab4867159..468583f7f 100644 --- a/junifer/datareader/tests/test_default_reader.py +++ b/junifer/datareader/tests/test_default_reader.py @@ -1,150 +1,168 @@ """Provide tests for default data reader.""" # Authors: Federico Raimondo +# Synchon Mandal # License: AGPL from pathlib import Path -import tempfile -from numpy.testing import assert_array_equal + import nibabel as nib -from nibabel import testing as nib_testing import pandas as pd +import pytest +from nibabel import testing as nib_testing +from numpy.testing import assert_array_equal from pandas.testing import assert_frame_equal from junifer.datareader import DefaultDataReader -def test_validation(): - """Test validating input/output.""" - kinds = [ - ['T1w', 'BOLD', 'T2', 'dwi'], - [], - None, - ['whatever'] - ] +@pytest.mark.parametrize( + "kind", [["T1w", "BOLD", "T2", "dwi"], [], None, ["whatever"]] +) +def test_validation(kind) -> None: + """Test validating input/output. + Parameters + ---------- + kind : list of str or str or None + The parametrized kind of data. + + """ reader = DefaultDataReader() - - for t_kind in kinds: - assert reader.validate_input(t_kind) is None - assert reader.get_output_kind(t_kind) == t_kind - assert reader.validate(t_kind) == t_kind + assert reader.validate_input(kind) is None + assert reader.get_output_kind(kind) == kind + assert reader.validate(kind) == kind -def test_meta(): +def test_meta() -> None: """Test reader metadata.""" reader = DefaultDataReader() t_meta = reader.get_meta() - assert t_meta['class'] == 'DefaultDataReader' + assert t_meta["class"] == "DefaultDataReader" nib_data_path = Path(nib_testing.data_path) - t_path = nib_data_path / 'example4d.nii.gz' - input = {'bold': {'path': t_path}} + t_path = nib_data_path / "example4d.nii.gz" + input = {"bold": {"path": t_path}} output = reader.fit_transform(input) - assert 'meta' in output - assert 'datareader' in output['meta'] - assert 'class' in output['meta']['datareader'] - assert output['meta']['datareader']['class'] == 'DefaultDataReader' + assert "meta" in output + assert "datareader" in output["meta"] + assert "class" in output["meta"]["datareader"] + assert output["meta"]["datareader"]["class"] == "DefaultDataReader" -def test_read_nifti(): - """Test reading NIFTI files.""" +@pytest.mark.parametrize( + "fname", ["example4d.nii.gz", "reoriented_anat_moved.nii"] +) +def test_read_nifti(fname: str) -> None: + """Test reading NIFTI files. + + Parameters + ---------- + fname : str + The parametrized NIfTI file names for testing. + + """ reader = DefaultDataReader() nib_data_path = Path(nib_testing.data_path) - for fname in ['example4d.nii.gz', - 'reoriented_anat_moved.nii']: - t_path = nib_data_path / fname + t_path = nib_data_path / fname - input = {'bold': {'path': t_path}} - output = reader.fit_transform(input) + input = {"bold": {"path": t_path}} + output = reader.fit_transform(input) - assert isinstance(output, dict) - assert 'bold' in output - assert isinstance(output['bold'], dict) - assert 'path' in output['bold'] - assert 'data' in output['bold'] + assert isinstance(output, dict) + assert "bold" in output + assert isinstance(output["bold"], dict) + assert "path" in output["bold"] + assert "data" in output["bold"] - read_img = output['bold']['data'] + read_img = output["bold"]["data"] - t_read_img = nib.load(t_path) - assert_array_equal(read_img.get_fdata(), t_read_img.get_fdata()) + t_read_img = nib.load(t_path) + assert_array_equal(read_img.get_fdata(), t_read_img.get_fdata()) - input = {'bold': {'path': t_path.as_posix()}} - output2 = reader.fit_transform(input) - assert output['bold']['path'] == output2['bold']['path'] + input = {"bold": {"path": t_path.as_posix()}} + output2 = reader.fit_transform(input) + assert output["bold"]["path"] == output2["bold"]["path"] -def test_read_unknown(): +def test_read_unknown() -> None: """Test (not) reading unknown files.""" reader = DefaultDataReader() nib_data_path = Path(nib_testing.data_path) - anat_path = nib_data_path / 'reoriented_anat_moved.nii' - whatever_path = nib_data_path / 'unexistent.unkwnownextension' + anat_path = nib_data_path / "reoriented_anat_moved.nii" + whatever_path = nib_data_path / "unexistent.unkwnownextension" - input = {'anat': {'path': anat_path}, 'whatever': {'path': whatever_path}} + input = {"anat": {"path": anat_path}, "whatever": {"path": whatever_path}} output = reader.fit_transform(input) assert isinstance(output, dict) - assert 'anat' in output - assert isinstance(output['anat'], dict) - assert 'path' in output['anat'] - assert isinstance(output['anat']['path'], Path) - assert 'data' in output['anat'] - assert output['anat']['data'] is not None + assert "anat" in output + assert isinstance(output["anat"], dict) + assert "path" in output["anat"] + assert isinstance(output["anat"]["path"], Path) + assert "data" in output["anat"] + assert output["anat"]["data"] is not None - assert isinstance(output['whatever'], dict) - assert 'path' in output['whatever'] - assert isinstance(output['whatever']['path'], Path) - assert 'data' in output['whatever'] - assert output['whatever']['data'] is None + assert isinstance(output["whatever"], dict) + assert "path" in output["whatever"] + assert isinstance(output["whatever"]["path"], Path) + assert "data" in output["whatever"] + assert output["whatever"]["data"] is None -def test_read_csv(): - """Test reading CSV files.""" - d = {'col1': [1, 2, 3, 4, 5], 'col2': [3, 4, 5, 6, 7]} +def test_read_csv(tmp_path: Path) -> None: + """Test reading CSV files. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + d = {"col1": [1, 2, 3, 4, 5], "col2": [3, 4, 5, 6, 7]} df = pd.DataFrame(d) - with tempfile.TemporaryDirectory() as tmpdir: - tmpdir = Path(tmpdir) - df.to_csv(tmpdir / 'test.csv') - reader = DefaultDataReader() - input = {'csv': {'path': tmpdir / 'test.csv'}} - output = reader.fit_transform(input) + df.to_csv(tmp_path / "test_read_csv.csv") - assert isinstance(output, dict) - assert 'csv' in output - assert isinstance(output['csv'], dict) - assert 'path' in output['csv'] - assert 'data' in output['csv'] + reader = DefaultDataReader() + input = {"csv": {"path": tmp_path / "test_read_csv.csv"}} + output = reader.fit_transform(input) - read_df = output['csv']['data'][['col1', 'col2']] - assert_frame_equal(df, read_df) + assert isinstance(output, dict) + assert "csv" in output + assert isinstance(output["csv"], dict) + assert "path" in output["csv"] + assert "data" in output["csv"] - df.to_csv(tmpdir / 'test.csv', sep=';') - input = {'csv': {'path': tmpdir / 'test.csv'}} - params = {'csv': {'sep': ';'}} - output = reader.fit_transform(input, params) + read_df = output["csv"]["data"][["col1", "col2"]] + assert_frame_equal(df, read_df) - assert isinstance(output, dict) - assert 'csv' in output - assert isinstance(output['csv'], dict) - assert 'path' in output['csv'] - assert 'data' in output['csv'] + df.to_csv(tmp_path / "test_read_csv.csv", sep=";") + input = {"csv": {"path": tmp_path / "test_read_csv.csv"}} + params = {"csv": {"sep": ";"}} + output = reader.fit_transform(input, params) - read_df = output['csv']['data'][['col1', 'col2']] - assert_frame_equal(df, read_df) + assert isinstance(output, dict) + assert "csv" in output + assert isinstance(output["csv"], dict) + assert "path" in output["csv"] + assert "data" in output["csv"] - df.to_csv(tmpdir / 'test.tsv', sep='\t') - input = {'csv': {'path': tmpdir / 'test.tsv'}} - output = reader.fit_transform(input) + read_df = output["csv"]["data"][["col1", "col2"]] + assert_frame_equal(df, read_df) - assert isinstance(output, dict) - assert 'csv' in output - assert isinstance(output['csv'], dict) - assert 'path' in output['csv'] - assert 'data' in output['csv'] + df.to_csv(tmp_path / "test_read_csv.tsv", sep="\t") + input = {"csv": {"path": tmp_path / "test_read_csv.tsv"}} + output = reader.fit_transform(input) - read_df = output['csv']['data'][['col1', 'col2']] - assert_frame_equal(df, read_df) + assert isinstance(output, dict) + assert "csv" in output + assert isinstance(output["csv"], dict) + assert "path" in output["csv"] + assert "data" in output["csv"] + + read_df = output["csv"]["data"][["col1", "col2"]] + # Check if dataframes are equal + assert_frame_equal(df, read_df) -- 2.52.0 From 19114ca6c93a69c8cce08bf071d342840682dd0d Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 15 Aug 2022 16:17:19 +0200 Subject: [PATCH 194/287] refactor: add docstring for datareader package import module --- junifer/datareader/__init__.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/junifer/datareader/__init__.py b/junifer/datareader/__init__.py index 4f1b6e6ca..19212063c 100644 --- a/junifer/datareader/__init__.py +++ b/junifer/datareader/__init__.py @@ -1,5 +1,8 @@ +"""Provide imports for datareader sub-package.""" + # Authors: Federico Raimondo # Leonard Sasse +# Synchon Mandal # License: AGPL -from .default import DefaultDataReader \ No newline at end of file +from .default import DefaultDataReader -- 2.52.0 From 9a46ee2ca2d2e6b038f148776fe448f799a587e5 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 16 Aug 2022 16:45:03 +0200 Subject: [PATCH 195/287] refactor: prune, add docstrings and type annotations in datagrabber/base.py --- junifer/datagrabber/base.py | 550 ++++++++---------------------------- 1 file changed, 114 insertions(+), 436 deletions(-) diff --git a/junifer/datagrabber/base.py b/junifer/datagrabber/base.py index b4731e902..a4908977b 100644 --- a/junifer/datagrabber/base.py +++ b/junifer/datagrabber/base.py @@ -1,143 +1,147 @@ -"""Provide class and functions for base datagrabber.""" +"""Provide abstract base class for datagrabber.""" # Authors: Federico Raimondo # Leonard Sasse +# Synchon Mandal # License: AGPL -from pathlib import Path -import re -import tempfile - -import datalad.api as dl from abc import ABC, abstractmethod +from pathlib import Path +from typing import Dict, Iterator, List, Tuple, Union -from ..api.decorators import register_datagrabber -from ..utils.logging import logger, raise_error, warn - - -def _validate_types(types): - """Validate the types.""" - if not isinstance(types, list): - raise_error("types must be a list", TypeError) # type: ignore - if any(not isinstance(x, str) for x in types): - raise_error( - "types must be a list of strings", TypeError) # type: ignore - - -def _validate_replacements(replacements, patterns): - """Validate the replacements.""" - if not isinstance(replacements, list): - raise_error("replacements must be a list", TypeError) # type: ignore - if any(not isinstance(x, str) for x in replacements): - raise_error( - "replacements must be a list of strings", - TypeError) # type: ignore - - for x in replacements: - if all(x not in y for y in patterns.values()): - warn(f"Replacement {x} is not part of any pattern") - - -def _validate_patterns(types, patterns): - """Validate the patterns.""" - _validate_types(types) - if not isinstance(patterns, dict): - raise_error("patterns must be a dict", TypeError) # type: ignore - if len(types) != len(patterns): - raise_error("types and patterns must have the same length", ValueError) - - if any(x not in patterns for x in types): - raise_error("patterns must contain all types", ValueError) +from ..utils import logger, raise_error +from .utils import validate_types class BaseDataGrabber(ABC): - """Base DataGrabber class (abstract). + """Abstract base class for datagrabber. + + For every interface that is required, one needs to provide a concrete + implementation of this abstract class. + + Parameters + ---------- + types : list of str + The types of data to be grabbed. + datadir : str or pathlib.Path + The directory where the data is / will be stored. Attributes ---------- - datadir - types : list - List of data types to be grabbed. + datadir : pathlib.Path + The directory where the data is / will be stored. - Methods - ------- - __getitem__(element) : dict[str -> Path] - Returns a dictionary of paths for each type of data required for the - specified element. Use the element as a key to index the datagrabber. - __enter__() : self - Returns the object itself. Can be overridden by subclasses. - __exit__() : None - Does nothing. Can be overridden by subclasses to clean up after - `__enter__` """ - def __init__(self, types, datadir): - """Initialize a BaseDataGrabber object. - - Parameters - ---------- - types : list of str - The types of data to be grabbed. - datadir : str or Path - That directory where the data is/will be stored. - """ - _validate_types(types) + def __init__(self, types: List[str], datadir: Union[str, Path]) -> None: + """Initialize the class.""" + # Validate types + validate_types(types) + # Convert str to Path if not isinstance(datadir, Path): datadir = Path(datadir) - logger.debug('Initializing BaseDataGrabber') - logger.debug(f'\t_datadir = {datadir}') - logger.debug(f'\ttypes = {types}') + logger.debug("Initializing BaseDataGrabber") + logger.debug(f"\t_datadir = {datadir}") + logger.debug(f"\ttypes = {types}") self._datadir = datadir self.types = types - def get_types(self): - """Get types.""" - return self.types.copy() - - def get_meta(self): - """Get metadata.""" - t_meta = {} - t_meta['class'] = self.__class__.__name__ - for k, v in vars(self).items(): - if not k.startswith('_'): - t_meta[k] = v - return t_meta - - def get_element_keys(self): - """Get element keys.""" - return 'element' - - @property - def datadir(self): - """Get data directory path. - - Returns - ------- - Path to the data directory. Implemented as a property, can be - overridden by subclasses. - """ - return self._datadir - - def __iter__(self): - """Iterate over elements in the datagrabber. + def __iter__(self) -> Iterator: + """Enable iterable support. Yields ------ - element : object + object An element that can be indexed by the datagrabber. + """ for elem in self.get_elements(): yield elem - def __getitem__(self, element): - """Get item implementation.""" - logger.info(f'Getting element {element}') + # TODO: element does nothing, check + def __getitem__(self, element: Union[str, Tuple]) -> Dict[str, Path]: + """Enable indexing support. + + Parameters + ---------- + element : str or tuple + The element to be indexed. If one string is provided, it is + assumed to be a tuple with only one item. If a tuple is provided, + each item in the tuple is the value for the replacement string + specified in "replacements". + + Returns + ------- + dict + Dictionary of paths for each type of data required for the + specified element. + + """ + logger.info(f"Getting element {element}") out = {} - out['meta'] = dict(datagrabber=self.get_meta()) + out["meta"] = {"datagrabber": self.get_meta()} return out + def __enter__(self) -> "BaseDataGrabber": + """Context entry.""" + return self + + def __exit__(self, exc_type, exc_value, exc_traceback) -> None: + """Context exit.""" + return None + + def get_types(self) -> List[str]: + """Get types. + + Returns + ------- + list of str + The types of data to be grabbed. + + """ + return self.types.copy() + + def get_meta(self) -> Dict: + """Get metadata. + + Returns + ------- + dict + The metadata as dictionary. + + """ + t_meta = {} + t_meta["class"] = self.__class__.__name__ + for k, v in vars(self).items(): + if not k.startswith("_"): + t_meta[k] = v + return t_meta + + # TODO: what is the final functionality? + def get_element_keys(self) -> str: + """Get element keys. + + Returns + ------- + str + + """ + return "element" + + @property + def datadir(self) -> Path: + """Get data directory path. + + Returns + ------- + pathlib.Path + Path to the data directory. Can be overridden by subclasses. + + """ + return self._datadir + @abstractmethod - def get_elements(self): + def get_elements(self) -> List: """Get elements. Returns @@ -149,332 +153,6 @@ class BaseDataGrabber(ABC): """ raise_error( - 'get_elements not implemented', - NotImplementedError) # type: ignore - - def __enter__(self): - """Context entry implementation.""" - return self - - def __exit__(self, exc_type, exc_value, exc_traceback): - """Context exit implementation.""" - pass - - -@register_datagrabber -class PatternDataGrabber(BaseDataGrabber): - """Pattern DataGrabber class (abstract). - - Implements a DataGrabber that understands patterns to grab data. - - Attributes - ---------- - datadir: Path - Directory where the data is stored - types : list[str] - List of data types to be grabbed. - patterns : dict[str -> str] - Patterns for each type of data. - replacements: list[str] - Replacements in the patterns for each item in the `element` tuple. - - Methods - ------- - get_elements: list[str] - Returns a list of elements that can be grabbed. Each element is a - subject in the BIDS database. - __getitem__(str): dict[str -> dict] - Returns a dictionary of paths for each type of data required for the - specified element. Each occurrence of the string `{subject}` is - replaced by the indexed element - """ - - def __init__(self, types=None, patterns=None, replacements=None, **kwargs): - """Initialize a BaseDataGrabber object. - - Parameters - ---------- - types : list[str] - The types of data to be grabbed. - patterns : dict[str -> str] - Patterns for each type of data. The keys are the types and the - values are the patterns. Each occurrence of the string `{subject}` - in the pattern will be replaced by the indexed element. - datadir : str or Path - That directory where the data is/will be stored. - """ - _validate_patterns(types, patterns) - if not isinstance(replacements, list): - replacements = [replacements] - _validate_replacements(replacements, patterns) - super().__init__(types=types, **kwargs) - logger.debug('Initializing PatternDataGrabber') - logger.debug(f'\tpatterns = {patterns}') - logger.debug(f'\treplacements = {replacements}') - self.patterns = patterns - self.replacements = replacements - - def _replace_patterns_regex(self, pattern): - """Replace the patterns in the pattern with the named groups. - - It allows elements to be obtained from the filesystem. - - Parameters - ---------- - pattern : str - The pattern to be replaced. - - Returns - ------- - re_pattern : str - The regular expression with the named groups. - glob_pattern : str - The search pattern to be used with glob - - """ - re_pattern = pattern - glob_pattern = pattern - for t_r in self.replacements: - # Replace the first of each with a named group definition - re_pattern = re_pattern.replace( - f'{{{t_r}}}', f'(?P<{t_r}>.*)', 1) - - for t_r in self.replacements: - # Replace the second appearance of each with the named group - # back reference - re_pattern = re_pattern.replace(f'{{{t_r}}}', f'(?P={t_r})') - - for t_r in self.replacements: - glob_pattern = glob_pattern.replace(f'{{{t_r}}}', '*') - return re_pattern, glob_pattern - - def get_elements(self): - """Get the list of elements in the dataset. - - It will use regex to search for `replacements` in the `patterns` and - return the intersection of the results for each type. That is, build a - list of elements that have all the required types. - - Returns - ------- - elements : list - The list of elements in the dataset. - """ - elements = None - for t_type in self.types: - types_element = set() - t_pattern = self.patterns[t_type] # get the pattern - re_pattern, glob_pattern = self._replace_patterns_regex(t_pattern) - for fname in self.datadir.glob(glob_pattern): - suffix = fname.relative_to(self.datadir).as_posix() - m = re.match(re_pattern, suffix) - if m is not None: - t_element = tuple(m.group(k) for k in self.replacements) - if len(self.replacements) == 1: - t_element = t_element[0] - types_element.add(t_element) - if elements is None: - elements = types_element - else: - elements = elements.intersection(types_element) - - return list(elements) - - def _replace_patterns_glob(self, element, pattern): - """Replace patterns with the element so it can be globbed. - - Parameters - ---------- - element : tuple - The element to be used in the replacement. - pattern : str - The pattern to be replaced. - - Returns - ------- - str - The pattern with the element replaced. - """ - if len(element) != len(self.replacements): - raise_error( - f'The element length must be {len(self.replacements)}, ' - f'indicating {self.replacements}') - to_replace = dict(zip(self.replacements, element)) - return pattern.format(**to_replace) - - def __getitem__(self, element): - """Index one element in the database. - - Each occurrence of the strings in `replacements` is replaced by the - corresponding item in the element tuple. - - Parameters - ---------- - element : str or tuple - The element to be indexed. If one string is provided, it is - assumed to be a tuple with only one item. If a tuple is provided, - each item in the tuple is the value for the replacement string - specified in `replacements`. - - Returns - ------- - out : dict[str -> Path] - Dictionary of paths for each type of data required for the - specified element. - """ - out = super().__getitem__(element) - if not isinstance(element, tuple): - element = (element,) - for t_type in self.types: - t_pattern = self.patterns[t_type] # type: ignore - t_replace = self._replace_patterns_glob(element, t_pattern) - if '*' in t_replace: - t_matches = list(self.datadir.glob(t_replace)) - if len(t_matches) > 1: - raise_error( - f'More than one file matches for {element} / {t_type}:' - f' {t_matches}') - elif len(t_matches) == 0: - raise_error(f'No file matches for {element} / {t_type}') - t_out = t_matches[0] - else: - t_out = self.datadir / t_replace - out[t_type] = dict(path=t_out) - # Meta here is element and types - out['meta']['element'] = dict(zip(self.replacements, element)) - return out - - -@register_datagrabber -class DataladDataGrabber(BaseDataGrabber): - """Datalad DataGrabber class (abstract). - - Implements a DataGrabber that gets data from a datalad sibling. - - Attributes - ---------- - datadir - uri : str - URI of the datalad sibling. - - Methods - ------- - install: - Installs (clones) the datalad dataset into the datadir. This method - is called automatically when the datagrabber is used within a `with` - statement. - remove: - Remove the datalad dataset from the datadir. This method is called - automatically when the datagrabber is used within a `with` statement. - - Notes - ----- - By itself, this class is still abstract as the `__getitem__` method relies - on the parent class `BaseDataGrabber.__getitem__` which is not yet - implemented. This class is intended to be used as a superclass of a class - with multiple inheritance. See :class:`BIDSDataladDataGrabber` for a - concrete class implementation. - - """ - - def __init__(self, rootdir='.', datadir=None, uri=None, **kwargs): - """Initialize a DataladDataGrabber object. - - Parameters - ---------- - rootdir : str or Path - The path within the datalad dataset to the root directory. - datadir : str or Path - That directory where the datalad dataset will be cloned. If None, - (default), the datalad dataset will be cloned into a temporary - directory. - uri : str - URI of the datalad sibling. - """ - if uri is None: - raise_error('uri must be provided', ValueError) - if datadir is None: - logger.warning('datadir is None, creating a temporary directory') - datadir = tempfile.mkdtemp() - logger.info(f'datadir set to {datadir}') - super().__init__(datadir=datadir, **kwargs) - logger.debug('Initializing DataladDataGrabber') - logger.debug(f'\turi = {uri}') - logger.debug(f'\t_rootdir = {rootdir}') - self.uri = uri - self._rootdir = rootdir - - def __enter__(self): - """Context entry implementation.""" - self.install() - return self - - @property - def datadir(self): - """Get data directory path.""" - return super().datadir / self._rootdir - - def install(self): - """Install the datalad dataset into the datadir.""" - logger.debug(f'Installing dataset {self.uri} to {self._datadir}') - self._dataset = dl.install( # type: ignore - self._datadir, source=self.uri) - logger.debug('Dataset installed') - - def __exit__(self, exc_type, exc_value, exc_traceback): - """Context exit implementation.""" - logger.debug('Removing dataset') - self.remove() - logger.debug('Dataset removed') - - def remove(self): - """Remove the datalad dataset from the datadir.""" - self._dataset.remove(recursive=True) - - def _dataset_get(self, out): - for _, v in out.items(): - if 'path' in v: - logger.debug(f'Getting {v["path"]}') - self._dataset.get(v['path']) - logger.debug('Get done') - - # append the version of the dataset - out['meta']['datagrabber']['dataset_commit_id'] = \ - self._dataset.repo.get_hexsha( - self._dataset.repo.get_corresponding_branch()) - return out - - def __getitem__(self, element): - """Index one element in the Datalad database. - - It will first obtain the paths from the parent class and then - `datalad get` each of the files. - - This method only works with multiple inheritance. - - """ - out = super().__getitem__(element) - out = self._dataset_get(out) - return out - - -@register_datagrabber -class PatternDataladDataGrabber(DataladDataGrabber, PatternDataGrabber): - """Pattern-based Datalad DataGrabber class (abstract). - - Implements a DataGrabber that gets data from a datalad sibling, - interpreting patterns. - - - See Also - -------- - DataladDataGrabber - PatternDataGrabber - - """ - - def __init__(self, types=None, patterns=None, **kwargs): - """Initialize the class.""" - _validate_patterns(types, patterns) - super().__init__(types=types, patterns=patterns, **kwargs) - self.patterns = patterns + msg="Concrete classes need to implement get_elements().", + klass=NotImplementedError, + ) -- 2.52.0 From ae189e497380fe60ff28ab61a076bef4f6929c02 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 16 Aug 2022 16:45:49 +0200 Subject: [PATCH 196/287] refactor: move utility functions from datagrabber/base.py to separate module --- junifer/datagrabber/utils.py | 77 ++++++++++++++++++++++++++++++++++++ 1 file changed, 77 insertions(+) create mode 100644 junifer/datagrabber/utils.py diff --git a/junifer/datagrabber/utils.py b/junifer/datagrabber/utils.py new file mode 100644 index 000000000..d30c11c13 --- /dev/null +++ b/junifer/datagrabber/utils.py @@ -0,0 +1,77 @@ +"""Provide utility functions for the datagrabber sub-package.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from typing import Dict, List + +from ..utils import raise_error, warn_with_log + + +def validate_types(types: List[str]) -> None: + """Validate the types. + + Parameters + ---------- + types : list of str + The object to validate. + + """ + if not isinstance(types, list): + raise_error(msg="`types` must be a list", klass=TypeError) + if any(not isinstance(x, str) for x in types): + raise_error(msg="`types` must be a list of strings", klass=TypeError) + + +def validate_replacements( + replacements: List[str], patterns: Dict[str, str] +) -> None: + """Validate the replacements. + + Parameters + ---------- + replacements : list of str + The object to validate. + patterns : dict + The patterns to validate against. + + """ + if not isinstance(replacements, list): + raise_error(msg="`replacements` must be a list.", klass=TypeError) + if any(not isinstance(x, str) for x in replacements): + raise_error( + msg="`replacements` must be a list of strings.", klass=TypeError + ) + + for x in replacements: + if all(x not in y for y in patterns.values()): + warn_with_log(msg=f"Replacement {x} is not part of any pattern.") + + +def validate_patterns(types: List[str], patterns: Dict[str, str]) -> None: + """Validate the patterns. + + Parameters + ---------- + types : list of str + The types list. + patterns : dict + The object to validate. + + """ + # Validate the types + validate_types(types) + if not isinstance(patterns, dict): + raise_error(msg="`patterns` must be a dict.", klass=TypeError) + # Unequal length of objects + if len(types) != len(patterns): + raise_error( + msg="`types` and `patterns` must have the same length.", + klass=ValueError, + ) + + if any(x not in patterns for x in types): + raise_error( + msg="`patterns` must contain all `types`", klass=ValueError + ) -- 2.52.0 From 7811060833c91a4654b609caec30250cbe838cda Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 16 Aug 2022 16:46:24 +0200 Subject: [PATCH 197/287] refactor: move PatternDataGrabber to separate module --- junifer/datagrabber/pattern.py | 210 +++++++++++++++++++++++++++++++++ 1 file changed, 210 insertions(+) create mode 100644 junifer/datagrabber/pattern.py diff --git a/junifer/datagrabber/pattern.py b/junifer/datagrabber/pattern.py new file mode 100644 index 000000000..137b535f0 --- /dev/null +++ b/junifer/datagrabber/pattern.py @@ -0,0 +1,210 @@ +"""Provide concrete implementation for pattern-based datagrabber.""" + +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + +import re +from pathlib import Path +from typing import Dict, List, Optional, Tuple, Union + +from ..api.decorators import register_datagrabber +from ..utils import logger, raise_error +from .base import BaseDataGrabber +from .utils import validate_patterns, validate_replacements + + +@register_datagrabber +class PatternDataGrabber(BaseDataGrabber): + """Concrete implementation for data fetching using patterns. + + Implements a DataGrabber that understands patterns to grab data. + + Parameters + ---------- + types : list of str, optional + The types of data to be grabbed (default None). + patterns : dict, optional + Patterns for each type of data as a dictionary. The keys are the types + and the values are the patterns. Each occurrence of the string + `{subject}` in the pattern will be replaced by the indexed element + (default None). + replacements: list of str + Replacements in the patterns for each item in the "element" tuple. + datadir : str or pathlib.Path + The directory where the data is / will be stored. + **kwargs + Keyword arguments passed to superclass. + + See Also + -------- + BaseDataGrabber + + """ + + def __init__( + self, + types: Optional[List[str]] = None, + patterns: Optional[Dict[str, str]] = None, + replacements: Optional[List[str]] = None, + **kwargs, + ) -> None: + """Initialize the class.""" + # Validate patterns + validate_patterns(types=types, patterns=patterns) + + if not isinstance(replacements, list): + replacements = [replacements] + # Validate replacements + validate_replacements(replacements=replacements, patterns=patterns) + + super().__init__(types=types, **kwargs) + logger.debug("Initializing PatternDataGrabber") + logger.debug(f"\tpatterns = {patterns}") + logger.debug(f"\treplacements = {replacements}") + self.patterns = patterns + self.replacements = replacements + + def _replace_patterns_regex(self, pattern: str) -> Tuple[str, str]: + """Replace the patterns in `pattern` with the named groups. + + It allows elements to be obtained from the filesystem. + + Parameters + ---------- + pattern : str + The pattern to be replaced. + + Returns + ------- + re_pattern : str + The regular expression with the named groups. + glob_pattern : str + The search pattern to be used with glob. + + """ + # re_pattern = pattern + # glob_pattern = pattern + for t_r in self.replacements: + # Replace the first appearance of each with a named group + # definition and the second appearance of each with the named group + # back reference + re_pattern = pattern.replace( + f"{{{t_r}}}", f"(?P<{t_r}>.*)", 1 + ).replace(f"{{{t_r}}}", f"(?P={t_r})") + glob_pattern = pattern.replace(f"{{{t_r}}}", "*") + + # for t_r in self.replacements: + # # Replace the second appearance of each with the named group + # # back reference + # re_pattern = pattern.replace(f"{{{t_r}}}", f"(?P={t_r})") + + # for t_r in self.replacements: + # glob_pattern = glob_pattern.replace(f"{{{t_r}}}", "*") + return re_pattern, glob_pattern + + def _replace_patterns_glob(self, element: Tuple, pattern: str) -> str: + """Replace patterns with the element so it can be globbed. + + Parameters + ---------- + element : tuple + The element to be used in the replacement. + pattern : str + The pattern to be replaced. + + Returns + ------- + str + The pattern with the element replaced. + + """ + if len(element) != len(self.replacements): + raise_error( + f"The element length must be {len(self.replacements)}, " + f"indicating {self.replacements}." + ) + to_replace = dict(zip(self.replacements, element)) + return pattern.format(**to_replace) + + def __getitem__(self, element: Union[str, Tuple]) -> Dict[str, Path]: + """Implement single element indexing in the database. + + Each occurrence of the strings in "replacements" is replaced by the + corresponding item in the element tuple. + + Parameters + ---------- + element : str or tuple + The element to be indexed. If one string is provided, it is + assumed to be a tuple with only one item. If a tuple is provided, + each item in the tuple is the value for the replacement string + specified in "replacements". + + Returns + ------- + dict + Dictionary of paths for each type of data required for the + specified element. + + """ + out = super().__getitem__(element) + if not isinstance(element, tuple): + element = (element,) + for t_type in self.types: + t_pattern = self.patterns[t_type] + t_replace = self._replace_patterns_glob(element, t_pattern) + if "*" in t_replace: + t_matches = list(self.datadir.glob(t_replace)) + if len(t_matches) > 1: + raise_error( + f"More than one file matches for {element} / {t_type}:" + f" {t_matches}" + ) + elif len(t_matches) == 0: + raise_error(f"No file matches for {element} / {t_type}") + t_out = t_matches[0] + else: + t_out = self.datadir / t_replace + out[t_type] = {"path": t_out} + # Meta here is element and types + out["meta"]["element"] = dict(zip(self.replacements, element)) + return out + + def get_elements(self) -> List: + """Implement fetching list of elements in the dataset. + + It will use regex to search for "replacements" in the "patterns" and + return the intersection of the results for each type i.e., build a + list of elements that have all the required types. + + Returns + ------- + elements : list + The list of elements that can be grabbed in the dataset. Each + element is a subject in the BIDS database. + + """ + elements = None + for t_type in self.types: + types_element = set() + # Get the pattern + t_pattern = self.patterns[t_type] + # Replace the pattern + re_pattern, glob_pattern = self._replace_patterns_regex(t_pattern) + for fname in self.datadir.glob(glob_pattern): + suffix = str(fname.relative_to(self.datadir).absolute()) + m = re.match(re_pattern, suffix) + if m is not None: + t_element = tuple(m.group(k) for k in self.replacements) + if len(self.replacements) == 1: + t_element = t_element[0] + types_element.add(t_element) + # TODO: does this make sense as elements is always None + if elements is None: + elements = types_element + else: + elements = elements.intersection(types_element) + + return list(elements) -- 2.52.0 From bb0630703a44f56458eadaf16ef59b947e58a9db Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 16 Aug 2022 16:46:56 +0200 Subject: [PATCH 198/287] refactor: move DataladDataGrabber to separate module --- junifer/datagrabber/datalad_base.py | 166 ++++++++++++++++++++++++++++ 1 file changed, 166 insertions(+) create mode 100644 junifer/datagrabber/datalad_base.py diff --git a/junifer/datagrabber/datalad_base.py b/junifer/datagrabber/datalad_base.py new file mode 100644 index 000000000..3bee2a0d4 --- /dev/null +++ b/junifer/datagrabber/datalad_base.py @@ -0,0 +1,166 @@ +"""Provide abstract base class for datalad datagrabber.""" + +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + +import tempfile +from pathlib import Path +from typing import Dict, Optional, Tuple, Union + +import datalad.api as dl + +from ..api.decorators import register_datagrabber +from ..utils import logger +from .base import BaseDataGrabber +from .utils import raise_error + + +@register_datagrabber +class DataladDataGrabber(BaseDataGrabber): + """Abstract base class for data fetching via Datalad. + + Defines a DataGrabber that gets data from a datalad sibling. + + Parameters + ---------- + rootdir : str or Path, optional + The path within the datalad dataset to the root directory + (default "."). + datadir : str or Path, optional + That directory where the datalad dataset will be cloned. If None, + the datalad dataset will be cloned into a temporary directory + (default None). + uri : str, optional + URI of the datalad sibling (default None). + **kwargs + Keyword arguments passed to superclass. + + Methods + ------- + install + Installs (clones) the datalad dataset into the `datadir`. This method + is called automatically when the datagrabber is used within a context. + remove + Removes the datalad dataset from the `datadir`. This method is called + automatically when the datagrabber is used within a context. + + Notes + ----- + By itself, this class is still abstract as the `__getitem__` method relies + on the parent class `BaseDataGrabber.__getitem__` which is not yet + implemented. This class is intended to be used as a superclass of a class + with multiple inheritance. + + See Also + -------- + BaseDataGrabber + BIDSDataladDataGrabber + + """ + + def __init__( + self, + rootdir: Union[str, Path] = ".", + datadir: Union[str, Path, None] = None, + uri: Optional[str] = None, + **kwargs, + ): + """Initialize the class.""" + if datadir is None: + logger.warning("`datadir` is None, creating a temporary directory") + # Create temporary directory + datadir = tempfile.mkdtemp() + logger.info(f"`datadir` set to {datadir}") + # TODO: uri can be converted to a positional argument + if uri is None: + raise_error("`uri` must be provided") + + super().__init__(datadir=datadir, **kwargs) + logger.debug("Initializing DataladDataGrabber") + logger.debug(f"\turi = {uri}") + logger.debug(f"\t_rootdir = {rootdir}") + self.uri = uri + self._rootdir = rootdir + + @property + def datadir(self) -> Path: + """Get data directory path.""" + return super().datadir / self._rootdir + + def _dataset_get(self, out: Dict) -> Dict: + """Get the dataset found from the path in `out`. + + Parameters + ---------- + out : dict + The dictionary from which path need to be searched. + + Returns + ------- + dict + The modified dictionary with version appended. + + """ + for _, v in out.items(): + if "path" in v: + logger.debug(f"Getting {v['path']}") + self._dataset.get(v["path"]) + logger.debug("Get done") + + # append the version of the dataset + out["meta"]["datagrabber"][ + "dataset_commit_id" + ] = self._dataset.repo.get_hexsha( + self._dataset.repo.get_corresponding_branch() + ) + return out + + def install(self) -> None: + """Install the datalad dataset into the datadir.""" + logger.debug(f"Installing dataset {self.uri} to {self._datadir}") + self._dataset = dl.install(self._datadir, source=self.uri) + logger.debug("Dataset installed") + + def remove(self): + """Remove the datalad dataset from the datadir.""" + self._dataset.remove(recursive=True) + + def __getitem__(self, element: Union[str, Tuple]) -> Dict[str, Path]: + """Implement single element indexing in the Datalad database. + + It will first obtain the paths from the parent class and then + `datalad get` each of the files. + + This method only works with multiple inheritance. + + Parameters + ---------- + element : str or tuple + The element to be indexed. If one string is provided, it is + assumed to be a tuple with only one item. If a tuple is provided, + each item in the tuple is the value for the replacement string + specified in "replacements". + + Returns + ------- + dict + Dictionary of paths for each type of data required for the + specified element. + + """ + out = super().__getitem__(element) + out = self._dataset_get(out) + return out + + def __enter__(self): + """Implement context entry.""" + self.install() + return self + + def __exit__(self, exc_type, exc_value, exc_traceback): + """Implement context exit.""" + logger.debug("Removing dataset") + self.remove() + logger.debug("Dataset removed") -- 2.52.0 From 6aaa80862d38f09c4115021c69f8a2b927411fe0 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 16 Aug 2022 16:47:24 +0200 Subject: [PATCH 199/287] refactor: move PatternDataladDataGrabber to separate module --- junifer/datagrabber/pattern_datalad_base.py | 53 +++++++++++++++++++++ 1 file changed, 53 insertions(+) create mode 100644 junifer/datagrabber/pattern_datalad_base.py diff --git a/junifer/datagrabber/pattern_datalad_base.py b/junifer/datagrabber/pattern_datalad_base.py new file mode 100644 index 000000000..b36aed865 --- /dev/null +++ b/junifer/datagrabber/pattern_datalad_base.py @@ -0,0 +1,53 @@ +"""Provide abstract base class for pattern-based datalad datagrabber.""" + +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + +from typing import Dict, List, Optional + +from ..api.decorators import register_datagrabber +from .datalad_base import DataladDataGrabber +from .pattern import PatternDataGrabber +from .utils import validate_patterns + + +@register_datagrabber +class PatternDataladDataGrabber(DataladDataGrabber, PatternDataGrabber): + """Abstract base class for pattern-based data fetching via Datalad. + + Defines a DataGrabber that gets data from a datalad sibling, + interpreting patterns. + + Parameters + ---------- + types : list of str, optional + The types of data to be grabbed (default None). + patterns : dict, optional + Patterns for each type of data as a dictionary. The keys are the types + and the values are the patterns. Each occurrence of the string + `{subject}` in the pattern will be replaced by the indexed element + (default None). + **kwargs + Keyword arguments passed to superclass. + + See Also + -------- + DataladDataGrabber + PatternDataGrabber + + """ + + def __init__( + self, + types: Optional[List[str]] = None, + patterns: Optional[Dict[str, str]] = None, + **kwargs, + ) -> None: + """Initialize the class.""" + # Validate patterns + validate_patterns(types=types, patterns=patterns) + + super().__init__(types=types, patterns=patterns, **kwargs) + self.patterns = patterns -- 2.52.0 From 132c65d793df5164d63f0c7169cfc22657a91398 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 16 Aug 2022 16:48:20 +0200 Subject: [PATCH 200/287] refactor: add docstrings and type annotations in datagrabber/hcp.py --- junifer/datagrabber/hcp.py | 262 ++++++++++++++++++++++--------------- 1 file changed, 155 insertions(+), 107 deletions(-) diff --git a/junifer/datagrabber/hcp.py b/junifer/datagrabber/hcp.py index dd9b801f2..1ebfc0a73 100644 --- a/junifer/datagrabber/hcp.py +++ b/junifer/datagrabber/hcp.py @@ -1,98 +1,151 @@ -"""Provide classes for HCP data access.""" +"""Provide concrete implementations for HCP data access.""" from itertools import product +from pathlib import Path +from typing import Dict, List, Tuple, Union -from junifer.datagrabber.base import DataladDataGrabber +from junifer.datagrabber.datalad_base import DataladDataGrabber -from ..datagrabber import PatternDataGrabber from ..api.decorators import register_datagrabber +from ..utils import raise_error +from .pattern import PatternDataGrabber @register_datagrabber class HCP1200(PatternDataGrabber): - """PatternDataGrabber implementation for HCP1200.""" + """Concrete implementation for pattern-based data fetching of HCP1200. + + Parameters + ---------- + datadir : str or Path, optional + The directory where the datalad dataset will be cloned. If None, + the datalad dataset will be cloned into a temporary directory + (default None). + tasks : {"REST1", "REST2", "SOCIAL", "WM", "RELATIONAL", "EMOTION", + "LANGUAGE", "GAMBLING", "MOTOR"} or list of the options, optional + HCP task sessions. If None, all available task sessions are selected + (default None). + phase_encodings : {"LR", "RL"} or list of the options, optional + HCP phase encoding directions. If None, both will be used + (default None). + **kwargs + Keyword arguments passed to superclass. + + """ def __init__( - self, datadir=None, tasks=None, phase_encodings=None - ): - """Initialize a HCP object. + self, + datadir: Union[str, Path, None] = None, + tasks: Union[str, List[str], None] = None, + phase_encodings: Union[str, List[str], None] = None, + **kwargs, + ) -> None: + """Initialize the class.""" + # All tasks + all_tasks = [ + "REST1", + "REST2", + "SOCIAL", + "WM", + "RELATIONAL", + "EMOTION", + "LANGUAGE", + "GAMBLING", + "MOTOR", + ] + # Set default tasks + if tasks is None: + self.tasks = all_tasks + # Convert single task into list + if isinstance(tasks, str): + tasks = [tasks] + # Check for invalid task(s) + for task in tasks: + if task not in all_tasks: + raise_error( + f"'{task}' is not a valid HCP-YA fMRI task input. " + f"Valid task values can be any or all of {all_tasks}." + ) - Parameters - ---------- - datadir : str or Path - That directory where the datalad dataset will be cloned. If None, - (default), the datalad dataset will be cloned into a temporary - directory. - tasks : str or list of strings - HCP task sessions. If 'None' (default), all available task - sessions are selected. Can be 'REST1', 'REST2', 'SOCIAL', 'WM', - 'RELATIONAL', 'EMOTION', 'LANGUAGE', 'GAMBLING', 'MOTOR', or a - list consisting of these names. - phase_encoding : str or list of strings - HCP phase encoding directions. Can be 'LR' or 'RL'. If 'None' - (default) both will be used. + # All phase encodings + all_phase_encodings = ["LR", "RL"] + # Set phase encodings + if phase_encodings is None: + phase_encodings = all_phase_encodings + # Convert single phase encoding into list + if isinstance(phase_encodings, str): + phase_encodings = [phase_encodings] + # Check for invalid phase encoding(s) + for pe in phase_encodings: + if pe not in all_phase_encodings: + raise_error( + f"'{pe}' is not a valid HCP-YA phase encoding. " + "Valid phase encoding can be any or all of " + f"{all_phase_encodings}." + ) - """ - types = ['BOLD'] - - replacements = ['subject', 'task', 'phase_encoding'] + # The types of data + types = ["BOLD"] + # The patterns patterns = { - 'BOLD': ('{subject}/MNINonLinear/Results/' - '{task}_{phase_encoding}/' - '{task}_{phase_encoding}_hp2000_clean.nii.gz') + "BOLD": ( + "{subject}/MNINonLinear/Results/" + "{task}_{phase_encoding}/" + "{task}_{phase_encoding}_hp2000_clean.nii.gz" + ) } + # The replacements + replacements = ["subject", "task", "phase_encoding"] super().__init__( - types=types, datadir=datadir, patterns=patterns, - replacements=replacements + types=types, + datadir=datadir, + patterns=patterns, + replacements=replacements, ) - self.tasks = tasks self.phase_encodings = phase_encodings - if isinstance(self.tasks, str): - self.tasks = [self.tasks] - if isinstance(self.phase_encodings, str): - self.phase_encodings = [self.phase_encodings] + def __getitem__(self, element: Tuple[str, str, str]) -> Dict[str, Path]: + """Index one element in the dataset. - all_tasks = [ - 'REST1', - 'REST2', - 'SOCIAL', - 'WM', - 'RELATIONAL', - 'EMOTION', - 'LANGUAGE', - 'GAMBLING', - 'MOTOR', - ] - - if self.tasks is None: - self.tasks = all_tasks - - if self.phase_encodings is None: - self.phase_encodings = ['LR', 'RL'] - - for task in self.tasks: - if task not in all_tasks: - raise ValueError( - f'{task} not a valid HCP-YA fMRI task input! \n' - f'task can be any of {all_tasks}' - ) - - for pe in self.phase_encodings: - if pe not in ['LR', 'RL']: - raise ValueError( - f'{pe} not a valid HCP-YA phase encoding. \n' - 'phase_encoding can be LR or RL (or both)!' - ) - - def get_elements(self): - """Get the list of subjects in the dataset. + Parameters + ---------- + element : triple of str + The element to be indexed. First element in the tuple is the + subject, second element is the task, third element is the + phase encoding direction. Returns ------- - elements : list[str] + out : dict + Dictionary of paths for each type of data required for the + specified element. + + """ + sub, task, phase_encoding = element + + # Resting task + if "REST" in task: + new_task = f"rfMRI_{task}" + else: + new_task = f"tfMRI_{task}" + + out = super().__getitem__((sub, new_task, phase_encoding)) + out["meta"]["element"] = { + "subject": sub, + "task": task, + "phase_encoding": phase_encoding, + } + return out + + def get_elements(self) -> List: + """Implement fetching list of subjects in the dataset. + + Returns + ------- + elements : list of str The list of subjects in the dataset. + """ subjects = [x.name for x in self.datadir.iterdir() if x.is_dir()] elems = [] @@ -103,51 +156,46 @@ class HCP1200(PatternDataGrabber): return elems - def __getitem__(self, element): - """Index one element in the dataset. - - Parameters - ---------- - element : tuple[str, str] - The element to be indexed. First element in the tuple is the - subject, second element is the task, third element is the - phase encoding direction. - - Returns - ------- - out : dict[str -> Path] - Dictionary of paths for each type of data required for the - specified element. - """ - sub, task, phase_encoding = element - - if "REST" in task: - new_task = f'rfMRI_{task}' - else: - new_task = f'tfMRI_{task}' - - out = super().__getitem__((sub, new_task, phase_encoding)) - out['meta']['element'] = dict( - subject=sub, task=task, phase_encoding=phase_encoding - ) - - return out - @register_datagrabber class DataladHCP1200(DataladDataGrabber, HCP1200): - """DataladDataGrabber implementation for HCP1200.""" + """Concrete implementation for datalad-based data fetching of HCP1200. - def __init__(self, datadir=None, tasks=None, phase_encodings=None): + Parameters + ---------- + datadir : str or Path, optional + The directory where the datalad dataset will be cloned. If None, + the datalad dataset will be cloned into a temporary directory + (default None). + tasks : {"REST1", "REST2", "SOCIAL", "WM", "RELATIONAL", "EMOTION", + "LANGUAGE", "GAMBLING", "MOTOR"} or list of the options, optional + HCP task sessions. If None, all available task sessions are selected + (default None). + phase_encodings : {"LR", "RL"} or list of the options, optional + HCP phase encoding directions. If None, both will be used + (default None). + **kwargs + Keyword arguments passed to superclass. + + """ + + def __init__( + self, + datadir: Union[str, Path, None] = None, + tasks: Union[str, List[str], None] = None, + phase_encodings: Union[str, List[str], None] = None, + **kwargs, + ) -> None: """Initialize the class.""" uri = ( - 'https://github.com/datalad-datasets/' - 'human-connectome-project-openaccess.git' + "https://github.com/datalad-datasets/" + "human-connectome-project-openaccess.git" ) - rootdir = 'HCP1200' + rootdir = "HCP1200" super().__init__( datadir=datadir, tasks=tasks, phase_encodings=phase_encodings, - uri=uri, rootdir=rootdir - ) # type: ignore + uri=uri, + rootdir=rootdir, + ) -- 2.52.0 From 789b20fd8f247c685cd003c939fac486bba01c6e Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 16 Aug 2022 16:48:46 +0200 Subject: [PATCH 201/287] refactor: add docstrings and type annotations in datagrabber/meta.py --- junifer/datagrabber/meta.py | 97 ++++++++++++++++++++++++++++--------- 1 file changed, 74 insertions(+), 23 deletions(-) diff --git a/junifer/datagrabber/meta.py b/junifer/datagrabber/meta.py index 100475b65..7fef19d60 100644 --- a/junifer/datagrabber/meta.py +++ b/junifer/datagrabber/meta.py @@ -1,42 +1,93 @@ -"""Provide class for metadata collection.""" +"""Provide abstract base class for metadata collection.""" + +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + +from pathlib import Path +from typing import Dict, List, Tuple, Union from .base import BaseDataGrabber class MultipleDataGrabber(BaseDataGrabber): - """MultipleDataGrabber implementation.""" + """Abstract base class for metadata fetch. - def __init__(self, datagrabbers): + Defines a DataGrabber which can be used to fetch metadata from multiple + sources. + + Parameters + ---------- + datagrabbers : list of datagrabbers + The datagrabbers to use to fetch metadata using. + **kwargs + Keyword arguments passed to superclass. + + """ + + def __init__(self, datagrabbers: List[BaseDataGrabber], **kwargs) -> None: """Initialize the class.""" # TODO: Check datagrabbers consistency # - same element keys # - no overlapping types self._datagrabbers = datagrabbers - def get_types(self): - """Get types.""" - types = [x for dg in self._datagrabbers for x in dg.get_types()] - return types + def __getitem__(self, element: Union[str, Tuple]) -> Dict[str, Path]: + """Implement indexing. - def get_meta(self): - """Get metadata.""" - t_meta = {} - t_meta['class'] = self.__class__.__name__ - t_meta['datagrabbers'] = [dg.get_meta() for dg in self._datagrabbers] + Parameters + ---------- + element : str or tuple + The element to be indexed. If one string is provided, it is + assumed to be a tuple with only one item. If a tuple is provided, + each item in the tuple is the value for the replacement string + specified in "replacements". - def __enter__(self): - """Context entry implementation.""" - for dg in self._datagrabbers: - dg.__enter__() + Returns + ------- + dict + Dictionary of paths for each type of data required for the + specified element. - def __exit__(self, exc_type, exc_value, exc_traceback): - """Context exit implementation.""" - for dg in self._datagrabbers: - dg.__exit__(exc_type, exc_value, exc_traceback) - - def __getitem__(self, element): - """Get item implementation.""" + """ out = {} for dg in self._datagrabbers: t_out = dg[element] out.update(t_out) + return out + + def __enter__(self) -> "BaseDataGrabber": + """Implement context entry.""" + for dg in self._datagrabbers: + dg.__enter__() + + def __exit__(self, exc_type, exc_value, exc_traceback) -> None: + """Implement context exit.""" + for dg in self._datagrabbers: + dg.__exit__(exc_type, exc_value, exc_traceback) + + def get_types(self) -> List[List[str]]: + """Get types. + + Returns + ------- + list of list of str + The types of data to be grabbed. + + """ + types = [x for dg in self._datagrabbers for x in dg.get_types()] + return types + + def get_meta(self) -> Dict: + """Get metadata. + + Returns + ------- + dict + The metadata as dictionary. + + """ + t_meta = {} + t_meta["class"] = self.__class__.__name__ + t_meta["datagrabbers"] = [dg.get_meta() for dg in self._datagrabbers] -- 2.52.0 From 672619c82a80d3483e4a40d25a065b1a4ea20f6b Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 16 Aug 2022 16:50:47 +0200 Subject: [PATCH 202/287] refactor: docstrings, type annotations, move unit tests for BaseDataGrabber to separate module --- junifer/datagrabber/tests/test_base.py | 43 ++++++++++++++++++++++++++ 1 file changed, 43 insertions(+) create mode 100644 junifer/datagrabber/tests/test_base.py diff --git a/junifer/datagrabber/tests/test_base.py b/junifer/datagrabber/tests/test_base.py new file mode 100644 index 000000000..ba0caae61 --- /dev/null +++ b/junifer/datagrabber/tests/test_base.py @@ -0,0 +1,43 @@ +"""Provide tests for base.""" + +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + +from pathlib import Path + +import pytest + +from junifer.datagrabber.base import BaseDataGrabber + + +def test_BaseDataGrabber_abstractness() -> None: + """Test BaseDataGrabber is abstract base class.""" + with pytest.raises(TypeError, match=r"abstract"): + BaseDataGrabber(datadir="/tmp", types=["func"]) + + +def test_BaseDataGrabber() -> None: + """Test BaseDataGrabber.""" + # Create concrete class. + class MyDataGrabber(BaseDataGrabber): + def __getitem__(self, element): + return super().__getitem__(element) + + def get_elements(self): + return super().get_elements() + + dg = MyDataGrabber(datadir="/tmp", types=["func"]) + elem = dg["elem"] + assert "meta" in elem + assert "datagrabber" in elem["meta"] + assert "class" in elem["meta"]["datagrabber"] + assert MyDataGrabber.__name__ in elem["meta"]["datagrabber"]["class"] + + with pytest.raises(NotImplementedError): + dg.get_elements() + + with dg: + assert dg.datadir == Path("/tmp") + assert dg.types == ["func"] -- 2.52.0 From 2abdaaa61e7e65170e56aaa40a37b932ff93f730 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 16 Aug 2022 16:52:11 +0200 Subject: [PATCH 203/287] refactor: move PatternDataGrabber unit tests to separate module --- junifer/datagrabber/tests/test_pattern.py | 114 ++++++++++++++++++++++ 1 file changed, 114 insertions(+) create mode 100644 junifer/datagrabber/tests/test_pattern.py diff --git a/junifer/datagrabber/tests/test_pattern.py b/junifer/datagrabber/tests/test_pattern.py new file mode 100644 index 000000000..32eb09757 --- /dev/null +++ b/junifer/datagrabber/tests/test_pattern.py @@ -0,0 +1,114 @@ +"""Provide tests for pattern.""" + +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + +from junifer.datagrabber.pattern import PatternDataGrabber + +from pathlib import Path + +import pytest + + +def test_PatternDataGrabber() -> None: + """Test PatternDataGrabber.""" + # Create concrete class + class MyDataGrabber(PatternDataGrabber): + def get_elements(self): + return super().get_elements() + + with pytest.raises(TypeError, match=r"`types` must be a list"): + MyDataGrabber( + datadir="/tmp", + types="wrong", + patterns={"wrong": "pattern"}, + replacements="subject", + ) + + with pytest.raises(TypeError, match=r"`types` must be a list of strings"): + MyDataGrabber( + datadir="/tmp", + types=[1, 2, 3], + patterns={"1": "pattern", "2": "pattern", "3": "pattern"}, + replacements="subject", + ) + + with pytest.raises(ValueError, match=r"must have the same length"): + MyDataGrabber( + datadir="/tmp", + types=["func", "anat"], + patterns={"1": "pattern", "2": "pattern", "3": "pattern"}, + replacements=1, + ) + + with pytest.raises(TypeError, match=r"`patterns` must be a dict"): + MyDataGrabber( + datadir="/tmp", + types=["func", "anat"], + patterns="wrong", + replacements="subject", + ) + + with pytest.raises( + ValueError, match=r"`patterns` must have the same length" + ): + MyDataGrabber( + datadir="/tmp", + types=["func", "anat"], + patterns={"wrong": "pattern"}, + replacements="subject", + ) + + with pytest.raises( + ValueError, match=r"`patterns` must contain all `types`" + ): + MyDataGrabber( + datadir="/tmp", + types=["func", "anat"], + patterns={"wrong": "pattern", "func": "pattern"}, + replacements="subject", + ) + + with pytest.raises(TypeError, match=r"must be a list of strings"): + MyDataGrabber( + datadir="/tmp", + types=["func", "anat"], + patterns={"func": "func/test", "anat": "anat/test"}, + replacements=1, + ) + + with pytest.warns(RuntimeWarning, match=r"not part of any pattern"): + MyDataGrabber( + datadir="/tmp", + types=["func", "anat"], + patterns={ + "func": "func/{subject}.nii", + "anat": "anat/{subject}.nii", + }, + replacements=["subject", "wrong"], + ) + + datagrabber = MyDataGrabber( + datadir="/tmp/data", + types=["func", "anat"], + patterns={"func": "func/{subject}.nii", "anat": "anat/{subject}.nii"}, + replacements="subject", + ) + assert datagrabber.datadir == Path("/tmp/data") + assert datagrabber.types == ["func", "anat"] + assert datagrabber.replacements == ["subject"] + + datagrabber = MyDataGrabber( + datadir=Path("/tmp/data"), + types=["func", "anat"], + patterns={ + "func": "func/{subject}.nii", + "anat": "anat/{subject}_{session}.nii", + }, + replacements=["subject", "session"], + ) + assert datagrabber.datadir == Path("/tmp/data") + assert datagrabber.types == ["func", "anat"] + assert datagrabber.replacements == ["subject", "session"] -- 2.52.0 From d27f64e9141b810d4264b7a9cfc9313749452dcb Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 16 Aug 2022 17:06:48 +0200 Subject: [PATCH 204/287] refactor: move DataladDataGrabber to separate module --- .../datagrabber/tests/test_datalad_base.py | 21 +++++++++++++++++++ 1 file changed, 21 insertions(+) create mode 100644 junifer/datagrabber/tests/test_datalad_base.py diff --git a/junifer/datagrabber/tests/test_datalad_base.py b/junifer/datagrabber/tests/test_datalad_base.py new file mode 100644 index 000000000..b05ff4337 --- /dev/null +++ b/junifer/datagrabber/tests/test_datalad_base.py @@ -0,0 +1,21 @@ +"""Provide tests for datalad_base.""" + +# Authors: Synchon Mandal +# License: AGPL + +import pytest + +from junifer.datagrabber.datalad_base import DataladDataGrabber + + +def test_datalad_base_abstractness() -> None: + """Test datalad base is abstract.""" + with pytest.raises(TypeError, match=r"abstract"): + DataladDataGrabber() + +# def test_datalad_base_missing_uri() -> None: +# """Test proper check of missing URI in datalad base initialization.""" +# with pytest.raises(ValueError, match=r"`uri` must be provided"): +# DataladDataGrabber( + +# ) -- 2.52.0 From dc1af365792207e2deae6fa86b872ec4325f3ac5 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 16 Aug 2022 17:07:15 +0200 Subject: [PATCH 205/287] refactor: move PatternDataladDataGrabber to separate module --- .../tests/test_pattern_datalad_base.py | 178 ++++++++++++++++++ 1 file changed, 178 insertions(+) create mode 100644 junifer/datagrabber/tests/test_pattern_datalad_base.py diff --git a/junifer/datagrabber/tests/test_pattern_datalad_base.py b/junifer/datagrabber/tests/test_pattern_datalad_base.py new file mode 100644 index 000000000..3def6d295 --- /dev/null +++ b/junifer/datagrabber/tests/test_pattern_datalad_base.py @@ -0,0 +1,178 @@ +"""Provide tests for pattern_datalad_base.""" + +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + +from pathlib import Path + +import pytest + +from junifer.datagrabber.pattern_datalad_base import PatternDataladDataGrabber + + +_testing_dataset = { + "example_bids": { + "uri": "https://gin.g-node.org/juaml/datalad-example-bids", + "id": "e2ce149bd723088769a86c72e57eded009258c6b", + }, + "example_bids_ses": { + "uri": "https://gin.g-node.org/juaml/datalad-example-bids-ses", + "id": "3d08d55d1faad4f12ab64ac9497544a0d924d47a", + }, +} + + +def test_bids_pattern_datalad_datagrabber_missing_uri() -> None: + """Test check of missing URI in pattern datalad datagrabber.""" + with pytest.raises(ValueError, match=r"`uri` must be provided"): + PatternDataladDataGrabber( + datadir=None, + types=[], + patterns={}, + replacements=[], + ) + + +def test_bids_PatternDataladDataGrabber(tmp_path: Path) -> None: + """Test a subject-based BIDS datalad datagrabber. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + # Define types + types = ["T1w", "bold"] + # Define patterns + patterns = { + "T1w": "{subject}/anat/{subject}_T1w.nii.gz", + "bold": "{subject}/func/{subject}_task-rest_bold.nii.gz", + } + # Define replacements + replacements = ["subject"] + + repo_uri = _testing_dataset["example_bids"]["uri"] + rootdir = "example_bids" + repo_commit = _testing_dataset["example_bids"]["id"] + + with PatternDataladDataGrabber( + rootdir=rootdir, + uri=repo_uri, + types=types, + patterns=patterns, + replacements=replacements, + ) as dg: + subs = [x for x in dg] + expected_subs = [f"sub-{i:02d}" for i in range(1, 10)] + assert set(subs) == set(expected_subs) + + for elem in dg: + t_sub = dg[elem] + assert "path" in t_sub["T1w"] + assert t_sub["T1w"]["path"] == ( + dg.datadir / f"{elem}/anat/{elem}_T1w.nii.gz" + ) + assert "path" in t_sub["bold"] + assert t_sub["bold"]["path"] == ( + dg.datadir / f"{elem}/func/{elem}_task-rest_bold.nii.gz" + ) + + assert "meta" in t_sub + assert "datagrabber" in t_sub["meta"] + dg_meta = t_sub["meta"]["datagrabber"] + assert "class" in dg_meta + assert dg_meta["class"] == "PatternDataladDataGrabber" + assert "uri" in dg_meta + assert dg_meta["uri"] == repo_uri + assert "dataset_commit_id" in dg_meta + assert dg_meta["dataset_commit_id"] == repo_commit + + with open(t_sub["T1w"]["path"], "r") as f: + assert f.readlines()[0] == "placeholder" + + # datadir = tmp_path / "dataset" # Need this for testing + # patterns = { + # "T1w": "{subject}/anat/{subject}_T*w.nii.gz", + # "bold": "{subject}/func/{subject}_task-rest_*.nii.gz", + # } + # with PatternDataladDataGrabber( + # rootdir=rootdir, + # uri=repo_uri, + # types=types, + # patterns=patterns, + # datadir=datadir, + # replacements=replacements, + # ) as dg: + # assert dg.datadir == datadir / rootdir + # for elem in dg: + # t_sub = dg[elem] + # assert "path" in t_sub["T1w"] + # assert t_sub["T1w"]["path"] == ( + # dg.datadir / f"{elem}/anat/{elem}_T1w.nii.gz" + # ) + # assert "path" in t_sub["bold"] + # assert t_sub["bold"]["path"] == ( + # dg.datadir / f"{elem}/func/{elem}_task-rest_bold.nii.gz" + # ) + + +# def test_bids_datalad_PatternDataGrabber_session(): +# """Test a subject and session-based BIDS datalad datagrabber.""" +# types = ["T1w", "bold"] +# patterns = { +# "T1w": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz", +# "bold": "{subject}/{session}/func/" +# "{subject}_{session}_task-rest_bold.nii.gz", +# } +# replacements = ["subject", "session"] + +# with pytest.raises(ValueError, match=r"`uri` must be provided"): +# PatternDataladDataGrabber( +# datadir=None, +# types=types, +# patterns=patterns, +# replacements=replacements, +# ) + +# repo_uri = _testing_dataset["example_bids_ses"]["uri"] +# rootdir = "example_bids_ses" +# # repo_commit = _testing_dataset['example_bids_ses']['id'] + +# # With T1W and bold, only 2 sessions are available +# with PatternDataladDataGrabber( +# rootdir=rootdir, +# uri=repo_uri, +# types=types, +# patterns=patterns, +# replacements=replacements, +# ) as dg: +# subs = [x for x in dg] +# expected_subs = [ +# (f"sub-{i:02d}", f"ses-{j:02d}") +# for j in range(1, 3) +# for i in range(1, 10) +# ] +# assert set(subs) == set(expected_subs) + +# # Test with a different T1w only, it should have 3 sessions +# types = ["T1w"] +# patterns = { +# "T1w": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz", +# } +# with PatternDataladDataGrabber( +# rootdir=rootdir, +# uri=repo_uri, +# types=types, +# patterns=patterns, +# replacements=replacements, +# ) as dg: +# subs = [x for x in dg] +# expected_subs = [ +# (f"sub-{i:02d}", f"ses-{j:02d}") +# for j in range(1, 4) +# for i in range(1, 10) +# ] +# assert set(subs) == set(expected_subs) -- 2.52.0 From 3804bb736be34168e151c8c8b9e1cf9237a543c3 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 16 Aug 2022 17:07:36 +0200 Subject: [PATCH 206/287] chore: remove old aggregated unit tests for datagrabber sub-package --- .../tests/test_base_datagrabber.py | 227 ------------------ 1 file changed, 227 deletions(-) delete mode 100644 junifer/datagrabber/tests/test_base_datagrabber.py diff --git a/junifer/datagrabber/tests/test_base_datagrabber.py b/junifer/datagrabber/tests/test_base_datagrabber.py deleted file mode 100644 index 330687e0f..000000000 --- a/junifer/datagrabber/tests/test_base_datagrabber.py +++ /dev/null @@ -1,227 +0,0 @@ -"""Provide tests for base datagrabber.""" - -# Authors: Federico Raimondo -# Leonard Sasse -# License: AGPL - -import tempfile -import pytest -from pathlib import Path -from junifer.datagrabber.base import ( - PatternDataGrabber, - BaseDataGrabber, - PatternDataladDataGrabber, -) - - -_testing_dataset = { - 'example_bids': { - 'uri': 'https://gin.g-node.org/juaml/datalad-example-bids', - 'id': 'e2ce149bd723088769a86c72e57eded009258c6b' - }, - 'example_bids_ses': { - 'uri': 'https://gin.g-node.org/juaml/datalad-example-bids-ses', - 'id': '3d08d55d1faad4f12ab64ac9497544a0d924d47a' - } -} - - -def test_BaseDataGrabber(): - """Test BaseDataGrabber.""" - with pytest.raises(TypeError, match=r"abstract"): - BaseDataGrabber(datadir='/tmp', types=['func']) # type: ignore - - class MyDataGrabber(BaseDataGrabber): - def __getitem__(self, element): - return super().__getitem__(element) - - def get_elements(self): - return super().get_elements() - - dg = MyDataGrabber(datadir='/tmp', types=['func']) - elem = dg['elem'] - assert 'meta' in elem - assert 'datagrabber' in elem['meta'] - assert 'class' in elem['meta']['datagrabber'] - assert MyDataGrabber.__name__ in elem['meta']['datagrabber']['class'] - - with pytest.raises(NotImplementedError): - dg.get_elements() - - with dg: - assert dg.datadir == Path('/tmp') - assert dg.types == ['func'] - - -def test_PatternDataGrabber(): - """Test PatternDataGrabber.""" - class MyDataGrabber(PatternDataGrabber): - def get_elements(self): - return super().get_elements() - - with pytest.raises(TypeError, match=r"types must be a list"): - MyDataGrabber(datadir='/tmp', types='wrong', - patterns=dict(wrong='pattern'), - replacements='subject') - - with pytest.raises(TypeError, match=r"must be a list of strings"): - MyDataGrabber(datadir='/tmp', types=[1, 2, 3], - patterns={'1': 'pattern', '2': 'pattern', - '3': 'pattern'}, - replacements='subject') - - with pytest.raises(ValueError, match=r"must have the same length"): - MyDataGrabber(datadir='/tmp', types=['func', 'anat'], - patterns={'1': 'pattern', '2': 'pattern', - '3': 'pattern'}, - replacements=1) - - with pytest.raises(TypeError, match=r"patterns must be a dict"): - MyDataGrabber(datadir='/tmp', types=['func', 'anat'], - patterns='wrong', replacements='subject') - - with pytest.raises(ValueError, - match=r"patterns must have the same length"): - MyDataGrabber(datadir='/tmp', types=['func', 'anat'], - patterns={'wrong': 'pattern'}, replacements='subject') - - with pytest.raises(ValueError, match=r"patterns must contain all types"): - MyDataGrabber(datadir='/tmp', types=['func', 'anat'], - patterns={'wrong': 'pattern', 'func': 'pattern'}, - replacements='subject') - - with pytest.raises(TypeError, match=r"must be a list of strings"): - MyDataGrabber(datadir='/tmp', types=['func', 'anat'], - patterns={'func': 'func/test', 'anat': 'anat/test'}, - replacements=1) - - with pytest.warns(RuntimeWarning, match=r"not part of any pattern"): - MyDataGrabber(datadir='/tmp', types=['func', 'anat'], - patterns={'func': 'func/{subject}.nii', - 'anat': 'anat/{subject}.nii'}, - replacements=['subject', 'wrong']) - - datagrabber = MyDataGrabber( - datadir='/tmp/data', types=['func', 'anat'], - patterns={'func': 'func/{subject}.nii', - 'anat': 'anat/{subject}.nii'}, - replacements='subject') - assert datagrabber.datadir == Path('/tmp/data') - assert datagrabber.types == ['func', 'anat'] - assert datagrabber.replacements == ['subject'] - - datagrabber = MyDataGrabber( - datadir=Path('/tmp/data'), types=['func', 'anat'], - patterns={'func': 'func/{subject}.nii', - 'anat': 'anat/{subject}_{session}.nii'}, - replacements=['subject', 'session']) - assert datagrabber.datadir == Path('/tmp/data') - assert datagrabber.types == ['func', 'anat'] - assert datagrabber.replacements == ['subject', 'session'] - - -def test_bids_datalad_PatternDataGrabber(): - """Test a subject-based BIDS datalad datagrabber.""" - types = ['T1w', 'bold'] - patterns = { - 'T1w': '{subject}/anat/{subject}_T1w.nii.gz', - 'bold': '{subject}/func/{subject}_task-rest_bold.nii.gz' - } - replacements = ['subject'] - - with pytest.raises(ValueError, match=r"uri must be provided"): - PatternDataladDataGrabber(datadir=None, types=types, patterns=patterns, - replacements=replacements) - - repo_uri = _testing_dataset['example_bids']['uri'] - rootdir = 'example_bids' - repo_commit = _testing_dataset['example_bids']['id'] - - with PatternDataladDataGrabber( - rootdir=rootdir, uri=repo_uri, types=types, patterns=patterns, - replacements=replacements) as dg: - subs = [x for x in dg] - expected_subs = [f'sub-{i:02d}' for i in range(1, 10)] - assert set(subs) == set(expected_subs) - - for elem in dg: - t_sub = dg[elem] - assert 'path' in t_sub['T1w'] - assert t_sub['T1w']['path'] == \ - (dg.datadir / f'{elem}/anat/{elem}_T1w.nii.gz') - assert 'path' in t_sub['bold'] - assert t_sub['bold']['path'] == \ - (dg.datadir / f'{elem}/func/{elem}_task-rest_bold.nii.gz') - - assert 'meta' in t_sub - assert 'datagrabber' in t_sub['meta'] - dg_meta = t_sub['meta']['datagrabber'] - assert 'class' in dg_meta - assert dg_meta['class'] == 'PatternDataladDataGrabber' - assert 'uri' in dg_meta - assert dg_meta['uri'] == repo_uri - assert 'dataset_commit_id' in dg_meta - assert dg_meta['dataset_commit_id'] == repo_commit - - with open(t_sub['T1w']['path'], 'r') as f: - assert f.readlines()[0] == 'placeholder' - - with tempfile.TemporaryDirectory() as tmpdir: - datadir = Path(tmpdir) / 'dataset' # Need this for testing - patterns = { - 'T1w': '{subject}/anat/{subject}_T*w.nii.gz', - 'bold': '{subject}/func/{subject}_task-rest_*.nii.gz' - } - with PatternDataladDataGrabber( - rootdir=rootdir, uri=repo_uri, types=types, patterns=patterns, - datadir=datadir, replacements=replacements) as dg: - assert dg.datadir == datadir / rootdir - for elem in dg: - t_sub = dg[elem] - assert 'path' in t_sub['T1w'] - assert t_sub['T1w']['path'] == \ - (dg.datadir / f'{elem}/anat/{elem}_T1w.nii.gz') - assert 'path' in t_sub['bold'] - assert t_sub['bold']['path'] == \ - (dg.datadir / f'{elem}/func/{elem}_task-rest_bold.nii.gz') - - -def test_bids_datalad_PatternDataGrabber_session(): - """Test a subject and session-based BIDS datalad datagrabber.""" - types = ['T1w', 'bold'] - patterns = { - 'T1w': '{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz', - 'bold': '{subject}/{session}/func/' - '{subject}_{session}_task-rest_bold.nii.gz' - } - replacements = ['subject', 'session'] - - with pytest.raises(ValueError, match=r"uri must be provided"): - PatternDataladDataGrabber(datadir=None, types=types, patterns=patterns, - replacements=replacements) - - repo_uri = _testing_dataset['example_bids_ses']['uri'] - rootdir = 'example_bids_ses' - # repo_commit = _testing_dataset['example_bids_ses']['id'] - - # With T1W and bold, only 2 sessions are available - with PatternDataladDataGrabber( - rootdir=rootdir, uri=repo_uri, types=types, patterns=patterns, - replacements=replacements) as dg: - subs = [x for x in dg] - expected_subs = [(f'sub-{i:02d}', f'ses-{j:02d}') for j in range(1, 3) - for i in range(1, 10)] - assert set(subs) == set(expected_subs) - - # Test with a different T1w only, it should have 3 sessions - types = ['T1w'] - patterns = { - 'T1w': '{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz', - } - with PatternDataladDataGrabber( - rootdir=rootdir, uri=repo_uri, types=types, patterns=patterns, - replacements=replacements) as dg: - subs = [x for x in dg] - expected_subs = [(f'sub-{i:02d}', f'ses-{j:02d}') for j in range(1, 4) - for i in range(1, 10)] - assert set(subs) == set(expected_subs) -- 2.52.0 From f53e48e74bd4155cbc0cced8a41eceb079d8315f Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 16 Aug 2022 17:08:10 +0200 Subject: [PATCH 207/287] refactor: improve imports for datagrabber sub-package --- junifer/datagrabber/__init__.py | 11 +++++++++-- 1 file changed, 9 insertions(+), 2 deletions(-) diff --git a/junifer/datagrabber/__init__.py b/junifer/datagrabber/__init__.py index ce54e9f35..f742d61bb 100644 --- a/junifer/datagrabber/__init__.py +++ b/junifer/datagrabber/__init__.py @@ -1,5 +1,12 @@ +"""Provide imports for datagrabber sub-package.""" + # Authors: Federico Raimondo # Leonard Sasse # License: AGPL -from .base import (DataladDataGrabber, PatternDataGrabber, - PatternDataladDataGrabber) \ No newline at end of file + +from .base import BaseDataGrabber +from .datalad_base import DataladDataGrabber +from .hcp import DataladHCP1200, HCP1200 +from .meta import MultipleDataGrabber +from .pattern import PatternDataGrabber +from .pattern_datalad_base import PatternDataladDataGrabber -- 2.52.0 From 55ace6363a8296ef53f32f109e2e36b44a43eecc Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 1 Sep 2022 12:35:48 +0200 Subject: [PATCH 208/287] refactor: rename datagrabber/meta.py to datagrabber/multiple_base.py --- junifer/datagrabber/__init__.py | 2 +- junifer/datagrabber/{meta.py => multiple_base.py} | 0 2 files changed, 1 insertion(+), 1 deletion(-) rename junifer/datagrabber/{meta.py => multiple_base.py} (100%) diff --git a/junifer/datagrabber/__init__.py b/junifer/datagrabber/__init__.py index f742d61bb..c224480b4 100644 --- a/junifer/datagrabber/__init__.py +++ b/junifer/datagrabber/__init__.py @@ -7,6 +7,6 @@ from .base import BaseDataGrabber from .datalad_base import DataladDataGrabber from .hcp import DataladHCP1200, HCP1200 -from .meta import MultipleDataGrabber +from .multiple_base import MultipleDataGrabber from .pattern import PatternDataGrabber from .pattern_datalad_base import PatternDataladDataGrabber diff --git a/junifer/datagrabber/meta.py b/junifer/datagrabber/multiple_base.py similarity index 100% rename from junifer/datagrabber/meta.py rename to junifer/datagrabber/multiple_base.py -- 2.52.0 From ebb546883cdda20bf42de84f4260d3dda6d2af30 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 1 Sep 2022 12:36:44 +0200 Subject: [PATCH 209/287] refactor: update docstrings for datagrabber/multiple_base.py --- junifer/datagrabber/multiple_base.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/junifer/datagrabber/multiple_base.py b/junifer/datagrabber/multiple_base.py index 7fef19d60..6e0e28daf 100644 --- a/junifer/datagrabber/multiple_base.py +++ b/junifer/datagrabber/multiple_base.py @@ -1,4 +1,4 @@ -"""Provide abstract base class for metadata collection.""" +"""Provide abstract base class for multiple source datagrabber.""" # Authors: Federico Raimondo # Leonard Sasse @@ -12,15 +12,15 @@ from .base import BaseDataGrabber class MultipleDataGrabber(BaseDataGrabber): - """Abstract base class for metadata fetch. + """Abstract base class for data fetching from multiple sources. - Defines a DataGrabber which can be used to fetch metadata from multiple + Defines a DataGrabber which can be used to fetch data from multiple sources. Parameters ---------- datagrabbers : list of datagrabbers - The datagrabbers to use to fetch metadata using. + The datagrabbers to use to fetch data using. **kwargs Keyword arguments passed to superclass. -- 2.52.0 From 4261ec5205e42911d8388ad4d71562c9943d7662 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 17 Aug 2022 13:51:16 +0200 Subject: [PATCH 210/287] refactor: add docstrings and type annotations in configs/juseless.py --- junifer/configs/juseless.py | 52 ++++++++++++++++++++++--------------- 1 file changed, 31 insertions(+), 21 deletions(-) diff --git a/junifer/configs/juseless.py b/junifer/configs/juseless.py index f7f81493a..fa16361d3 100644 --- a/junifer/configs/juseless.py +++ b/junifer/configs/juseless.py @@ -1,7 +1,15 @@ -"""Provide class for juseless datagrabber.""" +"""Provide class for juseless datalad datagrabber.""" + +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + +from pathlib import Path +from typing import Union -from ..datagrabber import PatternDataladDataGrabber from ..api.decorators import register_datagrabber +from ..datagrabber import PatternDataladDataGrabber @register_datagrabber @@ -10,25 +18,27 @@ class JuselessDataladUKBVBM(PatternDataladDataGrabber): Implements a DataGrabber to access the UKB VBM data in Juseless. + Parameters + ----------- + datadir : str or pathlib.Path, optional + The directory where the datalad dataset will be cloned. If None, + the datalad dataset will be cloned into a temporary directory + (default None). + """ - def __init__(self, datadir=None): - """Initialize a JuselessUKBVBM object. - - Parameters - ---------- - datadir : str or Path - That directory where the datalad dataset will be cloned. If None, - (default), the datalad dataset will be cloned into a temporary - directory. - """ - uri = 'ria+http://ukb.ds.inm7.de#~cat_m0wp1' - rootdir = 'm0wp1' - types = ['VBM_GM'] - replacements = ['subject', 'session'] - patterns = { - 'VBM_GM': 'm0wp1sub-{subject}_ses-{session}_T1w.nii.gz' - } + def __init__(self, datadir: Union[str, Path, None] = None) -> None: + """Initialize the class.""" + uri = "ria+http://ukb.ds.inm7.de#~cat_m0wp1" + rootdir = "m0wp1" + types = ["VBM_GM"] + replacements = ["subject", "session"] + patterns = {"VBM_GM": "m0wp1sub-{subject}_ses-{session}_T1w.nii.gz"} super().__init__( - types=types, datadir=datadir, uri=uri, rootdir=rootdir, - replacements=replacements, patterns=patterns) + types=types, + datadir=datadir, + uri=uri, + rootdir=rootdir, + replacements=replacements, + patterns=patterns, + ) -- 2.52.0 From e767dc9850b0b39f109fa96cee6d4fa2e826dc9b Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 17 Aug 2022 13:51:42 +0200 Subject: [PATCH 211/287] refactor: add docstrings and type annotations in test_juseless.py --- junifer/configs/tests/test_juseless.py | 38 ++++++++++++++++---------- 1 file changed, 24 insertions(+), 14 deletions(-) diff --git a/junifer/configs/tests/test_juseless.py b/junifer/configs/tests/test_juseless.py index 3234fb86a..e451c01c2 100644 --- a/junifer/configs/tests/test_juseless.py +++ b/junifer/configs/tests/test_juseless.py @@ -1,31 +1,41 @@ """Provide tests for juseless datagrabber.""" +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + import socket + import pytest from junifer.configs.juseless import JuselessDataladUKBVBM -from junifer.utils.logging import configure_logging from junifer.datagrabber.hcp import DataladHCP1200 - -if socket.gethostname() != 'juseless': - pytest.skip('This tests are only for juseless', allow_module_level=True) - -configure_logging(level='DEBUG') +from junifer.utils.logging import configure_logging -def test_juselessdataladukbvbm_datagrabber(): +# Check if the test is running on juseless +if socket.gethostname() != "juseless": + pytest.skip("These tests are only for juseless", allow_module_level=True) + +configure_logging(level="DEBUG") + + +def test_juselessdataladukbvbm_datagrabber() -> None: """Test datalad UKBVBM datagrabber.""" with JuselessDataladUKBVBM() as dg: all_elements = dg.get_elements() test_element = all_elements[0] out = dg[test_element] - assert 'VBM_GM' in out - assert out['VBM_GM']['path'].name == \ - f'm0wp1sub-{test_element[0]}_ses-{test_element[1]}_T1w.nii.gz' - assert out['VBM_GM']['path'].exists() + assert "VBM_GM" in out + assert ( + out["VBM_GM"]["path"].name + == f"m0wp1sub-{test_element[0]}_ses-{test_element[1]}_T1w.nii.gz" + ) + assert out["VBM_GM"]["path"].exists() -def test_juselessdataladhcp_datagrabber(): +def test_juselessdataladhcp_datagrabber() -> None: """Test datalad HCP datagrabber.""" with DataladHCP1200() as dg: all_elements = dg.get_elements() @@ -33,5 +43,5 @@ def test_juselessdataladhcp_datagrabber(): out = dg[test_element] - assert out['BOLD']['path'].exists() - assert out['BOLD']['path'].isfile() + assert out["BOLD"]["path"].exists() + assert out["BOLD"]["path"].isfile() -- 2.52.0 From e8a8da8ecf99123615b9af4f133f2d43f47f5c8c Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 17 Aug 2022 13:56:15 +0200 Subject: [PATCH 212/287] refactor: improve imports for configs sub-package --- junifer/configs/__init__.py | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/junifer/configs/__init__.py b/junifer/configs/__init__.py index e69de29bb..e1161687a 100644 --- a/junifer/configs/__init__.py +++ b/junifer/configs/__init__.py @@ -0,0 +1,8 @@ +"""Provide imports for configs sub-package.""" + +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + +from .juseless import JuselessDataladUKBVBM -- 2.52.0 From 580e3c427ad71ca45d40bba36058e26c958fccd4 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 2 Sep 2022 11:50:11 +0200 Subject: [PATCH 213/287] refactor: remove imports from configs/__init__.py --- junifer/configs/__init__.py | 8 -------- 1 file changed, 8 deletions(-) diff --git a/junifer/configs/__init__.py b/junifer/configs/__init__.py index e1161687a..e69de29bb 100644 --- a/junifer/configs/__init__.py +++ b/junifer/configs/__init__.py @@ -1,8 +0,0 @@ -"""Provide imports for configs sub-package.""" - -# Authors: Federico Raimondo -# Leonard Sasse -# Synchon Mandal -# License: AGPL - -from .juseless import JuselessDataladUKBVBM -- 2.52.0 From 0b8424b670100e91bc9304bc09f520b9e1385fbf Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 17 Aug 2022 09:15:51 +0200 Subject: [PATCH 214/287] refactor: add docstrings and type annotations in testing/datagrabbers.py --- junifer/testing/datagrabbers.py | 55 +++++++++++++++++++++++++-------- 1 file changed, 42 insertions(+), 13 deletions(-) diff --git a/junifer/testing/datagrabbers.py b/junifer/testing/datagrabbers.py index 549e84e7c..3c1e8a150 100644 --- a/junifer/testing/datagrabbers.py +++ b/junifer/testing/datagrabbers.py @@ -5,6 +5,8 @@ # License: AGPL import tempfile +from typing import Dict, List + from nilearn import datasets from ..datagrabber.base import BaseDataGrabber @@ -13,26 +15,53 @@ from ..datagrabber.base import BaseDataGrabber class OasisVBMTestingDatagrabber(BaseDataGrabber): """DataGrabber for Oasis VBM testing data.""" - def __init__(self): - """Initialize class.""" + def __init__(self) -> None: + """Initialize the class.""" + # Create temporary directory datadir = tempfile.mkdtemp() - types = ['VBM_GM'] + # Define types + types = ["VBM_GM"] super().__init__(types=types, datadir=datadir) - def get_elements(self): - """Get elements.""" - return [f'sub-{x:02d}' for x in list(range(1, 11))] + def __getitem__(self, element: str) -> Dict: + """Implement indexing support. - def __getitem__(self, element): - """Get item implementation.""" + Paramters + --------- + element : str + The element to retrieve. + + Returns + ------- + dict + The data along with the metadata. + + """ out = super().__getitem__(element) - i_sub = int(element.split('-')[1]) - 1 - out['VBM_GM'] = {'path': self._dataset.gray_matter_maps[i_sub]} + i_sub = int(element.split("-")[1]) - 1 + out["VBM_GM"] = {"path": self._dataset.gray_matter_maps[i_sub]} # Set the element accordingly - out['meta']['element'] = {'subject': element} + out["meta"]["element"] = {"subject": element} return out - def __enter__(self): - """Context enter implementation.""" + def __enter__(self) -> "OasisVBMTestingDatagrabber": + """Implement context entry. + + Returns + ------- + OasisVBMTestingDatagrabber + + """ self._dataset = datasets.fetch_oasis_vbm(n_subjects=10) return self + + def get_elements(self) -> List[str]: + """Get elements. + + Returns + ------- + list of str + List of elements that can be grabbed. + + """ + return [f"sub-{x:02d}" for x in list(range(1, 11))] -- 2.52.0 From 2dda0b240b66cd5b565a54bb6354ca0edcc1140f Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 17 Aug 2022 09:16:53 +0200 Subject: [PATCH 215/287] chore: update code style in testing/registry.py --- junifer/testing/registry.py | 16 ++++++++++++---- 1 file changed, 12 insertions(+), 4 deletions(-) diff --git a/junifer/testing/registry.py b/junifer/testing/registry.py index 0464c6c76..77024f767 100644 --- a/junifer/testing/registry.py +++ b/junifer/testing/registry.py @@ -1,8 +1,16 @@ """Provide testing registry.""" -from .datagrabbers import OasisVBMTestingDatagrabber -from ..api.registry import register +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL +from ..api.registry import register +from .datagrabbers import OasisVBMTestingDatagrabber + + +# Register testing datagrabber register( - 'datagrabber', 'OasisVBMTestingDatagrabber', - OasisVBMTestingDatagrabber) + step="datagrabber", + name="OasisVBMTestingDatagrabber", + klass=OasisVBMTestingDatagrabber, +) -- 2.52.0 From 509efe10f54807b955afecda13ffadd9c530d53c Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Wed, 17 Aug 2022 09:17:13 +0200 Subject: [PATCH 216/287] refactor: improve imports for testing sub-package --- junifer/testing/__init__.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/junifer/testing/__init__.py b/junifer/testing/__init__.py index 4b2ba56f9..bae0e931f 100644 --- a/junifer/testing/__init__.py +++ b/junifer/testing/__init__.py @@ -1,3 +1,7 @@ +"""Provide imports for testing sub-package.""" + # Authors: Federico Raimondo +# Synchon Mandal # License: AGPL -from . import datagrabbers \ No newline at end of file + +from .datagrabbers import datagrabbers -- 2.52.0 From ccaaa4e8ed0b09aedf782a84d2ff14794914496b Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Sun, 14 Aug 2022 14:18:24 +0200 Subject: [PATCH 217/287] refactor: add docstrings and type annotations in api/registry.py --- junifer/api/registry.py | 100 +++++++++++++++++++++++++++------------- 1 file changed, 67 insertions(+), 33 deletions(-) diff --git a/junifer/api/registry.py b/junifer/api/registry.py index e7c8f2412..e3f0d42e7 100644 --- a/junifer/api/registry.py +++ b/junifer/api/registry.py @@ -2,106 +2,140 @@ # Authors: Federico Raimondo # Leonard Sasse +# Synchon Mandal # License: AGPL -from ..utils.logging import raise_error, logger +from typing import Dict, List, Optional +from ..utils.logging import logger, raise_error + + +# Define valid steps for operation _valid_steps = [ - 'datagrabber', 'datareader', 'preprocessing', 'marker', 'storage'] + "datagrabber", + "datareader", + "preprocessing", + "marker", + "storage", +] +# Define registry for valid steps _registry = {x: {} for x in _valid_steps} -def register(step, name, klass): +def register(step: str, name: str, klass: type) -> None: """Register a function to be used in a pipeline step. Parameters ---------- step : str - Name of the step + Name of the step. name : str - Name of the function + Name of the function. klass : class - Class to be registered + Class to be registered. + """ + # Verify step if step not in _valid_steps: - raise_error(f'Invalid step: {step}', ValueError) - logger.info(f'Registering {name} in {step}') + raise_error(msg=f"Invalid step: {step}", klass=ValueError) + + logger.info(f"Registering {name} in {step}") _registry[step][name] = klass -def get_step_names(step): +def get_step_names(step: str) -> List: """Get the names of the registered functions for a given step. Parameters ---------- step : str - Name of the step + Name of the step. Returns ------- list - List of registered function names + List of registered function names. + """ + # Verify step if step not in _valid_steps: - raise_error(f'Invalid step: {step}', ValueError) + raise_error(msg=f"Invalid step: {step}", klass=ValueError) + return list(_registry[step].keys()) -def get(step, name): +def get_class(step: str, name: str) -> type: """Get the class of the registered function for a given step. Parameters ---------- step : str - Name of the step + Name of the step. name : str - Name of the function + Name of the function. Returns ------- class - Registered function class + Registered function class. + """ + # Verify step if step not in _valid_steps: - raise_error(f'Invalid step: {step}', ValueError) + raise_error(msg=f"Invalid step: {step}", klass=ValueError) + # Verify step name if name not in _registry[step]: - raise_error(f'Invalid name: {name}', ValueError) + raise_error(msg=f"Invalid name: {name}", klass=ValueError) + return _registry[step][name] -def build(step, name, baseclass, init_params=None): +def build( + step: str, + name: str, + baseclass: type, + init_params: Optional[Dict] = None, +) -> type: """Ensure that the given object is an instance of the given class. Parameters ---------- step : str - Name of the step + Name of the step. name : str Name of the function. baseclass : class - Class to be checked against - init_parms : dict - Parameters to pass to the class constructor + Class to be checked against. + init_parms : dict, optional + Parameters to pass to the base class constructor (default None). Returns ------- object - Object if it is an instance of the given class, otherwise a - ValueError is raised + An instance of the given base class. Raises ------ ValueError - If the name is not a string or the object is not an instance of the - baseclass parameter. + If the created object with the given name is not an instance of the + base class. + """ - klass = get(step, name) + # Set default init parameters if init_params is None: init_params = {} - object = klass(**init_params) - if not isinstance(object, baseclass): + # Get class of the registered function + klass = get_class(step=step, name=name) + # Create instance of the class + object_ = klass(**init_params) + # Verify created instance belongs to the base class + if not isinstance(object_, baseclass): raise_error( - f'Invalid {step} ({object.__class__.__name__}). ' - f'Must inherit from {baseclass.__name__}', ValueError) - return object + msg=( + f"Invalid {step} ({object_.__class__.__name__}). " + f"Must inherit from {baseclass.__name__}" + ), + klass=ValueError, + ) + return object_ -- 2.52.0 From 2ca98c3a0e541c8e4280ac6dbbd614b8bd373407 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Sun, 14 Aug 2022 14:19:05 +0200 Subject: [PATCH 218/287] refactor: prune, add docstrings and type annotations in test_registry.py --- junifer/api/tests/test_registry.py | 134 +++++++++++++++++++++++------ 1 file changed, 106 insertions(+), 28 deletions(-) diff --git a/junifer/api/tests/test_registry.py b/junifer/api/tests/test_registry.py index afded33d0..bb0382093 100644 --- a/junifer/api/tests/test_registry.py +++ b/junifer/api/tests/test_registry.py @@ -1,58 +1,136 @@ """Provide tests for registry.""" -import pytest +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + +import logging from abc import ABC -from junifer.api.registry import register, get_step_names, get, build +import pytest + +from junifer.api.registry import build, get_class, get_step_names, register +from junifer.datagrabber import PatternDataGrabber +from junifer.storage import SQLiteFeatureStorage -def test_register_error(): - """Test register error.""" - with pytest.raises(ValueError, match='Invalid ste'): - register('foo', 'bar', 'baz') +def test_register_invalid_step(): + """Test register invalid step name.""" + with pytest.raises(ValueError, match="Invalid step:"): + register(step="foo", name="bar", klass="baz") -def test_gets(): - """Test get.""" - with pytest.raises(ValueError, match='Invalid ste'): - get_step_names('foo') +# TODO: improve paramterization +@pytest.mark.parametrize( + "step, name, klass", + [ + ("datagrabber", "pattern-dg", PatternDataGrabber), + ("storage", "sqlite-storage", SQLiteFeatureStorage), + ], +) +def test_register( + caplog: pytest.LogCaptureFixture, step: str, name: str, klass: str +) -> None: + """Test register. - datagrabbers = get_step_names('datagrabber') - assert 'bar' not in datagrabbers - register('datagrabber', 'bar', 'baz') - datagrabbers = get_step_names('datagrabber') - assert 'bar' in datagrabbers + Parameters + ---------- + caplog : pytest.LogCaptureFixture + A pytest fixture to capture logging. + step : str + The parametrized name of the step. + name : str + The parametrized name of the function. + klass : str + The parametrized name of the base class. - with pytest.raises(ValueError, match='Invalid ste'): - get('foo', 'bar') - - with pytest.raises(ValueError, match='Invalid name'): - get('datagrabber', 'foo') - - obj = get('datagrabber', 'bar') - assert obj == 'baz' + """ + with caplog.at_level(logging.INFO): + # Register + register(step=step, name=name, klass=klass) + # Check logging message + assert "Registering" in caplog.text +def test_get_step_names_invalid_step() -> None: + """Test get step name invalid step name.""" + with pytest.raises(ValueError, match="Invalid step:"): + get_step_names(step="foo") + + +def test_get_step_names_absent() -> None: + """Test get step names for absent name.""" + # Get step names for datagrabber + datagrabbers = get_step_names(step="datagrabber") + # Check for datagrabber step name + assert "bar" not in datagrabbers + + +def test_get_step_names() -> None: + """Test get step names.""" + # Register datagrabber + register(step="datagrabber", name="bar", klass="baz") + # Get step names for datagrabber + datagrabbers = get_step_names(step="datagrabber") + # Check for datagrabber step name + assert "bar" in datagrabbers + + +def test_get_class_invalid_step() -> None: + """Test get class invalid step name.""" + with pytest.raises(ValueError, match="Invalid step:"): + get_class(step="foo", name="bar") + + +def test_get_class_invalid_name() -> None: + """Test get class invalid function name.""" + with pytest.raises(ValueError, match="Invalid name:"): + get_class(step="datagrabber", name="foo") + + +# TODO: enable paramterization +def test_get_class(): + """Test get class.""" + # Register datagrabber + register(step="datagrabber", name="bar", klass="baz") + # Get class + obj = get_class(step="datagrabber", name="bar") + assert obj == "baz" + + +# TODO: possible parametrization? def test_build(): """Test building objects from names.""" import numpy as np + # Define abstract base class class SuperClass(ABC): pass + # Define concrete class class ConcreteClass(SuperClass): def __init__(self, value=1): self.value = value - register('datagrabber', 'concrete', ConcreteClass) + # Register + register(step="datagrabber", name="concrete", klass=ConcreteClass) - obj = build('datagrabber', 'concrete', SuperClass) + # Build + obj = build(step="datagrabber", name="concrete", baseclass=SuperClass) assert isinstance(obj, ConcreteClass) assert obj.value == 1 - obj = build('datagrabber', 'concrete', SuperClass, {'value': 2}) + # Build + obj = build( + step="datagrabber", + name="concrete", + baseclass=SuperClass, + init_params={"value": 2}, + ) assert isinstance(obj, ConcreteClass) assert obj.value == 2 - with pytest.raises(ValueError, match='Must inherit'): - build('datagrabber', 'concrete', np.ndarray) + # Check error + with pytest.raises(ValueError, match="Must inherit"): + build(step="datagrabber", name="concrete", baseclass=np.ndarray) -- 2.52.0 From ae5dc9b396a8781ce2ae429d1589a0e65cb79439 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Sun, 14 Aug 2022 14:24:30 +0200 Subject: [PATCH 219/287] refactor: add docstrings and type annotations in api/parser.py --- junifer/api/parser.py | 47 ++++++++++++++++++++++++++++++++----------- 1 file changed, 35 insertions(+), 12 deletions(-) diff --git a/junifer/api/parser.py b/junifer/api/parser.py index 2cf7d1c33..470f39bed 100644 --- a/junifer/api/parser.py +++ b/junifer/api/parser.py @@ -1,28 +1,51 @@ """Provide functions for parser.""" -import yaml +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + import importlib from pathlib import Path +from typing import Dict, Union -from ..utils.logging import raise_error, logger +import yaml + +from ..utils.logging import logger, raise_error -def parse_yaml(filepath): - """Parse YAML.""" +def parse_yaml(filepath: Union[str, Path]) -> Dict: + """Parse YAML. + + Parameters + ---------- + filepath : str or pathlib.Path + The filepath to read from. + + Returns + ------- + dict + The contents represented as dictionary. + + """ + # Convert str to Path if not isinstance(filepath, Path): filepath = Path(filepath) - logger.info(f'Parsing yaml file: {filepath.as_posix()}') - if not filepath.exists(): - raise_error(f'File does not exist: {filepath.as_posix()}') - with open(filepath, 'r') as f: - contents = yaml.safe_load(f) - if 'with' in contents: - to_load = contents['with'] + logger.info(f"Parsing yaml file: {str(filepath.absolute())}") + # Filepath existence check + if not filepath.exists(): + raise_error(f"File does not exist: {str(filepath.absolute())}") + # Filepath reading + with open(filepath, "r") as f: + contents = yaml.safe_load(f) + # Autload modules + if "with" in contents: + to_load = contents["with"] + # Convert autload modules to list if not isinstance(to_load, list): to_load = [to_load] for t_module in to_load: - logger.info(f'Importing module {t_module}') + logger.info(f"Importing module: {t_module}") importlib.import_module(t_module) return contents -- 2.52.0 From 7bd6b6b838039c573f9c72b625d6ea727c5d773c Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Sun, 14 Aug 2022 14:24:50 +0200 Subject: [PATCH 220/287] refactor: prune, add docstrings and type annotations in test_parser.py --- junifer/api/tests/test_parser.py | 89 ++++++++++++++++++++++---------- 1 file changed, 61 insertions(+), 28 deletions(-) diff --git a/junifer/api/tests/test_parser.py b/junifer/api/tests/test_parser.py index 0d712340a..e61ff9fdb 100644 --- a/junifer/api/tests/test_parser.py +++ b/junifer/api/tests/test_parser.py @@ -1,44 +1,77 @@ """Provide tests for parser.""" +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + import sys from pathlib import Path -import tempfile + import pytest from junifer.api.parser import parse_yaml -def test_parse_yaml(): - """Test parse yaml.""" - with pytest.raises(ValueError, match='does not exist'): - parse_yaml('foo.yaml') +def test_parse_yaml_failure() -> None: + """Test YAML parsing failure.""" + with pytest.raises(ValueError, match="does not exist"): + parse_yaml("foo.yaml") - with tempfile.TemporaryDirectory() as _tmpdir: - fname = Path(_tmpdir) / 'test.yaml' - with open(fname, 'w') as f: - f.write('foo: bar\n') - contents = parse_yaml(fname) - assert 'foo' in contents - assert 'bar' == contents['foo'] +def test_parse_yaml_success(tmp_path: Path) -> None: + """Test YAML parsing success. - with open(fname, 'w') as f: - f.write('foo: bar\n') - f.write('with: numpy\n') + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. - contents = parse_yaml(fname) - assert 'foo' in contents - assert 'bar' == contents['foo'] - assert 'with' in contents - assert 'numpy' in contents['with'] + """ + # Write test file + fname = tmp_path / "test_parse_yaml_success.yaml" + fname.write_text("foo: bar") + # Check test file + contents = parse_yaml(fname) + assert "foo" in contents + assert contents["foo"] == "bar" - assert 'junifer.configs.wrong_config' not in sys.modules - with open(fname, 'w') as f: - f.write('foo: bar\n') - f.write('with:\n') - f.write(' - numpy\n') - f.write(' - junifer.testing.wrong_config\n') +def test_parse_yaml_success_with_module_autoload(tmp_path: Path) -> None: + """Test YAML parsing with single module autoload success. - with pytest.raises(ImportError, match='wrong_config'): - contents = parse_yaml(fname.as_posix()) + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + # Write test file + fname = tmp_path / "test_parse_yaml_with_single_module_autoload.yaml" + fname.write_text("foo: bar\nwith: numpy") + # Check test file + contents = parse_yaml(fname) + assert "foo" in contents + assert contents["foo"] == "bar" + assert "with" in contents + assert contents["with"] == "numpy" + assert "numpy" in sys.modules + assert "junifer.configs.wrong_config" not in sys.modules + + +def test_parse_yaml_failure_with_multi_module_autoload(tmp_path: Path) -> None: + """Test YAML parsing with multi module autoload failure. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + # Write test file + fname = tmp_path / "test_parse_yaml_with_multi_module_autoload.yaml" + fname.write_text( + "foo: bar\nwith:\n - numpy\n - junifer.testing.wrong_config" + ) + # Check test file + with pytest.raises(ImportError, match="wrong_config"): + parse_yaml(fname) -- 2.52.0 From bce678bd1d4cad06936fb9fb8ff2e83f41871171 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Sun, 14 Aug 2022 14:44:01 +0200 Subject: [PATCH 221/287] refactor: add docstrings and type annotations in api/functions.py --- junifer/api/functions.py | 520 ++++++++++++++++++++++++++++++--------- 1 file changed, 401 insertions(+), 119 deletions(-) diff --git a/junifer/api/functions.py b/junifer/api/functions.py index fd16a2b28..816a2dac0 100644 --- a/junifer/api/functions.py +++ b/junifer/api/functions.py @@ -2,76 +2,111 @@ # Authors: Federico Raimondo # Leonard Sasse +# Synchon Mandal # License: AGPL -from pathlib import Path import shutil -import yaml import subprocess +from pathlib import Path +from typing import Dict, List, Optional, Sequence, Union + +import yaml -from .registry import build -from ..utils import logger, raise_error -from ..utils.fs import make_executable from ..datagrabber.base import BaseDataGrabber from ..markers.base import BaseMarker -from ..storage.base import BaseFeatureStorage from ..markers.collection import MarkerCollection +from ..storage.base import BaseFeatureStorage +from ..utils import logger, raise_error +from ..utils.fs import make_executable +from .registry import build -def _get_datagrabber(datagrabber_config): +def _get_datagrabber(datagrabber_config: Dict) -> Dict: + """Get datagrabber. + + Parameters + ---------- + datagrabber_config : dict + The config to get the datagrabber using. + + Returns + ------- + dict + The datagrabber. + + """ datagrabber_params = datagrabber_config.copy() - datagrabber_kind = datagrabber_params.pop('kind') + datagrabber_kind = datagrabber_params.pop("kind") datagrabber = build( - 'datagrabber', datagrabber_kind, BaseDataGrabber, - init_params=datagrabber_params) + step="datagrabber", + name=datagrabber_kind, + baseclass=BaseDataGrabber, + init_params=datagrabber_params, + ) return datagrabber -def run(workdir, datagrabber, markers, storage, elements=None): +def run( + workdir: Union[str, Path], + datagrabber: Dict, + markers: List[Dict], + storage: Dict, + elements: Union[str, Sequence, None] = None, +) -> None: """Run the pipeline on the selected element. Parameters ---------- - workdir : str or path-like object - Directory where the pipeline will be executed + workdir : str or pathlib.Path + Directory where the pipeline will be executed. datagrabber : dict Datagrabber to use. Must have a key 'kind' with the kind of datagrabber to use. All other keys are passed to the datagrabber init function. - elements : str, tuple or list[str or tuple] - Element(s) to process. Will be used to index the datagrabber. markers : list of dict List of markers to extract. Each marker is a dict with at least two - keys: 'name' and 'kind'. The 'name' key is used to name the output - marker. The 'kind' key is used to specify the kind of marker to + keys: "name" and "kind". The "name" key is used to name the output + marker. The "kind" key is used to specify the kind of marker to extract. The rest of the keys are used to pass parameters to the marker calculation. storage : dict - Storage to use. Must have a key 'kind' with the kind of + Storage to use. Must have a key "kind" with the kind of storage to use. All other keys are passed to the storage init function. - """ - storage_params = storage.copy() - storage_kind = storage_params.pop('kind') + elements : str or tuple or list of str or tuple, optional + Element(s) to process. Will be used to index the datagrabber + (default None). + """ + # Convert str to Path if isinstance(workdir, str): workdir = Path(workdir) - + # Get datagrabber to use datagrabber = _get_datagrabber(datagrabber) # Copy to avoid changing the original dict _markers = [x.copy() for x in markers] built_markers = [] for t_marker in _markers: - kind = t_marker.pop('kind') - t_m = build('marker', kind, BaseMarker, init_params=t_marker) + kind = t_marker.pop("kind") + t_m = build( + step="marker", + name=kind, + baseclass=BaseMarker, + init_params=t_marker, + ) built_markers.append(t_m) - + # Get storage engine to use + storage_params = storage.copy() + storage_kind = storage_params.pop("kind") storage = build( - 'storage', storage_kind, BaseFeatureStorage, - init_params=storage_params) - - mc = MarkerCollection(built_markers, storage=storage) - + step="storage", + name=storage_kind, + baseclass=BaseFeatureStorage, + init_params=storage_params, + ) + # Create new marker collection + mc = MarkerCollection(markers=built_markers, storage=storage) + # Fit elements with datagrabber: if elements is not None: for t_element in elements: @@ -81,112 +116,209 @@ def run(workdir, datagrabber, markers, storage, elements=None): mc.fit(datagrabber[t_element]) -def collect(storage): - """Collect data.""" +def collect(storage: Dict) -> None: + """Collect and store data. + + Parameters + ---------- + storage : dict + Storage to use. Must have a key "kind" with the kind of + storage to use. All other keys are passed to the storage + init function. + + """ storage_params = storage.copy() - storage_kind = storage_params.pop('kind') - logger.info(f'Collecting data using {storage_kind}') - logger.debug(f'\tStorage params: {storage_params}') + storage_kind = storage_params.pop("kind") + logger.info(f"Collecting data using {storage_kind}") + logger.debug(f"\tStorage params: {storage_params}") storage = build( - 'storage', storage_kind, BaseFeatureStorage, - init_params=storage_params) - logger.debug('Running storage.collect()') + step="storage", + name=storage_kind, + baseclass=BaseFeatureStorage, + init_params=storage_params, + ) + logger.debug("Running storage.collect()") storage.collect() - logger.info('Collect done') + logger.info("Collect done") def queue( - config, - kind, - jobname='junifer_job', - overwrite=False, - elements=None, - **kwargs -): + config: Dict, + kind: str, + jobname: str = "junifer_job", + overwrite: bool = False, + elements: Union[str, Sequence, None] = None, + **kwargs: Union[str, int, bool], +) -> None: # pragma : no cover """Queue a job to be executed later. Parameters ---------- - kind : str - The kind of job to queue. + config : dict + The configuration to be used for queueing the job. + kind : {"HTCondor", "SLURM"} + The kind of job queue system to use. + jobname : str, optional + The name of the job (default "junifer_job"). + overwrite : bool, optional + Whether to overwrite if job directory already exists (default False). + elements : str or tuple or list of str or tuple, optional + Element(s) to process. Will be used to index the datagrabber + (default None). **kwargs : dict - The parameters to pass to the job. + The keyword arguments to pass to the job queue system. + + Raises + ------ + ValueError + If the value of `kind` is invalid. + """ # Create a folder within the CWD to store the job files / config cwd = Path.cwd() - jobdir = cwd / 'junifer_jobs' / jobname - logger.info(f'Creating job in {jobdir.as_posix()}') + jobdir = cwd / "junifer_jobs" / jobname + logger.info(f"Creating job in {str(jobdir.absolute())}") if jobdir.exists(): if overwrite is not True: - raise_error(f'Job folder for {jobname} already exists. ' - 'This error is raise to prevent overwriting files ' - 'of jobs that might be scheduled but yet not ' - 'executed. Either delete the directory ' - f'{jobdir.as_posix()} or set overwrite to True.') + raise_error( + f"Job folder for {jobname} already exists. " + "This error is raised to prevent overwriting job files " + "that might be scheduled but not yet executed. " + f"Either delete the directory {str(jobdir.absolute())} " + "or set overwrite=True." + ) else: logger.info( - f'Deleting previous job directory {jobdir.as_posix()}') + f"Deleting existing job directory at {str(jobdir.absolute())}" + ) shutil.rmtree(jobdir) jobdir.mkdir(exist_ok=True, parents=True) - yaml_config = jobdir / 'config.yaml' - logger.info(f'Writing YAML config to {yaml_config}') - with open(yaml_config, 'w') as f: + yaml_config = jobdir / "config.yaml" + logger.info(f"Writing YAML config to {str(yaml_config.absolute())}") + with open(yaml_config, "w") as f: f.write(yaml.dump(config)) # Get list of elements if elements is None: - if 'elements' in config: - elements = config['elements'] + if "elements" in config: + elements = config["elements"] else: # If no elements are specified, use all elements from the # datagrabber - datagrabber = _get_datagrabber(config['datagrabber']) + datagrabber = _get_datagrabber(config["datagrabber"]) with datagrabber as dg: elements = dg.get_elements() - if kind == 'HTCondor': - _queue_condor(jobname, jobdir, yaml_config, elements, **kwargs) - elif kind == 'SLURM': - _queue_slurm(jobname, jobdir, yaml_config, elements, **kwargs) + if kind == "HTCondor": + _queue_condor( + jobname=jobname, + jobdir=jobdir, + yaml_config=yaml_config, + elements=elements, + **kwargs, + ) + elif kind == "SLURM": + _queue_slurm( + jobname=jobname, + jobdir=jobdir, + yaml_config=yaml_config, + elements=elements, + **kwargs, + ) else: - raise ValueError(f'Unknown queue kind: {kind}') + raise ValueError(f"Unknown queue kind: {kind}") - logger.info('Queue done') + logger.info("Queue done") -def _queue_condor(jobname, jobdir, yaml_config, elements, env, mem='8G', - cpus=1, disk='1G', extra_preamble='', verbose='info', - collect=True, submit=False): - logger.debug('Creating HTCondor job') - run_junifer_args = (f'run {yaml_config.as_posix()} ' - f'--verbose {verbose} --element $(element)') - collect_junifer_args = \ - f'collect {yaml_config.as_posix()} --verbose {verbose} ' +def _queue_condor( + jobname: str, + jobdir: Path, + yaml_config: Path, + elements: Union[str, Sequence, None] = None, + env: Optional[Dict[str, str]] = None, + mem: str = "8G", + cpus: int = 1, + disk: str = "1G", + extra_preamble: str = "", + verbose: str = "info", + collect: bool = True, + submit: bool = False, +) -> None: + """Submit job to HTCondor. + + Parameters + ---------- + jobname : str + The name of the job. + jobdir : pathlib.Path + The path to the job directory. + yaml_config : pathlib.Path + The path to the YAML config file. + elements : str or tuple or list of str or tuple, optional + Element(s) to process. Will be used to index the datagrabber + (default None). + env : dict, optional + The environment variables passed as dictionary (default None). + mem : str, optional + The size of memory (RAM) to use (default "8G"). + cpus : int, optional + The number of CPU cores to use (default 1). + disk : str, optional + The size of disk (HDD or SSD) to use (default "1G"). + extra_preamble : str, optional + Extra commands to pass to HTCondor (default ""). + verbose : str, optional + The level of verbosity (default "info"). + collect : bool, optional + Whether to submit "collect" task for junifer (default True). + submit : bool, optional + Whether to submit the jobs. In any case, .dag files will be created + for submission (default False). + + Raises + ------ + ValueError + If the value of `env` is invalid. + + """ + logger.debug("Creating HTCondor job") + run_junifer_args = ( + f"run {str(yaml_config.absolute())} " + f"--verbose {verbose} --element $(element)" + ) + collect_junifer_args = ( + f"collect {str(yaml_config.absolute())} --verbose {verbose} " + ) # Set up the env_name, executable and arguments according to the # environment type if env is None: - env = {'kind': 'local'} - if env['kind'] == 'conda': - env_name = env['name'] - executable = 'run_conda.sh' - arguments = f'{env_name} junifer' + env = {"kind": "local"} + if env["kind"] == "conda": + env_name = env["name"] + executable = "run_conda.sh" + arguments = f"{env_name} junifer" # TODO: Copy run_conda.sh to jobdir exec_path = jobdir / executable - shutil.copy(Path(__file__).parent / 'res' / executable, exec_path) + shutil.copy(Path(__file__).parent / "res" / executable, exec_path) make_executable(exec_path) - elif env['kind'] == 'venv': - env_name = env['name'] - executable = 'run_venv.sh' - arguments = f'{env_name} junifer' + elif env["kind"] == "venv": + env_name = env["name"] + executable = "run_venv.sh" + arguments = f"{env_name} junifer" # TODO: Copy run_venv.sh to jobdir - elif env['kind'] == 'local': - executable = 'junifer' - arguments = '' + elif env["kind"] == "local": + executable = "junifer" + arguments = "" else: raise ValueError(f'Unknown env kind: {env["kind"]}') - log_dir = jobdir / 'logs' + + # Create log directory + log_dir = jobdir / "logs" log_dir.mkdir(exist_ok=True, parents=True) + + # Add preamble data run_preamble = f""" # The environment universe = vanilla @@ -198,7 +330,7 @@ def _queue_condor(jobname, jobdir, yaml_config, elements, env, mem='8G', request_disk = {disk} # Executable - initial_dir = {jobdir.as_posix()} + initial_dir = {str(jobdir.absolute())} executable = $(initial_dir)/{executable} transfer_executable = False @@ -207,19 +339,19 @@ def _queue_condor(jobname, jobdir, yaml_config, elements, env, mem='8G', {extra_preamble} # Logs - log = {log_dir.as_posix()}/junifer_run_$(element).log - output = {log_dir.as_posix()}/junifer_run_$(element).out - error = {log_dir.as_posix()}/junifer_run_$(element).err + log = {str(log_dir.absolute())}/junifer_run_$(element).log + output = {str(log_dir.absolute())}/junifer_run_$(element).out + error = {str(log_dir.absolute())}/junifer_run_$(element).err """ - submit_run_fname = jobdir / f'run_{jobname}.submit' - submit_collect_fname = jobdir / f'collect_{jobname}.submit' - dag_fname = jobdir / f'{jobname}.dag' + submit_run_fname = jobdir / f"run_{jobname}.submit" + submit_collect_fname = jobdir / f"collect_{jobname}.submit" + dag_fname = jobdir / f"{jobname}.dag" # Write to run submit files - with open(submit_run_fname, 'w') as submit_file: + with open(submit_run_fname, "w") as submit_file: submit_file.write(run_preamble) - submit_file.write('queue\n') + submit_file.write("queue\n") collect_preamble = f""" # The environment @@ -232,7 +364,7 @@ def _queue_condor(jobname, jobdir, yaml_config, elements, env, mem='8G', request_disk = {disk} # Executable - initial_dir = {jobdir.as_posix()} + initial_dir = {str(jobdir.absolute())} executable = $(initial_dir)/{executable} transfer_executable = False @@ -241,37 +373,187 @@ def _queue_condor(jobname, jobdir, yaml_config, elements, env, mem='8G', {extra_preamble} # Logs - log = {log_dir.as_posix()}/junifer_collect.log - output = {log_dir.as_posix()}/junifer_collect.out - error = {log_dir.as_posix()}/junifer_collect.err + log = {str(log_dir.absolute())}/junifer_collect.log + output = {str(log_dir.absolute())}/junifer_collect.out + error = {str(log_dir.absolute())}/junifer_collect.err """ # Now create the collect submit file - with open(submit_collect_fname, 'w') as submit_file: + with open(submit_collect_fname, "w") as submit_file: submit_file.write(collect_preamble) # Eval preamble here - submit_file.write('queue\n') + submit_file.write("queue\n") - with open(dag_fname, 'w') as dag_file: + with open(dag_fname, "w") as dag_file: # Get all subject and session names from file list for i_job, t_elem in enumerate(elements): - dag_file.write(f'JOB run{i_job} {submit_run_fname}\n') + dag_file.write(f"JOB run{i_job} {submit_run_fname}\n") dag_file.write(f'VARS run{i_job} element="{t_elem}"\n\n') if collect is True: - dag_file.write(f'JOB collect {submit_collect_fname}\n') - dag_file.write('PARENT ') + dag_file.write(f"JOB collect {submit_collect_fname}\n") + dag_file.write("PARENT ") for i_job, _t_elem in enumerate(elements): - dag_file.write(f'run{i_job} ') - dag_file.write('CHILD collect\n\n') + dag_file.write(f"run{i_job} ") + dag_file.write("CHILD collect\n\n") + # Submit job(s) if submit is True: - logger.info('Submitting HTCondor job') - subprocess.run(['condor_submit_dag', dag_fname]) - logger.info('HTCondor job submitted') + logger.info("Submitting HTCondor job") + subprocess.run(["condor_submit_dag", dag_fname]) + logger.info("HTCondor job submitted") else: - cmd = f'condor_submit_dag {dag_fname.as_posix()}' - logger.info('HTCondor job files created, to submit the job, ' - f'run "{cmd}"') + cmd = f"condor_submit_dag {str(dag_fname.absolute())}" + logger.info( + f"HTCondor job files created, to submit the job, run `{cmd}`" + ) -def _queue_slurm(jobname, jobdir, yaml_config, elements): +def _queue_slurm( + jobname: str, + jobdir: Path, + yaml_config: Path, + elements: Union[str, Sequence, None] = None, +) -> None: + """Submit job to SLURM. + + Parameters + ---------- + jobname : str + The name of the job. + jobdir : pathlib.Path + The path to the job directory. + yaml_config : pathlib.Path + The path to the YAML config file. + elements : str or tuple or list[str or tuple], optional + Element(s) to process. Will be used to index the datagrabber + (default None). + + """ pass + # logger.debug("Creating SLURM job") + # run_junifer_args = ( + # f"run {str(yaml_config.absolute())} " + # f"--verbose {verbose} --element $(element)" + # ) + # collect_junifer_args = \ + # f"collect {str(yaml_config.absolute())} --verbose {verbose} " + + # # Set up the env_name, executable and arguments according to the + # # environment type + # if env is None: + # env = { + # "kind": "local", + # } + # if env["kind"] == "conda": + # env_name = env["name"] + # executable = "run_conda.sh" + # arguments = f"{env_name} junifer" + # # TODO: Copy run_conda.sh to jobdir + # exec_path = jobdir / executable + # shutil.copy(Path(__file__).parent / "res" / executable, exec_path) + # make_executable(exec_path) + # elif env["kind"] == "venv": + # env_name = env["name"] + # executable = "run_venv.sh" + # arguments = f"{env_name} junifer" + # # TODO: Copy run_venv.sh to jobdir + # elif env["kind"] == "local": + # executable = "junifer" + # arguments = "" + # else: + # raise ValueError(f"Unknown env kind: {env['kind']}") + + # # Create log directory + # log_dir = jobdir / 'logs' + # log_dir.mkdir(exist_ok=True, parents=True) + + # # Add preamble data + # run_preamble = f""" + # #!/bin/bash + + # #SBATCH --job-name={} + # #SBATCH --account={} + # #SBATCH --partition={} + # #SBATCH --time={} + # #SBATCH --ntasks={} + # #SBATCH --cpus-per-task={cpus} + # #SBATCH --mem-per-cpu={mem} + # #SBATCH --mail-type={} + # #SBATCH --mail-user={} + # #SBATCH --output={} + # #SBATCH --error={} + + # # Executable + # initial_dir = {str(jobdir.absolute())} + # executable = $(initial_dir)/{executable} + # transfer_executable = False + + # arguments = {arguments} {run_junifer_args} + + # {extra_preamble} + + # # Logs + # log = {str(log_dir.absolute())}/junifer_run_$(element).log + # output = {str(log_dir.absolute())}/junifer_run_$(element).out + # error = {str(log_dir.absolute())}/junifer_run_$(element).err + # """ + + # submit_run_fname = jobdir / f'run_{jobname}.sh' + # submit_collect_fname = jobdir / f'collect_{jobname}.sh' + + # # Write to run submit files + # with open(submit_run_fname, 'w') as submit_file: + # submit_file.write(run_preamble) + # submit_file.write('queue\n') + + # collect_preamble = f""" + # # The environment + # universe = vanilla + # getenv = True + + # # Resources + # request_cpus = {cpus} + # request_memory = {mem} + # request_disk = {disk} + + # # Executable + # initial_dir = {str(jobdir.absolute())} + # executable = $(initial_dir)/{executable} + # transfer_executable = False + + # arguments = {arguments} {collect_junifer_args} + + # {extra_preamble} + + # # Logs + # log = {str(log_dir.absolute())}/junifer_collect.log + # output = {str(log_dir.absolute())}/junifer_collect.out + # error = {str(log_dir.absolute())}/junifer_collect.err + # """ + + # # Now create the collect submit file + # with open(submit_collect_fname, 'w') as submit_file: + # submit_file.write(collect_preamble) # Eval preamble here + # submit_file.write('queue\n') + + # with open(dag_fname, 'w') as dag_file: + # # Get all subject and session names from file list + # for i_job, t_elem in enumerate(elements): + # dag_file.write(f'JOB run{i_job} {submit_run_fname}\n') + # dag_file.write(f'VARS run{i_job} element="{t_elem}"\n\n') + # if collect is True: + # dag_file.write(f'JOB collect {submit_collect_fname}\n') + # dag_file.write('PARENT ') + # for i_job, _t_elem in enumerate(elements): + # dag_file.write(f'run{i_job} ') + # dag_file.write('CHILD collect\n\n') + + # # Submit job(s) + # if submit is True: + # logger.info('Submitting SLURM job') + # subprocess.run(['condor_submit_dag', dag_fname]) + # logger.info('HTCondor SLURM submitted') + # else: + # cmd = f"condor_submit_dag {str(dag_fname.absolute())}" + # logger.info( + # f"SLURM job files created, to submit the job, run `{cmd}`" + # ) -- 2.52.0 From 1b4030ecc674d518da46c684644e72563753d79f Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Sun, 14 Aug 2022 14:44:22 +0200 Subject: [PATCH 222/287] refactor: prune, add docstrings and type annotations in test_functions.py --- junifer/api/tests/test_functions.py | 204 ++++++++++++++++++---------- 1 file changed, 134 insertions(+), 70 deletions(-) diff --git a/junifer/api/tests/test_functions.py b/junifer/api/tests/test_functions.py index 9a00a5e8f..57f1cf7d6 100644 --- a/junifer/api/tests/test_functions.py +++ b/junifer/api/tests/test_functions.py @@ -1,93 +1,157 @@ """Provide tests for functions.""" -import tempfile +# Authors: Federico Raimondo +# Leonard Sasse +# Synchon Mandal +# License: AGPL + from pathlib import Path -from junifer.api.registry import build -from junifer.api.functions import run, collect -from junifer.datagrabber.base import BaseDataGrabber -import junifer.testing.registry # noqa: F401 +import pytest +import junifer.testing.registry # noqa: F401 +from junifer.api.functions import collect, run +from junifer.api.registry import build +from junifer.datagrabber.base import BaseDataGrabber + + +# Define datagrabber datagrabber = { - 'kind': 'OasisVBMTestingDatagrabber', + "kind": "OasisVBMTestingDatagrabber", } +# Define markers markers = [ - {'name': 'Schaefer1000x7_Mean', - 'kind': 'ParcelAggregation', - 'atlas': 'Schaefer1000x7', - 'method': 'mean'}, - {'name': 'Schaefer1000x7_Std', - 'kind': 'ParcelAggregation', - 'atlas': 'Schaefer1000x7', - 'method': 'std'} + { + "name": "Schaefer1000x7_Mean", + "kind": "ParcelAggregation", + "atlas": "Schaefer1000x7", + "method": "mean", + }, + { + "name": "Schaefer1000x7_Std", + "kind": "ParcelAggregation", + "atlas": "Schaefer1000x7", + "method": "std", + }, ] +# Define storage storage = { - 'kind': 'SQLiteFeatureStorage', + "kind": "SQLiteFeatureStorage", } -def test_run(): - """Test run function.""" - with tempfile.TemporaryDirectory() as tmpdir: - tmp_path = Path(tmpdir) - workdir = tmp_path / 'workdir' - workdir.mkdir() - outdir = tmp_path / 'out' - outdir.mkdir() - uri = outdir / 'test.db' - storage['uri'] = uri # type: ignore - run( - workdir=workdir, - datagrabber=datagrabber, - markers=markers, - storage=storage, - elements=['sub-01'] - ) +def test_run_single_element(tmp_path: Path) -> None: + """Test run function with single element. - files = list(outdir.glob('*.db')) - assert len(files) == 1 + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. - run( - workdir=workdir.as_posix(), - datagrabber=datagrabber, - markers=markers, - storage=storage, - elements=['sub-01', 'sub-03'] - ) - - files = list(outdir.glob('*.db')) - assert len(files) == 2 + """ + # Create working directory + workdir = tmp_path / "workdir_single" + workdir.mkdir() + # Create output directory + outdir = tmp_path / "out" + outdir.mkdir() + # Create storage + uri = outdir / "test.db" + storage["uri"] = uri + # Run operations + run( + workdir=workdir, + datagrabber=datagrabber, + markers=markers, + storage=storage, + elements=["sub-01"], + ) + # Check files + files = list(outdir.glob("*.db")) + assert len(files) == 1 -def test_collect(): - """Test run and collect functions.""" +def test_run_multi_element(tmp_path: Path) -> None: + """Test run function with multi element. - with tempfile.TemporaryDirectory() as tmpdir: - tmp_path = Path(tmpdir) - workdir = tmp_path / 'workdir' - workdir.mkdir() - outdir = tmp_path / 'out' - outdir.mkdir() - uri = outdir / 'test.db' - storage['uri'] = uri # type: ignore - run( - workdir=workdir, - datagrabber=datagrabber, - markers=markers, - storage=storage, - ) - dg = build('datagrabber', datagrabber['kind'], BaseDataGrabber) - elements = dg.get_elements() + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. - # This should create 10 files - files = list(outdir.glob('*.db')) - assert len(files) == len(elements) + """ + # Create working directory + workdir = tmp_path / "workdir_multi" + workdir.mkdir() + # Create output directory + outdir = tmp_path / "out" + outdir.mkdir() + # Create storage + uri = outdir / "test.db" + storage["uri"] = uri + # Run operations + run( + workdir=workdir, + datagrabber=datagrabber, + markers=markers, + storage=storage, + elements=["sub-01", "sub-03"], + ) + # Check files + files = list(outdir.glob("*.db")) + assert len(files) == 2 - # But the test.db file should not exist - assert not uri.exists() - collect(storage) - # Now the file exists - assert uri.exists() +def test_run_and_collect(tmp_path: Path) -> None: + """Test run and collect functions. + + Parameters + ---------- + tmp_path : pathlib.Path + The path to the test directory. + + """ + # Create working directory + workdir = tmp_path / "workdir" + workdir.mkdir() + # Create output directory + outdir = tmp_path / "out" + outdir.mkdir() + # Create storage + uri = outdir / "test.db" + storage["uri"] = uri + # Run operations + run( + workdir=workdir, + datagrabber=datagrabber, + markers=markers, + storage=storage, + ) + # Get datagrabber + dg = build( + step="datagrabber", name=datagrabber["kind"], baseclass=BaseDataGrabber + ) + elements = dg.get_elements() + # This should create 10 files + files = list(outdir.glob("*.db")) + assert len(files) == len(elements) + # But the test.db file should not exist + assert not uri.exists() + # Collect in storage + collect(storage) + # Now the file exists + assert uri.exists() + + +@pytest.mark.skip(reason="HTCondor not installed on system.") +def test_queue_condor() -> None: + """Test job queueing in HTCondor.""" + pass + + +@pytest.mark.skip(reason="SLURM not installed on system.") +def test_queue_slurm() -> None: + """Test job queueing in SLURM.""" + pass -- 2.52.0 From cd0b9b55974573de5894b43194229deae646fd4e Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Sun, 14 Aug 2022 14:49:26 +0200 Subject: [PATCH 223/287] refactor: add docstrings and type annotations in api/decorators.py --- junifer/api/decorators.py | 37 ++++++++++++++++++++++++++----------- 1 file changed, 26 insertions(+), 11 deletions(-) diff --git a/junifer/api/decorators.py b/junifer/api/decorators.py index ed624a58d..49011871d 100644 --- a/junifer/api/decorators.py +++ b/junifer/api/decorators.py @@ -8,8 +8,8 @@ from .registry import register -def register_datagrabber(klass): - """Datagrabber decorator. +def register_datagrabber(klass: type) -> type: + """Datagrabber registration decorator. Registers the datagrabber so it can be used by name. @@ -21,14 +21,19 @@ def register_datagrabber(klass): Returns ------- klass: class - The unmodified input class + The unmodified input class. + """ - register('datagrabber', klass.__name__, klass) + register( + step="datagrabber", + name=klass.__name__, + klass=klass, + ) return klass -def register_marker(klass): - """Marker decorator. +def register_marker(klass: type) -> type: + """Marker registration decorator. Registers the marker so it can be used by name. @@ -40,14 +45,19 @@ def register_marker(klass): Returns ------- klass: class - The unmodified input class + The unmodified input class. + """ - register('marker', klass.__name__, klass) + register( + step="marker", + name=klass.__name__, + klass=klass, + ) return klass def register_storage(klass): - """Storage decorator. + """Storage registration decorator. Registers the storage so it can be used by name. @@ -59,7 +69,12 @@ def register_storage(klass): Returns ------- klass: class - The unmodified input class + The unmodified input class. + """ - register('storage', klass.__name__, klass) + register( + step="storage", + name=klass.__name__, + klass=klass, + ) return klass -- 2.52.0 From 97ab7bd6e5007099944ce3b8308ca29206d930ca Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Sun, 14 Aug 2022 14:56:40 +0200 Subject: [PATCH 224/287] refactor: add docstrings and type annotations in api/cli.py --- junifer/api/cli.py | 200 ++++++++++++++++++++++++++++++++------------- 1 file changed, 142 insertions(+), 58 deletions(-) diff --git a/junifer/api/cli.py b/junifer/api/cli.py index e338968e7..078b3c4f2 100644 --- a/junifer/api/cli.py +++ b/junifer/api/cli.py @@ -1,95 +1,179 @@ """Provide functions for cli.""" +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from typing import Dict, List + import click -from .parser import parse_yaml -from .functions import run as api_run +from ..utils.logging import configure_logging, logger, warn_with_log from .functions import collect as api_collect from .functions import queue as api_queue -from ..utils.logging import configure_logging, logger, warn +from .functions import run as api_run +from .parser import parse_yaml -def _parse_elements(element, config): - logger.debug(f'Parsing elements: {element}') +def _parse_elements(element: str, config: Dict) -> List: + """Parse elements from cli. + + Parameters + ---------- + element : str + The element to operate on. + config : dict + The configuration to operate using. + + Returns + ------- + list + The element(s) as list. + + """ + logger.debug(f"Parsing elements: {element}") if len(element) == 0: return None # TODO: If len == 1, check if its a file, then parse elements from file - elements = [x.split(',') if ',' in x else x for x in element] - logger.debug(f'Parsed elements: {elements}') - if elements is not None and 'elements' in config: - warn('One or more elements have been specified in both the command ' - 'line and in the config file. The command line has precedence ' - 'over the configuration file. That is, the elements specified ' - 'in the command line will be used. The elements specified in ' - 'the configuration file will be ignored. To remove this warning, ' - 'please remove the "elements" item from the configuration file.') + elements = [x.split(",") if "," in x else x for x in element] + logger.debug(f"Parsed elements: {elements}") + if elements is not None and "elements" in config: + warn_with_log( + "One or more elements have been specified in both the command " + "line and in the config file. The command line has precedence " + "over the configuration file. That is, the elements specified " + "in the command line will be used. The elements specified in " + "the configuration file will be ignored. To remove this warning, " + 'please remove the "elements" item from the configuration file.' + ) elif elements is None: - elements = config.get('elements', None) + elements = config.get("elements", None) return elements @click.group() -def cli(): - """CLI wrapper.""" - pass +def cli() -> None: # pragma: no cover + """CLI for JUelich NeuroImaging FEature extractoR.""" @cli.command() @click.argument( - 'filepath', - type=click.Path(exists=True, readable=True, dir_okay=False)) -@click.option('-v', '--verbose', - type=click.Choice(['warning', 'info', 'debug'], - case_sensitive=False), - default='info') -@click.option('--element', type=str, multiple=True) -def run(filepath, element, verbose): - """Run command for CLI.""" + "filepath", type=click.Path(exists=True, readable=True, dir_okay=False) +) +@click.option("--element", type=str, multiple=True) +@click.option( + "-v", + "--verbose", + type=click.Choice(["warning", "info", "debug"], case_sensitive=False), + default="info", +) +def run(filepath: click.Path, element: str, verbose: click.Choice) -> None: + """Run command for CLI. + + Parameters + ---------- + filepath : click.Path + The filepath to the configuration file. + element : str + The element to operate using. + verbose : click.Choice + The verbosity level: warning, info or debug (default "info"). + + """ configure_logging(level=verbose.upper()) + # TODO: add validation config = parse_yaml(filepath) - workdir = config['workdir'] - datagrabber = config['datagrabber'] - markers = config['markers'] - storage = config['storage'] + workdir = config["workdir"] + datagrabber = config["datagrabber"] + markers = config["markers"] + storage = config["storage"] elements = _parse_elements(element, config) + # Perform operation api_run( - workdir=workdir, datagrabber=datagrabber, markers=markers, - storage=storage, elements=elements) + workdir=workdir, + datagrabber=datagrabber, + markers=markers, + storage=storage, + elements=elements, + ) @cli.command() @click.argument( - 'filepath', - type=click.Path(exists=True, readable=True, dir_okay=False)) -@click.option('-v', '--verbose', - type=click.Choice(['warning', 'info', 'debug'], - case_sensitive=False), - default='info') -def collect(filepath, verbose): - """Collect command for CLI.""" + "filepath", type=click.Path(exists=True, readable=True, dir_okay=False) +) +@click.option( + "-v", + "--verbose", + type=click.Choice(["warning", "info", "debug"], case_sensitive=False), + default="info", +) +def collect(filepath: click.Path, verbose: click.Choice) -> None: + """Collect command for CLI. + + Parameters + ---------- + filepath : click.Path + The filepath to the configuration file. + verbose : click.Choice + The verbosity level: warning, info or debug (default "info"). + + """ configure_logging(level=verbose.upper()) + # TODO: add validation config = parse_yaml(filepath) - storage = config['storage'] - api_collect(storage) + storage = config["storage"] + # Perform operation + api_collect(storage=storage) @cli.command() @click.argument( - 'filepath', - type=click.Path(exists=True, readable=True, dir_okay=False)) -@click.option('-v', '--verbose', - type=click.Choice(['warning', 'info', 'debug'], - case_sensitive=False), - default='info') -@click.option('--overwrite', is_flag=True) -@click.option('--submit', is_flag=True) -@click.option('--element', type=str, multiple=True) -def queue(filepath, element, overwrite, submit, verbose): - """Queue command for CLI.""" + "filepath", type=click.Path(exists=True, readable=True, dir_okay=False) +) +@click.option("--element", type=str, multiple=True) +@click.option("--overwrite", is_flag=True) +@click.option("--submit", is_flag=True) +@click.option( + "-v", + "--verbose", + type=click.Choice(["warning", "info", "debug"], case_sensitive=False), + default="info", +) +def queue( + filepath: click.Path, + element: str, + overwrite: bool, + submit: bool, + verbose: click.Choice, +) -> None: + """Queue command for CLI. + + Parameters + ---------- + filepath : click.Path + The filepath to the configuration file. + element : str + The element to operate using. + overwrite : bool + Whether to overwrite existing directory. + submit : bool + Whether to submit the job. + verbose : click.Choice + The verbosity level: warning, info or debug (default "info"). + + """ configure_logging(level=verbose.upper()) + # TODO: add validation config = parse_yaml(filepath) elements = _parse_elements(element, config) - queue_config = config.pop('queue') - kind = queue_config.pop('kind') - api_queue(config, kind=kind, overwrite=overwrite, submit=submit, - elements=elements, **queue_config) + queue_config = config.pop("queue") + kind = queue_config.pop("kind") + api_queue( + config=config, + kind=kind, + overwrite=overwrite, + elements=elements, + submit=submit, + **queue_config, + ) -- 2.52.0 From d69e240c9b382978f9c19bb80f05c681148f6f5e Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Sun, 14 Aug 2022 14:57:03 +0200 Subject: [PATCH 225/287] refactor: prune, add docstrings and type annotations in test_cli.py --- junifer/api/tests/test_cli.py | 110 ++++++++++++++++++---------------- 1 file changed, 59 insertions(+), 51 deletions(-) diff --git a/junifer/api/tests/test_cli.py b/junifer/api/tests/test_cli.py index db99ab9e1..380196f12 100644 --- a/junifer/api/tests/test_cli.py +++ b/junifer/api/tests/test_cli.py @@ -1,62 +1,70 @@ """Provide tests for cli.""" -from pathlib import Path -import tempfile -import yaml -from junifer.api.cli import run, collect +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL +from pathlib import Path +from typing import Tuple + +import pytest +import yaml from click.testing import CliRunner +from junifer.api.cli import collect, run + +# Create click test runner runner = CliRunner() -def _modify_path(tmpdir, in_file): - """Modify the path to use the temporary directory.""" - if not isinstance(tmpdir, Path): - tmpdir = Path(tmpdir) - with open(in_file, 'r') as f: +# TODO: adapt elements to take arrays +@pytest.mark.parametrize( + "elements", + [ + ("sub-01", "sub-02", "sub-03"), + ("sub-01", "sub-02", "sub-04"), + ], +) +def test_run_and_collect_commands( + tmp_path: Path, elements: Tuple[str, ...] +) -> None: + """Test run and collect commands.""" + # Get test config + infile = Path(__file__).parent / "data" / "gmd_mean.yaml" + # Read test config + with open(infile, mode="r") as f: contents = yaml.safe_load(f) - outfile = tmpdir / 'in.yaml' - outdir = tmpdir / 'out' - workdir = tmpdir / 'work' - contents['storage']['uri'] = outdir.as_posix() - contents['workdir'] = workdir.as_posix() - with open(outfile, 'w') as f: + # Working directory + workdir = tmp_path / "workdir" + contents["workdir"] = str(workdir.absolute()) + # Output directory + outdir = tmp_path / "outdir" + # Storage + contents["storage"]["uri"] = str(outdir.absolute()) + # Write new test config + outfile = tmp_path / "in.yaml" + with open(outfile, mode="w") as f: yaml.dump(contents, f) - return outfile - - -def test_run_collect(): - """Test run and collect.""" - infile = Path(__file__).parent / 'data' / 'gmd_mean.yaml' - with tempfile.TemporaryDirectory() as _tmpdir: - runfile = _modify_path(_tmpdir, infile) - run_args = [runfile.as_posix(), '--verbose', 'debug', '--element', - 'sub-01', '--element', 'sub-02', '--element', 'sub-03'] - response = runner.invoke(run, run_args) - assert response.exit_code == 0 - - # TODO: Check that there are 3 files in the output directory - - collect_args = [runfile.as_posix(), '--verbose', 'debug'] - response = runner.invoke(collect, collect_args) - assert response.exit_code == 0 - - # TODO: Check that there are 4 files in the output directory - # TODO: Check that the collected file has the correct number of rows - - run_args = [runfile.as_posix(), '--verbose', 'debug', '--element', - 'sub-01', '--element', 'sub-02', '--element', 'sub-04'] - response = runner.invoke(run, run_args) - assert response.exit_code == 0 - - # TODO: Check that there are 5 files in the output directory - # TODO: Check that the collected file has 3 rows (sub-04 is not there) - - collect_args = [runfile.as_posix(), '--verbose', 'debug'] - response = runner.invoke(collect, collect_args) - assert response.exit_code == 0 - - # TODO: Check that there are 5 files in the output directory - # TODO: Check that the collected file has 4 rows (sub-04 is there) + # Run command arguments + run_args = [ + str(outfile.absolute()), + "--verbose", + "debug", + "--element", + elements[0], + "--element", + elements[1], + "--element", + elements[2], + ] + # Invoke run command + run_result = runner.invoke(run, run_args) + # Check + assert run_result.exit_code == 0 + # Collect command arguments + collect_args = [str(outfile.absolute()), "--verbose", "debug"] + # Invoke collect command + collect_result = runner.invoke(collect, collect_args) + # Check + assert collect_result.exit_code == 0 -- 2.52.0 From 2c1ce69aee541682107281c5bf5c71991b53c5d5 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Sun, 14 Aug 2022 15:02:01 +0200 Subject: [PATCH 226/287] refactor: improve imports for api sub-package --- junifer/api/__init__.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/junifer/api/__init__.py b/junifer/api/__init__.py index ecd3d2c69..9961e2a4a 100644 --- a/junifer/api/__init__.py +++ b/junifer/api/__init__.py @@ -1,2 +1,8 @@ +"""Provide imports for api sub-package.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + +from .cli import cli from .functions import run, collect -from .import cli \ No newline at end of file -- 2.52.0 From 6089c2a36753e79414ef024b660023e32e20e311 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 22 Aug 2022 15:01:20 +0200 Subject: [PATCH 227/287] refactor: add docstrings and type annotations in test_main.py --- junifer/tests/test_main.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/junifer/tests/test_main.py b/junifer/tests/test_main.py index ea7f2b76c..70943eb00 100644 --- a/junifer/tests/test_main.py +++ b/junifer/tests/test_main.py @@ -1,11 +1,12 @@ -"""Provide tests for package.""" +"""Provide tests for junifer package.""" # Authors: Federico Raimondo # Leonard Sasse +# Synchon Mandal # License: AGPL -def test_import(): +def test_import() -> None: """Test junifer import.""" import junifer print(junifer.__version__) -- 2.52.0 From 6bd8da72a92aac239d83b195e1f44f8e5e711a01 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 2 Sep 2022 15:19:15 +0200 Subject: [PATCH 228/287] fix: rename test files to ease pytest file search --- junifer/markers/tests/{test_base.py => test_markers_base.py} | 0 junifer/storage/tests/{test_base.py => test_storage_base.py} | 0 2 files changed, 0 insertions(+), 0 deletions(-) rename junifer/markers/tests/{test_base.py => test_markers_base.py} (100%) rename junifer/storage/tests/{test_base.py => test_storage_base.py} (100%) diff --git a/junifer/markers/tests/test_base.py b/junifer/markers/tests/test_markers_base.py similarity index 100% rename from junifer/markers/tests/test_base.py rename to junifer/markers/tests/test_markers_base.py diff --git a/junifer/storage/tests/test_base.py b/junifer/storage/tests/test_storage_base.py similarity index 100% rename from junifer/storage/tests/test_base.py rename to junifer/storage/tests/test_storage_base.py -- 2.52.0 From 1834d8deae1f3c7470885d164903abf7a4e2785c Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 2 Sep 2022 17:23:31 +0200 Subject: [PATCH 229/287] chore: fix codespell checks --- junifer/api/tests/test_registry.py | 4 ++-- junifer/storage/base.py | 2 +- junifer/storage/sqlite.py | 8 ++++---- junifer/storage/tests/test_sqlite.py | 2 +- junifer/testing/datagrabbers.py | 4 ++-- 5 files changed, 10 insertions(+), 10 deletions(-) diff --git a/junifer/api/tests/test_registry.py b/junifer/api/tests/test_registry.py index bb0382093..ff4b9264c 100644 --- a/junifer/api/tests/test_registry.py +++ b/junifer/api/tests/test_registry.py @@ -21,7 +21,7 @@ def test_register_invalid_step(): register(step="foo", name="bar", klass="baz") -# TODO: improve paramterization +# TODO: improve parametrization @pytest.mark.parametrize( "step, name, klass", [ @@ -89,7 +89,7 @@ def test_get_class_invalid_name() -> None: get_class(step="datagrabber", name="foo") -# TODO: enable paramterization +# TODO: enable parametrization def test_get_class(): """Test get class.""" # Register datagrabber diff --git a/junifer/storage/base.py b/junifer/storage/base.py index 39e43e1e9..9130eeb02 100644 --- a/junifer/storage/base.py +++ b/junifer/storage/base.py @@ -194,7 +194,7 @@ class BaseFeatureStorage(ABC): columns : list or tuple of str, optional The columns (default None). rows_col_name : str, optional - The column name ot use in case number of rows greater than 1. + The column name to use in case number of rows greater than 1. If None and number of rows greater than 1, then the name will be "index" (default None). diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py index c5a5690d0..79b9457d7 100644 --- a/junifer/storage/sqlite.py +++ b/junifer/storage/sqlite.py @@ -223,7 +223,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): columns : list or tuple of str, optional The columns (default None). rows_col_name : str, optional - The column name ot use in case number of rows greater than 1. + The column name to use in case number of rows greater than 1. If None and number of rows greater than 1, then the name will be "index" (default None). @@ -313,7 +313,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): ------ ValueError If parameter values are invalid or feature is not found or - mulitple features are found. + multiple features are found. """ # Get sqlalchemy engine @@ -417,7 +417,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): col_names : list or tuple of str, optional The column names (default None). rows_col_name : str, optional - The column name ot use in case number of rows greater than 1. + The column name to use in case number of rows greater than 1. If None and number of rows greater than 1, then the name will be "index" (default None). @@ -445,7 +445,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): columns : list or tuple of str, optional The columns (default None). rows_col_name : str, optional - The column name ot use in case number of rows greater than 1. + The column name to use in case number of rows greater than 1. If None and number of rows greater than 1, then the name will be "index" (default None). diff --git a/junifer/storage/tests/test_sqlite.py b/junifer/storage/tests/test_sqlite.py index e5c251553..05f7993f7 100644 --- a/junifer/storage/tests/test_sqlite.py +++ b/junifer/storage/tests/test_sqlite.py @@ -241,7 +241,7 @@ def test_upsert_update(tmp_path: Path) -> None: def test_upsert_invalid_option(tmp_path: Path) -> None: - """Test dataframe store wtih invalid option for upsert. + """Test dataframe store with invalid option for upsert. Parameters ---------- diff --git a/junifer/testing/datagrabbers.py b/junifer/testing/datagrabbers.py index 3c1e8a150..9d0f29415 100644 --- a/junifer/testing/datagrabbers.py +++ b/junifer/testing/datagrabbers.py @@ -26,8 +26,8 @@ class OasisVBMTestingDatagrabber(BaseDataGrabber): def __getitem__(self, element: str) -> Dict: """Implement indexing support. - Paramters - --------- + Parameters + ---------- element : str The element to retrieve. -- 2.52.0 From 4f97b9bf790a749e9d7ee9d8e454abcb0f7a0798 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 2 Sep 2022 17:24:43 +0200 Subject: [PATCH 230/287] chore: fix isort and black checks --- junifer/datagrabber/tests/test_datalad_base.py | 1 + junifer/datagrabber/tests/test_pattern.py | 6 +++--- junifer/tests/test_main.py | 1 + junifer/utils/tests/test_logging.py | 2 +- setup.py | 1 + 5 files changed, 7 insertions(+), 4 deletions(-) diff --git a/junifer/datagrabber/tests/test_datalad_base.py b/junifer/datagrabber/tests/test_datalad_base.py index b05ff4337..2197fc097 100644 --- a/junifer/datagrabber/tests/test_datalad_base.py +++ b/junifer/datagrabber/tests/test_datalad_base.py @@ -13,6 +13,7 @@ def test_datalad_base_abstractness() -> None: with pytest.raises(TypeError, match=r"abstract"): DataladDataGrabber() + # def test_datalad_base_missing_uri() -> None: # """Test proper check of missing URI in datalad base initialization.""" # with pytest.raises(ValueError, match=r"`uri` must be provided"): diff --git a/junifer/datagrabber/tests/test_pattern.py b/junifer/datagrabber/tests/test_pattern.py index 32eb09757..48b08bc53 100644 --- a/junifer/datagrabber/tests/test_pattern.py +++ b/junifer/datagrabber/tests/test_pattern.py @@ -5,12 +5,12 @@ # Synchon Mandal # License: AGPL -from junifer.datagrabber.pattern import PatternDataGrabber - from pathlib import Path import pytest +from junifer.datagrabber.pattern import PatternDataGrabber + def test_PatternDataGrabber() -> None: """Test PatternDataGrabber.""" @@ -62,7 +62,7 @@ def test_PatternDataGrabber() -> None: ) with pytest.raises( - ValueError, match=r"`patterns` must contain all `types`" + ValueError, match=r"`patterns` must contain all `types`" ): MyDataGrabber( datadir="/tmp", diff --git a/junifer/tests/test_main.py b/junifer/tests/test_main.py index 70943eb00..e29855087 100644 --- a/junifer/tests/test_main.py +++ b/junifer/tests/test_main.py @@ -9,4 +9,5 @@ def test_import() -> None: """Test junifer import.""" import junifer + print(junifer.__version__) diff --git a/junifer/utils/tests/test_logging.py b/junifer/utils/tests/test_logging.py index c4ab534bc..a815f8658 100644 --- a/junifer/utils/tests/test_logging.py +++ b/junifer/utils/tests/test_logging.py @@ -11,11 +11,11 @@ from pathlib import Path import pytest from junifer.utils.logging import ( - logger, _close_handlers, configure_logging, get_versions, log_versions, + logger, raise_error, warn_with_log, ) diff --git a/setup.py b/setup.py index c6e4450ca..e0709284e 100644 --- a/setup.py +++ b/setup.py @@ -7,5 +7,6 @@ from setuptools import setup + if __name__ == "__main__": setup() -- 2.52.0 From b858c3fd8399646e39241956b47c94236d3e4e11 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 2 Sep 2022 17:37:42 +0200 Subject: [PATCH 231/287] chore: fix flake8 checks --- tox.ini | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tox.ini b/tox.ini index 4a7812cdd..3eb2b751f 100644 --- a/tox.ini +++ b/tox.ini @@ -84,7 +84,8 @@ known_third_party = exclude = __init__.py max-line-length = 79 -ignore = +extend-ignore = + B024 # abstract class with no abstract methods D202 E201 # whitespace after ‘(’ E202 # whitespace before ‘)’ -- 2.52.0 From 402f00c48ebccf624d1c7c69faad63930a12d80e Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Sat, 3 Sep 2022 11:51:37 +0200 Subject: [PATCH 232/287] refactor: add docstrings and type annotations in junifer/stats.py --- junifer/stats.py | 64 +++++++++++++++++++++++++++++++----------------- 1 file changed, 41 insertions(+), 23 deletions(-) diff --git a/junifer/stats.py b/junifer/stats.py index a4ea428c2..e0a30c0ce 100644 --- a/junifer/stats.py +++ b/junifer/stats.py @@ -1,13 +1,20 @@ """Provide functions for statistics.""" +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + from functools import partial +from typing import Callable, Dict + import numpy as np -from scipy.stats.mstats import winsorize from scipy.stats import trim_mean +from scipy.stats.mstats import winsorize + from .utils import logger, raise_error -def get_aggfunc_by_name(name, func_params): +def get_aggfunc_by_name(name: str, func_params: Dict) -> Callable: """Get an aggregation function by its name. Parameters @@ -15,62 +22,73 @@ def get_aggfunc_by_name(name, func_params): name : str Name to identify the function. Currently supported names and corresponding functions are: - 'winsorized_mean' -> scipy.stats.mstats.winsorize - 'mean' -> np.mean - 'std' -> np.std - 'trim_mean' -> scipy.stats.trim_mean - + - 'winsorized_mean' -> scipy.stats.mstats.winsorize + - 'mean' -> numpy.mean + - 'std' -> numpy.std + - 'trim_mean' -> scipy.stats.trim_mean func_params : dict Parameters to pass to the function. E.g. for 'winsorized_mean': func_params = {'limits': [0.1, 0.1]} Returns ------- - func : function + function Respective function with `func_params` parameter set. + """ # check validity of names - _valid_func_names = {'winsorized_mean', 'mean', 'std', 'trim_mean'} + _valid_func_names = {"winsorized_mean", "mean", "std", "trim_mean"} # apply functions - if name == 'winsorized_mean': + if name == "winsorized_mean": # check validity of func_params - limits = func_params.get('limits') + limits = func_params.get("limits") if all((lim >= 0.0 and lim <= 1) for lim in limits): - logger.info(f'Limits for winsorized mean are set to {limits}.') + logger.info(f"Limits for winsorized mean are set to {limits}.") else: raise_error( - 'Limits for the winsorized mean must be between 0 and 1.') + "Limits for the winsorized mean must be between 0 and 1." + ) # partially interpret func_params func = partial(winsorized_mean, **func_params) - elif name == 'mean': + elif name == "mean": func = np.mean - elif name == 'std': + elif name == "std": func = np.std - elif name == 'trim_mean': - func = partial(trim_mean, **func_params) + elif name == "trim_mean": + if func_params is None: + func = trim_mean + else: + func = partial(trim_mean, **func_params) else: - raise_error(f'Function {name} unknown. Please provide any of ' - f'{_valid_func_names}') + raise_error( + f"Function {name} unknown. Please provide any of " + f"{_valid_func_names}" + ) return func -def winsorized_mean(data, axis=None, **win_params): +def winsorized_mean( + data: np.ndarray, axis: int = None, **win_params +) -> np.ndarray: """Compute a winsorized mean by chaining winsorization and mean. Parameters ---------- - data : array + data : numpy.ndarray Data to calculate winsorized mean on. - win_params : dict + axis : int, optional + The axis to calculate winsorized mean on (default None). + **win_params : dict Dictionary containing the keyword arguments for the winsorize function. E.g. {'limits': [0.1, 0.1]} Returns ------- - win_mean : np.ndarray + numpy.ndarray Winsorized mean of the inputted data with the winsorize settings applied as specified in win_params. + """ win_dat = winsorize(data, axis=axis, **win_params) win_mean = win_dat.mean(axis=axis) -- 2.52.0 From 9758a334945dc8505ae675d5793c4766c417c8de Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Sat, 3 Sep 2022 13:23:34 +0200 Subject: [PATCH 233/287] update: add partial unit tests for junifer/stats.py --- junifer/tests/test_stats.py | 30 ++++++++++++++++++++++++++++++ 1 file changed, 30 insertions(+) create mode 100644 junifer/tests/test_stats.py diff --git a/junifer/tests/test_stats.py b/junifer/tests/test_stats.py new file mode 100644 index 000000000..cd1daa60d --- /dev/null +++ b/junifer/tests/test_stats.py @@ -0,0 +1,30 @@ +"""Provide tests for stats.""" + +# Authors: Synchon Mandal +# License: AGPL + +from typing import Dict, Optional + +import pytest + +from junifer.stats import get_aggfunc_by_name + + +@pytest.mark.parametrize( + "name, params", + [ + ("winsorized_mean", {"limits": [0.2, 0.7]}), + ("mean", None), + ("std", None), + ("trim_mean", None), + ], +) +def test_get_aggfunc_by_name(name: str, params: Optional[Dict]) -> None: + """Test aggregation function retrieval by name.""" + get_aggfunc_by_name(name=name, func_params=params) + + +@pytest.mark.skip(reason="test not implemented") +def test_winsorized_mean() -> None: + """Test winsorized mean computation.""" + ... -- 2.52.0 From 7b6c743c6d1658e96c9f856fe3438a658d7b9493 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Sat, 3 Sep 2022 13:33:12 +0200 Subject: [PATCH 234/287] update: add black config in pyproject.toml --- pyproject.toml | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/pyproject.toml b/pyproject.toml index 49bd06177..38d20190e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -79,3 +79,7 @@ packages = ["junifer"] version_scheme = "python-simplified-semver" local_scheme = "no-local-version" write_to = "junifer/_version.py" + +[tool.black] +line-length = 79 +target-version = ["py38"] \ No newline at end of file -- 2.52.0 From 35275c6a7c5f96a62e896d58cb5f63872c3558cd Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Sat, 3 Sep 2022 13:33:31 +0200 Subject: [PATCH 235/287] update: add black env in tox.ini --- tox.ini | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/tox.ini b/tox.ini index 3eb2b751f..10f890c2e 100644 --- a/tox.ini +++ b/tox.ini @@ -1,5 +1,5 @@ [tox] -envlist = isort, flake8, test, coverage, codespell, py3{8,9,10} +envlist = isort, black, flake8, test, coverage, codespell, py3{8,9,10} isolated_build = true [gh-actions] @@ -25,6 +25,13 @@ deps = commands = isort --check-only --diff {toxinidir}/junifer {toxinidir}/setup.py +[testenv:black] +skip_install = true +deps = + black +commands = + black --check --diff {toxinidir}/junifer {toxinidir}/setup.py + [testenv:flake8] skip_install = true deps = -- 2.52.0 From b8a2c5ae95b0d7514ea20ac20b7bc621a3c71842 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Sat, 3 Sep 2022 15:44:12 +0200 Subject: [PATCH 236/287] refactor: rename test_atlas.py to test_atlases.py to match module name --- junifer/data/tests/{test_atlas.py => test_atlases.py} | 0 1 file changed, 0 insertions(+), 0 deletions(-) rename junifer/data/tests/{test_atlas.py => test_atlases.py} (100%) diff --git a/junifer/data/tests/test_atlas.py b/junifer/data/tests/test_atlases.py similarity index 100% rename from junifer/data/tests/test_atlas.py rename to junifer/data/tests/test_atlases.py -- 2.52.0 From 748dfc235d9ffef6e10ca5fabdba58dd2eab3c3c Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 2 Sep 2022 14:51:35 +0200 Subject: [PATCH 237/287] refactor: add docstrings and type annotations in data/atlases.py --- junifer/data/atlases.py | 712 +++++++++++++++++++++++++--------------- 1 file changed, 446 insertions(+), 266 deletions(-) diff --git a/junifer/data/atlases.py b/junifer/data/atlases.py index d429bbf82..71e327e92 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,252 @@ 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 = ( + 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 +658,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 -- 2.52.0 From e4ff6649e57e07ccc9c72e8b887081c0100be415 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 5 Sep 2022 14:48:29 +0200 Subject: [PATCH 238/287] refactor: prune, add docstrings and type annotations in test_atlases.py --- junifer/data/tests/test_atlases.py | 622 ++++++++++++++++++++--------- 1 file changed, 440 insertions(+), 182 deletions(-) diff --git a/junifer/data/tests/test_atlases.py b/junifer/data/tests/test_atlases.py index 0e4cd17e5..6b9eea365 100644 --- a/junifer/data/tests/test_atlases.py +++ b/junifer/data/tests/test_atlases.py @@ -1,223 +1,481 @@ """Provide tests for atlas.""" -import tempfile -import pytest +# Authors: Federico Raimondo +# Vera Komeyer +# Synchon Mandal +# License: AGPL + from pathlib import Path -from numpy.testing import assert_array_equal, assert_array_almost_equal +from typing import List + +import pytest +from numpy.testing import assert_array_almost_equal, assert_array_equal from junifer.data.atlases import ( - register_atlas, list_atlases, - load_atlas, + _retrieve_atlas, _retrieve_schaefer, _retrieve_suit, - _retrieve_atlas, _retrieve_tian, + list_atlases, + load_atlas, + register_atlas, ) -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']) - +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('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 + register_atlas( + name="SUITxSUIT", + atlas_path="testatlas.nii.gz", + atl_labels=["1", "2", "3"], + overwrite=True, + ) -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.""" - +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}' + 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' + # Define atlas file names + 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') + # Load atlas + img, lbl, fname = load_atlas( + name="Schaefer100x7", atlas_dir=str(tmp_path.absolute()) + ) + # Check atlas values assert img is not None - home_dir = Path().home() / 'junifer' / 'data' / 'atlas' - assert home_dir in fname.parents # type: ignore + 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_suit(): - """Test SUIT atlas.""" +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 + 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]) + # 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]) - 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]) + # 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]) - 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') + # 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_tian(): - """Test TIAN atlas.""" +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 + 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]) - 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 +@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. - 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]) + 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. - 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]) + """ + 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]) - 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]) +@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. - with pytest.raises(ValueError, match=r"The parameter `space`"): - _retrieve_tian(tmpdir, resolution=1, scale=1, space='wrong') + 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. - with pytest.raises(ValueError, match=r"The parameter `magneticfield`"): - _retrieve_tian( - tmpdir, resolution=1, scale=1, magneticfield='wrong') + """ + 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", + ) -- 2.52.0 From c26237f3c72766164122734bf2ae3346c06166b3 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 5 Sep 2022 14:48:52 +0200 Subject: [PATCH 239/287] refactor: improve imports for data sub-package --- junifer/data/__init__.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) 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 -- 2.52.0 From 19ef4b9e68fe5c134efdaa9ce7197eec4e9d6f24 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Tue, 6 Sep 2022 10:39:43 +0200 Subject: [PATCH 240/287] Fix for tian test + testing import --- junifer/data/atlases.py | 1 + junifer/testing/__init__.py | 2 +- 2 files changed, 2 insertions(+), 1 deletion(-) diff --git a/junifer/data/atlases.py b/junifer/data/atlases.py index 71e327e92..7fe70a46a 100644 --- a/junifer/data/atlases.py +++ b/junifer/data/atlases.py @@ -559,6 +559,7 @@ def _retrieve_tian( ) 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" 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 -- 2.52.0 From 4dd848dab103159451f7af25588f598f5acce8a9 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Tue, 6 Sep 2022 10:52:46 +0200 Subject: [PATCH 241/287] chore: add module docstring for junifer package import --- junifer/__init__.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/junifer/__init__.py b/junifer/__init__.py index b6d36c58b..e8b301aad 100644 --- a/junifer/__init__.py +++ b/junifer/__init__.py @@ -1,3 +1,9 @@ +"""Provide imports for junifer package.""" + +# Authors: Federico Raimondo +# Synchon Mandal +# License: AGPL + from ._version import __version__ from . import api from . import utils -- 2.52.0 From d563eed2a20ad6c429c4e39e363a71f99f86fb50 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Thu, 8 Sep 2022 10:22:34 +0200 Subject: [PATCH 242/287] Fix datagrabbers + add basic test for multiple datagrabber --- junifer/datagrabber/__init__.py | 2 +- junifer/datagrabber/base.py | 2 +- .../{multiple_base.py => multiple.py} | 24 +++- junifer/datagrabber/pattern.py | 51 ++++----- junifer/datagrabber/pattern_datalad_base.py | 4 +- junifer/datagrabber/tests/test_multiple.py | 89 +++++++++++++++ junifer/datagrabber/tests/test_pattern.py | 24 ++-- .../tests/test_pattern_datalad_base.py | 106 +++++++++--------- junifer/testing/datagrabbers.py | 3 +- junifer/utils/logging.py | 4 +- 10 files changed, 204 insertions(+), 105 deletions(-) rename junifer/datagrabber/{multiple_base.py => multiple.py} (78%) create mode 100644 junifer/datagrabber/tests/test_multiple.py diff --git a/junifer/datagrabber/__init__.py b/junifer/datagrabber/__init__.py index c224480b4..8f0872fce 100644 --- a/junifer/datagrabber/__init__.py +++ b/junifer/datagrabber/__init__.py @@ -7,6 +7,6 @@ from .base import BaseDataGrabber from .datalad_base import DataladDataGrabber from .hcp import DataladHCP1200, HCP1200 -from .multiple_base import MultipleDataGrabber +from .multiple import MultipleDataGrabber from .pattern import PatternDataGrabber from .pattern_datalad_base import PatternDataladDataGrabber diff --git a/junifer/datagrabber/base.py b/junifer/datagrabber/base.py index a4908977b..8f7065e46 100644 --- a/junifer/datagrabber/base.py +++ b/junifer/datagrabber/base.py @@ -59,7 +59,7 @@ class BaseDataGrabber(ABC): yield elem # TODO: element does nothing, check - def __getitem__(self, element: Union[str, Tuple]) -> Dict[str, Path]: + def __getitem__(self, element: Union[str, Tuple]) -> Dict[str, Dict]: """Enable indexing support. Parameters diff --git a/junifer/datagrabber/multiple_base.py b/junifer/datagrabber/multiple.py similarity index 78% rename from junifer/datagrabber/multiple_base.py rename to junifer/datagrabber/multiple.py index 6e0e28daf..606fc82bc 100644 --- a/junifer/datagrabber/multiple_base.py +++ b/junifer/datagrabber/multiple.py @@ -12,10 +12,10 @@ from .base import BaseDataGrabber class MultipleDataGrabber(BaseDataGrabber): - """Abstract base class for data fetching from multiple sources. + """Datagrabber class for data fetching from multiple sources. Defines a DataGrabber which can be used to fetch data from multiple - sources. + datagrabbers. Parameters ---------- @@ -61,13 +61,30 @@ class MultipleDataGrabber(BaseDataGrabber): """Implement context entry.""" for dg in self._datagrabbers: dg.__enter__() + return self def __exit__(self, exc_type, exc_value, exc_traceback) -> None: """Implement context exit.""" for dg in self._datagrabbers: dg.__exit__(exc_type, exc_value, exc_traceback) - def get_types(self) -> List[List[str]]: + def get_elements(self) -> List: + """Get elements. + + Returns + ------- + elements : list + The list of elements that can be grabbed in the dataset. It + corresponds to the elements that are present in all the + related datagrabbers. + """ + all_elements = [dg.get_elements() for dg in self._datagrabbers] + elements = set(all_elements[0]) + for s in all_elements[1:]: + elements.intersection_update(s) + return list(elements) + + def get_types(self) -> List[str]: """Get types. Returns @@ -91,3 +108,4 @@ class MultipleDataGrabber(BaseDataGrabber): t_meta = {} t_meta["class"] = self.__class__.__name__ t_meta["datagrabbers"] = [dg.get_meta() for dg in self._datagrabbers] + return t_meta diff --git a/junifer/datagrabber/pattern.py b/junifer/datagrabber/pattern.py index 137b535f0..ad1690ba5 100644 --- a/junifer/datagrabber/pattern.py +++ b/junifer/datagrabber/pattern.py @@ -6,8 +6,7 @@ # License: AGPL import re -from pathlib import Path -from typing import Dict, List, Optional, Tuple, Union +from typing import Dict, List, Tuple, Union from ..api.decorators import register_datagrabber from ..utils import logger, raise_error @@ -23,13 +22,12 @@ class PatternDataGrabber(BaseDataGrabber): Parameters ---------- - types : list of str, optional + types : list of str The types of data to be grabbed (default None). - patterns : dict, optional + patterns : dict Patterns for each type of data as a dictionary. The keys are the types and the values are the patterns. Each occurrence of the string - `{subject}` in the pattern will be replaced by the indexed element - (default None). + `{subject}` in the pattern will be replaced by the indexed element. replacements: list of str Replacements in the patterns for each item in the "element" tuple. datadir : str or pathlib.Path @@ -45,9 +43,9 @@ class PatternDataGrabber(BaseDataGrabber): def __init__( self, - types: Optional[List[str]] = None, - patterns: Optional[Dict[str, str]] = None, - replacements: Optional[List[str]] = None, + types: List[str], + patterns: Dict[str, str], + replacements: List[str], **kwargs, ) -> None: """Initialize the class.""" @@ -84,24 +82,20 @@ class PatternDataGrabber(BaseDataGrabber): The search pattern to be used with glob. """ - # re_pattern = pattern - # glob_pattern = pattern + re_pattern = pattern + glob_pattern = pattern for t_r in self.replacements: - # Replace the first appearance of each with a named group - # definition and the second appearance of each with the named group + # Replace the first of each with a named group definition + re_pattern = re_pattern.replace( + f"{{{t_r}}}", f"(?P<{t_r}>.*)", 1) + + for t_r in self.replacements: + # Replace the second appearance of each with the named group # back reference - re_pattern = pattern.replace( - f"{{{t_r}}}", f"(?P<{t_r}>.*)", 1 - ).replace(f"{{{t_r}}}", f"(?P={t_r})") - glob_pattern = pattern.replace(f"{{{t_r}}}", "*") + re_pattern = re_pattern.replace(f"{{{t_r}}}", f"(?P={t_r})") - # for t_r in self.replacements: - # # Replace the second appearance of each with the named group - # # back reference - # re_pattern = pattern.replace(f"{{{t_r}}}", f"(?P={t_r})") - - # for t_r in self.replacements: - # glob_pattern = glob_pattern.replace(f"{{{t_r}}}", "*") + for t_r in self.replacements: + glob_pattern = glob_pattern.replace(f"{{{t_r}}}", "*") return re_pattern, glob_pattern def _replace_patterns_glob(self, element: Tuple, pattern: str) -> str: @@ -128,7 +122,7 @@ class PatternDataGrabber(BaseDataGrabber): to_replace = dict(zip(self.replacements, element)) return pattern.format(**to_replace) - def __getitem__(self, element: Union[str, Tuple]) -> Dict[str, Path]: + def __getitem__(self, element: Union[str, Tuple]) -> Dict[str, Dict]: """Implement single element indexing in the database. Each occurrence of the strings in "replacements" is replaced by the @@ -145,7 +139,7 @@ class PatternDataGrabber(BaseDataGrabber): Returns ------- dict - Dictionary of paths for each type of data required for the + Dictionary of dictionaries for each type of data required for the specified element. """ @@ -194,7 +188,7 @@ class PatternDataGrabber(BaseDataGrabber): # Replace the pattern re_pattern, glob_pattern = self._replace_patterns_regex(t_pattern) for fname in self.datadir.glob(glob_pattern): - suffix = str(fname.relative_to(self.datadir).absolute()) + suffix = fname.relative_to(self.datadir).as_posix() m = re.match(re_pattern, suffix) if m is not None: t_element = tuple(m.group(k) for k in self.replacements) @@ -206,5 +200,6 @@ class PatternDataGrabber(BaseDataGrabber): elements = types_element else: elements = elements.intersection(types_element) - + if elements is None: + elements = set() return list(elements) diff --git a/junifer/datagrabber/pattern_datalad_base.py b/junifer/datagrabber/pattern_datalad_base.py index b36aed865..5a43e6ec6 100644 --- a/junifer/datagrabber/pattern_datalad_base.py +++ b/junifer/datagrabber/pattern_datalad_base.py @@ -1,4 +1,4 @@ -"""Provide abstract base class for pattern-based datalad datagrabber.""" +"""Provide base class for pattern-based datalad datagrabber.""" # Authors: Federico Raimondo # Leonard Sasse @@ -15,7 +15,7 @@ from .utils import validate_patterns @register_datagrabber class PatternDataladDataGrabber(DataladDataGrabber, PatternDataGrabber): - """Abstract base class for pattern-based data fetching via Datalad. + """Base class for pattern-based data fetching via Datalad. Defines a DataGrabber that gets data from a datalad sibling, interpreting patterns. diff --git a/junifer/datagrabber/tests/test_multiple.py b/junifer/datagrabber/tests/test_multiple.py new file mode 100644 index 000000000..f0db96f4b --- /dev/null +++ b/junifer/datagrabber/tests/test_multiple.py @@ -0,0 +1,89 @@ +"""Provide tests for multiple.""" + +# Authors: Federico Raimondo +# License: AGPL + +from junifer.datagrabber import PatternDataladDataGrabber, MultipleDataGrabber + +_testing_dataset = { + "example_bids": { + "uri": "https://gin.g-node.org/juaml/datalad-example-bids", + "id": "e2ce149bd723088769a86c72e57eded009258c6b", + }, + "example_bids_ses": { + "uri": "https://gin.g-node.org/juaml/datalad-example-bids-ses", + "id": "3d08d55d1faad4f12ab64ac9497544a0d924d47a", + }, +} + + +def test_multiple() -> None: + + repo_uri = _testing_dataset["example_bids_ses"]["uri"] + rootdir = "example_bids_ses" + replacements = ["subject", "session"] + pattern1 = { + "T1w": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz", + } + pattern2 = { + "bold": "{subject}/{session}/func/" + "{subject}_{session}_task-rest_bold.nii.gz", + } + dg1 = PatternDataladDataGrabber( + rootdir=rootdir, + uri=repo_uri, + types=["T1w"], + patterns=pattern1, + replacements=replacements) + + dg2 = PatternDataladDataGrabber( + rootdir=rootdir, + uri=repo_uri, + types=["bold"], + patterns=pattern2, + replacements=replacements) + + dg = MultipleDataGrabber([dg1, dg2]) + expected_subs = [ + (f"sub-{i:02d}", f"ses-{j:02d}") + for j in range(1, 3) + for i in range(1, 10) + ] + + with dg: + subs = [x for x in dg] + assert set(subs) == set(expected_subs) + + +def test_multiple_no_intersection() -> None: + + repo_uri1 = _testing_dataset["example_bids"]["uri"] + repo_uri2 = _testing_dataset["example_bids_ses"]["uri"] + rootdir = "example_bids_ses" + replacements = ["subject", "session"] + pattern1 = { + "T1w": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz", + } + pattern2 = { + "bold": "{subject}/{session}/func/" + "{subject}_{session}_task-rest_bold.nii.gz", + } + dg1 = PatternDataladDataGrabber( + rootdir=rootdir, + uri=repo_uri1, + types=["T1w"], + patterns=pattern1, + replacements=replacements) + + dg2 = PatternDataladDataGrabber( + rootdir=rootdir, + uri=repo_uri2, + types=["bold"], + patterns=pattern2, + replacements=replacements) + + dg = MultipleDataGrabber([dg1, dg2]) + expected_subs = set() + with dg: + subs = [x for x in dg] + assert set(subs) == set(expected_subs) diff --git a/junifer/datagrabber/tests/test_pattern.py b/junifer/datagrabber/tests/test_pattern.py index 48b08bc53..530b7e447 100644 --- a/junifer/datagrabber/tests/test_pattern.py +++ b/junifer/datagrabber/tests/test_pattern.py @@ -14,13 +14,9 @@ from junifer.datagrabber.pattern import PatternDataGrabber def test_PatternDataGrabber() -> None: """Test PatternDataGrabber.""" - # Create concrete class - class MyDataGrabber(PatternDataGrabber): - def get_elements(self): - return super().get_elements() with pytest.raises(TypeError, match=r"`types` must be a list"): - MyDataGrabber( + PatternDataGrabber( datadir="/tmp", types="wrong", patterns={"wrong": "pattern"}, @@ -28,7 +24,7 @@ def test_PatternDataGrabber() -> None: ) with pytest.raises(TypeError, match=r"`types` must be a list of strings"): - MyDataGrabber( + PatternDataGrabber( datadir="/tmp", types=[1, 2, 3], patterns={"1": "pattern", "2": "pattern", "3": "pattern"}, @@ -36,7 +32,7 @@ def test_PatternDataGrabber() -> None: ) with pytest.raises(ValueError, match=r"must have the same length"): - MyDataGrabber( + PatternDataGrabber( datadir="/tmp", types=["func", "anat"], patterns={"1": "pattern", "2": "pattern", "3": "pattern"}, @@ -44,7 +40,7 @@ def test_PatternDataGrabber() -> None: ) with pytest.raises(TypeError, match=r"`patterns` must be a dict"): - MyDataGrabber( + PatternDataGrabber( datadir="/tmp", types=["func", "anat"], patterns="wrong", @@ -54,7 +50,7 @@ def test_PatternDataGrabber() -> None: with pytest.raises( ValueError, match=r"`patterns` must have the same length" ): - MyDataGrabber( + PatternDataGrabber( datadir="/tmp", types=["func", "anat"], patterns={"wrong": "pattern"}, @@ -64,7 +60,7 @@ def test_PatternDataGrabber() -> None: with pytest.raises( ValueError, match=r"`patterns` must contain all `types`" ): - MyDataGrabber( + PatternDataGrabber( datadir="/tmp", types=["func", "anat"], patterns={"wrong": "pattern", "func": "pattern"}, @@ -72,7 +68,7 @@ def test_PatternDataGrabber() -> None: ) with pytest.raises(TypeError, match=r"must be a list of strings"): - MyDataGrabber( + PatternDataGrabber( datadir="/tmp", types=["func", "anat"], patterns={"func": "func/test", "anat": "anat/test"}, @@ -80,7 +76,7 @@ def test_PatternDataGrabber() -> None: ) with pytest.warns(RuntimeWarning, match=r"not part of any pattern"): - MyDataGrabber( + PatternDataGrabber( datadir="/tmp", types=["func", "anat"], patterns={ @@ -90,7 +86,7 @@ def test_PatternDataGrabber() -> None: replacements=["subject", "wrong"], ) - datagrabber = MyDataGrabber( + datagrabber = PatternDataGrabber( datadir="/tmp/data", types=["func", "anat"], patterns={"func": "func/{subject}.nii", "anat": "anat/{subject}.nii"}, @@ -100,7 +96,7 @@ def test_PatternDataGrabber() -> None: assert datagrabber.types == ["func", "anat"] assert datagrabber.replacements == ["subject"] - datagrabber = MyDataGrabber( + datagrabber = PatternDataGrabber( datadir=Path("/tmp/data"), types=["func", "anat"], patterns={ diff --git a/junifer/datagrabber/tests/test_pattern_datalad_base.py b/junifer/datagrabber/tests/test_pattern_datalad_base.py index 3def6d295..0d445c9e8 100644 --- a/junifer/datagrabber/tests/test_pattern_datalad_base.py +++ b/junifer/datagrabber/tests/test_pattern_datalad_base.py @@ -119,60 +119,60 @@ def test_bids_PatternDataladDataGrabber(tmp_path: Path) -> None: # ) -# def test_bids_datalad_PatternDataGrabber_session(): -# """Test a subject and session-based BIDS datalad datagrabber.""" -# types = ["T1w", "bold"] -# patterns = { -# "T1w": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz", -# "bold": "{subject}/{session}/func/" -# "{subject}_{session}_task-rest_bold.nii.gz", -# } -# replacements = ["subject", "session"] +def test_bids_PatternDataladDataGrabber_session(): + """Test a subject and session-based BIDS datalad datagrabber.""" + types = ["T1w", "bold"] + patterns = { + "T1w": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz", + "bold": "{subject}/{session}/func/" + "{subject}_{session}_task-rest_bold.nii.gz", + } + replacements = ["subject", "session"] -# with pytest.raises(ValueError, match=r"`uri` must be provided"): -# PatternDataladDataGrabber( -# datadir=None, -# types=types, -# patterns=patterns, -# replacements=replacements, -# ) + with pytest.raises(ValueError, match=r"`uri` must be provided"): + PatternDataladDataGrabber( + datadir=None, + types=types, + patterns=patterns, + replacements=replacements, + ) -# repo_uri = _testing_dataset["example_bids_ses"]["uri"] -# rootdir = "example_bids_ses" -# # repo_commit = _testing_dataset['example_bids_ses']['id'] + repo_uri = _testing_dataset["example_bids_ses"]["uri"] + rootdir = "example_bids_ses" + # repo_commit = _testing_dataset['example_bids_ses']['id'] -# # With T1W and bold, only 2 sessions are available -# with PatternDataladDataGrabber( -# rootdir=rootdir, -# uri=repo_uri, -# types=types, -# patterns=patterns, -# replacements=replacements, -# ) as dg: -# subs = [x for x in dg] -# expected_subs = [ -# (f"sub-{i:02d}", f"ses-{j:02d}") -# for j in range(1, 3) -# for i in range(1, 10) -# ] -# assert set(subs) == set(expected_subs) + # With T1W and bold, only 2 sessions are available + with PatternDataladDataGrabber( + rootdir=rootdir, + uri=repo_uri, + types=types, + patterns=patterns, + replacements=replacements, + ) as dg: + subs = [x for x in dg] + expected_subs = [ + (f"sub-{i:02d}", f"ses-{j:02d}") + for j in range(1, 3) + for i in range(1, 10) + ] + assert set(subs) == set(expected_subs) -# # Test with a different T1w only, it should have 3 sessions -# types = ["T1w"] -# patterns = { -# "T1w": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz", -# } -# with PatternDataladDataGrabber( -# rootdir=rootdir, -# uri=repo_uri, -# types=types, -# patterns=patterns, -# replacements=replacements, -# ) as dg: -# subs = [x for x in dg] -# expected_subs = [ -# (f"sub-{i:02d}", f"ses-{j:02d}") -# for j in range(1, 4) -# for i in range(1, 10) -# ] -# assert set(subs) == set(expected_subs) + # Test with a different T1w only, it should have 3 sessions + types = ["T1w"] + patterns = { + "T1w": "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz", + } + with PatternDataladDataGrabber( + rootdir=rootdir, + uri=repo_uri, + types=types, + patterns=patterns, + replacements=replacements, + ) as dg: + subs = [x for x in dg] + expected_subs = [ + (f"sub-{i:02d}", f"ses-{j:02d}") + for j in range(1, 4) + for i in range(1, 10) + ] + assert set(subs) == set(expected_subs) diff --git a/junifer/testing/datagrabbers.py b/junifer/testing/datagrabbers.py index 9d0f29415..664a7800a 100644 --- a/junifer/testing/datagrabbers.py +++ b/junifer/testing/datagrabbers.py @@ -39,7 +39,8 @@ class OasisVBMTestingDatagrabber(BaseDataGrabber): """ out = super().__getitem__(element) i_sub = int(element.split("-")[1]) - 1 - out["VBM_GM"] = {"path": self._dataset.gray_matter_maps[i_sub]} + out["VBM_GM"] = { + "path": self._dataset.gray_matter_maps[i_sub]} # Set the element accordingly out["meta"]["element"] = {"subject": element} return out diff --git a/junifer/utils/logging.py b/junifer/utils/logging.py index 53431187b..bb4e5627e 100644 --- a/junifer/utils/logging.py +++ b/junifer/utils/logging.py @@ -9,7 +9,7 @@ import sys from distutils.version import LooseVersion from pathlib import Path from subprocess import PIPE, Popen, TimeoutExpired -from typing import Dict, NoReturn, Optional, Union +from typing import Dict, NoReturn, Optional, Type, Union from warnings import warn @@ -257,7 +257,7 @@ def configure_logging( log_versions() # log versions of installed packages -def raise_error(msg: str, klass: Exception = ValueError) -> NoReturn: +def raise_error(msg: str, klass: Type[Exception] = ValueError) -> NoReturn: """Raise error, but first log it. Parameters -- 2.52.0 From b61186b1ce189e28c4a06a9d8815b76571a880c3 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Thu, 8 Sep 2022 11:06:19 +0200 Subject: [PATCH 243/287] Fix docstrings --- junifer/data/atlases.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/junifer/data/atlases.py b/junifer/data/atlases.py index 7fe70a46a..c08f7d10d 100644 --- a/junifer/data/atlases.py +++ b/junifer/data/atlases.py @@ -61,7 +61,7 @@ for scale in range(1, 5): } t_name = f"TianxS{scale}x3TxMNI6thgeneration" _available_atlases[t_name] = { - "family": "Tian", + "family": "Tian",r "scale": scale, "magneticfield": "3T", "space": "MNI6thgeneration", -- 2.52.0 From 6807871f3418f6b81accaacd18112576489f6e75 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Thu, 8 Sep 2022 11:07:23 +0200 Subject: [PATCH 244/287] Fix docstrings --- junifer/datagrabber/tests/test_multiple.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/junifer/datagrabber/tests/test_multiple.py b/junifer/datagrabber/tests/test_multiple.py index f0db96f4b..eedd94965 100644 --- a/junifer/datagrabber/tests/test_multiple.py +++ b/junifer/datagrabber/tests/test_multiple.py @@ -18,7 +18,7 @@ _testing_dataset = { def test_multiple() -> None: - + """Test a multiple datagrabber.""" repo_uri = _testing_dataset["example_bids_ses"]["uri"] rootdir = "example_bids_ses" replacements = ["subject", "session"] @@ -56,7 +56,7 @@ def test_multiple() -> None: def test_multiple_no_intersection() -> None: - + """Test a multiple datagrabber without intersection (0 elements).""" repo_uri1 = _testing_dataset["example_bids"]["uri"] repo_uri2 = _testing_dataset["example_bids_ses"]["uri"] rootdir = "example_bids_ses" -- 2.52.0 From 705b37647a2f5e9c66f441cbfc4d37062e55daa6 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Thu, 8 Sep 2022 11:07:29 +0200 Subject: [PATCH 245/287] Revert "Fix docstrings" This reverts commit b61186b1ce189e28c4a06a9d8815b76571a880c3. --- junifer/data/atlases.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/junifer/data/atlases.py b/junifer/data/atlases.py index c08f7d10d..7fe70a46a 100644 --- a/junifer/data/atlases.py +++ b/junifer/data/atlases.py @@ -61,7 +61,7 @@ for scale in range(1, 5): } t_name = f"TianxS{scale}x3TxMNI6thgeneration" _available_atlases[t_name] = { - "family": "Tian",r + "family": "Tian", "scale": scale, "magneticfield": "3T", "space": "MNI6thgeneration", -- 2.52.0 From fc66c07ae2a0d3acbacbcd7fed373f1b56420511 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 8 Sep 2022 11:23:05 +0200 Subject: [PATCH 246/287] chore: black formatting for pattern.py --- junifer/datagrabber/pattern.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/junifer/datagrabber/pattern.py b/junifer/datagrabber/pattern.py index ad1690ba5..e32dc45af 100644 --- a/junifer/datagrabber/pattern.py +++ b/junifer/datagrabber/pattern.py @@ -86,8 +86,7 @@ class PatternDataGrabber(BaseDataGrabber): glob_pattern = pattern for t_r in self.replacements: # Replace the first of each with a named group definition - re_pattern = re_pattern.replace( - f"{{{t_r}}}", f"(?P<{t_r}>.*)", 1) + re_pattern = re_pattern.replace(f"{{{t_r}}}", f"(?P<{t_r}>.*)", 1) for t_r in self.replacements: # Replace the second appearance of each with the named group -- 2.52.0 From ce9b063e27d4064ee8adfc29894d2d8f296d187d Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 8 Sep 2022 11:23:22 +0200 Subject: [PATCH 247/287] chore: isort and black fix for test_multiple.py --- junifer/datagrabber/tests/test_multiple.py | 15 ++++++++++----- 1 file changed, 10 insertions(+), 5 deletions(-) diff --git a/junifer/datagrabber/tests/test_multiple.py b/junifer/datagrabber/tests/test_multiple.py index eedd94965..5d591c23a 100644 --- a/junifer/datagrabber/tests/test_multiple.py +++ b/junifer/datagrabber/tests/test_multiple.py @@ -3,7 +3,8 @@ # Authors: Federico Raimondo # License: AGPL -from junifer.datagrabber import PatternDataladDataGrabber, MultipleDataGrabber +from junifer.datagrabber import MultipleDataGrabber, PatternDataladDataGrabber + _testing_dataset = { "example_bids": { @@ -34,14 +35,16 @@ def test_multiple() -> None: uri=repo_uri, types=["T1w"], patterns=pattern1, - replacements=replacements) + replacements=replacements, + ) dg2 = PatternDataladDataGrabber( rootdir=rootdir, uri=repo_uri, types=["bold"], patterns=pattern2, - replacements=replacements) + replacements=replacements, + ) dg = MultipleDataGrabber([dg1, dg2]) expected_subs = [ @@ -73,14 +76,16 @@ def test_multiple_no_intersection() -> None: uri=repo_uri1, types=["T1w"], patterns=pattern1, - replacements=replacements) + replacements=replacements, + ) dg2 = PatternDataladDataGrabber( rootdir=rootdir, uri=repo_uri2, types=["bold"], patterns=pattern2, - replacements=replacements) + replacements=replacements, + ) dg = MultipleDataGrabber([dg1, dg2]) expected_subs = set() -- 2.52.0 From 5a29b4a20a3a14940b04500121b950cd797732e8 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 8 Sep 2022 11:23:42 +0200 Subject: [PATCH 248/287] fix: correct argument name for _save_upsert() in sqlite.py --- junifer/storage/sqlite.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py index 79b9457d7..f15903982 100644 --- a/junifer/storage/sqlite.py +++ b/junifer/storage/sqlite.py @@ -391,7 +391,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): # Process metadata meta_md5, t_meta_row = process_meta(t_meta) # Get sqlalchemy engine - engine = self.get_engine(t_meta) + engine = self.get_engine(meta=t_meta) if meta_md5 not in inspect(engine).get_table_names(): # Convert metadata to dataframe meta_df = self._meta_row(meta=t_meta_row, meta_md5=meta_md5) @@ -553,7 +553,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): table_name = f"meta_{meta_md5}" t_df = in_storage.read_df(feature_md5=meta_md5) # Save data - out_storage._save_upsert(t_df, table_name, if_exist="nocheck") + out_storage._save_upsert(t_df, table_name, if_exists="nocheck") # TODO: refactor -- 2.52.0 From a09bb8587d6e80d5f21152cf647018af7f847daa Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 8 Sep 2022 11:24:28 +0200 Subject: [PATCH 249/287] fix: correct argument name for _save_upsert() in test_sqlite.py --- junifer/storage/tests/test_sqlite.py | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/junifer/storage/tests/test_sqlite.py b/junifer/storage/tests/test_sqlite.py index 05f7993f7..65b970b9a 100644 --- a/junifer/storage/tests/test_sqlite.py +++ b/junifer/storage/tests/test_sqlite.py @@ -152,7 +152,7 @@ def test_upsert_replace(tmp_path: Path) -> None: # Metadata to store meta = {"element": "test", "version": "0.0.1"} # Save to database - storage.store_df(df1, meta) + storage.store_df(df=df1, meta=meta) # Store metadata table_name = storage.store_metadata(meta) # Read stored table @@ -162,7 +162,7 @@ def test_upsert_replace(tmp_path: Path) -> None: # Check if dataframes are equal assert_frame_equal(df1, c_df1) # Upsert using replace - storage._save_upsert(df2, table_name, if_exist="replace") + storage._save_upsert(df=df2, name=table_name, if_exists="replace") # Read stored table c_df2 = _read_sql( table_name=table_name, uri=uri, index_col=["element", "pk2"] @@ -188,9 +188,9 @@ def test_upsert_ignore(tmp_path: Path) -> None: # Metadata to store meta = {"element": "test", "version": "0.0.1"} # Save to database - storage.store_df(df1, meta) + storage.store_df(df=df1, meta=meta) # Store metadata - table_name = storage.store_metadata(meta) + table_name = storage.store_metadata(meta=meta) # Read stored table c_df1 = _read_sql( table_name=table_name, uri=uri, index_col=["element", "pk2"] @@ -206,7 +206,7 @@ def test_upsert_ignore(tmp_path: Path) -> None: assert_frame_equal(c_dfignore, df_ignore) # Check for error with pytest.raises(ValueError, match=r"already exists"): - storage._save_upsert(df2, table_name, if_exist="fail") + storage._save_upsert(df2, table_name, if_exists="fail") def test_upsert_update(tmp_path: Path) -> None: -- 2.52.0 From 2d328201cb8fec77b640174c5701b52491c2b263 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 8 Sep 2022 11:24:52 +0200 Subject: [PATCH 250/287] fix: revert test_store_table() test for test_sqlite.py --- junifer/storage/tests/test_sqlite.py | 30 +++++++--------------------- 1 file changed, 7 insertions(+), 23 deletions(-) diff --git a/junifer/storage/tests/test_sqlite.py b/junifer/storage/tests/test_sqlite.py index 65b970b9a..aa33b9a8d 100644 --- a/junifer/storage/tests/test_sqlite.py +++ b/junifer/storage/tests/test_sqlite.py @@ -375,39 +375,23 @@ def test_store_table(tmp_path: Path) -> None: # Check if dataframes are equal assert_frame_equal(df, c_df) - -def test_store_table_check_warning(tmp_path: Path) -> None: - """Test table store and check warning. - - Parameters - ---------- - tmp_path : pathlib.Path - The path to the test directory. - - """ - uri = tmp_path / "test_store_table_check_warning.db" - storage = SQLiteFeatureStorage(uri=uri, single_output=True) - # Metadata to store - meta = {"element": "test", "version": "0.0.1", "marker": {"name": "fc"}} - # Data to store - data = [[1, 10], [2, 20], [3, 300], [4, 40], [5, 50], [6, 600]] + # New data to store + data_new = [[1, 10], [2, 20], [3, 300], [4, 40], [5, 50], [6, 600]] # Convert element to index - idx = element_to_index(meta, n_rows=6, rows_col_name="scan") + idx_new = element_to_index(meta, n_rows=6, rows_col_name="scan") # Create dataframe - df = pd.DataFrame(data, columns=["f1", "f2"], index=idx) + df_new = pd.DataFrame(data_new, columns=["f1", "f2"], index=idx_new) # Check warning with pytest.warns(RuntimeWarning, match=r"Some rows"): storage.store_table( - data, meta, columns=["f1", "f2"], rows_col_name="scan" + data_new, meta, columns=["f1", "f2"], rows_col_name="scan" ) - # Store metadata - table_name = storage.store_metadata(meta) # Read stored table - c_df = _read_sql( + c_df_new = _read_sql( table_name=table_name, uri=uri, index_col=["element", "scan"] ) # Check if dataframes are equal - assert_frame_equal(df, c_df) + assert_frame_equal(df_new, c_df_new) # TODO: can the test be parametrized? -- 2.52.0 From c62132abdd85081be3665c4dfe6da6c97e17c921 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 8 Sep 2022 11:25:12 +0200 Subject: [PATCH 251/287] fix: revert process_meta() in storage/utils.py --- junifer/storage/utils.py | 14 ++++++++------ 1 file changed, 8 insertions(+), 6 deletions(-) diff --git a/junifer/storage/utils.py b/junifer/storage/utils.py index bc0bd8faf..b707e80ea 100644 --- a/junifer/storage/utils.py +++ b/junifer/storage/utils.py @@ -63,21 +63,23 @@ def process_meta(meta: Dict) -> Tuple[str, Dict]: """ if meta is None: raise_error(msg="`meta` must be a dict (currently is None)") + # Copy the metadata + t_meta = meta.copy() # Remove key "element" - element = meta.pop("element", None) + element = t_meta.pop("element", None) if element is None: - if "_element_keys" not in meta: + if "_element_keys" not in t_meta: raise_error( msg="`meta` must contain the key 'element' or '_element_keys'" ) else: if isinstance(element, dict): - meta["_element_keys"] = list(element.keys()) + t_meta["_element_keys"] = list(element.keys()) else: - meta["_element_keys"] = ["element"] + t_meta["_element_keys"] = ["element"] # MD5 hash of the metadata - md5_hash = _meta_hash(meta) - return md5_hash, meta + md5_hash = _meta_hash(t_meta) + return md5_hash, t_meta def element_to_prefix(element: Union[Tuple, Dict, str, int]) -> str: -- 2.52.0 From 771eca9993bd45b3d3fe02cc113082dfb19294d2 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 8 Sep 2022 11:25:32 +0200 Subject: [PATCH 252/287] chore: black formatting for datagrabbers.py --- junifer/testing/datagrabbers.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/junifer/testing/datagrabbers.py b/junifer/testing/datagrabbers.py index 664a7800a..9d0f29415 100644 --- a/junifer/testing/datagrabbers.py +++ b/junifer/testing/datagrabbers.py @@ -39,8 +39,7 @@ class OasisVBMTestingDatagrabber(BaseDataGrabber): """ out = super().__getitem__(element) i_sub = int(element.split("-")[1]) - 1 - out["VBM_GM"] = { - "path": self._dataset.gray_matter_maps[i_sub]} + out["VBM_GM"] = {"path": self._dataset.gray_matter_maps[i_sub]} # Set the element accordingly out["meta"]["element"] = {"subject": element} return out -- 2.52.0 From 25a234445828587b4aa7db37f99c38f24f8c9900 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 8 Sep 2022 12:05:45 +0200 Subject: [PATCH 253/287] fix: pass autosummary check for datalad_base.py --- junifer/datagrabber/datalad_base.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/junifer/datagrabber/datalad_base.py b/junifer/datagrabber/datalad_base.py index 3bee2a0d4..28cc2c6d1 100644 --- a/junifer/datagrabber/datalad_base.py +++ b/junifer/datagrabber/datalad_base.py @@ -39,10 +39,10 @@ class DataladDataGrabber(BaseDataGrabber): Methods ------- - install + install: Installs (clones) the datalad dataset into the `datadir`. This method is called automatically when the datagrabber is used within a context. - remove + remove: Removes the datalad dataset from the `datadir`. This method is called automatically when the datagrabber is used within a context. -- 2.52.0 From 12a740573812b71c1c0d41fb296e55019aee92e6 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 8 Sep 2022 12:06:15 +0200 Subject: [PATCH 254/287] fix: correct references for api.rst --- docs/api.rst | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/docs/api.rst b/docs/api.rst index a2789d71d..3aceb9ac7 100644 --- a/docs/api.rst +++ b/docs/api.rst @@ -8,10 +8,10 @@ Data Grabbers ^^^^^^^^^^^^^ .. autoclass:: junifer.datagrabber.base.BaseDataGrabber - :members: -.. autoclass:: junifer.datagrabber.base.PatternDataGrabber - :members: -.. autoclass:: junifer.datagrabber.base.DataladDataGrabber - :members: -.. autoclass:: junifer.datagrabber.base.PatternDataladDataGrabber - :members: \ No newline at end of file + :members: +.. autoclass:: junifer.datagrabber.pattern.PatternDataGrabber + :members: +.. autoclass:: junifer.datagrabber.datalad_base.DataladDataGrabber + :members: +.. autoclass:: junifer.datagrabber.pattern_datalad_base.PatternDataladDataGrabber + :members: -- 2.52.0 From cd89dc71842cd09b2bd9d85619dce080cd07a3d0 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 8 Sep 2022 12:20:52 +0200 Subject: [PATCH 255/287] fix: correct import for BIDS Datalad datagrabber example --- examples/run_datagrabber_bids_datalad.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/examples/run_datagrabber_bids_datalad.py b/examples/run_datagrabber_bids_datalad.py index 28a1eeb77..1727afed9 100644 --- a/examples/run_datagrabber_bids_datalad.py +++ b/examples/run_datagrabber_bids_datalad.py @@ -10,7 +10,7 @@ Authors: Federico Raimondo License: BSD 3 clause """ -from junifer.datagrabber.base import PatternDataladDataGrabber +from junifer.datagrabber import PatternDataladDataGrabber from junifer.utils import configure_logging ############################################################################### -- 2.52.0 From d776730f32be08d86c1082c52209fa5d5c3e3fd5 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Thu, 8 Sep 2022 15:14:10 +0200 Subject: [PATCH 256/287] Fix typing + rename base classes to concrete + fix some tests --- examples/norun_hcpfc_pearson.py | 4 +- junifer/api/cli.py | 27 +++--- junifer/api/functions.py | 52 +++++++----- junifer/api/registry.py | 9 +- junifer/api/tests/test_functions.py | 8 +- junifer/api/tests/test_registry.py | 11 +-- junifer/data/atlases.py | 83 +++++++++++-------- junifer/data/tests/test_atlases.py | 25 +++--- junifer/datagrabber/__init__.py | 2 +- junifer/datagrabber/datalad_base.py | 3 +- junifer/datagrabber/hcp.py | 24 +++--- ...ern_datalad_base.py => pattern_datalad.py} | 10 +-- junifer/datagrabber/tests/test_base.py | 2 +- ...atalad_base.py => test_pattern_datalad.py} | 4 +- junifer/datareader/default.py | 4 +- junifer/markers/base.py | 23 ++--- junifer/markers/collection.py | 15 +++- junifer/markers/parcel.py | 22 +++-- junifer/markers/pipeline_mixin.py | 2 +- junifer/markers/tests/test_collection.py | 7 +- junifer/markers/tests/test_markers_base.py | 12 +-- junifer/markers/tests/test_pipeline_mixin.py | 6 +- junifer/preprocess/confounds.py | 9 +- junifer/preprocess/tests/test_confounds.py | 24 +++--- junifer/stats.py | 23 ++++- junifer/storage/base.py | 2 +- junifer/storage/sqlite.py | 29 ++++--- junifer/storage/tests/test_sqlite.py | 45 ++++++---- junifer/storage/tests/test_storage_base.py | 4 +- junifer/storage/tests/test_utils.py | 2 +- junifer/storage/utils.py | 4 +- junifer/utils/logging.py | 12 +-- 32 files changed, 302 insertions(+), 207 deletions(-) rename junifer/datagrabber/{pattern_datalad_base.py => pattern_datalad.py} (86%) rename junifer/datagrabber/tests/{test_pattern_datalad_base.py => test_pattern_datalad.py} (97%) diff --git a/examples/norun_hcpfc_pearson.py b/examples/norun_hcpfc_pearson.py index fe2632e92..6875e23b9 100644 --- a/examples/norun_hcpfc_pearson.py +++ b/examples/norun_hcpfc_pearson.py @@ -1,6 +1,8 @@ -""" +"""extraction of FC from HCP data. + HCP FC Extraction ====================== + Authors: Leonard Sasse License: BSD 3 clause """ diff --git a/junifer/api/cli.py b/junifer/api/cli.py index 078b3c4f2..6bcba4392 100644 --- a/junifer/api/cli.py +++ b/junifer/api/cli.py @@ -4,9 +4,10 @@ # Synchon Mandal # License: AGPL -from typing import Dict, List +from typing import Dict, List, Union import click +import pathlib from ..utils.logging import configure_logging, logger, warn_with_log from .functions import collect as api_collect @@ -15,7 +16,7 @@ from .functions import run as api_run from .parser import parse_yaml -def _parse_elements(element: str, config: Dict) -> List: +def _parse_elements(element: str, config: Dict) -> Union[List, None]: """Parse elements from cli. Parameters @@ -58,7 +59,8 @@ def cli() -> None: # pragma: no cover @cli.command() @click.argument( - "filepath", type=click.Path(exists=True, readable=True, dir_okay=False) + "filepath", type=click.Path( + exists=True, readable=True, dir_okay=False, path_type=pathlib.Path) ) @click.option("--element", type=str, multiple=True) @click.option( @@ -80,9 +82,9 @@ def run(filepath: click.Path, element: str, verbose: click.Choice) -> None: The verbosity level: warning, info or debug (default "info"). """ - configure_logging(level=verbose.upper()) + configure_logging(level=str(verbose).upper()) # TODO: add validation - config = parse_yaml(filepath) + config = parse_yaml(filepath) # type: ignore workdir = config["workdir"] datagrabber = config["datagrabber"] markers = config["markers"] @@ -100,8 +102,8 @@ def run(filepath: click.Path, element: str, verbose: click.Choice) -> None: @cli.command() @click.argument( - "filepath", type=click.Path(exists=True, readable=True, dir_okay=False) -) + "filepath", type=click.Path( + exists=True, readable=True, dir_okay=False, path_type=pathlib.Path)) @click.option( "-v", "--verbose", @@ -119,9 +121,9 @@ def collect(filepath: click.Path, verbose: click.Choice) -> None: The verbosity level: warning, info or debug (default "info"). """ - configure_logging(level=verbose.upper()) + configure_logging(level=str(verbose).upper()) # TODO: add validation - config = parse_yaml(filepath) + config = parse_yaml(filepath) # type: ignore storage = config["storage"] # Perform operation api_collect(storage=storage) @@ -129,7 +131,8 @@ def collect(filepath: click.Path, verbose: click.Choice) -> None: @cli.command() @click.argument( - "filepath", type=click.Path(exists=True, readable=True, dir_okay=False) + "filepath", type=click.Path( + exists=True, readable=True, dir_okay=False, path_type=pathlib.Path) ) @click.option("--element", type=str, multiple=True) @click.option("--overwrite", is_flag=True) @@ -163,9 +166,9 @@ def queue( The verbosity level: warning, info or debug (default "info"). """ - configure_logging(level=verbose.upper()) + configure_logging(level=str(verbose).upper()) # TODO: add validation - config = parse_yaml(filepath) + config = parse_yaml(filepath) # type: ignore elements = _parse_elements(element, config) queue_config = config.pop("queue") kind = queue_config.pop("kind") diff --git a/junifer/api/functions.py b/junifer/api/functions.py index 816a2dac0..3c4bc7e4c 100644 --- a/junifer/api/functions.py +++ b/junifer/api/functions.py @@ -8,7 +8,8 @@ import shutil import subprocess from pathlib import Path -from typing import Dict, List, Optional, Sequence, Union +from typing import Dict, List, Optional, Union, Tuple +import typing import yaml @@ -21,7 +22,7 @@ from ..utils.fs import make_executable from .registry import build -def _get_datagrabber(datagrabber_config: Dict) -> Dict: +def _get_datagrabber(datagrabber_config: Dict) -> BaseDataGrabber: """Get datagrabber. Parameters @@ -43,6 +44,7 @@ def _get_datagrabber(datagrabber_config: Dict) -> Dict: baseclass=BaseDataGrabber, init_params=datagrabber_params, ) + datagrabber = typing.cast(BaseDataGrabber, datagrabber) return datagrabber @@ -51,7 +53,7 @@ def run( datagrabber: Dict, markers: List[Dict], storage: Dict, - elements: Union[str, Sequence, None] = None, + elements: Union[str, List[Union[str, Tuple]], Tuple, None] = None, ) -> None: """Run the pipeline on the selected element. @@ -81,8 +83,10 @@ def run( # Convert str to Path if isinstance(workdir, str): workdir = Path(workdir) + if not isinstance(elements, List) and elements is not None: + elements = [elements] # Get datagrabber to use - datagrabber = _get_datagrabber(datagrabber) + datagrabber_object = _get_datagrabber(datagrabber) # Copy to avoid changing the original dict _markers = [x.copy() for x in markers] built_markers = [] @@ -98,22 +102,23 @@ def run( # Get storage engine to use storage_params = storage.copy() storage_kind = storage_params.pop("kind") - storage = build( + storage_object = build( step="storage", name=storage_kind, baseclass=BaseFeatureStorage, init_params=storage_params, ) + storage_object = typing.cast(BaseFeatureStorage, storage_object) # Create new marker collection - mc = MarkerCollection(markers=built_markers, storage=storage) + mc = MarkerCollection(markers=built_markers, storage=storage_object) # Fit elements - with datagrabber: + with datagrabber_object: if elements is not None: for t_element in elements: - mc.fit(datagrabber[t_element]) + mc.fit(datagrabber_object[t_element]) else: - for t_element in datagrabber: - mc.fit(datagrabber[t_element]) + for t_element in datagrabber_object: + mc.fit(datagrabber_object[t_element]) def collect(storage: Dict) -> None: @@ -131,14 +136,15 @@ def collect(storage: Dict) -> None: storage_kind = storage_params.pop("kind") logger.info(f"Collecting data using {storage_kind}") logger.debug(f"\tStorage params: {storage_params}") - storage = build( + storage_object = build( step="storage", name=storage_kind, baseclass=BaseFeatureStorage, init_params=storage_params, ) + storage_object = typing.cast(BaseFeatureStorage, storage_object) logger.debug("Running storage.collect()") - storage.collect() + storage_object.collect() logger.info("Collect done") @@ -147,7 +153,7 @@ def queue( kind: str, jobname: str = "junifer_job", overwrite: bool = False, - elements: Union[str, Sequence, None] = None, + elements: Union[str, List[Union[str, Tuple]], Tuple, None] = None, **kwargs: Union[str, int, bool], ) -> None: # pragma : no cover """Queue a job to be executed later. @@ -209,12 +215,19 @@ def queue( datagrabber = _get_datagrabber(config["datagrabber"]) with datagrabber as dg: elements = dg.get_elements() + + # TODO: Fix typing of elements + if not isinstance(elements, List): + elements = [elements] # type: ignore + + typing.cast(List[Union[str, Tuple]], elements) + if kind == "HTCondor": _queue_condor( jobname=jobname, jobdir=jobdir, yaml_config=yaml_config, - elements=elements, + elements=elements, # type: ignore **kwargs, ) elif kind == "SLURM": @@ -222,7 +235,7 @@ def queue( jobname=jobname, jobdir=jobdir, yaml_config=yaml_config, - elements=elements, + elements=elements, # type: ignore **kwargs, ) else: @@ -235,7 +248,7 @@ def _queue_condor( jobname: str, jobdir: Path, yaml_config: Path, - elements: Union[str, Sequence, None] = None, + elements: List[Union[str, Tuple]], env: Optional[Dict[str, str]] = None, mem: str = "8G", cpus: int = 1, @@ -255,9 +268,8 @@ def _queue_condor( The path to the job directory. yaml_config : pathlib.Path The path to the YAML config file. - elements : str or tuple or list of str or tuple, optional - Element(s) to process. Will be used to index the datagrabber - (default None). + elements : list of str or tuple + Element(s) to process. Will be used to index the datagrabber. env : dict, optional The environment variables passed as dictionary (default None). mem : str, optional @@ -411,7 +423,7 @@ def _queue_slurm( jobname: str, jobdir: Path, yaml_config: Path, - elements: Union[str, Sequence, None] = None, + elements: List[Union[str, Tuple]], ) -> None: """Submit job to SLURM. diff --git a/junifer/api/registry.py b/junifer/api/registry.py index e3f0d42e7..1d567a82a 100644 --- a/junifer/api/registry.py +++ b/junifer/api/registry.py @@ -5,10 +5,13 @@ # Synchon Mandal # License: AGPL -from typing import Dict, List, Optional +from typing import TYPE_CHECKING, Dict, List, Optional, Union from ..utils.logging import logger, raise_error - +if TYPE_CHECKING: + from ..datagrabber.base import BaseDataGrabber + from ..markers.base import PipelineStepMixin + from ..storage.base import BaseFeatureStorage # Define valid steps for operation _valid_steps = [ @@ -96,7 +99,7 @@ def build( name: str, baseclass: type, init_params: Optional[Dict] = None, -) -> type: +) -> Union["BaseDataGrabber", "PipelineStepMixin", "BaseFeatureStorage"]: """Ensure that the given object is an instance of the given class. Parameters diff --git a/junifer/api/tests/test_functions.py b/junifer/api/tests/test_functions.py index 57f1cf7d6..3b175c901 100644 --- a/junifer/api/tests/test_functions.py +++ b/junifer/api/tests/test_functions.py @@ -59,7 +59,7 @@ def test_run_single_element(tmp_path: Path) -> None: outdir.mkdir() # Create storage uri = outdir / "test.db" - storage["uri"] = uri + storage["uri"] = uri # type: ignore # Run operations run( workdir=workdir, @@ -90,7 +90,7 @@ def test_run_multi_element(tmp_path: Path) -> None: outdir.mkdir() # Create storage uri = outdir / "test.db" - storage["uri"] = uri + storage["uri"] = uri # type: ignore # Run operations run( workdir=workdir, @@ -121,7 +121,7 @@ def test_run_and_collect(tmp_path: Path) -> None: outdir.mkdir() # Create storage uri = outdir / "test.db" - storage["uri"] = uri + storage["uri"] = uri # type: ignore # Run operations run( workdir=workdir, @@ -133,7 +133,7 @@ def test_run_and_collect(tmp_path: Path) -> None: dg = build( step="datagrabber", name=datagrabber["kind"], baseclass=BaseDataGrabber ) - elements = dg.get_elements() + elements = dg.get_elements() # type: ignore # This should create 10 files files = list(outdir.glob("*.db")) assert len(files) == len(elements) diff --git a/junifer/api/tests/test_registry.py b/junifer/api/tests/test_registry.py index ff4b9264c..7ac9a88ca 100644 --- a/junifer/api/tests/test_registry.py +++ b/junifer/api/tests/test_registry.py @@ -4,6 +4,7 @@ # Leonard Sasse # Synchon Mandal # License: AGPL +from typing import Type import logging from abc import ABC @@ -18,7 +19,7 @@ from junifer.storage import SQLiteFeatureStorage def test_register_invalid_step(): """Test register invalid step name.""" with pytest.raises(ValueError, match="Invalid step:"): - register(step="foo", name="bar", klass="baz") + register(step="foo", name="bar", klass=str) # TODO: improve parametrization @@ -30,7 +31,7 @@ def test_register_invalid_step(): ], ) def test_register( - caplog: pytest.LogCaptureFixture, step: str, name: str, klass: str + caplog: pytest.LogCaptureFixture, step: str, name: str, klass: Type ) -> None: """Test register. @@ -70,7 +71,7 @@ def test_get_step_names_absent() -> None: def test_get_step_names() -> None: """Test get step names.""" # Register datagrabber - register(step="datagrabber", name="bar", klass="baz") + register(step="datagrabber", name="bar", klass=str) # Get step names for datagrabber datagrabbers = get_step_names(step="datagrabber") # Check for datagrabber step name @@ -93,10 +94,10 @@ def test_get_class_invalid_name() -> None: def test_get_class(): """Test get class.""" # Register datagrabber - register(step="datagrabber", name="bar", klass="baz") + register(step="datagrabber", name="bar", klass=str) # Get class obj = get_class(step="datagrabber", name="bar") - assert obj == "baz" + assert obj == str # TODO: possible parametrization? diff --git a/junifer/data/atlases.py b/junifer/data/atlases.py index 7fe70a46a..d2aeeb0f6 100644 --- a/junifer/data/atlases.py +++ b/junifer/data/atlases.py @@ -10,7 +10,7 @@ import shutil import tempfile import zipfile from pathlib import Path -from typing import TYPE_CHECKING, List, Optional, Tuple, Union +from typing import TYPE_CHECKING, List, Optional, Tuple, Union, Any, Dict import nibabel as nib import numpy as np @@ -37,10 +37,11 @@ Optional keys: """ # TODO: have separate dictionary for built-in -_available_atlases = { +_available_atlases: Dict[str, Dict[Any, Any]] = { "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]: @@ -150,7 +151,7 @@ def list_atlases() -> List[str]: def load_atlas( name: str, atlas_dir: Union[str, Path, None] = None, - resolution: Optional[int] = None, + resolution: Optional[float] = None, path_only: bool = False, ) -> Tuple[Optional["Nifti1Image"], List[str], Path]: """Load a brain atlas (including a label file). @@ -165,7 +166,7 @@ def load_atlas( 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 + resolution : float, 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. By default, will load the highest one @@ -238,7 +239,7 @@ def load_atlas( def _retrieve_atlas( family: str, atlas_dir: Union[str, Path, None] = None, - resolution: Optional[int] = None, + resolution: Optional[float] = None, **kwargs, ) -> Tuple[Path, List[str]]: """Retrieve a brain atlas object from nilearn or a specified online source. @@ -254,7 +255,7 @@ def _retrieve_atlas( 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 + resolution : float, 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. By default, will load the highest one @@ -326,21 +327,21 @@ def _retrieve_atlas( def _closest_resolution( - resolution: int, - valid_resolution: Union[List[int], np.ndarray], -) -> int: + resolution: Optional[float], + valid_resolution: Union[List[float], List[int], np.ndarray], +) -> Union[float, int]: """Find the closest resolution. Parameters ---------- - resolution : int + resolution : float The given resolution. - valid_resolution : list of int or np.ndarray + valid_resolution : list of float or np.ndarray The array of valid resolutions. Returns ------- - int + float The closest valid resolution. """ @@ -363,7 +364,7 @@ def _closest_resolution( def _retrieve_schaefer( atlas_dir: Path, - resolution: int, + resolution: Optional[float] = None, n_rois: Optional[int] = None, yeo_networks: int = 7, ) -> Tuple[Path, List[str]]: @@ -373,8 +374,11 @@ def _retrieve_schaefer( ---------- atlas_dir : pathlib.Path The path to the atlas data directory. - resolution : {1, 2} - The resolution of the atlas to load. + resolution : float, 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. By default, will load the highest one + (default None). Available resolutions for this atlas are 1mm and 2mm. 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 @@ -440,7 +444,7 @@ def _retrieve_schaefer( datasets.fetch_atlas_schaefer_2018( n_rois=n_rois, yeo_networks=yeo_networks, - resolution_mm=resolution, + resolution_mm=resolution, # type: ignore we know it's 1 or 2 data_dir=str(atlas_dir.absolute()), ) @@ -462,7 +466,7 @@ def _retrieve_schaefer( def _retrieve_tian( atlas_dir: Path, - resolution: int, + resolution: Optional[float] = None, scale: Optional[int] = None, space: str = "MNI6thgeneration", magneticfield: str = "3T", @@ -473,8 +477,12 @@ def _retrieve_tian( ---------- atlas_dir : pathlib.Path The path to the atlas data directory. - resolution : {1, 2} - The resolution of the atlas to load. + resolution : float, 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. By default, will load the highest one + (default None). Available resolutions for this atlas depend on the + space and magnetic field. scale : {1, 2, 3, 4}, optional Scale of atlas (defines granularity) (default None). space : {"MNI6thgeneration", "MNInonlinear2009cAsym"}, optional @@ -505,32 +513,34 @@ def _retrieve_tian( logger.info(f"\tresolution: {resolution}") # check validity of atlas parameters _valid_scales = [1, 2, 3, 4] - _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}" ) - if magneticfield not in _valid_fields: - raise_error( - f"The parameter `magneticfield` ({magneticfield}) needs to be " - f"one of the following: {_valid_fields}" - ) + _valid_resolutions = [] # avoid pylance error if magneticfield == "3T": _valid_spaces = ["MNI6thgeneration", "MNInonlinear2009cAsym"] if space == "MNI6thgeneration": _valid_resolutions = [1, 2] elif space == "MNInonlinear2009cAsym": _valid_resolutions = [2] + else: + raise_error( + f"The parameter `space` ({space}) for 3T needs to be one of " + f"the following: {_valid_spaces}" + ) elif magneticfield == "7T": - _valid_spaces = ["MNI6thgeneration"] _valid_resolutions = [1.6] - - if space not in _valid_spaces: + if space != "MNI6thgeneration": + raise_error( + f"The parameter `space` ({space}) for 7T needs to be " + f"MNI6thgeneration") + else: raise_error( - f"The parameter `space` ({space}) needs to be one of " - f"the following: {_valid_spaces}" + f"The parameter `magneticfield` ({magneticfield}) needs to be " + f"one of the following: 3T or 7T" ) resolution = _closest_resolution(resolution, _valid_resolutions) @@ -581,6 +591,8 @@ def _retrieve_tian( "Currently there are no labels provided for the 7T Tian atlas. " "A simple numbering scheme for distinction was therefore used." ) + else: # pragma: no cover + raise_error('This should not happen. Please report this error.') # check existence of atlas if not (atlas_fname.exists() and atlas_lname.exists()): @@ -613,7 +625,9 @@ def _retrieve_tian( def _retrieve_suit( - atlas_dir: Path, resolution: int, space: str = "MNI" + atlas_dir: Path, + resolution: Optional[float], + space: str = "MNI" ) -> Tuple[Path, List[str]]: """Retrieve SUIT atlas. @@ -621,8 +635,11 @@ def _retrieve_suit( ---------- atlas_dir : pathlib.Path The path to the atlas data directory. - resolution : {1, 2} - The resolution of the atlas to load. + resolution : float, 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. By default, will load the highest one + (default None). Available resolutions for this atlas are 1mm and 2mm. space : {"MNI", "SUIT"}, optional Space of atlas (default "MNI"). (For more information see http://www.diedrichsenlab.org/imaging/suit.htm). diff --git a/junifer/data/tests/test_atlases.py b/junifer/data/tests/test_atlases.py index 6b9eea365..df70002fb 100644 --- a/junifer/data/tests/test_atlases.py +++ b/junifer/data/tests/test_atlases.py @@ -166,7 +166,7 @@ def test_schaefer_atlas(tmp_path: Path) -> None: assert img is not None assert fname.name == fname1 assert len(lbl) == 100 - assert_array_equal(img.header["pixdim"][1:4], [1, 1, 1]) + assert_array_equal(img.header["pixdim"][1:4], [1, 1, 1]) # type: ignore # Test with Path img, lbl, fname = load_atlas(name="Schaefer100x7", atlas_dir=tmp_path) @@ -180,7 +180,7 @@ def test_schaefer_atlas(tmp_path: Path) -> None: assert fname.name == fname2 assert len(lbl) == 100 assert img2 is not None - assert_array_equal(img2.header["pixdim"][1:4], [2, 2, 2]) + assert_array_equal(img2.header["pixdim"][1:4], [2, 2, 2]) # type: ignore # Load atlas img2, lbl, fname = load_atlas( "Schaefer100x7", @@ -191,7 +191,7 @@ def test_schaefer_atlas(tmp_path: Path) -> None: assert fname.name == fname2 assert len(lbl) == 100 assert img2 is not None - assert_array_equal(img2.header["pixdim"][1:4], [2, 2, 2]) + assert_array_equal(img2.header["pixdim"][1:4], [2, 2, 2]) # type: ignore # Load atlas img2, lbl, fname = load_atlas( "Schaefer100x7", @@ -202,7 +202,7 @@ def test_schaefer_atlas(tmp_path: Path) -> None: assert fname.name == fname1 assert len(lbl) == 100 assert img2 is not None - assert_array_equal(img2.header["pixdim"][1:4], [1, 1, 1]) + assert_array_equal(img2.header["pixdim"][1:4], [1, 1, 1]) # type: ignore # Load atlas img2, lbl, fname = load_atlas( "Schaefer100x7", @@ -213,7 +213,7 @@ def test_schaefer_atlas(tmp_path: Path) -> None: assert fname.name == fname1 assert len(lbl) == 100 assert img2 is not None - assert_array_equal(img2.header["pixdim"][1:4], [1, 1, 1]) + assert_array_equal(img2.header["pixdim"][1:4], [1, 1, 1]) # type: ignore def test_load_atlas_schaefer() -> None: @@ -274,7 +274,7 @@ def test_suit(tmp_path: Path) -> None: assert img is not None assert fname.name == fname1 assert len(lbl) == 34 - assert_array_equal(img.header["pixdim"][1:4], [1, 1, 1]) + assert_array_equal(img.header["pixdim"][1:4], [1, 1, 1]) # type: ignore # Load atlas img, lbl, fname = load_atlas(name="SUITxSUIT", atlas_dir=tmp_path) @@ -282,7 +282,7 @@ def test_suit(tmp_path: Path) -> None: assert img is not None assert fname.name == fname1 assert len(lbl) == 34 - assert_array_equal(img.header["pixdim"][1:4], [1, 1, 1]) + assert_array_equal(img.header["pixdim"][1:4], [1, 1, 1]) # type: ignore # Load atlas img, lbl, fname = load_atlas(name="SUITxMNI", atlas_dir=tmp_path) @@ -290,7 +290,7 @@ def test_suit(tmp_path: Path) -> None: assert img is not None assert fname.name == fname1 assert len(lbl) == 34 - assert_array_equal(img.header["pixdim"][1:4], [1, 1, 1]) + assert_array_equal(img.header["pixdim"][1:4], [1, 1, 1]) # type: ignore def test_retrieve_suit_incorrect_space(tmp_path: Path) -> None: @@ -346,7 +346,7 @@ def test_tian_3T_6thgeneration( 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]) + assert_array_equal(img.header["pixdim"][1:4], [1, 1, 1]) # type: ignore # Load atlas img, lbl, fname = load_atlas( name=f"TianxS{scale}x3TxMNI6thgeneration", @@ -357,7 +357,7 @@ def test_tian_3T_6thgeneration( 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]) + assert_array_equal(img.header["pixdim"][1:4], [2, 2, 2]) # type: ignore @pytest.mark.parametrize( @@ -400,7 +400,7 @@ def test_tian_3T_nonlinear2009cAsym( 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]) + assert_array_equal(img.header["pixdim"][1:4], [2, 2, 2]) # type: ignore @pytest.mark.parametrize( @@ -442,7 +442,8 @@ def test_tian_7T_6thgeneration( 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]) + assert_array_almost_equal( + img.header["pixdim"][1:4], [1.6, 1.6, 1.6]) # type: ignore def test_retrieve_tian_incorrect_space(tmp_path: Path) -> None: diff --git a/junifer/datagrabber/__init__.py b/junifer/datagrabber/__init__.py index 8f0872fce..5975d9b5d 100644 --- a/junifer/datagrabber/__init__.py +++ b/junifer/datagrabber/__init__.py @@ -9,4 +9,4 @@ from .datalad_base import DataladDataGrabber from .hcp import DataladHCP1200, HCP1200 from .multiple import MultipleDataGrabber from .pattern import PatternDataGrabber -from .pattern_datalad_base import PatternDataladDataGrabber +from .pattern_datalad import PatternDataladDataGrabber diff --git a/junifer/datagrabber/datalad_base.py b/junifer/datagrabber/datalad_base.py index 28cc2c6d1..54664a1ff 100644 --- a/junifer/datagrabber/datalad_base.py +++ b/junifer/datagrabber/datalad_base.py @@ -120,7 +120,8 @@ class DataladDataGrabber(BaseDataGrabber): def install(self) -> None: """Install the datalad dataset into the datadir.""" logger.debug(f"Installing dataset {self.uri} to {self._datadir}") - self._dataset = dl.install(self._datadir, source=self.uri) + self._dataset = dl.install( # type: ignore because of datalad + self._datadir, source=self.uri) logger.debug("Dataset installed") def remove(self): diff --git a/junifer/datagrabber/hcp.py b/junifer/datagrabber/hcp.py index 1ebfc0a73..9db59031b 100644 --- a/junifer/datagrabber/hcp.py +++ b/junifer/datagrabber/hcp.py @@ -55,18 +55,19 @@ class HCP1200(PatternDataGrabber): ] # Set default tasks if tasks is None: - self.tasks = all_tasks + self.tasks: List[str] = all_tasks # Convert single task into list - if isinstance(tasks, str): - tasks = [tasks] - # Check for invalid task(s) - for task in tasks: - if task not in all_tasks: - raise_error( - f"'{task}' is not a valid HCP-YA fMRI task input. " - f"Valid task values can be any or all of {all_tasks}." - ) - + else: + if not isinstance(tasks, List): + tasks = [tasks] + # Check for invalid task(s) + for task in tasks: + if task not in all_tasks: + raise_error( + f"'{task}' is not a valid HCP-YA fMRI task input. " + f"Valid task values can be any or all of {all_tasks}." + ) + self.tasks: List[str] = tasks # All phase encodings all_phase_encodings = ["LR", "RL"] # Set phase encodings @@ -102,7 +103,6 @@ class HCP1200(PatternDataGrabber): patterns=patterns, replacements=replacements, ) - self.tasks = tasks self.phase_encodings = phase_encodings def __getitem__(self, element: Tuple[str, str, str]) -> Dict[str, Path]: diff --git a/junifer/datagrabber/pattern_datalad_base.py b/junifer/datagrabber/pattern_datalad.py similarity index 86% rename from junifer/datagrabber/pattern_datalad_base.py rename to junifer/datagrabber/pattern_datalad.py index 5a43e6ec6..07e5e296d 100644 --- a/junifer/datagrabber/pattern_datalad_base.py +++ b/junifer/datagrabber/pattern_datalad.py @@ -5,7 +5,7 @@ # Synchon Mandal # License: AGPL -from typing import Dict, List, Optional +from typing import Dict, List from ..api.decorators import register_datagrabber from .datalad_base import DataladDataGrabber @@ -22,8 +22,8 @@ class PatternDataladDataGrabber(DataladDataGrabber, PatternDataGrabber): Parameters ---------- - types : list of str, optional - The types of data to be grabbed (default None). + types : list of str + The types of data to be grabbed. patterns : dict, optional Patterns for each type of data as a dictionary. The keys are the types and the values are the patterns. Each occurrence of the string @@ -41,8 +41,8 @@ class PatternDataladDataGrabber(DataladDataGrabber, PatternDataGrabber): def __init__( self, - types: Optional[List[str]] = None, - patterns: Optional[Dict[str, str]] = None, + types: List[str], + patterns: Dict[str, str], **kwargs, ) -> None: """Initialize the class.""" diff --git a/junifer/datagrabber/tests/test_base.py b/junifer/datagrabber/tests/test_base.py index ba0caae61..91d622d62 100644 --- a/junifer/datagrabber/tests/test_base.py +++ b/junifer/datagrabber/tests/test_base.py @@ -15,7 +15,7 @@ from junifer.datagrabber.base import BaseDataGrabber def test_BaseDataGrabber_abstractness() -> None: """Test BaseDataGrabber is abstract base class.""" with pytest.raises(TypeError, match=r"abstract"): - BaseDataGrabber(datadir="/tmp", types=["func"]) + BaseDataGrabber(datadir="/tmp", types=["func"]) # type: ignore def test_BaseDataGrabber() -> None: diff --git a/junifer/datagrabber/tests/test_pattern_datalad_base.py b/junifer/datagrabber/tests/test_pattern_datalad.py similarity index 97% rename from junifer/datagrabber/tests/test_pattern_datalad_base.py rename to junifer/datagrabber/tests/test_pattern_datalad.py index 0d445c9e8..105918c44 100644 --- a/junifer/datagrabber/tests/test_pattern_datalad_base.py +++ b/junifer/datagrabber/tests/test_pattern_datalad.py @@ -1,4 +1,4 @@ -"""Provide tests for pattern_datalad_base.""" +"""Provide tests for pattern_datalad.""" # Authors: Federico Raimondo # Leonard Sasse @@ -9,7 +9,7 @@ from pathlib import Path import pytest -from junifer.datagrabber.pattern_datalad_base import PatternDataladDataGrabber +from junifer.datagrabber.pattern_datalad import PatternDataladDataGrabber _testing_dataset = { diff --git a/junifer/datareader/default.py b/junifer/datareader/default.py index 34996f6de..7bddd2cc1 100644 --- a/junifer/datareader/default.py +++ b/junifer/datareader/default.py @@ -5,7 +5,7 @@ # License: AGPL from pathlib import Path -from typing import Dict +from typing import Dict, List import nibabel as nib import pandas as pd @@ -33,7 +33,7 @@ class DefaultDataReader(PipelineStepMixin): """Mixin class for default data reader.""" # TODO: complete type annotations - def validate_input(self, input): + def validate_input(self, input: List[str]) -> None: """Validate input. Parameters diff --git a/junifer/markers/base.py b/junifer/markers/base.py index 1fef738bf..d232f1405 100644 --- a/junifer/markers/base.py +++ b/junifer/markers/base.py @@ -4,7 +4,7 @@ # Synchon Mandal # License: AGPL -from typing import Dict, List, Optional +from typing import Dict, List, Optional, Union from ..utils import logger, raise_error from .pipeline_mixin import PipelineStepMixin @@ -22,7 +22,11 @@ class BaseMarker(PipelineStepMixin): """ - def __init__(self, on: List, name: Optional[str] = None) -> None: + def __init__( + self, + on: Union[List[str], str], + name: Optional[str] = None + ) -> None: """Initialize the class.""" if not isinstance(on, list): on = [on] @@ -78,24 +82,23 @@ class BaseMarker(PipelineStepMixin): Parameters ---------- input : list of str - The input to the pipeline step. The list must contain the + The input to the marker. The list must contain the available Junifer Data dictionary keys. Returns ------- list of str - The updated list of available Junifer Data dictionary keys after - the pipeline step. + The updated list of output kinds, as storage possibilities. """ - pass + raise_error(msg="compute() not implemented", klass=NotImplementedError) - def compute(self, input: List[str]) -> Dict: + def compute(self, input: Dict) -> Dict: """Compute. Parameters ---------- - input : list of str + input : Dict[str, Dict] The input to the pipeline step. The list must contain the available Junifer Data dictionary keys. @@ -108,7 +111,7 @@ class BaseMarker(PipelineStepMixin): raise_error(msg="compute() not implemented", klass=NotImplementedError) # TODO: complete type annotations - def store(self, input: List[str], out: Dict, storage) -> None: + def store(self, kind: str, out: Dict, storage) -> None: """Store. Parameters @@ -120,7 +123,7 @@ class BaseMarker(PipelineStepMixin): raise_error(msg="store() not implemented", klass=NotImplementedError) # TODO: complete type annotations - def fit_transform(self, input: List[str], storage=None) -> Dict: + def fit_transform(self, input: Dict[str, Dict], storage=None) -> Dict: """Fit and transform. Parameters diff --git a/junifer/markers/collection.py b/junifer/markers/collection.py index f658e1b9f..98fc57599 100644 --- a/junifer/markers/collection.py +++ b/junifer/markers/collection.py @@ -5,10 +5,14 @@ # License: AGPL from collections import Counter -from typing import Dict, Optional +from typing import Dict, Optional, List + +from junifer.markers.pipeline_mixin import PipelineStepMixin from ..datareader import DefaultDataReader from ..utils import logger +from .base import BaseMarker +from ..storage.base import BaseFeatureStorage class MarkerCollection: @@ -24,7 +28,11 @@ class MarkerCollection: """ def __init__( - self, markers, datareader=None, preprocessing=None, storage=None + self, + markers: List[BaseMarker], + datareader: Optional[PipelineStepMixin] = None, + preprocessing: Optional[PipelineStepMixin] = None, + storage: Optional[BaseFeatureStorage] = None ): """Initialize the class.""" # Check that the markers have different names @@ -42,8 +50,7 @@ class MarkerCollection: self._preprocessing = preprocessing self._storage = storage - # TODO: complete type annotations - def fit(self, input) -> Optional[Dict]: + def fit(self, input: Dict[str, Dict]) -> Optional[Dict]: """Fit the pipeline. Parameters diff --git a/junifer/markers/parcel.py b/junifer/markers/parcel.py index 6b217d2c5..755c12f15 100644 --- a/junifer/markers/parcel.py +++ b/junifer/markers/parcel.py @@ -42,7 +42,7 @@ class ParcelAggregation(BaseMarker): on = ["T1w", "BOLD", "VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"] super().__init__(on=on, name=name) - def get_output_kind(self, input: List[str]) -> str: + def get_output_kind(self, input: List[str]) -> List[str]: """Get output kind. Parameters @@ -56,13 +56,18 @@ class ParcelAggregation(BaseMarker): The kind of output. """ - if input in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]: - return "table" - if input in ["BOLD"]: - return "timeseries" + outputs = [] + for t_input in input: + if t_input in ["VBM_GM", "VBM_WM", "fALFF", "GCOR", "LCOR"]: + outputs.append("table") + elif input in ["BOLD"]: + outputs.append("timeseries") + else: + raise ValueError(f"Unknown input kind for {t_input}") + return outputs # TODO: complete type annotations - def store(self, kind: List[str], out, storage) -> None: + def store(self, kind: str, out, storage) -> None: """Store. Parameters @@ -98,7 +103,7 @@ class ParcelAggregation(BaseMarker): self.method, func_params=self.method_params ) # Get the min of the voxels sizes and use it as the resolution - resolution = np.min(t_input.header.get_zooms()[:3]) + resolution = np.min(t_input.header.get_zooms()[:3]) # type: ignore t_atlas, t_labels, _ = load_atlas(self.atlas, resolution=resolution) atlas_img_res = resample_to_img( t_atlas, @@ -110,7 +115,8 @@ class ParcelAggregation(BaseMarker): img=atlas_img_res, ) logger.debug("Masking") - masker = NiftiMasker(atlas_bin, target_affine=t_input.affine) + masker = NiftiMasker( + atlas_bin, target_affine=t_input.affine) # type: ignore # Mask the input data and the atlas data = masker.fit_transform(t_input) diff --git a/junifer/markers/pipeline_mixin.py b/junifer/markers/pipeline_mixin.py index 22cd52da4..915458810 100644 --- a/junifer/markers/pipeline_mixin.py +++ b/junifer/markers/pipeline_mixin.py @@ -91,7 +91,7 @@ class PipelineStepMixin: self.validate_input(input=input) return self.get_output_kind(input=input) - def fit_transform(self, input: List[str]) -> None: + def fit_transform(self, input: Dict[str, Dict]) -> Dict[str, Dict]: """Fit and transform. Parameters diff --git a/junifer/markers/tests/test_collection.py b/junifer/markers/tests/test_collection.py index 4645ab7ae..224bc6f04 100644 --- a/junifer/markers/tests/test_collection.py +++ b/junifer/markers/tests/test_collection.py @@ -86,11 +86,13 @@ def test_marker_collection(): with dg: input = dg["sub-01"] out2 = mc2.fit(input) + assert out2 is not None for t_marker in markers: t_name = t_marker.name assert_array_equal( - out[t_name]["VBM_GM"]["data"], out2[t_name]["VBM_GM"]["data"] - ) # type: ignore + out[t_name]["VBM_GM"]["data"], + out2[t_name]["VBM_GM"]["data"] + ) def test_MarkerCollection_storage(tmp_path) -> None: @@ -125,6 +127,7 @@ def test_MarkerCollection_storage(tmp_path) -> None: markers=markers, storage=storage, datareader=DefaultDataReader() ) mc.validate(dg) + assert mc._storage is not None assert mc._storage.uri == storage.uri with dg: input = dg["sub-01"] diff --git a/junifer/markers/tests/test_markers_base.py b/junifer/markers/tests/test_markers_base.py index f91ac3eb0..c6459d625 100644 --- a/junifer/markers/tests/test_markers_base.py +++ b/junifer/markers/tests/test_markers_base.py @@ -51,19 +51,19 @@ def test_BaseMarker() -> None: """Test base class.""" base = BaseMarker(on=["bold", "dwi"], name="mymarker") input_ = {"bold": {"path": "test"}, "t2": {"path": "test"}} - base.validate_input(input_) + base.validate_input(list(input_.keys())) wrong_input = {"t2": {"path": "test"}} with pytest.raises(ValueError): - base.validate_input(wrong_input) + base.validate_input(list(wrong_input.keys())) - output = base.get_output_kind(input_) - assert output is None + with pytest.raises(NotImplementedError): + base.get_output_kind(list(wrong_input.keys())) with pytest.raises(NotImplementedError): base.fit_transform(input_) - base.compute = lambda x: {"data": 1} + base.compute = lambda x: {"data": 1} # type: ignore out = base.fit_transform(input_) assert out["bold"]["data"] == 1 @@ -71,7 +71,7 @@ def test_BaseMarker() -> None: assert out["bold"]["meta"]["marker"]["class"] == "BaseMarker" base2 = BaseMarker(on="bold", name="mymarker") - base2.compute = lambda x: {"data": 1} + base2.compute = lambda x: {"data": 1} # type: ignore out2 = base2.fit_transform(input_) assert out2["bold"]["data"] == 1 assert out2["bold"]["meta"]["marker"]["name"] == "bold_mymarker" diff --git a/junifer/markers/tests/test_pipeline_mixin.py b/junifer/markers/tests/test_pipeline_mixin.py index ace380e9a..7086a7805 100644 --- a/junifer/markers/tests/test_pipeline_mixin.py +++ b/junifer/markers/tests/test_pipeline_mixin.py @@ -13,11 +13,11 @@ def test_PipelineStepMixin() -> None: """Test PipelineStepMixin.""" mixin = PipelineStepMixin() with pytest.raises(NotImplementedError): - mixin.validate_input(None) + mixin.validate_input([]) with pytest.raises(NotImplementedError): - mixin.get_output_kind(None) + mixin.get_output_kind([]) with pytest.raises(NotImplementedError): - mixin.fit_transform(None) + mixin.fit_transform({}) def test_pipeline_step_mixin_meta(): diff --git a/junifer/preprocess/confounds.py b/junifer/preprocess/confounds.py index 869097b08..e8f413947 100644 --- a/junifer/preprocess/confounds.py +++ b/junifer/preprocess/confounds.py @@ -5,7 +5,7 @@ # Synchon Mandal # License: AGPL -from typing import TYPE_CHECKING, Dict, List, Optional +from typing import TYPE_CHECKING, Dict, List, Optional, Union import numpy as np import pandas as pd @@ -18,7 +18,7 @@ from ..utils import logger, raise_error if TYPE_CHECKING: - from nibabel import Nifti1Image + from nibabel import Nifti1Image, Nifti2Image, MGHImage class BaseConfoundRemover(PipelineStepMixin): @@ -218,7 +218,8 @@ class BaseConfoundRemover(PipelineStepMixin): def _remove_confounds( self, bold_img: "Nifti1Image", confounds_df: pd.DataFrame - ) -> "Nifti1Image": + ) -> Union["Nifti1Image", "Nifti2Image", "MGHImage", List]: + """Remove confounds from the BOLD data.""" """Remove confounds from the BOLD image. Parameters @@ -242,7 +243,7 @@ class BaseConfoundRemover(PipelineStepMixin): t_r = self.t_r if t_r is None: logger.info("No `t_r` specified, using t_r from nifti header") - zooms = bold_img.header.get_zooms() + zooms = bold_img.header.get_zooms() # type: ignore t_r = zooms[3] logger.info( f"Read t_r from nifti header: {t_r}", diff --git a/junifer/preprocess/tests/test_confounds.py b/junifer/preprocess/tests/test_confounds.py index 01463cdd9..8050e1d2d 100644 --- a/junifer/preprocess/tests/test_confounds.py +++ b/junifer/preprocess/tests/test_confounds.py @@ -140,8 +140,8 @@ def test_baseconfoundremover() -> None: cr = BaseConfoundRemover( strategy=strat1, spike=None, mask_img=simsk, t_r=0.75 ) - cr.validate_input(input_data_obj.keys()) - out_type = cr.get_output_kind(input_data_obj.keys()) + cr.validate_input(list(input_data_obj.keys())) + out_type = cr.get_output_kind(list(input_data_obj.keys())) assert "BOLD" in out_type @@ -175,8 +175,8 @@ def test_baseconfoundremover() -> None: cr = BaseConfoundRemover( strategy=strat2, spike=None, mask_img=simsk, t_r=0.75 ) - cr.validate_input(input_data_obj.keys()) - out_type = cr.get_output_kind(input_data_obj.keys()) + cr.validate_input(list(input_data_obj.keys())) + out_type = cr.get_output_kind(list(input_data_obj.keys())) assert "BOLD" in out_type @@ -207,8 +207,8 @@ def test_baseconfoundremover() -> None: cr = BaseConfoundRemover( strategy=strat3, spike=None, mask_img=simsk, t_r=0.75 ) - cr.validate_input(input_data_obj.keys()) - out_type = cr.get_output_kind(input_data_obj.keys()) + cr.validate_input(list(input_data_obj.keys())) + out_type = cr.get_output_kind(list(input_data_obj.keys())) assert "BOLD" in out_type @@ -240,8 +240,8 @@ def test_baseconfoundremover() -> None: cr = BaseConfoundRemover( strategy=strat4, spike=None, mask_img=simsk, t_r=0.75 ) - cr.validate_input(input_data_obj.keys()) - out_type = cr.get_output_kind(input_data_obj.keys()) + cr.validate_input(list(input_data_obj.keys())) + out_type = cr.get_output_kind(list(input_data_obj.keys())) assert "BOLD" in out_type @@ -273,8 +273,8 @@ def test_baseconfoundremover() -> None: cr = BaseConfoundRemover( strategy=strat5, spike=None, mask_img=simsk, t_r=0.75 ) - cr.validate_input(input_data_obj.keys()) - out_type = cr.get_output_kind(input_data_obj.keys()) + cr.validate_input(list(input_data_obj.keys())) + out_type = cr.get_output_kind(list(input_data_obj.keys())) assert "BOLD" in out_type @@ -320,8 +320,8 @@ def test_baseconfoundremover() -> None: cr = BaseConfoundRemover( strategy=strat6, spike=0.75, mask_img=simsk, t_r=0.75 ) - cr.validate_input(input_data_obj.keys()) - out_type = cr.get_output_kind(input_data_obj.keys()) + cr.validate_input(list(input_data_obj.keys())) + out_type = cr.get_output_kind(list(input_data_obj.keys())) assert "BOLD" in out_type diff --git a/junifer/stats.py b/junifer/stats.py index e0a30c0ce..aa55e7475 100644 --- a/junifer/stats.py +++ b/junifer/stats.py @@ -5,7 +5,7 @@ # License: AGPL from functools import partial -from typing import Callable, Dict +from typing import Callable, Dict, Any, Optional import numpy as np from scipy.stats import trim_mean @@ -14,7 +14,9 @@ from scipy.stats.mstats import winsorize from .utils import logger, raise_error -def get_aggfunc_by_name(name: str, func_params: Dict) -> Callable: +def get_aggfunc_by_name( + name: str, + func_params: Optional[Dict[str, Any]]) -> Callable: """Get an aggregation function by its name. Parameters @@ -38,11 +40,24 @@ def get_aggfunc_by_name(name: str, func_params: Dict) -> Callable: """ # check validity of names _valid_func_names = {"winsorized_mean", "mean", "std", "trim_mean"} - + if func_params is None: + func_params = {} # apply functions if name == "winsorized_mean": # check validity of func_params limits = func_params.get("limits") + if limits is None or not isinstance(limits, list): + raise_error( + "func_params must contain a list of limits for " + "winsorized_mean", + ValueError, + ) + if len(limits) != 2: + raise_error( + "func_params must contain a list of two limits for " + "winsorized_mean", + ValueError, + ) if all((lim >= 0.0 and lim <= 1) for lim in limits): logger.info(f"Limits for winsorized mean are set to {limits}.") else: @@ -69,7 +84,7 @@ def get_aggfunc_by_name(name: str, func_params: Dict) -> Callable: def winsorized_mean( - data: np.ndarray, axis: int = None, **win_params + data: np.ndarray, axis: Optional[int] = None, **win_params ) -> np.ndarray: """Compute a winsorized mean by chaining winsorization and mean. diff --git a/junifer/storage/base.py b/junifer/storage/base.py index 9130eeb02..46a2f0e40 100644 --- a/junifer/storage/base.py +++ b/junifer/storage/base.py @@ -182,7 +182,7 @@ class BaseFeatureStorage(ABC): data, meta: Dict, columns: Optional[Iterable[str]] = None, - rows_col_name: str = None, + rows_col_name: Optional[str] = None, ) -> None: """Store table. diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py index f15903982..b08213323 100644 --- a/junifer/storage/sqlite.py +++ b/junifer/storage/sqlite.py @@ -20,7 +20,7 @@ from .utils import element_to_index, element_to_prefix, process_meta if TYPE_CHECKING: - from sqlalchemy import Engine + from sqlalchemy.engine import Engine @register_storage @@ -103,17 +103,18 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): meta = {} # Retrieve element key from metadata element = meta.get("element", None) - # Functionality check - if self.single_output is False and element is None: - raise_error( - msg="element must be specified when single_output is False." - ) # Prefixed elements prefix = "" if self.single_output is False: + if element is None: + raise_error( + msg="element must be specified when" + "single_output is False." + ) prefix = element_to_prefix(element) # Format URI for engine creation - uri = f"sqlite:///{self.uri.parent}/{prefix}{self.uri.name}" + uri = "sqlite:///" \ + f"{self.uri.parent}/{prefix}{self.uri.name}" # type: ignore return create_engine(uri, echo=False) def _save_upsert( @@ -211,7 +212,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): data, meta: Dict, columns: Optional[Iterable[str]] = None, - rows_col_name: str = None, + rows_col_name: Optional[str] = None, ) -> None: """Store 2D dataframe. @@ -234,7 +235,8 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): meta=meta, n_rows=n_rows, rows_col_name=rows_col_name ) # Prepare new dataframe - data_df = pd.DataFrame(data, columns=columns, index=idx) + data_df = pd.DataFrame( + data, columns=columns, index=idx) # type: ignore # Store dataframe self.store_df(df=data_df, meta=meta) @@ -405,7 +407,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): data, meta: Dict, col_names: Optional[Iterable[str]] = None, - rows_col_name: str = None, + rows_col_name: Optional[str] = None, ) -> None: """Implement 2D matrix storing. @@ -433,7 +435,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): data, meta: Dict, columns: Optional[Iterable[str]] = None, - rows_col_name: str = None, + rows_col_name: Optional[str] = None, ) -> None: """Implement table storing. @@ -528,13 +530,14 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): """ if self.single_output is True: raise_error(msg="collect() is not implemented for single output.") - logger.info(f"Collecting data from {self.uri.parent}/*{self.uri.name}") + logger.info("Collecting data from " + f"{self.uri.parent}/*{self.uri.name}") # type: ignore # Create new instance out_storage = SQLiteFeatureStorage( uri=self.uri, single_output=True, upsert="ignore" ) # Glob files - files = self.uri.parent.glob(f"*{self.uri.name}") + files = self.uri.parent.glob(f"*{self.uri.name}") # type: ignore for elem in tqdm(files, desc="file"): logger.debug(f"Reading from {str(elem.absolute())}") in_storage = SQLiteFeatureStorage(uri=elem, single_output=True) diff --git a/junifer/storage/tests/test_sqlite.py b/junifer/storage/tests/test_sqlite.py index aa33b9a8d..542b27632 100644 --- a/junifer/storage/tests/test_sqlite.py +++ b/junifer/storage/tests/test_sqlite.py @@ -4,6 +4,8 @@ # Synchon Mandal # License: AGPL +from typing import Union, List + from pathlib import Path import numpy as np @@ -57,7 +59,8 @@ df_ignore = pd.DataFrame( ).set_index(["element", "pk2"]) -def _read_sql(table_name: str, uri: str, index_col: str) -> pd.DataFrame: +def _read_sql(table_name: str, uri: str, + index_col: Union[str, List[str]]) -> pd.DataFrame: """Read database table into a pandas DataFrame. Parameters @@ -157,7 +160,8 @@ def test_upsert_replace(tmp_path: Path) -> None: table_name = storage.store_metadata(meta) # Read stored table c_df1 = _read_sql( - table_name=table_name, uri=uri, index_col=["element", "pk2"] + table_name=table_name, uri=uri.as_posix(), + index_col=["element", "pk2"] ) # Check if dataframes are equal assert_frame_equal(df1, c_df1) @@ -165,7 +169,8 @@ def test_upsert_replace(tmp_path: Path) -> None: storage._save_upsert(df=df2, name=table_name, if_exists="replace") # Read stored table c_df2 = _read_sql( - table_name=table_name, uri=uri, index_col=["element", "pk2"] + table_name=table_name, uri=uri.as_posix(), + index_col=["element", "pk2"] ) # Check if dataframes are equal assert_frame_equal(df2, c_df2) @@ -193,7 +198,9 @@ def test_upsert_ignore(tmp_path: Path) -> None: table_name = storage.store_metadata(meta=meta) # Read stored table c_df1 = _read_sql( - table_name=table_name, uri=uri, index_col=["element", "pk2"] + table_name=table_name, + uri=uri.as_posix(), + index_col=["element", "pk2"] ) # Check if dataframes are equal assert_frame_equal(df1, c_df1) @@ -201,7 +208,8 @@ def test_upsert_ignore(tmp_path: Path) -> None: with pytest.warns(RuntimeWarning, match="are already present"): storage.store_df(df2, meta) # Read stored table - c_dfignore = _read_sql(table_name, uri=uri, index_col=["element", "pk2"]) + c_dfignore = _read_sql( + table_name, uri=uri.as_posix(), index_col=["element", "pk2"]) # Check if dataframes are equal assert_frame_equal(c_dfignore, df_ignore) # Check for error @@ -228,14 +236,17 @@ def test_upsert_update(tmp_path: Path) -> None: table_name = storage.store_metadata(meta) # Read stored table c_df1 = _read_sql( - table_name=table_name, uri=uri, index_col=["element", "pk2"] + table_name=table_name, uri=uri.as_posix(), + index_col=["element", "pk2"] ) # Check if dataframes are equal assert_frame_equal(df1, c_df1) # Save to database storage.store_df(df2, meta) # Read stored table - c_dfupdate = _read_sql(table_name, uri=uri, index_col=["element", "pk2"]) + c_dfupdate = _read_sql( + table_name, uri=uri.as_posix(), + index_col=["element", "pk2"]) # Check if dataframes are equal assert_frame_equal(c_dfupdate, df_update) @@ -370,7 +381,8 @@ def test_store_table(tmp_path: Path) -> None: table_name = storage.store_metadata(meta) # Read stored table c_df = _read_sql( - table_name=table_name, uri=uri, index_col=["element", "scan"] + table_name=table_name, uri=uri.as_posix(), + index_col=["element", "scan"] ) # Check if dataframes are equal assert_frame_equal(df, c_df) @@ -388,7 +400,8 @@ def test_store_table(tmp_path: Path) -> None: ) # Read stored table c_df_new = _read_sql( - table_name=table_name, uri=uri, index_col=["element", "scan"] + table_name=table_name, uri=uri.as_posix(), + index_col=["element", "scan"] ) # Check if dataframes are equal assert_frame_equal(df_new, c_df_new) @@ -484,9 +497,9 @@ def test_store_multiple_output(tmp_path: Path): # Set index columns cols = ["subject", "session", "scan"] # Read stored tables - cdf1 = _read_sql(table_name, uri1, index_col=cols) - cdf2 = _read_sql(table_name, uri2, index_col=cols) - cdf3 = _read_sql(table_name, uri3, index_col=cols) + cdf1 = _read_sql(table_name, uri1.as_posix(), index_col=cols) + cdf2 = _read_sql(table_name, uri2.as_posix(), index_col=cols) + cdf3 = _read_sql(table_name, uri3.as_posix(), index_col=cols) # Check if dataframes are equal assert_frame_equal(df1, cdf1) assert_frame_equal(df2, cdf2) @@ -566,10 +579,10 @@ def test_collect(tmp_path: Path) -> None: # Store metadata table_name = storage.store_metadata(meta1) # Read stored tables - all_df = _read_sql(table_name, uri, index_col=cols) - cdf1 = _read_sql(table_name, uri1, index_col=cols) - cdf2 = _read_sql(table_name, uri2, index_col=cols) - cdf3 = _read_sql(table_name, uri3, index_col=cols) + all_df = _read_sql(table_name, uri.as_posix(), index_col=cols) + cdf1 = _read_sql(table_name, uri1.as_posix(), index_col=cols) + cdf2 = _read_sql(table_name, uri2.as_posix(), index_col=cols) + cdf3 = _read_sql(table_name, uri3.as_posix(), index_col=cols) # Operate on retrieved tables all_cdf = pd.concat([cdf1, cdf2, cdf3]) all_df.sort_index(level=cols, inplace=True) diff --git a/junifer/storage/tests/test_storage_base.py b/junifer/storage/tests/test_storage_base.py index 6a69d9503..c9af9db04 100644 --- a/junifer/storage/tests/test_storage_base.py +++ b/junifer/storage/tests/test_storage_base.py @@ -12,7 +12,7 @@ from junifer.storage.base import BaseFeatureStorage def test_BaseFeatureStorage_abstractness() -> None: """Test BaseFeatureStorage is abstract base class.""" with pytest.raises(TypeError, match=r"abstract"): - BaseFeatureStorage(uri="/tmp") + BaseFeatureStorage(uri="/tmp") # type: ignore def test_BaseFeatureStorage() -> None: @@ -74,7 +74,7 @@ def test_BaseFeatureStorage() -> None: st.store_table(None, None) with pytest.raises(NotImplementedError): - st.store_df(None, None) + st.store_df(None, None) # type: ignore with pytest.raises(NotImplementedError): st.store_timeseries(None, None) diff --git a/junifer/storage/tests/test_utils.py b/junifer/storage/tests/test_utils.py index 751410b85..a5b2b466d 100644 --- a/junifer/storage/tests/test_utils.py +++ b/junifer/storage/tests/test_utils.py @@ -19,7 +19,7 @@ def test_process_meta_invalid_metadata_type() -> None: """Test invalid metadata type check for metadata hash processing.""" meta = None with pytest.raises(ValueError, match=r"`meta` must be a dict"): - process_meta(meta) + process_meta(meta) # type: ignore # TODO: parameterize diff --git a/junifer/storage/utils.py b/junifer/storage/utils.py index b707e80ea..1d7585787 100644 --- a/junifer/storage/utils.py +++ b/junifer/storage/utils.py @@ -4,6 +4,8 @@ # Synchon Mandal # License: AGPL +from typing import Any + import hashlib import json from typing import Dict, Optional, Tuple, Union @@ -155,7 +157,7 @@ def element_to_index( # Check rows_col_name if rows_col_name is None: rows_col_name = "idx" - elem_idx = {k: [v] * n_rows for k, v in element.items()} + elem_idx: Dict[Any, Any] = {k: [v] * n_rows for k, v in element.items()} elem_idx[rows_col_name] = np.arange(n_rows) # Create index index = pd.MultiIndex.from_frame( diff --git a/junifer/utils/logging.py b/junifer/utils/logging.py index bb4e5627e..0dc6c0ccd 100644 --- a/junifer/utils/logging.py +++ b/junifer/utils/logging.py @@ -23,7 +23,7 @@ _logging_types = { } -class WrapStdOut: +class WrapStdOut(logging.StreamHandler): """ Dynamically wrap to sys.stdout. @@ -67,7 +67,7 @@ def _get_git_head(path: Path) -> str: ) try: stdout, _ = process.communicate(timeout=10) - proc_stdout = stdout.strip() + proc_stdout = stdout.strip().decode() except TimeoutExpired: process.kill() proc_stdout = "" @@ -95,7 +95,7 @@ def get_versions() -> Dict: if module_version is None: module_version = None elif "git" in module_version: - git_path = Path(module.__file__).resolve().parent + git_path = Path(module.__file__).resolve().parent # type: ignore head = _get_git_head(git_path) module_version += f"-HEAD:{head}" @@ -240,7 +240,7 @@ def configure_logging( mode = "w" if overwrite else "a" lh = logging.FileHandler(fname, mode=mode) else: - lh = logging.StreamHandler(WrapStdOut()) + lh = logging.StreamHandler(WrapStdOut()) # type: ignore # Set logging format if output_format is None: @@ -272,7 +272,9 @@ def raise_error(msg: str, klass: Type[Exception] = ValueError) -> NoReturn: raise klass(msg) -def warn_with_log(msg: str, category: Warning = RuntimeWarning) -> None: +def warn_with_log( + msg: str, + category: Optional[Type[Warning]] = RuntimeWarning) -> None: """Warn, but first log it. Parameters -- 2.52.0 From 63a993f6fcb7b408811c11f70a091d6f9bbd99a5 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Thu, 8 Sep 2022 15:20:26 +0200 Subject: [PATCH 257/287] Fix doc --- docs/api.rst | 2 +- examples/norun_hcpfc_pearson.py | 3 +-- 2 files changed, 2 insertions(+), 3 deletions(-) diff --git a/docs/api.rst b/docs/api.rst index 3aceb9ac7..a20d04331 100644 --- a/docs/api.rst +++ b/docs/api.rst @@ -13,5 +13,5 @@ Data Grabbers :members: .. autoclass:: junifer.datagrabber.datalad_base.DataladDataGrabber :members: -.. autoclass:: junifer.datagrabber.pattern_datalad_base.PatternDataladDataGrabber +.. autoclass:: junifer.datagrabber.pattern_datalad.PatternDataladDataGrabber :members: diff --git a/examples/norun_hcpfc_pearson.py b/examples/norun_hcpfc_pearson.py index 6875e23b9..e3c98e1ca 100644 --- a/examples/norun_hcpfc_pearson.py +++ b/examples/norun_hcpfc_pearson.py @@ -1,5 +1,4 @@ -"""extraction of FC from HCP data. - +""" HCP FC Extraction ====================== -- 2.52.0 From 888bb00012c80f4a6d89e78d2e868e62677435ae Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Thu, 8 Sep 2022 15:21:49 +0200 Subject: [PATCH 258/287] Dman flake --- junifer/utils/logging.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/junifer/utils/logging.py b/junifer/utils/logging.py index 0dc6c0ccd..bcb79ef3e 100644 --- a/junifer/utils/logging.py +++ b/junifer/utils/logging.py @@ -273,7 +273,7 @@ def raise_error(msg: str, klass: Type[Exception] = ValueError) -> NoReturn: def warn_with_log( - msg: str, + msg: str, category: Optional[Type[Warning]] = RuntimeWarning) -> None: """Warn, but first log it. -- 2.52.0 From ece51af1b02a86f7dc6acc57fd2aa70e7cb4e23b Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Thu, 8 Sep 2022 16:48:05 +0200 Subject: [PATCH 259/287] docs: update README.md with shields and content --- README.md | 54 ++++++++++++++++++++++++++++++++++++++++++++++++++---- 1 file changed, 50 insertions(+), 4 deletions(-) diff --git a/README.md b/README.md index 64d11f49a..4c4773c31 100644 --- a/README.md +++ b/README.md @@ -1,5 +1,18 @@ -# python-library-mockup -JUelich NeuroImaging FEature extractoR +# junifer - JUelich NeuroImaging FEature extractoR + +![PyPI](https://img.shields.io/pypi/v/junifer?style=flat-square) +![PyPI - Python Version](https://img.shields.io/pypi/pyversions/junifer?style=flat-square) +![PyPI - Wheel](https://img.shields.io/pypi/wheel/junifer?style=flat-square) +![GitHub](https://img.shields.io/github/license/juaml/junifer?style=flat-square) +[![codecov](https://codecov.io/gh/juaml/junifer/branch/main/graph/badge.svg?token=5H21JuZXMw)](https://codecov.io/gh/juaml/junifer) + +## About + +junifer is a data handling and feature extraction library targeted towards neuroimaging data specifically functional MRI data. + +It is curently being developed and maintained at the [Applied Machine Learning](https://www.fz-juelich.de/en/inm/inm-7/research-groups/applied-machine-learning-aml) group at [Forschungszentrum Juelich](https://www.fz-juelich.de/en), Germany. Although the library is designed for people working at [Institute of Neuroscience and Medicine - Brain and Behaviour (INM-7)](https://www.fz-juelich.de/en/inm/inm-7), it is designed to be as modular as possible thus enabling others to extend it easily. + +The documentation is available at [https://juaml.github.io/junifer](https://juaml.github.io/junifer/main/index.html). ## Repository Organization @@ -15,6 +28,39 @@ JUelich NeuroImaging FEature extractoR * `preprocess`: Preprocessing module. * `storage`: Storage module. * `utils`: Utilities module (e.g. logging) - - \ No newline at end of file + +## Installation + +Use `pip` to install from PyPI like so: + +``` +pip install junifer +``` + +## Citation + +If you use junifer in a scientific publication, we would appreciate if you cite our work. Currently, we do not have a publication, so feel free to use the project [URL](https://juaml.github.io/junifer). + +## Contribution + +Contributions are welcome and greatly appreciated. Please read the [guidelines](https://juaml.github.io/junifer/main/contributing.html) to get started. + +## License + +junifer is released under the AGPL v3 license: + +julearn, FZJuelich AML neuroimaging feature extraction library. +Copyright (C) 2022, authors of junifer. + +This program is free software: you can redistribute it and/or modify +it under the terms of the GNU Affero General Public License as published by +the Free Software Foundation, either version 3 of the License, or any later version. + +This program is distributed in the hope that it will be useful, +but WITHOUT ANY WARRANTY; without even the implied warranty of +MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +GNU Affero General Public License for more details. + +You should have received a copy of the GNU Affero General Public License +along with this program. If not, see . -- 2.52.0 From c183095f3151a637135dce2e7dcd8c2220affa5b Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 9 Sep 2022 15:10:22 +0200 Subject: [PATCH 260/287] docs: add about for junifer in index.rst --- docs/index.rst | 13 +++++++++++-- 1 file changed, 11 insertions(+), 2 deletions(-) diff --git a/docs/index.rst b/docs/index.rst index 3deeb2c34..1c44c3944 100644 --- a/docs/index.rst +++ b/docs/index.rst @@ -1,8 +1,17 @@ .. include:: links.inc -Welcome to the documentation! -============================= +Welcome to junifer's documentation! +=================================== +junifer (JUelich NeuroImaging FEature extractoR) is a data handling and feature +extraction library targeted towards neuroimaging data specifically functional +MRI data. + +It is curently being developed and maintained at the Applied Machine Learning +(`AML`_) group at Forschungszentrum Juelich, Germany. Although the library is +designed for people working at Institute of Neuroscience and Medicine - Brain +and Behaviour (`INM-7`_), it is designed to be as modular as possible thus +enabling others to extend it easily. .. toctree:: :maxdepth: 2 -- 2.52.0 From 562189b1a7de0bcc48e2002095129cf8d3b25489 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 9 Sep 2022 16:02:25 +0200 Subject: [PATCH 261/287] docs: update installation.rst --- docs/installation.rst | 86 ++++++++++++------------------------------- 1 file changed, 24 insertions(+), 62 deletions(-) diff --git a/docs/installation.rst b/docs/installation.rst index 99ddee684..d257b6042 100644 --- a/docs/installation.rst +++ b/docs/installation.rst @@ -1,91 +1,53 @@ .. include:: links.inc -Installing -========== +Installing junifer +================== Requirements ^^^^^^^^^^^^ -junifer requires the following packages: +junifer is compatible with `Python`_ >= 3.8 and requires the following packages: -Running the examples requires: +* click>=8.1.3,<8.2 +* numpy>=1.22,<1.23 +* datalad>=0.15.4,<0.18 +* pandas>=1.4.0,<1.5 +* nibabel>=3.2.0,<4.1 +* nilearn>=0.9.0,<1.0 +* sqlalchemy>=1.4.27,<= 1.5.0 +* pyyaml>=5.1.2,<7.0 -Depending on the installation method, this packages might be installed -automatically. +Depending on the installation method, these packages might be installed automatically. -Installing -^^^^^^^^^^ -There are different ways to install junifer: +Installation +^^^^^^^^^^^^ +Depending on your use-case, junifer can be installed differently. * Install the :ref:`install_latest_release`. This is the most suitable approach - for most end users. -* Install the :ref:`install_latest_development`. This version will have the - latest features. However, it is still under development and not yet - officially released. Some features might still change before the next stable - release. -* Install from :ref:`install_development_git`. This is mostly suitable for - developers that want to have the latest version and yet edit the code. + for end users. +* Install from :ref:`install_development_git`. This is mostly suitable approach + for developers. -Either way, we strongly recommend using virtual environments: - -* `venv`_ -* `conda env`_ +Either way, we strongly recommend using `virtual environments `_. .. _install_latest_release: -Latest release +Stable release -------------- -We have packaged junifer and published it in PyPi, so you can just install it -with `pip`. +Use `pip` to install julearn from PyPI, like so: .. code-block:: bash pip install -U junifer -.. _install_latest_development: - -Latest Development Version --------------------------- -First, make sure that you have all the dependencies installed: - -Then, install junifer from TestPypi - -.. code-block:: bash - - pip install -U junifer --pre - - .. _install_development_git: -Local git repository (for developers) -------------------------------------- -First, make sure that you have all the dependencies installed: +Local Git repository +-------------------- -Then, clone `junifer Github`_ repository in a folder of your choice: - -.. code-block:: bash - - git clone https://github.com/juaml/junifer.git - -Install development mode requirements: - -.. code-block:: bash - - cd junifer - pip install -r dev-requirements.txt - -Finally, install in development mode: - -.. code-block:: bash - - python setup.py develop - -.. note:: Every time that you run ``setup.py develop``, the version is going to - be automatically set based on the git history. Nevertheless, this change - should not be committed (changes to ``_version.py``). Running ``git stash`` - at this point will forget the local changes to ``_version.py``. +Follow the `detailed contribution guidelines `_. -- 2.52.0 From 5ad92fa2f527b6e1c027faf69f3b91992db5e939 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 9 Sep 2022 16:03:03 +0200 Subject: [PATCH 262/287] docs: add contribution.rst --- docs/contribution.rst | 187 ++++++++++++++++++++++++++++++++++++++++++ docs/index.rst | 1 + 2 files changed, 188 insertions(+) create mode 100644 docs/contribution.rst diff --git a/docs/contribution.rst b/docs/contribution.rst new file mode 100644 index 000000000..e68c8123d --- /dev/null +++ b/docs/contribution.rst @@ -0,0 +1,187 @@ +.. include:: links.inc + +Contributing to junifer +======================= + + +Setting up the local development environment +-------------------------------------------- + +1. Fork the https://github.com/juaml/junifer repository on GitHub. If you + have never done this before, `follow the official guide + `_. +2. Clone your fork locally as described in the same guide. +3. Install your local copy into a Python virtual environment. You can `read + this guide to learn more + `_ about them + and how to create one. + + .. code-block:: console + + pip install -e ".[dev]" + +4. Create a branch for local development using the ``dev`` branch as a + starting point. Use ``fix``, ``refactor``, or ``feat`` as a prefix. + + .. code-block:: console + + git checkout dev + git checkout -b / + + Now you can make your changes locally. + +5. When making changes locally, it is helpful to ``git commit`` your work + regularly. On one hand to save your work and on the other hand, the smaller + the steps, the easier it is to review your work later. Please use `semantic + commit messages + `_. + + .. code-block:: console + + git add . + git commit -m ": " + +6. When you're done making changes, check that your changes pass our test suite. + This is all included with ``tox``. + + .. code-block:: console + + tox + + You can also run all ``tox`` tests in parallel. As of ``tox 3.7``, you can run:: + + .. code-block:: console + + tox --parallel + + +7. Push your branch to GitHub. + + .. code-block:: console + + git push origin / + +8. Open the link displayed in the message when pushing your new branch in order + to submit a pull request. Please follow the template presented to you in the + web interface to complete your pull request. + + +GitHub Pull Request guidelines +------------------------------ + +Before you submit a pull request, check that it meets these guidelines: + +1. The pull request should include tests in the respective ``tests`` directory. + Except in rare circumstances, code coverage must not decrease (as reported + by codecov which runs automatically when you submit your pull request). +2. If the pull request adds functionality, the docs should be + updated. Consider creating a Python file that demonstrates the usage in + ``examples`` directory. +3. The pull request should also include a short one-liner of your contribution + in `docs/changes/latest.inc`. If it's your first contribution, also add + yourself to `docs/changes/contributors.inc`. +3. The pull request will be tested for several different Python versions. +4. Someone from the core team will review your work and guide you to a successful + contribution. + + +Running unit tests +------------------ + +junifer uses `pytest `_ for its +unit-tests and new features should in general always come with new +tests that make sure that the code runs as intended. + +To run all tests:: + +.. code-block:: console + + tox -e test + + +Adding and building documentation +--------------------------------- + +Building the documentation requires some extra packages and can be installed by:: + +.. code-block:: console + + pip install -e ".[docs]" + +To build the docs:: + +.. code-block:: bash + + cd docs + make html + +To view the documentation, open `docs/_build/html/index.html`. + +In case you remove some files or change their filenames, you can run into +errors when using ``make html``. In this situation you can use ``make clean`` +to clean up the already build files and then re-run ``make html``. + + +Writing Examples +---------------- + +Examples are run and displayed in HTML format using `sphinx gallery`_. To add an +example, just create a ``.py`` file that starts either with ``plot_`` or ``run_``, +dependending on whether the example generates a figure or not. + +The first lines of the example should be a python block comment with a title, +a description of the example an the following include directive to be able to +use the links. + +The format used for text is reST. Check the `sphinx reST reference`_ for more +details. + +Example of the first lines: + + +.. code-block:: python + + """ + Simple Binary Classification + ============================ + + This example uses the 'iris' dataset and performs a simple binary + classification using a Support Vector Machine classifier. + + .. include:: ../../links.inc + """ + + +The rest of the script will be executed as normal Python code. In order to +render the output and embed formatted text within the code, you need to add +a 79 ``#`` (a full line) at the point in which you want to render and add text. +Each line of text shall be preceded with ``#``. The code that is not +commented will be executed. + +The following example will create 3 texts and render the output between the +texts. + +.. code-block:: python + + ############################################################################### + # Imports needed for the example + from seaborn import load_dataset + from julearn import run_cross_validation + from julearn.utils import configure_logging + + ############################################################################### + df_iris = load_dataset('iris') + + ############################################################################### + # The dataset has three kind of species. We will keep two to perform a binary + # classification. + df_iris = df_iris[df_iris['species'].isin(['versicolor', 'virginica'])] + + +Finally, when the example is done, you can run as a normal Python script. +To generate the HTML, just build the docs: + +.. code-block:: bash + + cd docs + make html diff --git a/docs/index.rst b/docs/index.rst index 1c44c3944..7e0aa03ef 100644 --- a/docs/index.rst +++ b/docs/index.rst @@ -21,6 +21,7 @@ enabling others to extend it easily. data api auto_examples/index.rst + contribution maintaining whats_new -- 2.52.0 From c3c503f34f1fe1f9c830f25bfdeb82848787a12c Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 9 Sep 2022 16:14:30 +0200 Subject: [PATCH 263/287] docs: fix numbering in contribution.rst --- docs/contribution.rst | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/docs/contribution.rst b/docs/contribution.rst index e68c8123d..1e77d21c0 100644 --- a/docs/contribution.rst +++ b/docs/contribution.rst @@ -80,8 +80,8 @@ Before you submit a pull request, check that it meets these guidelines: 3. The pull request should also include a short one-liner of your contribution in `docs/changes/latest.inc`. If it's your first contribution, also add yourself to `docs/changes/contributors.inc`. -3. The pull request will be tested for several different Python versions. -4. Someone from the core team will review your work and guide you to a successful +4. The pull request will be tested against several Python versions. +5. Someone from the core team will review your work and guide you to a successful contribution. -- 2.52.0 From b47dd8a473af4fd4c364cc9fc9f8b552477d5f09 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Fri, 9 Sep 2022 16:14:48 +0200 Subject: [PATCH 264/287] docs: fix links.inc --- docs/links.inc | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/links.inc b/docs/links.inc index 27414754b..444454691 100644 --- a/docs/links.inc +++ b/docs/links.inc @@ -30,4 +30,4 @@ .. _`setuptools_scm`: https://github.com/pypa/setuptools_scm/ .. _`sphinx gallery`: https://sphinx-gallery.github.io/stable/index.html -.. _`sphinx RST reference`: https://www.sphinx-doc.org/en/master/usage/restructuredtext/basics.html#inline-markup \ No newline at end of file +.. _`sphinx reST reference`: https://www.sphinx-doc.org/en/master/usage/restructuredtext/basics.html#inline-markup -- 2.52.0 From 08dfd3331e9fd403ef6e211f3c3c2ccd4ea42370 Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 9 Sep 2022 18:50:10 +0200 Subject: [PATCH 265/287] Started restructuring the DOC --- docs/api.rst | 17 -------------- docs/api/api.rst | 7 ++++++ docs/api/datagrabbers.rst | 6 +++++ docs/api/datareaders.rst | 7 ++++++ docs/api/index.rst | 27 ++++++++++++++++++++++ docs/api/markers.rst | 7 ++++++ docs/api/preprocessing.rst | 7 ++++++ docs/api/storage.rst | 7 ++++++ docs/api/testing.rst | 6 +++++ docs/api/utils.rst | 7 ++++++ docs/index.rst | 4 ++-- docs/links.inc | 1 + docs/{ => understanding}/data.rst | 36 ++++++++++++++--------------- docs/understanding/index.rst | 25 ++++++++++++++++++++ junifer/datagrabber/datalad_base.py | 1 - junifer/preprocess/confounds.py | 2 +- junifer/storage/sqlite.py | 12 ++++------ 17 files changed, 131 insertions(+), 48 deletions(-) delete mode 100644 docs/api.rst create mode 100644 docs/api/api.rst create mode 100644 docs/api/datagrabbers.rst create mode 100644 docs/api/datareaders.rst create mode 100644 docs/api/index.rst create mode 100644 docs/api/markers.rst create mode 100644 docs/api/preprocessing.rst create mode 100644 docs/api/storage.rst create mode 100644 docs/api/testing.rst create mode 100644 docs/api/utils.rst rename docs/{ => understanding}/data.rst (93%) create mode 100644 docs/understanding/index.rst diff --git a/docs/api.rst b/docs/api.rst deleted file mode 100644 index a20d04331..000000000 --- a/docs/api.rst +++ /dev/null @@ -1,17 +0,0 @@ -# Authors: Federico Raimondo -# License: AGPL -Reference -========= -.. include:: links.inc - -Data Grabbers -^^^^^^^^^^^^^ - -.. autoclass:: junifer.datagrabber.base.BaseDataGrabber - :members: -.. autoclass:: junifer.datagrabber.pattern.PatternDataGrabber - :members: -.. autoclass:: junifer.datagrabber.datalad_base.DataladDataGrabber - :members: -.. autoclass:: junifer.datagrabber.pattern_datalad.PatternDataladDataGrabber - :members: diff --git a/docs/api/api.rst b/docs/api/api.rst new file mode 100644 index 000000000..b1a96e243 --- /dev/null +++ b/docs/api/api.rst @@ -0,0 +1,7 @@ + +API Functions +^^^^^^^^^^^^^^ + +.. automodule:: junifer.api + :members: + :imported-members: diff --git a/docs/api/datagrabbers.rst b/docs/api/datagrabbers.rst new file mode 100644 index 000000000..974bd303b --- /dev/null +++ b/docs/api/datagrabbers.rst @@ -0,0 +1,6 @@ +Data Grabbers +^^^^^^^^^^^^^ + +.. automodule:: junifer.datagrabber + :members: + :imported-members: diff --git a/docs/api/datareaders.rst b/docs/api/datareaders.rst new file mode 100644 index 000000000..3dee2aa11 --- /dev/null +++ b/docs/api/datareaders.rst @@ -0,0 +1,7 @@ + +Data Readers +^^^^^^^^^^^^ + +.. automodule:: junifer.datareader + :members: + :imported-members: diff --git a/docs/api/index.rst b/docs/api/index.rst new file mode 100644 index 000000000..4a16b9f78 --- /dev/null +++ b/docs/api/index.rst @@ -0,0 +1,27 @@ +API Reference +============= + +Pipeline Elements +^^^^^^^^^^^^^^^^^ + +.. toctree:: + :maxdepth: 2 + :caption: Contents: + + datagrabbers + datareaders + preprocessing + markers + storage + + +Utilities +^^^^^^^^^ + +.. toctree:: + :maxdepth: 2 + :caption: Contents: + + api + utils + testing \ No newline at end of file diff --git a/docs/api/markers.rst b/docs/api/markers.rst new file mode 100644 index 000000000..7c6a34b8d --- /dev/null +++ b/docs/api/markers.rst @@ -0,0 +1,7 @@ + +Markers +^^^^^^^ + +.. automodule:: junifer.markers + :members: + :imported-members: diff --git a/docs/api/preprocessing.rst b/docs/api/preprocessing.rst new file mode 100644 index 000000000..571bc7b6e --- /dev/null +++ b/docs/api/preprocessing.rst @@ -0,0 +1,7 @@ + +Pre-processing +^^^^^^^^^^^^^^ + +.. automodule:: junifer.preprocess + :members: + :imported-members: diff --git a/docs/api/storage.rst b/docs/api/storage.rst new file mode 100644 index 000000000..4e785ce60 --- /dev/null +++ b/docs/api/storage.rst @@ -0,0 +1,7 @@ + +Storage +^^^^^^^ + +.. automodule:: junifer.storage + :members: + :imported-members: diff --git a/docs/api/testing.rst b/docs/api/testing.rst new file mode 100644 index 000000000..0220bd917 --- /dev/null +++ b/docs/api/testing.rst @@ -0,0 +1,6 @@ + +Testing +^^^^^^^ + +.. automodule:: junifer.testing.datagrabbers + :members: diff --git a/docs/api/utils.rst b/docs/api/utils.rst new file mode 100644 index 000000000..999a180d3 --- /dev/null +++ b/docs/api/utils.rst @@ -0,0 +1,7 @@ + +Utils +^^^^^ + +.. automodule:: junifer.utils + :members: + :imported-members: diff --git a/docs/index.rst b/docs/index.rst index 7e0aa03ef..34c26f5eb 100644 --- a/docs/index.rst +++ b/docs/index.rst @@ -18,9 +18,9 @@ enabling others to extend it easily. :caption: Contents: installation - data - api + understanding/index.rst auto_examples/index.rst + api/index.rst contribution maintaining whats_new diff --git a/docs/links.inc b/docs/links.inc index 444454691..47d3ad6a4 100644 --- a/docs/links.inc +++ b/docs/links.inc @@ -11,6 +11,7 @@ .. _`AML`: https://www.fz-juelich.de/inm/inm-7/EN/Forschung/Applied%20Machine%20Learning/_node.html .. _`INM-7`: https://www.fz-juelich.de/inm/inm-7/EN/Home/home_node.html +.. _`julearn`: https://juaml.github.io/julearn .. _`pandas`: https://pandas.pydata.org .. _`pandas.DataFrame` : https://pandas.pydata.org/pandas-docs/stable/reference/api/pandas.DataFrame.html diff --git a/docs/data.rst b/docs/understanding/data.rst similarity index 93% rename from docs/data.rst rename to docs/understanding/data.rst index 00e387178..138d29fa8 100644 --- a/docs/data.rst +++ b/docs/understanding/data.rst @@ -1,11 +1,25 @@ -.. include:: links.inc +.. include:: ../links.inc The Data Object =============== -Introduction -^^^^^^^^^^^^ +Description +^^^^^^^^^^^ + +This is the *object* that traverses the steps of the pipeline. It is indeed a +dictionary of dictionaries. The first level of keys are the *data types* and a +special key named 'meta' that contains all the information on the data object +including source and previous transformation steps. + +The second level of keys are the actual data. So far, there are two keys used: +- `path`: path to the file containing the data. +- `data`: the data loaded in memory. + +The *DataGrabber* step will only fill the `path` value. The `data` value will +be filled by the *DataReader* step, if it is one of the possible file types +that the datareader can read. + Data types ^^^^^^^^^^ @@ -30,19 +44,3 @@ Data types - VBM White Matter segmentation (3D) - CAT output (`m0wp2` images) - -The Data Object -^^^^^^^^^^^^^^^ - -This is the *object* that traverses the steps of the pipeline. It is indeed a -dictionary of dictionaries. The first level of keys are the *data types* and a -special key named 'meta' that contains all the information on the data object -including source and previous transformation steps. - -The second level of keys are the actual data. So far, there are two keys used: -- `path`: path to the file containing the data. -- `data`: the data loaded in memory. - -The *DataGrabber* step will only fill the `path` value. The `data` value will -be filled by the *DataReader* step, if it is one of the possible file types -that the datareader can read. diff --git a/docs/understanding/index.rst b/docs/understanding/index.rst new file mode 100644 index 000000000..7a382d639 --- /dev/null +++ b/docs/understanding/index.rst @@ -0,0 +1,25 @@ +.. include:: ../links.inc + +Understanding junifer +===================== + +Before you start, you should understand how junifer works. Junifer is a +tool conceived to extract features from neuroimaging data in a easy-to-use +manner, with minimal coding and minimal user expertise in the internal aspects. + +Unlike other tools like FSL, SPM, AFNI, etc., junifer is not a toolbox to +preprocess data, but a toolbox to extract features from previously-preprocessed +data. + +The main idea is that you have a set of images (e.g. a set of functional fMRI, +structural MRI, diffusion MRI, etc.) and you want to extract features to be +later used in stastical analysis or machine learning (for example using +julearn_). + + + +.. toctree:: + :maxdepth: 2 + :caption: Contents: + + data \ No newline at end of file diff --git a/junifer/datagrabber/datalad_base.py b/junifer/datagrabber/datalad_base.py index 54664a1ff..2c298ada1 100644 --- a/junifer/datagrabber/datalad_base.py +++ b/junifer/datagrabber/datalad_base.py @@ -56,7 +56,6 @@ class DataladDataGrabber(BaseDataGrabber): See Also -------- BaseDataGrabber - BIDSDataladDataGrabber """ diff --git a/junifer/preprocess/confounds.py b/junifer/preprocess/confounds.py index e8f413947..98e3db2ef 100644 --- a/junifer/preprocess/confounds.py +++ b/junifer/preprocess/confounds.py @@ -30,7 +30,7 @@ class BaseConfoundRemover(PipelineStepMixin): Confound removal is based on `nilearn.image.clean_img`. Parameters - ----------- + ---------- strategy : dict, optional The keys of the dictionary should correspond to names of noise components to include: diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py index b08213323..0881de6cb 100644 --- a/junifer/storage/sqlite.py +++ b/junifer/storage/sqlite.py @@ -41,6 +41,8 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): upsert : {"ignore", "update"}, optional Upsert mode. If "ignore" is used, the existing elements are ignored. If "update", the existing elements are updated (default "update"). + **kwargs : dict + The keyword arguments passed to the superclass. See Also -------- @@ -56,12 +58,6 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): **kwargs: str, ) -> None: """Initialize the class. - - Extra Parameters - ---------------- - **kwargs : dict - The keyword arguments passed to the superclass. - """ if upsert not in ["update", "ignore"]: raise_error( @@ -373,7 +369,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): return df def store_metadata(self, meta: Dict) -> str: - """Implement metadata storing in the storage. + r"""Implement metadata storing in the storage. Parameters ---------- @@ -383,7 +379,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): Returns ------- str - The MD5 hash of the metadata prefixed with "meta_" . + The MD5 hash of the metadata prefixed with "meta\_". """ # Copy metadata -- 2.52.0 From 83dda3cb72759ce1359a8ee42b75c16c5d30891c Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Fri, 9 Sep 2022 19:34:40 +0200 Subject: [PATCH 266/287] flake --- junifer/storage/sqlite.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py index 0881de6cb..0942e024b 100644 --- a/junifer/storage/sqlite.py +++ b/junifer/storage/sqlite.py @@ -57,8 +57,7 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): upsert: str = "update", **kwargs: str, ) -> None: - """Initialize the class. - """ + """Initialize the class.""" if upsert not in ["update", "ignore"]: raise_error( msg=( -- 2.52.0 From 1ce35aae1ddcbbfd8f469f02001e4f8a247dcfda Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Sun, 11 Sep 2022 14:08:50 +0200 Subject: [PATCH 267/287] chore: lint examples --- examples/norun_hcpfc_pearson.py | 95 +++++++++++++----------- examples/norun_ukbvm_gmd.py | 43 ++++++----- examples/run_compute_parcel_mean.py | 16 ++-- examples/run_datagrabber_bids_datalad.py | 37 +++++---- examples/run_run_gmd_mean.py | 45 ++++++----- 5 files changed, 134 insertions(+), 102 deletions(-) diff --git a/examples/norun_hcpfc_pearson.py b/examples/norun_hcpfc_pearson.py index e3c98e1ca..d6a460440 100644 --- a/examples/norun_hcpfc_pearson.py +++ b/examples/norun_hcpfc_pearson.py @@ -9,62 +9,73 @@ License: BSD 3 clause from junifer.api import run + datagrabber = { - 'kind': 'HCPOpenAccess', - 'modality': 'fMRI', - 'preprocessed': 'ICA+FIX', - 'space': 'volumetric', + "kind": "HCPOpenAccess", + "modality": "fMRI", + "preprocessed": "ICA+FIX", + "space": "volumetric", } custom_confound_strategy = { - 'filter': 'butterworth', - 'detrend': True, - 'high_pass': 0.01, - 'low_pass': 0.08, - 'standardize': True, - 'confounds': ['csf', 'wm', 'gsr'], - 'derivatives': True, - 'squares': True, - 'other': [] + "filter": "butterworth", + "detrend": True, + "high_pass": 0.01, + "low_pass": 0.08, + "standardize": True, + "confounds": ["csf", "wm", "gsr"], + "derivatives": True, + "squares": True, + "other": [], } markers = [ - {'name': 'Power264_FCPearson', - 'kind': 'FunctionalConnectivity', - 'atlas': 'Power264', - 'method': 'Pearson', - 'confound_strategy': 'Params36'}, - {'name': 'Schaefer400x17_FCPearson', - 'kind': 'FunctionalConnectivity', - 'atlas': 'Schaefer400x17', - 'method': 'Pearson', - 'confound_strategy': 'Params24'}, - {'name': 'Power264_FCSpearman', - 'kind': 'FunctionalConnectivity', - 'atlas': 'Power264', - 'method': 'Spearman', - 'confound_strategy': 'ICAAROMA'}, - {'name': 'Schaefer400x17_FCSpearman', - 'kind': 'FunctionalConnectivity', - 'atlas': 'Schaefer400x17', - 'method': 'Spearman', - 'confound_strategy': 'path/to/predefined/confound_file.tsv'}, - {'name': 'Schaefer400x17_FCSpearman', - 'kind': 'FunctionalConnectivity', - 'atlas': 'Schaefer400x17', - 'method': 'Spearman', - 'confound_strategy': custom_confound_strategy} + { + "name": "Power264_FCPearson", + "kind": "FunctionalConnectivity", + "atlas": "Power264", + "method": "Pearson", + "confound_strategy": "Params36", + }, + { + "name": "Schaefer400x17_FCPearson", + "kind": "FunctionalConnectivity", + "atlas": "Schaefer400x17", + "method": "Pearson", + "confound_strategy": "Params24", + }, + { + "name": "Power264_FCSpearman", + "kind": "FunctionalConnectivity", + "atlas": "Power264", + "method": "Spearman", + "confound_strategy": "ICAAROMA", + }, + { + "name": "Schaefer400x17_FCSpearman", + "kind": "FunctionalConnectivity", + "atlas": "Schaefer400x17", + "method": "Spearman", + "confound_strategy": "path/to/predefined/confound_file.tsv", + }, + { + "name": "Schaefer400x17_FCSpearman", + "kind": "FunctionalConnectivity", + "atlas": "Schaefer400x17", + "method": "Spearman", + "confound_strategy": custom_confound_strategy, + }, ] storage = { - 'kind': 'SQLiteFeatureStorage', - 'uri': '/data/project/juniferexample' + "kind": "SQLiteFeatureStorage", + "uri": "/data/project/juniferexample", } run( - workdir='/tmp', + workdir="/tmp", datagrabber=datagrabber, - elements=[('100408', 'REST1', "LR")], + elements=[("100408", "REST1", "LR")], markers=markers, storage=storage, ) diff --git a/examples/norun_ukbvm_gmd.py b/examples/norun_ukbvm_gmd.py index 2cde565c0..dc4ca5c4a 100644 --- a/examples/norun_ukbvm_gmd.py +++ b/examples/norun_ukbvm_gmd.py @@ -10,27 +10,34 @@ License: BSD 3 clause from junifer.api import run + markers = [ - {'name': 'Schaefer1000x7_TrimMean80', - 'kind': 'ParcelAggregation', - 'atlas': 'Schaefer1000x7', - 'method': 'trim_mean', - 'method_params': {'proportiontocut': 0.2}}, - {'name': 'Schaefer1000x7_Mean', - 'kind': 'ParcelAggregation', - 'atlas': 'Schaefer1000x7', - 'method': 'mean'}, - {'name': 'Schaefer1000x7_Std', - 'kind': 'ParcelAggregation', - 'atlas': 'Schaefer1000x7', - 'method': 'std'} + { + "name": "Schaefer1000x7_TrimMean80", + "kind": "ParcelAggregation", + "atlas": "Schaefer1000x7", + "method": "trim_mean", + "method_params": {"proportiontocut": 0.2}, + }, + { + "name": "Schaefer1000x7_Mean", + "kind": "ParcelAggregation", + "atlas": "Schaefer1000x7", + "method": "mean", + }, + { + "name": "Schaefer1000x7_Std", + "kind": "ParcelAggregation", + "atlas": "Schaefer1000x7", + "method": "std", + }, ] run( - workdir='/tmp', - datagrabber='JuselessUKBVBM', - elements=('sub-1627474', 'ses-2'), + workdir="/tmp", + datagrabber="JuselessUKBVBM", + elements=("sub-1627474", "ses-2"), markers=markers, - storage='SQLDataFrameStorage', - storage_params={'outpath': '/data/project/juniferexample'}, + storage="SQLDataFrameStorage", + storage_params={"outpath": "/data/project/juniferexample"}, ) diff --git a/examples/run_compute_parcel_mean.py b/examples/run_compute_parcel_mean.py index b5c7ddc41..1f76a88f2 100644 --- a/examples/run_compute_parcel_mean.py +++ b/examples/run_compute_parcel_mean.py @@ -12,12 +12,13 @@ License: BSD 3 clause import nilearn -from junifer.utils import configure_logging from junifer.markers.parcel import ParcelAggregation +from junifer.utils import configure_logging + ############################################################################### # Set the logging level to info to see extra information -configure_logging(level='INFO') +configure_logging(level="INFO") ############################################################################### @@ -36,14 +37,11 @@ fmri_img = nilearn.image.concat_imgs(s_func_data.func) ############################################################################### # Define the marker -marker = ParcelAggregation(atlas='Schaefer100x7', method='mean') +marker = ParcelAggregation(atlas="Schaefer100x7", method="mean") ############################################################################### # Prepare the input -input = { - 'BOLD': {'data': fmri_img}, - 'VBM_GM': {'data': vbm_img} -} +input = {"BOLD": {"data": fmri_img}, "VBM_GM": {"data": vbm_img}} ############################################################################### # Fit transform the data @@ -53,7 +51,7 @@ out = marker.fit_transform(input) # Check the results print(out.keys()) -print(out['VBM_GM']['data'].shape) # Shape is (1 x parcels) +print(out["VBM_GM"]["data"].shape) # Shape is (1 x parcels) print(out.keys()) -print(out['BOLD']['data'].shape) # Shape is (timepoints x parcels) +print(out["BOLD"]["data"].shape) # Shape is (timepoints x parcels) diff --git a/examples/run_datagrabber_bids_datalad.py b/examples/run_datagrabber_bids_datalad.py index 1727afed9..eb14892a6 100644 --- a/examples/run_datagrabber_bids_datalad.py +++ b/examples/run_datagrabber_bids_datalad.py @@ -13,34 +13,39 @@ License: BSD 3 clause from junifer.datagrabber import PatternDataladDataGrabber from junifer.utils import configure_logging + ############################################################################### # Set the logging level to info to see extra information -configure_logging(level='INFO') +configure_logging(level="INFO") ############################################################################### # The BIDS datagrabber requires three parameters: the types of data we want, # the specific pattern that matches each type, and the variables that will be # replaced int he patterns. -types = ['T1w', 'bold'] +types = ["T1w", "bold"] patterns = { - 'T1w': '{subject}/anat/{subject}_T1w.nii.gz', - 'bold': '{subject}/func/{subject}_task-rest_bold.nii.gz' + "T1w": "{subject}/anat/{subject}_T1w.nii.gz", + "bold": "{subject}/func/{subject}_task-rest_bold.nii.gz", } -replacements = ['subject'] +replacements = ["subject"] ############################################################################### # Additionally, a datalad datagrabber requires the URI of the remote sibling # and the location of the dataset within the remote sibling. -repo_uri = 'https://gin.g-node.org/juaml/datalad-example-bids' -rootdir = 'example_bids' +repo_uri = "https://gin.g-node.org/juaml/datalad-example-bids" +rootdir = "example_bids" ############################################################################### # Now we can use the datagrabber within a `with` context # One thing we can do with any datagrabber is iterate over the elements. # In this case, each element of the datagrabber is one session. -with PatternDataladDataGrabber(rootdir=rootdir, types=types, - patterns=patterns, uri=repo_uri, - replacements=replacements) as dg: +with PatternDataladDataGrabber( + rootdir=rootdir, + types=types, + patterns=patterns, + uri=repo_uri, + replacements=replacements, +) as dg: for elem in dg: print(elem) @@ -48,8 +53,12 @@ with PatternDataladDataGrabber(rootdir=rootdir, types=types, # Another feature of the datagrabber is the ability to get a specific # element by its name. In this case, we index `sub-01` and we get the file # paths for the two types of data we want (T1w and bold). -with PatternDataladDataGrabber(rootdir=rootdir, types=types, - patterns=patterns, uri=repo_uri, - replacements=replacements) as dg: - sub01 = dg['sub-01'] +with PatternDataladDataGrabber( + rootdir=rootdir, + types=types, + patterns=patterns, + uri=repo_uri, + replacements=replacements, +) as dg: + sub01 = dg["sub-01"] print(sub01) diff --git a/examples/run_run_gmd_mean.py b/examples/run_run_gmd_mean.py index 6259530b7..2b2ee79f8 100644 --- a/examples/run_run_gmd_mean.py +++ b/examples/run_run_gmd_mean.py @@ -8,38 +8,45 @@ License: BSD 3 clause """ import tempfile -from junifer.api import run import junifer.testing.registry # noqa: F401 +from junifer.api import run + datagrabber = { - 'kind': 'OasisVBMTestingDatagrabber', + "kind": "OasisVBMTestingDatagrabber", } markers = [ - {'name': 'Schaefer1000x7_TrimMean80', - 'kind': 'ParcelAggregation', - 'atlas': 'Schaefer1000x7', - 'method': 'trim_mean', - 'method_params': {'proportiontocut': 0.2}}, - {'name': 'Schaefer1000x7_Mean', - 'kind': 'ParcelAggregation', - 'atlas': 'Schaefer1000x7', - 'method': 'mean'}, - {'name': 'Schaefer1000x7_Std', - 'kind': 'ParcelAggregation', - 'atlas': 'Schaefer1000x7', - 'method': 'std'} + { + "name": "Schaefer1000x7_TrimMean80", + "kind": "ParcelAggregation", + "atlas": "Schaefer1000x7", + "method": "trim_mean", + "method_params": {"proportiontocut": 0.2}, + }, + { + "name": "Schaefer1000x7_Mean", + "kind": "ParcelAggregation", + "atlas": "Schaefer1000x7", + "method": "mean", + }, + { + "name": "Schaefer1000x7_Std", + "kind": "ParcelAggregation", + "atlas": "Schaefer1000x7", + "method": "std", + }, ] storage = { - 'kind': 'SQLiteFeatureStorage', + "kind": "SQLiteFeatureStorage", } with tempfile.TemporaryDirectory() as tmpdir: - uri = f'{tmpdir}/test.db' - storage['uri'] = uri + uri = f"{tmpdir}/test.db" + storage["uri"] = uri run( - workdir='/tmp', + workdir="/tmp", datagrabber=datagrabber, markers=markers, storage=storage, -- 2.52.0 From 316259969b9e17830d19c23469f4b878ab4962cf Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Sun, 11 Sep 2022 14:10:37 +0200 Subject: [PATCH 268/287] chore: improve installation.rst --- docs/installation.rst | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/docs/installation.rst b/docs/installation.rst index d257b6042..da70233d8 100644 --- a/docs/installation.rst +++ b/docs/installation.rst @@ -22,11 +22,11 @@ Depending on the installation method, these packages might be installed automati Installation ^^^^^^^^^^^^ -Depending on your use-case, junifer can be installed differently. +Depending on your use-case, junifer can be installed differently: * Install the :ref:`install_latest_release`. This is the most suitable approach for end users. -* Install from :ref:`install_development_git`. This is mostly suitable approach +* Install from :ref:`install_development_git`. This is the most suitable approach for developers. @@ -38,7 +38,7 @@ Either way, we strongly recommend using `virtual environments `_, like so: .. code-block:: bash -- 2.52.0 From 5048af05d8aa63fb5d82f3b7ec9871989a64aa35 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Sun, 11 Sep 2022 14:12:23 +0200 Subject: [PATCH 269/287] chore: improve contribution.rst --- docs/contribution.rst | 35 ++++++++++++++--------------------- 1 file changed, 14 insertions(+), 21 deletions(-) diff --git a/docs/contribution.rst b/docs/contribution.rst index 1e77d21c0..229481fbe 100644 --- a/docs/contribution.rst +++ b/docs/contribution.rst @@ -48,7 +48,7 @@ Setting up the local development environment tox - You can also run all ``tox`` tests in parallel. As of ``tox 3.7``, you can run:: + You can also run all ``tox`` tests in parallel. As of ``tox 3.7``, you can run .. code-block:: console @@ -76,10 +76,10 @@ Before you submit a pull request, check that it meets these guidelines: by codecov which runs automatically when you submit your pull request). 2. If the pull request adds functionality, the docs should be updated. Consider creating a Python file that demonstrates the usage in - ``examples`` directory. + ``examples/`` directory. 3. The pull request should also include a short one-liner of your contribution - in `docs/changes/latest.inc`. If it's your first contribution, also add - yourself to `docs/changes/contributors.inc`. + in ``docs/changes/latest.inc``. If it's your first contribution, also add + yourself to ``docs/changes/contributors.inc``. 4. The pull request will be tested against several Python versions. 5. Someone from the core team will review your work and guide you to a successful contribution. @@ -92,7 +92,7 @@ junifer uses `pytest `_ for its unit-tests and new features should in general always come with new tests that make sure that the code runs as intended. -To run all tests:: +To run all tests .. code-block:: console @@ -102,20 +102,20 @@ To run all tests:: Adding and building documentation --------------------------------- -Building the documentation requires some extra packages and can be installed by:: +Building the documentation requires some extra packages and can be installed by .. code-block:: console pip install -e ".[docs]" -To build the docs:: +To build the docs .. code-block:: bash cd docs make html -To view the documentation, open `docs/_build/html/index.html`. +To view the documentation, open ``docs/_build/html/index.html``. In case you remove some files or change their filenames, you can run into errors when using ``make html``. In this situation you can use ``make clean`` @@ -125,19 +125,15 @@ to clean up the already build files and then re-run ``make html``. Writing Examples ---------------- -Examples are run and displayed in HTML format using `sphinx gallery`_. To add an +The format used for text is reST. Check the `sphinx reST reference`_ for more +details. The examples are run and displayed in HTML format using `sphinx gallery`_. To add an example, just create a ``.py`` file that starts either with ``plot_`` or ``run_``, dependending on whether the example generates a figure or not. -The first lines of the example should be a python block comment with a title, -a description of the example an the following include directive to be able to -use the links. - -The format used for text is reST. Check the `sphinx reST reference`_ for more -details. - -Example of the first lines: +The first lines of the example should be a Python block comment with a title, +a description of the example, authors and license name. +The following is an example of how to start an example .. code-block:: python @@ -158,7 +154,7 @@ a 79 ``#`` (a full line) at the point in which you want to render and add text. Each line of text shall be preceded with ``#``. The code that is not commented will be executed. -The following example will create 3 texts and render the output between the +The following example will create texts and render the output between the texts. .. code-block:: python @@ -181,7 +177,4 @@ texts. Finally, when the example is done, you can run as a normal Python script. To generate the HTML, just build the docs: -.. code-block:: bash - cd docs - make html -- 2.52.0 From eb74c517738eea0a5867da1592455e13c65d564f Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Sun, 11 Sep 2022 14:12:38 +0200 Subject: [PATCH 270/287] chore: replace contribution.rst example with junifer one --- docs/contribution.rst | 73 ++++++++++++++++++++++++++++++++----------- 1 file changed, 55 insertions(+), 18 deletions(-) diff --git a/docs/contribution.rst b/docs/contribution.rst index 229481fbe..0816d5a7a 100644 --- a/docs/contribution.rst +++ b/docs/contribution.rst @@ -138,16 +138,17 @@ The following is an example of how to start an example .. code-block:: python """ - Simple Binary Classification - ============================ + Generic BIDS datagrabber for datalad. + ===================================== - This example uses the 'iris' dataset and performs a simple binary - classification using a Support Vector Machine classifier. + This example uses a generic BIDS datagraber to get the data from a BIDS dataset + store in a datalad remote sibling. - .. include:: ../../links.inc + Authors: Federico Raimondo + + License: BSD 3 clause """ - The rest of the script will be executed as normal Python code. In order to render the output and embed formatted text within the code, you need to add a 79 ``#`` (a full line) at the point in which you want to render and add text. @@ -159,22 +160,58 @@ texts. .. code-block:: python - ############################################################################### - # Imports needed for the example - from seaborn import load_dataset - from julearn import run_cross_validation - from julearn.utils import configure_logging + from junifer.datagrabber import PatternDataladDataGrabber + from junifer.utils import configure_logging + ############################################################################### - df_iris = load_dataset('iris') + # Set the logging level to info to see extra information + configure_logging(level="INFO") + ############################################################################### - # The dataset has three kind of species. We will keep two to perform a binary - # classification. - df_iris = df_iris[df_iris['species'].isin(['versicolor', 'virginica'])] + # The BIDS datagrabber requires three parameters: the types of data we want, + # the specific pattern that matches each type, and the variables that will be + # replaced int he patterns. + types = ["T1w", "bold"] + patterns = { + "T1w": "{subject}/anat/{subject}_T1w.nii.gz", + "bold": "{subject}/func/{subject}_task-rest_bold.nii.gz", + } + replacements = ["subject"] + ############################################################################### + # Additionally, a datalad datagrabber requires the URI of the remote sibling + # and the location of the dataset within the remote sibling. + repo_uri = "https://gin.g-node.org/juaml/datalad-example-bids" + rootdir = "example_bids" + ############################################################################### + # Now we can use the datagrabber within a `with` context + # One thing we can do with any datagrabber is iterate over the elements. + # In this case, each element of the datagrabber is one session. + with PatternDataladDataGrabber( + rootdir=rootdir, + types=types, + patterns=patterns, + uri=repo_uri, + replacements=replacements, + ) as dg: + for elem in dg: + print(elem) -Finally, when the example is done, you can run as a normal Python script. -To generate the HTML, just build the docs: - + ############################################################################### + # Another feature of the datagrabber is the ability to get a specific + # element by its name. In this case, we index `sub-01` and we get the file + # paths for the two types of data we want (T1w and bold). + with PatternDataladDataGrabber( + rootdir=rootdir, + types=types, + patterns=patterns, + uri=repo_uri, + replacements=replacements, + ) as dg: + sub01 = dg["sub-01"] + print(sub01) +Finally, when the example is done, you can run it as a normal Python script. +To generate the HTML, just build the docs. -- 2.52.0 From d81d90f073c89f0e82eb26098520f75b9dcd5b50 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Sun, 11 Sep 2022 14:13:21 +0200 Subject: [PATCH 271/287] chore: make toctree numbered --- docs/index.rst | 1 + 1 file changed, 1 insertion(+) diff --git a/docs/index.rst b/docs/index.rst index 34c26f5eb..805de8e18 100644 --- a/docs/index.rst +++ b/docs/index.rst @@ -14,6 +14,7 @@ and Behaviour (`INM-7`_), it is designed to be as modular as possible thus enabling others to extend it easily. .. toctree:: + :numbered: :maxdepth: 2 :caption: Contents: -- 2.52.0 From 216a5192a67cedf6d32eed7a2c88bbe2f1446658 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Sun, 11 Sep 2022 14:14:55 +0200 Subject: [PATCH 272/287] chore: remove -U from pip install in installation.rst --- docs/installation.rst | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/installation.rst b/docs/installation.rst index da70233d8..90209b8e5 100644 --- a/docs/installation.rst +++ b/docs/installation.rst @@ -42,7 +42,7 @@ Use ``pip`` to install julearn from `PyPI `_, like so: .. code-block:: bash - pip install -U junifer + pip install junifer .. _install_development_git: -- 2.52.0 From c6b546d162f0a0a1d685e406465003512f81026a Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Mon, 12 Sep 2022 09:41:46 +0200 Subject: [PATCH 273/287] Fix instructions for building docs locally --- docs/contribution.rst | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/docs/contribution.rst b/docs/contribution.rst index 0816d5a7a..881b6b73a 100644 --- a/docs/contribution.rst +++ b/docs/contribution.rst @@ -113,13 +113,13 @@ To build the docs .. code-block:: bash cd docs - make html + make local To view the documentation, open ``docs/_build/html/index.html``. In case you remove some files or change their filenames, you can run into -errors when using ``make html``. In this situation you can use ``make clean`` -to clean up the already build files and then re-run ``make html``. +errors when using ``make local``. In this situation you can use ``make clean`` +to clean up the already build files and then re-run ``make local``. Writing Examples -- 2.52.0 From 71717c9aa0ae930a999ea729dd9e993855486bab Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Mon, 12 Sep 2022 13:32:00 +0200 Subject: [PATCH 274/287] Adding more structure to the docs --- docs/builtin.rst | 112 +++++++++++++++++++++++++++++ docs/index.rst | 1 + docs/sphinxext/gh_substitutions.py | 2 +- 3 files changed, 114 insertions(+), 1 deletion(-) create mode 100644 docs/builtin.rst diff --git a/docs/builtin.rst b/docs/builtin.rst new file mode 100644 index 000000000..e5ece4e49 --- /dev/null +++ b/docs/builtin.rst @@ -0,0 +1,112 @@ + +Available pipeline steps +======================== + + +Data Grabbers +^^^^^^^^^^^^^ + +.. + Provide a list of the DataGrabbers that are implemented or planned. + Access: Valid options are + - Open + - Open with registration + - Restricted + + Type/config: this should mention weather the class is built-in in the + core of junifer or needs to be imported from a specific configuration in + the `junifer.configs` module. + + State: this should indicate the state of the dataset. Valid options are + - Planned + - In Progress + - Done + + Version added: If the status is "Done", the Junifer version in which the + dataset was added. Else, a link to the Github issue or pull request + implementing the dataset. Links to github can be added by using the + following syntax: :gh:`` + +.. list-table:: Available data grabbers + :widths: auto + :header-rows: 1 + + * - Class + - Description + - Access + - Type/Config + - State + - Version Added + * - `DataladHCP1200` + - `HCP OpenAccess dataset `_ + - Open with registration + - Built-in + - In Progress + - :gh:`4` + * - `JuselessDataladUKBVBM` + - UKB VBM dataset preprocessed with CAT. Available for Juseless only + - Restricted + - `junifer.configs.juseless` + - Done + - 0.0.1 + + + +Markers +^^^^^^^ + +.. + Provide a list of the Markers that are implemented or planned. + + State: this should indicate the state of the dataset. Valid options are + - Planned + - In Progress + - Done + + Version added: If the status is "Done", the Junifer version in which the + dataset was added. Else, a link to the Github issue or pull request + implementing the dataset. Links to github can be added by using the + following syntax: :gh:`` + +.. list-table:: Available data grabbers + :widths: auto + :header-rows: 1 + + * - Class + - Description + - State + - Version Added + * - :class:`junifer.markers.ParcelAggregation` + - Apply parcellation and perform aggregation function + - Done + - 0.0.1 + + + +Available Atlases and Coordinates +================================= + ++------------------+-----------------------+-----------------------------+---------------+ +| Name | Options | Keys | Version Added | ++==================+=======================+=============================+===============+ +| Schaefer | `n_rois` | `Schaefer100x7` | 0.0.1 | +| | `yeo_networks` | `Schaefer200x7` | | +| | | `Schaefer300x7` | | +| | | `Schaefer400x7` | | +| | | `Schaefer500x7` | | +| | | `Schaefer600x7` | | +| | | `Schaefer700x7` | | +| | | `Schaefer800x7` | | +| | | `Schaefer900x7` | | +| | | `Schaefer1000x7` | | +| | | `Schaefer100x17` | | +| | | `Schaefer200x17` | | +| | | `Schaefer300x17` | | +| | | `Schaefer400x17` | | +| | | `Schaefer500x17` | | +| | | `Schaefer600x17` | | +| | | `Schaefer700x17` | | +| | | `Schaefer800x17` | | +| | | `Schaefer900x17` | | +| | | `Schaefer1000x17` | | ++------------------+-----------------------+-----------------------------+---------------+ diff --git a/docs/index.rst b/docs/index.rst index 805de8e18..a17394b0d 100644 --- a/docs/index.rst +++ b/docs/index.rst @@ -20,6 +20,7 @@ enabling others to extend it easily. installation understanding/index.rst + builtin auto_examples/index.rst api/index.rst contribution diff --git a/docs/sphinxext/gh_substitutions.py b/docs/sphinxext/gh_substitutions.py index 47b5e86f6..e9e651fd9 100644 --- a/docs/sphinxext/gh_substitutions.py +++ b/docs/sphinxext/gh_substitutions.py @@ -19,7 +19,7 @@ def gh_role(name, rawtext, text, lineno, inliner, options={}, content=[]): else: slug = 'issues/' + text text = '#' + text - ref = 'https://github.com/juaml/julearn/' + slug + ref = 'https://github.com/juaml/junifer/' + slug set_classes(options) node = reference(rawtext, text, refuri=ref, **options) return [node], [] -- 2.52.0 From 2d213fb5a2c56c7c892cf9386c11cdfc69306c54 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 12 Sep 2022 13:59:55 +0200 Subject: [PATCH 275/287] docs: add faq.rst --- docs/faq.rst | 4 ++++ docs/index.rst | 1 + 2 files changed, 5 insertions(+) create mode 100644 docs/faq.rst diff --git a/docs/faq.rst b/docs/faq.rst new file mode 100644 index 000000000..7e49b63a5 --- /dev/null +++ b/docs/faq.rst @@ -0,0 +1,4 @@ +.. include:: links.inc + +FAQs +==== diff --git a/docs/index.rst b/docs/index.rst index a17394b0d..cef0f0606 100644 --- a/docs/index.rst +++ b/docs/index.rst @@ -25,6 +25,7 @@ enabling others to extend it easily. api/index.rst contribution maintaining + faq whats_new -- 2.52.0 From 79788609a8385f7f714d9ed4ca2a7e6bc6e5dd09 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 12 Sep 2022 14:39:22 +0200 Subject: [PATCH 276/287] docs: update base branch for PR in contribution.rst --- docs/contribution.rst | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/docs/contribution.rst b/docs/contribution.rst index 881b6b73a..606341056 100644 --- a/docs/contribution.rst +++ b/docs/contribution.rst @@ -20,7 +20,7 @@ Setting up the local development environment pip install -e ".[dev]" -4. Create a branch for local development using the ``dev`` branch as a +4. Create a branch for local development using the ``main`` branch as a starting point. Use ``fix``, ``refactor``, or ``feat`` as a prefix. .. code-block:: console -- 2.52.0 From 8a1134a2b2a88173017b1dc5189e26b5ff6a2054 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 12 Sep 2022 15:45:13 +0200 Subject: [PATCH 277/287] refactor: move MarkerCollection and PipelineStepMixin to junifer/pipeline --- junifer/markers/__init__.py | 2 -- junifer/pipeline/__init__.py | 7 +++++++ junifer/{markers => pipeline}/collection.py | 0 junifer/{markers => pipeline}/pipeline_mixin.py | 0 junifer/{markers => pipeline}/tests/test_collection.py | 0 junifer/{markers => pipeline}/tests/test_pipeline_mixin.py | 0 6 files changed, 7 insertions(+), 2 deletions(-) create mode 100644 junifer/pipeline/__init__.py rename junifer/{markers => pipeline}/collection.py (100%) rename junifer/{markers => pipeline}/pipeline_mixin.py (100%) rename junifer/{markers => pipeline}/tests/test_collection.py (100%) rename junifer/{markers => pipeline}/tests/test_pipeline_mixin.py (100%) diff --git a/junifer/markers/__init__.py b/junifer/markers/__init__.py index 823d27bcd..6747c60ba 100644 --- a/junifer/markers/__init__.py +++ b/junifer/markers/__init__.py @@ -5,6 +5,4 @@ # License: AGPL from .base import BaseMarker -from .collection import MarkerCollection from .parcel import ParcelAggregation -from .pipeline_mixin import PipelineStepMixin diff --git a/junifer/pipeline/__init__.py b/junifer/pipeline/__init__.py new file mode 100644 index 000000000..14a45b679 --- /dev/null +++ b/junifer/pipeline/__init__.py @@ -0,0 +1,7 @@ +"""Provide imports for pipeline sub-package.""" + +# Authors: Synchon Mandal +# License: AGPL + +from .collection import MarkerCollection +from .pipeline_mixin import PipelineStepMixin diff --git a/junifer/markers/collection.py b/junifer/pipeline/collection.py similarity index 100% rename from junifer/markers/collection.py rename to junifer/pipeline/collection.py diff --git a/junifer/markers/pipeline_mixin.py b/junifer/pipeline/pipeline_mixin.py similarity index 100% rename from junifer/markers/pipeline_mixin.py rename to junifer/pipeline/pipeline_mixin.py diff --git a/junifer/markers/tests/test_collection.py b/junifer/pipeline/tests/test_collection.py similarity index 100% rename from junifer/markers/tests/test_collection.py rename to junifer/pipeline/tests/test_collection.py diff --git a/junifer/markers/tests/test_pipeline_mixin.py b/junifer/pipeline/tests/test_pipeline_mixin.py similarity index 100% rename from junifer/markers/tests/test_pipeline_mixin.py rename to junifer/pipeline/tests/test_pipeline_mixin.py -- 2.52.0 From c20eef11b1eb6b48bdcbf3a066b58e67166edff3 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 12 Sep 2022 17:59:46 +0200 Subject: [PATCH 278/287] refactor: revert moving of MarkerCollection --- junifer/markers/__init__.py | 1 + junifer/{pipeline => markers}/collection.py | 13 ++++++------- .../{pipeline => markers}/tests/test_collection.py | 5 ++--- junifer/pipeline/__init__.py | 1 - 4 files changed, 9 insertions(+), 11 deletions(-) rename junifer/{pipeline => markers}/collection.py (93%) rename junifer/{pipeline => markers}/tests/test_collection.py (97%) diff --git a/junifer/markers/__init__.py b/junifer/markers/__init__.py index 6747c60ba..c73e6f8d5 100644 --- a/junifer/markers/__init__.py +++ b/junifer/markers/__init__.py @@ -5,4 +5,5 @@ # License: AGPL from .base import BaseMarker +from .collection import MarkerCollection from .parcel import ParcelAggregation diff --git a/junifer/pipeline/collection.py b/junifer/markers/collection.py similarity index 93% rename from junifer/pipeline/collection.py rename to junifer/markers/collection.py index 98fc57599..1897783f7 100644 --- a/junifer/pipeline/collection.py +++ b/junifer/markers/collection.py @@ -5,14 +5,13 @@ # License: AGPL from collections import Counter -from typing import Dict, Optional, List +from typing import Dict, List, Optional -from junifer.markers.pipeline_mixin import PipelineStepMixin - -from ..datareader import DefaultDataReader -from ..utils import logger -from .base import BaseMarker +from ..datareader.default import DefaultDataReader +from ..markers.base import BaseMarker +from ..pipeline import PipelineStepMixin from ..storage.base import BaseFeatureStorage +from ..utils import logger class MarkerCollection: @@ -32,7 +31,7 @@ class MarkerCollection: markers: List[BaseMarker], datareader: Optional[PipelineStepMixin] = None, preprocessing: Optional[PipelineStepMixin] = None, - storage: Optional[BaseFeatureStorage] = None + storage: Optional[BaseFeatureStorage] = None, ): """Initialize the class.""" # Check that the markers have different names diff --git a/junifer/pipeline/tests/test_collection.py b/junifer/markers/tests/test_collection.py similarity index 97% rename from junifer/pipeline/tests/test_collection.py rename to junifer/markers/tests/test_collection.py index 224bc6f04..8ef21777d 100644 --- a/junifer/pipeline/tests/test_collection.py +++ b/junifer/markers/tests/test_collection.py @@ -9,7 +9,7 @@ from numpy.testing import assert_array_equal from junifer.datareader.default import DefaultDataReader from junifer.markers import MarkerCollection, ParcelAggregation -from junifer.markers.base import PipelineStepMixin +from junifer.pipeline import PipelineStepMixin from junifer.storage import SQLiteFeatureStorage from junifer.testing.datagrabbers import OasisVBMTestingDatagrabber @@ -90,8 +90,7 @@ def test_marker_collection(): for t_marker in markers: t_name = t_marker.name assert_array_equal( - out[t_name]["VBM_GM"]["data"], - out2[t_name]["VBM_GM"]["data"] + out[t_name]["VBM_GM"]["data"], out2[t_name]["VBM_GM"]["data"] ) diff --git a/junifer/pipeline/__init__.py b/junifer/pipeline/__init__.py index 14a45b679..6d57d7b78 100644 --- a/junifer/pipeline/__init__.py +++ b/junifer/pipeline/__init__.py @@ -3,5 +3,4 @@ # Authors: Synchon Mandal # License: AGPL -from .collection import MarkerCollection from .pipeline_mixin import PipelineStepMixin -- 2.52.0 From bee5b151d842750037b7bf9bbb9bd77a7501e3db Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Mon, 12 Sep 2022 18:04:23 +0200 Subject: [PATCH 279/287] Update latest.inc comments --- docs/changes/latest.inc | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/docs/changes/latest.inc b/docs/changes/latest.inc index 7c826721b..2a6627ca6 100644 --- a/docs/changes/latest.inc +++ b/docs/changes/latest.inc @@ -8,6 +8,10 @@ - "Bugs" for bug fixes - "API changes" for backward-incompatible changes +.. NOTE: add the contributors and reference to the github issue/PR at the end + Example: + - Implemented feature X (:gh:`151` by `Sami Hamdan`_). + .. _current: Current (0.0.0.dev) -- 2.52.0 From 83d9c06bdb99cdae8d99804364902f9a8720b738 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 12 Sep 2022 18:01:08 +0200 Subject: [PATCH 280/287] fix: import fixes --- junifer/api/registry.py | 4 +++- junifer/datareader/default.py | 2 +- junifer/markers/base.py | 2 +- junifer/pipeline/tests/test_pipeline_mixin.py | 2 +- junifer/preprocess/confounds.py | 2 +- 5 files changed, 7 insertions(+), 5 deletions(-) diff --git a/junifer/api/registry.py b/junifer/api/registry.py index 1d567a82a..c1398d211 100644 --- a/junifer/api/registry.py +++ b/junifer/api/registry.py @@ -8,9 +8,11 @@ from typing import TYPE_CHECKING, Dict, List, Optional, Union from ..utils.logging import logger, raise_error + + if TYPE_CHECKING: from ..datagrabber.base import BaseDataGrabber - from ..markers.base import PipelineStepMixin + from ..pipeline import PipelineStepMixin from ..storage.base import BaseFeatureStorage # Define valid steps for operation diff --git a/junifer/datareader/default.py b/junifer/datareader/default.py index 7bddd2cc1..dd3877d2f 100644 --- a/junifer/datareader/default.py +++ b/junifer/datareader/default.py @@ -10,7 +10,7 @@ from typing import Dict, List import nibabel as nib import pandas as pd -from ..markers.base import PipelineStepMixin +from ..pipeline.pipeline_mixin import PipelineStepMixin from ..utils.logging import logger, warn_with_log diff --git a/junifer/markers/base.py b/junifer/markers/base.py index d232f1405..7530fcfee 100644 --- a/junifer/markers/base.py +++ b/junifer/markers/base.py @@ -6,8 +6,8 @@ from typing import Dict, List, Optional, Union +from ..pipeline.pipeline_mixin import PipelineStepMixin from ..utils import logger, raise_error -from .pipeline_mixin import PipelineStepMixin class BaseMarker(PipelineStepMixin): diff --git a/junifer/pipeline/tests/test_pipeline_mixin.py b/junifer/pipeline/tests/test_pipeline_mixin.py index 7086a7805..afae53362 100644 --- a/junifer/pipeline/tests/test_pipeline_mixin.py +++ b/junifer/pipeline/tests/test_pipeline_mixin.py @@ -6,7 +6,7 @@ import pytest -from junifer.markers.pipeline_mixin import PipelineStepMixin +from junifer.pipeline.pipeline_mixin import PipelineStepMixin def test_PipelineStepMixin() -> None: diff --git a/junifer/preprocess/confounds.py b/junifer/preprocess/confounds.py index 98e3db2ef..bf7fed2e2 100644 --- a/junifer/preprocess/confounds.py +++ b/junifer/preprocess/confounds.py @@ -13,7 +13,7 @@ from nilearn._utils.niimg_conversions import check_niimg_4d from nilearn.image import clean_img from nilearn.masking import compute_brain_mask -from ..markers import PipelineStepMixin +from ..pipeline import PipelineStepMixin from ..utils import logger, raise_error -- 2.52.0 From 0b061a200351944dcc58da446d87cbf1533235ef Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 12 Sep 2022 18:05:24 +0200 Subject: [PATCH 281/287] chore: lint --- junifer/api/cli.py | 27 +++++++++++++++----- junifer/api/functions.py | 4 +-- junifer/api/tests/test_registry.py | 3 +-- junifer/data/atlases.py | 11 ++++---- junifer/data/tests/test_atlases.py | 3 ++- junifer/datagrabber/datalad_base.py | 3 ++- junifer/markers/base.py | 4 +-- junifer/markers/parcel.py | 3 ++- junifer/preprocess/confounds.py | 2 +- junifer/stats.py | 6 ++--- junifer/storage/sqlite.py | 13 ++++++---- junifer/storage/tests/test_sqlite.py | 38 +++++++++++++--------------- junifer/storage/utils.py | 4 +-- junifer/utils/logging.py | 4 +-- 14 files changed, 68 insertions(+), 57 deletions(-) diff --git a/junifer/api/cli.py b/junifer/api/cli.py index 6bcba4392..2402de4b9 100644 --- a/junifer/api/cli.py +++ b/junifer/api/cli.py @@ -4,10 +4,10 @@ # Synchon Mandal # License: AGPL +import pathlib from typing import Dict, List, Union import click -import pathlib from ..utils.logging import configure_logging, logger, warn_with_log from .functions import collect as api_collect @@ -59,8 +59,10 @@ def cli() -> None: # pragma: no cover @cli.command() @click.argument( - "filepath", type=click.Path( - exists=True, readable=True, dir_okay=False, path_type=pathlib.Path) + "filepath", + type=click.Path( + exists=True, readable=True, dir_okay=False, path_type=pathlib.Path + ), ) @click.option("--element", type=str, multiple=True) @click.option( @@ -102,8 +104,11 @@ def run(filepath: click.Path, element: str, verbose: click.Choice) -> None: @cli.command() @click.argument( - "filepath", type=click.Path( - exists=True, readable=True, dir_okay=False, path_type=pathlib.Path)) + "filepath", + type=click.Path( + exists=True, readable=True, dir_okay=False, path_type=pathlib.Path + ), +) @click.option( "-v", "--verbose", @@ -131,8 +136,10 @@ def collect(filepath: click.Path, verbose: click.Choice) -> None: @cli.command() @click.argument( - "filepath", type=click.Path( - exists=True, readable=True, dir_okay=False, path_type=pathlib.Path) + "filepath", + type=click.Path( + exists=True, readable=True, dir_okay=False, path_type=pathlib.Path + ), ) @click.option("--element", type=str, multiple=True) @click.option("--overwrite", is_flag=True) @@ -180,3 +187,9 @@ def queue( submit=submit, **queue_config, ) + + +@cli.command() +def selftest() -> None: + """Selftest command for CLI.""" + pass diff --git a/junifer/api/functions.py b/junifer/api/functions.py index 3c4bc7e4c..f8ba24aa8 100644 --- a/junifer/api/functions.py +++ b/junifer/api/functions.py @@ -7,9 +7,9 @@ import shutil import subprocess -from pathlib import Path -from typing import Dict, List, Optional, Union, Tuple import typing +from pathlib import Path +from typing import Dict, List, Optional, Tuple, Union import yaml diff --git a/junifer/api/tests/test_registry.py b/junifer/api/tests/test_registry.py index 7ac9a88ca..7a45c0d5f 100644 --- a/junifer/api/tests/test_registry.py +++ b/junifer/api/tests/test_registry.py @@ -4,10 +4,9 @@ # Leonard Sasse # Synchon Mandal # License: AGPL -from typing import Type - import logging from abc import ABC +from typing import Type import pytest diff --git a/junifer/data/atlases.py b/junifer/data/atlases.py index d2aeeb0f6..9b9bfc3b3 100644 --- a/junifer/data/atlases.py +++ b/junifer/data/atlases.py @@ -10,7 +10,7 @@ import shutil import tempfile import zipfile from pathlib import Path -from typing import TYPE_CHECKING, List, Optional, Tuple, Union, Any, Dict +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union import nibabel as nib import numpy as np @@ -536,7 +536,8 @@ def _retrieve_tian( if space != "MNI6thgeneration": raise_error( f"The parameter `space` ({space}) for 7T needs to be " - f"MNI6thgeneration") + f"MNI6thgeneration" + ) else: raise_error( f"The parameter `magneticfield` ({magneticfield}) needs to be " @@ -592,7 +593,7 @@ def _retrieve_tian( "A simple numbering scheme for distinction was therefore used." ) else: # pragma: no cover - raise_error('This should not happen. Please report this error.') + raise_error("This should not happen. Please report this error.") # check existence of atlas if not (atlas_fname.exists() and atlas_lname.exists()): @@ -625,9 +626,7 @@ def _retrieve_tian( def _retrieve_suit( - atlas_dir: Path, - resolution: Optional[float], - space: str = "MNI" + atlas_dir: Path, resolution: Optional[float], space: str = "MNI" ) -> Tuple[Path, List[str]]: """Retrieve SUIT atlas. diff --git a/junifer/data/tests/test_atlases.py b/junifer/data/tests/test_atlases.py index df70002fb..fb6190f88 100644 --- a/junifer/data/tests/test_atlases.py +++ b/junifer/data/tests/test_atlases.py @@ -443,7 +443,8 @@ def test_tian_7T_6thgeneration( assert fname.name == fname1 assert len(lbl) == n_label assert_array_almost_equal( - img.header["pixdim"][1:4], [1.6, 1.6, 1.6]) # type: ignore + img.header["pixdim"][1:4], [1.6, 1.6, 1.6] + ) # type: ignore def test_retrieve_tian_incorrect_space(tmp_path: Path) -> None: diff --git a/junifer/datagrabber/datalad_base.py b/junifer/datagrabber/datalad_base.py index 2c298ada1..a0b004323 100644 --- a/junifer/datagrabber/datalad_base.py +++ b/junifer/datagrabber/datalad_base.py @@ -120,7 +120,8 @@ class DataladDataGrabber(BaseDataGrabber): """Install the datalad dataset into the datadir.""" logger.debug(f"Installing dataset {self.uri} to {self._datadir}") self._dataset = dl.install( # type: ignore because of datalad - self._datadir, source=self.uri) + self._datadir, source=self.uri + ) logger.debug("Dataset installed") def remove(self): diff --git a/junifer/markers/base.py b/junifer/markers/base.py index 7530fcfee..d08348868 100644 --- a/junifer/markers/base.py +++ b/junifer/markers/base.py @@ -23,9 +23,7 @@ class BaseMarker(PipelineStepMixin): """ def __init__( - self, - on: Union[List[str], str], - name: Optional[str] = None + self, on: Union[List[str], str], name: Optional[str] = None ) -> None: """Initialize the class.""" if not isinstance(on, list): diff --git a/junifer/markers/parcel.py b/junifer/markers/parcel.py index 755c12f15..386ad8dbb 100644 --- a/junifer/markers/parcel.py +++ b/junifer/markers/parcel.py @@ -116,7 +116,8 @@ class ParcelAggregation(BaseMarker): ) logger.debug("Masking") masker = NiftiMasker( - atlas_bin, target_affine=t_input.affine) # type: ignore + atlas_bin, target_affine=t_input.affine + ) # type: ignore # Mask the input data and the atlas data = masker.fit_transform(t_input) diff --git a/junifer/preprocess/confounds.py b/junifer/preprocess/confounds.py index bf7fed2e2..f61f29644 100644 --- a/junifer/preprocess/confounds.py +++ b/junifer/preprocess/confounds.py @@ -18,7 +18,7 @@ from ..utils import logger, raise_error if TYPE_CHECKING: - from nibabel import Nifti1Image, Nifti2Image, MGHImage + from nibabel import MGHImage, Nifti1Image, Nifti2Image class BaseConfoundRemover(PipelineStepMixin): diff --git a/junifer/stats.py b/junifer/stats.py index aa55e7475..f5301d759 100644 --- a/junifer/stats.py +++ b/junifer/stats.py @@ -5,7 +5,7 @@ # License: AGPL from functools import partial -from typing import Callable, Dict, Any, Optional +from typing import Any, Callable, Dict, Optional import numpy as np from scipy.stats import trim_mean @@ -15,8 +15,8 @@ from .utils import logger, raise_error def get_aggfunc_by_name( - name: str, - func_params: Optional[Dict[str, Any]]) -> Callable: + name: str, func_params: Optional[Dict[str, Any]] +) -> Callable: """Get an aggregation function by its name. Parameters diff --git a/junifer/storage/sqlite.py b/junifer/storage/sqlite.py index 0942e024b..9522713b0 100644 --- a/junifer/storage/sqlite.py +++ b/junifer/storage/sqlite.py @@ -108,8 +108,9 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): ) prefix = element_to_prefix(element) # Format URI for engine creation - uri = "sqlite:///" \ - f"{self.uri.parent}/{prefix}{self.uri.name}" # type: ignore + uri = ( + "sqlite:///" f"{self.uri.parent}/{prefix}{self.uri.name}" + ) # type: ignore return create_engine(uri, echo=False) def _save_upsert( @@ -231,7 +232,8 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): ) # Prepare new dataframe data_df = pd.DataFrame( - data, columns=columns, index=idx) # type: ignore + data, columns=columns, index=idx + ) # type: ignore # Store dataframe self.store_df(df=data_df, meta=meta) @@ -525,8 +527,9 @@ class SQLiteFeatureStorage(PandasBaseFeatureStorage): """ if self.single_output is True: raise_error(msg="collect() is not implemented for single output.") - logger.info("Collecting data from " - f"{self.uri.parent}/*{self.uri.name}") # type: ignore + logger.info( + "Collecting data from " f"{self.uri.parent}/*{self.uri.name}" + ) # type: ignore # Create new instance out_storage = SQLiteFeatureStorage( uri=self.uri, single_output=True, upsert="ignore" diff --git a/junifer/storage/tests/test_sqlite.py b/junifer/storage/tests/test_sqlite.py index 542b27632..af5c4a3e0 100644 --- a/junifer/storage/tests/test_sqlite.py +++ b/junifer/storage/tests/test_sqlite.py @@ -4,9 +4,8 @@ # Synchon Mandal # License: AGPL -from typing import Union, List - from pathlib import Path +from typing import List, Union import numpy as np import pandas as pd @@ -59,8 +58,9 @@ df_ignore = pd.DataFrame( ).set_index(["element", "pk2"]) -def _read_sql(table_name: str, uri: str, - index_col: Union[str, List[str]]) -> pd.DataFrame: +def _read_sql( + table_name: str, uri: str, index_col: Union[str, List[str]] +) -> pd.DataFrame: """Read database table into a pandas DataFrame. Parameters @@ -160,8 +160,7 @@ def test_upsert_replace(tmp_path: Path) -> None: table_name = storage.store_metadata(meta) # Read stored table c_df1 = _read_sql( - table_name=table_name, uri=uri.as_posix(), - index_col=["element", "pk2"] + table_name=table_name, uri=uri.as_posix(), index_col=["element", "pk2"] ) # Check if dataframes are equal assert_frame_equal(df1, c_df1) @@ -169,8 +168,7 @@ def test_upsert_replace(tmp_path: Path) -> None: storage._save_upsert(df=df2, name=table_name, if_exists="replace") # Read stored table c_df2 = _read_sql( - table_name=table_name, uri=uri.as_posix(), - index_col=["element", "pk2"] + table_name=table_name, uri=uri.as_posix(), index_col=["element", "pk2"] ) # Check if dataframes are equal assert_frame_equal(df2, c_df2) @@ -198,9 +196,7 @@ def test_upsert_ignore(tmp_path: Path) -> None: table_name = storage.store_metadata(meta=meta) # Read stored table c_df1 = _read_sql( - table_name=table_name, - uri=uri.as_posix(), - index_col=["element", "pk2"] + table_name=table_name, uri=uri.as_posix(), index_col=["element", "pk2"] ) # Check if dataframes are equal assert_frame_equal(df1, c_df1) @@ -209,7 +205,8 @@ def test_upsert_ignore(tmp_path: Path) -> None: storage.store_df(df2, meta) # Read stored table c_dfignore = _read_sql( - table_name, uri=uri.as_posix(), index_col=["element", "pk2"]) + table_name, uri=uri.as_posix(), index_col=["element", "pk2"] + ) # Check if dataframes are equal assert_frame_equal(c_dfignore, df_ignore) # Check for error @@ -236,8 +233,7 @@ def test_upsert_update(tmp_path: Path) -> None: table_name = storage.store_metadata(meta) # Read stored table c_df1 = _read_sql( - table_name=table_name, uri=uri.as_posix(), - index_col=["element", "pk2"] + table_name=table_name, uri=uri.as_posix(), index_col=["element", "pk2"] ) # Check if dataframes are equal assert_frame_equal(df1, c_df1) @@ -245,8 +241,8 @@ def test_upsert_update(tmp_path: Path) -> None: storage.store_df(df2, meta) # Read stored table c_dfupdate = _read_sql( - table_name, uri=uri.as_posix(), - index_col=["element", "pk2"]) + table_name, uri=uri.as_posix(), index_col=["element", "pk2"] + ) # Check if dataframes are equal assert_frame_equal(c_dfupdate, df_update) @@ -381,8 +377,9 @@ def test_store_table(tmp_path: Path) -> None: table_name = storage.store_metadata(meta) # Read stored table c_df = _read_sql( - table_name=table_name, uri=uri.as_posix(), - index_col=["element", "scan"] + table_name=table_name, + uri=uri.as_posix(), + index_col=["element", "scan"], ) # Check if dataframes are equal assert_frame_equal(df, c_df) @@ -400,8 +397,9 @@ def test_store_table(tmp_path: Path) -> None: ) # Read stored table c_df_new = _read_sql( - table_name=table_name, uri=uri.as_posix(), - index_col=["element", "scan"] + table_name=table_name, + uri=uri.as_posix(), + index_col=["element", "scan"], ) # Check if dataframes are equal assert_frame_equal(df_new, c_df_new) diff --git a/junifer/storage/utils.py b/junifer/storage/utils.py index 1d7585787..0a6fa58df 100644 --- a/junifer/storage/utils.py +++ b/junifer/storage/utils.py @@ -4,11 +4,9 @@ # Synchon Mandal # License: AGPL -from typing import Any - import hashlib import json -from typing import Dict, Optional, Tuple, Union +from typing import Any, Dict, Optional, Tuple, Union import numpy as np import pandas as pd diff --git a/junifer/utils/logging.py b/junifer/utils/logging.py index bcb79ef3e..169a74276 100644 --- a/junifer/utils/logging.py +++ b/junifer/utils/logging.py @@ -273,8 +273,8 @@ def raise_error(msg: str, klass: Type[Exception] = ValueError) -> NoReturn: def warn_with_log( - msg: str, - category: Optional[Type[Warning]] = RuntimeWarning) -> None: + msg: str, category: Optional[Type[Warning]] = RuntimeWarning +) -> None: """Warn, but first log it. Parameters -- 2.52.0 From 19bd73a2c32e843411cab672323da72fa5085df6 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 12 Sep 2022 18:32:23 +0200 Subject: [PATCH 282/287] chore: update project structure in README --- README.md | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/README.md b/README.md index 4c4773c31..31ed94d14 100644 --- a/README.md +++ b/README.md @@ -18,15 +18,17 @@ The documentation is available at [https://juaml.github.io/junifer](https://juam * `docs`: Documentation, built using sphinx. * `examples`: Examples, using sphinx-gallery. File names of examples that create visual output must start with `plot_`, otherwise, with `run_`. -* `junifer`: Main library directory - * `api`: User API module - * `data`: Module that handles data required for the library to work (e.g. atlases) +* `junifer`: Main library directory. + * `api`: User API module. + * `configs`: Module for pre-defined configs for most used computing clusters. + * `data`: Module that handles data required for the library to work (e.g. atlases). * `datagrabber`: DataGrabber module. * `datareader`: DataReader module. * `markers`: Markers module. * `pipeline`: Pipeline module. * `preprocess`: Preprocessing module. * `storage`: Storage module. + * `testing`: Testing components module. * `utils`: Utilities module (e.g. logging) -- 2.52.0 From fedab6239c8def4f1a493df0a36b65cfdd85eefc Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 12 Sep 2022 18:47:13 +0200 Subject: [PATCH 283/287] chore: add missing imports in junifer/__init__.py --- junifer/__init__.py | 11 ++++++++--- 1 file changed, 8 insertions(+), 3 deletions(-) diff --git a/junifer/__init__.py b/junifer/__init__.py index e8b301aad..5d46af23a 100644 --- a/junifer/__init__.py +++ b/junifer/__init__.py @@ -6,7 +6,12 @@ from ._version import __version__ from . import api -from . import utils -from . import datagrabber -from . import markers from . import configs +from . import data +from . import datagrabber +from . import datareader +from . import markers +from . import pipeline +from . import preprocess +from . import storage +from . import utils -- 2.52.0 From 53b3c81226dd2f173413ed9f26bd1aa8b8be61bf Mon Sep 17 00:00:00 2001 From: Fede Raimondo Date: Mon, 12 Sep 2022 18:49:31 +0200 Subject: [PATCH 284/287] Add dev env yaml file for conda --- conda-env.yml | 32 ++++++++++++++++++++++++++++++++ 1 file changed, 32 insertions(+) create mode 100644 conda-env.yml diff --git a/conda-env.yml b/conda-env.yml new file mode 100644 index 000000000..2801c724a --- /dev/null +++ b/conda-env.yml @@ -0,0 +1,32 @@ +name: junifer-dev +channels: + - conda-forge + - defaults +dependencies: + - python=3.10 + - click>=8.1.3,<8.2 + - numpy>=1.22,<1.23 + - datalad>=0.15.4,<0.18 + - pandas>=1.4.0,<1.5 + - nibabel>=3.2.0,<4.1 + - nilearn>=0.9.0,<1.0 + - sqlalchemy>=1.4.27,<= 1.5.0 + - pyyaml>=5.1.2,<7.0 + - seaborn>=0.11.2,<0.12 + - Sphinx>=5.0.2,<5.1 + - sphinx-gallery>=0.10.1,<0.11 + - numpydoc>=1.4.0,<1.5 + - tox + - ipykernel + - isort + - pytest-cov + - pytest + - black + - flake8 + - flake8-docstrings + - flake8-bugbear + - codespell + - pip + - pip: + - sphinx-rtd-theme>=1.0.0,<1.1 + - sphinx-multiversion>=0.2.4,<0.3 -- 2.52.0 From e789350c62c80efa9a5bb6868e395f700be8aed7 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 12 Sep 2022 19:49:45 +0200 Subject: [PATCH 285/287] docs: add new files for understanding/ --- docs/understanding/datagrabber.rst | 6 ++++++ docs/understanding/datareader.rst | 6 ++++++ docs/understanding/marker.rst | 4 ++++ docs/understanding/storage.rst | 4 ++++ 4 files changed, 20 insertions(+) create mode 100644 docs/understanding/datagrabber.rst create mode 100644 docs/understanding/datareader.rst create mode 100644 docs/understanding/marker.rst create mode 100644 docs/understanding/storage.rst diff --git a/docs/understanding/datagrabber.rst b/docs/understanding/datagrabber.rst new file mode 100644 index 000000000..1a23ed3fa --- /dev/null +++ b/docs/understanding/datagrabber.rst @@ -0,0 +1,6 @@ +.. include:: ../links.inc + +.. _datagrabber: + +Data Grabber +============ diff --git a/docs/understanding/datareader.rst b/docs/understanding/datareader.rst new file mode 100644 index 000000000..8a3521efd --- /dev/null +++ b/docs/understanding/datareader.rst @@ -0,0 +1,6 @@ +.. include:: ../links.inc + +.. _datareader: + +Data Reader +=========== diff --git a/docs/understanding/marker.rst b/docs/understanding/marker.rst new file mode 100644 index 000000000..3bf018ba7 --- /dev/null +++ b/docs/understanding/marker.rst @@ -0,0 +1,4 @@ +.. include:: ../links.inc + +Marker +====== diff --git a/docs/understanding/storage.rst b/docs/understanding/storage.rst new file mode 100644 index 000000000..ef296e1d4 --- /dev/null +++ b/docs/understanding/storage.rst @@ -0,0 +1,4 @@ +.. include:: ../links.inc + +Storage +======= -- 2.52.0 From 4e7c7711c8d61eb2ab537886b8924ef6a80be7a7 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 12 Sep 2022 19:50:38 +0200 Subject: [PATCH 286/287] docs: improve understanding/index.rst --- docs/understanding/index.rst | 15 +++++++++------ 1 file changed, 9 insertions(+), 6 deletions(-) diff --git a/docs/understanding/index.rst b/docs/understanding/index.rst index 7a382d639..e24c92883 100644 --- a/docs/understanding/index.rst +++ b/docs/understanding/index.rst @@ -8,18 +8,21 @@ tool conceived to extract features from neuroimaging data in a easy-to-use manner, with minimal coding and minimal user expertise in the internal aspects. Unlike other tools like FSL, SPM, AFNI, etc., junifer is not a toolbox to -preprocess data, but a toolbox to extract features from previously-preprocessed +pre-process data, but a toolbox to extract features from previously pre-processed data. -The main idea is that you have a set of images (e.g. a set of functional fMRI, -structural MRI, diffusion MRI, etc.) and you want to extract features to be -later used in stastical analysis or machine learning (for example using +The main idea is that you have a set of images (e.g. a set of functional MRI, +structural MRI, diffusion MRI, etc.) and you want to extract features to +later use in stastical analyses or machine learning (for example, using julearn_). - .. toctree:: :maxdepth: 2 :caption: Contents: - data \ No newline at end of file + data + datagrabber + datareader + marker + storage -- 2.52.0 From 1543215157c9aca9e02ea1d2ea284cf7650fcb71 Mon Sep 17 00:00:00 2001 From: Synchon Mandal Date: Mon, 12 Sep 2022 19:51:01 +0200 Subject: [PATCH 287/287] docs: improve understanding/data.rst --- docs/understanding/data.rst | 28 ++++++++++++++-------------- 1 file changed, 14 insertions(+), 14 deletions(-) diff --git a/docs/understanding/data.rst b/docs/understanding/data.rst index 138d29fa8..1010720a8 100644 --- a/docs/understanding/data.rst +++ b/docs/understanding/data.rst @@ -1,4 +1,3 @@ - .. include:: ../links.inc The Data Object @@ -8,18 +7,20 @@ Description ^^^^^^^^^^^ This is the *object* that traverses the steps of the pipeline. It is indeed a -dictionary of dictionaries. The first level of keys are the *data types* and a -special key named 'meta' that contains all the information on the data object -including source and previous transformation steps. +dictionary of dictionaries. The first level of keys are the :ref:`data_types` +and a special key named ``meta`` that contains all the information on the data +object including source and previous transformation steps. The second level of keys are the actual data. So far, there are two keys used: -- `path`: path to the file containing the data. -- `data`: the data loaded in memory. -The *DataGrabber* step will only fill the `path` value. The `data` value will -be filled by the *DataReader* step, if it is one of the possible file types +- ``path``: path to the file containing the data. +- ``data``: the data loaded in memory. + +The :ref:`datagrabber` step will only fill the ``path`` value. +The ``data`` value will be filled by the :ref:`datareader` step, if it is one of the possible file types that the datareader can read. +.. _data_types: Data types ^^^^^^^^^^ @@ -30,17 +31,16 @@ Data types * - Name - Description - - Example - * - `T1w` + - Example + * - ``T1w`` - T1w image (3D) - Preprocessed or Raw T1w image - * - `BOLD` + * - ``BOLD`` - BOLD image (4D) - Preprocessed/Denoised BOLD image (fmriprep output) - * - `VBM_GM` + * - ``VBM_GM`` - VBM Gray Matter segmentation (3D) - CAT output (`m0wp1` images) - * - `VBM_WM` + * - ``VBM_WM`` - VBM White Matter segmentation (3D) - CAT output (`m0wp2` images) - -- 2.52.0