Skip to main content

zsfm_tirex/
config.rs

1/// TiRex model configuration (fixed from checkpoint hyper_parameters).
2#[derive(Clone)]
3pub struct TiRexConfig {
4    pub patch_size: usize,
5    pub num_blocks: usize,
6    pub embedding_dim: usize,
7    pub num_heads: usize,
8    pub input_ff_dim: usize,
9    pub ffn_up_dim: usize,
10    pub train_ctx_len: usize,
11    pub quantiles: Vec<f32>,
12    pub num_quantiles: usize,
13}
14
15impl TiRexConfig {
16    pub fn default_from_ckpt() -> Self {
17        let quantiles = vec![0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9];
18        let num_quantiles = quantiles.len();
19        Self {
20            patch_size: 32,
21            num_blocks: 12,
22            embedding_dim: 512,
23            num_heads: 4,
24            input_ff_dim: 2048,
25            ffn_up_dim: 1408, // round_up(512 * 2.6667, 64)
26            train_ctx_len: 2048,
27            quantiles,
28            num_quantiles,
29        }
30    }
31
32    pub fn head_dim(&self) -> usize {
33        self.embedding_dim / self.num_heads
34    }
35
36    pub fn num_patches(&self) -> usize {
37        self.train_ctx_len / self.patch_size
38    }
39
40    pub fn output_dim(&self) -> usize {
41        self.num_quantiles * self.patch_size
42    }
43
44    pub fn input_dim(&self) -> usize {
45        self.patch_size * 2 // values + mask
46    }
47}