1use std::collections::HashMap;
8use std::sync::Mutex;
9use std::io::{BufReader, Read, Seek};
10use std::path::Path;
11
12use anyhow::{Context, Result};
13use candle_core::quantized::gguf_file;
14use candle_core::{DType, Device, Tensor, D};
15use simdeez::prelude::*;
16
17use crate::config::LagLlamaConfig;
18
19struct TransformerBlock {
24 rms1_w: Tensor,
25 rms2_w: Tensor,
26 qkv_w: Tensor, c_w: Tensor,
28 fc1_w: Tensor,
29 fc2_w: Tensor,
30 proj_w: Tensor,
31}
32
33struct RawBlockWeights {
34 rms1: Vec<f32>, rms2: Vec<f32>, qkv: Vec<f32>, c: Vec<f32>, fc1: Vec<f32>, fc2: Vec<f32>, proj: Vec<f32>, }
42
43pub struct LagLlamaModel {
44 device: Device,
45 config: LagLlamaConfig,
46 rope_cos: Tensor,
47 rope_sin: Tensor,
48 causal_mask_cache: Mutex<HashMap<usize, Tensor>>,
49 wte_w: Tensor,
50 wte_b: Tensor,
51 blocks: Vec<TransformerBlock>,
52 raw_blocks: Vec<RawBlockWeights>,
54 rope_cos_raw: Vec<f32>, rope_sin_raw: Vec<f32>,
56 norm_f_raw: Vec<f32>, wte_w_raw: Vec<f32>, wte_b_raw: Vec<f32>, mu_w_raw: Vec<f32>, mu_b_raw: Vec<f32>, }
62
63fn load_t(
68 content: &gguf_file::Content,
69 reader: &mut (impl Read + Seek),
70 name: &str,
71 device: &Device,
72) -> Result<Tensor> {
73 zsfm_nn::load_tensor(content, reader, name, device, DType::F32)
74}
75
76impl LagLlamaModel {
77 pub fn load(gguf_path: &Path, config: LagLlamaConfig) -> Result<Self> {
78 let device = Device::Cpu;
79 let file = std::fs::File::open(gguf_path)
80 .with_context(|| format!("open {}", gguf_path.display()))?;
81 let mut reader = BufReader::with_capacity(zsfm_gguf::READ_BUF_CAPACITY, file);
82 let content = gguf_file::Content::read(&mut reader).context("parse GGUF header")?;
83
84 let wte_w = load_t(&content, &mut reader, "enc.wte.weight", &device)?;
85 let wte_b = load_t(&content, &mut reader, "enc.wte.bias", &device)?;
86
87 let mut blocks = Vec::with_capacity(config.n_layer);
88 for n in 0..config.n_layer {
89 let p = |s: &str| format!("blk.{n}.{s}");
90 let q_w = load_t(&content, &mut reader, &p("attn_q.weight"), &device)?;
91 let kv_w = load_t(&content, &mut reader, &p("attn_kv.weight"), &device)?;
92 blocks.push(TransformerBlock {
93 rms1_w: load_t(&content, &mut reader, &p("rms1.weight"), &device)?,
94 rms2_w: load_t(&content, &mut reader, &p("rms2.weight"), &device)?,
95 qkv_w: Tensor::cat(&[&q_w, &kv_w], 0)?,
96 c_w: load_t(&content, &mut reader, &p("attn_c.weight"), &device)?,
97 fc1_w: load_t(&content, &mut reader, &p("mlp_fc1.weight"), &device)?,
98 fc2_w: load_t(&content, &mut reader, &p("mlp_fc2.weight"), &device)?,
99 proj_w: load_t(&content, &mut reader, &p("mlp_proj.weight"), &device)?,
100 });
101 }
102
103 let norm_f_w = load_t(&content, &mut reader, "norm_f.weight", &device)?;
104 let mu_w_t = load_t(&content, &mut reader, "head.mu.weight", &device)?;
105 let mu_b_t = load_t(&content, &mut reader, "head.mu.bias", &device)?;
106
107 let head_dim = config.n_embd_per_head;
108 let half = head_dim / 2;
109 let inv_freq: Vec<f32> = (0..half)
110 .map(|i| 1.0_f32 / 10000_f32.powf(2.0 * i as f32 / head_dim as f32))
111 .collect();
112
113 let max_pos = config.max_context_length + 4096;
114 let mut cos_vals = vec![0.0f32; max_pos * half];
115 let mut sin_vals = vec![0.0f32; max_pos * half];
116 for p in 0..max_pos {
117 let pos = p as f32;
118 for i in 0..half {
119 let theta = pos * inv_freq[i];
120 cos_vals[p * half + i] = theta.cos();
121 sin_vals[p * half + i] = theta.sin();
122 }
123 }
124
125 let rope_cos_raw = cos_vals.clone();
127 let rope_sin_raw = sin_vals.clone();
128 let rope_cos = Tensor::from_vec(cos_vals, (max_pos, half), &device)?;
129 let rope_sin = Tensor::from_vec(sin_vals, (max_pos, half), &device)?;
130
131 let norm_f_raw = norm_f_w.flatten_all()?.to_vec1::<f32>()?;
133 let wte_w_raw = wte_w.flatten_all()?.to_vec1::<f32>()?;
134 let wte_b_raw = wte_b.flatten_all()?.to_vec1::<f32>()?;
135 let mu_w_raw = mu_w_t.flatten_all()?.to_vec1::<f32>()?;
136 let mu_b_raw = mu_b_t.flatten_all()?.to_vec1::<f32>()?;
137
138 let mut raw_blocks = Vec::with_capacity(config.n_layer);
140 for blk in &blocks {
141 raw_blocks.push(RawBlockWeights {
142 rms1: blk.rms1_w.flatten_all()?.to_vec1::<f32>()?,
143 rms2: blk.rms2_w.flatten_all()?.to_vec1::<f32>()?,
144 qkv: blk.qkv_w.flatten_all()?.to_vec1::<f32>()?,
145 c: blk.c_w.flatten_all()?.to_vec1::<f32>()?,
146 fc1: blk.fc1_w.flatten_all()?.to_vec1::<f32>()?,
147 fc2: blk.fc2_w.flatten_all()?.to_vec1::<f32>()?,
148 proj: blk.proj_w.flatten_all()?.to_vec1::<f32>()?,
149 });
150 }
151
152 Ok(Self {
153 device,
154 config,
155 rope_cos,
156 rope_sin,
157 causal_mask_cache: Mutex::new(HashMap::new()),
158 wte_w,
159 wte_b,
160 blocks,
161 raw_blocks,
162 rope_cos_raw,
163 rope_sin_raw,
164 norm_f_raw,
165 wte_w_raw,
166 wte_b_raw,
167 mu_w_raw,
168 mu_b_raw,
169 })
170 }
171
172 pub fn forecast(&self, context: &[f32], horizon: usize) -> Result<Vec<f32>> {
177 let cfg = &self.config;
178 let max_lag = *cfg.lags_seq.iter().max().unwrap_or(&0);
179
180 let (loc, scale) = robust_stats(context);
181 let scale = scale.max(1e-8);
182
183 let mut hist: Vec<f32> = vec![0.0; max_lag + 1];
184 for &v in context {
185 hist.push((v - loc) / scale);
186 }
187
188 let ctx_buf_end = hist.len();
189 let ctx_buf_start = ctx_buf_end.saturating_sub(cfg.max_context_length);
190 let seq_len = ctx_buf_end - ctx_buf_start;
191
192 let feat_ctx = build_feature_matrix(&hist, ctx_buf_start, seq_len, cfg);
194 let x = Tensor::from_vec(feat_ctx, (seq_len, cfg.feature_size), &self.device)?;
195 let mut h = zsfm_nn::linear_bias(&x, &self.wte_w, &self.wte_b)?;
196
197 let mut kv_caches: Vec<(Tensor, Tensor)> = Vec::with_capacity(cfg.n_layer);
198 for blk in &self.blocks {
199 let (h_out, k, v) = self.prefill_block(&h, blk, seq_len, ctx_buf_start)?;
200 h = h_out;
201 kv_caches.push((k, v));
202 }
203
204 let mut h_raw: Vec<f32> = h.get(seq_len - 1)?.to_vec1()?;
206
207 let n_head = cfg.n_head;
208 let head_dim = cfg.n_embd_per_head;
209 let n_embd = cfg.n_embd;
210 let mlp_hidden = cfg.mlp_hidden;
211 let feat_size = cfg.feature_size;
212
213 let mut kv_raw = extract_kv_caches_raw(&kv_caches, n_head, head_dim, horizon)?;
215 let mut kv_len = seq_len;
216
217 let max_kv_len = seq_len + horizon;
219 let mut h_tmp = vec![0.0f32; n_embd];
220 let mut qkv_buf = vec![0.0f32; 3 * n_embd];
221 let mut proj_buf = vec![0.0f32; n_embd];
222 let mut mlp_gate = vec![0.0f32; mlp_hidden];
223 let mut mlp_up = vec![0.0f32; mlp_hidden];
224 let mut mlp_out = vec![0.0f32; n_embd];
225 let mut attn_out = vec![0.0f32; n_embd];
226 let mut scores_scratch = vec![0.0f32; n_head * max_kv_len];
227
228 let mut preds = Vec::with_capacity(horizon);
229 let mut rope_offset = ctx_buf_start + seq_len;
230
231 for step in 0..horizon {
232 h_tmp.copy_from_slice(&h_raw);
234 rms_norm_raw(&mut h_tmp, &self.norm_f_raw, 1e-5);
235 let pred_scaled = raw_dot(&h_tmp, &self.mu_w_raw) + self.mu_b_raw[0];
236 preds.push(pred_scaled * scale + loc);
237
238 if step == horizon - 1 { break; }
239
240 hist.push(pred_scaled);
242
243 let abs_t = hist.len() - 1;
245 let feat_one = build_one_feature(&hist, abs_t, cfg);
246 raw_gemv_bias(&feat_one, &self.wte_w_raw, &self.wte_b_raw,
247 n_embd, feat_size, &mut h_raw);
248
249 let new_kv_len = kv_len + 1;
251 for li in 0..cfg.n_layer {
252 let blk = &self.raw_blocks[li];
253 let (ref mut k_heads, ref mut v_heads) = kv_raw[li];
254
255 h_tmp.copy_from_slice(&h_raw);
257 rms_norm_raw(&mut h_tmp, &blk.rms1, 1e-5);
258 raw_gemv(&h_tmp, &blk.qkv, 3 * n_embd, n_embd, &mut qkv_buf);
259
260 rope_single_inplace(&mut qkv_buf[..n_embd],
262 rope_offset, &self.rope_cos_raw, &self.rope_sin_raw,
263 n_head, head_dim);
264 rope_single_inplace(&mut qkv_buf[n_embd..2 * n_embd],
265 rope_offset, &self.rope_cos_raw, &self.rope_sin_raw,
266 n_head, head_dim);
267
268 for hi in 0..n_head {
270 k_heads[hi].extend_from_slice(
271 &qkv_buf[n_embd + hi * head_dim..n_embd + (hi + 1) * head_dim]);
272 v_heads[hi].extend_from_slice(
273 &qkv_buf[2 * n_embd + hi * head_dim..2 * n_embd + (hi + 1) * head_dim]);
274 }
275
276 mha_decode_raw(&qkv_buf[..n_embd], k_heads, v_heads,
277 n_head, head_dim, new_kv_len,
278 &mut scores_scratch, &mut attn_out);
279
280 raw_gemv(&attn_out, &blk.c, n_embd, n_embd, &mut proj_buf);
281 for i in 0..n_embd { h_raw[i] += proj_buf[i]; }
282
283 h_tmp.copy_from_slice(&h_raw);
285 rms_norm_raw(&mut h_tmp, &blk.rms2, 1e-5);
286 silu_mlp_raw(&h_tmp, &blk.fc1, &blk.fc2, &blk.proj,
287 n_embd, mlp_hidden,
288 &mut mlp_gate, &mut mlp_up, &mut mlp_out);
289 for i in 0..n_embd { h_raw[i] += mlp_out[i]; }
290 }
291
292 kv_len = new_kv_len;
293 rope_offset += 1;
294 }
295
296 Ok(preds)
297 }
298
299 fn prefill_block(
304 &self,
305 hidden: &Tensor,
306 blk: &TransformerBlock,
307 seq_len: usize,
308 rope_start: usize,
309 ) -> Result<(Tensor, Tensor, Tensor)> {
310 let res = hidden;
311 let h = zsfm_nn::rms_norm(hidden, Some(&blk.rms1_w), 1e-5)?;
312 let (attn_out, k, v) = self.prefill_attn(&h, blk, seq_len, rope_start)?;
313 let h = (attn_out + res)?;
314
315 let res2 = h.clone();
316 let h2 = zsfm_nn::rms_norm(&h, Some(&blk.rms2_w), 1e-5)?;
317 let h2 = silu_mlp(&h2, &blk.fc1_w, &blk.fc2_w, &blk.proj_w)?;
318 Ok(((h2 + res2)?, k, v))
319 }
320
321 fn prefill_attn(
322 &self,
323 hidden: &Tensor,
324 blk: &TransformerBlock,
325 seq_len: usize,
326 rope_start: usize,
327 ) -> Result<(Tensor, Tensor, Tensor)> {
328 let cfg = &self.config;
329 let n_head = cfg.n_head;
330 let head_dim = cfg.n_embd_per_head;
331 let n_embd = cfg.n_embd;
332
333 let qkv = zsfm_nn::linear_nobias(hidden, &blk.qkv_w)?;
334 let q = qkv.narrow(1, 0, n_embd)?;
335 let k = qkv.narrow(1, n_embd, n_embd)?;
336 let v = qkv.narrow(1, 2 * n_embd, n_embd)?;
337
338 let q = q.reshape((seq_len, n_head, head_dim))?.permute((1, 0, 2))?.contiguous()?;
339 let k = k.reshape((seq_len, n_head, head_dim))?.permute((1, 0, 2))?.contiguous()?;
340 let v = v.reshape((seq_len, n_head, head_dim))?.permute((1, 0, 2))?.contiguous()?;
341
342 let q = apply_rope(&q, rope_start, seq_len, &self.rope_cos, &self.rope_sin)?;
343 let k = apply_rope(&k, rope_start, seq_len, &self.rope_cos, &self.rope_sin)?;
344
345 let scale = (head_dim as f64).sqrt();
346 let attn_weights = (q.matmul(&k.permute((0, 2, 1))?)? / scale)?;
347 let attn_weights = self.apply_causal_mask_ll(attn_weights, seq_len)?;
348 let attn = candle_nn::ops::softmax_last_dim(&attn_weights)?;
349
350 let out = attn.matmul(&v)?;
351 let out = out.permute((1, 0, 2))?.contiguous()?.reshape((seq_len, n_embd))?;
352 let out = zsfm_nn::linear_nobias(&out, &blk.c_w)?;
353
354 Ok((out, k, v))
355 }
356
357 fn apply_causal_mask_ll(&self, attn: Tensor, seq_len: usize) -> Result<Tensor> {
358 let mut cache = self.causal_mask_cache.lock().unwrap();
359 if !cache.contains_key(&seq_len) {
360 let mut mask_data = vec![0.0f32; seq_len * seq_len];
361 for i in 0..seq_len {
362 for j in (i + 1)..seq_len {
363 mask_data[i * seq_len + j] = f32::NEG_INFINITY;
364 }
365 }
366 let mask = Tensor::from_vec(mask_data, (1usize, seq_len, seq_len), &self.device)?;
367 cache.insert(seq_len, mask);
368 }
369 Ok(attn.broadcast_add(&cache[&seq_len])?)
370 }
371}
372
373fn build_feature_matrix(
378 hist: &[f32],
379 start: usize,
380 seq_len: usize,
381 cfg: &LagLlamaConfig,
382) -> Vec<f32> {
383 let mut feat = vec![0.0f32; seq_len * cfg.feature_size];
384 for t in 0..seq_len {
385 let abs_t = start + t;
386 for (li, &lag) in cfg.lags_seq.iter().enumerate() {
387 let src = abs_t as isize - lag as isize;
388 if src >= 0 && (src as usize) < hist.len() {
389 feat[t * cfg.feature_size + li] = hist[src as usize];
390 }
391 }
392 }
393 feat
394}
395
396fn build_one_feature(hist: &[f32], abs_t: usize, cfg: &LagLlamaConfig) -> Vec<f32> {
397 let mut feat = vec![0.0f32; cfg.feature_size];
398 for (li, &lag) in cfg.lags_seq.iter().enumerate() {
399 let src = abs_t as isize - lag as isize;
400 if src >= 0 && (src as usize) < hist.len() {
401 feat[li] = hist[src as usize];
402 }
403 }
404 feat
405}
406
407fn silu_mlp(x: &Tensor, fc1_w: &Tensor, fc2_w: &Tensor, proj_w: &Tensor) -> Result<Tensor> {
415 zsfm_nn::swiglu_ffn(x, fc1_w, proj_w, fc2_w)
416}
417
418fn apply_rope(
419 x: &Tensor,
420 offset: usize,
421 seq_len: usize,
422 cos_table: &Tensor,
423 sin_table: &Tensor,
424) -> Result<Tensor> {
425 let half = cos_table.dim(1)?;
426 let cos_t = cos_table.narrow(0, offset, seq_len)?.unsqueeze(0)?;
427 let sin_t = sin_table.narrow(0, offset, seq_len)?.unsqueeze(0)?;
428
429 let x1 = x.narrow(D::Minus1, 0, half)?.contiguous()?;
430 let x2 = x.narrow(D::Minus1, half, half)?.contiguous()?;
431 let rot1 = (x1.broadcast_mul(&cos_t)? - x2.broadcast_mul(&sin_t)?)?;
432 let rot2 = (x1.broadcast_mul(&sin_t)? + x2.broadcast_mul(&cos_t)?)?;
433 Ok(Tensor::cat(&[&rot1, &rot2], D::Minus1)?.contiguous()?)
434}
435
436simd_runtime_generate!(
441 fn simd_sq_sum(row: &[f32]) -> f32 {
442 let mut r = &row[..];
443 let mut acc = S::Vf32::zeroes();
444 while r.len() >= S::Vf32::WIDTH {
445 let v = S::Vf32::load_from_slice(r);
446 acc = v.mul_add(v, acc);
447 r = &r[S::Vf32::WIDTH..];
448 }
449 let mut sum = acc.horizontal_add();
450 for &x in r { sum += x * x; }
451 sum
452 }
453);
454
455simd_runtime_generate!(
456 fn simd_dot(a: &[f32], b: &[f32]) -> f32 {
457 let mut aa = &a[..];
458 let mut bb = &b[..];
459 let mut acc = S::Vf32::zeroes();
460 while aa.len() >= S::Vf32::WIDTH {
461 let va = S::Vf32::load_from_slice(aa);
462 let vb = S::Vf32::load_from_slice(bb);
463 acc = va.mul_add(vb, acc);
464 aa = &aa[S::Vf32::WIDTH..];
465 bb = &bb[S::Vf32::WIDTH..];
466 }
467 let mut sum = acc.horizontal_add();
468 for (&x, &y) in aa.iter().zip(bb.iter()) { sum += x * y; }
469 sum
470 }
471);
472
473#[inline(always)]
476fn fast_exp_f32(x: f32) -> f32 {
477 let x = x.max(-87.3365_f32); f32::from_bits(((x * 12102203.0_f32) as i32 + 1064866805_i32) as u32)
479}
480
481#[inline(always)]
482fn raw_dot(a: &[f32], b: &[f32]) -> f32 {
483 simd_dot(a, b)
484}
485
486fn rms_norm_raw(x: &mut [f32], w: &[f32], eps: f32) {
487 let n = x.len() as f32;
488 let rms = (simd_sq_sum(x) / n + eps).sqrt();
489 for i in 0..x.len() {
490 x[i] = x[i] / rms * w[i];
491 }
492}
493
494fn raw_gemv(x: &[f32], w: &[f32], n_out: usize, n_in: usize, out: &mut [f32]) {
495 for i in 0..n_out {
496 out[i] = raw_dot(x, &w[i * n_in..(i + 1) * n_in]);
497 }
498}
499
500fn raw_gemv_bias(x: &[f32], w: &[f32], b: &[f32], n_out: usize, n_in: usize, out: &mut [f32]) {
501 for i in 0..n_out {
502 out[i] = raw_dot(x, &w[i * n_in..(i + 1) * n_in]) + b[i];
503 }
504}
505
506fn rope_single_inplace(
508 qk: &mut [f32],
509 pos: usize,
510 cos: &[f32],
511 sin: &[f32],
512 n_head: usize,
513 head_dim: usize,
514) {
515 let half = head_dim / 2;
516 let cos_row = &cos[pos * half..(pos + 1) * half];
517 let sin_row = &sin[pos * half..(pos + 1) * half];
518 for h in 0..n_head {
519 let base = h * head_dim;
520 for i in 0..half {
521 let x1 = qk[base + i];
522 let x2 = qk[base + half + i];
523 qk[base + i] = x1 * cos_row[i] - x2 * sin_row[i];
524 qk[base + half + i] = x1 * sin_row[i] + x2 * cos_row[i];
525 }
526 }
527}
528
529fn mha_decode_raw(
535 q: &[f32],
536 k_heads: &[Vec<f32>],
537 v_heads: &[Vec<f32>],
538 n_head: usize,
539 head_dim: usize,
540 kv_len: usize,
541 scores_scratch: &mut [f32],
542 out: &mut [f32],
543) {
544 let scale_inv = 1.0 / (head_dim as f32).sqrt();
545 for h in 0..n_head {
546 let q_h = &q[h * head_dim..(h + 1) * head_dim];
547 let k_h = &k_heads[h];
548 let v_h = &v_heads[h];
549 let sc = &mut scores_scratch[h * kv_len..(h + 1) * kv_len];
550
551 let k_ptr = k_h.as_ptr();
553 for j in 0..kv_len {
554 let kj = unsafe { std::slice::from_raw_parts(k_ptr.add(j * head_dim), head_dim) };
555 let mut dot = 0.0f32;
556 for d in 0..head_dim { dot += q_h[d] * kj[d]; }
557 sc[j] = dot * scale_inv;
558 }
559
560 let max_s = sc.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
562 let mut sum = 0.0f32;
563 for s in sc.iter_mut() { *s = fast_exp_f32(*s - max_s); sum += *s; }
564 let inv_sum = 1.0 / sum;
565 for s in sc.iter_mut() { *s *= inv_sum; }
566
567 let out_h = &mut out[h * head_dim..(h + 1) * head_dim];
568 out_h.iter_mut().for_each(|v| *v = 0.0);
569 let v_ptr = v_h.as_ptr();
572 for j in 0..kv_len {
573 let sc_j = sc[j];
574 let v_j = unsafe { std::slice::from_raw_parts(v_ptr.add(j * head_dim), head_dim) };
575 for d in 0..head_dim {
576 out_h[d] += sc_j * v_j[d];
577 }
578 }
579 }
580}
581
582fn silu_mlp_raw(
585 x: &[f32],
586 fc1: &[f32],
587 fc2: &[f32],
588 proj: &[f32],
589 n_embd: usize,
590 mlp_hidden: usize,
591 gate: &mut [f32],
592 up: &mut [f32],
593 out: &mut [f32],
594) {
595 for i in 0..mlp_hidden {
596 let v = raw_dot(x, &fc1[i * n_embd..(i + 1) * n_embd]);
597 gate[i] = v / (1.0 + fast_exp_f32(-v)); }
599 for i in 0..mlp_hidden {
600 up[i] = raw_dot(x, &fc2[i * n_embd..(i + 1) * n_embd]);
601 }
602 for j in 0..mlp_hidden { gate[j] *= up[j]; } raw_gemv(gate, proj, n_embd, mlp_hidden, out);
604}
605
606fn extract_kv_caches_raw(
609 kv_caches: &[(Tensor, Tensor)],
610 n_head: usize,
611 head_dim: usize,
612 horizon: usize,
613) -> Result<Vec<(Vec<Vec<f32>>, Vec<Vec<f32>>)>> {
614 let mut result = Vec::with_capacity(kv_caches.len());
615 for (k_t, v_t) in kv_caches {
616 let k_flat = k_t.flatten_all()?.to_vec1::<f32>()?;
617 let v_flat = v_t.flatten_all()?.to_vec1::<f32>()?;
618 let tokens_per_head = k_flat.len() / n_head; let cap = tokens_per_head + horizon * head_dim;
620 let mut k_heads: Vec<Vec<f32>> = Vec::with_capacity(n_head);
621 let mut v_heads: Vec<Vec<f32>> = Vec::with_capacity(n_head);
622 for h in 0..n_head {
623 let start = h * tokens_per_head;
624 let end = start + tokens_per_head;
625 let mut kh = Vec::with_capacity(cap);
626 kh.extend_from_slice(&k_flat[start..end]);
627 let mut vh = Vec::with_capacity(cap);
628 vh.extend_from_slice(&v_flat[start..end]);
629 k_heads.push(kh);
630 v_heads.push(vh);
631 }
632 result.push((k_heads, v_heads));
633 }
634 Ok(result)
635}
636
637fn robust_stats(x: &[f32]) -> (f32, f32) {
642 let n = x.len();
643 if n == 0 { return (0.0, 1.0); }
644 let mut sorted = x.to_vec();
645 sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
646 let median = if n % 2 == 0 {
647 (sorted[n / 2 - 1] + sorted[n / 2]) / 2.0
648 } else {
649 sorted[n / 2]
650 };
651 let mut devs: Vec<f32> = sorted.iter().map(|&v| (v - median).abs()).collect();
652 devs.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
653 let mad = if n % 2 == 0 {
654 (devs[n / 2 - 1] + devs[n / 2]) / 2.0
655 } else {
656 devs[n / 2]
657 };
658 (median, mad)
659}
660
661impl zsfm_core::Forecaster for LagLlamaModel {
666 type Config = LagLlamaConfig;
667
668 fn load(gguf_path: &Path, config: LagLlamaConfig) -> Result<Self> {
669 LagLlamaModel::load(gguf_path, config)
670 }
671
672 fn forecast(
674 &self,
675 context: &[Vec<f32>],
676 _mask: &[Vec<bool>],
677 horizon: usize,
678 ) -> Result<zsfm_core::QuantileMatrix> {
679 anyhow::ensure!(context.len() == 1, "LagLlamaModel only supports univariate forecasting (1 variate)");
680 let point = LagLlamaModel::forecast(self, &context[0], horizon)?;
681 Ok(vec![vec![point]])
682 }
683}