[ENH]: Simplify PipelineStepMixin interface #476
7 changed files with 3 additions and 89 deletions
1
docs/changes/newsfragments/476.enh
Normal file
1
docs/changes/newsfragments/476.enh
Normal file
|
|
@ -0,0 +1 @@
|
||||||
|
Simplify :class:`.PipelineStepMixin` interface by `Synchon Mandal`_
|
||||||
|
|
@ -56,23 +56,6 @@ class DefaultDataReader(PipelineStepMixin, UpdateMetaMixin):
|
||||||
# Nothing to validate, any input is fine
|
# Nothing to validate, any input is fine
|
||||||
return input
|
return input
|
||||||
|
|
||||||
def get_output_type(self, input_type: str) -> str:
|
|
||||||
"""Get output type.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
input_type : str
|
|
||||||
The data type input to the reader.
|
|
||||||
|
|
||||||
Returns
|
|
||||||
-------
|
|
||||||
str
|
|
||||||
The data type output by the reader.
|
|
||||||
|
|
||||||
"""
|
|
||||||
# It will output the same type of data as the input
|
|
||||||
return input_type
|
|
||||||
|
|
||||||
def _fit_transform(
|
def _fit_transform(
|
||||||
self,
|
self,
|
||||||
input: dict[str, dict],
|
input: dict[str, dict],
|
||||||
|
|
|
||||||
|
|
@ -30,7 +30,6 @@ def test_DefaultDataReader_validation(type_) -> None:
|
||||||
"""
|
"""
|
||||||
reader = DefaultDataReader()
|
reader = DefaultDataReader()
|
||||||
assert reader.validate_input(type_) == type_
|
assert reader.validate_input(type_) == type_
|
||||||
assert reader.get_output_type(type_) == type_
|
|
||||||
assert reader.validate(type_) == type_
|
assert reader.validate(type_) == type_
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -53,25 +53,6 @@ class PipelineStepMixin:
|
||||||
klass=NotImplementedError,
|
klass=NotImplementedError,
|
||||||
) # pragma: no cover
|
) # pragma: no cover
|
||||||
|
|
||||||
def get_output_type(self, input_type: str) -> str:
|
|
||||||
"""Get output type.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
input_type : str
|
|
||||||
The data type input to the marker.
|
|
||||||
|
|
||||||
Returns
|
|
||||||
-------
|
|
||||||
str
|
|
||||||
The storage type output by the marker.
|
|
||||||
|
|
||||||
"""
|
|
||||||
raise_error(
|
|
||||||
msg="Concrete classes need to implement get_output_type().",
|
|
||||||
klass=NotImplementedError,
|
|
||||||
) # pragma: no cover
|
|
||||||
|
|
||||||
def _fit_transform(
|
def _fit_transform(
|
||||||
self,
|
self,
|
||||||
input: dict[str, dict],
|
input: dict[str, dict],
|
||||||
|
|
@ -223,8 +204,9 @@ class PipelineStepMixin:
|
||||||
for val in self._MARKER_INOUT_MAPPINGS[t_input].values()
|
for val in self._MARKER_INOUT_MAPPINGS[t_input].values()
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
# Only for datareader and preprocessor
|
||||||
else:
|
else:
|
||||||
outputs = [self.get_output_type(t_input) for t_input in fit_input]
|
outputs = fit_input
|
||||||
return outputs
|
return outputs
|
||||||
|
|
||||||
def fit_transform(
|
def fit_transform(
|
||||||
|
|
|
||||||
|
|
@ -29,9 +29,6 @@ def test_PipelineStepMixin_correct_dependencies() -> None:
|
||||||
def validate_input(self, input: list[str]) -> list[str]:
|
def validate_input(self, input: list[str]) -> list[str]:
|
||||||
return input
|
return input
|
||||||
|
|
||||||
def get_output_type(self, input_type: str) -> str:
|
|
||||||
return input_type
|
|
||||||
|
|
||||||
def _fit_transform(self, input: dict[str, dict]) -> dict[str, dict]:
|
def _fit_transform(self, input: dict[str, dict]) -> dict[str, dict]:
|
||||||
return {"input": input}
|
return {"input": input}
|
||||||
|
|
||||||
|
|
@ -50,9 +47,6 @@ def test_PipelineStepMixin_incorrect_dependencies() -> None:
|
||||||
def validate_input(self, input: list[str]) -> list[str]:
|
def validate_input(self, input: list[str]) -> list[str]:
|
||||||
return input
|
return input
|
||||||
|
|
||||||
def get_output_type(self, input_type: str) -> str:
|
|
||||||
return input_type
|
|
||||||
|
|
||||||
def _fit_transform(self, input: dict[str, dict]) -> dict[str, dict]:
|
def _fit_transform(self, input: dict[str, dict]) -> dict[str, dict]:
|
||||||
return {"input": input}
|
return {"input": input}
|
||||||
|
|
||||||
|
|
@ -75,9 +69,6 @@ def test_PipelineStepMixin_correct_ext_dependencies() -> None:
|
||||||
def validate_input(self, input: list[str]) -> list[str]:
|
def validate_input(self, input: list[str]) -> list[str]:
|
||||||
return input
|
return input
|
||||||
|
|
||||||
def get_output_type(self, input_type: str) -> str:
|
|
||||||
return input_type
|
|
||||||
|
|
||||||
def _fit_transform(self, input: dict[str, dict]) -> dict[str, dict]:
|
def _fit_transform(self, input: dict[str, dict]) -> dict[str, dict]:
|
||||||
return {"input": input}
|
return {"input": input}
|
||||||
|
|
||||||
|
|
@ -101,9 +92,6 @@ def test_PipelineStepMixin_ext_deps_correct_commands() -> None:
|
||||||
def validate_input(self, input: list[str]) -> list[str]:
|
def validate_input(self, input: list[str]) -> list[str]:
|
||||||
return input
|
return input
|
||||||
|
|
||||||
def get_output_type(self, input_type: str) -> str:
|
|
||||||
return input_type
|
|
||||||
|
|
||||||
def _fit_transform(self, input: dict[str, dict]) -> dict[str, dict]:
|
def _fit_transform(self, input: dict[str, dict]) -> dict[str, dict]:
|
||||||
return {"input": input}
|
return {"input": input}
|
||||||
|
|
||||||
|
|
@ -129,9 +117,6 @@ def test_PipelineStepMixin_ext_deps_incorrect_commands() -> None:
|
||||||
def validate_input(self, input: list[str]) -> list[str]:
|
def validate_input(self, input: list[str]) -> list[str]:
|
||||||
return input
|
return input
|
||||||
|
|
||||||
def get_output_type(self, input_type: str) -> str:
|
|
||||||
return input_type
|
|
||||||
|
|
||||||
def _fit_transform(self, input: dict[str, dict]) -> dict[str, dict]:
|
def _fit_transform(self, input: dict[str, dict]) -> dict[str, dict]:
|
||||||
return {"input": input}
|
return {"input": input}
|
||||||
|
|
||||||
|
|
@ -153,9 +138,6 @@ def test_PipelineStepMixin_incorrect_ext_dependencies() -> None:
|
||||||
def validate_input(self, input: list[str]) -> list[str]:
|
def validate_input(self, input: list[str]) -> list[str]:
|
||||||
return input
|
return input
|
||||||
|
|
||||||
def get_output_type(self, input_type: str) -> str:
|
|
||||||
return input_type
|
|
||||||
|
|
||||||
def _fit_transform(self, input: dict[str, dict]) -> dict[str, dict]:
|
def _fit_transform(self, input: dict[str, dict]) -> dict[str, dict]:
|
||||||
return {"input": input}
|
return {"input": input}
|
||||||
|
|
||||||
|
|
@ -185,9 +167,6 @@ def test_PipelineStepMixin_correct_conditional_dependencies() -> None:
|
||||||
def validate_input(self, input: list[str]) -> list[str]:
|
def validate_input(self, input: list[str]) -> list[str]:
|
||||||
return input
|
return input
|
||||||
|
|
||||||
def get_output_type(self, input_type: str) -> str:
|
|
||||||
return input_type
|
|
||||||
|
|
||||||
def _fit_transform(self, input: dict[str, dict]) -> dict[str, dict]:
|
def _fit_transform(self, input: dict[str, dict]) -> dict[str, dict]:
|
||||||
return {"input": input}
|
return {"input": input}
|
||||||
|
|
||||||
|
|
@ -214,9 +193,6 @@ def test_PipelineStepMixin_incorrect_conditional_dependencies() -> None:
|
||||||
def validate_input(self, input: list[str]) -> list[str]:
|
def validate_input(self, input: list[str]) -> list[str]:
|
||||||
return input
|
return input
|
||||||
|
|
||||||
def get_output_type(self, input_type: str) -> str:
|
|
||||||
return input_type
|
|
||||||
|
|
||||||
def _fit_transform(self, input: dict[str, dict]) -> dict[str, dict]:
|
def _fit_transform(self, input: dict[str, dict]) -> dict[str, dict]:
|
||||||
return {"input": input}
|
return {"input": input}
|
||||||
|
|
||||||
|
|
@ -249,9 +225,6 @@ def test_PipelineStepMixin_correct_conditional_ext_dependencies() -> None:
|
||||||
def validate_input(self, input: list[str]) -> list[str]:
|
def validate_input(self, input: list[str]) -> list[str]:
|
||||||
return input
|
return input
|
||||||
|
|
||||||
def get_output_type(self, input_type: str) -> str:
|
|
||||||
return input_type
|
|
||||||
|
|
||||||
def _fit_transform(self, input: dict[str, dict]) -> dict[str, dict]:
|
def _fit_transform(self, input: dict[str, dict]) -> dict[str, dict]:
|
||||||
return {"input": input}
|
return {"input": input}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -115,23 +115,6 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin):
|
||||||
"""
|
"""
|
||||||
return list(self._VALID_DATA_TYPES)
|
return list(self._VALID_DATA_TYPES)
|
||||||
|
|
||||||
def get_output_type(self, input_type: str) -> str:
|
|
||||||
"""Get output type.
|
|
||||||
|
|
||||||
Parameters
|
|
||||||
----------
|
|
||||||
input_type : str
|
|
||||||
The data type input to the preprocessor.
|
|
||||||
|
|
||||||
Returns
|
|
||||||
-------
|
|
||||||
str
|
|
||||||
The data type output by the preprocessor.
|
|
||||||
|
|
||||||
"""
|
|
||||||
# Does not add any new keys
|
|
||||||
return input_type
|
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def preprocess(
|
def preprocess(
|
||||||
self,
|
self,
|
||||||
|
|
|
||||||
|
|
@ -65,13 +65,6 @@ def test_fMRIPrepConfoundRemover_get_valid_inputs() -> None:
|
||||||
assert confound_remover.get_valid_inputs() == ["BOLD"]
|
assert confound_remover.get_valid_inputs() == ["BOLD"]
|
||||||
|
|
||||||
|
|
||||||
def test_fMRIPrepConfoundRemover_get_output_type() -> None:
|
|
||||||
"""Test fMRIPrepConfoundRemover get_output_type."""
|
|
||||||
confound_remover = fMRIPrepConfoundRemover()
|
|
||||||
# Confound remover works in place
|
|
||||||
assert confound_remover.get_output_type("BOLD") == "BOLD"
|
|
||||||
|
|
||||||
|
|
||||||
def test_fMRIPrepConfoundRemover__map_adhoc_to_fmriprep() -> None:
|
def test_fMRIPrepConfoundRemover__map_adhoc_to_fmriprep() -> None:
|
||||||
"""Test fMRIPrepConfoundRemover adhoc to fmriprep spec mapping."""
|
"""Test fMRIPrepConfoundRemover adhoc to fmriprep spec mapping."""
|
||||||
confound_remover = fMRIPrepConfoundRemover()
|
confound_remover = fMRIPrepConfoundRemover()
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue