1use serde::Deserialize;
2
3#[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 #[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 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 }