diff --git a/docs/changes/newsfragments/476.enh b/docs/changes/newsfragments/476.enh new file mode 100644 index 000000000..834ef69e4 --- /dev/null +++ b/docs/changes/newsfragments/476.enh @@ -0,0 +1 @@ +Simplify :class:`.PipelineStepMixin` interface by `Synchon Mandal`_ diff --git a/junifer/datareader/default.py b/junifer/datareader/default.py index 9c3a86f6c..8859a0faf 100644 --- a/junifer/datareader/default.py +++ b/junifer/datareader/default.py @@ -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], diff --git a/junifer/datareader/tests/test_default_reader.py b/junifer/datareader/tests/test_default_reader.py index b81522014..663a89bb4 100644 --- a/junifer/datareader/tests/test_default_reader.py +++ b/junifer/datareader/tests/test_default_reader.py @@ -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_ diff --git a/junifer/pipeline/pipeline_step_mixin.py b/junifer/pipeline/pipeline_step_mixin.py index 151948216..926e486e5 100644 --- a/junifer/pipeline/pipeline_step_mixin.py +++ b/junifer/pipeline/pipeline_step_mixin.py @@ -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( diff --git a/junifer/pipeline/tests/test_pipeline_step_mixin.py b/junifer/pipeline/tests/test_pipeline_step_mixin.py index a5bd72621..d9b6c208b 100644 --- a/junifer/pipeline/tests/test_pipeline_step_mixin.py +++ b/junifer/pipeline/tests/test_pipeline_step_mixin.py @@ -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} diff --git a/junifer/preprocess/base.py b/junifer/preprocess/base.py index f2200b212..4cc51e5fb 100644 --- a/junifer/preprocess/base.py +++ b/junifer/preprocess/base.py @@ -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, diff --git a/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py b/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py index ab026701e..d8c8aa3da 100644 --- a/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py +++ b/junifer/preprocess/confounds/tests/test_fmriprep_confound_remover.py @@ -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()