Skip to main content

zsfm_tabdpt/
config.rs

1/// TabDPT model configuration. One checkpoint (`Layer6/TabDPT`) serves both classification and
2/// regression — the head produces `max_num_classes + regression_bin_count` outputs; callers
3/// slice whichever half they need.
4#[derive(Clone, Debug)]
5pub struct TabDptConfig {
6    pub dim: usize,                  // 512 (emsize / ninp)
7    pub n_layers: usize,             // 32
8    pub n_heads: usize,              // 8
9    pub ff_dim: usize,               // 512 (nhid)
10    pub y_encoder_dim: usize,        // 128
11    pub max_num_classes: usize,      // 16 (n_out)
12    pub regression_bin_count: usize, // 2048
13    pub regression_bin_min: f32,     // -10.0
14    pub regression_bin_max: f32,     // 10.0
15    pub max_num_features: usize,     // 128
16    pub base_len: usize,             // 64 (min_eval_context)
17    pub max_len: usize,              // 1_048_576 (max_eval_context)
18    pub n_thinking_rows: usize,      // 64
19}
20
21impl TabDptConfig {
22    /// Matches `Layer6/TabDPT`'s `tabdpt1_2.safetensors` embedded config exactly.
23    pub fn default_v1_2() -> Self {
24        Self {
25            dim: 512,
26            n_layers: 32,
27            n_heads: 8,
28            ff_dim: 512,
29            y_encoder_dim: 128,
30            max_num_classes: 16,
31            regression_bin_count: 2048,
32            regression_bin_min: -10.0,
33            regression_bin_max: 10.0,
34            max_num_features: 128,
35            base_len: 64,
36            max_len: 1_048_576,
37            n_thinking_rows: 64,
38        }
39    }
40
41    /// `kappa = (sqrt(head_dim) - 1) / ln(max_len / base_len)`, used by every layer's attention
42    /// temperature scaling. `None` (scaling disabled) when `base_len == max_len`.
43    pub fn kappa(&self) -> Option<f64> {
44        if self.base_len == self.max_len {
45            return None;
46        }
47        let head_dim = (self.dim / self.n_heads) as f64;
48        Some((head_dim.sqrt() - 1.0) / (self.max_len as f64 / self.base_len as f64).ln())
49    }
50}