Skip to main content

zsfm_lag_llama/infer/
mod.rs

1//! Lag-Llama inference engine with KV caching.
2//!
3//! Two-phase inference:
4//! 1. Prefill: full forward pass over context tokens, collects K/V cache per layer.
5//! 2. Decode:  per-step single-token forward pass in pure Rust (zero Candle overhead).
6
7use 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
19// ---------------------------------------------------------------------------
20// Weight structs
21// ---------------------------------------------------------------------------
22
23struct TransformerBlock {
24    rms1_w: Tensor,
25    rms2_w: Tensor,
26    qkv_w:  Tensor,  // fused [3*n_embd, n_embd]
27    c_w:    Tensor,
28    fc1_w:  Tensor,
29    fc2_w:  Tensor,
30    proj_w: Tensor,
31}
32
33struct RawBlockWeights {
34    rms1: Vec<f32>,  // [n_embd]
35    rms2: Vec<f32>,  // [n_embd]
36    qkv:  Vec<f32>,  // [3*n_embd, n_embd] row-major
37    c:    Vec<f32>,  // [n_embd, n_embd]
38    fc1:  Vec<f32>,  // [mlp_hidden, n_embd]
39    fc2:  Vec<f32>,  // [mlp_hidden, n_embd]
40    proj: Vec<f32>,  // [n_embd, mlp_hidden]
41}
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 arrays for zero-overhead decode loop
53    raw_blocks:   Vec<RawBlockWeights>,
54    rope_cos_raw: Vec<f32>,  // [max_pos * half_head_dim]
55    rope_sin_raw: Vec<f32>,
56    norm_f_raw:   Vec<f32>,  // [n_embd]
57    wte_w_raw:    Vec<f32>,  // [n_embd, feature_size]
58    wte_b_raw:    Vec<f32>,  // [n_embd]
59    mu_w_raw:     Vec<f32>,  // [n_embd] (flattened from head weight)
60    mu_b_raw:     Vec<f32>,  // [1]
61}
62
63// ---------------------------------------------------------------------------
64// GGUF loading
65// ---------------------------------------------------------------------------
66
67fn 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        // Keep raw copies before Tensor::from_vec consumes the Vecs
126        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        // Extract global raw weights
132        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        // Extract per-block raw weights
139        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    // -----------------------------------------------------------------------
173    // Forecasting: Candle prefill + raw-array decode loop
174    // -----------------------------------------------------------------------
175
176    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        // --- Phase 1: Candle prefill (full sequence, once) ---
193        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        // Extract last hidden state + KV caches into raw arrays (one-time cost)
205        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        // Per-layer, per-head KV buffers; pre-reserve full decode capacity
214        let mut kv_raw = extract_kv_caches_raw(&kv_caches, n_head, head_dim, horizon)?;
215        let mut kv_len = seq_len;
216
217        // --- Phase 2: pure-Rust decode loop (zero Candle ops per step) ---
218        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            // Predict from current hidden state
233            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            // Append normalized prediction to history
241            hist.push(pred_scaled);
242
243            // Build feature for next token and embed it
244            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            // Run 8 transformer layers in raw Rust
250            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                // Attention sublayer
256                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 on Q (qkv_buf[0..n_embd]) and K (qkv_buf[n_embd..2*n_embd])
261                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                // Append this token's K and V into per-head buffers
269                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                // FFN sublayer
284                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    // -----------------------------------------------------------------------
300    // Prefill (Candle path — runs once per window)
301    // -----------------------------------------------------------------------
302
303    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
373// ---------------------------------------------------------------------------
374// Feature construction
375// ---------------------------------------------------------------------------
376
377fn 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
407// ---------------------------------------------------------------------------
408// Candle ops (used by prefill path)
409// ---------------------------------------------------------------------------
410
411/// `silu(fc1(x)) * fc2(x)`, then `proj`. Same shape as [`zsfm_nn::swiglu_ffn`]'s
412/// `silu(A(x)) * C(x)` then `B(h)` — here `B` (final proj) is the 3rd param, not the
413/// 2nd, so the last two args are passed to it swapped.
414fn 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
436// ---------------------------------------------------------------------------
437// Raw ops (used by decode path — zero Candle overhead)
438// ---------------------------------------------------------------------------
439
440simd_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// Schraudolph fast exp: ~0.2% relative error — sufficient for softmax and SiLU.
474// ~3–5× faster than libm exp() by bypassing PLT dispatch and IEEE corner-case handling.
475#[inline(always)]
476fn fast_exp_f32(x: f32) -> f32 {
477    let x = x.max(-87.3365_f32); // clamp underflow to 0.0 output
478    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
506// Apply RoPE in-place to a flat [n_head, head_dim] buffer for a single token at `pos`.
507fn 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
529// Single-query multi-head attention against per-head KV buffers.
530// q:             [n_head * head_dim] — Q for the new token
531// k_heads[h]:    [kv_len * head_dim] — all K tokens for head h
532// scores_scratch: [n_head * kv_len_max] — temporary scores (stride = kv_len)
533// out:           [n_head * head_dim]
534fn 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        // Inline 16-element dot without function-pointer dispatch (LLVM auto-vecs to fmla.4s).
552        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        // Numerically stable softmax — one division, rest multiply
561        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        // Safety: v_h.len() == kv_len * head_dim (invariant: prefill extraction +
570        // extend_from_slice(head_dim) per step). Avoids bounds-checked slice per j.
571        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
582// SiLU-gated MLP: out = proj(silu(fc1(x)) * fc2(x)).
583// gate and up are scratch buffers of length mlp_hidden.
584fn 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)); // silu
598    }
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]; }  // gate now holds h = silu(fc1) * fc2
603    raw_gemv(gate, proj, n_embd, mlp_hidden, out);
604}
605
606// Extract Candle KV caches ([n_head, seq_len, head_dim]) into per-head Vecs.
607// Pre-allocates capacity for the full decode horizon to avoid realloc.
608fn 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; // seq_len * head_dim
619        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
637// ---------------------------------------------------------------------------
638// Robust scaler
639// ---------------------------------------------------------------------------
640
641fn 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
661// ---------------------------------------------------------------------------
662// zsfm-core::Forecaster
663// ---------------------------------------------------------------------------
664
665impl 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    /// Lag-Llama is univariate-only and point-forecast-only; `mask` is unused.
673    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}