diff --git a/rdt/transformers/ordinal.py b/rdt/transformers/ordinal.py index e88a1c0d..848df65e 100644 --- a/rdt/transformers/ordinal.py +++ b/rdt/transformers/ordinal.py @@ -85,10 +85,10 @@ def _get_order(self, data): order = order.sort_values(key=lambda x: x.astype(str)) if pd.isna(data).any(): - order = np.append(order, [np.nan]) + order = pd.Series(np.append(order, [None])) if self.missing_value_encoding is None: - order = self.order[~pd.isna(self.order)] + order = self.order.dropna() return order diff --git a/tests/integration/transformers/test_ordinal.py b/tests/integration/transformers/test_ordinal.py index 8a0229bc..fae06a77 100644 --- a/tests/integration/transformers/test_ordinal.py +++ b/tests/integration/transformers/test_ordinal.py @@ -14,6 +14,27 @@ class TestOrderedUniformEncoder: """Test class for the OrderedUniformEncoder.""" + def test_end_to_end(self): + """Test that transformer works end-to-end with default alphanumeric ordering.""" + # Setup + data = pd.DataFrame({'column_name': [1.0, 2.0, 3.0, 2.0, np.nan, 1.0, 1.0, 0.0]}) + transformer = OrderedUniformEncoder() + column = 'column_name' + + # Run + with warnings.catch_warnings(): + warnings.filterwarnings('error', module='rdt.transformers.categorical') + transformer.fit(data, column) + learned_order = transformer._get_order(data['column_name']) + transformed = transformer.transform(data) + reverse = transformer.reverse_transform(transformed) + expected_order = pd.Series([0, 1, 2, 3, None], dtype='object') + + # Asserts + pd.testing.assert_series_equal(reverse[column], data[column]) + pd.testing.assert_series_equal(learned_order, expected_order) + assert set(transformer.intervals.keys()) == {0, 1, 2, 3, None} + def test_order(self): """Test that the ``order`` parameter is respected.""" # Setup diff --git a/tests/unit/transformers/test_ordinal.py b/tests/unit/transformers/test_ordinal.py index e70b72d4..6bb6d06a 100644 --- a/tests/unit/transformers/test_ordinal.py +++ b/tests/unit/transformers/test_ordinal.py @@ -32,7 +32,7 @@ def test___init__(self): transformer = OrderedUniformEncoder(order=['b', 'c', 'a', None]) # Asserts - pd.testing.assert_series_equal(transformer.order, pd.Series(['b', 'c', 'a', np.nan])) + pd.testing.assert_series_equal(transformer.order, pd.Series(['b', 'c', 'a', None])) def test___init___duplicate_categories(self): """Test the ``__init__`` method errors if duplicate categories provided. @@ -108,7 +108,7 @@ def test__get_order_numerical(self): ordered = transformer._get_order(arr) # Assert - np.testing.assert_array_equal(ordered, np.array([-2.5, 3.11, 5, 67.8, 100, np.nan])) + np.testing.assert_array_equal(ordered, np.array([-2.5, 3.11, 5, 67.8, 100, None])) def test__order_warns_mixed_dtypes(self): """Test transformer warns if the data contains mixed dtypes."""