Looking for usage documentation? Check out Interpretability and Classification.
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 decoder_readout.class_vote
View source 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.
decoder_readout.plot_decoder_readout
View source 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.