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
1 change: 1 addition & 0 deletions changelog/369.fixed.md
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
`TabPFNRegressor.predict(output_type="full")` now matches a local prediction: masked-out bars come back as `-inf` rather than `NaN`, and test sets above the server's full-output row cap are split across calls instead of raising.
150 changes: 126 additions & 24 deletions src/tabpfn_client/estimator.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@
from tabpfn_client.utils import model_limit_from_version, model_version_from_path
from tabpfn_client.service_wrapper import InferenceClient
from tabpfn_client.api_models import (
ModelLimit,
RegressorTabPFNConfig,
ClassifierTabPFNConfig,
RegressorPredictParams,
Expand Down Expand Up @@ -794,15 +795,25 @@ def predict(
predict_params=predict_params,
)

# A full-output response carries one logit per histogram bar for every
# test row, so the server caps the rows a single response may cover.
# Splitting the call here keeps the returned arrays identical to what
# one unrestricted call would have produced.
rows_per_call = (
_full_output_row_limit(tabpfn_config.model_path)
if output_type == "full"
else None
)
chunked = rows_per_call is not None and X.shape[0] > rows_per_call
validate_test_set(
X,
output_type,
tabpfn_config.model_path,
train_rows=self._last_train_X.shape[0]
if self._last_train_X is not None
else None,
split_full_output=chunked,
)
X_clean = _clean_text_features(X)

# NOTE(@trace_id)
# If this instance reuses a previous fit via a directly-assigned
Expand All @@ -816,34 +827,58 @@ def predict(
):
self.client_options.headers["sentry-trace"] = self._last_trace_id

def predict_task() -> PredictionResult:
return InferenceClient.predict(
X_clean,
fitted_train_set_id=self.model_id_,
task_config=task_config,
client_options=self.client_options,
)
def predict_rows(X_rows: Any) -> PredictionResult:
X_clean = _clean_text_features(X_rows)

result = run_task(predict_task, "Predicting")
# Unpack and store metadata
self._last_meta = result.metadata
def predict_task() -> PredictionResult:
return InferenceClient.predict(
X_clean,
fitted_train_set_id=self.model_id_,
task_config=task_config,
client_options=self.client_options,
)

return run_task(predict_task, "Predicting")

if chunked:
rows_per_call = cast(int, rows_per_call)
results = [
predict_rows(_row_slice(X, start, start + rows_per_call))
for start in range(0, X.shape[0], rows_per_call)
]
# Metadata describes the request, and every chunk shares the same
# config; the last one stands for the whole prediction.
self._last_meta = results[-1].metadata
output = _merge_full_outputs(
[cast("dict[str, np.ndarray]", r.y_pred) for r in results]
)
else:
result = predict_rows(X)
# Unpack and store metadata
self._last_meta = result.metadata
output = result.y_pred

output = result.y_pred
if output_type == "quantiles" and isinstance(output, np.ndarray):
return list(output) if output.ndim == 2 else [output]
if output_type == "full":
# `criterion` is a bar distribution rather than an array, so the
# full output is looser than the declared `dict[str, np.ndarray]`.
full = cast("dict[str, Any]", output)
if "logits" in full:
full["logits"] = _restore_masked_logits(full["logits"])
try:
from tabpfn.regressor import FullSupportBarDistribution # type: ignore
import torch # type: ignore

output["criterion"] = FullSupportBarDistribution(
borders=torch.tensor(output["borders"])
full["criterion"] = FullSupportBarDistribution(
borders=torch.tensor(full["borders"])
)
except ImportError:
logger.warning(
"Optional dependencies 'tabpfn' and 'torch' are required to "
"construct the criterion when output_type='full'. Skipping criterion."
)
return full

return output

Expand Down Expand Up @@ -906,23 +941,39 @@ def validate_train_set(
)


def _limit_for_model_path(model_path: str | None) -> ModelLimit | None:
"""Return the row/cell caps for `model_path`, or None if unknown."""
api_settings = ServiceClient.get_settings()
if api_settings is None:
return None
if not model_path:
return api_settings.model_limits[api_settings.default_model_version]
model_version = model_version_from_path(model_path)
return model_limit_from_version(model_version, api_settings.model_limits)


def _full_output_row_limit(model_path: str | None) -> int | None:
"""Rows a single `output_type="full"` response may cover, None if unknown."""
limit = _limit_for_model_path(model_path)
return limit.test_set_max_rows_w_full_regression_output if limit else None


def validate_test_set(
X: pd.DataFrame | np.ndarray,
output_type: str | None,
model_path: str | None = None,
train_rows: int | None = None,
split_full_output: bool = False,
):
"""Check the integrity of the test data."""
"""Check the integrity of the test data.

api_settings = ServiceClient.get_settings()
if api_settings is None:
return
`split_full_output` marks that the caller will honour the full-output row
cap by splitting the request, so that cap is not enforced here.
"""

if not model_path:
limit = api_settings.model_limits[api_settings.default_model_version]
else:
model_version = model_version_from_path(model_path)
limit = model_limit_from_version(model_version, api_settings.model_limits)
limit = _limit_for_model_path(model_path)
if limit is None:
return

max_rows = limit.test_set_max_rows
if train_rows:
Expand All @@ -944,14 +995,65 @@ def validate_test_set(
f"The number of test cells ({n_cells}) exceeds the maximum of {limit.test_set_max_cells}. "
"Split the test set across multiple calls to reduce the number of cells."
)
if output_type == "full":
if output_type == "full" and not split_full_output:
if X.shape[0] > limit.test_set_max_rows_w_full_regression_output:
raise ValueError(
f"The number of test rows ({X.shape[0]}) exceeds the maximum of {limit.test_set_max_rows_w_full_regression_output} "
"for full regression output."
)


def _restore_masked_logits(logits: np.ndarray) -> np.ndarray:
"""Return `logits` with bars outside a row's support back at -inf.

The response encoding carries no representation for -inf, so those bars
arrive as null and land in the array as NaN.
"""
logits = np.asarray(logits, dtype=float)
return np.where(np.isnan(logits), -np.inf, logits)


def _row_slice(X: Any, start: int, stop: int) -> Any:
"""Return rows `[start:stop)` of `X`, keeping its container type."""
if isinstance(X, pd.DataFrame):
return X.iloc[start:stop]
return X[start:stop]


# Axis along which each `output_type="full"` array runs over test rows. None
# marks a row-independent array: `borders` describes the histogram of the
# fitted target, so every chunk returns the same one.
_FULL_OUTPUT_ROW_AXIS: dict[str, int | None] = {
"mean": 0,
"median": 0,
"mode": 0,
"logits": 0,
"quantiles": 1,
"borders": None,
}


def _merge_full_outputs(parts: list[dict[str, np.ndarray]]) -> dict[str, np.ndarray]:
"""Stitch per-chunk full-output predictions back into one result."""
unknown = sorted(set(parts[0]) - set(_FULL_OUTPUT_ROW_AXIS))
if unknown:
raise RuntimeError(
f"Cannot combine full regression output across calls: the server "
f"returned unrecognised field(s) {unknown}. Upgrade tabpfn-client, "
f"or split the test set yourself and merge the results."
)
merged: dict[str, np.ndarray] = {}
for key, axis in _FULL_OUTPUT_ROW_AXIS.items():
if key not in parts[0]:
continue
merged[key] = (
parts[0][key]
if axis is None
else np.concatenate([p[key] for p in parts], axis=axis)
)
return merged


@overload
def _clean_text_features(X: pd.DataFrame) -> pd.DataFrame: ...
@overload
Expand Down
129 changes: 129 additions & 0 deletions tests/unit/test_tabpfn_regressor.py
Original file line number Diff line number Diff line change
Expand Up @@ -982,3 +982,132 @@ def test_paper_version_behavior(self, mock_predict, mock_fit):
tabpfn_false.fit(X, y)
y_pred_false = tabpfn_false.predict(test_X)
self.assertIsNotNone(y_pred_false)


class TestFullOutputChunking(unittest.TestCase):
"""`output_type="full"` splits into per-chunk calls above the server cap."""

FULL_OUTPUT_MAX_ROWS = 4
N_BARS = 6
N_QUANTILES = 9

def setUp(self):
# skip init
config.Config.is_initialized = True
payload = _api_settings_payload()
for limit in [payload["max_model_limit"], *payload["model_limits"].values()]:
limit["test_set_max_rows_w_full_regression_output"] = (
self.FULL_OUTPUT_MAX_ROWS
)
ServiceClient._api_settings = GetSettingsResponse(**payload)
ServiceClient._api_settings_ts = time.monotonic()

self.borders = np.linspace(0.0, 1.0, self.N_BARS + 1)
self.regressor = TabPFNRegressor(model_path="v2.5_default")
self.regressor.model_id_ = UUID("00000000-0000-0000-0000-000000000000")
self.regressor._last_train_X = np.random.randn(5, 2)

def tearDown(self):
ServiceClient._api_settings = None
ServiceClient._api_settings_ts = 0.0
# undo setUp
config.reset()

def _full_output(self, n_rows: int, offset: int) -> dict[str, np.ndarray]:
"""A full-output payload whose values encode each row's global index."""
rows = np.arange(offset, offset + n_rows, dtype=float)
return {
"mean": rows,
"median": rows + 0.1,
"mode": rows + 0.2,
"quantiles": rows + np.arange(self.N_QUANTILES)[:, None],
"logits": rows[:, None] + np.arange(self.N_BARS)[None, :],
"borders": self.borders,
}

def _predict(self, n_test_rows: int):
"""Predict `n_test_rows` rows, serving each request from its own chunk."""
served = []

def fake_predict(X, **kwargs):
offset = sum(len(x) for x in served)
served.append(X)
return PredictionResult(
y_pred=cast(Any, self._full_output(len(X), offset)), metadata={}
)

with patch.object(InferenceClient, "predict", side_effect=fake_predict):
output = cast(
"dict[str, Any]",
self.regressor.predict(
np.random.randn(n_test_rows, 2), output_type="full"
),
)
return output, served

def test_stays_a_single_call_at_the_limit(self):
_, served = self._predict(self.FULL_OUTPUT_MAX_ROWS)
self.assertEqual([len(x) for x in served], [self.FULL_OUTPUT_MAX_ROWS])

def test_splits_into_chunks_of_at_most_the_limit(self):
_, served = self._predict(self.FULL_OUTPUT_MAX_ROWS * 2 + 1)
self.assertEqual(
[len(x) for x in served],
[self.FULL_OUTPUT_MAX_ROWS, self.FULL_OUTPUT_MAX_ROWS, 1],
)

def test_merged_output_matches_a_single_unsplit_call(self):
n_rows = self.FULL_OUTPUT_MAX_ROWS * 2 + 1
output, _ = self._predict(n_rows)
expected = self._full_output(n_rows, 0)
for key, value in expected.items():
np.testing.assert_allclose(output[key], value, err_msg=key)

def test_rejects_output_it_cannot_merge(self):
def fake_predict(X, **kwargs):
payload = self._full_output(len(X), 0)
payload["surprise"] = np.zeros(len(X))
return PredictionResult(y_pred=cast(Any, payload), metadata={})

with patch.object(InferenceClient, "predict", side_effect=fake_predict):
with self.assertRaisesRegex(RuntimeError, "surprise"):
self.regressor.predict(
np.random.randn(self.FULL_OUTPUT_MAX_ROWS + 1, 2),
output_type="full",
)

def test_masked_bars_come_back_as_negative_infinity(self):
"""Null logits mark bars outside a row's support, i.e. -inf."""

def fake_predict(X, **kwargs):
payload = self._full_output(len(X), 0)
payload["logits"] = payload["logits"].copy()
payload["logits"][:, -2:] = np.nan
return PredictionResult(y_pred=cast(Any, payload), metadata={})

with patch.object(InferenceClient, "predict", side_effect=fake_predict):
output = cast(
"dict[str, Any]",
self.regressor.predict(
np.random.randn(self.FULL_OUTPUT_MAX_ROWS + 1, 2),
output_type="full",
),
)

logits = output["logits"]
self.assertFalse(np.isnan(logits).any())
self.assertTrue(np.isneginf(logits[:, -2:]).all())
self.assertTrue(np.isfinite(logits[:, :-2]).all())

def test_other_output_types_keep_the_row_limit(self):
"""The cap only applies to `full`, so `mean` is never split."""
served = []

def fake_predict(X, **kwargs):
served.append(X)
return PredictionResult(y_pred=np.zeros(len(X)), metadata={})

n_rows = self.FULL_OUTPUT_MAX_ROWS * 3
with patch.object(InferenceClient, "predict", side_effect=fake_predict):
self.regressor.predict(np.random.randn(n_rows, 2), output_type="mean")
self.assertEqual([len(x) for x in served], [n_rows])
Loading