diff --git a/changelog/1159.added.md b/changelog/1159.added.md new file mode 100644 index 000000000..63398da68 --- /dev/null +++ b/changelog/1159.added.md @@ -0,0 +1 @@ +`fit()` now warns when a column of `X` looks like free text. diff --git a/src/tabpfn/preprocessing/modality_detection.py b/src/tabpfn/preprocessing/modality_detection.py index 25b77bebd..ab7eb7c8c 100644 --- a/src/tabpfn/preprocessing/modality_detection.py +++ b/src/tabpfn/preprocessing/modality_detection.py @@ -4,6 +4,7 @@ from __future__ import annotations +import warnings from collections.abc import Sequence from typing import TYPE_CHECKING @@ -11,6 +12,7 @@ from tabpfn.errors import TabPFNUserError from tabpfn.preprocessing.datamodel import ( + INPUT_FEATURE_PREFIX, Feature, FeatureModality, FeatureSchema, @@ -22,6 +24,10 @@ _EARLY_EXIT_PREFIX_ROWS = 1024 +#: Cap on how many column names the likely-text warning lists, so a wide frame of +#: text columns does not produce an unreadable multi-kilobyte message. +_MAX_TEXT_COLUMNS_IN_WARNING = 10 + def detect_feature_modalities( X: np.ndarray, @@ -75,7 +81,70 @@ def detect_feature_modalities( big_enough_n_to_infer_cat=big_enough_n_to_infer_cat, ) features.append(Feature(name=feature_name, modality=feat_modality)) - return FeatureSchema(features=features) + feature_schema = FeatureSchema(features=features) + _warn_if_text_features( + feature_schema, + declared_categorical_indices=provided_categorical_indices, + ) + return feature_schema + + +def _warn_if_text_features( + feature_schema: FeatureSchema, + *, + declared_categorical_indices: Sequence[int] | None = None, +) -> None: + """Warn when input columns look like free text rather than categoricals. + + High-cardinality string columns are labelled `FeatureModality.TEXT` by + `detect_feature_modalities`, but this package has no text handling: they are swept + into the same `OrdinalEncoder` as real categoricals, which selects columns by dtype + (see `get_ordinal_encoder`). That turns near-unique text into near-unique integer + codes, i.e. noise rather than signal, without any error to hint at it. + + Called by `detect_feature_modalities` while the schema still carries the TEXT + labels, i.e. before the first preprocessing step that rebuilds it, since + `FeatureSchema.from_only_categorical_indices` collapses TEXT into NUMERICAL. + + Args: + feature_schema: The schema produced by `detect_feature_modalities`. + declared_categorical_indices: Positional indices the caller passed as + `categorical_features_indices`. These are never reported: declaring a + column categorical states that the user already knows it holds + non-numeric values and intends them as categories, so warning about it + would be noise. + """ + declared = set(declared_categorical_indices or ()) + text_names = [ + feature.name.removeprefix(INPUT_FEATURE_PREFIX) + for index, feature in enumerate(feature_schema.features) + if feature.modality is FeatureModality.TEXT and index not in declared + ] + if not text_names: + return + + shown = text_names[:_MAX_TEXT_COLUMNS_IN_WARNING] + column_names_to_print = ", ".join(repr(name) for name in shown) + if len(text_names) > len(shown): + column_names_to_print += f" (and {len(text_names) - len(shown)} more)" + + warnings.warn( + f"These columns look like free text and are being ordinal-encoded as " + f"high-cardinality categoricals, which usually adds noise rather than " + f"signal: {column_names_to_print}.\n" + "If such a column holds numbers stored as strings, convert it to a numeric " + "dtype. If it holds genuine text, this package has no text handling -- " + "consider the tabpfn-client API, which embeds text natively: " + "https://github.com/PriorLabs/tabpfn-client \n" + "To silence this for a column that is genuinely a high-cardinality category, " + "pass its index in `categorical_features_indices`.", + UserWarning, + # Points at a direct `estimator.fit(X, y)` call site. Six frames out: this + # function, `detect_feature_modalities`, `_initialize_dataset_preprocessing`, + # `fit`, and the contextlib wrapper added by the `@config_context(...)` + # decorator on `fit`. Pinned by the `warning.filename` asserts in the tests. + stacklevel=6, + ) def _detect_feature_modality( diff --git a/tests/test_preprocessing/test_modality_detection.py b/tests/test_preprocessing/test_modality_detection.py index 3989981c9..80b00b48f 100644 --- a/tests/test_preprocessing/test_modality_detection.py +++ b/tests/test_preprocessing/test_modality_detection.py @@ -4,16 +4,25 @@ from __future__ import annotations +import warnings from typing import Any import numpy as np import pandas as pd import pytest -from tabpfn.preprocessing.datamodel import FeatureModality +from tabpfn import TabPFNClassifier, TabPFNRegressor +from tabpfn.preprocessing.datamodel import ( + INPUT_FEATURE_PREFIX, + Feature, + FeatureModality, + FeatureSchema, +) from tabpfn.preprocessing.modality_detection import ( _EARLY_EXIT_PREFIX_ROWS, + _MAX_TEXT_COLUMNS_IN_WARNING, _detect_feature_modality, + _warn_if_text_features, detect_feature_modalities, ) from tabpfn.preprocessing.type_detection import infer_categorical_features @@ -427,3 +436,221 @@ def test__early_exit_not_fooled_by_uninformative_prefix(): np.concatenate([np.zeros(_EARLY_EXIT_PREFIX_ROWS), np.arange(1.0, 4000.0)]) ) assert _for_test_detect_with_defaults(s) == FeatureModality.NUMERICAL + + +def _text_schema(*names: str) -> FeatureSchema: + """Schema of TEXT features with the `input_` prefix real input names carry.""" + return FeatureSchema( + features=[ + Feature(name=f"{INPUT_FEATURE_PREFIX}{name}", modality=FeatureModality.TEXT) + for name in names + ] + ) + + +class TestWarnIfTextFeatures: + """Schema-level unit tests for `warn_if_text_features`.""" + + def test__no_text_features__does_not_warn(self) -> None: + schema = FeatureSchema( + features=[ + Feature(name="input_a", modality=FeatureModality.NUMERICAL), + Feature(name="input_b", modality=FeatureModality.CATEGORICAL), + ] + ) + + with warnings.catch_warnings(): + warnings.simplefilter("error") + _warn_if_text_features(schema) + + def test__text_features__warn_with_column_names_and_remedies(self) -> None: + with pytest.warns(UserWarning, match="look like free text") as record: + _warn_if_text_features(_text_schema("review")) + + message = str(record[0].message) + # Column names are shown as the user wrote them, without the input_ prefix. + assert "'review'" in message + assert INPUT_FEATURE_PREFIX not in message + # The message must state all remedies. + assert "numeric dtype" in message + assert "https://github.com/PriorLabs/tabpfn-client" in message + assert "categorical_features_indices" in message + + def test__declared_categorical_indices__are_not_reported(self) -> None: + schema = _text_schema("sku", "review") + + with pytest.warns(UserWarning, match="look like free text") as record: + _warn_if_text_features(schema, declared_categorical_indices=[0]) + message = str(record[0].message) + assert "'review'" in message + assert "'sku'" not in message + + with warnings.catch_warnings(): + warnings.simplefilter("error") + _warn_if_text_features(schema, declared_categorical_indices=[0, 1]) + + def test__many_text_columns__message_is_truncated(self) -> None: + n_extra = 5 + n_columns = _MAX_TEXT_COLUMNS_IN_WARNING + n_extra + schema = _text_schema(*(f"t{i}" for i in range(n_columns))) + + with pytest.warns(UserWarning, match="look like free text") as record: + _warn_if_text_features(schema) + + message = str(record[0].message) + assert f"(and {n_extra} more)" in message + assert f"'t{_MAX_TEXT_COLUMNS_IN_WARNING - 1}'" in message + assert f"'t{_MAX_TEXT_COLUMNS_IN_WARNING}'" not in message + + +class TestDetectFeatureModalitiesWarnsOnText: + """`detect_feature_modalities` emits the text warning over real columns. + + The warning is now produced inside `detect_feature_modalities`, so these + exercise the whole path: which columns actually get labelled TEXT and thus + reach the warning, which the schema-level tests above cannot (they build + schemas by hand). + """ + + n_rows = 200 + + def _numeric_column(self) -> np.ndarray: + return np.random.default_rng(0).normal(size=self.n_rows) + + def _detect( + self, X: pd.DataFrame, declared: list[int] | None = None + ) -> FeatureSchema: + """Run modality detection over a frame, as `fit()` does.""" + return detect_feature_modalities( + X=X.to_numpy(dtype=object), + feature_names=list(X.columns), + provided_categorical_indices=declared, + min_samples_for_inference=100, + max_unique_for_category=30, + min_unique_for_numerical=4, + ) + + def test__free_text_column__warns(self) -> None: + X = pd.DataFrame( + { + "num": self._numeric_column(), + "review": [f"review {i}, a fairly long sentence" for i in range(200)], + } + ) + + with pytest.warns(UserWarning, match="look like free text") as record: + self._detect(X) + + assert "'review'" in str(record[0].message) + + def test__ordinary_columns__do_not_warn(self) -> None: + """Neither low-cardinality strings nor fully numeric strings are TEXT. + + The former are ordinary categoricals and the latter are + detected NUMERICAL. + """ + values = np.random.default_rng(1).normal(size=200) + X = pd.DataFrame( + { + "num": self._numeric_column(), + "color": ["red", "green", "blue"] * 66 + ["red", "red"], + "as_str": [str(round(float(v), 4)) for v in values], + } + ) + + with warnings.catch_warnings(): + warnings.simplefilter("error") + self._detect(X) + + def test__numeric_column_with_one_stray_token__warns(self) -> None: + """A single non-numeric token flips a whole numeric column to TEXT. + + `_is_numeric_pandas_series` requires *every* value to be coercible, so one + stray "N/A" makes the column ordinal-encoded as a near-unique categorical. + Warning here is the point of the feature: the fix is a numeric dtype. + """ + values = np.random.default_rng(2).normal(size=200) + mostly_numeric = [str(round(float(v), 4)) for v in values] + mostly_numeric[7] = "N/A" + X = pd.DataFrame( + {"num": self._numeric_column(), "mostly_numeric": mostly_numeric} + ) + + with pytest.warns(UserWarning, match="look like free text") as record: + self._detect(X) + + assert "'mostly_numeric'" in str(record[0].message) + + def test__declared_categorical_columns__do_not_warn(self) -> None: + """Declaring a column categorical states intent, so it must stay quiet. + + Covers both a plain string column and an explicit pandas `category` + dtype, each above the cardinality threshold. + """ + X = pd.DataFrame( + { + "num": self._numeric_column(), + "sku": [f"sku_{i % 60}" for i in range(200)], + "sku_cat": pd.Series( + [f"sku_{i % 60}" for i in range(200)], dtype="category" + ), + } + ) + declared = [1, 2] + + # Without the declaration the columns really are detected as TEXT and warn. + with pytest.warns(UserWarning, match="look like free text"): + self._detect(X) + + # Declaring them silences the warning; the columns are still labelled TEXT. + with warnings.catch_warnings(): + warnings.simplefilter("error") + schema = self._detect(X, declared) + assert schema.indices_for(FeatureModality.TEXT) == declared + + +@pytest.mark.parametrize("estimator_cls", [TabPFNClassifier, TabPFNRegressor]) +def test__fit_with_text_column__warns_at_call_site(estimator_cls: type) -> None: + """`fit` runs `detect_feature_modalities`, so a free-text column warns. + + Both estimators share the detection path, so one parametrized test pins the + estimator-level behaviour: `fit` emits the warning naming the column and + blaming this file's `fit` call (the stacklevel), declaring the column in + `categorical_features_indices` silences it, and `predict` stays quiet. + """ + n = 120 + rng = np.random.default_rng(seed=42) + X = pd.DataFrame( + { + "num": rng.normal(size=n), + "review": [f"review {i}, a fairly long sentence" for i in range(n)], + } + ) + y = ( + rng.integers(0, 2, size=n) + if estimator_cls is TabPFNClassifier + else rng.normal(size=n) + ) + + model = estimator_cls(n_estimators=1, device="cpu") + with pytest.warns(UserWarning, match="look like free text") as record: + model.fit(X, y) + assert "'review'" in str(record[0].message) + # Pins the stacklevel: the warning must blame this file's `fit` call, not a + # frame inside tabpfn or the contextlib wrapper around `fit`. + assert record[0].filename == __file__ + + # Only `fit` runs modality detection, so `predict` must not warn again. + # catch_warnings collects any warning instead of failing on unrelated ones. + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + model.predict(X) + assert not [w for w in caught if "look like free text" in str(w.message)] + + model = estimator_cls( + n_estimators=1, device="cpu", categorical_features_indices=[1] + ) + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + model.fit(X, y) + assert not [w for w in caught if "look like free text" in str(w.message)]