[ENH]: Enable extendable data registries #450

Merged
synchon merged 11 commits from feat/extendable-data-registry into main 2025-07-11 08:43:01 +00:00
16 changed files with 371 additions and 72 deletions

View file

@ -0,0 +1 @@
Add documentation on extending data registries by `Synchon Mandal`_

View file

@ -0,0 +1 @@
Enable data registries to be extended by exposing :class:`.BasePipelineDataRegistry` and introducing :func:`.register_data_registry` decorator by `Synchon Mandal`_

View 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.

View file

@ -5,8 +5,8 @@
Creating a ``junifer`` extension
================================
``junifer`` is designed to be easily extensible. Through the use of a registry
and decorators, you can easily add new functionality to ``junifer`` during
``junifer`` is designed to be easily extensible. Through the use of data registries,
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
running ``junifer``.

View file

@ -31,3 +31,4 @@ DataGrabbers, Preprocessors, Markers, etc., following the *junifer* way.
coordinates
masks
plugins
data_registries

View file

@ -5,11 +5,19 @@
# Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from ..data import DataDispatcher
from ..pipeline import PipelineComponentRegistry
from ..typing import DataGrabberLike, MarkerLike, PreprocessorLike, StorageLike
from ..typing import (
DataGrabberLike,
DataRegistryLike,
MarkerLike,
PreprocessorLike,
StorageLike,
)
__all__ = [
"register_data_registry",
"register_datagrabber",
"register_datareader",
"register_marker",
@ -30,7 +38,7 @@ def register_datagrabber(klass: DataGrabberLike) -> DataGrabberLike:
Returns
-------
klass: class
class
The unmodified input class.
Notes
@ -57,7 +65,7 @@ def register_datareader(klass: type) -> type:
Returns
-------
klass: class
class
The unmodified input class.
Notes
@ -84,7 +92,7 @@ def register_preprocessor(klass: PreprocessorLike) -> PreprocessorLike:
Returns
-------
klass: class
class
The unmodified input class.
"""
@ -107,7 +115,7 @@ def register_marker(klass: MarkerLike) -> MarkerLike:
Returns
-------
klass: class
class
The unmodified input class.
"""
@ -130,7 +138,7 @@ def register_storage(klass: StorageLike) -> StorageLike:
Returns
-------
klass: class
class
The unmodified input class.
"""
@ -139,3 +147,40 @@ def register_storage(klass: StorageLike) -> StorageLike:
klass=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

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

View file

@ -1,5 +1,7 @@
__all__ = [
"BasePipelineDataRegistry",
"CoordinatesRegistry",
"DataDispatcher",
"ParcellationRegistry",
"MaskRegistry",
"get_data",
@ -12,11 +14,13 @@ __all__ = [
"utils",
]
from .pipeline_data_registry_base import BasePipelineDataRegistry
from .coordinates import CoordinatesRegistry
from .parcellations import ParcellationRegistry
from .masks import MaskRegistry
from ._dispatch import (
DataDispatcher,
get_data,
list_data,
load_data,

View file

@ -3,6 +3,7 @@
# Authors: Synchon Mandal <s.mandal@fz-juelich.de>
# License: AGPL
from collections.abc import Iterator, MutableMapping
from pathlib import Path
from typing import (
TYPE_CHECKING,
@ -18,13 +19,14 @@ from ..utils import raise_error
from .coordinates import CoordinatesRegistry
from .masks import MaskRegistry
from .parcellations import ParcellationRegistry
from .pipeline_data_registry_base import BasePipelineDataRegistry
if TYPE_CHECKING:
from nibabel.nifti1 import Nifti1Image
__all__ = [
"DataDispatcher",
"deregister_data",
"get_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(
kind: str,
names: Union[
@ -76,27 +149,16 @@ def get_data(
If ``kind`` is invalid value.
"""
if kind == "coordinates":
return CoordinatesRegistry().get(
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
try:
registry = DataDispatcher()[kind]
except KeyError:
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]:
@ -119,14 +181,12 @@ def list_data(kind: str) -> list[str]:
"""
if kind == "coordinates":
return CoordinatesRegistry().list
elif kind == "parcellation":
return ParcellationRegistry().list
elif kind == "mask":
return MaskRegistry().list
else: # pragma: no cover
try:
registry = DataDispatcher()[kind]
except KeyError:
raise_error(f"Unknown data kind: {kind}")
else:
return registry().list
def load_data(
@ -165,15 +225,12 @@ def load_data(
If ``kind`` is invalid value.
"""
if kind == "coordinates":
return CoordinatesRegistry().load(name=name)
elif kind == "parcellation":
return ParcellationRegistry().load(name=name, **kwargs)
elif kind == "mask":
return MaskRegistry().load(name=name, **kwargs)
else: # pragma: no cover
try:
registry = DataDispatcher()[kind]
except KeyError:
raise_error(f"Unknown data kind: {kind}")
else:
return registry().load(name, **kwargs)
def register_data(
@ -205,20 +262,14 @@ def register_data(
"""
if kind == "coordinates":
return CoordinatesRegistry().register(
name=name, space=space, overwrite=overwrite, **kwargs
)
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
try:
registry = DataDispatcher()[kind]
except KeyError:
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:
@ -238,11 +289,9 @@ def deregister_data(kind: str, name: str) -> None:
"""
if kind == "coordinates":
return CoordinatesRegistry().deregister(name=name)
elif kind == "parcellation":
return ParcellationRegistry().deregister(name=name)
elif kind == "mask":
return MaskRegistry().deregister(name=name)
else: # pragma: no cover
try:
registry = DataDispatcher()[kind]
except KeyError:
raise_error(f"Unknown data kind: {kind}")
else:
return registry().deregister(name=name)

View file

@ -13,7 +13,6 @@ from junifer_data import get
from numpy.typing import ArrayLike
from ...utils import logger, raise_error
from ...utils.singleton import Singleton
from ..pipeline_data_registry_base import BasePipelineDataRegistry
from ..utils import JUNIFER_DATA_PARAMS, get_dataset_path, get_native_warper
from ._ants_coordinates_warper import ANTsCoordinatesWarper
@ -23,7 +22,7 @@ from ._fsl_coordinates_warper import FSLCoordinatesWarper
__all__ = ["CoordinatesRegistry"]
class CoordinatesRegistry(BasePipelineDataRegistry, metaclass=Singleton):
class CoordinatesRegistry(BasePipelineDataRegistry):
"""Class for coordinates data registry.
This class is a singleton and is used for managing available coordinates

View file

@ -24,7 +24,6 @@ from nilearn.masking import (
)
from ...utils import logger, raise_error
from ...utils.singleton import Singleton
from ..pipeline_data_registry_base import BasePipelineDataRegistry
from ..template_spaces import get_template
from ..utils import (
@ -216,7 +215,7 @@ def compute_brain_mask(
return nimg.new_img_like(target_data["data"], mask) # type: ignore
class MaskRegistry(BasePipelineDataRegistry, metaclass=Singleton):
class MaskRegistry(BasePipelineDataRegistry):
"""Class for mask data registry.
This class is a singleton and is used for managing available mask

View file

@ -16,7 +16,6 @@ import pandas as pd
from junifer_data import get
from ...utils import logger, raise_error, warn_with_log
from ...utils.singleton import Singleton
from ..pipeline_data_registry_base import BasePipelineDataRegistry
from ..utils import (
JUNIFER_DATA_PARAMS,
@ -38,7 +37,7 @@ __all__ = [
]
class ParcellationRegistry(BasePipelineDataRegistry, metaclass=Singleton):
class ParcellationRegistry(BasePipelineDataRegistry):
"""Class for parcellation data registry.
This class is a singleton and is used for managing available parcellation

View file

@ -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 junifer.data import (
deregister_data,
get_data,
list_data,
load_data,
@ -1245,3 +1246,9 @@ def test_get_multi_different_space() -> None:
],
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")

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

View file

@ -1,5 +1,6 @@
__all__ = [
"DataGrabberLike",
"DataRegistryLike",
"PreprocessorLike",
"MarkerLike",
"StorageLike",
@ -16,6 +17,7 @@ __all__ = [
from ._typing import (
DataGrabberLike,
DataRegistryLike,
PreprocessorLike,
MarkerLike,
StorageLike,

View file

@ -11,6 +11,7 @@ from typing import (
if TYPE_CHECKING:
from ..data import BasePipelineDataRegistry
from ..datagrabber import BaseDataGrabber
from ..datareader import DefaultDataReader
from ..markers import BaseMarker
@ -23,6 +24,7 @@ __all__ = [
"ConfigVal",
"DataGrabberLike",
"DataGrabberPatterns",
"DataRegistryLike",
"Dependencies",
"Element",
"Elements",
@ -35,6 +37,7 @@ __all__ = [
]
DataRegistryLike = type["BasePipelineDataRegistry"]
DataGrabberLike = type["BaseDataGrabber"]
PreprocessorLike = type["BasePreprocessor"]
MarkerLike = type["BaseMarker"]