Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 7 additions & 1 deletion benchmark/tabular/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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]] = {
Expand Down
55 changes: 53 additions & 2 deletions sdm/stype.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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):
Expand All @@ -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 (
Expand Down Expand Up @@ -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):
Expand All @@ -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 (
Expand All @@ -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):
Expand Down Expand Up @@ -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
35 changes: 35 additions & 0 deletions test/test_stype.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
{
Expand Down
Loading