[ENH]: Enable extendable data registries #450
16 changed files with 371 additions and 72 deletions
1
docs/changes/newsfragments/450.doc
Normal file
1
docs/changes/newsfragments/450.doc
Normal file
|
|
@ -0,0 +1 @@
|
||||||
|
Add documentation on extending data registries by `Synchon Mandal`_
|
||||||
1
docs/changes/newsfragments/450.feature
Normal file
1
docs/changes/newsfragments/450.feature
Normal file
|
|
@ -0,0 +1 @@
|
||||||
|
Enable data registries to be extended by exposing :class:`.BasePipelineDataRegistry` and introducing :func:`.register_data_registry` decorator by `Synchon Mandal`_
|
||||||
70
docs/extending/data_registries.rst
Normal file
70
docs/extending/data_registries.rst
Normal file
|
|
@ -0,0 +1,70 @@
|
||||||
|
.. include:: ../links.inc
|
||||||
|
|
||||||
|
.. _data_registries:
|
||||||
|
|
||||||
|
Creating a Data Registry
|
||||||
|
========================
|
||||||
|
|
||||||
|
What is a data registry
|
||||||
|
-----------------------
|
||||||
|
|
||||||
|
Data Registry is an object which manages pipeline data (like parcellations,
|
||||||
|
coordinates and masks) based on the scope. ``junifer`` comes with the
|
||||||
|
following in-built data registries:
|
||||||
|
|
||||||
|
* :class:`.ParcellationRegistry` for parcellations
|
||||||
|
* :class:`.CoordinatesRegistry` for coordinates
|
||||||
|
* :class:`.MaskRegistry` for masks
|
||||||
|
|
||||||
|
You interact with them indirectly via :func:`.get_data`, :func:`.load_data`,
|
||||||
|
:func:`.list_data`, :func:`.register_data` and :func:`.deregister_data`. The
|
||||||
|
``kind`` parameter in them directs which data registry to interact with. These
|
||||||
|
in turn supply preprocessors and markers with their necessary data for computation.
|
||||||
|
|
||||||
|
How to make a data registry
|
||||||
|
---------------------------
|
||||||
|
|
||||||
|
Ideally you would not need to create a custom data registry if you work with fMRI data.
|
||||||
|
In case you work with other modalities like EEG, you might want to create a
|
||||||
|
data registry for montages. Here is how you would go about it:
|
||||||
|
|
||||||
|
#. Check :ref:`extending junifer <extending_extension>` on how to create a
|
||||||
|
*junifer extension* if you have not done so.
|
||||||
|
#. Create the data registry in the *extension script* like so:
|
||||||
|
|
||||||
|
.. code-block:: python
|
||||||
|
|
||||||
|
from junifer.api.decorators import register_data_registry
|
||||||
|
from junifer.data import BasePipelineDataRegistry
|
||||||
|
|
||||||
|
|
||||||
|
@register_data_registry("montage")
|
||||||
|
class MontageDataRegistry(BasePipelineDataRegistry):
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
def register(self):
|
||||||
|
pass
|
||||||
|
|
||||||
|
def deregister(self):
|
||||||
|
pass
|
||||||
|
|
||||||
|
def load(self):
|
||||||
|
pass
|
||||||
|
|
||||||
|
def get(self):
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
* :func:`.register_data_registry` registers a class with the name passed in
|
||||||
|
the argument, ``"montage"`` in this case.
|
||||||
|
* Inheriting from :class:`.BasePipelineDataRegistry` takes care of the class
|
||||||
|
acting as a data registry.
|
||||||
|
* :meth:`.BasePipelineDataRegistry.register`,
|
||||||
|
:meth:`.BasePipelineDataRegistry.deregister`,
|
||||||
|
:meth:`.BasePipelineDataRegistry.load` and
|
||||||
|
:meth:`.BasePipelineDataRegistry.get` need to be implemented
|
||||||
|
(check other registries for reference implementations).
|
||||||
|
|
||||||
|
#. Pass ``kind="montage"`` in ``*_data`` functions to
|
||||||
|
verify if your data registry is set up properly.
|
||||||
|
|
@ -5,8 +5,8 @@
|
||||||
Creating a ``junifer`` extension
|
Creating a ``junifer`` extension
|
||||||
================================
|
================================
|
||||||
|
|
||||||
``junifer`` is designed to be easily extensible. Through the use of a registry
|
``junifer`` is designed to be easily extensible. Through the use of data registries,
|
||||||
and decorators, you can easily add new functionality to ``junifer`` during
|
a component registry and decorators, you can easily add new functionality to ``junifer`` during
|
||||||
runtime. This is done by creating a new Python module and importing it before
|
runtime. This is done by creating a new Python module and importing it before
|
||||||
running ``junifer``.
|
running ``junifer``.
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -31,3 +31,4 @@ DataGrabbers, Preprocessors, Markers, etc., following the *junifer* way.
|
||||||
coordinates
|
coordinates
|
||||||
masks
|
masks
|
||||||
plugins
|
plugins
|
||||||
|
data_registries
|
||||||
|
|
|
||||||
|
|
@ -5,11 +5,19 @@
|
||||||
# Synchon Mandal <s.mandal@fz-juelich.de>
|
# Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
|
from ..data import DataDispatcher
|
||||||
from ..pipeline import PipelineComponentRegistry
|
from ..pipeline import PipelineComponentRegistry
|
||||||
from ..typing import DataGrabberLike, MarkerLike, PreprocessorLike, StorageLike
|
from ..typing import (
|
||||||
|
DataGrabberLike,
|
||||||
|
DataRegistryLike,
|
||||||
|
MarkerLike,
|
||||||
|
PreprocessorLike,
|
||||||
|
StorageLike,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
|
"register_data_registry",
|
||||||
"register_datagrabber",
|
"register_datagrabber",
|
||||||
"register_datareader",
|
"register_datareader",
|
||||||
"register_marker",
|
"register_marker",
|
||||||
|
|
@ -30,7 +38,7 @@ def register_datagrabber(klass: DataGrabberLike) -> DataGrabberLike:
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
klass: class
|
class
|
||||||
The unmodified input class.
|
The unmodified input class.
|
||||||
|
|
||||||
Notes
|
Notes
|
||||||
|
|
@ -57,7 +65,7 @@ def register_datareader(klass: type) -> type:
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
klass: class
|
class
|
||||||
The unmodified input class.
|
The unmodified input class.
|
||||||
|
|
||||||
Notes
|
Notes
|
||||||
|
|
@ -84,7 +92,7 @@ def register_preprocessor(klass: PreprocessorLike) -> PreprocessorLike:
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
klass: class
|
class
|
||||||
The unmodified input class.
|
The unmodified input class.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
@ -107,7 +115,7 @@ def register_marker(klass: MarkerLike) -> MarkerLike:
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
klass: class
|
class
|
||||||
The unmodified input class.
|
The unmodified input class.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
@ -130,7 +138,7 @@ def register_storage(klass: StorageLike) -> StorageLike:
|
||||||
|
|
||||||
Returns
|
Returns
|
||||||
-------
|
-------
|
||||||
klass: class
|
class
|
||||||
The unmodified input class.
|
The unmodified input class.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
@ -139,3 +147,40 @@ def register_storage(klass: StorageLike) -> StorageLike:
|
||||||
klass=klass,
|
klass=klass,
|
||||||
)
|
)
|
||||||
return klass
|
return klass
|
||||||
|
|
||||||
|
|
||||||
|
def register_data_registry(name: str) -> DataRegistryLike:
|
||||||
|
"""Registry registration decorator.
|
||||||
|
|
||||||
|
Registers the data registry as ``name``.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
name : str
|
||||||
|
The name of the data registry.
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
class
|
||||||
|
The unmodified input class.
|
||||||
|
|
||||||
|
"""
|
||||||
|
|
||||||
|
def decorator(klass: DataRegistryLike) -> DataRegistryLike:
|
||||||
|
"""Actual decorator.
|
||||||
|
|
||||||
|
Parameters
|
||||||
|
----------
|
||||||
|
klass : class
|
||||||
|
The class of the data registry to register.
|
||||||
|
|
||||||
|
Returns
|
||||||
|
-------
|
||||||
|
class
|
||||||
|
The unmodified input class.
|
||||||
|
|
||||||
|
"""
|
||||||
|
DataDispatcher()[name] = klass
|
||||||
|
return klass
|
||||||
|
|
||||||
|
return decorator
|
||||||
|
|
|
||||||
32
junifer/api/tests/test_decorators.py
Normal file
32
junifer/api/tests/test_decorators.py
Normal file
|
|
@ -0,0 +1,32 @@
|
||||||
|
"""Provide tests for public decorators."""
|
||||||
|
|
||||||
|
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
|
# License: AGPL
|
||||||
|
|
||||||
|
from junifer.api.decorators import register_data_registry
|
||||||
|
from junifer.data import BasePipelineDataRegistry, DataDispatcher
|
||||||
|
|
||||||
|
|
||||||
|
def test_register_data_registry() -> None:
|
||||||
|
"""Test data registry registration."""
|
||||||
|
|
||||||
|
@register_data_registry("dumb")
|
||||||
|
class DumDum(BasePipelineDataRegistry):
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
def register(self):
|
||||||
|
pass
|
||||||
|
|
||||||
|
def deregister(self):
|
||||||
|
pass
|
||||||
|
|
||||||
|
def load(self):
|
||||||
|
pass
|
||||||
|
|
||||||
|
def get(self):
|
||||||
|
pass
|
||||||
|
|
||||||
|
assert "dumb" in DataDispatcher()
|
||||||
|
_ = DataDispatcher().pop("dumb")
|
||||||
|
assert "dumb" not in DataDispatcher()
|
||||||
|
|
@ -1,5 +1,7 @@
|
||||||
__all__ = [
|
__all__ = [
|
||||||
|
"BasePipelineDataRegistry",
|
||||||
"CoordinatesRegistry",
|
"CoordinatesRegistry",
|
||||||
|
"DataDispatcher",
|
||||||
"ParcellationRegistry",
|
"ParcellationRegistry",
|
||||||
"MaskRegistry",
|
"MaskRegistry",
|
||||||
"get_data",
|
"get_data",
|
||||||
|
|
@ -12,11 +14,13 @@ __all__ = [
|
||||||
"utils",
|
"utils",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
from .pipeline_data_registry_base import BasePipelineDataRegistry
|
||||||
from .coordinates import CoordinatesRegistry
|
from .coordinates import CoordinatesRegistry
|
||||||
from .parcellations import ParcellationRegistry
|
from .parcellations import ParcellationRegistry
|
||||||
from .masks import MaskRegistry
|
from .masks import MaskRegistry
|
||||||
|
|
||||||
from ._dispatch import (
|
from ._dispatch import (
|
||||||
|
DataDispatcher,
|
||||||
get_data,
|
get_data,
|
||||||
list_data,
|
list_data,
|
||||||
load_data,
|
load_data,
|
||||||
|
|
|
||||||
|
|
@ -3,6 +3,7 @@
|
||||||
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
|
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
|
from collections.abc import Iterator, MutableMapping
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import (
|
from typing import (
|
||||||
TYPE_CHECKING,
|
TYPE_CHECKING,
|
||||||
|
|
@ -18,13 +19,14 @@ from ..utils import raise_error
|
||||||
from .coordinates import CoordinatesRegistry
|
from .coordinates import CoordinatesRegistry
|
||||||
from .masks import MaskRegistry
|
from .masks import MaskRegistry
|
||||||
from .parcellations import ParcellationRegistry
|
from .parcellations import ParcellationRegistry
|
||||||
|
from .pipeline_data_registry_base import BasePipelineDataRegistry
|
||||||
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from nibabel.nifti1 import Nifti1Image
|
from nibabel.nifti1 import Nifti1Image
|
||||||
|
|
||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
|
"DataDispatcher",
|
||||||
"deregister_data",
|
"deregister_data",
|
||||||
"get_data",
|
"get_data",
|
||||||
"list_data",
|
"list_data",
|
||||||
|
|
@ -33,6 +35,77 @@ __all__ = [
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
class DataDispatcher(MutableMapping):
|
||||||
|
"""Class for helping dynamic data dispatch."""
|
||||||
|
|
||||||
|
_instance = None
|
||||||
|
|
||||||
|
def __new__(cls):
|
||||||
|
# Make class singleton
|
||||||
|
if cls._instance is None:
|
||||||
|
cls._instance = super().__new__(cls)
|
||||||
|
# Set registries
|
||||||
|
cls._registries: dict[str, type[BasePipelineDataRegistry]] = {}
|
||||||
|
cls._builtin: dict[str, type[BasePipelineDataRegistry]] = {}
|
||||||
|
cls._external: dict[str, type[BasePipelineDataRegistry]] = {}
|
||||||
|
cls._builtin.update(
|
||||||
|
{
|
||||||
|
"coordinates": CoordinatesRegistry,
|
||||||
|
"parcellation": ParcellationRegistry,
|
||||||
|
"mask": MaskRegistry,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
cls._registries.update(cls._builtin)
|
||||||
|
return cls._instance
|
||||||
|
|
||||||
|
def __getitem__(self, key: str) -> type[BasePipelineDataRegistry]:
|
||||||
|
return self._registries[key]
|
||||||
|
|
||||||
|
def __iter__(self) -> Iterator[str]:
|
||||||
|
return iter(self._registries)
|
||||||
|
|
||||||
|
def __len__(self) -> int:
|
||||||
|
return len(self._registries)
|
||||||
|
|
||||||
|
def __delitem__(self, key: str) -> None:
|
||||||
|
# Internal check
|
||||||
|
if key in self._builtin:
|
||||||
|
raise_error(f"Cannot delete in-built key: {key}")
|
||||||
|
# Non-existing key
|
||||||
|
if key not in self._external:
|
||||||
|
raise_error(klass=KeyError, msg=key)
|
||||||
|
# Update external
|
||||||
|
_ = self._external.pop(key)
|
||||||
|
# Update global
|
||||||
|
_ = self._registries.pop(key)
|
||||||
|
|
||||||
|
def __setitem__(
|
||||||
|
self, key: str, value: type[BasePipelineDataRegistry]
|
||||||
|
) -> None:
|
||||||
|
# Internal check
|
||||||
|
if key in self._builtin:
|
||||||
|
raise_error(f"Cannot set value for in-built key: {key}")
|
||||||
|
# Value type check
|
||||||
|
if not issubclass(value, BasePipelineDataRegistry):
|
||||||
|
raise_error(f"Invalid value type: {type(value)}")
|
||||||
|
# Update external
|
||||||
|
self._external[key] = value
|
||||||
|
# Update global
|
||||||
|
self._registries[key] = value
|
||||||
|
|
||||||
|
def popitem():
|
||||||
|
"""Not implemented."""
|
||||||
|
pass
|
||||||
|
|
||||||
|
def clear(self):
|
||||||
|
"""Not implemented."""
|
||||||
|
pass
|
||||||
|
|
||||||
|
def setdefault(self, key: str, value=None):
|
||||||
|
"""Not implemented."""
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
def get_data(
|
def get_data(
|
||||||
kind: str,
|
kind: str,
|
||||||
names: Union[
|
names: Union[
|
||||||
|
|
@ -76,27 +149,16 @@ def get_data(
|
||||||
If ``kind`` is invalid value.
|
If ``kind`` is invalid value.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
try:
|
||||||
if kind == "coordinates":
|
registry = DataDispatcher()[kind]
|
||||||
return CoordinatesRegistry().get(
|
except KeyError:
|
||||||
coords=names,
|
|
||||||
target_data=target_data,
|
|
||||||
extra_input=extra_input,
|
|
||||||
)
|
|
||||||
elif kind == "parcellation":
|
|
||||||
return ParcellationRegistry().get(
|
|
||||||
parcellations=names,
|
|
||||||
target_data=target_data,
|
|
||||||
extra_input=extra_input,
|
|
||||||
)
|
|
||||||
elif kind == "mask":
|
|
||||||
return MaskRegistry().get(
|
|
||||||
masks=names,
|
|
||||||
target_data=target_data,
|
|
||||||
extra_input=extra_input,
|
|
||||||
)
|
|
||||||
else: # pragma: no cover
|
|
||||||
raise_error(f"Unknown data kind: {kind}")
|
raise_error(f"Unknown data kind: {kind}")
|
||||||
|
else:
|
||||||
|
return registry().get(
|
||||||
|
names,
|
||||||
|
target_data=target_data,
|
||||||
|
extra_input=extra_input,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def list_data(kind: str) -> list[str]:
|
def list_data(kind: str) -> list[str]:
|
||||||
|
|
@ -119,14 +181,12 @@ def list_data(kind: str) -> list[str]:
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
if kind == "coordinates":
|
try:
|
||||||
return CoordinatesRegistry().list
|
registry = DataDispatcher()[kind]
|
||||||
elif kind == "parcellation":
|
except KeyError:
|
||||||
return ParcellationRegistry().list
|
|
||||||
elif kind == "mask":
|
|
||||||
return MaskRegistry().list
|
|
||||||
else: # pragma: no cover
|
|
||||||
raise_error(f"Unknown data kind: {kind}")
|
raise_error(f"Unknown data kind: {kind}")
|
||||||
|
else:
|
||||||
|
return registry().list
|
||||||
|
|
||||||
|
|
||||||
def load_data(
|
def load_data(
|
||||||
|
|
@ -165,15 +225,12 @@ def load_data(
|
||||||
If ``kind`` is invalid value.
|
If ``kind`` is invalid value.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
try:
|
||||||
if kind == "coordinates":
|
registry = DataDispatcher()[kind]
|
||||||
return CoordinatesRegistry().load(name=name)
|
except KeyError:
|
||||||
elif kind == "parcellation":
|
|
||||||
return ParcellationRegistry().load(name=name, **kwargs)
|
|
||||||
elif kind == "mask":
|
|
||||||
return MaskRegistry().load(name=name, **kwargs)
|
|
||||||
else: # pragma: no cover
|
|
||||||
raise_error(f"Unknown data kind: {kind}")
|
raise_error(f"Unknown data kind: {kind}")
|
||||||
|
else:
|
||||||
|
return registry().load(name, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
def register_data(
|
def register_data(
|
||||||
|
|
@ -205,20 +262,14 @@ def register_data(
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
if kind == "coordinates":
|
try:
|
||||||
return CoordinatesRegistry().register(
|
registry = DataDispatcher()[kind]
|
||||||
name=name, space=space, overwrite=overwrite, **kwargs
|
except KeyError:
|
||||||
)
|
|
||||||
elif kind == "parcellation":
|
|
||||||
return ParcellationRegistry().register(
|
|
||||||
name=name, space=space, overwrite=overwrite, **kwargs
|
|
||||||
)
|
|
||||||
elif kind == "mask":
|
|
||||||
return MaskRegistry().register(
|
|
||||||
name=name, space=space, overwrite=overwrite, **kwargs
|
|
||||||
)
|
|
||||||
else: # pragma: no cover
|
|
||||||
raise_error(f"Unknown data kind: {kind}")
|
raise_error(f"Unknown data kind: {kind}")
|
||||||
|
else:
|
||||||
|
return registry().register(
|
||||||
|
name=name, space=space, overwrite=overwrite, **kwargs
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def deregister_data(kind: str, name: str) -> None:
|
def deregister_data(kind: str, name: str) -> None:
|
||||||
|
|
@ -238,11 +289,9 @@ def deregister_data(kind: str, name: str) -> None:
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
if kind == "coordinates":
|
try:
|
||||||
return CoordinatesRegistry().deregister(name=name)
|
registry = DataDispatcher()[kind]
|
||||||
elif kind == "parcellation":
|
except KeyError:
|
||||||
return ParcellationRegistry().deregister(name=name)
|
|
||||||
elif kind == "mask":
|
|
||||||
return MaskRegistry().deregister(name=name)
|
|
||||||
else: # pragma: no cover
|
|
||||||
raise_error(f"Unknown data kind: {kind}")
|
raise_error(f"Unknown data kind: {kind}")
|
||||||
|
else:
|
||||||
|
return registry().deregister(name=name)
|
||||||
|
|
|
||||||
|
|
@ -13,7 +13,6 @@ from junifer_data import get
|
||||||
from numpy.typing import ArrayLike
|
from numpy.typing import ArrayLike
|
||||||
|
|
||||||
from ...utils import logger, raise_error
|
from ...utils import logger, raise_error
|
||||||
from ...utils.singleton import Singleton
|
|
||||||
from ..pipeline_data_registry_base import BasePipelineDataRegistry
|
from ..pipeline_data_registry_base import BasePipelineDataRegistry
|
||||||
from ..utils import JUNIFER_DATA_PARAMS, get_dataset_path, get_native_warper
|
from ..utils import JUNIFER_DATA_PARAMS, get_dataset_path, get_native_warper
|
||||||
from ._ants_coordinates_warper import ANTsCoordinatesWarper
|
from ._ants_coordinates_warper import ANTsCoordinatesWarper
|
||||||
|
|
@ -23,7 +22,7 @@ from ._fsl_coordinates_warper import FSLCoordinatesWarper
|
||||||
__all__ = ["CoordinatesRegistry"]
|
__all__ = ["CoordinatesRegistry"]
|
||||||
|
|
||||||
|
|
||||||
class CoordinatesRegistry(BasePipelineDataRegistry, metaclass=Singleton):
|
class CoordinatesRegistry(BasePipelineDataRegistry):
|
||||||
"""Class for coordinates data registry.
|
"""Class for coordinates data registry.
|
||||||
|
|
||||||
This class is a singleton and is used for managing available coordinates
|
This class is a singleton and is used for managing available coordinates
|
||||||
|
|
|
||||||
|
|
@ -24,7 +24,6 @@ from nilearn.masking import (
|
||||||
)
|
)
|
||||||
|
|
||||||
from ...utils import logger, raise_error
|
from ...utils import logger, raise_error
|
||||||
from ...utils.singleton import Singleton
|
|
||||||
from ..pipeline_data_registry_base import BasePipelineDataRegistry
|
from ..pipeline_data_registry_base import BasePipelineDataRegistry
|
||||||
from ..template_spaces import get_template
|
from ..template_spaces import get_template
|
||||||
from ..utils import (
|
from ..utils import (
|
||||||
|
|
@ -216,7 +215,7 @@ def compute_brain_mask(
|
||||||
return nimg.new_img_like(target_data["data"], mask) # type: ignore
|
return nimg.new_img_like(target_data["data"], mask) # type: ignore
|
||||||
|
|
||||||
|
|
||||||
class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
|
class MaskRegistry(BasePipelineDataRegistry):
|
||||||
"""Class for mask data registry.
|
"""Class for mask data registry.
|
||||||
|
|
||||||
This class is a singleton and is used for managing available mask
|
This class is a singleton and is used for managing available mask
|
||||||
|
|
|
||||||
|
|
@ -16,7 +16,6 @@ import pandas as pd
|
||||||
from junifer_data import get
|
from junifer_data import get
|
||||||
|
|
||||||
from ...utils import logger, raise_error, warn_with_log
|
from ...utils import logger, raise_error, warn_with_log
|
||||||
from ...utils.singleton import Singleton
|
|
||||||
from ..pipeline_data_registry_base import BasePipelineDataRegistry
|
from ..pipeline_data_registry_base import BasePipelineDataRegistry
|
||||||
from ..utils import (
|
from ..utils import (
|
||||||
JUNIFER_DATA_PARAMS,
|
JUNIFER_DATA_PARAMS,
|
||||||
|
|
@ -38,7 +37,7 @@ __all__ = [
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton):
|
class ParcellationRegistry(BasePipelineDataRegistry):
|
||||||
"""Class for parcellation data registry.
|
"""Class for parcellation data registry.
|
||||||
|
|
||||||
This class is a singleton and is used for managing available parcellation
|
This class is a singleton and is used for managing available parcellation
|
||||||
|
|
|
||||||
|
|
@ -14,6 +14,7 @@ from nilearn.image import new_img_like, resample_to_img
|
||||||
from numpy.testing import assert_array_almost_equal, assert_array_equal
|
from numpy.testing import assert_array_almost_equal, assert_array_equal
|
||||||
|
|
||||||
from junifer.data import (
|
from junifer.data import (
|
||||||
|
deregister_data,
|
||||||
get_data,
|
get_data,
|
||||||
list_data,
|
list_data,
|
||||||
load_data,
|
load_data,
|
||||||
|
|
@ -1245,3 +1246,9 @@ def test_get_multi_different_space() -> None:
|
||||||
],
|
],
|
||||||
target_data=element_data["VBM_GM"],
|
target_data=element_data["VBM_GM"],
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_deregister() -> None:
|
||||||
|
"""Test parcellation deregistration."""
|
||||||
|
deregister_data(kind="parcellation", name="testparc_3")
|
||||||
|
assert "testparc_3" not in list_data(kind="parcellation")
|
||||||
|
|
|
||||||
87
junifer/data/tests/test_dispatch.py
Normal file
87
junifer/data/tests/test_dispatch.py
Normal file
|
|
@ -0,0 +1,87 @@
|
||||||
|
"""Provide tests for data dispatching."""
|
||||||
|
|
||||||
|
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
|
# License: AGPL
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from junifer.data import (
|
||||||
|
BasePipelineDataRegistry,
|
||||||
|
deregister_data,
|
||||||
|
get_data,
|
||||||
|
list_data,
|
||||||
|
load_data,
|
||||||
|
register_data,
|
||||||
|
)
|
||||||
|
from junifer.data._dispatch import DataDispatcher
|
||||||
|
|
||||||
|
|
||||||
|
def test_dispatcher_addition_errors() -> None:
|
||||||
|
"""Test registry addition errors."""
|
||||||
|
with pytest.raises(ValueError, match="Cannot set"):
|
||||||
|
DataDispatcher()["mask"] = dict
|
||||||
|
|
||||||
|
with pytest.raises(ValueError, match="Invalid"):
|
||||||
|
DataDispatcher()["masks"] = dict
|
||||||
|
|
||||||
|
|
||||||
|
def test_dispatcher_removal_errors() -> None:
|
||||||
|
"""Test registry removal errors."""
|
||||||
|
with pytest.raises(ValueError, match="Cannot delete"):
|
||||||
|
_ = DataDispatcher().pop("mask")
|
||||||
|
|
||||||
|
with pytest.raises(KeyError, match="masks"):
|
||||||
|
del DataDispatcher()["masks"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_dispatcher() -> None:
|
||||||
|
"""Test registry addition and removal."""
|
||||||
|
|
||||||
|
class DumDum(BasePipelineDataRegistry):
|
||||||
|
def register():
|
||||||
|
pass
|
||||||
|
|
||||||
|
def deregister():
|
||||||
|
pass
|
||||||
|
|
||||||
|
def load():
|
||||||
|
pass
|
||||||
|
|
||||||
|
def get():
|
||||||
|
pass
|
||||||
|
|
||||||
|
DataDispatcher().update({"masks": DumDum})
|
||||||
|
assert "masks" in DataDispatcher()
|
||||||
|
|
||||||
|
_ = DataDispatcher().pop("masks")
|
||||||
|
assert "masks" not in DataDispatcher()
|
||||||
|
|
||||||
|
|
||||||
|
def test_get_data_error() -> None:
|
||||||
|
"""Test error for get_data()."""
|
||||||
|
with pytest.raises(ValueError, match="Unknown data kind"):
|
||||||
|
get_data(kind="planet", names="neptune", target_data={})
|
||||||
|
|
||||||
|
|
||||||
|
def test_list_data_error() -> None:
|
||||||
|
"""Test error for list_data()."""
|
||||||
|
with pytest.raises(ValueError, match="Unknown data kind"):
|
||||||
|
list_data(kind="planet")
|
||||||
|
|
||||||
|
|
||||||
|
def test_load_data_error() -> None:
|
||||||
|
"""Test error for load_data()."""
|
||||||
|
with pytest.raises(ValueError, match="Unknown data kind"):
|
||||||
|
load_data(kind="planet", name="neptune")
|
||||||
|
|
||||||
|
|
||||||
|
def test_register_data_error() -> None:
|
||||||
|
"""Test error for register_data()."""
|
||||||
|
with pytest.raises(ValueError, match="Unknown data kind"):
|
||||||
|
register_data(kind="planet", name="neptune", space="milkyway")
|
||||||
|
|
||||||
|
|
||||||
|
def test_deregister_data_error() -> None:
|
||||||
|
"""Test error for deregister_data()."""
|
||||||
|
with pytest.raises(ValueError, match="Unknown data kind"):
|
||||||
|
deregister_data(kind="planet", name="neptune")
|
||||||
|
|
@ -1,5 +1,6 @@
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"DataGrabberLike",
|
"DataGrabberLike",
|
||||||
|
"DataRegistryLike",
|
||||||
"PreprocessorLike",
|
"PreprocessorLike",
|
||||||
"MarkerLike",
|
"MarkerLike",
|
||||||
"StorageLike",
|
"StorageLike",
|
||||||
|
|
@ -16,6 +17,7 @@ __all__ = [
|
||||||
|
|
||||||
from ._typing import (
|
from ._typing import (
|
||||||
DataGrabberLike,
|
DataGrabberLike,
|
||||||
|
DataRegistryLike,
|
||||||
PreprocessorLike,
|
PreprocessorLike,
|
||||||
MarkerLike,
|
MarkerLike,
|
||||||
StorageLike,
|
StorageLike,
|
||||||
|
|
|
||||||
|
|
@ -11,6 +11,7 @@ from typing import (
|
||||||
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
|
from ..data import BasePipelineDataRegistry
|
||||||
from ..datagrabber import BaseDataGrabber
|
from ..datagrabber import BaseDataGrabber
|
||||||
from ..datareader import DefaultDataReader
|
from ..datareader import DefaultDataReader
|
||||||
from ..markers import BaseMarker
|
from ..markers import BaseMarker
|
||||||
|
|
@ -23,6 +24,7 @@ __all__ = [
|
||||||
"ConfigVal",
|
"ConfigVal",
|
||||||
"DataGrabberLike",
|
"DataGrabberLike",
|
||||||
"DataGrabberPatterns",
|
"DataGrabberPatterns",
|
||||||
|
"DataRegistryLike",
|
||||||
"Dependencies",
|
"Dependencies",
|
||||||
"Element",
|
"Element",
|
||||||
"Elements",
|
"Elements",
|
||||||
|
|
@ -35,6 +37,7 @@ __all__ = [
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
DataRegistryLike = type["BasePipelineDataRegistry"]
|
||||||
DataGrabberLike = type["BaseDataGrabber"]
|
DataGrabberLike = type["BaseDataGrabber"]
|
||||||
PreprocessorLike = type["BasePreprocessor"]
|
PreprocessorLike = type["BasePreprocessor"]
|
||||||
MarkerLike = type["BaseMarker"]
|
MarkerLike = type["BaseMarker"]
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue