Skip to main content

zsfm_tabfm/
config.rs

1use serde::Deserialize;
2
3/// `{task}/config.json` from `google/tabfm-1.0.0-pytorch`. Mirrors the upstream Python loader's
4/// `Config` dataclass (`tabfm/src/pytorch/tabfm_v1_0_0.py`) field-for-field.
5#[derive(Debug, Deserialize)]
6pub struct TabFMConfig {
7    pub embed_dim: u32,
8    pub max_classes: u32,
9    pub col_num_blocks: u32,
10    pub col_nhead: u32,
11    pub col_num_inds: u32,
12    pub row_num_blocks: u32,
13    pub row_nhead: u32,
14    pub row_num_cls: u32,
15    pub icl_num_blocks: u32,
16    pub icl_nhead: u32,
17    pub ff_factor: u32,
18    pub feature_group_size: u32,
19    pub is_classifier: bool,
20    #[serde(default = "default_num_freq")]
21    pub num_freq: u32,
22    /// `null` in the JSON means "use the `TabFM.__init__` default of `icl_dim * 2`".
23    #[serde(default)]
24    pub decoder_hidden: Option<u32>,
25    #[serde(default = "default_norm_eps")]
26    pub norm_eps: f64,
27}
28
29impl TabFMConfig {
30    pub fn from_json(s: &str) -> anyhow::Result<Self> {
31        Ok(serde_json::from_str(s)?)
32    }
33
34    /// `d_model` fed into the ICL stage: row_num_cls CLS-token embeddings, concatenated.
35    pub fn icl_dim(&self) -> u32 {
36        self.embed_dim * self.row_num_cls
37    }
38
39    pub fn col_dim_ff(&self) -> u32 {
40        self.embed_dim * self.ff_factor
41    }
42
43    pub fn icl_dim_ff(&self) -> u32 {
44        self.icl_dim() * self.ff_factor
45    }
46
47    pub fn decoder_hidden(&self) -> u32 {
48        self.decoder_hidden.unwrap_or(self.icl_dim() * 2)
49    }
50}
51
52fn default_num_freq() -> u32 { 32 }
53fn default_norm_eps() -> f64 { 1e-6 }