Skip to content

Commit 2139da7

Browse files
immu4989thinkall
andauthored
refactor: use sklearn OrdinalEncoder as source of truth for DataTransformer categorical encoding (#1564) (#1569)
* refactor: use sklearn OrdinalEncoder as source of truth for DataTransformer categorical encoding (#1564) * address review: dict lookup for encoder columns, accurate transform comment * fix: preserve categorical encoding order and compatibility --------- Co-authored-by: Li Jiang <bnujli@gmail.com>
1 parent fbf104e commit 2139da7

2 files changed

Lines changed: 192 additions & 16 deletions

File tree

flaml/automl/data.py

Lines changed: 44 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -315,7 +315,8 @@ def fit_transform(self, X: Union[DataFrame, np.ndarray], y, task: Union[str, "Ta
315315
elif X[column].dtype.name == "category":
316316
current_categories = X[column].cat.categories
317317
if "__NAN__" not in current_categories:
318-
X[column] = X[column].cat.add_categories("__NAN__").fillna("__NAN__")
318+
X[column] = X[column].cat.add_categories("__NAN__")
319+
X[column] = X[column].fillna("__NAN__")
319320
cat_columns.append(column)
320321
else:
321322
X[column] = X[column].fillna("__NAN__")
@@ -351,17 +352,16 @@ def fit_transform(self, X: Union[DataFrame, np.ndarray], y, task: Union[str, "Ta
351352
X.insert(0, TS_TIMESTAMP_COL, ds_col)
352353
if cat_columns:
353354
X[cat_columns] = X[cat_columns].astype("category")
354-
# Pin the per-column category list seen at fit time so
355-
# `transform()` produces the same integer codes for the same
356-
# values regardless of what is passed at predict time (see
357-
# issue #1101). "__NAN__" is reserved as the sentinel slot
358-
# used for values unseen at fit time.
359-
self._cat_categories = {}
360-
for col in cat_columns:
361-
cats = list(X[col].cat.categories)
362-
if "__NAN__" not in cats:
363-
cats.append("__NAN__")
364-
self._cat_categories[col] = cats
355+
from sklearn.preprocessing import OrdinalEncoder
356+
357+
categories = [X[column].cat.categories.to_numpy(dtype=object) for column in cat_columns]
358+
self._ordinal_encoder = OrdinalEncoder(
359+
categories=[np.arange(len(values)) for values in categories],
360+
handle_unknown="use_encoded_value",
361+
unknown_value=-1,
362+
)
363+
self._ordinal_encoder.fit(np.column_stack([X[column].cat.codes for column in cat_columns]))
364+
self._ordinal_encoder.categories_ = categories
365365
if num_columns:
366366
X_num = X[num_columns]
367367
try:
@@ -459,16 +459,44 @@ def transform(self, X: Union[DataFrame, np.array]):
459459
elif X[column].dtype.name == "category":
460460
current_categories = X[column].cat.categories
461461
if "__NAN__" not in current_categories:
462-
X[column] = X[column].cat.add_categories("__NAN__").fillna("__NAN__")
462+
X[column] = X[column].cat.add_categories("__NAN__")
463+
X[column] = X[column].fillna("__NAN__")
463464
if cat_columns:
464465
X[cat_columns] = X[cat_columns].astype("category")
465466
# Pin codes to the categories seen at fit time so they do not
466467
# drift when the predict-time column has a different value
467468
# distribution than the fit-time column (see issue #1101).
468-
# Older pickles without `_cat_categories` fall back to
469-
# whatever `astype("category")` inferred above.
469+
# Three-tier fallback for cross-version pickle compatibility:
470+
# 1. `_ordinal_encoder` (post-#1564) — sklearn OrdinalEncoder is
471+
# the source of truth for allowed categories;
472+
# 2. `_cat_categories` (post-#1561) — ad-hoc dict from the
473+
# defensive patch that landed before this refactor;
474+
# 3. neither (pre-#1561 pickle) — fall through to
475+
# whatever `astype("category")` inferred above.
476+
encoder = getattr(self, "_ordinal_encoder", None)
470477
saved_cats_map = getattr(self, "_cat_categories", None)
471-
if saved_cats_map:
478+
if encoder is not None:
479+
for col_idx, column in enumerate(cat_columns):
480+
known_cats = list(encoder.categories_[col_idx])
481+
# Include "__NAN__" as the sentinel slot even if the
482+
# fit-time data did not contain missing values.
483+
pinned_cats = list(known_cats)
484+
if "__NAN__" not in pinned_cats:
485+
pinned_cats.append("__NAN__")
486+
current = X[column].astype(object)
487+
unseen_mask = ~current.isin(pinned_cats) & current.notna()
488+
if unseen_mask.any():
489+
samples = sorted({str(v) for v in current[unseen_mask].unique()})[:5]
490+
warnings.warn(
491+
f"Column '{column}' contains values unseen at fit time "
492+
f"(e.g. {samples}); these rows will be encoded as '__NAN__' "
493+
"and predictions may be unreliable.",
494+
UserWarning,
495+
stacklevel=2,
496+
)
497+
current = current.where(~unseen_mask, "__NAN__")
498+
X[column] = pd.Categorical(current, categories=pinned_cats)
499+
elif saved_cats_map:
472500
for column in cat_columns:
473501
saved_cats = saved_cats_map.get(column)
474502
if saved_cats is None:

test/automl/test_preprocess_api.py

Lines changed: 148 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,10 @@
11
"""Tests for the public preprocessor APIs."""
22
import unittest
3+
from unittest.mock import patch
34

45
import numpy as np
56
import pandas as pd
7+
import pytest
68
from sklearn.datasets import load_breast_cancer, load_diabetes
79

810
from flaml import AutoML
@@ -292,5 +294,151 @@ def test_unseen_categories_emit_warning_and_map_to_sentinel(self):
292294
self.assertTrue((unseen_rows == nan_code).all())
293295

294296

297+
class TestOrdinalEncoderBackedTransform(unittest.TestCase):
298+
"""Coverage for the categorical-encoding refactor in #1564: DataTransformer
299+
uses sklearn's OrdinalEncoder as the source of truth for the per-column
300+
category list at fit time, and the three-tier backward-compat fallback
301+
(`_ordinal_encoder` → `_cat_categories` → legacy) at transform time so that
302+
pickles from any recent FLAML version continue to load."""
303+
304+
def _fit_simple(self):
305+
from flaml.automl.data import DataTransformer
306+
from flaml.automl.task.factory import task_factory
307+
308+
rng = np.random.RandomState(0)
309+
n = 100
310+
fit_df = pd.DataFrame({"a": rng.randn(n), "gender": rng.choice(["M", "F"], n)})
311+
fit_y = pd.Series(rng.randn(n))
312+
transformer = DataTransformer()
313+
task = task_factory("regression", fit_df, fit_y)
314+
X_fit, _ = transformer.fit_transform(fit_df.copy(), fit_y, task)
315+
return transformer, X_fit
316+
317+
def test_fit_transform_installs_ordinal_encoder(self):
318+
from sklearn.preprocessing import OrdinalEncoder
319+
320+
transformer, _ = self._fit_simple()
321+
self.assertTrue(hasattr(transformer, "_ordinal_encoder"))
322+
self.assertIsInstance(transformer._ordinal_encoder, OrdinalEncoder)
323+
self.assertIn("gender", transformer._cat_columns)
324+
325+
def test_ordinal_encoder_path_matches_1561_semantics(self):
326+
"""Refactored path preserves the observable behavior from #1561:
327+
- known-category codes are stable across fit/predict distributions,
328+
- unseen values are remapped to the "__NAN__" sentinel code,
329+
- a `UserWarning` is emitted on unseen values."""
330+
import warnings
331+
332+
transformer, X_fit = self._fit_simple()
333+
fit_M_code = int(X_fit["gender"].cat.codes[X_fit["gender"] == "M"].iloc[0])
334+
335+
# (a) Stability under a smaller predict-time value set
336+
predict_df = pd.DataFrame({"a": np.zeros(20), "gender": ["M"] * 20})
337+
X_pred = transformer.transform(predict_df.copy())
338+
pred_M_code = int(X_pred["gender"].cat.codes[X_pred["gender"] == "M"].iloc[0])
339+
self.assertEqual(fit_M_code, pred_M_code)
340+
341+
# (b) Unseen-category warning + (c) sentinel remap
342+
predict_df2 = pd.DataFrame({"a": np.zeros(5), "gender": ["M", "F", "X", "M", "Y"]})
343+
with warnings.catch_warnings(record=True) as caught:
344+
warnings.simplefilter("always")
345+
X_pred2 = transformer.transform(predict_df2.copy())
346+
unseen_warnings = [
347+
w for w in caught if issubclass(w.category, UserWarning) and "unseen at fit time" in str(w.message)
348+
]
349+
self.assertEqual(len(unseen_warnings), 1)
350+
nan_code = list(X_pred2["gender"].cat.categories).index("__NAN__")
351+
unseen_rows = X_pred2["gender"].cat.codes[predict_df2["gender"].isin(["X", "Y"]).values]
352+
self.assertTrue((unseen_rows == nan_code).all())
353+
354+
def test_transform_falls_back_to_cat_categories_when_encoder_missing(self):
355+
"""Pickles produced between #1561 and #1564 only have `_cat_categories`.
356+
`transform()` must still work correctly for them."""
357+
transformer, X_fit = self._fit_simple()
358+
# Simulate a #1561-era pickle by removing the encoder and installing the
359+
# ad-hoc dict the older code produced.
360+
del transformer._ordinal_encoder
361+
transformer._cat_categories = {
362+
"gender": list(X_fit["gender"].cat.categories) + ["__NAN__"],
363+
}
364+
365+
fit_M_code = int(X_fit["gender"].cat.codes[X_fit["gender"] == "M"].iloc[0])
366+
predict_df = pd.DataFrame({"a": np.zeros(20), "gender": ["M"] * 20})
367+
X_pred = transformer.transform(predict_df.copy())
368+
pred_M_code = int(X_pred["gender"].cat.codes[X_pred["gender"] == "M"].iloc[0])
369+
self.assertEqual(fit_M_code, pred_M_code)
370+
371+
def test_transform_legacy_pickle_without_either_attribute(self):
372+
"""Pickles from before #1561 have neither `_ordinal_encoder` nor
373+
`_cat_categories`. `transform()` must not raise; it falls through to the
374+
legacy `astype("category")` path. Drift is possible on those pickles
375+
(that is the bug that #1561 and #1564 fix), but load-and-predict must
376+
continue to work without upgrading users' pickle files."""
377+
transformer, _ = self._fit_simple()
378+
del transformer._ordinal_encoder # no `_cat_categories` was ever set
379+
380+
predict_df = pd.DataFrame({"a": np.zeros(20), "gender": ["M"] * 20})
381+
# Legacy path just does astype("category") — no exception, no warning.
382+
X_pred = transformer.transform(predict_df.copy())
383+
self.assertEqual(str(X_pred["gender"].dtype), "category")
384+
385+
386+
@pytest.mark.parametrize("columns", [["first", "second"], [0, 1], [0, "second"]])
387+
@pytest.mark.parametrize("categories", [["z", "a", "unused"], [3, 1, 2], [1, "1", "unused"]])
388+
def test_category_order_and_column_labels(columns, categories):
389+
from flaml.automl.data import DataTransformer
390+
from flaml.automl.task.factory import task_factory
391+
392+
values = [categories[0], categories[1], None] * 4
393+
X = pd.DataFrame(
394+
{
395+
columns[0]: pd.Categorical(values, categories=categories, ordered=True),
396+
columns[1]: pd.Categorical(["b", "a"] * 6, categories=["b", "a"]),
397+
}
398+
)
399+
y = pd.Series(np.arange(len(X), dtype=float))
400+
transformer = DataTransformer()
401+
fitted, _ = transformer.fit_transform(X.copy(), y, task_factory("regression", X, y))
402+
predicted = transformer.transform(X.copy())
403+
for column in columns:
404+
assert list(fitted[column].cat.categories) == list(predicted[column].cat.categories)
405+
np.testing.assert_array_equal(fitted[column].cat.codes, predicted[column].cat.codes)
406+
assert list(fitted[columns[0]].cat.categories) == categories + ["__NAN__"]
407+
np.testing.assert_array_equal(predicted[columns[0]].cat.codes, [0, 1, 3] * 4)
408+
subset = X.iloc[[1, 0]].astype(object)
409+
np.testing.assert_array_equal(transformer.transform(subset)[columns[0]].cat.codes, [1, 0])
410+
subset.loc[subset.index[0], columns[0]] = "new"
411+
with pytest.warns(UserWarning, match="unseen at fit time"):
412+
unseen = transformer.transform(subset)
413+
assert unseen[columns[0]].iloc[0] == "__NAN__"
414+
415+
416+
def test_encoder_uses_sklearn_10_constructor():
417+
from sklearn.preprocessing import OrdinalEncoder
418+
419+
def legacy_encoder(*, categories="auto", dtype=np.float64, handle_unknown="error", unknown_value=None):
420+
return OrdinalEncoder(
421+
categories=categories, dtype=dtype, handle_unknown=handle_unknown, unknown_value=unknown_value
422+
)
423+
424+
with patch("sklearn.preprocessing.OrdinalEncoder", side_effect=legacy_encoder):
425+
transformer, fitted = TestOrdinalEncoderBackedTransform()._fit_simple()
426+
predicted = transformer.transform(fitted.copy())
427+
np.testing.assert_array_equal(fitted["gender"].cat.codes, predicted["gender"].cat.codes)
428+
429+
430+
def test_existing_missing_sentinel_is_filled():
431+
from flaml.automl.data import DataTransformer
432+
from flaml.automl.task.factory import task_factory
433+
434+
X = pd.DataFrame({"cat": pd.Categorical([1, "one", None] * 4, categories=[1, "one", "__NAN__"])})
435+
y = pd.Series(np.arange(len(X), dtype=float))
436+
transformer = DataTransformer()
437+
fitted, _ = transformer.fit_transform(X.copy(), y, task_factory("regression", X, y))
438+
predicted = transformer.transform(X.copy())
439+
np.testing.assert_array_equal(fitted["cat"].cat.codes, [0, 1, 2] * 4)
440+
np.testing.assert_array_equal(predicted["cat"].cat.codes, fitted["cat"].cat.codes)
441+
442+
295443
if __name__ == "__main__":
296444
unittest.main()

0 commit comments

Comments
 (0)