Skip to main content

zsfm_ttm/infer/
mod.rs

1//! TinyTimeMixer inference engine.
2//!
3//! Architecture: MLP-Mixer with adaptive patching.
4//! - StdScaler → Patchify → Linear patcher → 3 adaptive levels → Linear adapter →
5//!   2 decoder layers → Flatten → Linear head → inverse scale
6//! - No attention. Each "mixer layer" = PatchMixerBlock + FeatureMixerBlock.
7//! - GatedAttention: softmax(linear(x)) * x applied after each MLP.
8
9use std::io::{BufReader, Read, Seek};
10use std::path::{Path, PathBuf};
11
12use anyhow::{Context, Result};
13use candle_core::quantized::gguf_file;
14use candle_core::{DType, Device, Tensor};
15use zsfm_nn::{layer_norm, linear};
16
17use crate::config::TtmConfig;
18
19// ---------------------------------------------------------------------------
20// Weight structs
21// ---------------------------------------------------------------------------
22
23struct MixerLayer {
24    patch_norm_w: Tensor,
25    patch_norm_b: Tensor,
26    patch_fc1_w: Tensor,
27    patch_fc1_b: Tensor,
28    patch_fc2_w: Tensor,
29    patch_fc2_b: Tensor,
30    patch_gate_w: Tensor,
31    patch_gate_b: Tensor,
32    feat_norm_w: Tensor,
33    feat_norm_b: Tensor,
34    feat_fc1_w: Tensor,
35    feat_fc1_b: Tensor,
36    feat_fc2_w: Tensor,
37    feat_fc2_b: Tensor,
38    feat_gate_w: Tensor,
39    feat_gate_b: Tensor,
40}
41
42struct AdaptiveLevel {
43    factor: usize,
44    layers: Vec<MixerLayer>,
45}
46
47pub struct TtmModel {
48    device: Device,
49    config: TtmConfig,
50    patcher_w: Tensor,
51    patcher_b: Tensor,
52    enc_levels: Vec<AdaptiveLevel>,
53    dec_adapter_w: Tensor,
54    dec_adapter_b: Tensor,
55    dec_layers: Vec<MixerLayer>,
56    head_w: Tensor,
57    head_b: Tensor,
58}
59
60// ---------------------------------------------------------------------------
61// Builder
62// ---------------------------------------------------------------------------
63
64/// Fluent constructor for [`TtmModel`]. `TtmConfig` (from a parsed `config.json`) doubles as
65/// this model's load-time config — there's no separate `InferConfig` to map onto.
66///
67/// ```no_run
68/// use zsfm_ttm::TtmModel;
69///
70/// # fn main() -> anyhow::Result<()> {
71/// let model = TtmModel::builder("ttm.gguf")
72///     .config_json(&std::fs::read_to_string("config.json")?)?
73///     .build()?;
74/// # Ok(()) }
75/// ```
76pub struct TtmModelBuilder {
77    gguf_path: PathBuf,
78    config: Option<TtmConfig>,
79}
80
81impl TtmModelBuilder {
82    fn new(gguf_path: impl Into<PathBuf>) -> Self {
83        Self { gguf_path: gguf_path.into(), config: None }
84    }
85
86    pub fn config(mut self, config: TtmConfig) -> Self {
87        self.config = Some(config);
88        self
89    }
90
91    pub fn config_json(mut self, s: &str) -> Result<Self> {
92        self.config = Some(TtmConfig::from_json(s)?);
93        Ok(self)
94    }
95
96    pub fn build(self) -> Result<TtmModel> {
97        let config = self
98            .config
99            .context("TtmModelBuilder: no config set — call .config(...) or .config_json(...)")?;
100        TtmModel::load(&self.gguf_path, config)
101    }
102}
103
104// ---------------------------------------------------------------------------
105// GGUF loading
106// ---------------------------------------------------------------------------
107
108fn load_t(
109    content: &gguf_file::Content,
110    reader: &mut (impl Read + Seek),
111    name: &str,
112    device: &Device,
113) -> Result<Tensor> {
114    zsfm_nn::load_tensor(content, reader, name, device, DType::F32)
115}
116
117fn load_mixer_layer(
118    content: &gguf_file::Content,
119    reader: &mut (impl Read + Seek),
120    prefix: &str,
121    device: &Device,
122) -> Result<MixerLayer> {
123    let mut t = |s: &str| -> Result<Tensor> {
124        load_t(content, reader, &format!("{prefix}.{s}"), device)
125    };
126    Ok(MixerLayer {
127        patch_norm_w: t("patch_norm.weight")?,
128        patch_norm_b: t("patch_norm.bias")?,
129        patch_fc1_w:  t("patch_fc1.weight")?,
130        patch_fc1_b:  t("patch_fc1.bias")?,
131        patch_fc2_w:  t("patch_fc2.weight")?,
132        patch_fc2_b:  t("patch_fc2.bias")?,
133        patch_gate_w: t("patch_gate.weight")?,
134        patch_gate_b: t("patch_gate.bias")?,
135        feat_norm_w:  t("feat_norm.weight")?,
136        feat_norm_b:  t("feat_norm.bias")?,
137        feat_fc1_w:   t("feat_fc1.weight")?,
138        feat_fc1_b:   t("feat_fc1.bias")?,
139        feat_fc2_w:   t("feat_fc2.weight")?,
140        feat_fc2_b:   t("feat_fc2.bias")?,
141        feat_gate_w:  t("feat_gate.weight")?,
142        feat_gate_b:  t("feat_gate.bias")?,
143    })
144}
145
146impl TtmModel {
147    /// Start building a [`TtmModel`] — see [`TtmModelBuilder`].
148    pub fn builder(gguf_path: impl Into<PathBuf>) -> TtmModelBuilder {
149        TtmModelBuilder::new(gguf_path)
150    }
151
152    pub fn load(gguf_path: &Path, config: TtmConfig) -> Result<Self> {
153        let device = Device::Cpu;
154        let file = std::fs::File::open(gguf_path)
155            .with_context(|| format!("open {}", gguf_path.display()))?;
156        let mut reader = BufReader::with_capacity(zsfm_gguf::READ_BUF_CAPACITY, file);
157        let content = gguf_file::Content::read(&mut reader).context("parse GGUF header")?;
158
159        let patcher_w = load_t(&content, &mut reader, "enc.patcher.weight", &device)?;
160        let patcher_b = load_t(&content, &mut reader, "enc.patcher.bias", &device)?;
161
162        let n_levels = config.adaptive_patching_levels;
163        let n_enc_layers = config.num_layers;
164        let mut enc_levels = Vec::with_capacity(n_levels);
165        for l in 0..n_levels {
166            // mixers[0] was created with adapt_patch_level = n_levels-1, factor = 2^(n_levels-1)
167            let factor = 1usize << (n_levels - 1 - l);
168            let mut layers = Vec::with_capacity(n_enc_layers);
169            for n in 0..n_enc_layers {
170                let prefix = format!("enc.blk.{l}.layer.{n}");
171                layers.push(load_mixer_layer(&content, &mut reader, &prefix, &device)?);
172            }
173            enc_levels.push(AdaptiveLevel { factor, layers });
174        }
175
176        let dec_adapter_w = load_t(&content, &mut reader, "dec.adapter.weight", &device)?;
177        let dec_adapter_b = load_t(&content, &mut reader, "dec.adapter.bias", &device)?;
178
179        let n_dec_layers = config.decoder_num_layers;
180        let mut dec_layers = Vec::with_capacity(n_dec_layers);
181        for n in 0..n_dec_layers {
182            let prefix = format!("dec.blk.{n}");
183            dec_layers.push(load_mixer_layer(&content, &mut reader, &prefix, &device)?);
184        }
185
186        let head_w = load_t(&content, &mut reader, "head.weight", &device)?;
187        let head_b = load_t(&content, &mut reader, "head.bias", &device)?;
188
189        Ok(Self {
190            device,
191            config,
192            patcher_w,
193            patcher_b,
194            enc_levels,
195            dec_adapter_w,
196            dec_adapter_b,
197            dec_layers,
198            head_w,
199            head_b,
200        })
201    }
202
203    // -----------------------------------------------------------------------
204    // Inference
205    // -----------------------------------------------------------------------
206
207    /// Forecast univariate time series.
208    ///
209    /// `context` must be at least `config.patch_length` long.
210    /// Returns `config.prediction_length` forecast values.
211    pub fn forecast(&self, context: &[f32]) -> Result<Vec<f32>> {
212        let cfg = &self.config;
213
214        // 1. StdScaler
215        let (scaled, mean, std) = std_scale(context);
216
217        // 2. Patchify → [num_patches, patch_length]
218        let patches = patchify(&scaled, cfg.patch_length, cfg.patch_stride, cfg.num_patches);
219        let flat: Vec<f32> = patches.into_iter().flatten().collect();
220        let mut h = Tensor::from_vec(flat, (cfg.num_patches, cfg.patch_length), &self.device)?;
221
222        // 3. Patcher Linear(patch_length → d_model) → [num_patches, d_model]
223        h = linear(&h, &self.patcher_w, Some(&self.patcher_b))?;
224
225        // 4. Encoder adaptive patching (3 levels)
226        for level in &self.enc_levels {
227            h = self.forward_adaptive_level(&h, level)?;
228        }
229
230        // 5. Decoder adapter Linear(d_model → decoder_d_model) → [num_patches, decoder_d_model]
231        h = linear(&h, &self.dec_adapter_w, Some(&self.dec_adapter_b))?;
232
233        // 6. Decoder block (regular mixer layers)
234        for layer in &self.dec_layers {
235            h = forward_mixer_layer(&h, layer, cfg.norm_eps)?;
236        }
237
238        // 7. Flatten → [num_patches * decoder_d_model]
239        let flat_dim = cfg.num_patches * cfg.decoder_d_model;
240        h = h.reshape((flat_dim,))?;
241
242        // 8. Head Linear(flat_dim → prediction_length) → [prediction_length]
243        h = h.unsqueeze(0)?; // [1, flat_dim]
244        h = linear(&h, &self.head_w, Some(&self.head_b))?; // [1, prediction_length]
245        h = h.squeeze(0)?; // [prediction_length]
246
247        // 9. Inverse scale
248        let forecast: Vec<f32> = h.to_vec1()?;
249        Ok(forecast.iter().map(|&v| v * std + mean).collect())
250    }
251
252    fn forward_adaptive_level(&self, hidden: &Tensor, level: &AdaptiveLevel) -> Result<Tensor> {
253        let factor = level.factor;
254        let (p, f) = (hidden.dim(0)?, hidden.dim(1)?);
255
256        // Reshape: [P, F] → [P*factor, F/factor]
257        let mut h = if factor > 1 {
258            hidden.reshape((p * factor, f / factor))?
259        } else {
260            hidden.clone()
261        };
262
263        for layer in &level.layers {
264            h = forward_mixer_layer(&h, layer, self.config.norm_eps)?;
265        }
266
267        // Reshape back: [P*factor, F/factor] → [P, F]
268        if factor > 1 { Ok(h.reshape((p, f))?) } else { Ok(h) }
269    }
270}
271
272// ---------------------------------------------------------------------------
273// MLP-Mixer forward
274// ---------------------------------------------------------------------------
275
276fn forward_mixer_layer(hidden: &Tensor, layer: &MixerLayer, norm_eps: f64) -> Result<Tensor> {
277    let h = forward_patch_mixer(hidden, layer, norm_eps)?;
278    forward_feat_mixer(&h, layer, norm_eps)
279}
280
281/// PatchMixerBlock: normalize features, transpose, MLP on patch dim, gate, transpose back.
282fn forward_patch_mixer(hidden: &Tensor, w: &MixerLayer, eps: f64) -> Result<Tensor> {
283    // LayerNorm on feature dim (last dim)
284    let h = layer_norm(hidden, &w.patch_norm_w, &w.patch_norm_b, eps)?;
285    // Transpose [P, F] → [F, P]
286    let h = h.t()?.contiguous()?;
287    // MLP on last dim (P)
288    let h = linear(&h, &w.patch_fc1_w, Some(&w.patch_fc1_b))?.gelu_erf()?;
289    let h = linear(&h, &w.patch_fc2_w, Some(&w.patch_fc2_b))?;
290    // Gated attention: softmax(linear(h)) * h
291    let gate = candle_nn::ops::softmax_last_dim(
292        &linear(&h, &w.patch_gate_w, Some(&w.patch_gate_b))?
293    )?;
294    let h = (h * gate)?;
295    // Transpose back [F, P] → [P, F]
296    let h = h.t()?.contiguous()?;
297    // Residual
298    Ok((h + hidden)?)
299}
300
301/// FeatureMixerBlock: normalize features, MLP on feature dim, gate, add residual.
302fn forward_feat_mixer(hidden: &Tensor, w: &MixerLayer, eps: f64) -> Result<Tensor> {
303    let h = layer_norm(hidden, &w.feat_norm_w, &w.feat_norm_b, eps)?;
304    let h = linear(&h, &w.feat_fc1_w, Some(&w.feat_fc1_b))?.gelu_erf()?;
305    let h = linear(&h, &w.feat_fc2_w, Some(&w.feat_fc2_b))?;
306    let gate = candle_nn::ops::softmax_last_dim(
307        &linear(&h, &w.feat_gate_w, Some(&w.feat_gate_b))?
308    )?;
309    let h = (h * gate)?;
310    Ok((h + hidden)?)
311}
312
313// ---------------------------------------------------------------------------
314// Primitive ops
315// ---------------------------------------------------------------------------
316
317// ---------------------------------------------------------------------------
318// Preprocessing
319// ---------------------------------------------------------------------------
320
321/// Compute mean and std over context, return (normalized, mean, std).
322fn std_scale(x: &[f32]) -> (Vec<f32>, f32, f32) {
323    let n = x.len() as f64;
324    let mean = x.iter().map(|&v| v as f64).sum::<f64>() / n;
325    let var = x.iter().map(|&v| (v as f64 - mean).powi(2)).sum::<f64>() / n;
326    // minimum_scale = 1e-5 (matches Python TinyTimeMixerStdScaler)
327    let std = (var + 1e-5).sqrt() as f32;
328    let scaled: Vec<f32> = x.iter().map(|&v| (v as f32 - mean as f32) / std).collect();
329    (scaled, mean as f32, std)
330}
331
332/// Extract `num_patches` patches of length `patch_length` with stride `patch_stride`.
333/// Uses the last `patch_length + patch_stride * (num_patches - 1)` timesteps.
334/// Left-pads with the first value when the context is shorter than needed.
335fn patchify(x: &[f32], patch_length: usize, patch_stride: usize, num_patches: usize) -> Vec<Vec<f32>> {
336    let new_seq_len = patch_length + patch_stride * (num_patches - 1);
337    let padded: Vec<f32> = if x.len() < new_seq_len {
338        let pad_val = x.first().copied().unwrap_or(0.0);
339        let mut v = vec![pad_val; new_seq_len - x.len()];
340        v.extend_from_slice(x);
341        v
342    } else {
343        x[x.len() - new_seq_len..].to_vec()
344    };
345    (0..num_patches)
346        .map(|i| padded[i * patch_stride..i * patch_stride + patch_length].to_vec())
347        .collect()
348}
349
350// ---------------------------------------------------------------------------
351// zsfm-core::Forecaster
352// ---------------------------------------------------------------------------
353
354impl zsfm_core::Forecaster for TtmModel {
355    type Config = TtmConfig;
356
357    fn load(gguf_path: &Path, config: TtmConfig) -> Result<Self> {
358        TtmModel::load(gguf_path, config)
359    }
360
361    /// TTM is univariate-only and point-forecast-only; `mask` is unused.
362    fn forecast(
363        &self,
364        context: &[Vec<f32>],
365        _mask: &[Vec<bool>],
366        horizon: usize,
367    ) -> Result<zsfm_core::QuantileMatrix> {
368        anyhow::ensure!(context.len() == 1, "TtmModel only supports univariate forecasting (1 variate)");
369        let raw = TtmModel::forecast(self, &context[0])?;
370        let point: Vec<f32> = raw.into_iter().take(horizon).collect();
371        Ok(vec![vec![point]])
372    }
373}