Skip to main content

zsfm_flowstate/
config.rs

1use serde::Deserialize;
2
3/// Top-level FlowState `config.json`.
4#[derive(Debug, Deserialize, Clone)]
5pub struct FlowStateConfig {
6    pub context_length: u32,
7    pub decoder_dim: u32,
8    pub decoder_patch_len: u32,
9    pub decoder_type: String,
10    pub embedding_feature_dim: u32,
11    pub encoder_num_hippo_blocks: u32,
12    pub encoder_num_layers: u32,
13    pub encoder_state_dim: u32,
14    pub quantiles: Vec<f32>,
15    #[serde(default = "default_bool_true")]
16    pub with_missing: bool,
17    #[serde(default = "default_bool_true")]
18    pub use_freq: bool,
19    #[serde(default = "default_bool_true")]
20    pub init_processing: bool,
21    #[serde(default = "default_u32_2048")]
22    pub min_context: u32,
23}
24
25impl FlowStateConfig {
26    pub fn from_json(s: &str) -> anyhow::Result<Self> {
27        Ok(serde_json::from_str(s)?)
28    }
29
30    pub fn n_quantiles(&self) -> u32 {
31        self.quantiles.len() as u32
32    }
33
34    /// Number of input channels (value + missing mask).
35    pub fn n_inputs(&self) -> u32 {
36        if self.with_missing { 2 } else { 1 }
37    }
38
39    /// Legendre basis range for "legs" / "hlegs" decoder.
40    pub fn basis_range(&self) -> [f32; 2] {
41        let dt = self.decoder_type.to_lowercase();
42        if dt == "hlegs" {
43            [0.0, 0.95]
44        } else {
45            [-1.0, 0.95]
46        }
47    }
48}
49
50fn default_bool_true() -> bool { true }
51fn default_u32_2048() -> u32 { 2048 }