Skip to main content

zsfm_tabicl/
config.rs

1/// TabICL v2 configuration (`jingang/TabICL`'s `tabicl-classifier-v2-*.ckpt`). Matches the
2/// checkpoint's embedded `config` dict exactly. Classification only — regression's
3/// quantile-distribution head and the >10-class mixed-radix/hierarchical paths are out of
4/// scope (see crate docs).
5#[derive(Clone, Debug)]
6pub struct TabIclConfig {
7    pub max_classes: usize,     // 10
8    pub embed_dim: usize,       // 128
9    pub col_num_blocks: usize,  // 3
10    pub col_nhead: usize,       // 8
11    pub col_num_inds: usize,    // 128
12    pub feature_group_size: usize, // 3
13    pub row_num_blocks: usize,  // 3
14    pub row_nhead: usize,       // 8
15    pub row_num_cls: usize,     // 4
16    pub row_rope_base: f64,     // 100000
17    pub icl_num_blocks: usize,  // 12
18    pub icl_nhead: usize,       // 8
19    pub ff_factor: usize,       // 2
20}
21
22impl TabIclConfig {
23    pub fn v2() -> Self {
24        Self {
25            max_classes: 10,
26            embed_dim: 128,
27            col_num_blocks: 3,
28            col_nhead: 8,
29            col_num_inds: 128,
30            feature_group_size: 3,
31            row_num_blocks: 3,
32            row_nhead: 8,
33            row_num_cls: 4,
34            row_rope_base: 100_000.0,
35            icl_num_blocks: 12,
36            icl_nhead: 8,
37            ff_factor: 2,
38        }
39    }
40
41    pub fn icl_dim(&self) -> usize {
42        self.embed_dim * self.row_num_cls
43    }
44}