1#[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 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 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}