|
1 | 1 | """Tests for the public preprocessor APIs.""" |
2 | 2 | import unittest |
| 3 | +from unittest.mock import patch |
3 | 4 |
|
4 | 5 | import numpy as np |
5 | 6 | import pandas as pd |
| 7 | +import pytest |
6 | 8 | from sklearn.datasets import load_breast_cancer, load_diabetes |
7 | 9 |
|
8 | 10 | from flaml import AutoML |
@@ -292,5 +294,151 @@ def test_unseen_categories_emit_warning_and_map_to_sentinel(self): |
292 | 294 | self.assertTrue((unseen_rows == nan_code).all()) |
293 | 295 |
|
294 | 296 |
|
| 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 | + |
295 | 443 | if __name__ == "__main__": |
296 | 444 | unittest.main() |
0 commit comments