Skip to content

Log usage of TabPFN estimators for accounts that opt in to analytics [ENG-890] - #1359

Open
safaricd wants to merge 6 commits into
mainfrom
ENG-890
Open

safaricd wants to merge 6 commits into
mainfrom
ENG-890

Conversation

@safaricd

@safaricd safaricd commented Oct 3, 2026 •

Copy link
Copy Markdown
Collaborator

Changes

  • log_usage logs each outermost call to an estimator's fit, predict and embedding methods as one usage event of data sizes and settings, never the data.
  • The 16 entry points of TabPFNClassifier, TabPFNRegressor and the fine-tuned estimators are decorated; calls they make internally (tuning holdouts, fine-tuning fits, forward) are not logged again.
  • A batched call (predict_proba_batched, predict_batched) is logged as one event with num_datasets.
  • Events log the GPU the estimator ran on (gpu_type), looked up once per device.
  • tabpfn/analytics/collector.py, modelled on PostHog's client, queues events in memory and sends them from a background thread in batches of up to 100, once 10 are waiting or every 5 s.
  • A batch the API does not take stays in memory and is sent again after 60 s, while later events wait in the queue (at most 10,000); nothing is written to disk.
  • The first logged call starts collecting if there is an API key, and the thread asks GET /account/telemetry before sending anything.
  • When that check says no or cannot answer, or POST /telemetry returns 403, collecting stops and the queued events are dropped.
  • check_telemetry_enabled now calls /account/telemetry, returns True/False/None like check_license_accepted, and never raises.
  • A forked process logs nothing, since on macOS its first network request crashes it.
  • The free-text warning's stacklevel steps past the log_usage wrapper, so it still names the user's fit call.
  • tests/conftest.py turns logging off for the test suite.

Motivation

Usage analytics for licensed accounts whose contract covers it, sent to gapi's POST /telemetry. It is on only if an API key exists (TABPFN_TOKEN, ~/.cache/tabpfn/auth_token or ~/.tabpfn/token) and GET /account/telemetry returns {"enabled": true}, which Prior Labs sets per account with the customer's consent. No environment variable turns it on or off; the API is the only switch. For everyone else nothing is sent, and no thread outlives the check. No call waits on the network, and exit waits at most 2 s. Requires gapi's num_datasets / optional num_rows schema change to be deployed before release.

How this was tested

tests/test_analytics.py: decorator and events (40 tests)
  • A predict that runs the logged forward is logged once.
  • A fit that fits other estimators internally is logged once.
  • Fit, predict and embedding events carry exactly the fields gapi's schema allows.
  • Two estimators used in turn are each logged from their own state.
  • Concurrent calls on threads are each logged from their own estimator.
  • Concurrent asyncio tasks are each logged once.
  • Only setting values gapi accepts are logged; others are logged as None.
  • Settings are logged as passed to the constructor, including torch.device objects.
  • Every logged setting is a constructor parameter of the TabPFN estimators.
  • Every logged call parameter is an argument of the method it is read from.
  • A default estimator logs each accepted setting unchanged in value and type.
  • Only published checkpoint names are logged; any other path is logged as other.
  • model_path="auto" logs the default model version.
  • A failed fit is logged as failed, without the previous fit's state.
  • A batched call is one prediction event with its number of datasets.
  • A batched call with mismatched shapes is logged without a shape.
  • Regressor predictions log output_type, quantiles and the regression task.
  • Fine-tuning fits are logged only on torchrun rank 0.
  • A sink that raises leaves the call's result unaffected.
  • A sink installed mid-call still does not log the nested calls.
  • The method is logged with the class that defines it, e.g. TabPFNClassifier.predict.
  • A fine-tuned estimator is logged from the estimator it fitted.
  • Quantiles are logged from arrays, and only between 0 and 1.
  • A failed call whose input size cannot be read is logged only for predict.
  • An embedding call with a data source gapi rejects is not logged.
  • A sink that calls a logged method does not log it again.
  • The TabPFN version is logged without a local build suffix.
  • gpu_type is the CUDA device's name, looked up once per device.
  • gpu_type is None on CPU and mps on Apple GPUs.
  • gpu_type is None when CUDA cannot start.
  • gpu_type is None for a device name gapi would reject.
  • Decorating a generator or coroutine is rejected.
  • Decorating a method without its data arguments is rejected.
  • A decorated method keeps its name and signature.
  • Shapes are read as rows and columns, and None where unreadable.
  • Exactly the 16 entry points are decorated.
  • A real TabPFNClassifier logs one event per entry-point call.
  • A real TabPFNClassifier with tuning_config logs one fit.
  • A real TabPFNRegressor logs one event per entry-point call.
  • A real fine-tuning run logs one fit for all of fine-tuning (slow).
tests/test_analytics_collector.py: delivery against a local fake gapi (25 tests)
  • _post returns the response status (204, 401, 403, 422, 503).
  • _post returns None when the connection is refused or the URL is malformed.
  • _post returns None when the API answers too slowly.
  • 250 events arrive in batches of at most 100.
  • Nine events wait, and the tenth sends all ten at once.
  • stop() sends queued events without waiting for the flush interval.
  • When the check says not enabled or cannot answer, the sink is removed, the queue is emptied and nothing is sent.
  • When the check raises, events are dropped without a traceback on the user's terminal.
  • stop() returns quietly when the thread never started.
  • Stopping before the check answers sends nothing, even once it answers.
  • A batch that got no response is kept and sent once the API answers.
  • A batch that got a 307, 408, 429, 500 or 503 is kept and sent again.
  • Later events wait behind a kept batch and are all delivered, the kept batch first.
  • Batches refused with 400, 401 or 422 are dropped after one request.
  • A 403 while sending stops collecting, removes the sink and drops the events, including those still in flight.
  • stop() while waiting to send a batch again returns at once.
  • stop() while a request hangs returns within its timeout.
  • A full queue drops events instead of blocking the caller.
  • An event that is not JSON (numpy int, NaN) is dropped alone.
  • The first logged call with an API key starts collecting, and later calls reuse its check.
  • Without an API key, the first call removes the sink and no request is made.
  • A slow check does not slow down the call that started it.
  • start() and stop() deliver a decorated estimator's events end to end.
  • A forked child logs nothing and exits cleanly, and the parent's events are sent once.
  • A child forked before the first logged call logs nothing and makes no request.
tests/test_browser_auth.py::TestCheckTelemetryEnabled (5 tests)
  • Only {"enabled": true} counts as enabled; other bodies are False and invalid JSON is None.
  • 401 and 403 are False, and 500 is None.
  • An unreachable server is None.
  • The URL is /account/telemetry, with or without a trailing slash on the base URL.
  • The answer is cached, so a process asks once.
Schema, regressions and static checks (6 checks)
  • 16 events from the current code (every entry point of a real classifier and regressor, batched calls and a failed fit) pass gapi's schema, alone and as one batch.
  • 62 events from an earlier run over all entry points, settings and checkpoint names passed gapi's schema once config fields matched it.
  • The full test suite had the same 108 failures before and after decorating the estimators (license-gated downloads).
  • Pyright reported the same 225 errors on the estimator modules before and after decorating them.
  • test__fit_with_text_column__transform_text_off__warns_at_call_site pins the free-text warning to the user's fit call through the wrapper.
  • pre-commit (ruff, ruff format, mypy) passes on every changed file, and Pyright reports no errors in tabpfn/analytics.
E2E smoke tests: fresh processes against a fake gapi, on this version (54 scenarios)

Background thread

  • Without an API key no thread starts, no request is made, and the sink is removed after the first call.
  • With analytics not enabled the thread ends about 20 ms after the first call, and nothing is sent.
  • A check answered with 401 ends the thread the same way.
  • A check answered with 403 ends the thread the same way.
  • A check answered with 500 ends the thread the same way.
  • An unreachable API ends the thread the same way.
  • A check that never answers leaves a daemon thread waiting and delays exit by at most 2.4 s.
  • With analytics enabled the thread runs until exit and every event is delivered.
  • A 403 on send mid-session ends the thread and removes the sink.

Latency

  • With the API answering, collecting adds about 47 µs per call (p99 73 µs), and the first call costs 0.3 ms.
  • With the API never answering, collecting adds about 50 µs per call (p99 129 µs).
  • With the check never answering, collecting adds about 49 µs per call (p99 108 µs).
  • Without an API key, a logged call costs under 1 µs more than an undecorated one.
  • Real fit+predict (about 110 ms) with the API answering differs by +4.7 ms collecting vs not, within noise.
  • Real fit+predict with the API never answering differs by −0.9 ms.
  • Real fit+predict with the API down differs by −0.3 ms.
  • Exit takes 0.38 s with the API answering.
  • Exit takes 0.37 s with analytics not enabled.
  • Exit takes 2.37 s with the API never answering.
  • Exit takes 2.4 s with the check never answering.
  • Exit takes 0.35 s without an API key.

Exceptions (results compared with a run without analytics)

  • Event building raising leaves results unchanged.
  • A sink raising leaves results unchanged.
  • Reading the API key raising leaves results unchanged.
  • The opt-in check raising leaves results unchanged and prints nothing.
  • The thread failing to start leaves results unchanged and exit quiet.
  • Sending raising leaves results unchanged, but the thread dies with a traceback on stderr (open).
  • An event that is not JSON leaves results unchanged.
  • The GPU-name lookup raising leaves results unchanged.
  • A user error in fit reaches the caller unchanged.
  • KeyboardInterrupt reaches the caller.

Concurrency

  • 64 threads × 20 fit+predict deliver 2,560 events exactly once.
  • 64 threads making their first call at once start one collector and one check.
  • 300 asyncio tasks on threads plus 300 on the loop deliver every event once.
  • Real estimators on 8 threads, then with n_preprocessing_jobs=4, deliver exactly 18 events.
  • 20 forks while another thread logs leave no child hung, and every parent event arrives once.
  • A fork while the first call holds the start lock leaves the child exiting cleanly.
  • A spawn multiprocessing.Pool delivers every worker's events.
  • A fork multiprocessing.Pool finishes, and its workers log nothing.
  • A forkserver pool finishes, and its workers log nothing.
  • A terminated spawn pool loses its workers' events (accepted, like SIGTERM).
  • joblib loky workers deliver every event.
  • A logged call that logs on a thread it starts counts both calls (known).

Unreachable API (events are kept in memory only)

  • With the API answering throughout, 392/392 events are delivered and no file is written.
  • A 2 s outage in a 10 s session loses 452 of 482 events with the default 60 s retry (open).
  • The same outage loses none with a 1 s retry.
  • An API flapping every 0.4 s for 8 s loses all 392 events with the default 60 s retry (open).
  • The same flapping loses all 394 events with a 1 s retry, as one hung request blocks sending for its 10 s timeout (open).
  • An API down for the whole session loses its 32 events at exit.
  • An API that never answers loses its 32 events, with a 2.4 s exit.
  • An API returning 503 for the whole session loses its 32 events.
  • An API returning 429 for the whole session loses its 32 events.
  • Offline at a process's first event drops its events, as opt-in is unconfirmed (by design).
  • Exiting before the check answers drops the events, as opt-in is unconfirmed (by design).
E2E smoke tests on the earlier on-disk version (47 scenarios)
  • The background-thread, latency, exceptions and concurrency scenarios above gave the same results, after fixing a thread traceback when the check raised, an atexit error when the thread could not start, a race that could re-install the sink after "not enabled", and a macOS crash and pool hang in forked processes.
  • With batches saved to disk, an API that was down, never answered, returned 503 or 429, or flapped lost no events: the next process delivered all of them.
Earlier harness runs, before delivery moved to memory only

Real-estimator break harness (16 scenarios; tests the decorator, unchanged since)

  • Every classifier entry point called once gives one event each.
  • A classifier with tuning_config gives one fit despite its hidden holdout fits.
  • Regressor output types and quantiles are logged, including quantiles given as a numpy array (fixed during testing).
  • Uniform, ragged and regressor batched calls give one event each.
  • Four threads with four estimators give eight correctly attributed events.
  • scikit-learn cross_val_score logs 3 fits and 3 predicts, also on 3 threads, and loky workers log their own calls.
  • Malformed inputs are logged as failed calls, and column names, values and labels never appear in events.
  • Configuration values and checkpoint names are logged as gapi expects.
  • Fine-tuning gives one fit despite its training, validation and final fits, and **kwargs options are logged (fixed during testing).
  • The torchrun rank filter logs fine-tuning only on rank 0.
  • A pickled and restored estimator keeps logging.
  • A logged call that starts a thread calling another logged method counts that call too (known; TabPFN itself never does this).
  • A sink that calls a logged method no longer logs recursively (fixed during testing).
  • The user's exceptions and KeyboardInterrupt propagate unchanged and are logged as failed.
  • Calls from asyncio tasks are each logged once.
  • Logging off adds 0.6 µs per call, and building one real event takes about 37 µs.

Delivery break harness (33 scenarios, against the earlier on-disk backlog)

  • 307 and 308 responses keep the batch instead of counting it delivered (fixed during testing).
  • 301 and 302 are followed by urllib as a GET, so they only matter for a misconfigured URL.
  • A connection closed without a response keeps the batch.
  • 200, 201 and 202 count as delivered.
  • One numpy value no longer loses the other events of its batch (fixed during testing).
  • With the API down, the other events of that batch are kept.
  • A NaN value no longer loses its batch (fixed during testing).
  • 16 threads × 600 events are delivered exactly once.
  • A 20k-event burst against a slow API never blocks collect() (worst 0.3 ms), and events beyond the 10k queue are dropped.
  • An unwritable save directory lost the events without crashing the thread or stop().
  • A save directory that is a file lost the events without raising.
  • stop() before start() returns quietly (fixed during testing).
  • Calling stop() twice is harmless.
  • A directory named *.json among saved batches no longer blocked the rest (fixed during testing).
  • An api_url without a scheme keeps the batch instead of losing it (fixed during testing).
  • 500 saved batches (50k events) were resent in 0.3 s locally.
  • Four threads saving into one directory kept it within the limit, with no temp files left.
  • A 403 while other threads produce stops collecting and keeps nothing (fixed during testing).
  • Ending with exit(), sys.exit(3), an uncaught exception or Ctrl-C delivers every event.
  • SIGTERM and os._exit() lose the queued events, as atexit does not run (accepted, as in PostHog).
  • Exiting while the API never answered saved every event and took 2.3 s.
  • The next process sent what the previous one saved.
  • Fork-based multiprocessing.Pool workers lose their events, as they exit without atexit (accepted).
  • Four processes resending one backlog at once delivered every event, with duplicates PostHog removes by event_id.

Root cause of the macOS fork crash

  • A standalone script showed that a forked process's first urllib request aborts it on macOS, even without torch, unless the parent made a request before forking.

🤖 Generated with Claude Code

safaricd and others added 3 commits October 1, 2026 22:35
Add `tabpfn.telemetry`, whose `log_usage` decorator turns each outermost
call to an estimator's fit, predict and embedding methods into a usage
event in the format of the telemetry API. Calls a logged method makes
internally, such as the holdout fits of `tuning_config` or the training
and validation fits of fine-tuning, are not logged on their own. A batched
call is one event with the number of datasets it scored.

Events record the sizes of the data and the estimator's settings, never
the data itself, and only settings and values the API accepts. They go
nowhere until a sink is installed with `set_sink`, which only happens for
accounts that have opted in; delivery to the API follows separately.

Wrap the entry points of TabPFNClassifier, TabPFNRegressor and the
fine-tuned estimators with it.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Add `tabpfn.telemetry.collector`, modelled on the queue of PostHog's
Python client: usage events go on a bounded in-memory queue, and a
background thread sends them in batches of up to 100, once 10 are
waiting or every 5 seconds. A batch that cannot be sent is saved under
~/.cache/tabpfn/telemetry and sent once the API answers again, from this
process or a later one; events keep their id, so PostHog deduplicates a
batch sent twice.

The first logged call starts collecting if there is an API key. The
background thread then asks `GET /account/telemetry` and sends or saves
nothing unless the account opted in; if it did not, the thread ends and
nothing more is logged. A 403 from `POST /telemetry` stops collecting the
same way. No call waits on the network, and exit waits at most 2 seconds.

Fix `check_telemetry_enabled` to call the route gapi defines, to return
None when it cannot tell, and never to raise. Log the GPU each call ran
on. A forked process logs nothing: on macOS its first network request
can crash it.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
safaricd and others added 2 commits October 5, 2026 09:24
The warning that a column looks like free text names the user's
`estimator.fit(X, y)` call by counting frames up to it. `log_usage` wraps
`fit` in one more frame, so the warning named the wrapper in
`tabpfn/telemetry/decorator.py` instead.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Saved batches were named after `time.time_ns()`, whose clock on Windows
ticks only every 15ms, so batches saved within one tick sorted by their
random suffix, and the oldest were not always the ones deleted at the
limit. Number the batches a process saves, after the time.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>

@brendan-priorlabs brendan-priorlabs left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thank you, @safaricd. Left a couple comments for you. Requesting changes just to be cautious

test_shapes = {_shape(X) for X in arguments["X_test_list"]}
fields: dict[str, Any] = {
# An empty list fails the call, and the API only accepts a positive count.
"num_datasets": len(arguments["X_train_list"]) or None,

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Claude thinks there's a bug here with "batched" requests. Seems serious. Implies the need for and e2e with the server in the loop.

  1. Batched predictions are always rejected by the server, and take other
    events down with them
  • events.py:120 puts a num_datasets field on batched events (from
    predict_proba_batched and predict_batched).
  • The server's event models don't define that field, and its base model is
    extra="forbid" (libs/prior/models/pydantic.py:13). So the server rejects the
    request with a 422.
  • The server accepts or rejects a batch as a whole
    (routers/telemetry/telemetry.py). On the client, _send (collector.py:540)
    treats a 4xx response as "done" and throws the batch away.
  • Result: every batched call silently discards up to 99 good events that were
    sent with it.
  • The client tests assert num_datasets exists (test_telemetry.py:497), so they
    never test against the server's schema.
  • Fix: add num_datasets to PredictCalled on the server, or drop it on the
    client. Also add a test that checks client events against the server's
    models

Comment thread src/tabpfn/analytics/events.py
Comment thread src/tabpfn/telemetry/collector.py Outdated
SHUTDOWN_TIMEOUT = 2.0
"""Seconds that sending the queued events may delay the exit of the process."""

SAVED_BATCHES_ROOT = CACHE_DIR / "telemetry"

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

How about we don't write anything to disk? We can use an exit handler to write everything at exit and set the batch size to be small enough that we don't lose too much.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Dropped - events not dumped to disk any longer but rather only using the in-memory queue with a lower flush period.

Comment thread src/tabpfn/telemetry/collector.py Outdated
if not enabled:
# The saved batches are kept if the API could not say, to be sent
# by a process that hears that the account opted in.
self._disable(keep_saved=enabled is None)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We should err on the side of caution. If the API can't say, we don't record anything.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Done ✅



@lru_cache(maxsize=1)
def check_telemetry_enabled(token: str, api_url: str) -> bool | None:

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I have a question here regarding opt-in/out, etc. Will follow up directly.

Rename `tabpfn.telemetry` to `tabpfn.analytics`, and its wording with
it. The API routes keep their names, `/telemetry` and
`/account/telemetry`.

Drop the batches saved to disk. A batch the API does not take is kept in
memory and sent again after `RETRY_DELAY`, while later events wait in the
bounded queue; events not sent when the process ends are dropped. A stop
ends the wait to send again, so that the exit is not delayed.

Record nothing when the API cannot say whether usage analytics is enabled
for the account, as when it says it is not.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
@safaricd safaricd changed the title Log usage of TabPFN estimators for accounts that opt in to telemetry [ENG-890] Log usage of TabPFN estimators for accounts that opt in to analytics [ENG-890] Oct 6, 2026

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants