[BUG]: Fix metadata update for MultipleDataGrabber #398

Merged
synchon merged 7 commits from fix/multiple-dg-meta-update into main 2024-11-21 13:14:10 +00:00
11 changed files with 129 additions and 62 deletions

View file

@ -0,0 +1 @@
Fix metadata update for :class:`.MultipleDataGrabber` and adjust :meth:`.PatternDataGrabber.get_elements` to check ``list``-like data type values by `Fede Raimondo`_ and `Synchon Mandal`_

View file

@ -111,8 +111,12 @@ class MultipleDataGrabber(BaseDataGrabber):
# Update all the metas again # Update all the metas again
for kind in out: for kind in out:
self.update_meta(out[kind], "datagrabber") to_update = out[kind]
out[kind]["meta"]["datagrabber"]["datagrabbers"] = metas if not isinstance(to_update, list):
to_update = [to_update]
for t_kind in to_update:
self.update_meta(t_kind, "datagrabber")
t_kind["meta"]["datagrabber"]["datagrabbers"] = metas
return out return out
def __enter__(self) -> "MultipleDataGrabber": def __enter__(self) -> "MultipleDataGrabber":

View file

@ -13,6 +13,7 @@ from typing import Optional, Union
import numpy as np import numpy as np
from ..api.decorators import register_datagrabber from ..api.decorators import register_datagrabber
from ..typing import DataGrabberPatterns
from ..utils import logger, raise_error from ..utils import logger, raise_error
from .base import BaseDataGrabber from .base import BaseDataGrabber
from .pattern_validation_mixin import PatternValidationMixin from .pattern_validation_mixin import PatternValidationMixin
@ -171,7 +172,7 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin):
def __init__( def __init__(
self, self,
types: list[str], types: list[str],
patterns: dict[str, dict[str, str]], patterns: DataGrabberPatterns,
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,
@ -478,58 +479,63 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin):
t_type = self.types[t_idx] t_type = self.types[t_idx]
types_element = set() types_element = set()
# Get the pattern dict # Data type dictionary
t_pattern = self.patterns[t_type] patterns = self.patterns[t_type]
# Conditional fetch of base pattern for getting elements # Conditional for list dtype vals like Warp
pattern = None if not isinstance(patterns, list):
# Try for data type pattern patterns = [patterns]
pattern = t_pattern.get("pattern") for t_pattern in patterns:
# Try for nested data type pattern # Conditional fetch of base pattern for getting elements
if pattern is None and self.partial_pattern_ok: pattern = None
for v in t_pattern.values(): # Try for data type pattern
if isinstance(v, dict) and "pattern" in v: pattern = t_pattern.get("pattern")
pattern = v["pattern"] # Try for nested data type pattern
break 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(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)
if m is not None: if m is not None:
# Find the groups of replacements present in the pattern # Find the groups of replacements present in the
# If one replacement is not present, set it to None. # pattern. If one replacement is not present, set it
# We will take care of this in the intersection # to None. We will take care of this in the
t_element = tuple([m.group(k) for k in t_replacements]) # intersection.
if len(self.replacements) == 1: t_element = tuple([m.group(k) for k in t_replacements])
t_element = t_element[0] if len(self.replacements) == 1:
types_element.add(t_element) t_element = t_element[0]
# TODO: does this make sense as elements is always None types_element.add(t_element)
if elements is None: # TODO: does this make sense as elements is always None
elements = types_element if elements is None:
else: elements = types_element
# Do the intersection by filtering out elements in which
# the replacements are not None
if t_replacements == self.replacements:
elements.intersection(types_element)
else: else:
t_repl_idx = [ # Do the intersection by filtering out elements in which
i # the replacements are not None
for i, v in enumerate(self.replacements) if t_replacements == self.replacements:
if v in t_replacements elements.intersection(types_element)
] else:
new_elements = set() t_repl_idx = [
for t_element in elements: i
if ( for i, v in enumerate(self.replacements)
tuple(np.array(t_element)[t_repl_idx]) if v in t_replacements
in types_element ]
): new_elements = set()
new_elements.add(t_element) for t_element in elements:
elements = new_elements if (
tuple(np.array(t_element)[t_repl_idx])
in types_element
):
new_elements.add(t_element)
elements = new_elements
if elements is None: if elements is None:
elements = set() elements = set()
return list(elements) return list(elements)

View file

@ -3,8 +3,8 @@
# Authors: Synchon Mandal <s.mandal@fz-juelich.de> # Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL # License: AGPL
from typing import Union
from ..typing import DataGrabberPatterns
from ..utils import logger, raise_error, warn_with_log from ..utils import logger, raise_error, warn_with_log
@ -96,7 +96,7 @@ class PatternValidationMixin:
def _validate_replacements( def _validate_replacements(
self, self,
replacements: list[str], replacements: list[str],
patterns: dict[str, Union[dict[str, str], list[dict[str, str]]]], patterns: DataGrabberPatterns,
partial_pattern_ok: bool, partial_pattern_ok: bool,
) -> None: ) -> None:
"""Validate the replacements. """Validate the replacements.
@ -263,7 +263,7 @@ class PatternValidationMixin:
self, self,
types: list[str], types: list[str],
replacements: list[str], replacements: list[str],
patterns: dict[str, Union[dict[str, str], list[dict[str, str]]]], patterns: DataGrabberPatterns,
partial_pattern_ok: bool = False, partial_pattern_ok: bool = False,
) -> None: ) -> None:
"""Validate the patterns. """Validate the patterns.

View file

@ -14,8 +14,8 @@ from junifer.datagrabber import DataladDataGrabber
_testing_dataset = { _testing_dataset = {
"example_bids": { "example_bids": {
"uri": "https://gin.g-node.org/juaml/datalad-example-bids", "uri": "https://gin.g-node.org/juaml/datalad-example-bids",
"commit": "b87897cbe51bf0ee5514becaa5c7dd76491db5ad", "commit": "3f288c8725207ae0c9b3616e093e78cda192b570",
"id": "8fddff30-6993-420a-9d1e-b5b028c59468", "id": "582b9696-f13f-42e4-9587-b4e62aa2a8e7",
}, },
"example_bids_ses": { "example_bids_ses": {
"uri": "https://gin.g-node.org/juaml/datalad-example-bids-ses", "uri": "https://gin.g-node.org/juaml/datalad-example-bids-ses",

View file

@ -29,7 +29,7 @@ def test_MultipleDataGrabber() -> None:
dg1 = PatternDataladDataGrabber( dg1 = PatternDataladDataGrabber(
rootdir=rootdir, rootdir=rootdir,
uri=repo_uri, uri=repo_uri,
types=["T1w"], types=["T1w", "Warp"],
patterns={ patterns={
"T1w": { "T1w": {
"pattern": ( "pattern": (
@ -44,6 +44,28 @@ def test_MultipleDataGrabber() -> None:
"space": "native", "space": "native",
}, },
}, },
"Warp": [
{
"pattern": (
"{subject}/{session}/anat/"
"{subject}_{session}_from-MNI152NLin2009cAsym_to-T1w_"
"xfm.h5"
),
"src": "MNI152NLin2009cAsym",
"dst": "native",
"warper": "ants",
},
{
"pattern": (
"{subject}/{session}/anat/"
"{subject}_{session}_from-T1w_to-MNI152NLin2009cAsym_"
"xfm.h5"
),
"src": "native",
"dst": "MNI152NLin2009cAsym",
"warper": "ants",
},
],
}, },
replacements=replacements, replacements=replacements,
) )
@ -75,6 +97,7 @@ def test_MultipleDataGrabber() -> None:
types = dg.get_types() types = dg.get_types()
assert "T1w" in types assert "T1w" in types
assert "Warp" in types
assert "BOLD" in types assert "BOLD" in types
expected_subs = [ expected_subs = [
@ -90,6 +113,7 @@ def test_MultipleDataGrabber() -> None:
elem = dg[("sub-01", "ses-01")] elem = dg[("sub-01", "ses-01")]
# Check data types # Check data types
assert "T1w" in elem assert "T1w" in elem
assert "Warp" in elem
assert "BOLD" in elem assert "BOLD" in elem
# Check meta # Check meta
assert "meta" in elem["BOLD"] assert "meta" in elem["BOLD"]
@ -111,7 +135,7 @@ def test_MultipleDataGrabber_no_intersection() -> None:
dg1 = PatternDataladDataGrabber( dg1 = PatternDataladDataGrabber(
rootdir=rootdir, rootdir=rootdir,
uri=_testing_dataset["example_bids"]["uri"], uri=_testing_dataset["example_bids"]["uri"],
types=["T1w"], types=["T1w", "Warp"],
patterns={ patterns={
"T1w": { "T1w": {
"pattern": ( "pattern": (
@ -119,6 +143,28 @@ def test_MultipleDataGrabber_no_intersection() -> None:
), ),
"space": "native", "space": "native",
}, },
"Warp": [
{
"pattern": (
"{subject}/{session}/anat/"
"{subject}_{session}_from-MNI152NLin2009cAsym_to-T1w_"
"xfm.h5"
),
"src": "MNI152NLin2009cAsym",
"dst": "native",
"warper": "ants",
},
{
"pattern": (
"{subject}/{session}/anat/"
"{subject}_{session}_from-T1w_to-MNI152NLin2009cAsym_"
"xfm.h5"
),
"src": "native",
"dst": "MNI152NLin2009cAsym",
"warper": "ants",
},
],
}, },
replacements=replacements, replacements=replacements,
) )

View file

@ -15,7 +15,7 @@ from junifer.datagrabber import PatternDataladDataGrabber
_testing_dataset = { _testing_dataset = {
"example_bids": { "example_bids": {
"uri": "https://gin.g-node.org/juaml/datalad-example-bids", "uri": "https://gin.g-node.org/juaml/datalad-example-bids",
"commit": "b87897cbe51bf0ee5514becaa5c7dd76491db5ad", "commit": "3f288c8725207ae0c9b3616e093e78cda192b570",
"id": "8fddff30-6993-420a-9d1e-b5b028c59468", "id": "8fddff30-6993-420a-9d1e-b5b028c59468",
}, },
"example_bids_ses": { "example_bids_ses": {

View file

@ -8,6 +8,7 @@ __all__ = [
"ConditionalDependencies", "ConditionalDependencies",
"ExternalDependencies", "ExternalDependencies",
"MarkerInOutMappings", "MarkerInOutMappings",
"DataGrabberPatterns",
] ]
from ._typing import ( from ._typing import (
@ -20,4 +21,5 @@ from ._typing import (
ConditionalDependencies, ConditionalDependencies,
ExternalDependencies, ExternalDependencies,
MarkerInOutMappings, MarkerInOutMappings,
DataGrabberPatterns,
) )

View file

@ -28,6 +28,7 @@ __all__ = [
"ConditionalDependencies", "ConditionalDependencies",
"ExternalDependencies", "ExternalDependencies",
"MarkerInOutMappings", "MarkerInOutMappings",
"DataGrabberPatterns",
] ]
@ -56,3 +57,6 @@ ConditionalDependencies = Sequence[
] ]
ExternalDependencies = Sequence[MutableMapping[str, Union[str, Sequence[str]]]] ExternalDependencies = Sequence[MutableMapping[str, Union[str, Sequence[str]]]]
MarkerInOutMappings = MutableMapping[str, MutableMapping[str, str]] MarkerInOutMappings = MutableMapping[str, MutableMapping[str, str]]
DataGrabberPatterns = dict[
str, Union[dict[str, str], Sequence[dict[str, str]]]
]

View file

@ -26,6 +26,8 @@ with TemporaryDirectory() as tmpdir_name:
fnames = [ fnames = [
f"anat/{t_sub}_T1w.nii.gz", f"anat/{t_sub}_T1w.nii.gz",
f"anat/{t_sub}_brain_mask.nii.gz", f"anat/{t_sub}_brain_mask.nii.gz",
f"anat/{t_sub}_from-MNI152NLin2009cAsym_to-T1w_xfm.h5",
f"anat/{t_sub}_from-T1w_to-MNI152NLin2009cAsym_xfm.h5",
f"func/{t_sub}_task-rest_bold.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_bold.json",
f"func/{t_sub}_task-rest_brain_mask.nii.gz", f"func/{t_sub}_task-rest_brain_mask.nii.gz",

View file

@ -31,6 +31,8 @@ with TemporaryDirectory() as tmpdir_name:
fnames = [ fnames = [
f"anat/{t_sub}_{t_ses}_T1w.nii.gz", f"anat/{t_sub}_{t_ses}_T1w.nii.gz",
f"anat/{t_sub}_{t_ses}_brain_mask.nii.gz", f"anat/{t_sub}_{t_ses}_brain_mask.nii.gz",
f"anat/{t_sub}_{t_ses}_from-MNI152NLin2009cAsym_to-T1w_xfm.h5",
f"anat/{t_sub}_{t_ses}_from-T1w_to-MNI152NLin2009cAsym_xfm.h5",
] ]
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(