[ENH]: Adapt MultipleDataGrabber to patterns with nested data types #351
22 changed files with 877 additions and 563 deletions
1
docs/changes/newsfragments/351.change
Normal file
1
docs/changes/newsfragments/351.change
Normal file
|
|
@ -0,0 +1 @@
|
||||||
|
Add ``partial_pattern_ok`` argument to :class:`.PatternDataGrabber` to not raise error on missing mandatory key checks for data types by `Synchon Mandal`_
|
||||||
1
docs/changes/newsfragments/351.enh
Normal file
1
docs/changes/newsfragments/351.enh
Normal file
|
|
@ -0,0 +1 @@
|
||||||
|
Adapt :class:`.MultipleDataGrabber` to handle "nested types" introduced in :gh:`341` by `Synchon Mandal`_
|
||||||
1
docs/changes/newsfragments/351.feature
Normal file
1
docs/changes/newsfragments/351.feature
Normal file
|
|
@ -0,0 +1 @@
|
||||||
|
Introduce :class:`.PatternValidationMixin` to simplify validation for pattern-based DataGrabbers and :func:`.deep_update` for updating dictionary with varying width and depth by `Synchon Mandal`_
|
||||||
1
docs/changes/newsfragments/351.misc
Normal file
1
docs/changes/newsfragments/351.misc
Normal file
|
|
@ -0,0 +1 @@
|
||||||
|
Integrate ``warnings`` with ``logging`` respecting filters by `Fede Raimondo`_
|
||||||
|
|
@ -17,6 +17,7 @@ from .hcp1200 import HCP1200, DataladHCP1200
|
||||||
from .multiple import MultipleDataGrabber
|
from .multiple import MultipleDataGrabber
|
||||||
from .dmcc13_benchmark import DMCC13Benchmark
|
from .dmcc13_benchmark import DMCC13Benchmark
|
||||||
|
|
||||||
|
from .pattern_validation_mixin import PatternValidationMixin
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"BaseDataGrabber",
|
"BaseDataGrabber",
|
||||||
|
|
@ -30,4 +31,5 @@ __all__ = [
|
||||||
"DataladHCP1200",
|
"DataladHCP1200",
|
||||||
"MultipleDataGrabber",
|
"MultipleDataGrabber",
|
||||||
"DMCC13Benchmark",
|
"DMCC13Benchmark",
|
||||||
|
"PatternValidationMixin",
|
||||||
]
|
]
|
||||||
|
|
|
||||||
|
|
@ -11,7 +11,6 @@ from typing import Dict, Iterator, List, Tuple, Union
|
||||||
|
|
||||||
from ..pipeline import UpdateMetaMixin
|
from ..pipeline import UpdateMetaMixin
|
||||||
from ..utils import logger, raise_error
|
from ..utils import logger, raise_error
|
||||||
from .utils import validate_types
|
|
||||||
|
|
||||||
|
|
||||||
__all__ = ["BaseDataGrabber"]
|
__all__ = ["BaseDataGrabber"]
|
||||||
|
|
@ -30,16 +29,21 @@ class BaseDataGrabber(ABC, UpdateMetaMixin):
|
||||||
datadir : str or pathlib.Path
|
datadir : str or pathlib.Path
|
||||||
The directory where the data is / will be stored.
|
The directory where the data is / will be stored.
|
||||||
|
|
||||||
Attributes
|
Raises
|
||||||
----------
|
------
|
||||||
datadir : pathlib.Path
|
TypeError
|
||||||
The directory where the data is / will be stored.
|
If ``types`` is not a list or if the values are not string.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, types: List[str], datadir: Union[str, Path]) -> None:
|
def __init__(self, types: List[str], datadir: Union[str, Path]) -> None:
|
||||||
# Validate types
|
# Validate types
|
||||||
validate_types(types)
|
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
|
||||||
|
)
|
||||||
self.types = types
|
self.types = types
|
||||||
|
|
||||||
# Convert str to Path
|
# Convert str to Path
|
||||||
|
|
|
||||||
|
|
@ -10,8 +10,8 @@ from pathlib import Path
|
||||||
from typing import Dict, List, Union
|
from typing import Dict, List, Union
|
||||||
|
|
||||||
from ...api.decorators import register_datagrabber
|
from ...api.decorators import register_datagrabber
|
||||||
|
from ...utils import raise_error
|
||||||
from ..pattern import PatternDataGrabber
|
from ..pattern import PatternDataGrabber
|
||||||
from ..utils import raise_error
|
|
||||||
|
|
||||||
|
|
||||||
__all__ = ["HCP1200"]
|
__all__ = ["HCP1200"]
|
||||||
|
|
|
||||||
|
|
@ -7,13 +7,15 @@
|
||||||
|
|
||||||
from typing import Dict, List, Tuple, Union
|
from typing import Dict, List, Tuple, Union
|
||||||
|
|
||||||
from ..utils import raise_error
|
from ..api.decorators import register_datagrabber
|
||||||
|
from ..utils import deep_update, raise_error
|
||||||
from .base import BaseDataGrabber
|
from .base import BaseDataGrabber
|
||||||
|
|
||||||
|
|
||||||
__all__ = ["MultipleDataGrabber"]
|
__all__ = ["MultipleDataGrabber"]
|
||||||
|
|
||||||
|
|
||||||
|
@register_datagrabber
|
||||||
class MultipleDataGrabber(BaseDataGrabber):
|
class MultipleDataGrabber(BaseDataGrabber):
|
||||||
"""Concrete implementation for multi sourced data fetching.
|
"""Concrete implementation for multi sourced data fetching.
|
||||||
|
|
||||||
|
|
@ -27,19 +29,53 @@ class MultipleDataGrabber(BaseDataGrabber):
|
||||||
**kwargs
|
**kwargs
|
||||||
Keyword arguments passed to superclass.
|
Keyword arguments passed to superclass.
|
||||||
|
|
||||||
|
Raises
|
||||||
|
------
|
||||||
|
RuntimeError
|
||||||
|
If ``datagrabbers`` have different element keys or
|
||||||
|
overlapping data types or nested data types.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, datagrabbers: List[BaseDataGrabber], **kwargs) -> None:
|
def __init__(self, datagrabbers: List[BaseDataGrabber], **kwargs) -> None:
|
||||||
# Check datagrabbers consistency
|
# Check datagrabbers consistency
|
||||||
# 1) same element keys
|
# Check for same element keys
|
||||||
first_keys = datagrabbers[0].get_element_keys()
|
first_keys = datagrabbers[0].get_element_keys()
|
||||||
for dg in datagrabbers[1:]:
|
for dg in datagrabbers[1:]:
|
||||||
if dg.get_element_keys() != first_keys:
|
if dg.get_element_keys() != first_keys:
|
||||||
raise_error("DataGrabbers have different element keys.")
|
raise_error(
|
||||||
# 2) no overlapping types
|
msg="DataGrabbers have different element keys",
|
||||||
|
klass=RuntimeError,
|
||||||
|
)
|
||||||
|
# Check for no overlapping types (and nested data types)
|
||||||
types = [x for dg in datagrabbers for x in dg.get_types()]
|
types = [x for dg in datagrabbers for x in dg.get_types()]
|
||||||
if len(types) != len(set(types)):
|
if len(types) != len(set(types)):
|
||||||
raise_error("DataGrabbers have overlapping types.")
|
if all(hasattr(dg, "patterns") for dg in datagrabbers):
|
||||||
|
first_patterns = datagrabbers[0].patterns
|
||||||
|
for dg in datagrabbers[1:]:
|
||||||
|
for data_type in set(types):
|
||||||
|
dtype_pattern = dg.patterns.get(data_type)
|
||||||
|
if dtype_pattern is None:
|
||||||
|
continue
|
||||||
|
# Check if first-level keys of data type are same
|
||||||
|
if (
|
||||||
|
dtype_pattern.keys()
|
||||||
|
== first_patterns[data_type].keys()
|
||||||
|
):
|
||||||
|
raise_error(
|
||||||
|
msg=(
|
||||||
|
"DataGrabbers have overlapping mandatory "
|
||||||
|
"and / or optional key(s) for data type: "
|
||||||
|
f"`{data_type}`"
|
||||||
|
),
|
||||||
|
klass=RuntimeError,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# Can't check further
|
||||||
|
raise_error(
|
||||||
|
msg="DataGrabbers have overlapping types",
|
||||||
|
klass=RuntimeError,
|
||||||
|
)
|
||||||
self._datagrabbers = datagrabbers
|
self._datagrabbers = datagrabbers
|
||||||
|
|
||||||
def __getitem__(self, element: Union[str, Tuple]) -> Dict:
|
def __getitem__(self, element: Union[str, Tuple]) -> Dict:
|
||||||
|
|
@ -65,7 +101,7 @@ class MultipleDataGrabber(BaseDataGrabber):
|
||||||
metas = []
|
metas = []
|
||||||
for dg in self._datagrabbers:
|
for dg in self._datagrabbers:
|
||||||
t_out = dg[element]
|
t_out = dg[element]
|
||||||
out.update(t_out)
|
deep_update(out, t_out)
|
||||||
# Now get the meta for this datagrabber
|
# Now get the meta for this datagrabber
|
||||||
t_meta = {}
|
t_meta = {}
|
||||||
dg.update_meta(t_meta, "datagrabber")
|
dg.update_meta(t_meta, "datagrabber")
|
||||||
|
|
|
||||||
|
|
@ -15,7 +15,7 @@ import numpy as np
|
||||||
from ..api.decorators import register_datagrabber
|
from ..api.decorators import register_datagrabber
|
||||||
from ..utils import logger, raise_error
|
from ..utils import logger, raise_error
|
||||||
from .base import BaseDataGrabber
|
from .base import BaseDataGrabber
|
||||||
from .utils import validate_patterns, validate_replacements
|
from .pattern_validation_mixin import PatternValidationMixin
|
||||||
|
|
||||||
|
|
||||||
__all__ = ["PatternDataGrabber"]
|
__all__ = ["PatternDataGrabber"]
|
||||||
|
|
@ -26,7 +26,7 @@ _CONFOUNDS_FORMATS = ("fmriprep", "adhoc")
|
||||||
|
|
||||||
|
|
||||||
@register_datagrabber
|
@register_datagrabber
|
||||||
class PatternDataGrabber(BaseDataGrabber):
|
class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin):
|
||||||
"""Concrete implementation for pattern-based data fetching.
|
"""Concrete implementation for pattern-based data fetching.
|
||||||
|
|
||||||
Implements a DataGrabber that understands patterns to grab data.
|
Implements a DataGrabber that understands patterns to grab data.
|
||||||
|
|
@ -142,6 +142,13 @@ class PatternDataGrabber(BaseDataGrabber):
|
||||||
The directory where the data is / will be stored.
|
The directory where the data is / will be stored.
|
||||||
confounds_format : {"fmriprep", "adhoc"} or None, optional
|
confounds_format : {"fmriprep", "adhoc"} or None, optional
|
||||||
The format of the confounds for the dataset (default None).
|
The format of the confounds for the dataset (default None).
|
||||||
|
partial_pattern_ok : bool, optional
|
||||||
|
Whether to raise error if partial pattern for a data type is found.
|
||||||
|
This allows to bypass mandatory key check and issue a warning
|
||||||
|
instead of raising error. This allows one to have a DataGrabber
|
||||||
|
with data types without the corresponding mandatory keys and is
|
||||||
|
powerful when used with :class:`.MultipleDataGrabber`
|
||||||
|
(default True).
|
||||||
|
|
||||||
Raises
|
Raises
|
||||||
------
|
------
|
||||||
|
|
@ -157,17 +164,21 @@ class PatternDataGrabber(BaseDataGrabber):
|
||||||
replacements: Union[List[str], str],
|
replacements: Union[List[str], str],
|
||||||
datadir: Union[str, Path],
|
datadir: Union[str, Path],
|
||||||
confounds_format: Optional[str] = None,
|
confounds_format: Optional[str] = None,
|
||||||
|
partial_pattern_ok: bool = False,
|
||||||
) -> None:
|
) -> None:
|
||||||
# Validate patterns
|
|
||||||
validate_patterns(types=types, patterns=patterns)
|
|
||||||
self.patterns = patterns
|
|
||||||
|
|
||||||
# Convert replacements to list if not already
|
# Convert replacements to list if not already
|
||||||
if not isinstance(replacements, list):
|
if not isinstance(replacements, list):
|
||||||
replacements = [replacements]
|
replacements = [replacements]
|
||||||
# Validate replacements
|
# Validate patterns
|
||||||
validate_replacements(replacements=replacements, patterns=patterns)
|
self.validate_patterns(
|
||||||
|
types=types,
|
||||||
|
replacements=replacements,
|
||||||
|
patterns=patterns,
|
||||||
|
partial_pattern_ok=partial_pattern_ok,
|
||||||
|
)
|
||||||
self.replacements = replacements
|
self.replacements = replacements
|
||||||
|
self.patterns = patterns
|
||||||
|
self.partial_pattern_ok = partial_pattern_ok
|
||||||
|
|
||||||
# Validate confounds format
|
# Validate confounds format
|
||||||
if (
|
if (
|
||||||
|
|
@ -436,14 +447,26 @@ class PatternDataGrabber(BaseDataGrabber):
|
||||||
for t_idx in reversed(order):
|
for t_idx in reversed(order):
|
||||||
t_type = self.types[t_idx]
|
t_type = self.types[t_idx]
|
||||||
types_element = set()
|
types_element = set()
|
||||||
# Get the pattern
|
|
||||||
|
# Get the pattern dict
|
||||||
t_pattern = self.patterns[t_type]
|
t_pattern = self.patterns[t_type]
|
||||||
|
# Conditional fetch of base pattern for getting elements
|
||||||
|
pattern = None
|
||||||
|
# Try for data type pattern
|
||||||
|
pattern = t_pattern.get("pattern")
|
||||||
|
# Try for nested data type pattern
|
||||||
|
if pattern is None and self.partial_pattern_ok:
|
||||||
|
for v in t_pattern.values():
|
||||||
|
if isinstance(v, dict) and "pattern" in v:
|
||||||
|
pattern = v["pattern"]
|
||||||
|
break
|
||||||
|
|
||||||
# Replace the pattern
|
# Replace the pattern
|
||||||
(
|
(
|
||||||
re_pattern,
|
re_pattern,
|
||||||
glob_pattern,
|
glob_pattern,
|
||||||
t_replacements,
|
t_replacements,
|
||||||
) = self._replace_patterns_regex(t_pattern["pattern"])
|
) = self._replace_patterns_regex(pattern)
|
||||||
for fname in self.datadir.glob(glob_pattern):
|
for fname in self.datadir.glob(glob_pattern):
|
||||||
suffix = fname.relative_to(self.datadir).as_posix()
|
suffix = fname.relative_to(self.datadir).as_posix()
|
||||||
m = re.match(re_pattern, suffix)
|
m = re.match(re_pattern, suffix)
|
||||||
|
|
|
||||||
388
junifer/datagrabber/pattern_validation_mixin.py
Normal file
388
junifer/datagrabber/pattern_validation_mixin.py
Normal file
|
|
@ -0,0 +1,388 @@
|
||||||
|
"""Provide mixin validation class for pattern-based DataGrabber."""
|
||||||
|
|
||||||
|
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
|
# License: AGPL
|
||||||
|
|
||||||
|
from typing import Dict, List
|
||||||
|
|
||||||
|
from ..utils import logger, raise_error, warn_with_log
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = ["PatternValidationMixin"]
|
||||||
|
|
||||||
|
|
||||||
|
# Define schema for pattern-based datagrabber's patterns
|
||||||
|
PATTERNS_SCHEMA = {
|
||||||
|
"T1w": {
|
||||||
|
"mandatory": ["pattern", "space"],
|
||||||
|
"optional": {
|
||||||
|
"mask": {"mandatory": ["pattern", "space"], "optional": []},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"T2w": {
|
||||||
|
"mandatory": ["pattern", "space"],
|
||||||
|
"optional": {
|
||||||
|
"mask": {"mandatory": ["pattern", "space"], "optional": []},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"BOLD": {
|
||||||
|
"mandatory": ["pattern", "space"],
|
||||||
|
"optional": {
|
||||||
|
"mask": {"mandatory": ["pattern", "space"], "optional": []},
|
||||||
|
"confounds": {
|
||||||
|
"mandatory": ["pattern", "format"],
|
||||||
|
"optional": ["mappings"],
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
"Warp": {
|
||||||
|
"mandatory": ["pattern", "src", "dst"],
|
||||||
|
"optional": {},
|
||||||
|
},
|
||||||
|
"VBM_GM": {
|
||||||
|
"mandatory": ["pattern", "space"],
|
||||||
|
"optional": {},
|
||||||
|
},
|
||||||
|
"VBM_WM": {
|
||||||
|
"mandatory": ["pattern", "space"],
|
||||||
|
"optional": {},
|
||||||
|
},
|
||||||
|
"VBM_CSF": {
|
||||||
|
"mandatory": ["pattern", "space"],
|
||||||
|
"optional": {},
|
||||||
|
},
|
||||||
|
"DWI": {
|
||||||
|
"mandatory": ["pattern"],
|
||||||
|
"optional": {},
|
||||||
|
},
|
||||||
|
"FreeSurfer": {
|
||||||
|
"mandatory": ["pattern"],
|
||||||
|
"optional": {
|
||||||
|
"aseg": {"mandatory": ["pattern"], "optional": []},
|
||||||
|
"norm": {"mandatory": ["pattern"], "optional": []},
|
||||||
|
"lh_white": {"mandatory": ["pattern"], "optional": []},
|
||||||
|
"rh_white": {"mandatory": ["pattern"], "optional": []},
|
||||||
|
"lh_pial": {"mandatory": ["pattern"], "optional": []},
|
||||||
|
"rh_pial": {"mandatory": ["pattern"], "optional": []},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class PatternValidationMixin:
|
||||||
|
"""Mixin class for pattern validation."""
|
||||||
|
|
||||||
|
def _validate_types(self, types: List[str]) -> None:
|
||||||
|
"""Validate the types.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
types : list of str
|
||||||
|
The data types to validate.
|
||||||
|
|
||||||
|
Raises
|
||||||
|
------
|
||||||
|
TypeError
|
||||||
|
If ``types`` is not a list or if the values are not string.
|
||||||
|
|
||||||
|
"""
|
||||||
|
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(
|
||||||
|
self,
|
||||||
|
replacements: List[str],
|
||||||
|
patterns: Dict[str, Dict[str, str]],
|
||||||
|
partial_pattern_ok: bool,
|
||||||
|
) -> None:
|
||||||
|
"""Validate the replacements.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
replacements : list of str
|
||||||
|
The replacements to validate.
|
||||||
|
patterns : dict
|
||||||
|
The patterns to validate replacements against.
|
||||||
|
partial_pattern_ok : bool
|
||||||
|
Whether to raise error if partial pattern for a data type is found.
|
||||||
|
|
||||||
|
Raises
|
||||||
|
------
|
||||||
|
TypeError
|
||||||
|
If ``replacements`` is not a list or if the values are not string.
|
||||||
|
ValueError
|
||||||
|
If a value in ``replacements`` is not part of a data type pattern
|
||||||
|
and ``partial_pattern_ok=False`` or
|
||||||
|
if no data type patterns contain all values in ``replacements`` and
|
||||||
|
``partial_pattern_ok=False``.
|
||||||
|
|
||||||
|
Warns
|
||||||
|
-----
|
||||||
|
RuntimeWarning
|
||||||
|
If a value in ``replacements`` is not part of the data type pattern
|
||||||
|
and ``partial_pattern_ok=True``.
|
||||||
|
|
||||||
|
"""
|
||||||
|
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 [
|
||||||
|
data_type_val.get("pattern", "")
|
||||||
|
for data_type_val in patterns.values()
|
||||||
|
]
|
||||||
|
):
|
||||||
|
if partial_pattern_ok:
|
||||||
|
warn_with_log(
|
||||||
|
f"Replacement: `{x}` is not part of any pattern, "
|
||||||
|
"things might not work as expected if you are unsure "
|
||||||
|
"of what you are doing"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
raise_error(
|
||||||
|
msg=f"Replacement: {x} is not part of any pattern."
|
||||||
|
)
|
||||||
|
|
||||||
|
# Check that at least one pattern has all the replacements
|
||||||
|
at_least_one = False
|
||||||
|
for data_type_val in patterns.values():
|
||||||
|
if all(
|
||||||
|
x in data_type_val.get("pattern", "") for x in replacements
|
||||||
|
):
|
||||||
|
at_least_one = True
|
||||||
|
if not at_least_one and not partial_pattern_ok:
|
||||||
|
raise_error(
|
||||||
|
msg="At least one pattern must contain all replacements."
|
||||||
|
)
|
||||||
|
|
||||||
|
def _validate_mandatory_keys(
|
||||||
|
self,
|
||||||
|
keys: List[str],
|
||||||
|
schema: List[str],
|
||||||
|
data_type: str,
|
||||||
|
partial_pattern_ok: bool = False,
|
||||||
|
) -> None:
|
||||||
|
"""Validate mandatory keys.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
keys : list of str
|
||||||
|
The keys to validate.
|
||||||
|
schema : list of str
|
||||||
|
The schema to validate against.
|
||||||
|
data_type : str
|
||||||
|
The data type being validated.
|
||||||
|
partial_pattern_ok : bool, optional
|
||||||
|
Whether to raise error if partial pattern for a data type is found
|
||||||
|
(default True).
|
||||||
|
|
||||||
|
Raises
|
||||||
|
------
|
||||||
|
KeyError
|
||||||
|
If any mandatory key is missing for a data type and
|
||||||
|
``partial_pattern_ok=False``.
|
||||||
|
|
||||||
|
Warns
|
||||||
|
-----
|
||||||
|
RuntimeWarning
|
||||||
|
If any mandatory key is missing for a data type and
|
||||||
|
``partial_pattern_ok=True``.
|
||||||
|
|
||||||
|
"""
|
||||||
|
for key in schema:
|
||||||
|
if key not in keys:
|
||||||
|
if partial_pattern_ok:
|
||||||
|
warn_with_log(
|
||||||
|
f"Mandatory key: `{key}` not found for {data_type}, "
|
||||||
|
"things might not work as expected if you are unsure "
|
||||||
|
"of what you are doing"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
raise_error(
|
||||||
|
msg=f"Mandatory key: `{key}` missing for {data_type}",
|
||||||
|
klass=KeyError,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logger.debug(f"Mandatory key: `{key}` found for {data_type}")
|
||||||
|
|
||||||
|
def _identify_stray_keys(
|
||||||
|
self, keys: List[str], schema: List[str], data_type: str
|
||||||
|
) -> None:
|
||||||
|
"""Identify stray keys.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
keys : list of str
|
||||||
|
The keys to check.
|
||||||
|
schema : list of str
|
||||||
|
The schema to check against.
|
||||||
|
data_type : str
|
||||||
|
The data type being checked.
|
||||||
|
|
||||||
|
Raises
|
||||||
|
------
|
||||||
|
RuntimeError
|
||||||
|
If an unknown key is found for a data type.
|
||||||
|
|
||||||
|
"""
|
||||||
|
for key in keys:
|
||||||
|
if key not in schema:
|
||||||
|
raise_error(
|
||||||
|
msg=(
|
||||||
|
f"Key: {key} not accepted for {data_type} "
|
||||||
|
"pattern, remove it to proceed"
|
||||||
|
),
|
||||||
|
klass=RuntimeError,
|
||||||
|
)
|
||||||
|
|
||||||
|
def validate_patterns(
|
||||||
|
self,
|
||||||
|
types: List[str],
|
||||||
|
replacements: List[str],
|
||||||
|
patterns: Dict[str, Dict[str, str]],
|
||||||
|
partial_pattern_ok: bool = False,
|
||||||
|
) -> None:
|
||||||
|
"""Validate the patterns.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
types : list of str
|
||||||
|
The data types to check patterns of.
|
||||||
|
replacements : list of str
|
||||||
|
The replacements to be replaced in the patterns.
|
||||||
|
patterns : dict
|
||||||
|
The patterns to validate.
|
||||||
|
partial_pattern_ok : bool, optional
|
||||||
|
Whether to raise error if partial pattern for a data type is found.
|
||||||
|
If False, a warning is issued instead of raising an error
|
||||||
|
(default False).
|
||||||
|
|
||||||
|
Raises
|
||||||
|
------
|
||||||
|
TypeError
|
||||||
|
If ``patterns`` is not a dictionary.
|
||||||
|
ValueError
|
||||||
|
If length of ``types`` and ``patterns`` are different or
|
||||||
|
if ``patterns`` is missing entries from ``types`` or
|
||||||
|
if unknown data type is found in ``patterns`` or
|
||||||
|
if data type pattern key contains '*' as value.
|
||||||
|
|
||||||
|
"""
|
||||||
|
# Validate types
|
||||||
|
self._validate_types(types=types)
|
||||||
|
|
||||||
|
# Validate patterns
|
||||||
|
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="Length of `types` more than that of `patterns`",
|
||||||
|
klass=ValueError,
|
||||||
|
)
|
||||||
|
# Missing type in patterns
|
||||||
|
if any(x not in patterns for x in types):
|
||||||
|
raise_error(
|
||||||
|
msg="`patterns` must contain all `types`", klass=ValueError
|
||||||
|
)
|
||||||
|
# Check against schema
|
||||||
|
for data_type_key, data_type_val in patterns.items():
|
||||||
|
# Check if valid data type is provided
|
||||||
|
if data_type_key not in PATTERNS_SCHEMA:
|
||||||
|
raise_error(
|
||||||
|
f"Unknown data type: {data_type_key}, "
|
||||||
|
f"should be one of: {list(PATTERNS_SCHEMA.keys())}"
|
||||||
|
)
|
||||||
|
# Check mandatory keys for data type
|
||||||
|
self._validate_mandatory_keys(
|
||||||
|
keys=list(data_type_val),
|
||||||
|
schema=PATTERNS_SCHEMA[data_type_key]["mandatory"],
|
||||||
|
data_type=data_type_key,
|
||||||
|
partial_pattern_ok=partial_pattern_ok,
|
||||||
|
)
|
||||||
|
# Check optional keys for data type
|
||||||
|
for optional_key, optional_val in PATTERNS_SCHEMA[data_type_key][
|
||||||
|
"optional"
|
||||||
|
].items():
|
||||||
|
if optional_key not in data_type_val:
|
||||||
|
logger.debug(
|
||||||
|
f"Optional key: `{optional_key}` missing for "
|
||||||
|
f"{data_type_key}"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logger.debug(
|
||||||
|
f"Optional key: `{optional_key}` found for "
|
||||||
|
f"{data_type_key}"
|
||||||
|
)
|
||||||
|
# Set nested type name for easier access
|
||||||
|
nested_data_type = f"{data_type_key}.{optional_key}"
|
||||||
|
nested_mandatory_keys_schema = PATTERNS_SCHEMA[
|
||||||
|
data_type_key
|
||||||
|
]["optional"][optional_key]["mandatory"]
|
||||||
|
nested_optional_keys_schema = PATTERNS_SCHEMA[
|
||||||
|
data_type_key
|
||||||
|
]["optional"][optional_key]["optional"]
|
||||||
|
# Check mandatory keys for nested type
|
||||||
|
self._validate_mandatory_keys(
|
||||||
|
keys=list(optional_val["mandatory"]),
|
||||||
|
schema=nested_mandatory_keys_schema,
|
||||||
|
data_type=nested_data_type,
|
||||||
|
partial_pattern_ok=partial_pattern_ok,
|
||||||
|
)
|
||||||
|
# Check optional keys for nested type
|
||||||
|
for nested_optional_key in nested_optional_keys_schema:
|
||||||
|
if nested_optional_key not in optional_val["optional"]:
|
||||||
|
logger.debug(
|
||||||
|
f"Optional key: `{nested_optional_key}` "
|
||||||
|
f"missing for {nested_data_type}"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logger.debug(
|
||||||
|
f"Optional key: `{nested_optional_key}` found "
|
||||||
|
f"for {nested_data_type}"
|
||||||
|
)
|
||||||
|
# Check stray key for nested data type
|
||||||
|
self._identify_stray_keys(
|
||||||
|
keys=optional_val["mandatory"]
|
||||||
|
+ optional_val["optional"],
|
||||||
|
schema=nested_mandatory_keys_schema
|
||||||
|
+ nested_optional_keys_schema,
|
||||||
|
data_type=nested_data_type,
|
||||||
|
)
|
||||||
|
# Check stray key for data type
|
||||||
|
self._identify_stray_keys(
|
||||||
|
keys=list(data_type_val.keys()),
|
||||||
|
schema=(
|
||||||
|
PATTERNS_SCHEMA[data_type_key]["mandatory"]
|
||||||
|
+ list(PATTERNS_SCHEMA[data_type_key]["optional"].keys())
|
||||||
|
),
|
||||||
|
data_type=data_type_key,
|
||||||
|
)
|
||||||
|
# Wildcard check in patterns
|
||||||
|
if "}*" in data_type_val.get("pattern", ""):
|
||||||
|
raise_error(
|
||||||
|
msg=(
|
||||||
|
f"`{data_type_key}.pattern` must not contain `*` "
|
||||||
|
"following a replacement"
|
||||||
|
),
|
||||||
|
klass=ValueError,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Validate replacements
|
||||||
|
self._validate_replacements(
|
||||||
|
replacements=replacements,
|
||||||
|
patterns=patterns,
|
||||||
|
partial_pattern_ok=partial_pattern_ok,
|
||||||
|
)
|
||||||
|
|
@ -25,28 +25,26 @@ def test_MultipleDataGrabber() -> None:
|
||||||
repo_uri = _testing_dataset["example_bids_ses"]["uri"]
|
repo_uri = _testing_dataset["example_bids_ses"]["uri"]
|
||||||
rootdir = "example_bids_ses"
|
rootdir = "example_bids_ses"
|
||||||
replacements = ["subject", "session"]
|
replacements = ["subject", "session"]
|
||||||
pattern1 = {
|
|
||||||
"T1w": {
|
|
||||||
"pattern": (
|
|
||||||
"{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz"
|
|
||||||
),
|
|
||||||
"space": "native",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
pattern2 = {
|
|
||||||
"BOLD": {
|
|
||||||
"pattern": (
|
|
||||||
"{subject}/{session}/func/"
|
|
||||||
"{subject}_{session}_task-rest_bold.nii.gz"
|
|
||||||
),
|
|
||||||
"space": "MNI152NLin6Asym",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
dg1 = PatternDataladDataGrabber(
|
dg1 = PatternDataladDataGrabber(
|
||||||
rootdir=rootdir,
|
rootdir=rootdir,
|
||||||
uri=repo_uri,
|
uri=repo_uri,
|
||||||
types=["T1w"],
|
types=["T1w"],
|
||||||
patterns=pattern1,
|
patterns={
|
||||||
|
"T1w": {
|
||||||
|
"pattern": (
|
||||||
|
"{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz"
|
||||||
|
),
|
||||||
|
"space": "native",
|
||||||
|
"mask": {
|
||||||
|
"pattern": (
|
||||||
|
"{subject}/{session}/anat/{subject}_{session}_"
|
||||||
|
"brain_mask.nii.gz"
|
||||||
|
),
|
||||||
|
"space": "native",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
replacements=replacements,
|
replacements=replacements,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -54,7 +52,22 @@ def test_MultipleDataGrabber() -> None:
|
||||||
rootdir=rootdir,
|
rootdir=rootdir,
|
||||||
uri=repo_uri,
|
uri=repo_uri,
|
||||||
types=["BOLD"],
|
types=["BOLD"],
|
||||||
patterns=pattern2,
|
patterns={
|
||||||
|
"BOLD": {
|
||||||
|
"pattern": (
|
||||||
|
"{subject}/{session}/func/"
|
||||||
|
"{subject}_{session}_task-rest_bold.nii.gz"
|
||||||
|
),
|
||||||
|
"space": "MNI152NLin6Asym",
|
||||||
|
"mask": {
|
||||||
|
"pattern": (
|
||||||
|
"{subject}/{session}/func/"
|
||||||
|
"{subject}_{session}_task-rest_brain_mask.nii.gz"
|
||||||
|
),
|
||||||
|
"space": "MNI152NLin6Asym",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
replacements=replacements,
|
replacements=replacements,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -73,14 +86,17 @@ def test_MultipleDataGrabber() -> None:
|
||||||
with dg:
|
with dg:
|
||||||
subs = list(dg)
|
subs = list(dg)
|
||||||
assert set(subs) == set(expected_subs)
|
assert set(subs) == set(expected_subs)
|
||||||
|
# Check data type
|
||||||
elem = dg[("sub-01", "ses-01")]
|
elem = dg[("sub-01", "ses-01")]
|
||||||
|
# Check data types
|
||||||
assert "T1w" in elem
|
assert "T1w" in elem
|
||||||
assert "BOLD" in elem
|
assert "BOLD" in elem
|
||||||
|
# Check meta
|
||||||
assert "meta" in elem["BOLD"]
|
assert "meta" in elem["BOLD"]
|
||||||
meta = elem["BOLD"]["meta"]["datagrabber"]
|
meta = elem["BOLD"]["meta"]["datagrabber"]
|
||||||
assert "class" in meta
|
assert "class" in meta
|
||||||
assert meta["class"] == "MultipleDataGrabber"
|
assert meta["class"] == "MultipleDataGrabber"
|
||||||
|
# Check datagrabbers
|
||||||
assert "datagrabbers" in meta
|
assert "datagrabbers" in meta
|
||||||
assert len(meta["datagrabbers"]) == 2
|
assert len(meta["datagrabbers"]) == 2
|
||||||
assert meta["datagrabbers"][0]["class"] == "PatternDataladDataGrabber"
|
assert meta["datagrabbers"][0]["class"] == "PatternDataladDataGrabber"
|
||||||
|
|
@ -89,40 +105,37 @@ def test_MultipleDataGrabber() -> None:
|
||||||
|
|
||||||
def test_MultipleDataGrabber_no_intersection() -> None:
|
def test_MultipleDataGrabber_no_intersection() -> None:
|
||||||
"""Test MultipleDataGrabber without intersection (0 elements)."""
|
"""Test MultipleDataGrabber without intersection (0 elements)."""
|
||||||
repo_uri1 = _testing_dataset["example_bids"]["uri"]
|
|
||||||
repo_uri2 = _testing_dataset["example_bids_ses"]["uri"]
|
|
||||||
rootdir = "example_bids_ses"
|
rootdir = "example_bids_ses"
|
||||||
replacements = ["subject", "session"]
|
replacements = ["subject", "session"]
|
||||||
pattern1 = {
|
|
||||||
"T1w": {
|
|
||||||
"pattern": (
|
|
||||||
"{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz"
|
|
||||||
),
|
|
||||||
"space": "native",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
pattern2 = {
|
|
||||||
"BOLD": {
|
|
||||||
"pattern": (
|
|
||||||
"{subject}/{session}/func/"
|
|
||||||
"{subject}_{session}_task-rest_bold.nii.gz"
|
|
||||||
),
|
|
||||||
"space": "MNI152NLin6Asym",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
dg1 = PatternDataladDataGrabber(
|
dg1 = PatternDataladDataGrabber(
|
||||||
rootdir=rootdir,
|
rootdir=rootdir,
|
||||||
uri=repo_uri1,
|
uri=_testing_dataset["example_bids"]["uri"],
|
||||||
types=["T1w"],
|
types=["T1w"],
|
||||||
patterns=pattern1,
|
patterns={
|
||||||
|
"T1w": {
|
||||||
|
"pattern": (
|
||||||
|
"{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz"
|
||||||
|
),
|
||||||
|
"space": "native",
|
||||||
|
},
|
||||||
|
},
|
||||||
replacements=replacements,
|
replacements=replacements,
|
||||||
)
|
)
|
||||||
|
|
||||||
dg2 = PatternDataladDataGrabber(
|
dg2 = PatternDataladDataGrabber(
|
||||||
rootdir=rootdir,
|
rootdir=rootdir,
|
||||||
uri=repo_uri2,
|
uri=_testing_dataset["example_bids_ses"]["uri"],
|
||||||
types=["BOLD"],
|
types=["BOLD"],
|
||||||
patterns=pattern2,
|
patterns={
|
||||||
|
"BOLD": {
|
||||||
|
"pattern": (
|
||||||
|
"{subject}/{session}/func/"
|
||||||
|
"{subject}_{session}_task-rest_bold.nii.gz"
|
||||||
|
),
|
||||||
|
"space": "MNI152NLin6Asym",
|
||||||
|
},
|
||||||
|
},
|
||||||
replacements=replacements,
|
replacements=replacements,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -135,23 +148,19 @@ def test_MultipleDataGrabber_no_intersection() -> None:
|
||||||
|
|
||||||
def test_MultipleDataGrabber_get_item() -> None:
|
def test_MultipleDataGrabber_get_item() -> None:
|
||||||
"""Test MultipleDataGrabber get_item() error."""
|
"""Test MultipleDataGrabber get_item() error."""
|
||||||
repo_uri1 = _testing_dataset["example_bids"]["uri"]
|
|
||||||
rootdir = "example_bids_ses"
|
|
||||||
replacements = ["subject", "session"]
|
|
||||||
pattern1 = {
|
|
||||||
"T1w": {
|
|
||||||
"pattern": (
|
|
||||||
"{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz"
|
|
||||||
),
|
|
||||||
"space": "native",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
dg1 = PatternDataladDataGrabber(
|
dg1 = PatternDataladDataGrabber(
|
||||||
rootdir=rootdir,
|
rootdir="example_bids_ses",
|
||||||
uri=repo_uri1,
|
uri=_testing_dataset["example_bids"]["uri"],
|
||||||
types=["T1w"],
|
types=["T1w"],
|
||||||
patterns=pattern1,
|
patterns={
|
||||||
replacements=replacements,
|
"T1w": {
|
||||||
|
"pattern": (
|
||||||
|
"{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz"
|
||||||
|
),
|
||||||
|
"space": "native",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
replacements=["subject", "session"],
|
||||||
)
|
)
|
||||||
|
|
||||||
dg = MultipleDataGrabber([dg1])
|
dg = MultipleDataGrabber([dg1])
|
||||||
|
|
@ -161,43 +170,111 @@ def test_MultipleDataGrabber_get_item() -> None:
|
||||||
|
|
||||||
def test_MultipleDataGrabber_validation() -> None:
|
def test_MultipleDataGrabber_validation() -> None:
|
||||||
"""Test MultipleDataGrabber init validation."""
|
"""Test MultipleDataGrabber init validation."""
|
||||||
repo_uri1 = _testing_dataset["example_bids"]["uri"]
|
|
||||||
repo_uri2 = _testing_dataset["example_bids_ses"]["uri"]
|
|
||||||
rootdir = "example_bids_ses"
|
rootdir = "example_bids_ses"
|
||||||
replacement1 = ["subject", "session"]
|
|
||||||
replacement2 = ["subject"]
|
|
||||||
pattern1 = {
|
|
||||||
"T1w": {
|
|
||||||
"pattern": (
|
|
||||||
"{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz"
|
|
||||||
),
|
|
||||||
"space": "native",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
pattern2 = {
|
|
||||||
"BOLD": {
|
|
||||||
"pattern": "{subject}/func/{subject}_task-rest_bold.nii.gz",
|
|
||||||
"space": "MNI152NLin6Asym",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
dg1 = PatternDataladDataGrabber(
|
dg1 = PatternDataladDataGrabber(
|
||||||
rootdir=rootdir,
|
rootdir=rootdir,
|
||||||
uri=repo_uri1,
|
uri=_testing_dataset["example_bids"]["uri"],
|
||||||
types=["T1w"],
|
types=["T1w"],
|
||||||
patterns=pattern1,
|
patterns={
|
||||||
replacements=replacement1,
|
"T1w": {
|
||||||
|
"pattern": (
|
||||||
|
"{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz"
|
||||||
|
),
|
||||||
|
"space": "native",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
replacements=["subject", "session"],
|
||||||
)
|
)
|
||||||
|
|
||||||
dg2 = PatternDataladDataGrabber(
|
dg2 = PatternDataladDataGrabber(
|
||||||
rootdir=rootdir,
|
rootdir=rootdir,
|
||||||
uri=repo_uri2,
|
uri=_testing_dataset["example_bids_ses"]["uri"],
|
||||||
types=["BOLD"],
|
types=["BOLD"],
|
||||||
patterns=pattern2,
|
patterns={
|
||||||
replacements=replacement2,
|
"BOLD": {
|
||||||
|
"pattern": "{subject}/func/{subject}_task-rest_bold.nii.gz",
|
||||||
|
"space": "MNI152NLin6Asym",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
replacements=["subject"],
|
||||||
)
|
)
|
||||||
|
|
||||||
with pytest.raises(ValueError, match="different element key"):
|
with pytest.raises(RuntimeError, match="have different element keys"):
|
||||||
MultipleDataGrabber([dg1, dg2])
|
MultipleDataGrabber([dg1, dg2])
|
||||||
|
|
||||||
with pytest.raises(ValueError, match="overlapping types"):
|
with pytest.raises(RuntimeError, match="have overlapping mandatory"):
|
||||||
MultipleDataGrabber([dg1, dg1])
|
MultipleDataGrabber([dg1, dg1])
|
||||||
|
|
||||||
|
|
||||||
|
def test_MultipleDataGrabber_partial_pattern() -> None:
|
||||||
|
"""Test MultipleDataGrabber partial pattern."""
|
||||||
|
repo_uri = _testing_dataset["example_bids_ses"]["uri"]
|
||||||
|
rootdir = "example_bids_ses"
|
||||||
|
replacements = ["subject", "session"]
|
||||||
|
|
||||||
|
dg1 = PatternDataladDataGrabber(
|
||||||
|
rootdir=rootdir,
|
||||||
|
uri=repo_uri,
|
||||||
|
types=["BOLD"],
|
||||||
|
patterns={
|
||||||
|
"BOLD": {
|
||||||
|
"pattern": (
|
||||||
|
"{subject}/{session}/func/"
|
||||||
|
"{subject}_{session}_task-rest_bold.nii.gz"
|
||||||
|
),
|
||||||
|
"space": "MNI152NLin6Asym",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
replacements=replacements,
|
||||||
|
)
|
||||||
|
|
||||||
|
dg2 = PatternDataladDataGrabber(
|
||||||
|
rootdir=rootdir,
|
||||||
|
uri=repo_uri,
|
||||||
|
types=["BOLD"],
|
||||||
|
patterns={
|
||||||
|
"BOLD": {
|
||||||
|
"confounds": {
|
||||||
|
"pattern": (
|
||||||
|
"{subject}/{session}/func/"
|
||||||
|
"{subject}_{session}_task-rest_"
|
||||||
|
"confounds_regressors.tsv"
|
||||||
|
),
|
||||||
|
"format": "fmriprep",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
},
|
||||||
|
replacements=["subject", "session"],
|
||||||
|
partial_pattern_ok=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
dg = MultipleDataGrabber([dg1, dg2])
|
||||||
|
|
||||||
|
types = dg.get_types()
|
||||||
|
assert "BOLD" in types
|
||||||
|
|
||||||
|
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 = list(dg)
|
||||||
|
assert set(subs) == set(expected_subs)
|
||||||
|
# Fetch element
|
||||||
|
elem = dg[("sub-01", "ses-01")]
|
||||||
|
# Check data type and nested data type
|
||||||
|
assert "BOLD" in elem
|
||||||
|
assert "confounds" in elem["BOLD"]
|
||||||
|
# Check meta
|
||||||
|
assert "meta" in elem["BOLD"]
|
||||||
|
meta = elem["BOLD"]["meta"]["datagrabber"]
|
||||||
|
assert "class" in meta
|
||||||
|
assert meta["class"] == "MultipleDataGrabber"
|
||||||
|
# Check datagrabbers
|
||||||
|
assert "datagrabbers" in meta
|
||||||
|
assert len(meta["datagrabbers"]) == 2
|
||||||
|
assert meta["datagrabbers"][0]["class"] == "PatternDataladDataGrabber"
|
||||||
|
assert meta["datagrabbers"][1]["class"] == "PatternDataladDataGrabber"
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,7 @@
|
||||||
"""Provide tests for utils."""
|
"""Provide tests for PatternValidationMixin."""
|
||||||
|
|
||||||
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
||||||
|
# Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from contextlib import nullcontext
|
from contextlib import nullcontext
|
||||||
|
|
@ -8,136 +9,57 @@ from typing import ContextManager, Dict, List, Union
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from junifer.datagrabber.utils import (
|
from junifer.datagrabber.pattern_validation_mixin import PatternValidationMixin
|
||||||
validate_patterns,
|
|
||||||
validate_replacements,
|
|
||||||
validate_types,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
"types, expect",
|
"types, replacements, patterns, expect",
|
||||||
[
|
|
||||||
("wrong", pytest.raises(TypeError, match="must be a list")),
|
|
||||||
([1], pytest.raises(TypeError, match="must be a list of strings")),
|
|
||||||
(["T1w", "BOLD"], nullcontext()),
|
|
||||||
],
|
|
||||||
)
|
|
||||||
def test_validate_types(
|
|
||||||
types: Union[str, List[str], List[int]],
|
|
||||||
expect: ContextManager,
|
|
||||||
) -> None:
|
|
||||||
"""Test validation of types.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
types : str, list of int or str
|
|
||||||
The parametrized data types to validate.
|
|
||||||
expect : typing.ContextManager
|
|
||||||
The parametrized ContextManager object.
|
|
||||||
|
|
||||||
"""
|
|
||||||
with expect:
|
|
||||||
validate_types(types) # type: ignore
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
|
||||||
"replacements, patterns, expect",
|
|
||||||
[
|
[
|
||||||
(
|
(
|
||||||
"wrong",
|
"wrong",
|
||||||
"also wrong",
|
[],
|
||||||
pytest.raises(TypeError, match="must be a list"),
|
{},
|
||||||
|
pytest.raises(TypeError, match="`types` must be a list"),
|
||||||
),
|
),
|
||||||
(
|
(
|
||||||
[1],
|
[1],
|
||||||
{
|
[],
|
||||||
"T1w": {"pattern": "{subject}/anat/{subject}_T1w.nii.gz"},
|
{},
|
||||||
"BOLD": {
|
pytest.raises(
|
||||||
"pattern": "{subject}/func/{subject}_task-rest_bold.nii.gz"
|
TypeError, match="`types` must be a list of strings"
|
||||||
},
|
),
|
||||||
},
|
|
||||||
pytest.raises(TypeError, match="must be a list of strings"),
|
|
||||||
),
|
),
|
||||||
(
|
(
|
||||||
["session"],
|
["BOLD"],
|
||||||
{
|
[],
|
||||||
"T1w": {"pattern": "{subject}/anat/{subject}_T1w.nii.gz"},
|
|
||||||
"BOLD": {
|
|
||||||
"pattern": "{subject}/func/{subject}_task-rest_bold.nii.gz"
|
|
||||||
},
|
|
||||||
},
|
|
||||||
pytest.raises(ValueError, match="is not part of"),
|
|
||||||
),
|
|
||||||
(
|
|
||||||
["subject", "session"],
|
|
||||||
{
|
|
||||||
"T1w": {"pattern": "{subject}/anat/_T1w.nii.gz"},
|
|
||||||
"BOLD": {"pattern": "{session}/func/_task-rest_bold.nii.gz"},
|
|
||||||
},
|
|
||||||
pytest.raises(ValueError, match="At least one pattern"),
|
|
||||||
),
|
|
||||||
(
|
|
||||||
["subject"],
|
|
||||||
{
|
|
||||||
"T1w": {"pattern": "{subject}/anat/{subject}_T1w.nii.gz"},
|
|
||||||
"BOLD": {
|
|
||||||
"pattern": "{subject}/func/{subject}_task-rest_bold.nii.gz"
|
|
||||||
},
|
|
||||||
},
|
|
||||||
nullcontext(),
|
|
||||||
),
|
|
||||||
],
|
|
||||||
)
|
|
||||||
def test_validate_replacements(
|
|
||||||
replacements: Union[str, List[str], List[int]],
|
|
||||||
patterns: Union[str, Dict[str, Dict[str, str]]],
|
|
||||||
expect: ContextManager,
|
|
||||||
) -> None:
|
|
||||||
"""Test validation of replacements.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
replacements : str, list of str or int
|
|
||||||
The parametrized pattern replacements to validate.
|
|
||||||
patterns : str, dict
|
|
||||||
The parametrized patterns to validate against.
|
|
||||||
expect : typing.ContextManager
|
|
||||||
The parametrized ContextManager object.
|
|
||||||
|
|
||||||
"""
|
|
||||||
with expect:
|
|
||||||
validate_replacements(replacements=replacements, patterns=patterns) # type: ignore
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
|
||||||
"types, patterns, expect",
|
|
||||||
[
|
|
||||||
(
|
|
||||||
["T1w", "BOLD"],
|
|
||||||
"wrong",
|
"wrong",
|
||||||
pytest.raises(TypeError, match="must be a dict"),
|
pytest.raises(TypeError, match="`patterns` must be a dict"),
|
||||||
),
|
),
|
||||||
(
|
(
|
||||||
["T1w", "BOLD"],
|
["T1w", "BOLD"],
|
||||||
|
"",
|
||||||
{
|
{
|
||||||
"T1w": {"pattern": "{subject}/anat/{subject}_T1w.nii.gz"},
|
"T1w": {"pattern": "{subject}/anat/{subject}_T1w.nii.gz"},
|
||||||
},
|
},
|
||||||
pytest.raises(
|
pytest.raises(
|
||||||
ValueError,
|
ValueError,
|
||||||
match="Length of `types` more than that of `patterns`.",
|
match="Length of `types` more than that of `patterns`",
|
||||||
),
|
),
|
||||||
),
|
),
|
||||||
(
|
(
|
||||||
["T1w", "BOLD"],
|
["T1w", "BOLD"],
|
||||||
|
"",
|
||||||
{
|
{
|
||||||
"T1w": {"pattern": "{subject}/anat/{subject}_T1w.nii.gz"},
|
"T1w": {"pattern": "{subject}/anat/{subject}_T1w.nii.gz"},
|
||||||
"T2w": {"pattern": "{subject}/anat/{subject}_T2w.nii.gz"},
|
"T2w": {"pattern": "{subject}/anat/{subject}_T2w.nii.gz"},
|
||||||
},
|
},
|
||||||
pytest.raises(ValueError, match="contain all"),
|
pytest.raises(
|
||||||
|
ValueError, match="`patterns` must contain all `types`"
|
||||||
|
),
|
||||||
),
|
),
|
||||||
(
|
(
|
||||||
["T3w"],
|
["T3w"],
|
||||||
|
"",
|
||||||
{
|
{
|
||||||
"T3w": {"pattern": "{subject}/anat/{subject}_T3w.nii.gz"},
|
"T3w": {"pattern": "{subject}/anat/{subject}_T3w.nii.gz"},
|
||||||
},
|
},
|
||||||
|
|
@ -145,6 +67,7 @@ def test_validate_replacements(
|
||||||
),
|
),
|
||||||
(
|
(
|
||||||
["BOLD"],
|
["BOLD"],
|
||||||
|
"",
|
||||||
{
|
{
|
||||||
"BOLD": {"patterns": "{subject}/func/{subject}_BOLD.nii.gz"},
|
"BOLD": {"patterns": "{subject}/func/{subject}_BOLD.nii.gz"},
|
||||||
},
|
},
|
||||||
|
|
@ -152,6 +75,7 @@ def test_validate_replacements(
|
||||||
),
|
),
|
||||||
(
|
(
|
||||||
["BOLD"],
|
["BOLD"],
|
||||||
|
"",
|
||||||
{
|
{
|
||||||
"BOLD": {
|
"BOLD": {
|
||||||
"pattern": (
|
"pattern": (
|
||||||
|
|
@ -169,6 +93,7 @@ def test_validate_replacements(
|
||||||
),
|
),
|
||||||
(
|
(
|
||||||
["T1w"],
|
["T1w"],
|
||||||
|
"",
|
||||||
{
|
{
|
||||||
"T1w": {
|
"T1w": {
|
||||||
"pattern": "{subject}/anat/{subject}*.nii",
|
"pattern": "{subject}/anat/{subject}*.nii",
|
||||||
|
|
@ -177,8 +102,65 @@ def test_validate_replacements(
|
||||||
},
|
},
|
||||||
pytest.raises(ValueError, match="following a replacement"),
|
pytest.raises(ValueError, match="following a replacement"),
|
||||||
),
|
),
|
||||||
|
(
|
||||||
|
["T1w"],
|
||||||
|
"wrong",
|
||||||
|
{
|
||||||
|
"T1w": {
|
||||||
|
"pattern": "{subject}/anat/{subject}_T1w.nii",
|
||||||
|
"space": "native",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
pytest.raises(TypeError, match="`replacements` must be a list"),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
["T1w"],
|
||||||
|
[1],
|
||||||
|
{
|
||||||
|
"T1w": {
|
||||||
|
"pattern": "{subject}/anat/{subject}_T1w.nii",
|
||||||
|
"space": "native",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
pytest.raises(
|
||||||
|
TypeError, match="`replacements` must be a list of strings"
|
||||||
|
),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
["T1w", "BOLD"],
|
||||||
|
["subject", "session"],
|
||||||
|
{
|
||||||
|
"T1w": {
|
||||||
|
"pattern": "{subject}/anat/{subject}_T1w.nii.gz",
|
||||||
|
"space": "native",
|
||||||
|
},
|
||||||
|
"BOLD": {
|
||||||
|
"pattern": (
|
||||||
|
"{subject}/func/{subject}_task-rest_bold.nii.gz"
|
||||||
|
),
|
||||||
|
"space": "MNI152NLin6Asym",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
pytest.raises(ValueError, match="is not part of any pattern"),
|
||||||
|
),
|
||||||
|
(
|
||||||
|
["BOLD"],
|
||||||
|
["subject", "session"],
|
||||||
|
{
|
||||||
|
"T1w": {
|
||||||
|
"pattern": "{subject}/anat/_T1w.nii.gz",
|
||||||
|
"space": "native",
|
||||||
|
},
|
||||||
|
"BOLD": {
|
||||||
|
"pattern": "{session}/func/_task-rest_bold.nii.gz",
|
||||||
|
"space": "MNI152NLin6Asym",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
pytest.raises(ValueError, match="At least one pattern"),
|
||||||
|
),
|
||||||
(
|
(
|
||||||
["T1w", "T2w", "BOLD"],
|
["T1w", "T2w", "BOLD"],
|
||||||
|
["subject"],
|
||||||
{
|
{
|
||||||
"T1w": {
|
"T1w": {
|
||||||
"pattern": "{subject}/anat/{subject}_T1w.nii.gz",
|
"pattern": "{subject}/anat/{subject}_T1w.nii.gz",
|
||||||
|
|
@ -190,7 +172,7 @@ def test_validate_replacements(
|
||||||
},
|
},
|
||||||
"BOLD": {
|
"BOLD": {
|
||||||
"pattern": (
|
"pattern": (
|
||||||
"{subject}/func/{subject}_task-rest_bold.nii.gz"
|
"{subject}/func/{session}/{subject}_task-rest_bold.nii.gz"
|
||||||
),
|
),
|
||||||
"space": "MNI152NLin6Asym",
|
"space": "MNI152NLin6Asym",
|
||||||
"confounds": {
|
"confounds": {
|
||||||
|
|
@ -203,22 +185,65 @@ def test_validate_replacements(
|
||||||
),
|
),
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
def test_validate_patterns(
|
def test_PatternValidationMixin(
|
||||||
types: List[str],
|
types: Union[str, List[str], List[int]],
|
||||||
|
replacements: Union[str, List[str], List[int]],
|
||||||
patterns: Union[str, Dict[str, Dict[str, str]]],
|
patterns: Union[str, Dict[str, Dict[str, str]]],
|
||||||
expect: ContextManager,
|
expect: ContextManager,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Test validation of patterns.
|
"""Test validation.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
----------
|
----------
|
||||||
types : list of str
|
types : str, list of int or str
|
||||||
The parametrized data types.
|
The parametrized data types to validate.
|
||||||
|
replacements : str, list of str or int
|
||||||
|
The parametrized pattern replacements to validate.
|
||||||
patterns : str, dict
|
patterns : str, dict
|
||||||
The patterns to validate.
|
The parametrized patterns to validate against.
|
||||||
expect : typing.ContextManager
|
expect : typing.ContextManager
|
||||||
The parametrized ContextManager object.
|
The parametrized ContextManager object.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
class MockDataGrabber(PatternValidationMixin):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
types,
|
||||||
|
replacements,
|
||||||
|
patterns,
|
||||||
|
) -> None:
|
||||||
|
self.types = types
|
||||||
|
self.replacements = replacements
|
||||||
|
self.patterns = patterns
|
||||||
|
|
||||||
|
def validate(self) -> None:
|
||||||
|
self.validate_patterns(
|
||||||
|
types=self.types,
|
||||||
|
replacements=self.replacements,
|
||||||
|
patterns=self.patterns,
|
||||||
|
)
|
||||||
|
|
||||||
|
dg = MockDataGrabber(types, replacements, patterns)
|
||||||
with expect:
|
with expect:
|
||||||
validate_patterns(types=types, patterns=patterns) # type: ignore
|
dg.validate()
|
||||||
|
|
||||||
|
|
||||||
|
# This test is kept separate as bool doesn't support context manager protocol,
|
||||||
|
# used in the earlier test
|
||||||
|
def test_PatternValidationMixin_partial_pattern_check() -> None:
|
||||||
|
"""Test validation for partial patterns."""
|
||||||
|
with pytest.warns(RuntimeWarning, match="might not work as expected"):
|
||||||
|
PatternValidationMixin().validate_patterns(
|
||||||
|
types=["BOLD"],
|
||||||
|
replacements=["subject"],
|
||||||
|
patterns={
|
||||||
|
"BOLD": {
|
||||||
|
"mask": {
|
||||||
|
"pattern": "{subject}/func/{subject}_BOLD.nii.gz",
|
||||||
|
"space": "MNI152NLin6Asym",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}, # type: ignore
|
||||||
|
partial_pattern_ok=True,
|
||||||
|
)
|
||||||
|
|
@ -1,317 +0,0 @@
|
||||||
"""Provide utility functions for the datagrabber sub-package."""
|
|
||||||
|
|
||||||
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
|
||||||
# Synchon Mandal <s.mandal@fz-juelich.de>
|
|
||||||
# License: AGPL
|
|
||||||
|
|
||||||
from typing import Dict, List
|
|
||||||
|
|
||||||
from ..utils import logger, raise_error
|
|
||||||
|
|
||||||
|
|
||||||
__all__ = ["validate_types", "validate_replacements", "validate_patterns"]
|
|
||||||
|
|
||||||
|
|
||||||
# Define schema for pattern-based datagrabber's patterns
|
|
||||||
PATTERNS_SCHEMA = {
|
|
||||||
"T1w": {
|
|
||||||
"mandatory": ["pattern", "space"],
|
|
||||||
"optional": {
|
|
||||||
"mask": {"mandatory": ["pattern", "space"], "optional": []},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
"T2w": {
|
|
||||||
"mandatory": ["pattern", "space"],
|
|
||||||
"optional": {
|
|
||||||
"mask": {"mandatory": ["pattern", "space"], "optional": []},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
"BOLD": {
|
|
||||||
"mandatory": ["pattern", "space"],
|
|
||||||
"optional": {
|
|
||||||
"mask": {"mandatory": ["pattern", "space"], "optional": []},
|
|
||||||
"confounds": {
|
|
||||||
"mandatory": ["pattern", "format"],
|
|
||||||
"optional": ["mappings"],
|
|
||||||
},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
"Warp": {
|
|
||||||
"mandatory": ["pattern", "src", "dst"],
|
|
||||||
"optional": {},
|
|
||||||
},
|
|
||||||
"VBM_GM": {
|
|
||||||
"mandatory": ["pattern", "space"],
|
|
||||||
"optional": {},
|
|
||||||
},
|
|
||||||
"VBM_WM": {
|
|
||||||
"mandatory": ["pattern", "space"],
|
|
||||||
"optional": {},
|
|
||||||
},
|
|
||||||
"VBM_CSF": {
|
|
||||||
"mandatory": ["pattern", "space"],
|
|
||||||
"optional": {},
|
|
||||||
},
|
|
||||||
"DWI": {
|
|
||||||
"mandatory": ["pattern"],
|
|
||||||
"optional": {},
|
|
||||||
},
|
|
||||||
"FreeSurfer": {
|
|
||||||
"mandatory": ["pattern"],
|
|
||||||
"optional": {
|
|
||||||
"aseg": {"mandatory": ["pattern"], "optional": []},
|
|
||||||
"norm": {"mandatory": ["pattern"], "optional": []},
|
|
||||||
"lh_white": {"mandatory": ["pattern"], "optional": []},
|
|
||||||
"rh_white": {"mandatory": ["pattern"], "optional": []},
|
|
||||||
"lh_pial": {"mandatory": ["pattern"], "optional": []},
|
|
||||||
"rh_pial": {"mandatory": ["pattern"], "optional": []},
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
def validate_types(types: List[str]) -> None:
|
|
||||||
"""Validate the types.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
types : list of str
|
|
||||||
The object to validate.
|
|
||||||
|
|
||||||
Raises
|
|
||||||
------
|
|
||||||
TypeError
|
|
||||||
If ``types`` is not a list or if the values are not string.
|
|
||||||
|
|
||||||
"""
|
|
||||||
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, Dict[str, str]]
|
|
||||||
) -> None:
|
|
||||||
"""Validate the replacements.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
replacements : list of str
|
|
||||||
The object to validate.
|
|
||||||
patterns : dict
|
|
||||||
The patterns to validate against.
|
|
||||||
|
|
||||||
Raises
|
|
||||||
------
|
|
||||||
TypeError
|
|
||||||
If ``replacements`` is not a list or if the values are not string.
|
|
||||||
ValueError
|
|
||||||
If a value in ``replacements`` is not part of a data type pattern or
|
|
||||||
if no data type patterns contain all values in ``replacements``.
|
|
||||||
|
|
||||||
"""
|
|
||||||
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 [
|
|
||||||
data_type_val["pattern"] for data_type_val in patterns.values()
|
|
||||||
]
|
|
||||||
):
|
|
||||||
raise_error(msg=f"Replacement: {x} is not part of any pattern.")
|
|
||||||
|
|
||||||
# Check that at least one pattern has all the replacements
|
|
||||||
at_least_one = False
|
|
||||||
for data_type_val in patterns.values():
|
|
||||||
if all(x in data_type_val["pattern"] for x in replacements):
|
|
||||||
at_least_one = True
|
|
||||||
if at_least_one is False:
|
|
||||||
raise_error(msg="At least one pattern must contain all replacements.")
|
|
||||||
|
|
||||||
|
|
||||||
def _validate_mandatory_keys(
|
|
||||||
keys: List[str], schema: List[str], data_type: str
|
|
||||||
) -> None:
|
|
||||||
"""Validate mandatory keys.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
keys : list of str
|
|
||||||
The keys to validate.
|
|
||||||
schema : list of str
|
|
||||||
The schema to validate against.
|
|
||||||
data_type : str
|
|
||||||
The data type being validated.
|
|
||||||
|
|
||||||
Raises
|
|
||||||
------
|
|
||||||
KeyError
|
|
||||||
If any mandatory key is missing for a data type.
|
|
||||||
|
|
||||||
"""
|
|
||||||
for key in schema:
|
|
||||||
if key not in keys:
|
|
||||||
raise_error(
|
|
||||||
msg=f"Mandatory key: `{key}` missing for {data_type}",
|
|
||||||
klass=KeyError,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
logger.debug(f"Mandatory key: `{key}` found for {data_type}")
|
|
||||||
|
|
||||||
|
|
||||||
def _identify_stray_keys(
|
|
||||||
keys: List[str], schema: List[str], data_type: str
|
|
||||||
) -> None:
|
|
||||||
"""Identify stray keys.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
keys : list of str
|
|
||||||
The keys to check.
|
|
||||||
schema : list of str
|
|
||||||
The schema to check against.
|
|
||||||
data_type : str
|
|
||||||
The data type being checked.
|
|
||||||
|
|
||||||
Raises
|
|
||||||
------
|
|
||||||
RuntimeError
|
|
||||||
If an unknown key is found for a data type.
|
|
||||||
|
|
||||||
"""
|
|
||||||
for key in keys:
|
|
||||||
if key not in schema:
|
|
||||||
raise_error(
|
|
||||||
msg=(
|
|
||||||
f"Key: {key} not accepted for {data_type} "
|
|
||||||
"pattern, remove it to proceed"
|
|
||||||
),
|
|
||||||
klass=RuntimeError,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def validate_patterns(
|
|
||||||
types: List[str], patterns: Dict[str, Dict[str, str]]
|
|
||||||
) -> None:
|
|
||||||
"""Validate the patterns.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
types : list of str
|
|
||||||
The types list.
|
|
||||||
patterns : dict
|
|
||||||
The object to validate.
|
|
||||||
|
|
||||||
Raises
|
|
||||||
------
|
|
||||||
TypeError
|
|
||||||
If ``patterns`` is not a dictionary.
|
|
||||||
ValueError
|
|
||||||
If length of ``types`` and ``patterns`` are different or
|
|
||||||
if ``patterns`` is missing entries from ``types`` or
|
|
||||||
if unknown data type is found in ``patterns`` or
|
|
||||||
if data type pattern key contains '*' as value.
|
|
||||||
|
|
||||||
"""
|
|
||||||
# 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="Length of `types` more than that of `patterns`.",
|
|
||||||
klass=ValueError,
|
|
||||||
)
|
|
||||||
# Missing type in patterns
|
|
||||||
if any(x not in patterns for x in types):
|
|
||||||
raise_error(
|
|
||||||
msg="`patterns` must contain all `types`", klass=ValueError
|
|
||||||
)
|
|
||||||
# Check against schema
|
|
||||||
for data_type_key, data_type_val in patterns.items():
|
|
||||||
# Check if valid data type is provided
|
|
||||||
if data_type_key not in PATTERNS_SCHEMA:
|
|
||||||
raise_error(
|
|
||||||
f"Unknown data type: {data_type_key}, "
|
|
||||||
f"should be one of: {list(PATTERNS_SCHEMA.keys())}"
|
|
||||||
)
|
|
||||||
# Check mandatory keys for data type
|
|
||||||
_validate_mandatory_keys(
|
|
||||||
keys=list(data_type_val),
|
|
||||||
schema=PATTERNS_SCHEMA[data_type_key]["mandatory"],
|
|
||||||
data_type=data_type_key,
|
|
||||||
)
|
|
||||||
# Check optional keys for data type
|
|
||||||
for optional_key, optional_val in PATTERNS_SCHEMA[data_type_key][
|
|
||||||
"optional"
|
|
||||||
].items():
|
|
||||||
if optional_key not in data_type_val:
|
|
||||||
logger.debug(
|
|
||||||
f"Optional key: `{optional_key}` missing for "
|
|
||||||
f"{data_type_key}"
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
logger.debug(
|
|
||||||
f"Optional key: `{optional_key}` found for "
|
|
||||||
f"{data_type_key}"
|
|
||||||
)
|
|
||||||
# Set nested type name for easier access
|
|
||||||
nested_data_type = f"{data_type_key}.{optional_key}"
|
|
||||||
nested_mandatory_keys_schema = PATTERNS_SCHEMA[data_type_key][
|
|
||||||
"optional"
|
|
||||||
][optional_key]["mandatory"]
|
|
||||||
nested_optional_keys_schema = PATTERNS_SCHEMA[data_type_key][
|
|
||||||
"optional"
|
|
||||||
][optional_key]["optional"]
|
|
||||||
# Check mandatory keys for nested type
|
|
||||||
_validate_mandatory_keys(
|
|
||||||
keys=list(optional_val["mandatory"]),
|
|
||||||
schema=nested_mandatory_keys_schema,
|
|
||||||
data_type=nested_data_type,
|
|
||||||
)
|
|
||||||
# Check optional keys for nested type
|
|
||||||
for nested_optional_key in nested_optional_keys_schema:
|
|
||||||
if nested_optional_key not in optional_val["optional"]:
|
|
||||||
logger.debug(
|
|
||||||
f"Optional key: `{nested_optional_key}` missing "
|
|
||||||
f"for {nested_data_type}"
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
logger.debug(
|
|
||||||
f"Optional key: `{nested_optional_key}` found for "
|
|
||||||
f"{nested_data_type}"
|
|
||||||
)
|
|
||||||
# Check stray key for nested data type
|
|
||||||
_identify_stray_keys(
|
|
||||||
keys=optional_val["mandatory"] + optional_val["optional"],
|
|
||||||
schema=nested_mandatory_keys_schema
|
|
||||||
+ nested_optional_keys_schema,
|
|
||||||
data_type=nested_data_type,
|
|
||||||
)
|
|
||||||
# Check stray key for data type
|
|
||||||
_identify_stray_keys(
|
|
||||||
keys=list(data_type_val.keys()),
|
|
||||||
schema=(
|
|
||||||
PATTERNS_SCHEMA[data_type_key]["mandatory"]
|
|
||||||
+ list(PATTERNS_SCHEMA[data_type_key]["optional"].keys())
|
|
||||||
),
|
|
||||||
data_type=data_type_key,
|
|
||||||
)
|
|
||||||
# Wildcard check in patterns
|
|
||||||
if "}*" in data_type_val["pattern"]:
|
|
||||||
raise_error(
|
|
||||||
msg=(
|
|
||||||
f"`{data_type_key}.pattern` must not contain `*` "
|
|
||||||
"following a replacement"
|
|
||||||
),
|
|
||||||
klass=ValueError,
|
|
||||||
)
|
|
||||||
|
|
@ -6,7 +6,7 @@
|
||||||
|
|
||||||
from .fs import make_executable
|
from .fs import make_executable
|
||||||
from .logging import configure_logging, logger, raise_error, warn_with_log
|
from .logging import configure_logging, logger, raise_error, warn_with_log
|
||||||
from .helpers import run_ext_cmd
|
from .helpers import run_ext_cmd, deep_update
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
|
|
@ -16,4 +16,5 @@ __all__ = [
|
||||||
"raise_error",
|
"raise_error",
|
||||||
"warn_with_log",
|
"warn_with_log",
|
||||||
"run_ext_cmd",
|
"run_ext_cmd",
|
||||||
|
"deep_update",
|
||||||
]
|
]
|
||||||
|
|
|
||||||
|
|
@ -3,13 +3,14 @@
|
||||||
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
|
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
|
import collections.abc
|
||||||
import subprocess
|
import subprocess
|
||||||
from typing import List
|
from typing import Dict, List
|
||||||
|
|
||||||
from .logging import logger, raise_error
|
from .logging import logger, raise_error
|
||||||
|
|
||||||
|
|
||||||
__all__ = ["run_ext_cmd"]
|
__all__ = ["run_ext_cmd", "deep_update"]
|
||||||
|
|
||||||
|
|
||||||
def run_ext_cmd(name: str, cmd: List[str]) -> None:
|
def run_ext_cmd(name: str, cmd: List[str]) -> None:
|
||||||
|
|
@ -54,3 +55,30 @@ def run_ext_cmd(name: str, cmd: List[str]) -> None:
|
||||||
),
|
),
|
||||||
klass=RuntimeError,
|
klass=RuntimeError,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def deep_update(d: Dict, u: Dict) -> Dict:
|
||||||
|
"""Deep update `d` with `u`.
|
||||||
|
|
||||||
|
From: "https://stackoverflow.com/questions/3232943/update-value-of-a-nested
|
||||||
|
-dictionary-of-varying-depth"
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
d : dict
|
||||||
|
The dictionary to deep-update.
|
||||||
|
u : dict
|
||||||
|
The dictionary to deep-update `d` with.
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
dict
|
||||||
|
The updated dictionary.
|
||||||
|
|
||||||
|
"""
|
||||||
|
for k, v in u.items():
|
||||||
|
if isinstance(v, collections.abc.Mapping):
|
||||||
|
d[k] = deep_update(d.get(k, {}), v)
|
||||||
|
else:
|
||||||
|
d[k] = v
|
||||||
|
return d
|
||||||
|
|
|
||||||
|
|
@ -13,6 +13,7 @@ else: # pragma: no cover
|
||||||
from looseversion import LooseVersion
|
from looseversion import LooseVersion
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
import warnings
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from subprocess import PIPE, Popen, TimeoutExpired
|
from subprocess import PIPE, Popen, TimeoutExpired
|
||||||
from typing import Dict, NoReturn, Optional, Type, Union
|
from typing import Dict, NoReturn, Optional, Type, Union
|
||||||
|
|
@ -44,6 +45,23 @@ _logging_types = {
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
# Copied over from stdlib and tweaked to our use-case.
|
||||||
|
def _showwarning(message, category, filename, lineno, file=None, line=None):
|
||||||
|
s = warnings.formatwarning(message, category, filename, lineno, line)
|
||||||
|
logger.warning(str(s))
|
||||||
|
|
||||||
|
|
||||||
|
# Overwrite warnings display to integrate with logging
|
||||||
|
|
||||||
|
|
||||||
|
def capture_warnings():
|
||||||
|
"""Capture warnings and log them."""
|
||||||
|
warnings.showwarning = _showwarning
|
||||||
|
|
||||||
|
|
||||||
|
capture_warnings()
|
||||||
|
|
||||||
|
|
||||||
class WrapStdOut(logging.StreamHandler):
|
class WrapStdOut(logging.StreamHandler):
|
||||||
"""Dynamically wrap to sys.stdout.
|
"""Dynamically wrap to sys.stdout.
|
||||||
|
|
||||||
|
|
@ -325,5 +343,4 @@ def warn_with_log(
|
||||||
The warning subclass (default RuntimeWarning).
|
The warning subclass (default RuntimeWarning).
|
||||||
|
|
||||||
"""
|
"""
|
||||||
logger.warning(msg)
|
|
||||||
warn(msg, category=category, stacklevel=2)
|
warn(msg, category=category, stacklevel=2)
|
||||||
|
|
|
||||||
|
|
@ -145,8 +145,16 @@ def test_log_file(tmp_path: Path) -> None:
|
||||||
assert any("Warn3 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)
|
assert any("Error3 message" in line for line in lines)
|
||||||
|
|
||||||
|
# This should raise a warning (test that it was raised)
|
||||||
with pytest.warns(RuntimeWarning, match=r"Warn raised"):
|
with pytest.warns(RuntimeWarning, match=r"Warn raised"):
|
||||||
warn_with_log("Warn raised")
|
warn_with_log("Warn raised")
|
||||||
|
|
||||||
|
# This should log the warning (workaround for pytest messing with logging)
|
||||||
|
from junifer.utils.logging import capture_warnings
|
||||||
|
|
||||||
|
capture_warnings()
|
||||||
|
|
||||||
|
warn_with_log("Warn raised 2")
|
||||||
with pytest.raises(ValueError, match=r"Error raised"):
|
with pytest.raises(ValueError, match=r"Error raised"):
|
||||||
raise_error("Error raised")
|
raise_error("Error raised")
|
||||||
with open(tmp_path / "test4.log") as f:
|
with open(tmp_path / "test4.log") as f:
|
||||||
|
|
|
||||||
|
|
@ -4,11 +4,12 @@
|
||||||
# Vera Komeyer <v.komeyer@fz-juelich.de>
|
# Vera Komeyer <v.komeyer@fz-juelich.de>
|
||||||
# Xuan Li <xu.li@fz-juelich.de>
|
# Xuan Li <xu.li@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
from tempfile import TemporaryDirectory
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from tempfile import TemporaryDirectory
|
||||||
|
|
||||||
import datalad.api as dl
|
import datalad.api as dl
|
||||||
|
|
||||||
|
|
||||||
# repo has to be created on gin manually beforehand if not owner
|
# repo has to be created on gin manually beforehand if not owner
|
||||||
dst = "git@gin.g-node.org:/juaml/datalad-example-aomic1000.git"
|
dst = "git@gin.g-node.org:/juaml/datalad-example-aomic1000.git"
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -1,4 +1,5 @@
|
||||||
"""Create an example/testing dataset for PIOP1 with mock data."""
|
"""Create an example/testing dataset for PIOP1 with mock data."""
|
||||||
|
|
||||||
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
||||||
# Vera Komeyer <v.komeyer@fz-juelich.de>
|
# Vera Komeyer <v.komeyer@fz-juelich.de>
|
||||||
# Xuan Li <xu.li@fz-juelich.de>
|
# Xuan Li <xu.li@fz-juelich.de>
|
||||||
|
|
|
||||||
|
|
@ -1,4 +1,5 @@
|
||||||
"""Create an example/testing dataset for PIOP2 with mock data."""
|
"""Create an example/testing dataset for PIOP2 with mock data."""
|
||||||
|
|
||||||
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
||||||
# Vera Komeyer <v.komeyer@fz-juelich.de>
|
# Vera Komeyer <v.komeyer@fz-juelich.de>
|
||||||
# Xuan Li <xu.li@fz-juelich.de>
|
# Xuan Li <xu.li@fz-juelich.de>
|
||||||
|
|
|
||||||
|
|
@ -1,34 +1,40 @@
|
||||||
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
from tempfile import TemporaryDirectory
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from tempfile import TemporaryDirectory
|
||||||
|
|
||||||
import datalad.api as dl
|
import datalad.api as dl
|
||||||
|
|
||||||
dst = 'git@gin.g-node.org:/juaml/datalad-example-bids.git'
|
|
||||||
|
dst = "git@gin.g-node.org:/juaml/datalad-example-bids.git"
|
||||||
|
|
||||||
with TemporaryDirectory() as tmpdir_name:
|
with TemporaryDirectory() as tmpdir_name:
|
||||||
tmpdir = Path(tmpdir_name)
|
tmpdir = Path(tmpdir_name)
|
||||||
ds = dl.create(tmpdir) # type: ignore
|
ds = dl.create(tmpdir) # type: ignore
|
||||||
|
|
||||||
base_dir = tmpdir / 'example_bids'
|
base_dir = tmpdir / "example_bids"
|
||||||
base_dir.mkdir()
|
base_dir.mkdir()
|
||||||
|
|
||||||
for i_sub in range(1, 10):
|
for i_sub in range(1, 10):
|
||||||
t_sub = f'sub-{i_sub:02d}'
|
t_sub = f"sub-{i_sub:02d}"
|
||||||
sub_dir = base_dir / t_sub
|
sub_dir = base_dir / t_sub
|
||||||
sub_dir.mkdir()
|
sub_dir.mkdir()
|
||||||
|
|
||||||
for dname in ['anat', 'func']:
|
for dname in ["anat", "func"]:
|
||||||
(sub_dir / dname).mkdir()
|
(sub_dir / dname).mkdir()
|
||||||
|
|
||||||
fnames = [f'anat/{t_sub}_T1w.nii.gz',
|
fnames = [
|
||||||
f'func/{t_sub}_task-rest_bold.nii.gz',
|
f"anat/{t_sub}_T1w.nii.gz",
|
||||||
f'func/{t_sub}_task-rest_bold.json']
|
f"anat/{t_sub}_brain_mask.nii.gz",
|
||||||
|
f"func/{t_sub}_task-rest_bold.nii.gz",
|
||||||
|
f"func/{t_sub}_task-rest_bold.json",
|
||||||
|
f"func/{t_sub}_task-rest_brain_mask.nii.gz",
|
||||||
|
f"func/{t_sub}_task-rest_confounds_regressors.tsv",
|
||||||
|
]
|
||||||
for fname in fnames:
|
for fname in fnames:
|
||||||
with open(sub_dir / fname, 'w') as f:
|
with open(sub_dir / fname, "w") as f:
|
||||||
f.write(f'placeholder-{fname}')
|
f.write(f"placeholder-{fname}")
|
||||||
|
|
||||||
ds.save(recursive=True)
|
ds.save(recursive=True)
|
||||||
ds.siblings('add', name='gin', url=dst)
|
ds.siblings("add", name="gin", url=dst)
|
||||||
ds.push(to='gin', force='all')
|
ds.push(to="gin", force="all")
|
||||||
|
|
|
||||||
|
|
@ -1,41 +1,50 @@
|
||||||
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
# Authors: Federico Raimondo <f.raimondo@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
from tempfile import TemporaryDirectory
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
from tempfile import TemporaryDirectory
|
||||||
|
|
||||||
import datalad.api as dl
|
import datalad.api as dl
|
||||||
|
|
||||||
dst = 'git@gin.g-node.org:/juaml/datalad-example-bids-ses.git'
|
|
||||||
|
dst = "git@gin.g-node.org:/juaml/datalad-example-bids-ses.git"
|
||||||
|
|
||||||
with TemporaryDirectory() as tmpdir_name:
|
with TemporaryDirectory() as tmpdir_name:
|
||||||
tmpdir = Path(tmpdir_name)
|
tmpdir = Path(tmpdir_name)
|
||||||
ds = dl.create(tmpdir) # type: ignore
|
ds = dl.create(tmpdir) # type: ignore
|
||||||
|
|
||||||
base_dir = tmpdir / 'example_bids_ses'
|
base_dir = tmpdir / "example_bids_ses"
|
||||||
base_dir.mkdir()
|
base_dir.mkdir()
|
||||||
|
|
||||||
for i_sub in range(1, 10):
|
for i_sub in range(1, 10):
|
||||||
t_sub = f'sub-{i_sub:02d}'
|
t_sub = f"sub-{i_sub:02d}"
|
||||||
sub_dir = base_dir / t_sub
|
sub_dir = base_dir / t_sub
|
||||||
sub_dir.mkdir()
|
sub_dir.mkdir()
|
||||||
|
|
||||||
for i_ses in range(1, 4):
|
for i_ses in range(1, 4):
|
||||||
t_ses = f'ses-{i_ses:02d}'
|
t_ses = f"ses-{i_ses:02d}"
|
||||||
ses_dir = sub_dir / t_ses
|
ses_dir = sub_dir / t_ses
|
||||||
ses_dir.mkdir()
|
ses_dir.mkdir()
|
||||||
|
|
||||||
for dname in ['anat', 'func']:
|
for dname in ["anat", "func"]:
|
||||||
(ses_dir / dname).mkdir()
|
(ses_dir / dname).mkdir()
|
||||||
|
|
||||||
fnames = [f'anat/{t_sub}_{t_ses}_T1w.nii.gz']
|
fnames = [
|
||||||
|
f"anat/{t_sub}_{t_ses}_T1w.nii.gz",
|
||||||
|
f"anat/{t_sub}_{t_ses}_brain_mask.nii.gz",
|
||||||
|
]
|
||||||
if i_ses != 3: # Session 3 does not have functional data
|
if i_ses != 3: # Session 3 does not have functional data
|
||||||
fnames.extend([
|
fnames.extend(
|
||||||
f'func/{t_sub}_{t_ses}_task-rest_bold.nii.gz',
|
[
|
||||||
f'func/{t_sub}_{t_ses}_task-rest_bold.json'])
|
f"func/{t_sub}_{t_ses}_task-rest_bold.nii.gz",
|
||||||
|
f"func/{t_sub}_{t_ses}_task-rest_bold.json",
|
||||||
|
f"func/{t_sub}_{t_ses}_task-rest_brain_mask.nii.gz",
|
||||||
|
f"func/{t_sub}_{t_ses}_task-rest_confounds_regressors.tsv",
|
||||||
|
]
|
||||||
|
)
|
||||||
for fname in fnames:
|
for fname in fnames:
|
||||||
with open(ses_dir / fname, 'w') as f:
|
with open(ses_dir / fname, "w") as f:
|
||||||
f.write('placeholder-{fname}')
|
f.write("placeholder-{fname}")
|
||||||
|
|
||||||
ds.save(recursive=True)
|
ds.save(recursive=True)
|
||||||
ds.siblings('add', name='gin', url=dst)
|
ds.siblings("add", name="gin", url=dst)
|
||||||
ds.push(to='gin', force='all')
|
ds.push(to="gin", force="all")
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue