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}