[ENH]: Simplify PipelineStepMixin interface #476

Merged
synchon merged 4 commits from refactor/pipeline-step-mixin into main 2025-11-12 16:58:51 +00:00
7 changed files with 3 additions and 89 deletions

View file

@ -0,0 +1 @@
Simplify :class:`.PipelineStepMixin` interface by `Synchon Mandal`_

View file

@ -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],

View file

@ -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_

View file

@ -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(

View file

@ -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}

View file

@ -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,

View file

@ -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()