[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
|
||||
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(
|
||||
self,
|
||||
input: dict[str, dict],
|
||||
|
|
|
|||
|
|
@ -30,7 +30,6 @@ def test_DefaultDataReader_validation(type_) -> None:
|
|||
"""
|
||||
reader = DefaultDataReader()
|
||||
assert reader.validate_input(type_) == type_
|
||||
assert reader.get_output_type(type_) == type_
|
||||
assert reader.validate(type_) == type_
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -53,25 +53,6 @@ class PipelineStepMixin:
|
|||
klass=NotImplementedError,
|
||||
) # 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(
|
||||
self,
|
||||
input: dict[str, dict],
|
||||
|
|
@ -223,8 +204,9 @@ class PipelineStepMixin:
|
|||
for val in self._MARKER_INOUT_MAPPINGS[t_input].values()
|
||||
}
|
||||
)
|
||||
# Only for datareader and preprocessor
|
||||
else:
|
||||
outputs = [self.get_output_type(t_input) for t_input in fit_input]
|
||||
outputs = fit_input
|
||||
return outputs
|
||||
|
||||
def fit_transform(
|
||||
|
|
|
|||
|
|
@ -29,9 +29,6 @@ def test_PipelineStepMixin_correct_dependencies() -> None:
|
|||
def validate_input(self, input: list[str]) -> list[str]:
|
||||
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]:
|
||||
return {"input": input}
|
||||
|
||||
|
|
@ -50,9 +47,6 @@ def test_PipelineStepMixin_incorrect_dependencies() -> None:
|
|||
def validate_input(self, input: list[str]) -> list[str]:
|
||||
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]:
|
||||
return {"input": input}
|
||||
|
||||
|
|
@ -75,9 +69,6 @@ def test_PipelineStepMixin_correct_ext_dependencies() -> None:
|
|||
def validate_input(self, input: list[str]) -> list[str]:
|
||||
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]:
|
||||
return {"input": input}
|
||||
|
||||
|
|
@ -101,9 +92,6 @@ def test_PipelineStepMixin_ext_deps_correct_commands() -> None:
|
|||
def validate_input(self, input: list[str]) -> list[str]:
|
||||
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]:
|
||||
return {"input": input}
|
||||
|
||||
|
|
@ -129,9 +117,6 @@ def test_PipelineStepMixin_ext_deps_incorrect_commands() -> None:
|
|||
def validate_input(self, input: list[str]) -> list[str]:
|
||||
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]:
|
||||
return {"input": input}
|
||||
|
||||
|
|
@ -153,9 +138,6 @@ def test_PipelineStepMixin_incorrect_ext_dependencies() -> None:
|
|||
def validate_input(self, input: list[str]) -> list[str]:
|
||||
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]:
|
||||
return {"input": input}
|
||||
|
||||
|
|
@ -185,9 +167,6 @@ def test_PipelineStepMixin_correct_conditional_dependencies() -> None:
|
|||
def validate_input(self, input: list[str]) -> list[str]:
|
||||
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]:
|
||||
return {"input": input}
|
||||
|
||||
|
|
@ -214,9 +193,6 @@ def test_PipelineStepMixin_incorrect_conditional_dependencies() -> None:
|
|||
def validate_input(self, input: list[str]) -> list[str]:
|
||||
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]:
|
||||
return {"input": input}
|
||||
|
||||
|
|
@ -249,9 +225,6 @@ def test_PipelineStepMixin_correct_conditional_ext_dependencies() -> None:
|
|||
def validate_input(self, input: list[str]) -> list[str]:
|
||||
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]:
|
||||
return {"input": input}
|
||||
|
||||
|
|
|
|||
|
|
@ -115,23 +115,6 @@ class BasePreprocessor(ABC, PipelineStepMixin, UpdateMetaMixin):
|
|||
"""
|
||||
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
|
||||
def preprocess(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -65,13 +65,6 @@ def test_fMRIPrepConfoundRemover_get_valid_inputs() -> None:
|
|||
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:
|
||||
"""Test fMRIPrepConfoundRemover adhoc to fmriprep spec mapping."""
|
||||
confound_remover = fMRIPrepConfoundRemover()
|
||||
|
|
|
|||
Loading…
Reference in a new issue