1use serde::Deserialize;
2
3#[derive(Debug, Deserialize)]
6pub struct Chronos2Config {
7 pub d_model: u32,
8 #[serde(default = "default_u32::<64>")]
9 pub d_kv: u32,
10 pub d_ff: u32,
11 pub num_layers: u32,
12 pub num_heads: u32,
13 #[serde(default = "default_layer_norm_eps")]
14 pub layer_norm_epsilon: f64,
15 #[serde(default = "default_f64_ten_thousand")]
16 pub rope_theta: f64,
17 #[serde(default = "default_str_relu")]
18 pub feed_forward_proj: String,
19 pub chronos_config: ChronosInnerConfig,
20}
21
22#[derive(Debug, Deserialize)]
24#[allow(dead_code)]
25pub struct ChronosInnerConfig {
26 pub context_length: u32,
27 pub input_patch_size: u32,
28 pub output_patch_size: u32,
29 pub input_patch_stride: u32,
30 pub quantiles: Vec<f32>,
31 #[serde(default)]
32 pub use_reg_token: bool,
33 #[serde(default)]
34 pub use_arcsinh: bool,
35 #[serde(default = "default_u32::<1>")]
36 pub max_output_patches: u32,
37 pub time_encoding_scale: Option<u32>,
38}
39
40impl Chronos2Config {
41 pub fn from_json(s: &str) -> anyhow::Result<Self> {
42 Ok(serde_json::from_str(s)?)
43 }
44
45 pub fn time_encoding_scale(&self) -> u32 {
47 self.chronos_config
48 .time_encoding_scale
49 .unwrap_or(self.chronos_config.context_length)
50 }
51
52 pub fn dense_act_fn(&self) -> &str {
54 self.feed_forward_proj.split('-').last().unwrap_or("relu")
57 }
58}
59
60fn default_u32<const N: u32>() -> u32 { N }
61fn default_f64_ten_thousand() -> f64 { 10000.0 }
62fn default_layer_norm_eps() -> f64 { 1e-6 }
63fn default_str_relu() -> String { "relu".into() }