From 725dbf8bc262744364d8e60d5babed773004d302 Mon Sep 17 00:00:00 2001 From: Akihiro Nitta Date: Sat, 26 Sep 2026 00:53:55 +0000 Subject: [PATCH] Reland Add experimental `infer_stypes(_low_cardinality: bool)` Restore `infer_stypes(_low_cardinality=...)` from #987 and enable it for `SDMKumoTabularModel`. The category pinning from #987 is not relanded, as `AlignCategories` is back in the recipe (#988). Co-authored-by: Jingang Qu --- benchmark/tabular/model.py | 8 +++++- sdm/stype.py | 55 ++++++++++++++++++++++++++++++++++++-- test/test_stype.py | 35 ++++++++++++++++++++++++ 3 files changed, 95 insertions(+), 3 deletions(-) diff --git a/benchmark/tabular/model.py b/benchmark/tabular/model.py index 4a84ad96e..55533e6df 100644 --- a/benchmark/tabular/model.py +++ b/benchmark/tabular/model.py @@ -43,6 +43,7 @@ class SDMModel(AbstractTorchModel, abc.ABC): default_num_estimators: ClassVar[int] autocast_dtype: ClassVar[torch.dtype] + low_cardinality: ClassVar[Literal["off", "infer"]] = "off" @staticmethod @abc.abstractmethod @@ -87,7 +88,10 @@ def _fit( ) X = self.preprocess(X, y=y) - self.stypes = sdm.infer_stypes(X) + self.stypes = sdm.infer_stypes( + X, + _low_cardinality=self.low_cardinality, + ) x_context = sdm.TableTensor.from_pandas( df=X, stypes=self.stypes, @@ -295,6 +299,8 @@ class SDMKumoTabularModel(SDMModel): size: ClassVar[KumoTabularSize] default_num_estimators = 16 autocast_dtype = torch.float16 + # AutoGluon's feature generator hands binary columns over as integers. + low_cardinality = "infer" # Bagged children are fit one at a time in this process, so they share the # pretrained network of their task through AutoGluon's registry. _default_ag_args_ensemble_extra: ClassVar[dict[str, Any]] = { diff --git a/sdm/stype.py b/sdm/stype.py index c9d2aabd2..5f6e5130d 100644 --- a/sdm/stype.py +++ b/sdm/stype.py @@ -58,6 +58,7 @@ def infer_stypes( *, text: Literal["off", "infer", "drop"] = "off", id: Literal["off", "infer", "drop"] = "off", + _low_cardinality: Literal["off", "infer"] = "off", unsupported: Literal["error", "warn", "drop"] = "error", ) -> dict[str, StypeLike]: r"""Infer semantic types from raw data statistics. @@ -103,7 +104,13 @@ def infer_stypes( """ overrides = overrides or {} - fn: Callable[[str, object, Policy, Policy], Stype | None] | None = None + fn: ( + Callable[ + [str, object, Policy, Policy, Literal["off", "infer"]], + Stype | None, + ] + | None + ) = None columns: Iterable[tuple[Hashable, object]] | None = None if isinstance(table, pa.Table): fn = _infer_arrow_stype @@ -136,7 +143,7 @@ def infer_stypes( continue try: - stype = fn(name, column, text, id) + stype = fn(name, column, text, id, _low_cardinality) except TypeError: if unsupported == "error": raise @@ -169,6 +176,7 @@ def _infer_arrow_stype( array: object, text: Policy, id: Policy, + low_cardinality: Literal["off", "infer"], ) -> Stype | None: assert isinstance(array, pa.Array | pa.ChunkedArray) dtype = array.type @@ -189,6 +197,8 @@ def _infer_arrow_stype( or pa.types.is_floating(dtype) or pa.types.is_decimal(dtype) ): + if low_cardinality != "off" and _is_arrow_low_cardinality(array): + return Stype.categorical return Stype.numerical if pa.types.is_boolean(dtype) or pa.types.is_dictionary(dtype): @@ -210,6 +220,7 @@ def _infer_pandas_stype( ser: object, text: Policy, id: Policy, + low_cardinality: Literal["off", "infer"], ) -> Stype | None: import pandas as pd from pandas.api.types import ( @@ -237,6 +248,8 @@ def _infer_pandas_stype( return None if id == "drop" else Stype.id if is_integer_dtype(dtype) or is_float_dtype(dtype): + if low_cardinality != "off" and _is_series_low_cardinality(ser): + return Stype.categorical return Stype.numerical if is_bool_dtype(dtype) or isinstance(dtype, pd.CategoricalDtype): @@ -258,6 +271,7 @@ def _infer_cudf_stype( ser: object, text: Policy, id: Policy, + low_cardinality: Literal["off", "infer"], ) -> Stype | None: import cudf from cudf.api.types import ( @@ -284,6 +298,8 @@ def _infer_cudf_stype( or is_float_dtype(dtype) or is_decimal_dtype(dtype) ): + if low_cardinality != "off" and _is_series_low_cardinality(ser): + return Stype.categorical return Stype.numerical if is_bool_dtype(dtype) or isinstance(dtype, cudf.CategoricalDtype): @@ -343,3 +359,38 @@ def _is_cudf_text(ser: cudf.Series) -> bool: unique = ser.dropna().unique() avg_words = unique.str.token_count().mean() return avg_words >= _TEXT_MIN_AVERAGE_WORD_COUNT + + +_LOW_CARDINALITY_MIN_ROWS = 151 +_LOW_CARDINALITY_MAX_UNIQUE_VALUES = 3 +_LOW_CARDINALITY_PREFIX_ROWS = 1024 + + +def _is_arrow_low_cardinality(array: pa.Array | pa.ChunkedArray) -> bool: + if len(array) < _LOW_CARDINALITY_MIN_ROWS: + return False + + # A prefix holds a subset of the distinct values, so most columns are + # ruled out without a full pass. + options = pc.CountOptions(mode="all") + prefix = array.slice(0, _LOW_CARDINALITY_PREFIX_ROWS) + num_unique = pc.call_function("count_distinct", [prefix], options).as_py() + if num_unique > _LOW_CARDINALITY_MAX_UNIQUE_VALUES: + return False + + num_unique = pc.call_function("count_distinct", [array], options).as_py() + return 1 < num_unique <= _LOW_CARDINALITY_MAX_UNIQUE_VALUES + + +def _is_series_low_cardinality(ser: pd.Series | cudf.Series) -> bool: + if len(ser) < _LOW_CARDINALITY_MIN_ROWS: + return False + + # A prefix holds a subset of the distinct values, so most columns are + # ruled out without a full pass. + prefix = ser.iloc[:_LOW_CARDINALITY_PREFIX_ROWS] + if prefix.nunique(dropna=False) > _LOW_CARDINALITY_MAX_UNIQUE_VALUES: + return False + + num_unique = ser.nunique(dropna=False) + return 1 < num_unique <= _LOW_CARDINALITY_MAX_UNIQUE_VALUES diff --git a/test/test_stype.py b/test/test_stype.py index f6534fc04..f38f8ac25 100644 --- a/test/test_stype.py +++ b/test/test_stype.py @@ -145,6 +145,41 @@ def test_id_detection() -> None: } +@pytest.mark.parametrize("backend", _BACKENDS) +def test_low_cardinality_detection(backend: str) -> None: + def make_table(num_rows: int) -> pa.Table | pd.DataFrame | cudf.DataFrame: + data = { + "binary": [i % 2 for i in range(num_rows)], + "ternary": [(0.5, 1.5, None)[i % 3] for i in range(num_rows)], + "count": [(0, 1, 2, None)[i % 4] for i in range(num_rows)], + # Two more values in the last rows only. + "late": [ + i % 2 if i < num_rows - 2 else i for i in range(num_rows) + ], + "constant": [1] * num_rows, + } + if backend == "pandas": + return pd.DataFrame(data) + if backend == "arrow": + return pa.table(data) + cudf = pytest.importorskip("cudf") + return cudf.DataFrame(data) + + expected = dict.fromkeys( + ["binary", "ternary", "count", "late", "constant"], + Stype.numerical, + ) + assert infer_stypes(make_table(2048)) == expected + assert infer_stypes(make_table(150), _low_cardinality="infer") == expected + for num_rows in (151, 2048): + table = make_table(num_rows) + assert infer_stypes(table, _low_cardinality="infer") == { + **expected, + "binary": Stype.categorical, + "ternary": Stype.categorical, + } + + def test_overrides() -> None: table = pa.table( {