From 2b655efaf26f7802e4e41f2e64e1b9abdcaa6cd2 Mon Sep 17 00:00:00 2001 From: "Thomas J. Fan" Date: Wed, 19 Aug 2020 04:40:34 -0400 Subject: [PATCH] FIX Adds remainder in column transformer repr_html (#18167) --- doc/whats_new/v0.24.rst | 3 ++ sklearn/compose/_column_transformer.py | 17 +++++- .../compose/tests/test_column_transformer.py | 52 +++++++++++++++++++ 3 files changed, 71 insertions(+), 1 deletion(-) diff --git a/doc/whats_new/v0.24.rst b/doc/whats_new/v0.24.rst index 7feaf21d3d3..aaf86a2f057 100644 --- a/doc/whats_new/v0.24.rst +++ b/doc/whats_new/v0.24.rst @@ -86,6 +86,9 @@ Changelog column selector is a list of bools that are False. :pr:`17616` by `Thomas Fan`_. +- |FIX| :class:`compose.ColumnTransformer` now displays the remainder in the + diagram display. :pr:`18167` by `Thomas Fan`_. + :mod:`sklearn.covariance` ......................... diff --git a/sklearn/compose/_column_transformer.py b/sklearn/compose/_column_transformer.py index 66c155f1f82..7ee8f9d0271 100644 --- a/sklearn/compose/_column_transformer.py +++ b/sklearn/compose/_column_transformer.py @@ -639,7 +639,22 @@ class ColumnTransformer(TransformerMixin, _BaseComposition): return np.hstack(Xs) def _sk_visual_block_(self): - names, transformers, name_details = zip(*self.transformers) + if isinstance(self.remainder, str) and self.remainder == 'drop': + transformers = self.transformers + elif hasattr(self, "_remainder"): + remainder_columns = self._remainder[2] + if hasattr(self, '_df_columns'): + remainder_columns = ( + self._df_columns[remainder_columns].tolist() + ) + transformers = chain(self.transformers, + [('remainder', self.remainder, + remainder_columns)]) + else: + transformers = chain(self.transformers, + [('remainder', self.remainder, '')]) + + names, transformers, name_details = zip(*transformers) return _VisualBlock('parallel', transformers, names=names, name_details=name_details) diff --git a/sklearn/compose/tests/test_column_transformer.py b/sklearn/compose/tests/test_column_transformer.py index bfbd43ccf89..4e58769e244 100644 --- a/sklearn/compose/tests/test_column_transformer.py +++ b/sklearn/compose/tests/test_column_transformer.py @@ -1381,3 +1381,55 @@ def test_feature_names_empty_columns(empty_col): ct.fit(df) assert ct.get_feature_names() == ['ohe__x0_a', 'ohe__x0_b', 'ohe__x1_z'] + + +@pytest.mark.parametrize('remainder', ["passthrough", StandardScaler()]) +def test_sk_visual_block_remainder(remainder): + # remainder='passthrough' or an estimator will be shown in repr_html + ohe = OneHotEncoder() + ct = ColumnTransformer(transformers=[('ohe', ohe, ["col1", "col2"])], + remainder=remainder) + visual_block = ct._sk_visual_block_() + assert visual_block.names == ('ohe', 'remainder') + assert visual_block.name_details == (['col1', 'col2'], '') + assert visual_block.estimators == (ohe, remainder) + + +def test_sk_visual_block_remainder_drop(): + # remainder='drop' is not shown in repr_html + ohe = OneHotEncoder() + ct = ColumnTransformer(transformers=[('ohe', ohe, ["col1", "col2"])]) + visual_block = ct._sk_visual_block_() + assert visual_block.names == ('ohe',) + assert visual_block.name_details == (['col1', 'col2'],) + assert visual_block.estimators == (ohe,) + + +@pytest.mark.parametrize('remainder', ["passthrough", StandardScaler()]) +def test_sk_visual_block_remainder_fitted_pandas(remainder): + # Remainder shows the columns after fitting + pd = pytest.importorskip('pandas') + ohe = OneHotEncoder() + ct = ColumnTransformer(transformers=[('ohe', ohe, ["col1", "col2"])], + remainder=remainder) + df = pd.DataFrame({"col1": ["a", "b", "c"], "col2": ["z", "z", "z"], + "col3": [1, 2, 3], "col4": [3, 4, 5]}) + ct.fit(df) + visual_block = ct._sk_visual_block_() + assert visual_block.names == ('ohe', 'remainder') + assert visual_block.name_details == (['col1', 'col2'], ['col3', 'col4']) + assert visual_block.estimators == (ohe, remainder) + + +@pytest.mark.parametrize('remainder', ["passthrough", StandardScaler()]) +def test_sk_visual_block_remainder_fitted_numpy(remainder): + # Remainder shows the indices after fitting + X = np.array([[1, 2, 3], [4, 5, 6]], dtype=float) + scaler = StandardScaler() + ct = ColumnTransformer(transformers=[('scale', scaler, [0, 2])], + remainder=remainder) + ct.fit(X) + visual_block = ct._sk_visual_block_() + assert visual_block.names == ('scale', 'remainder') + assert visual_block.name_details == ([0, 2], [1]) + assert visual_block.estimators == (scaler, remainder)