[ENH] Move junifer.api.registry to junifer.pipeline #50
8 changed files with 26 additions and 21 deletions
|
|
@ -7,4 +7,3 @@
|
||||||
from .cli import cli
|
from .cli import cli
|
||||||
from .functions import run, collect
|
from .functions import run, collect
|
||||||
from . import decorators
|
from . import decorators
|
||||||
from . import registry
|
|
||||||
|
|
|
||||||
|
|
@ -5,7 +5,7 @@
|
||||||
# Synchon Mandal <s.mandal@fz-juelich.de>
|
# Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from .registry import register
|
from ..pipeline.registry import register
|
||||||
|
|
||||||
|
|
||||||
def register_datagrabber(klass: type) -> type:
|
def register_datagrabber(klass: type) -> type:
|
||||||
|
|
|
||||||
|
|
@ -16,10 +16,10 @@ import yaml
|
||||||
from ..datagrabber.base import BaseDataGrabber
|
from ..datagrabber.base import BaseDataGrabber
|
||||||
from ..markers.base import BaseMarker
|
from ..markers.base import BaseMarker
|
||||||
from ..markers.collection import MarkerCollection
|
from ..markers.collection import MarkerCollection
|
||||||
|
from ..pipeline.registry import build
|
||||||
from ..storage.base import BaseFeatureStorage
|
from ..storage.base import BaseFeatureStorage
|
||||||
from ..utils import logger, raise_error
|
from ..utils import logger, raise_error
|
||||||
from ..utils.fs import make_executable
|
from ..utils.fs import make_executable
|
||||||
from .registry import build
|
|
||||||
|
|
||||||
|
|
||||||
def _get_datagrabber(datagrabber_config: Dict) -> BaseDataGrabber:
|
def _get_datagrabber(datagrabber_config: Dict) -> BaseDataGrabber:
|
||||||
|
|
|
||||||
|
|
@ -11,8 +11,8 @@ import pytest
|
||||||
|
|
||||||
import junifer.testing.registry # noqa: F401
|
import junifer.testing.registry # noqa: F401
|
||||||
from junifer.api.functions import collect, run
|
from junifer.api.functions import collect, run
|
||||||
from junifer.api.registry import build
|
|
||||||
from junifer.datagrabber.base import BaseDataGrabber
|
from junifer.datagrabber.base import BaseDataGrabber
|
||||||
|
from junifer.pipeline.registry import build
|
||||||
|
|
||||||
|
|
||||||
# Define datagrabber
|
# Define datagrabber
|
||||||
|
|
|
||||||
|
|
@ -11,12 +11,13 @@ from ..utils.logging import logger, raise_error
|
||||||
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from ..datagrabber.base import BaseDataGrabber
|
from ..datagrabber import BaseDataGrabber
|
||||||
from ..pipeline import PipelineStepMixin
|
from ..storage import BaseFeatureStorage
|
||||||
from ..storage.base import BaseFeatureStorage
|
from .pipeline_step_mixin import PipelineStepMixin
|
||||||
|
|
||||||
|
|
||||||
# Define valid steps for operation
|
# Define valid steps for operation
|
||||||
_valid_steps = [
|
_VALID_STEPS: List[str] = [
|
||||||
"datagrabber",
|
"datagrabber",
|
||||||
"datareader",
|
"datareader",
|
||||||
"preprocessing",
|
"preprocessing",
|
||||||
|
|
@ -25,7 +26,7 @@ _valid_steps = [
|
||||||
]
|
]
|
||||||
|
|
||||||
# Define registry for valid steps
|
# Define registry for valid steps
|
||||||
_registry = {x: {} for x in _valid_steps}
|
_REGISTRY: Dict[str, Dict[str, type]] = {x: {} for x in _VALID_STEPS}
|
||||||
|
|
||||||
|
|
||||||
def register(step: str, name: str, klass: type) -> None:
|
def register(step: str, name: str, klass: type) -> None:
|
||||||
|
|
@ -42,14 +43,14 @@ def register(step: str, name: str, klass: type) -> None:
|
||||||
|
|
||||||
"""
|
"""
|
||||||
# Verify step
|
# Verify step
|
||||||
if step not in _valid_steps:
|
if step not in _VALID_STEPS:
|
||||||
raise_error(msg=f"Invalid step: {step}", klass=ValueError)
|
raise_error(msg=f"Invalid step: {step}", klass=ValueError)
|
||||||
|
|
||||||
logger.info(f"Registering {name} in {step}")
|
logger.info(f"Registering {name} in {step}")
|
||||||
_registry[step][name] = klass
|
_REGISTRY[step][name] = klass
|
||||||
|
|
||||||
|
|
||||||
def get_step_names(step: str) -> List:
|
def get_step_names(step: str) -> List[str]:
|
||||||
"""Get the names of the registered functions for a given step.
|
"""Get the names of the registered functions for a given step.
|
||||||
|
|
||||||
Parameters
|
Parameters
|
||||||
|
|
@ -64,10 +65,10 @@ def get_step_names(step: str) -> List:
|
||||||
|
|
||||||
"""
|
"""
|
||||||
# Verify step
|
# Verify step
|
||||||
if step not in _valid_steps:
|
if step not in _VALID_STEPS:
|
||||||
raise_error(msg=f"Invalid step: {step}", klass=ValueError)
|
raise_error(msg=f"Invalid step: {step}", klass=ValueError)
|
||||||
|
|
||||||
return list(_registry[step].keys())
|
return list(_REGISTRY[step].keys())
|
||||||
|
|
||||||
|
|
||||||
def get_class(step: str, name: str) -> type:
|
def get_class(step: str, name: str) -> type:
|
||||||
|
|
@ -87,13 +88,13 @@ def get_class(step: str, name: str) -> type:
|
||||||
|
|
||||||
"""
|
"""
|
||||||
# Verify step
|
# Verify step
|
||||||
if step not in _valid_steps:
|
if step not in _VALID_STEPS:
|
||||||
raise_error(msg=f"Invalid step: {step}", klass=ValueError)
|
raise_error(msg=f"Invalid step: {step}", klass=ValueError)
|
||||||
# Verify step name
|
# Verify step name
|
||||||
if name not in _registry[step]:
|
if name not in _REGISTRY[step]:
|
||||||
raise_error(msg=f"Invalid name: {name}", klass=ValueError)
|
raise_error(msg=f"Invalid name: {name}", klass=ValueError)
|
||||||
|
|
||||||
return _registry[step][name]
|
return _REGISTRY[step][name]
|
||||||
|
|
||||||
|
|
||||||
def build(
|
def build(
|
||||||
|
|
@ -10,8 +10,13 @@ from typing import Type
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from junifer.api.registry import build, get_class, get_step_names, register
|
|
||||||
from junifer.datagrabber import PatternDataGrabber
|
from junifer.datagrabber import PatternDataGrabber
|
||||||
|
from junifer.pipeline.registry import (
|
||||||
|
build,
|
||||||
|
get_class,
|
||||||
|
get_step_names,
|
||||||
|
register,
|
||||||
|
)
|
||||||
from junifer.storage import SQLiteFeatureStorage
|
from junifer.storage import SQLiteFeatureStorage
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -4,7 +4,7 @@
|
||||||
# Synchon Mandal <s.mandal@fz-juelich.de>
|
# Synchon Mandal <s.mandal@fz-juelich.de>
|
||||||
# License: AGPL
|
# License: AGPL
|
||||||
|
|
||||||
from ..api.registry import register
|
from ..pipeline.registry import register
|
||||||
from .datagrabbers import (
|
from .datagrabbers import (
|
||||||
OasisVBMTestingDatagrabber,
|
OasisVBMTestingDatagrabber,
|
||||||
SPMAuditoryTestingDatagrabber,
|
SPMAuditoryTestingDatagrabber,
|
||||||
|
|
|
||||||
|
|
@ -2,14 +2,14 @@
|
||||||
|
|
||||||
import importlib
|
import importlib
|
||||||
|
|
||||||
from junifer.api.registry import get_step_names
|
from junifer.pipeline.registry import get_step_names
|
||||||
|
|
||||||
|
|
||||||
def test_testing_registry() -> None:
|
def test_testing_registry() -> None:
|
||||||
"""Test testing registry."""
|
"""Test testing registry."""
|
||||||
import junifer
|
import junifer
|
||||||
|
|
||||||
importlib.reload(junifer.api.registry)
|
importlib.reload(junifer.pipeline.registry)
|
||||||
importlib.reload(junifer)
|
importlib.reload(junifer)
|
||||||
|
|
||||||
assert "OasisVBMTestingDatagrabber" not in get_step_names("datagrabber")
|
assert "OasisVBMTestingDatagrabber" not in get_step_names("datagrabber")
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue