1use 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
19struct 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
60pub 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
104fn 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 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 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 pub fn forecast(&self, context: &[f32]) -> Result<Vec<f32>> {
212 let cfg = &self.config;
213
214 let (scaled, mean, std) = std_scale(context);
216
217 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 h = linear(&h, &self.patcher_w, Some(&self.patcher_b))?;
224
225 for level in &self.enc_levels {
227 h = self.forward_adaptive_level(&h, level)?;
228 }
229
230 h = linear(&h, &self.dec_adapter_w, Some(&self.dec_adapter_b))?;
232
233 for layer in &self.dec_layers {
235 h = forward_mixer_layer(&h, layer, cfg.norm_eps)?;
236 }
237
238 let flat_dim = cfg.num_patches * cfg.decoder_d_model;
240 h = h.reshape((flat_dim,))?;
241
242 h = h.unsqueeze(0)?; h = linear(&h, &self.head_w, Some(&self.head_b))?; h = h.squeeze(0)?; 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 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 if factor > 1 { Ok(h.reshape((p, f))?) } else { Ok(h) }
269 }
270}
271
272fn 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
281fn forward_patch_mixer(hidden: &Tensor, w: &MixerLayer, eps: f64) -> Result<Tensor> {
283 let h = layer_norm(hidden, &w.patch_norm_w, &w.patch_norm_b, eps)?;
285 let h = h.t()?.contiguous()?;
287 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 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 let h = h.t()?.contiguous()?;
297 Ok((h + hidden)?)
299}
300
301fn 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
313fn 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 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
332fn 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
350impl 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 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}