Skip to main content

zsfm_sundial/
config.rs

1use serde::Deserialize;
2
3#[derive(Debug, Clone, Deserialize)]
4pub struct SundialConfig {
5    pub hidden_size: usize,
6    pub intermediate_size: usize,
7    pub num_hidden_layers: usize,
8    pub num_attention_heads: usize,
9    pub input_token_len: usize,
10    #[serde(default = "default_output_token_lens")]
11    pub output_token_lens: Vec<usize>,
12    pub rope_theta: f64,
13    #[serde(default = "default_flow_depth")]
14    pub flow_loss_depth: usize,
15    #[serde(default = "default_sampling_steps")]
16    pub num_sampling_steps: usize,
17}
18
19fn default_output_token_lens() -> Vec<usize> { vec![720] }
20fn default_flow_depth() -> usize { 3 }
21fn default_sampling_steps() -> usize { 50 }
22
23impl SundialConfig {
24    pub fn head_dim(&self) -> usize { self.hidden_size / self.num_attention_heads }
25    pub fn output_token_len(&self) -> usize { self.output_token_lens[0] }
26}
27
28impl Default for SundialConfig {
29    fn default() -> Self {
30        Self {
31            hidden_size: 768,
32            intermediate_size: 3072,
33            num_hidden_layers: 12,
34            num_attention_heads: 12,
35            input_token_len: 16,
36            output_token_lens: vec![720],
37            rope_theta: 10000.0,
38            flow_loss_depth: 3,
39            num_sampling_steps: 50,
40        }
41    }
42}