Skip to main content
Looking for usage documentation? Check out Interpretability and Classification.
Read out TabPFN’s classification head as a label-vote over training rows. TabPFN classifies with an attention-based retrieval head (ManyClassDecoder): each test row attends to the training rows, and the prediction is the average of their one-hot labels weighted by that attention. The prediction is therefore a weighted vote, and P(class c) for a test row is the sum of its attention weights over the training rows whose label is c. get_decoder_readout recovers those per-training-row attention weights, so you can see which training points drive a prediction and by how much. The weights come from the head’s own ManyClassDecoder.attention_weights (tabpfn>=8.3.0), read off a forward pre-hook during predict; forward itself fuses the attention into a single kernel and never materializes them. For each test row the weights sum to 1 (averaged over the decoder’s attention heads and over the ensemble members). Collapsing them by training label with class_vote reproduces the model’s predict_proba up to the head’s log-clamping when softmax_temperature=1.0 and balance_probabilities=False. Both are applied to the decoder’s logits after this readout, so at the library default softmax_temperature=0.9 the temperature sharpens the vote (per estimator, predict_proba ∝ vote ** (1 / T)); predict_proba then differs from the vote by up to ~2 percentage points for binary and ~6 at 10 classes, and balance_probabilities=True widens it further. Only the local tabpfn backend is supported: the client/API backend does not expose the model internals this reads from. Row subsampling (TabPFNClassifier(..., subsample_samples=...)) is not supported, since the weight columns would no longer align to a single set of training rows.

decoder_readout.get_decoder_readout

View source
Extract the decoder-head attention weights over training rows.
Parameters
Returns

decoder_readout.class_vote

View source
Collapse per-row readout weights into a per-class vote. Sums the attention weights within each training label, turning the readout into a class distribution. Averaged over the ensemble, this reproduces the model’s predict_proba up to the head’s log-clamping when softmax_temperature=1.0 and balance_probabilities=False. Both are applied downstream of this readout, so at the library default softmax_temperature=0.9 the vote is sharpened (per estimator, predict_proba ∝ vote ** (1 / T)), differing by up to ~2 percentage points for binary and ~6 at 10 classes. It is also exact only in full precision; reduced precision costs ~1e-2 relative, so fit with inference_precision=torch.float32 for a tight match.
Parameters
Returns

decoder_readout.plot_decoder_readout

View source
Draw the decoder readout for a set of queries over a 2D projection. Each panel places one query and its top_k most-attended training rows on a 2D projection of the rows, drawing a line from the query to each attended row colored by the row’s class and scaled by its vote weight; the query star is colored by the predicted class (argmax of the class votes). The projection uses embeddings (a (train_vecs, test_vecs) pair, e.g. TabPFN’s target-conditioned embeddings from get_embeddings) when given, else UMAP (falling back to PCA) over the raw train_features/test_features. Contrasting the two shows what the head keys on: distance in the embedding space, where votes concentrate on the query’s own class, versus the raw feature space, where that locality is weaker. Works for any number of classes. All test-row arrays (weights, test_features, the test vectors in embeddings, y_test) span the full test set; queries indexes into them to select the rows to draw.
Parameters
Returns