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