Skip to main content

zsfm_toto/
config.rs

1use serde::Deserialize;
2
3/// Mirrors the relevant fields from Datadog/Toto-2.0 config.json.
4/// Accepts the Python model field names (d_model, num_layers, …).
5/// Unknown fields are ignored so we stay forward-compatible.
6#[derive(Debug, Deserialize)]
7pub struct TotoConfig {
8    /// Number of transformer layers.
9    #[serde(alias = "num_layers", default = "default_u32::<48>")]
10    pub num_hidden_layers: u32,
11
12    /// Hidden / embedding dimension.
13    #[serde(alias = "d_model", default = "default_u32::<2048>")]
14    pub hidden_size: u32,
15
16    /// Number of query attention heads.
17    #[serde(alias = "num_heads", default = "default_u32::<32>")]
18    pub num_attention_heads: u32,
19
20    /// Number of KV groups (num_groups in Python = GQA groups).
21    #[serde(alias = "num_groups", default = "default_u32::<32>")]
22    pub num_key_value_heads: u32,
23
24    /// Per-head QK dimension.
25    #[serde(alias = "qk_dim", default = "default_u32::<64>")]
26    pub head_dim: u32,
27
28    /// Patch size in timesteps.
29    #[serde(default = "default_u32::<32>")]
30    pub patch_size: u32,
31
32    /// Number of output quantile levels.
33    #[serde(default = "default_u32::<9>")]
34    pub num_quantiles: u32,
35}
36
37impl TotoConfig {
38    pub fn from_json(s: &str) -> anyhow::Result<Self> {
39        Ok(serde_json::from_str(s)?)
40    }
41}
42
43fn default_u32<const N: u32>() -> u32 { N }