[ENH]: Enable adding data types #451

Merged
synchon merged 8 commits from feat/extendable-data-types into main 2025-07-18 10:06:20 +00:00
8 changed files with 388 additions and 88 deletions

View file

@ -0,0 +1 @@
Add documentation on adding data types by `Synchon Mandal`_

View file

@ -0,0 +1 @@
Enable data types to be added by introducing :func:`.register_data_type` by `Synchon Mandal`_

View file

@ -0,0 +1,57 @@
.. include:: ../links.inc
.. _adding_data_types:
Adding Data Types
=================
``junifer`` supports most of the :ref:`data types <data_types>` required for fMRI
but also provides a way to add custom data types in case you work with other
modalities like EEG.
How to add a data type
----------------------
#. Check :ref:`extending junifer <extending_extension>` on how to create a
*junifer extension* if you have not done so.
#. Define the data type schema in the *extension script* like so:
.. code-block:: python
from junifer.datagrabber import DataTypeSchema
dtype_schema: DataTypeSchema = {
"mandatory": ["pattern"],
"optional": {
"mask": {
"mandatory": ["pattern"],
"optional": [],
},
},
}
* The :obj:`.DataTypeSchema` has two mandatory keys:
* ``mandatory`` : list of str
* ``optional`` : dict of str and :obj:`.OptionalTypeSchema`
* ``mandatory`` defines the keys that must be present when defining a *pattern*
in a DataGrabber.
* ``optional`` defines the mapping from *sub-types* that are optional, to their patterns.
The patterns in turn require a ``mandatory`` key and an ``optional`` key both
just being the keys that must be there if the optional key is found. It's
possible that the *sub-type* (``mask`` in the example) can be absent from the dataset.
#. Register the data type before defining / using a DataGrabber like so:
.. code-block:: python
from junifer.datagrabber import register_data_type
...
# registers the data type as "dtype"
register_data_type(name="dtype", schema=dtype_schema)
...

View file

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

View file

@ -10,7 +10,11 @@ __all__ = [
"DataladHCP1200", "DataladHCP1200",
"MultipleDataGrabber", "MultipleDataGrabber",
"DMCC13Benchmark", "DMCC13Benchmark",
"DataTypeManager",
"DataTypeSchema",
"OptionalTypeSchema",
"PatternValidationMixin", "PatternValidationMixin",
"register_data_type",
] ]
# These 4 need to be in this order, otherwise it is a circular import # These 4 need to be in this order, otherwise it is a circular import
@ -24,4 +28,10 @@ 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 from .pattern_validation_mixin import (
DataTypeManager,
DataTypeSchema,
OptionalTypeSchema,
PatternValidationMixin,
register_data_type,
)

View file

@ -3,71 +3,199 @@
# 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 typing import TypedDict
from ..typing import DataGrabberPatterns from ..typing import DataGrabberPatterns
from ..utils import logger, raise_error, warn_with_log from ..utils import logger, raise_error, warn_with_log
__all__ = ["PatternValidationMixin"] __all__ = [
"DataTypeManager",
"DataTypeSchema",
"OptionalTypeSchema",
"PatternValidationMixin",
"register_data_type",
]
# Define schema for pattern-based datagrabber's patterns class OptionalTypeSchema(TypedDict):
PATTERNS_SCHEMA = { """Optional type schema."""
"T1w": {
"mandatory": ["pattern", "space"], mandatory: list[str]
"optional": { optional: list[str]
"mask": {"mandatory": ["pattern", "space"], "optional": []},
},
}, class DataTypeSchema(TypedDict):
"T2w": { """Data type schema."""
"mandatory": ["pattern", "space"],
"optional": { mandatory: list[str]
"mask": {"mandatory": ["pattern", "space"], "optional": []}, optional: dict[str, OptionalTypeSchema]
},
},
"BOLD": { class DataTypeManager(MutableMapping):
"mandatory": ["pattern", "space"], """Class for managing data types."""
"optional": {
"mask": {"mandatory": ["pattern", "space"], "optional": []}, _instance = None
"confounds": {
"mandatory": ["pattern", "format"], def __new__(cls):
"optional": ["mappings"], """Overridden to make the class singleton."""
}, # Make class singleton
"reference": {"mandatory": ["pattern"], "optional": []}, if cls._instance is None:
"prewarp_space": {"mandatory": [], "optional": []}, cls._instance = super().__new__(cls)
}, # Set global schema
}, cls._global: dict[str, DataTypeSchema] = {}
"Warp": { cls._builtin: dict[str, DataTypeSchema] = {}
"mandatory": ["pattern", "src", "dst", "warper"], cls._external: dict[str, DataTypeSchema] = {}
"optional": {}, cls._builtin.update(
}, {
"VBM_GM": { "T1w": {
"mandatory": ["pattern", "space"], "mandatory": ["pattern", "space"],
"optional": {}, "optional": {
}, "mask": {
"VBM_WM": { "mandatory": ["pattern", "space"],
"mandatory": ["pattern", "space"], "optional": [],
"optional": {}, },
}, },
"VBM_CSF": { },
"mandatory": ["pattern", "space"], "T2w": {
"optional": {}, "mandatory": ["pattern", "space"],
}, "optional": {
"DWI": { "mask": {
"mandatory": ["pattern"], "mandatory": ["pattern", "space"],
"optional": {}, "optional": [],
}, },
"FreeSurfer": { },
"mandatory": ["pattern"], },
"optional": { "BOLD": {
"aseg": {"mandatory": ["pattern"], "optional": []}, "mandatory": ["pattern", "space"],
"norm": {"mandatory": ["pattern"], "optional": []}, "optional": {
"lh_white": {"mandatory": ["pattern"], "optional": []}, "mask": {
"rh_white": {"mandatory": ["pattern"], "optional": []}, "mandatory": ["pattern", "space"],
"lh_pial": {"mandatory": ["pattern"], "optional": []}, "optional": [],
"rh_pial": {"mandatory": ["pattern"], "optional": []}, },
}, "confounds": {
}, "mandatory": ["pattern", "format"],
} "optional": ["mappings"],
},
"reference": {
"mandatory": ["pattern"],
"optional": [],
},
"prewarp_space": {"mandatory": [], "optional": []},
},
},
"Warp": {
"mandatory": ["pattern", "src", "dst", "warper"],
"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": [],
},
},
},
}
)
cls._global.update(cls._builtin)
return cls._instance
def __getitem__(self, key: str) -> DataTypeSchema:
"""Retrieve schema for ``key``."""
return self._global[key]
def __iter__(self) -> Iterator[str]:
"""Iterate over data types."""
return iter(self._global)
def __len__(self) -> int:
"""Get data type count."""
return len(self._global)
def __delitem__(self, key: str) -> None:
"""Remove schema for ``key``."""
# 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._global.pop(key)
def __setitem__(self, key: str, value: DataTypeSchema) -> None:
"""Update ``key`` with ``value``."""
# Internal check
if key in self._builtin:
raise_error(f"Cannot set value for in-built key: {key}")
# Value type check
if not isinstance(value, dict):
raise_error(f"Invalid value type: {type(value)}")
# Update external
self._external[key] = value
# Update global
self._global[key] = value
def popitem():
"""Not implemented."""
pass
def clear(self):
"""Not implemented."""
pass
def setdefault(self, key: str, value=None):
"""Not implemented."""
pass
def register_data_type(name: str, schema: DataTypeSchema) -> None:
"""Register custom data type.
Parameters
----------
name : str
The data type name.
schema : DataTypeSchema
The data type schema.
"""
DataTypeManager()[name] = schema
class PatternValidationMixin: class PatternValidationMixin:
@ -311,12 +439,13 @@ class PatternValidationMixin:
msg="`patterns` must contain all `types`", klass=ValueError msg="`patterns` must contain all `types`", klass=ValueError
) )
# Check against schema # Check against schema
dtype_mgr = DataTypeManager()
for dtype_key, dtype_val in patterns.items(): for dtype_key, dtype_val in patterns.items():
# Check if valid data type is provided # Check if valid data type is provided
if dtype_key not in PATTERNS_SCHEMA: if dtype_key not in dtype_mgr:
raise_error( raise_error(
f"Unknown data type: {dtype_key}, " f"Unknown data type: {dtype_key}, "
f"should be one of: {list(PATTERNS_SCHEMA.keys())}" f"should be one of: {list(dtype_mgr.keys())}"
) )
# Conditional for list dtype vals like Warp # Conditional for list dtype vals like Warp
if isinstance(dtype_val, list): if isinstance(dtype_val, list):
@ -324,14 +453,14 @@ class PatternValidationMixin:
# Check mandatory keys for data type # Check mandatory keys for data type
self._validate_mandatory_keys( self._validate_mandatory_keys(
keys=list(entry), keys=list(entry),
schema=PATTERNS_SCHEMA[dtype_key]["mandatory"], schema=dtype_mgr[dtype_key]["mandatory"],
data_type=f"{dtype_key}.{idx}", data_type=f"{dtype_key}.{idx}",
partial_pattern_ok=partial_pattern_ok, partial_pattern_ok=partial_pattern_ok,
) )
# Check optional keys for data type # Check optional keys for data type
for optional_key, optional_val in PATTERNS_SCHEMA[ for optional_key, optional_val in dtype_mgr[dtype_key][
dtype_key "optional"
]["optional"].items(): ].items():
if optional_key not in entry: if optional_key not in entry:
logger.debug( logger.debug(
f"Optional key: `{optional_key}` missing for " f"Optional key: `{optional_key}` missing for "
@ -344,12 +473,12 @@ class PatternValidationMixin:
) )
# Set nested type name for easier access # Set nested type name for easier access
nested_dtype = f"{dtype_key}.{idx}.{optional_key}" nested_dtype = f"{dtype_key}.{idx}.{optional_key}"
nested_mandatory_keys_schema = PATTERNS_SCHEMA[ nested_mandatory_keys_schema = dtype_mgr[
dtype_key dtype_key
]["optional"][optional_key]["mandatory"] ]["optional"][optional_key]["mandatory"]
nested_optional_keys_schema = PATTERNS_SCHEMA[ nested_optional_keys_schema = dtype_mgr[dtype_key][
dtype_key "optional"
]["optional"][optional_key]["optional"] ][optional_key]["optional"]
# Check mandatory keys for nested type # Check mandatory keys for nested type
self._validate_mandatory_keys( self._validate_mandatory_keys(
keys=list(optional_val["mandatory"]), keys=list(optional_val["mandatory"]),
@ -392,10 +521,8 @@ class PatternValidationMixin:
self._identify_stray_keys( self._identify_stray_keys(
keys=list(entry.keys()), keys=list(entry.keys()),
schema=( schema=(
PATTERNS_SCHEMA[dtype_key]["mandatory"] dtype_mgr[dtype_key]["mandatory"]
+ list( + list(dtype_mgr[dtype_key]["optional"].keys())
PATTERNS_SCHEMA[dtype_key]["optional"].keys()
)
), ),
data_type=dtype_key, data_type=dtype_key,
) )
@ -412,12 +539,12 @@ class PatternValidationMixin:
# Check mandatory keys for data type # Check mandatory keys for data type
self._validate_mandatory_keys( self._validate_mandatory_keys(
keys=list(dtype_val), keys=list(dtype_val),
schema=PATTERNS_SCHEMA[dtype_key]["mandatory"], schema=dtype_mgr[dtype_key]["mandatory"],
data_type=dtype_key, data_type=dtype_key,
partial_pattern_ok=partial_pattern_ok, partial_pattern_ok=partial_pattern_ok,
) )
# Check optional keys for data type # Check optional keys for data type
for optional_key, optional_val in PATTERNS_SCHEMA[dtype_key][ for optional_key, optional_val in dtype_mgr[dtype_key][
"optional" "optional"
].items(): ].items():
if optional_key not in dtype_val: if optional_key not in dtype_val:
@ -432,12 +559,12 @@ class PatternValidationMixin:
) )
# Set nested type name for easier access # Set nested type name for easier access
nested_dtype = f"{dtype_key}.{optional_key}" nested_dtype = f"{dtype_key}.{optional_key}"
nested_mandatory_keys_schema = PATTERNS_SCHEMA[ nested_mandatory_keys_schema = dtype_mgr[dtype_key][
dtype_key "optional"
]["optional"][optional_key]["mandatory"] ][optional_key]["mandatory"]
nested_optional_keys_schema = PATTERNS_SCHEMA[ nested_optional_keys_schema = dtype_mgr[dtype_key][
dtype_key "optional"
]["optional"][optional_key]["optional"] ][optional_key]["optional"]
# Check mandatory keys for nested type # Check mandatory keys for nested type
self._validate_mandatory_keys( self._validate_mandatory_keys(
keys=list(optional_val["mandatory"]), keys=list(optional_val["mandatory"]),
@ -476,8 +603,8 @@ class PatternValidationMixin:
self._identify_stray_keys( self._identify_stray_keys(
keys=list(dtype_val.keys()), keys=list(dtype_val.keys()),
schema=( schema=(
PATTERNS_SCHEMA[dtype_key]["mandatory"] dtype_mgr[dtype_key]["mandatory"]
+ list(PATTERNS_SCHEMA[dtype_key]["optional"].keys()) + list(dtype_mgr[dtype_key]["optional"].keys())
), ),
data_type=dtype_key, data_type=dtype_key,
) )

View file

@ -9,7 +9,110 @@ from typing import Union
import pytest import pytest
from junifer.datagrabber.pattern_validation_mixin import PatternValidationMixin from junifer.datagrabber.pattern_validation_mixin import (
DataTypeManager,
DataTypeSchema,
PatternValidationMixin,
register_data_type,
)
def test_dtype_mgr_addition_errors() -> None:
"""Test data type manager addition errors."""
with pytest.raises(ValueError, match="Cannot set"):
dtype_schema: DataTypeSchema = {
"mandatory": ["pattern"],
"optional": {},
}
DataTypeManager()["T1w"] = dtype_schema
with pytest.raises(ValueError, match="Invalid"):
DataTypeManager()["DType"] = ""
def test_dtype_mgr_removal_errors() -> None:
"""Test data type manager removal errors."""
with pytest.raises(ValueError, match="Cannot delete"):
_ = DataTypeManager().pop("T1w")
with pytest.raises(KeyError, match="DType"):
del DataTypeManager()["DType"]
@pytest.mark.parametrize(
"dtype",
[
{
"mandatory": ["pattern"],
"optional": {},
},
{
"mandatory": ["pattern"],
"optional": {
"subtype": {
"mandatory": [],
"optional": [],
}
},
},
{
"mandatory": ["pattern"],
"optional": {
"subtype": {
"mandatory": ["pattern"],
"optional": [],
}
},
},
{
"mandatory": ["pattern"],
"optional": {
"subtype": {
"mandatory": ["pattern"],
"optional": ["pattern"],
}
},
},
],
)
def test_dtype_mgr(dtype: DataTypeSchema) -> None:
"""Test data type manager addition and removal.
Parameters
----------
dtype : DataTypeSchema
The parametrized schema.
"""
DataTypeManager().update({"DType": dtype})
assert "DType" in DataTypeManager()
_ = DataTypeManager().pop("DType")
assert "DType" not in DataTypeManager()
def test_register_data_type() -> None:
"""Test data type registration."""
dtype_schema: DataTypeSchema = {
"mandatory": ["pattern"],
"optional": {
"mask": {
"mandatory": ["pattern"],
"optional": [],
},
},
}
register_data_type(
name="dtype",
schema=dtype_schema,
)
assert "dtype" in DataTypeManager()
_ = DataTypeManager().pop("dtype")
assert "dumb" not in DataTypeManager()
@pytest.mark.parametrize( @pytest.mark.parametrize(

View file

@ -64,7 +64,7 @@ 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[ DataGrabberPatterns = dict[
str, Union[dict[str, str], Sequence[dict[str, str]]] str, Union[dict[str, str], list[dict[str, str]]]
] ]
ConfigVal = Union[bool, int, float, str] ConfigVal = Union[bool, int, float, str]
Element = Union[str, tuple[str, ...]] Element = Union[str, tuple[str, ...]]