-
Notifications
You must be signed in to change notification settings - Fork 25
TabPFN interface for on-prem #342
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Merged
Changes from all commits
Commits
Show all changes
6 commits
Select commit
Hold shift + click to select a range
b75b89e
Include options for KV cache and model_path
safaricd e3f3a85
Minor fixes
safaricd e25c6fe
Leave sklearn interface intact
safaricd 03c33b1
Update client version
safaricd b261ff5
Merge branch 'main' of github.com:PriorLabs/tabpfn-client into ENG-871
safaricd 80968b9
Fix inconsistency in fit vs model_id reusability
safaricd File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,11 @@ | ||
| # Copyright (c) Prior Labs GmbH 2026. | ||
| # Licensed under the Apache License, Version 2.0 | ||
| """HTTP client for a self-hosted TabPFN inference container. | ||
|
|
||
| from tabpfn_client.hosted import TabPFNClassifier, TabPFNRegressor | ||
| """ | ||
|
|
||
| from tabpfn_client.hosted.estimator import TabPFNClassifier, TabPFNRegressor | ||
|
|
||
|
|
||
| __all__ = ["TabPFNClassifier", "TabPFNRegressor"] |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,277 @@ | ||
| # Copyright (c) Prior Labs GmbH 2026. | ||
| # Licensed under the Apache License, Version 2.0 | ||
| """scikit-learn estimators for a TabPFN inference container reached over HTTP. | ||
|
|
||
| Mirrors the `tabpfn_client.TabPFNClassifier` / `TabPFNRegressor` surface; | ||
| each `predict*` call POSTs to the user-supplied `endpoint_url` (the full | ||
| scoring URL, including the `/predict` path). `fit()` does not call the | ||
| endpoint — it just stores `X` / `y` on the estimator. The training data | ||
| is shipped to the endpoint on the next `predict*` call, where the actual | ||
| fit runs. | ||
|
|
||
| Auth is optional: when `api_key` is set it is sent as `Bearer <api_key>`, | ||
| otherwise no `Authorization` header is added. | ||
|
|
||
| Cached predict is opt-in via the `model_id` constructor argument. | ||
| When set (and `fit()` has not been called), `predict*` sends | ||
| `context.model_id` and omits the training data, letting the endpoint | ||
| reuse an already-fit model. Calling `fit()` on the estimator supersedes | ||
| the constructor `model_id` — the freshly-stored training data is | ||
| shipped on the next `predict*`, matching sklearn's re-fit semantics. | ||
| To make the endpoint build a reusable cache in the first | ||
| place, set `fit_mode="fit_with_cache"`; the endpoint then returns a | ||
| `model_id` in its response, captured on the fitted attribute | ||
| `self.model_id_` so callers can construct a follow-up estimator with | ||
| `model_id=prior.model_id_`. `model_path` is similarly optional and | ||
| only sent when set; some deployments reject overrides. | ||
|
|
||
| `model_id` lives on the constructor (not on `predict`) so the sklearn | ||
| estimator contract stays intact: `predict(X)` / `predict_proba(X)` work | ||
| with `Pipeline`, `GridSearchCV`, `cross_validate`, etc. | ||
| """ | ||
|
|
||
| from __future__ import annotations | ||
|
|
||
| from typing import Any, Dict, List, Literal, Optional | ||
|
|
||
| import httpx | ||
| import numpy as np | ||
| import pandas as pd | ||
| from sklearn.base import BaseEstimator, ClassifierMixin, RegressorMixin | ||
| from sklearn.utils.validation import check_is_fitted | ||
|
|
||
|
|
||
| def _to_jsonable(X: Any) -> list: | ||
| """Coerce numpy / pandas inputs to plain Python lists for JSON.""" | ||
| if isinstance(X, pd.DataFrame): | ||
| return X.values.tolist() | ||
| if isinstance(X, pd.Series): | ||
| return X.tolist() | ||
| if isinstance(X, list): | ||
| return X | ||
|
safaricd marked this conversation as resolved.
|
||
| return np.asarray(X).tolist() | ||
|
safaricd marked this conversation as resolved.
|
||
|
|
||
|
|
||
| class _HostedBase(BaseEstimator): | ||
| """Shared HTTP plumbing for the endpoint-backed TabPFN estimators.""" | ||
|
|
||
| _TASK: str = "" # overridden by subclasses | ||
|
|
||
| def __init__( | ||
| self, | ||
| endpoint_url: str, | ||
| api_key: Optional[str] = None, | ||
| extra_headers: Optional[Dict[str, str]] = None, | ||
| model_id: Optional[str] = None, | ||
| model_path: Optional[str] = None, | ||
| fit_mode: Optional[ | ||
| Literal["fit_preprocessors", "low_memory", "fit_with_cache", "batched"] | ||
| ] = None, | ||
| n_estimators: int = 8, | ||
| softmax_temperature: float = 0.9, | ||
| balance_probabilities: bool = False, | ||
| average_before_softmax: bool = False, | ||
| inference_precision: Literal["autocast", "auto"] = "auto", | ||
| random_state: Optional[int] = 0, | ||
| inference_config: Optional[Dict[str, Any]] = None, | ||
| n_preprocessing_jobs: int = 4, | ||
| memory_saving_mode: Optional[bool | Literal["auto"]] = None, | ||
| categorical_features_indices: Optional[List[int]] = None, | ||
| timeout_s: float = 300.0, | ||
| ): | ||
| self.endpoint_url = endpoint_url | ||
| self.api_key = api_key | ||
| self.extra_headers = extra_headers | ||
| self.model_id = model_id | ||
| self.model_path = model_path | ||
| self.fit_mode = fit_mode | ||
| self.n_estimators = n_estimators | ||
| self.softmax_temperature = softmax_temperature | ||
| self.balance_probabilities = balance_probabilities | ||
| self.average_before_softmax = average_before_softmax | ||
| self.inference_precision = inference_precision | ||
| self.random_state = random_state | ||
| self.inference_config = inference_config | ||
| self.n_preprocessing_jobs = n_preprocessing_jobs | ||
| self.memory_saving_mode = memory_saving_mode | ||
| self.categorical_features_indices = categorical_features_indices | ||
| self.timeout_s = timeout_s | ||
|
|
||
| def _build_tabpfn_config(self) -> Dict[str, Any]: | ||
| cfg: Dict[str, Any] = { | ||
| "n_estimators": self.n_estimators, | ||
| "softmax_temperature": self.softmax_temperature, | ||
| "average_before_softmax": self.average_before_softmax, | ||
| "inference_precision": self.inference_precision, | ||
| "random_state": self.random_state, | ||
| "inference_config": self.inference_config, | ||
| "n_preprocessing_jobs": self.n_preprocessing_jobs, | ||
| } | ||
| if self.memory_saving_mode is not None: | ||
| cfg["memory_saving_mode"] = self.memory_saving_mode | ||
| if self.categorical_features_indices is not None: | ||
| cfg["categorical_features_indices"] = self.categorical_features_indices | ||
| if self.model_path is not None: | ||
| cfg["model_path"] = self.model_path | ||
| if self.fit_mode is not None: | ||
| cfg["fit_mode"] = self.fit_mode | ||
| if self._TASK == "classification": | ||
| cfg["balance_probabilities"] = self.balance_probabilities | ||
| return cfg | ||
|
safaricd marked this conversation as resolved.
|
||
|
|
||
| def _headers(self) -> Dict[str, str]: | ||
| headers: Dict[str, str] = { | ||
| "Content-Type": "application/json", | ||
| "Accept": "application/json", | ||
| } | ||
| # The API key is optional because some deployments do not require it | ||
| if self.api_key is not None: | ||
| headers["Authorization"] = f"Bearer {self.api_key}" | ||
| if self.extra_headers: | ||
| headers.update(self.extra_headers) | ||
| return headers | ||
|
|
||
| def _http_client(self) -> httpx.Client: | ||
| # Cache the httpx.Client so repeated predict* calls reuse the TCP / | ||
| # TLS connection (keep-alive) instead of redoing the handshake on | ||
| # every request. | ||
| client = getattr(self, "_cached_client", None) | ||
| if client is not None: | ||
| return client | ||
| client = httpx.Client(timeout=self.timeout_s) | ||
| self._cached_client = client | ||
| return client | ||
|
|
||
| def __getstate__(self) -> Dict[str, Any]: | ||
| # httpx.Client isn't pickleable; strip the cache for sklearn pickling. | ||
| state = self.__dict__.copy() | ||
| state.pop("_cached_client", None) | ||
| return state | ||
|
safaricd marked this conversation as resolved.
|
||
|
|
||
| def fit(self, X: Any, y: Any) -> "_HostedBase": | ||
| X_arr = X if isinstance(X, pd.DataFrame) else np.asarray(X) | ||
| y_arr = y if isinstance(y, (pd.DataFrame, pd.Series)) else np.asarray(y) | ||
| if X_arr.shape[0] != y_arr.shape[0]: | ||
| raise ValueError( | ||
| f"X and y must have the same number of samples; " | ||
| f"got X={X_arr.shape}, y={y_arr.shape}" | ||
| ) | ||
| self.X_train_ = X_arr | ||
| self.y_train_ = y_arr | ||
| # New training data invalidates any cached model_id captured from a | ||
| # prior predict; otherwise `predict(X, model_id=self.model_id_)` would | ||
| # keep hitting the old server-side model. | ||
| self.model_id_ = None | ||
| if self._TASK == "classification": | ||
| self.classes_ = np.unique(y_arr) | ||
| return self | ||
|
safaricd marked this conversation as resolved.
|
||
|
|
||
| def _invoke( | ||
| self, | ||
| X_test: Any, | ||
| output_type: str, | ||
| predict_params: Optional[Dict[str, Any]] = None, | ||
| ) -> Dict[str, Any]: | ||
| params: Dict[str, Any] = {"output_type": output_type} | ||
| if predict_params: | ||
| params.update(predict_params) | ||
|
|
||
| body: Dict[str, Any] = { | ||
| "task_config": { | ||
| "task": self._TASK, | ||
| "tabpfn_config": self._build_tabpfn_config(), | ||
| "predict_params": params, | ||
| }, | ||
| "X_test": _to_jsonable(X_test), | ||
| } | ||
|
|
||
| # Precedence: a completed fit() wins over the constructor's model_id. | ||
| # Rationale — if the caller ran fit(new_X, new_y) on an estimator that | ||
| # was constructed with model_id set, they have re-trained; the cached | ||
| # id is now stale relative to the new data. Falls back to the cached | ||
| # path only when fit() was never called. | ||
| has_training_data = hasattr(self, "X_train_") and hasattr(self, "y_train_") | ||
| if has_training_data: | ||
| # y_train on the wire is 2D (n_samples, 1). | ||
| y_arr = np.asarray(self.y_train_) | ||
| if y_arr.ndim == 1: | ||
| y_arr = y_arr.reshape(-1, 1) | ||
| body["X_train"] = _to_jsonable(self.X_train_) | ||
| body["y_train"] = y_arr.tolist() | ||
|
cursor[bot] marked this conversation as resolved.
|
||
| elif self.model_id is not None: | ||
| body["context"] = {"model_id": self.model_id} | ||
| else: | ||
| # Neither path available — raise the standard sklearn error. | ||
| check_is_fitted(self, ["X_train_", "y_train_"]) | ||
|
|
||
| resp = self._http_client().post( | ||
| self.endpoint_url, | ||
| json=body, | ||
| headers=self._headers(), | ||
| ) | ||
| resp.raise_for_status() | ||
|
safaricd marked this conversation as resolved.
|
||
| payload = resp.json() | ||
|
|
||
| returned_id = payload.get("model_id") | ||
| if returned_id is not None: | ||
| self.model_id_ = returned_id | ||
|
|
||
| return payload | ||
|
|
||
|
|
||
| class TabPFNClassifier(_HostedBase, ClassifierMixin): | ||
| """TabPFN classifier backed by a self-hosted inference endpoint. | ||
|
|
||
| Example: | ||
| from tabpfn_client.hosted import TabPFNClassifier | ||
| clf = TabPFNClassifier( | ||
| endpoint_url="https://<your-endpoint>/predict", | ||
| api_key="<optional-bearer-token>", | ||
| ) | ||
| clf.fit(X_train, y_train) | ||
| clf.predict(X_test) | ||
| clf.predict_proba(X_test) | ||
| """ | ||
|
|
||
| _TASK = "classification" | ||
|
|
||
| def predict(self, X: Any) -> np.ndarray: | ||
| result = self._invoke(X, output_type="preds") | ||
| return np.asarray(result["prediction"]) | ||
|
|
||
| def predict_proba(self, X: Any) -> np.ndarray: | ||
| result = self._invoke(X, output_type="probas") | ||
| return np.asarray(result["prediction"]) | ||
|
|
||
|
|
||
| class TabPFNRegressor(_HostedBase, RegressorMixin): | ||
| """TabPFN regressor backed by a self-hosted inference endpoint. | ||
|
|
||
| Example: | ||
| from tabpfn_client.hosted import TabPFNRegressor | ||
| reg = TabPFNRegressor( | ||
| endpoint_url="https://<your-endpoint>/predict", | ||
|
safaricd marked this conversation as resolved.
|
||
| api_key="<optional-bearer-token>", | ||
| ) | ||
| reg.fit(X_train, y_train) | ||
| reg.predict(X_test) | ||
| reg.predict(X_test, output_type="quantiles", quantiles=[0.1, 0.5, 0.9]) | ||
| """ | ||
|
|
||
| _TASK = "regression" | ||
|
|
||
| def predict( | ||
| self, | ||
| X: Any, | ||
| output_type: str = "mean", | ||
| quantiles: Optional[list] = None, | ||
| ) -> np.ndarray: | ||
| predict_params: Dict[str, Any] = {} | ||
| if quantiles is not None: | ||
| predict_params["quantiles"] = quantiles | ||
| result = self._invoke( | ||
| X, | ||
| output_type=output_type, | ||
| predict_params=predict_params, | ||
| ) | ||
| return np.asarray(result["prediction"]) | ||
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
Uh oh!
There was an error while loading. Please reload this page.