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}