[ENH]: Adapt MultipleDataGrabber to patterns with nested data types #351

Merged
synchon merged 20 commits from fix/multiple-dg-access into main 2024-07-19 11:15:25 +00:00
22 changed files with 877 additions and 563 deletions

View 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`_

View file

@ -0,0 +1 @@
Adapt :class:`.MultipleDataGrabber` to handle "nested types" introduced in :gh:`341` by `Synchon Mandal`_

View 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`_

View file

@ -0,0 +1 @@
Integrate ``warnings`` with ``logging`` respecting filters by `Fede Raimondo`_

View file

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

View file

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

View file

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

View file

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

View file

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

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

View file

@ -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 = {
dg1 = PatternDataladDataGrabber(
rootdir=rootdir,
uri=repo_uri,
types=["T1w"],
patterns={
"T1w": { "T1w": {
"pattern": ( "pattern": (
"{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz" "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz"
), ),
"space": "native", "space": "native",
}, "mask": {
}
pattern2 = {
"BOLD": {
"pattern": ( "pattern": (
"{subject}/{session}/func/" "{subject}/{session}/anat/{subject}_{session}_"
"{subject}_{session}_task-rest_bold.nii.gz" "brain_mask.nii.gz"
), ),
"space": "MNI152NLin6Asym", "space": "native",
},
},
}, },
}
dg1 = PatternDataladDataGrabber(
rootdir=rootdir,
uri=repo_uri,
types=["T1w"],
patterns=pattern1,
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,19 +105,29 @@ 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 = {
dg1 = PatternDataladDataGrabber(
rootdir=rootdir,
uri=_testing_dataset["example_bids"]["uri"],
types=["T1w"],
patterns={
"T1w": { "T1w": {
"pattern": ( "pattern": (
"{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz" "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz"
), ),
"space": "native", "space": "native",
}, },
} },
pattern2 = { replacements=replacements,
)
dg2 = PatternDataladDataGrabber(
rootdir=rootdir,
uri=_testing_dataset["example_bids_ses"]["uri"],
types=["BOLD"],
patterns={
"BOLD": { "BOLD": {
"pattern": ( "pattern": (
"{subject}/{session}/func/" "{subject}/{session}/func/"
@ -109,20 +135,7 @@ def test_MultipleDataGrabber_no_intersection() -> None:
), ),
"space": "MNI152NLin6Asym", "space": "MNI152NLin6Asym",
}, },
} },
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, 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"] dg1 = PatternDataladDataGrabber(
rootdir = "example_bids_ses" rootdir="example_bids_ses",
replacements = ["subject", "session"] uri=_testing_dataset["example_bids"]["uri"],
pattern1 = { types=["T1w"],
patterns={
"T1w": { "T1w": {
"pattern": ( "pattern": (
"{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz" "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz"
), ),
"space": "native", "space": "native",
}, },
} },
dg1 = PatternDataladDataGrabber( replacements=["subject", "session"],
rootdir=rootdir,
uri=repo_uri1,
types=["T1w"],
patterns=pattern1,
replacements=replacements,
) )
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"] dg1 = PatternDataladDataGrabber(
pattern1 = { rootdir=rootdir,
uri=_testing_dataset["example_bids"]["uri"],
types=["T1w"],
patterns={
"T1w": { "T1w": {
"pattern": ( "pattern": (
"{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz" "{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz"
), ),
"space": "native", "space": "native",
}, },
}
pattern2 = {
"BOLD": {
"pattern": "{subject}/func/{subject}_task-rest_bold.nii.gz",
"space": "MNI152NLin6Asym",
}, },
} replacements=["subject", "session"],
dg1 = PatternDataladDataGrabber(
rootdir=rootdir,
uri=repo_uri1,
types=["T1w"],
patterns=pattern1,
replacements=replacement1,
) )
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"

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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