Skip to main content

zsfm_chronos/
config.rs

1use serde::Deserialize;
2
3/// Top-level Chronos-2 `config.json`.
4/// The `chronos_config` key holds a nested dict with forecasting-specific settings.
5#[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/// Nested `chronos_config` dict inside `config.json`.
23#[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    /// Effective time_encoding_scale: falls back to context_length if not set.
46    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    /// Dense activation function name (e.g. "relu").
53    pub fn dense_act_fn(&self) -> &str {
54        // feed_forward_proj may be "relu", "gelu", or "gated-gelu" etc.
55        // Chronos-2 asserts not gated, so just split on '-' and take the last part.
56        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() }