[ENH]: Raise/warn if element is not ran #319

Merged
synchon merged 8 commits from enh/raise-invalid-element-selector into main 2025-01-08 17:29:31 +00:00
11 changed files with 186 additions and 65 deletions

View file

@ -0,0 +1 @@
Raise error when partial or complete element selectors are invalid when running the pipeline by `Synchon Mandal`_

View file

@ -20,7 +20,13 @@ from ..pipeline import (
)
from ..preprocess import BasePreprocessor
from ..storage import BaseFeatureStorage
from ..typing import DataGrabberLike, MarkerLike, PreprocessorLike, StorageLike
from ..typing import (
DataGrabberLike,
Elements,
MarkerLike,
PreprocessorLike,
StorageLike,
)
from ..utils import logger, raise_error, warn_with_log, yaml
@ -121,7 +127,7 @@ def run(
markers: list[dict],
storage: dict,
preprocessors: Optional[list[dict]] = None,
elements: Optional[list[tuple[str, ...]]] = None,
elements: Optional[Elements] = None,
) -> None:
"""Run the pipeline on the selected element.
@ -147,7 +153,7 @@ def run(
List of preprocessors to use. Each preprocessor is a dict with at
least a key ``kind`` specifying the preprocessor to use. All other keys
are passed to the preprocessor constructor (default None).
elements : list of tuple or None, optional
elements : list or None, optional
Element(s) to process. Will be used to index the DataGrabber
(default None).
@ -155,6 +161,8 @@ def run(
------
ValueError
If ``workdir.cleanup=False`` when ``len(elements) > 1``.
RuntimeError
If invalid element selectors are found.
"""
# Conditional to handle workdir config
@ -208,10 +216,22 @@ def run(
# Fit elements
with datagrabber_object:
if elements is not None:
for t_element in datagrabber_object.filter(
elements # type: ignore
):
# Keep track of valid selectors
valid_elements = []
for t_element in datagrabber_object.filter(elements):
valid_elements.append(t_element)
mc.fit(datagrabber_object[t_element])
# Compute invalid selectors
invalid_elements = set(elements) - set(valid_elements)
# Report if invalid selectors are found
if invalid_elements:
raise_error(
msg=(
"The following element selectors are invalid:\n"
f"{invalid_elements}"
),
klass=RuntimeError,
)
else:
for t_element in datagrabber_object:
mc.fit(datagrabber_object[t_element])
@ -243,7 +263,7 @@ def queue(
kind: str,
jobname: str = "junifer_job",
fraimondo commented 2025-01-02 13:24:03 +00:00 (Migrated from github.com)

This logic won't work in case element is part of a tuple of elements.

The idea is that we can have subject/session/task/etc as the element but only pass the "subject" part an the filter function will return all of the elements for that subject.

I think that the logic should be to check if all items in the element array yielded at least one valid element.

This logic won't work in case element is part of a tuple of elements. The idea is that we can have subject/session/task/etc as the element but only pass the "subject" part an the filter function will return all of the elements for that subject. I think that the logic should be to check if all items in the `element` array yielded at least one valid element.
synchon commented 2025-01-02 15:34:04 +00:00 (Migrated from github.com)

Let me try to make it concrete and maybe it'll be easier to understand:

For this DataGrabber:

...
...
...
datagrabber:
  kind: DMCC13Benchmark
  types:
    - BOLD 
...
...
...

junifer list-elements <yaml> will give you:

...
...
...
sub-f9057kp,ses-wave1pro,Rest,AP,1
sub-f9057kp,ses-wave1pro,Rest,PA,2
sub-f9057kp,ses-wave1rea,Rest,AP,1
sub-f9057kp,ses-wave1rea,Rest,PA,2

Now if I do:

junifer run <yaml> --element sub-f9057kp

will run all "four" elements of the "subject".

Now, if I instead do:

junifer run <yaml> --element sub-f9057k

it'd just "exit" without the changes in this PR, but give you this:

...
...
...
RuntimeError: The following element selectors are invalid:
{'sub-f9057k'}

with the current PR. Also, doing this:

junifer run <yaml> --element sub-f9057kp,ses-wave1pr

would give you:

...
...
...
RuntimeError: The following element selectors are invalid:
{('sub-f9057kp', 'ses-wave1pr')}

Is this what you meant or something else?

Let me try to make it concrete and maybe it'll be easier to understand: For this DataGrabber: ```yaml ... ... ... datagrabber: kind: DMCC13Benchmark types: - BOLD ... ... ... ``` `junifer list-elements <yaml>` will give you: ```console ... ... ... sub-f9057kp,ses-wave1pro,Rest,AP,1 sub-f9057kp,ses-wave1pro,Rest,PA,2 sub-f9057kp,ses-wave1rea,Rest,AP,1 sub-f9057kp,ses-wave1rea,Rest,PA,2 ``` Now if I do: ```console junifer run <yaml> --element sub-f9057kp ``` will run all "four" elements of the "subject". Now, if I instead do: ```console junifer run <yaml> --element sub-f9057k ``` it'd just "exit" without the changes in this PR, but give you this: ```console ... ... ... RuntimeError: The following element selectors are invalid: {'sub-f9057k'} ``` with the current PR. Also, doing this: ```console junifer run <yaml> --element sub-f9057kp,ses-wave1pr ``` would give you: ```console ... ... ... RuntimeError: The following element selectors are invalid: {('sub-f9057kp', 'ses-wave1pr')} ``` Is this what you meant or something else?
fraimondo commented 2025-01-04 08:26:55 +00:00 (Migrated from github.com)

That's what I meant. I'm kind of rusty now, but I had the impression that the PR would not account for that logic.

That's what I meant. I'm kind of rusty now, but I had the impression that the PR would not account for that logic.
synchon commented 2025-01-06 10:34:47 +00:00 (Migrated from github.com)

Resolving it then.

Resolving it then.
overwrite: bool = False,
elements: Optional[list[tuple[str, ...]]] = None,
elements: Optional[Elements] = None,
**kwargs: Union[str, int, bool, dict, tuple, list],
) -> None:
"""Queue a job to be executed later.
@ -258,7 +278,7 @@ def queue(
The name of the job (default "junifer_job").
overwrite : bool, optional
Whether to overwrite if job directory already exists (default False).
elements : list of tuple or None, optional
elements : list or None, optional
Element(s) to process. Will be used to index the DataGrabber
(default None).
**kwargs : dict
@ -341,7 +361,7 @@ def queue(
elements = dg.get_elements()
# Listify elements
if not isinstance(elements, list):
elements: list[Union[str, tuple]] = [elements]
elements: Elements = [elements]
# Check job queueing system
adapter = None
@ -406,7 +426,7 @@ def reset(config: dict) -> None:
def list_elements(
datagrabber: dict,
elements: Optional[list[tuple[str, ...]]] = None,
elements: Optional[Elements] = None,
) -> str:
"""List elements of the datagrabber filtered using `elements`.
@ -416,7 +436,7 @@ def list_elements(
DataGrabber to index. Must have a key ``kind`` with the kind of
DataGrabber to use. All other keys are passed to the DataGrabber
constructor.
elements : list of tuple or None, optional
elements : list or None, optional
Element(s) to filter using. Will be used to index the DataGrabber
(default None).

View file

@ -6,8 +6,9 @@
import shutil
import textwrap
from pathlib import Path
from typing import Optional, Union
from typing import Optional
from ...typing import Elements
from ...utils import logger, make_executable, raise_error, run_ext_cmd
from .queue_context_adapter import QueueContextAdapter
@ -63,7 +64,7 @@ class GnuParallelLocalAdapter(QueueContextAdapter):
job_name: str,
job_dir: Path,
yaml_config_path: Path,
elements: list[Union[str, tuple]],
elements: Elements,
pre_run: Optional[str] = None,
pre_collect: Optional[str] = None,
env: Optional[dict[str, str]] = None,

View file

@ -6,8 +6,9 @@
import shutil
import textwrap
from pathlib import Path
from typing import Optional, Union
from typing import Optional
from ...typing import Elements
from ...utils import logger, make_executable, raise_error, run_ext_cmd
from .queue_context_adapter import QueueContextAdapter
@ -81,7 +82,7 @@ class HTCondorAdapter(QueueContextAdapter):
job_name: str,
job_dir: Path,
yaml_config_path: Path,
elements: list[Union[str, tuple]],
elements: Elements,
pre_run: Optional[str] = None,
pre_collect: Optional[str] = None,
env: Optional[dict[str, str]] = None,

View file

@ -6,16 +6,19 @@
# License: AGPL
import logging
from contextlib import AbstractContextManager, nullcontext
from pathlib import Path
from typing import Optional, Union
from typing import Any, Optional, Union
import pytest
from nibabel.filebasedimages import ImageFileError
from ruamel.yaml import YAML
import junifer.testing.registry # noqa: F401
from junifer.api import collect, list_elements, queue, reset, run
from junifer.datagrabber.base import BaseDataGrabber
from junifer.pipeline import PipelineComponentRegistry
from junifer.typing import Elements
# Configure YAML class
@ -25,12 +28,37 @@ yaml.allow_unicode = True
yaml.indent(mapping=2, sequence=4, offset=2)
# Kept for parametrizing
_datagrabber = {
"kind": "PartlyCloudyTestingDataGrabber",
}
_bids_ses_datagrabber = {
"kind": "PatternDataladDataGrabber",
"uri": "https://gin.g-node.org/juaml/datalad-example-bids-ses",
"types": ["T1w", "BOLD"],
"patterns": {
"T1w": {
"pattern": (
"{subject}/{session}/anat/{subject}_{session}_T1w.nii.gz"
),
"space": "MNI152NLin6Asym",
},
"BOLD": {
"pattern": (
"{subject}/{session}/func/{subject}_{session}_task-rest_bold.nii.gz"
),
"space": "MNI152NLin6Asym",
},
},
"replacements": ["subject", "session"],
"rootdir": "example_bids_ses",
}
@pytest.fixture
def datagrabber() -> dict[str, str]:
"""Return a datagrabber as a dictionary."""
return {
"kind": "PartlyCloudyTestingDataGrabber",
}
return _datagrabber.copy()
@pytest.fixture
@ -60,11 +88,48 @@ def storage() -> dict[str, str]:
}
@pytest.mark.parametrize(
"datagrabber, element, expect",
[
(
_datagrabber,
[("sub-01",)],
pytest.raises(RuntimeError, match="element selectors are invalid"),
),
(
_datagrabber,
["sub-01"],
nullcontext(),
),
(
_bids_ses_datagrabber,
["sub-01"],
pytest.raises(ImageFileError, match="is not a gzip file"),
),
(
_bids_ses_datagrabber,
[("sub-01", "ses-01")],
pytest.raises(ImageFileError, match="is not a gzip file"),
),
(
_bids_ses_datagrabber,
[("sub-01", "ses-100")],
pytest.raises(RuntimeError, match="element selectors are invalid"),
),
(
_bids_ses_datagrabber,
[("sub-100", "ses-01")],
pytest.raises(RuntimeError, match="element selectors are invalid"),
),
],
)
def test_run_single_element(
tmp_path: Path,
datagrabber: dict[str, str],
datagrabber: dict[str, Any],
markers: list[dict[str, str]],
storage: dict[str, str],
element: Elements,
expect: AbstractContextManager,
) -> None:
"""Test run function with single element.
@ -78,21 +143,26 @@ def test_run_single_element(
Testing markers as list of dictionary.
storage : dict
Testing storage as dictionary.
element : list of str or tuple
The parametrized element.
expect : typing.ContextManager
The parametrized ContextManager object.
"""
# Set storage
storage["uri"] = str((tmp_path / "out.sqlite").resolve())
# Run operations
run(
workdir=tmp_path,
datagrabber=datagrabber,
markers=markers,
storage=storage,
elements=[("sub-01",)],
)
# Check files
files = list(tmp_path.glob("*.sqlite"))
assert len(files) == 1
with expect:
run(
workdir=tmp_path,
datagrabber=datagrabber,
markers=markers,
storage=storage,
elements=element,
)
# Check files
files = list(tmp_path.glob("*.sqlite"))
assert len(files) == 1
def test_run_single_element_with_preprocessing(
@ -128,18 +198,30 @@ def test_run_single_element_with_preprocessing(
"kind": "fMRIPrepConfoundRemover",
}
],
elements=[("sub-01",)],
elements=["sub-01"],
)
# Check files
files = list(tmp_path.glob("*.sqlite"))
assert len(files) == 1
@pytest.mark.parametrize(
"element, expect",
[
(
[("sub-01",), ("sub-03",)],
pytest.raises(RuntimeError, match="element selectors are invalid"),
),
(["sub-01", "sub-03"], nullcontext()),
],
)
def test_run_multi_element_multi_output(
tmp_path: Path,
datagrabber: dict[str, str],
markers: list[dict[str, str]],
storage: dict[str, str],
element: Elements,
expect: AbstractContextManager,
) -> None:
"""Test run function with multi element and multi output.
@ -153,22 +235,27 @@ def test_run_multi_element_multi_output(
Testing markers as list of dictionary.
storage : dict
Testing storage as dictionary.
element : list of str or tuple
The parametrized element.
expect : typing.ContextManager
The parametrized ContextManager object.
"""
# Set storage
storage["uri"] = str((tmp_path / "out.sqlite").resolve())
storage["single_output"] = False # type: ignore
# Run operations
run(
workdir=tmp_path,
datagrabber=datagrabber,
markers=markers,
storage=storage,
elements=[("sub-01",), ("sub-03",)],
)
# Check files
files = list(tmp_path.glob("*.sqlite"))
assert len(files) == 2
with expect:
run(
workdir=tmp_path,
datagrabber=datagrabber,
markers=markers,
storage=storage,
elements=element,
)
# Check files
files = list(tmp_path.glob("*.sqlite"))
assert len(files) == 2
def test_run_multi_element_single_output(
@ -200,7 +287,7 @@ def test_run_multi_element_single_output(
datagrabber=datagrabber,
markers=markers,
storage=storage,
elements=[("sub-01",), ("sub-03",)],
elements=["sub-01", "sub-03"],
)
# Check files
files = list(tmp_path.glob("*.sqlite"))
@ -569,7 +656,7 @@ def test_reset_run(
datagrabber=datagrabber,
markers=markers,
storage=storage,
elements=[("sub-01",)],
elements=["sub-01"],
)
# Reset operation
reset(config={"storage": storage})

View file

@ -12,6 +12,7 @@ from typing import Union
import pandas as pd
from ..typing import Elements
from ..utils import logger, raise_error, warn_with_log, yaml
@ -142,7 +143,7 @@ def parse_yaml(filepath: Union[str, Path]) -> dict: # noqa: C901
def parse_elements(
element: tuple[str, ...], config: dict
) -> Union[list[tuple[str, ...]], None]:
) -> Union[Elements, None]:
"""Parse elements from cli.
Parameters
@ -203,7 +204,7 @@ def parse_elements(
return elements
def _parse_elements_file(filepath: Path) -> list[tuple[str, ...]]:
def _parse_elements_file(filepath: Path) -> Elements:
"""Parse elements from file.
Parameters
@ -213,7 +214,7 @@ def _parse_elements_file(filepath: Path) -> list[tuple[str, ...]]:
Returns
-------
list of tuple of str
list
The element(s) as list.
"""
@ -227,5 +228,8 @@ def _parse_elements_file(filepath: Path) -> list[tuple[str, ...]]:
)
# Remove trailing whitespace in cell entries
csv_df_trimmed = csv_df.apply(lambda x: x.str.strip())
# Convert to list of tuple of str
return list(map(tuple, csv_df_trimmed.to_numpy()))
# Convert to list of tuple of str if more than one column else flatten
if len(csv_df_trimmed.columns) == 1:
return csv_df_trimmed.to_numpy().flatten().tolist()
else:
return list(map(tuple, csv_df_trimmed.to_numpy()))

View file

@ -11,6 +11,7 @@ from pathlib import Path
from typing import Union
from ..pipeline import UpdateMetaMixin
from ..typing import Element, Elements
from ..utils import logger, raise_error
@ -67,9 +68,7 @@ class BaseDataGrabber(ABC, UpdateMetaMixin):
"""
yield from self.get_elements()
def __getitem__(
self, element: Union[str, tuple[str, ...]]
) -> dict[str, dict]:
def __getitem__(self, element: Element) -> dict[str, dict]:
"""Enable indexing support.
Parameters
@ -137,13 +136,13 @@ class BaseDataGrabber(ABC, UpdateMetaMixin):
"""
return self._datadir
def filter(self, selection: list[Union[str, tuple[str]]]) -> Iterator:
def filter(self, selection: Elements) -> Iterator:
"""Filter elements to be grabbed.
Parameters
----------
selection : list of str or tuple
The list of partial element key values to filter using.
selection : list
The list of partial or complete element selectors to filter using.
Yields
------
@ -152,7 +151,7 @@ class BaseDataGrabber(ABC, UpdateMetaMixin):
"""
def filter_func(element: Union[str, tuple[str]]) -> bool:
def filter_func(element: Element) -> bool:
"""Filter element based on selection.
Parameters
@ -201,15 +200,14 @@ class BaseDataGrabber(ABC, UpdateMetaMixin):
) # pragma: no cover
@abstractmethod
def get_elements(self) -> list[Union[str, tuple[str]]]:
def get_elements(self) -> Elements:
"""Get elements.
Returns
-------
list
List of elements that can be grabbed. The elements can be strings,
tuples or any object that will be then used as a key to index the
DataGrabber.
List of elements that can be grabbed. The elements can be strings
or tuples of strings to index the DataGrabber.
"""
raise_error(

View file

@ -17,6 +17,7 @@ from datalad.support.exceptions import IncompleteResultsError
from datalad.support.gitrepo import GitRepo
from ..pipeline import WorkDirManager
from ..typing import Element
from ..utils import config, logger, raise_error, warn_with_log
from .base import BaseDataGrabber
@ -312,7 +313,7 @@ class DataladDataGrabber(BaseDataGrabber):
logger.debug(f"Dropping {f}")
self._dataset.drop(f, result_renderer="disabled")
def __getitem__(self, element: Union[str, tuple]) -> dict:
def __getitem__(self, element: Element) -> dict:
"""Implement single element indexing in the Datalad database.
It will first obtain the paths from the parent class and then
@ -320,11 +321,11 @@ class DataladDataGrabber(BaseDataGrabber):
Parameters
----------
element : str or tuple
element : str or tuple of str
The element to be indexed. If one string is provided, it is
assumed to be a tuple with only one item. If a tuple is provided,
each item in the tuple is the value for the replacement string
specified in "replacements".
specified in ``"replacements"``.
Returns
-------

View file

@ -13,7 +13,7 @@ from typing import Optional, Union
import numpy as np
from ..api.decorators import register_datagrabber
from ..typing import DataGrabberPatterns
from ..typing import DataGrabberPatterns, Elements
from ..utils import logger, raise_error
from .base import BaseDataGrabber
from .pattern_validation_mixin import PatternValidationMixin
@ -367,7 +367,7 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin):
"""
return self.replacements
def get_item(self, **element: str) -> dict[str, dict]:
def get_item(self, **element: dict) -> dict[str, dict]:
"""Implement single element indexing for the datagrabber.
This method constructs a real path to the requested item's data, by
@ -450,7 +450,7 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin):
return out
def get_elements(self) -> list:
def get_elements(self) -> Elements:
"""Implement fetching list of elements in the dataset.
It will use regex to search for "replacements" in the "patterns" and

View file

@ -10,6 +10,8 @@ __all__ = [
"MarkerInOutMappings",
"DataGrabberPatterns",
"ConfigVal",
"Element",
"Elements",
]
from ._typing import (
@ -24,4 +26,6 @@ from ._typing import (
MarkerInOutMappings,
DataGrabberPatterns,
ConfigVal,
Element,
Elements,
)

View file

@ -24,6 +24,8 @@ __all__ = [
"DataGrabberLike",
"DataGrabberPatterns",
"Dependencies",
"Element",
"Elements",
"ExternalDependencies",
"MarkerInOutMappings",
"MarkerLike",
@ -62,3 +64,5 @@ DataGrabberPatterns = dict[
str, Union[dict[str, str], Sequence[dict[str, str]]]
]
ConfigVal = Union[bool, int, float]
Element = Union[str, tuple[str, ...]]
Elements = Sequence[Element]