Skip to main content

zsfm_tabpfn/
config.rs

1/// TabPFN-3 (Prior-Labs/tabpfn_3) architecture config — matches the real checkpoint's embedded
2/// `config` dict exactly (classifier and regressor share this shape; only `task_type`/heads
3/// differ, and only classification is implemented here).
4#[derive(Clone, Debug)]
5pub struct TabPfnConfig {
6    pub embed_dim: usize,
7    pub dist_embed_num_blocks: usize,
8    pub dist_embed_num_heads: usize,
9    pub dist_embed_num_inducing_points: usize,
10    pub feature_group_size: usize,
11    pub feat_agg_num_blocks: usize,
12    pub feat_agg_num_heads: usize,
13    pub feat_agg_num_cls_tokens: usize,
14    pub nlayers: usize,
15    pub icl_num_heads: usize,
16    pub icl_num_kv_heads_test: Option<usize>,
17    pub decoder_head_dim: usize,
18    pub decoder_num_heads: usize,
19    pub decoder_use_softmax_scaling: bool,
20    pub ff_factor: usize,
21    pub softmax_scaling_mlp_hidden_dim: usize,
22    pub max_num_classes: usize,
23    pub use_nan_indicators: bool,
24}
25
26impl TabPfnConfig {
27    /// Values from the real `tabpfn-v3-classifier-v3_default.ckpt`'s embedded config.
28    pub fn v3_default() -> Self {
29        Self {
30            embed_dim: 128,
31            dist_embed_num_blocks: 3,
32            dist_embed_num_heads: 8,
33            dist_embed_num_inducing_points: 128,
34            feature_group_size: 3,
35            feat_agg_num_blocks: 3,
36            feat_agg_num_heads: 8,
37            feat_agg_num_cls_tokens: 4,
38            nlayers: 24,
39            icl_num_heads: 8,
40            icl_num_kv_heads_test: Some(1),
41            decoder_head_dim: 64,
42            decoder_num_heads: 6,
43            decoder_use_softmax_scaling: true,
44            ff_factor: 2,
45            softmax_scaling_mlp_hidden_dim: 64,
46            max_num_classes: 160,
47            use_nan_indicators: true,
48        }
49    }
50
51    pub fn icl_dim(&self) -> usize {
52        self.embed_dim * self.feat_agg_num_cls_tokens
53    }
54
55    /// `x_embed`'s input width: grouped raw values, doubled if NaN indicators are concatenated.
56    pub fn cell_in_features(&self) -> usize {
57        if self.use_nan_indicators {
58            self.feature_group_size * 2
59        } else {
60            self.feature_group_size
61        }
62    }
63}