Skip to main content

ferrum_kernels/backend/
cpu.rs

1//! CPU backend using Accelerate (macOS) / portable fallback (Linux).
2//! Context = () — all ops execute immediately, no batching needed.
3
4use half::f16;
5
6use super::{AttnConfig, Backend};
7use ferrum_types::{FerrumError, Result};
8
9// ── Q4_K_M block layout ────────────────────────────────────────────────
10//
11// Mirrors GGML / candle's `BlockQ4K`. Used by `load_q4_k` to dequant raw
12// GGUF block bytes to fp32 row-major weights on CPU.
13
14const Q4_K_QK: usize = 256;
15const Q4_K_SCALE_SIZE: usize = 12;
16const Q4_K_BLOCK_BYTES: usize = 4 + Q4_K_SCALE_SIZE + Q4_K_QK / 2; // 144
17
18/// Bit-unpacker matching candle's `quantized::utils::get_scale_min_k4`.
19fn get_scale_min_k4(j: usize, q: &[u8]) -> (u8, u8) {
20    if j < 4 {
21        (q[j] & 63, q[j + 4] & 63)
22    } else {
23        let d = (q[j + 4] & 0xF) | ((q[j - 4] >> 6) << 4);
24        let m = (q[j + 4] >> 4) | ((q[j] >> 6) << 4);
25        (d, m)
26    }
27}
28
29/// Port of candle's CPU `BlockQ4K::to_float`. Bit-identical output for
30/// identical input — the test in `q4_k.rs` verifies our Metal kernel
31/// also matches.
32fn dequant_q4_k_cpu(bytes: &[u8], n_blocks: usize) -> Vec<f32> {
33    debug_assert_eq!(bytes.len(), n_blocks * Q4_K_BLOCK_BYTES);
34    let mut out = Vec::with_capacity(n_blocks * Q4_K_QK);
35    for b in 0..n_blocks {
36        let off = b * Q4_K_BLOCK_BYTES;
37        let d = f16::from_le_bytes([bytes[off], bytes[off + 1]]).to_f32();
38        let dmin = f16::from_le_bytes([bytes[off + 2], bytes[off + 3]]).to_f32();
39        let scales = &bytes[off + 4..off + 4 + Q4_K_SCALE_SIZE];
40        let qs = &bytes[off + 4 + Q4_K_SCALE_SIZE..off + Q4_K_BLOCK_BYTES];
41
42        let mut is = 0usize;
43        for j in (0..Q4_K_QK).step_by(64) {
44            let q_chunk = &qs[j / 2..j / 2 + 32];
45            let (sc1, mn1) = get_scale_min_k4(is, scales);
46            let d1 = d * sc1 as f32;
47            let m1 = dmin * mn1 as f32;
48            let (sc2, mn2) = get_scale_min_k4(is + 1, scales);
49            let d2 = d * sc2 as f32;
50            let m2 = dmin * mn2 as f32;
51            for q in q_chunk {
52                out.push(d1 * (q & 0xF) as f32 - m1);
53            }
54            for q in q_chunk {
55                out.push(d2 * (q >> 4) as f32 - m2);
56            }
57            is += 2;
58        }
59    }
60    out
61}
62
63#[allow(clippy::too_many_arguments)]
64fn validate_gated_delta_rule_shape(
65    query_len: usize,
66    key_len: usize,
67    value_len_actual: usize,
68    g_len: usize,
69    beta_len: usize,
70    initial_state_len: usize,
71    out_len: usize,
72    final_state_len: usize,
73    tokens: usize,
74    key_heads: usize,
75    value_heads: usize,
76    key_dim: usize,
77    value_dim: usize,
78) -> Result<()> {
79    if tokens == 0 || key_heads == 0 || value_heads == 0 || key_dim == 0 || value_dim == 0 {
80        return Err(FerrumError::model(format!(
81            "gated_delta_rule shape must be positive, got tokens={tokens} key_heads={key_heads} value_heads={value_heads} key_dim={key_dim} value_dim={value_dim}"
82        )));
83    }
84    if value_heads % key_heads != 0 {
85        return Err(FerrumError::model(format!(
86            "gated_delta_rule value_heads {value_heads} must be divisible by key_heads {key_heads}"
87        )));
88    }
89
90    for (label, actual, expected) in [
91        ("query", query_len, tokens * key_heads * key_dim),
92        ("key", key_len, tokens * key_heads * key_dim),
93        ("value", value_len_actual, tokens * value_heads * value_dim),
94        ("g", g_len, tokens * value_heads),
95        ("beta", beta_len, tokens * value_heads),
96        (
97            "initial_state",
98            initial_state_len,
99            value_heads * value_dim * key_dim,
100        ),
101        ("out", out_len, tokens * value_heads * value_dim),
102        (
103            "final_state",
104            final_state_len,
105            value_heads * value_dim * key_dim,
106        ),
107    ] {
108        if actual < expected {
109            return Err(FerrumError::model(format!(
110                "gated_delta_rule {label} length {actual} < expected {expected}"
111            )));
112        }
113    }
114    Ok(())
115}
116
117fn sigmoid(x: f32) -> f32 {
118    if x >= 0.0 {
119        let z = (-x).exp();
120        1.0 / (1.0 + z)
121    } else {
122        let z = x.exp();
123        z / (1.0 + z)
124    }
125}
126
127fn silu(x: f32) -> f32 {
128    x * sigmoid(x)
129}
130
131fn softplus(x: f32) -> f32 {
132    if x > 20.0 {
133        x
134    } else if x < -20.0 {
135        x.exp()
136    } else {
137        (1.0 + x.exp()).ln()
138    }
139}
140
141#[allow(clippy::too_many_arguments)]
142fn validate_linear_attention_prepare_shape(
143    mixed_qkv_raw_len: usize,
144    conv_weight_len: usize,
145    a_raw_len: usize,
146    b_raw_len: usize,
147    a_log_len: usize,
148    dt_bias_len: usize,
149    query_len: usize,
150    key_len: usize,
151    value_len_actual: usize,
152    g_len: usize,
153    beta_len: usize,
154    tokens: usize,
155    key_heads: usize,
156    value_heads: usize,
157    key_dim: usize,
158    value_dim: usize,
159    conv_kernel: usize,
160) -> Result<()> {
161    if tokens == 0
162        || key_heads == 0
163        || value_heads == 0
164        || key_dim == 0
165        || value_dim == 0
166        || conv_kernel == 0
167    {
168        return Err(FerrumError::model(format!(
169            "linear_attention_prepare shape must be positive, got tokens={tokens} key_heads={key_heads} value_heads={value_heads} key_dim={key_dim} value_dim={value_dim} conv_kernel={conv_kernel}"
170        )));
171    }
172
173    let qk_total = key_heads * key_dim;
174    let value_total = value_heads * value_dim;
175    let conv_channels = 2 * qk_total + value_total;
176    for (label, actual, expected) in [
177        ("mixed_qkv_raw", mixed_qkv_raw_len, tokens * conv_channels),
178        ("conv_weight", conv_weight_len, conv_channels * conv_kernel),
179        ("a_raw", a_raw_len, tokens * value_heads),
180        ("b_raw", b_raw_len, tokens * value_heads),
181        ("a_log", a_log_len, value_heads),
182        ("dt_bias", dt_bias_len, value_heads),
183        ("query", query_len, tokens * qk_total),
184        ("key", key_len, tokens * qk_total),
185        ("value", value_len_actual, tokens * value_total),
186        ("g", g_len, tokens * value_heads),
187        ("beta", beta_len, tokens * value_heads),
188    ] {
189        if actual < expected {
190            return Err(FerrumError::model(format!(
191                "linear_attention_prepare {label} length {actual} < expected {expected}"
192            )));
193        }
194    }
195    Ok(())
196}
197
198#[allow(clippy::too_many_arguments)]
199fn validate_linear_attention_decode_prepare_shape(
200    mixed_qkv_raw_len: usize,
201    conv_weight_len: usize,
202    conv_state_len: usize,
203    a_raw_len: usize,
204    b_raw_len: usize,
205    a_log_len: usize,
206    dt_bias_len: usize,
207    query_len: usize,
208    key_len: usize,
209    value_len_actual: usize,
210    g_len: usize,
211    beta_len: usize,
212    next_conv_state_len: usize,
213    key_heads: usize,
214    value_heads: usize,
215    key_dim: usize,
216    value_dim: usize,
217    conv_kernel: usize,
218) -> Result<()> {
219    if key_heads == 0 || value_heads == 0 || key_dim == 0 || value_dim == 0 || conv_kernel == 0 {
220        return Err(FerrumError::model(format!(
221            "linear_attention_decode_prepare shape must be positive, got key_heads={key_heads} value_heads={value_heads} key_dim={key_dim} value_dim={value_dim} conv_kernel={conv_kernel}"
222        )));
223    }
224
225    let qk_total = key_heads * key_dim;
226    let value_total = value_heads * value_dim;
227    let conv_channels = 2 * qk_total + value_total;
228    let conv_state_elements = conv_channels * conv_kernel.saturating_sub(1);
229    for (label, actual, expected) in [
230        ("mixed_qkv_raw", mixed_qkv_raw_len, conv_channels),
231        ("conv_weight", conv_weight_len, conv_channels * conv_kernel),
232        ("conv_state", conv_state_len, conv_state_elements),
233        ("a_raw", a_raw_len, value_heads),
234        ("b_raw", b_raw_len, value_heads),
235        ("a_log", a_log_len, value_heads),
236        ("dt_bias", dt_bias_len, value_heads),
237        ("query", query_len, qk_total),
238        ("key", key_len, qk_total),
239        ("value", value_len_actual, value_total),
240        ("g", g_len, value_heads),
241        ("beta", beta_len, value_heads),
242        ("next_conv_state", next_conv_state_len, conv_state_elements),
243    ] {
244        if actual < expected {
245            return Err(FerrumError::model(format!(
246                "linear_attention_decode_prepare {label} length {actual} < expected {expected}"
247            )));
248        }
249    }
250    Ok(())
251}
252
253fn validate_gated_rms_norm_shape(
254    core_len: usize,
255    z_len: usize,
256    weight_len: usize,
257    out_len: usize,
258    tokens: usize,
259    heads: usize,
260    dim: usize,
261) -> Result<()> {
262    if tokens == 0 || heads == 0 || dim == 0 {
263        return Err(FerrumError::model(format!(
264            "gated_rms_norm shape must be positive, got tokens={tokens} heads={heads} dim={dim}"
265        )));
266    }
267    let expected = tokens * heads * dim;
268    for (label, actual, expected) in [
269        ("core", core_len, expected),
270        ("z", z_len, expected),
271        ("weight", weight_len, dim),
272        ("out", out_len, expected),
273    ] {
274        if actual < expected {
275            return Err(FerrumError::model(format!(
276                "gated_rms_norm {label} length {actual} < expected {expected}"
277            )));
278        }
279    }
280    Ok(())
281}
282
283pub struct CpuBackend;
284
285#[cfg(target_os = "macos")]
286unsafe extern "C" {
287    unsafe fn cblas_sgemm(
288        order: i32,
289        transa: i32,
290        transb: i32,
291        m: i32,
292        n: i32,
293        k: i32,
294        alpha: f32,
295        a: *const f32,
296        lda: i32,
297        b: *const f32,
298        ldb: i32,
299        beta: f32,
300        c: *mut f32,
301        ldc: i32,
302    );
303    fn vDSP_dotpr(
304        a: *const f32,
305        a_stride: i32,
306        b: *const f32,
307        b_stride: i32,
308        result: *mut f32,
309        n: u64,
310    );
311}
312
313/// CPU-side GPTQ store — dequantized f32 weights in row-major [n, k] layout.
314/// Trades memory for simplicity: repack once at load, then run normal GEMM.
315pub struct CpuGptqStore {
316    pub weight_f32: Vec<f32>, // [n, k] row-major
317    pub k: usize,
318    pub n: usize,
319}
320
321/// CPU-side container for any GGUF k-quant flavour. Each variant holds
322/// the dense fp32 weights post-eager-dequant — CPU isn't the bench
323/// target so we don't pay the complexity of on-the-fly dequant here;
324/// the variant tag exists so `gemm_quant` can route consistently.
325///
326/// New k-quant types (Q5_K / Q6_K / Q8_0) become new variants — no
327/// trait churn, just a new arm in `load_quant` and `gemm_quant`.
328pub enum CpuQuantStore {
329    Q4K {
330        weights: Vec<f32>, // [n_rows, n_cols] row-major
331        n_rows: usize,
332        n_cols: usize,
333    },
334}
335
336impl Backend for CpuBackend {
337    type Buffer = Vec<f32>;
338    type Context = ();
339    // type GptqStore: removed in Phase C step 4e. CpuGptqStore is now
340    // a private (crate-internal) detail of CpuMarlinExpertStack.
341
342    type Timer = crate::backend::timer::CpuTimer;
343    fn make_timer() -> Self::Timer {
344        crate::backend::timer::CpuTimer::new()
345    }
346
347    fn new_context() -> Self::Context {}
348    fn sync(_ctx: &mut Self::Context) {}
349    fn activation_elem_size_bytes() -> usize {
350        std::mem::size_of::<f32>()
351    }
352
353    fn zero_buffer(_ctx: &mut Self::Context, buf: &mut Self::Buffer, len: usize) -> Result<()> {
354        if buf.len() < len {
355            return Err(FerrumError::model(format!(
356                "zero_buffer length {len} exceeds CPU buffer length {}",
357                buf.len()
358            )));
359        }
360        for value in &mut buf[..len] {
361            *value = 0.0;
362        }
363        Ok(())
364    }
365
366    /// Phase D step 2+3: typed alloc. CPU Buffer is Vec<f32> — bytes
367    /// are dtype-erased, so we size the underlying Vec to hold `n`
368    /// elements of `dtype` (bit-cast at read/write time).
369    fn alloc_typed(dtype: crate::backend::Dtype, n: usize) -> Self::Buffer {
370        // f32 storage; for i8 we round up to 4-byte boundary so the
371        // Vec<f32> length covers all i8 elements.
372        let bytes = n * dtype.bytes_per_elem();
373        let f32_len = bytes.div_ceil(4);
374        vec![0.0f32; f32_len]
375    }
376
377    /// Phase D step 2+3: typed upload. Bit-cast host data into f32
378    /// words (CPU buffer is dtype-erased Vec<f32>, see alloc_typed).
379    fn from_slice_typed<T: crate::backend::HostDtype>(data: &[T]) -> Self::Buffer {
380        let bytes = data.len() * std::mem::size_of::<T>();
381        let f32_len = bytes.div_ceil(4);
382        let mut out = vec![0.0f32; f32_len];
383        unsafe {
384            std::ptr::copy_nonoverlapping(
385                data.as_ptr() as *const u8,
386                out.as_mut_ptr() as *mut u8,
387                bytes,
388            );
389        }
390        out
391    }
392
393    /// Phase D step 2+3: typed in-place write. Bit-cast bytes into
394    /// the dtype-erased f32 storage.
395    fn write_typed<T: crate::backend::HostDtype>(
396        _ctx: &mut Self::Context,
397        dst: &mut Self::Buffer,
398        data: &[T],
399    ) {
400        let bytes = data.len() * std::mem::size_of::<T>();
401        debug_assert!(
402            bytes <= dst.len() * 4,
403            "CpuBackend::write_typed: src bytes {} > dst bytes {}",
404            bytes,
405            dst.len() * 4
406        );
407        unsafe {
408            std::ptr::copy_nonoverlapping(
409                data.as_ptr() as *const u8,
410                dst.as_mut_ptr() as *mut u8,
411                bytes,
412            );
413        }
414    }
415
416    fn fused_silu_mul_split_strided(
417        _ctx: &mut Self::Context,
418        gate_up: &Self::Buffer,
419        in_row_offset: usize,
420        out: &mut Self::Buffer,
421        out_row_offset: usize,
422        tokens: usize,
423        intermediate: usize,
424    ) {
425        let in_per_row = 2 * intermediate;
426        let in_start = in_row_offset * in_per_row;
427        let out_start = out_row_offset * intermediate;
428        for r in 0..tokens {
429            for c in 0..intermediate {
430                let g = gate_up[in_start + r * in_per_row + c];
431                let u = gate_up[in_start + r * in_per_row + intermediate + c];
432                let silu = g / (1.0 + (-g).exp());
433                out[out_start + r * intermediate + c] = silu * u;
434            }
435        }
436    }
437
438    fn gemm(
439        _ctx: &mut Self::Context,
440        a: &Self::Buffer,
441        b: &Self::Buffer,
442        out: &mut Self::Buffer,
443        m: usize,
444        n: usize,
445        k: usize,
446    ) {
447        assert!(
448            a.len() >= m * k,
449            "gemm: a too small len={} m={m} k={k}",
450            a.len()
451        );
452        assert!(
453            b.len() >= n * k,
454            "gemm: b too small len={} n={n} k={k}",
455            b.len()
456        );
457        assert!(
458            out.len() >= m * n,
459            "gemm: out too small len={} m={m} n={n}",
460            out.len()
461        );
462        #[cfg(target_os = "macos")]
463        unsafe {
464            cblas_sgemm(
465                101,
466                111,
467                112,
468                m as i32,
469                n as i32,
470                k as i32,
471                1.0,
472                a.as_ptr(),
473                k as i32,
474                b.as_ptr(),
475                k as i32,
476                0.0,
477                out.as_mut_ptr(),
478                n as i32,
479            );
480        }
481        #[cfg(not(target_os = "macos"))]
482        {
483            for i in 0..m {
484                for j in 0..n {
485                    let mut sum = 0.0f64;
486                    for p in 0..k {
487                        sum += a[i * k + p] as f64 * b[j * k + p] as f64;
488                    }
489                    out[i * n + j] = sum as f32;
490                }
491            }
492        }
493    }
494
495    fn rms_norm(
496        _ctx: &mut Self::Context,
497        x: &Self::Buffer,
498        w: &Self::Buffer,
499        eps: f32,
500        out: &mut Self::Buffer,
501        tokens: usize,
502        dim: usize,
503    ) {
504        for t in 0..tokens {
505            let row = &x[t * dim..(t + 1) * dim];
506            let o = &mut out[t * dim..(t + 1) * dim];
507            let sum_sq = dot_product(row, row);
508            let inv = 1.0f32 / (sum_sq / dim as f32 + eps).sqrt();
509            for i in 0..dim {
510                o[i] = row[i] * inv * w[i];
511            }
512        }
513    }
514
515    fn fused_add_rms_norm(
516        _ctx: &mut Self::Context,
517        residual: &mut Self::Buffer,
518        x: &Self::Buffer,
519        w: &Self::Buffer,
520        eps: f32,
521        out: &mut Self::Buffer,
522        tokens: usize,
523        dim: usize,
524    ) {
525        for t in 0..tokens {
526            let off = t * dim;
527            for i in 0..dim {
528                residual[off + i] += x[off + i];
529            }
530            let row = &residual[off..off + dim];
531            let o = &mut out[off..off + dim];
532            let sum_sq = dot_product(row, row);
533            let inv = 1.0f32 / (sum_sq / dim as f32 + eps).sqrt();
534            for i in 0..dim {
535                o[i] = row[i] * inv * w[i];
536            }
537        }
538    }
539
540    fn flash_attention(
541        _ctx: &mut Self::Context,
542        q: &Self::Buffer,
543        k: &Self::Buffer,
544        v: &Self::Buffer,
545        out: &mut Self::Buffer,
546        batch: usize,
547        q_len: usize,
548        kv_len: usize,
549        pos_offset: usize,
550        cfg: &AttnConfig,
551    ) {
552        cpu_attention(
553            q, k, v, out, batch, q_len, kv_len, cfg.causal, pos_offset, cfg,
554        );
555    }
556
557    #[allow(clippy::too_many_arguments)]
558    fn recurrent_gated_delta_rule_f32(
559        _ctx: &mut Self::Context,
560        query: &Self::Buffer,
561        key: &Self::Buffer,
562        value: &Self::Buffer,
563        g: &Self::Buffer,
564        beta: &Self::Buffer,
565        initial_state: &Self::Buffer,
566        out: &mut Self::Buffer,
567        final_state: &mut Self::Buffer,
568        tokens: usize,
569        key_heads: usize,
570        value_heads: usize,
571        key_dim: usize,
572        value_dim: usize,
573        use_qk_l2norm: bool,
574        scale: f32,
575    ) -> Result<()> {
576        validate_gated_delta_rule_shape(
577            query.len(),
578            key.len(),
579            value.len(),
580            g.len(),
581            beta.len(),
582            initial_state.len(),
583            out.len(),
584            final_state.len(),
585            tokens,
586            key_heads,
587            value_heads,
588            key_dim,
589            value_dim,
590        )?;
591
592        let repeat_factor = value_heads / key_heads;
593        let state_len = value_heads * value_dim * key_dim;
594        final_state[..state_len].copy_from_slice(&initial_state[..state_len]);
595
596        for token in 0..tokens {
597            for value_head in 0..value_heads {
598                let key_head = value_head / repeat_factor;
599                let mut q_inv = 1.0;
600                let mut k_inv = 1.0;
601                if use_qk_l2norm {
602                    let mut q_norm = 0.0;
603                    let mut k_norm = 0.0;
604                    for kd in 0..key_dim {
605                        let qk_idx = ((token * key_heads + key_head) * key_dim) + kd;
606                        q_norm += query[qk_idx] * query[qk_idx];
607                        k_norm += key[qk_idx] * key[qk_idx];
608                    }
609                    q_inv = (q_norm + 1e-6).sqrt().recip();
610                    k_inv = (k_norm + 1e-6).sqrt().recip();
611                }
612                let gate_idx = token * value_heads + value_head;
613                let decay = g[gate_idx].exp();
614                let beta_t = beta[gate_idx];
615                for vd in 0..value_dim {
616                    let state_base = (value_head * value_dim + vd) * key_dim;
617                    let mut kv_mem = 0.0;
618                    for kd in 0..key_dim {
619                        let qk_idx = ((token * key_heads + key_head) * key_dim) + kd;
620                        let state_idx = state_base + kd;
621                        final_state[state_idx] *= decay;
622                        kv_mem += final_state[state_idx] * (key[qk_idx] * k_inv);
623                    }
624                    let value_idx = ((token * value_heads + value_head) * value_dim) + vd;
625                    let delta = (value[value_idx] - kv_mem) * beta_t;
626                    for kd in 0..key_dim {
627                        let qk_idx = ((token * key_heads + key_head) * key_dim) + kd;
628                        final_state[state_base + kd] += delta * (key[qk_idx] * k_inv);
629                    }
630                    let mut acc = 0.0;
631                    for kd in 0..key_dim {
632                        let qk_idx = ((token * key_heads + key_head) * key_dim) + kd;
633                        acc += final_state[state_base + kd] * (query[qk_idx] * q_inv * scale);
634                    }
635                    out[value_idx] = acc;
636                }
637            }
638        }
639        Ok(())
640    }
641
642    #[allow(clippy::too_many_arguments)]
643    fn recurrent_gated_delta_rule_batch_f32(
644        _ctx: &mut Self::Context,
645        query: &Self::Buffer,
646        key: &Self::Buffer,
647        value: &Self::Buffer,
648        g: &Self::Buffer,
649        beta: &Self::Buffer,
650        initial_states: &Self::Buffer,
651        out: &mut Self::Buffer,
652        final_states: &mut Self::Buffer,
653        batch: usize,
654        key_heads: usize,
655        value_heads: usize,
656        key_dim: usize,
657        value_dim: usize,
658        use_qk_l2norm: bool,
659        scale: f32,
660    ) -> Result<()> {
661        if batch == 0 {
662            return Err(FerrumError::model(
663                "gated_delta_rule_batch batch must be positive",
664            ));
665        }
666        let state_len = value_heads * value_dim * key_dim;
667        validate_gated_delta_rule_shape(
668            query.len(),
669            key.len(),
670            value.len(),
671            g.len(),
672            beta.len(),
673            initial_states.len(),
674            out.len(),
675            final_states.len(),
676            batch,
677            key_heads,
678            value_heads,
679            key_dim,
680            value_dim,
681        )?;
682        if initial_states.len() < batch * state_len || final_states.len() < batch * state_len {
683            return Err(FerrumError::model(format!(
684                "gated_delta_rule_batch state length too small: initial={} final={} expected={}",
685                initial_states.len(),
686                final_states.len(),
687                batch * state_len
688            )));
689        }
690
691        let repeat_factor = value_heads / key_heads;
692        for row in 0..batch {
693            let row_state_base = row * state_len;
694            final_states[row_state_base..row_state_base + state_len]
695                .copy_from_slice(&initial_states[row_state_base..row_state_base + state_len]);
696            for value_head in 0..value_heads {
697                let key_head = value_head / repeat_factor;
698                let mut q_inv = 1.0;
699                let mut k_inv = 1.0;
700                if use_qk_l2norm {
701                    let mut q_norm = 0.0;
702                    let mut k_norm = 0.0;
703                    for kd in 0..key_dim {
704                        let qk_idx = ((row * key_heads + key_head) * key_dim) + kd;
705                        q_norm += query[qk_idx] * query[qk_idx];
706                        k_norm += key[qk_idx] * key[qk_idx];
707                    }
708                    q_inv = (q_norm + 1e-6).sqrt().recip();
709                    k_inv = (k_norm + 1e-6).sqrt().recip();
710                }
711                let gate_idx = row * value_heads + value_head;
712                let decay = g[gate_idx].exp();
713                let beta_t = beta[gate_idx];
714                for vd in 0..value_dim {
715                    let state_base = row_state_base + (value_head * value_dim + vd) * key_dim;
716                    let mut kv_mem = 0.0;
717                    for kd in 0..key_dim {
718                        let qk_idx = ((row * key_heads + key_head) * key_dim) + kd;
719                        let state_idx = state_base + kd;
720                        final_states[state_idx] *= decay;
721                        kv_mem += final_states[state_idx] * (key[qk_idx] * k_inv);
722                    }
723                    let value_idx = ((row * value_heads + value_head) * value_dim) + vd;
724                    let delta = (value[value_idx] - kv_mem) * beta_t;
725                    for kd in 0..key_dim {
726                        let qk_idx = ((row * key_heads + key_head) * key_dim) + kd;
727                        final_states[state_base + kd] += delta * (key[qk_idx] * k_inv);
728                    }
729                    let mut acc = 0.0;
730                    for kd in 0..key_dim {
731                        let qk_idx = ((row * key_heads + key_head) * key_dim) + kd;
732                        acc += final_states[state_base + kd] * (query[qk_idx] * q_inv * scale);
733                    }
734                    out[value_idx] = acc;
735                }
736            }
737        }
738        Ok(())
739    }
740
741    #[allow(clippy::too_many_arguments)]
742    fn recurrent_gated_delta_rule_varlen_f32(
743        _ctx: &mut Self::Context,
744        query: &Self::Buffer,
745        key: &Self::Buffer,
746        value: &Self::Buffer,
747        g: &Self::Buffer,
748        beta: &Self::Buffer,
749        initial_states: &Self::Buffer,
750        cu_seqlens: &Self::Buffer,
751        out: &mut Self::Buffer,
752        final_states: &mut Self::Buffer,
753        batch: usize,
754        total_tokens: usize,
755        key_heads: usize,
756        value_heads: usize,
757        key_dim: usize,
758        value_dim: usize,
759        use_qk_l2norm: bool,
760        scale: f32,
761    ) -> Result<()> {
762        if batch == 0
763            || total_tokens == 0
764            || key_heads == 0
765            || value_heads == 0
766            || key_dim == 0
767            || value_dim == 0
768        {
769            return Err(FerrumError::model(format!(
770                "gated_delta_rule_varlen shape must be positive, got batch={batch} total_tokens={total_tokens} key_heads={key_heads} value_heads={value_heads} key_dim={key_dim} value_dim={value_dim}"
771            )));
772        }
773        let cu = cpu_read_u32_buffer(cu_seqlens, batch + 1, "gated_delta_rule_varlen cu_seqlens")?;
774        if cu.first().copied() != Some(0) || cu.last().copied() != Some(total_tokens as u32) {
775            return Err(FerrumError::model(format!(
776                "gated_delta_rule_varlen cu_seqlens must start at 0 and end at total_tokens {total_tokens}, got first={:?} last={:?}",
777                cu.first(),
778                cu.last()
779            )));
780        }
781        for seq in 0..batch {
782            if cu[seq + 1] <= cu[seq] {
783                return Err(FerrumError::model(format!(
784                    "gated_delta_rule_varlen sequence {seq} has empty or non-monotonic range {}..{}",
785                    cu[seq],
786                    cu[seq + 1]
787                )));
788            }
789        }
790
791        let state_len = value_heads * value_dim * key_dim;
792        validate_gated_delta_rule_shape(
793            query.len(),
794            key.len(),
795            value.len(),
796            g.len(),
797            beta.len(),
798            state_len,
799            out.len(),
800            state_len,
801            total_tokens,
802            key_heads,
803            value_heads,
804            key_dim,
805            value_dim,
806        )?;
807        if initial_states.len() < batch * state_len || final_states.len() < batch * state_len {
808            return Err(FerrumError::model(format!(
809                "gated_delta_rule_varlen state length too small: initial={} final={} expected={}",
810                initial_states.len(),
811                final_states.len(),
812                batch * state_len
813            )));
814        }
815
816        let repeat_factor = value_heads / key_heads;
817        for seq in 0..batch {
818            let token_start = cu[seq] as usize;
819            let token_end = cu[seq + 1] as usize;
820            let row_state_base = seq * state_len;
821            final_states[row_state_base..row_state_base + state_len]
822                .copy_from_slice(&initial_states[row_state_base..row_state_base + state_len]);
823
824            for token in token_start..token_end {
825                for value_head in 0..value_heads {
826                    let key_head = value_head / repeat_factor;
827                    let mut q_inv = 1.0;
828                    let mut k_inv = 1.0;
829                    if use_qk_l2norm {
830                        let mut q_norm = 0.0;
831                        let mut k_norm = 0.0;
832                        for kd in 0..key_dim {
833                            let qk_idx = ((token * key_heads + key_head) * key_dim) + kd;
834                            q_norm += query[qk_idx] * query[qk_idx];
835                            k_norm += key[qk_idx] * key[qk_idx];
836                        }
837                        q_inv = (q_norm + 1e-6).sqrt().recip();
838                        k_inv = (k_norm + 1e-6).sqrt().recip();
839                    }
840                    let gate_idx = token * value_heads + value_head;
841                    let decay = g[gate_idx].exp();
842                    let beta_t = beta[gate_idx];
843                    for vd in 0..value_dim {
844                        let state_base = row_state_base + (value_head * value_dim + vd) * key_dim;
845                        let mut kv_mem = 0.0;
846                        for kd in 0..key_dim {
847                            let qk_idx = ((token * key_heads + key_head) * key_dim) + kd;
848                            let state_idx = state_base + kd;
849                            final_states[state_idx] *= decay;
850                            kv_mem += final_states[state_idx] * (key[qk_idx] * k_inv);
851                        }
852                        let value_idx = ((token * value_heads + value_head) * value_dim) + vd;
853                        let delta = (value[value_idx] - kv_mem) * beta_t;
854                        for kd in 0..key_dim {
855                            let qk_idx = ((token * key_heads + key_head) * key_dim) + kd;
856                            final_states[state_base + kd] += delta * (key[qk_idx] * k_inv);
857                        }
858                        let mut acc = 0.0;
859                        for kd in 0..key_dim {
860                            let qk_idx = ((token * key_heads + key_head) * key_dim) + kd;
861                            acc += final_states[state_base + kd] * (query[qk_idx] * q_inv * scale);
862                        }
863                        out[value_idx] = acc;
864                    }
865                }
866            }
867        }
868        Ok(())
869    }
870
871    #[allow(clippy::too_many_arguments)]
872    fn linear_attention_prepare_f32(
873        _ctx: &mut Self::Context,
874        mixed_qkv_raw: &Self::Buffer,
875        conv_weight: &Self::Buffer,
876        a_raw: &Self::Buffer,
877        b_raw: &Self::Buffer,
878        a_log: &Self::Buffer,
879        dt_bias: &Self::Buffer,
880        query: &mut Self::Buffer,
881        key: &mut Self::Buffer,
882        value: &mut Self::Buffer,
883        g: &mut Self::Buffer,
884        beta: &mut Self::Buffer,
885        tokens: usize,
886        key_heads: usize,
887        value_heads: usize,
888        key_dim: usize,
889        value_dim: usize,
890        conv_kernel: usize,
891        apply_qk_l2norm: bool,
892    ) -> Result<()> {
893        validate_linear_attention_prepare_shape(
894            mixed_qkv_raw.len(),
895            conv_weight.len(),
896            a_raw.len(),
897            b_raw.len(),
898            a_log.len(),
899            dt_bias.len(),
900            query.len(),
901            key.len(),
902            value.len(),
903            g.len(),
904            beta.len(),
905            tokens,
906            key_heads,
907            value_heads,
908            key_dim,
909            value_dim,
910            conv_kernel,
911        )?;
912
913        let qk_total = key_heads * key_dim;
914        let value_total = value_heads * value_dim;
915        let conv_channels = 2 * qk_total + value_total;
916        let pad = conv_kernel - 1;
917        for token in 0..tokens {
918            for channel in 0..conv_channels {
919                let mut acc = 0.0;
920                for kernel_idx in 0..conv_kernel {
921                    let padded = token + kernel_idx;
922                    if padded >= pad {
923                        let src_token = padded - pad;
924                        if src_token < tokens {
925                            acc += mixed_qkv_raw[src_token * conv_channels + channel]
926                                * conv_weight[channel * conv_kernel + kernel_idx];
927                        }
928                    }
929                }
930                let conv = silu(acc);
931                if channel < qk_total {
932                    query[token * qk_total + channel] = conv;
933                } else if channel < 2 * qk_total {
934                    key[token * qk_total + (channel - qk_total)] = conv;
935                } else {
936                    value[token * value_total + (channel - 2 * qk_total)] = conv;
937                }
938            }
939
940            for value_head in 0..value_heads {
941                let gate_idx = token * value_heads + value_head;
942                g[gate_idx] =
943                    -a_log[value_head].exp() * softplus(a_raw[gate_idx] + dt_bias[value_head]);
944                beta[gate_idx] = sigmoid(b_raw[gate_idx]);
945            }
946        }
947
948        if apply_qk_l2norm {
949            for row in 0..tokens * key_heads {
950                let base = row * key_dim;
951                let mut q_sum = 0.0;
952                let mut k_sum = 0.0;
953                for d in 0..key_dim {
954                    q_sum += query[base + d] * query[base + d];
955                    k_sum += key[base + d] * key[base + d];
956                }
957                let q_inv = (q_sum + 1e-6).sqrt().recip();
958                let k_inv = (k_sum + 1e-6).sqrt().recip();
959                for d in 0..key_dim {
960                    query[base + d] *= q_inv;
961                    key[base + d] *= k_inv;
962                }
963            }
964        }
965        Ok(())
966    }
967
968    #[allow(clippy::too_many_arguments)]
969    fn linear_attention_prepare_varlen_f32(
970        _ctx: &mut Self::Context,
971        mixed_qkv_raw: &Self::Buffer,
972        conv_weight: &Self::Buffer,
973        initial_conv_states: &Self::Buffer,
974        a_raw: &Self::Buffer,
975        b_raw: &Self::Buffer,
976        a_log: &Self::Buffer,
977        dt_bias: &Self::Buffer,
978        cu_seqlens: &Self::Buffer,
979        token_seq_indices: &Self::Buffer,
980        query: &mut Self::Buffer,
981        key: &mut Self::Buffer,
982        value: &mut Self::Buffer,
983        g: &mut Self::Buffer,
984        beta: &mut Self::Buffer,
985        final_conv_states: &mut Self::Buffer,
986        batch: usize,
987        total_tokens: usize,
988        key_heads: usize,
989        value_heads: usize,
990        key_dim: usize,
991        value_dim: usize,
992        conv_kernel: usize,
993        apply_qk_l2norm: bool,
994    ) -> Result<()> {
995        if batch == 0 {
996            return Err(FerrumError::model(
997                "linear_attention_prepare_varlen batch must be positive",
998            ));
999        }
1000        validate_linear_attention_prepare_shape(
1001            mixed_qkv_raw.len(),
1002            conv_weight.len(),
1003            a_raw.len(),
1004            b_raw.len(),
1005            a_log.len(),
1006            dt_bias.len(),
1007            query.len(),
1008            key.len(),
1009            value.len(),
1010            g.len(),
1011            beta.len(),
1012            total_tokens,
1013            key_heads,
1014            value_heads,
1015            key_dim,
1016            value_dim,
1017            conv_kernel,
1018        )?;
1019        let qk_total = key_heads * key_dim;
1020        let value_total = value_heads * value_dim;
1021        let conv_channels = 2 * qk_total + value_total;
1022        let state_len = conv_kernel.saturating_sub(1);
1023        let conv_state_len = conv_channels * state_len;
1024        for (label, actual, expected) in [
1025            (
1026                "initial_conv_states",
1027                initial_conv_states.len(),
1028                batch * conv_state_len,
1029            ),
1030            (
1031                "final_conv_states",
1032                final_conv_states.len(),
1033                batch * conv_state_len,
1034            ),
1035        ] {
1036            if actual < expected {
1037                return Err(FerrumError::model(format!(
1038                    "linear_attention_prepare_varlen {label} length {actual} < expected {expected}"
1039                )));
1040            }
1041        }
1042        let cu = cpu_read_u32_buffer(
1043            cu_seqlens,
1044            batch + 1,
1045            "linear_attention_prepare_varlen cu_seqlens",
1046        )?;
1047        let token_rows = cpu_read_u32_buffer(
1048            token_seq_indices,
1049            total_tokens,
1050            "linear_attention_prepare_varlen token_seq_indices",
1051        )?;
1052        if cu.first().copied() != Some(0) || cu.last().copied() != Some(total_tokens as u32) {
1053            return Err(FerrumError::model(format!(
1054                "linear_attention_prepare_varlen cu_seqlens must start at 0 and end at total_tokens {total_tokens}, got first={:?} last={:?}",
1055                cu.first(),
1056                cu.last()
1057            )));
1058        }
1059        for seq in 0..batch {
1060            if cu[seq + 1] <= cu[seq] {
1061                return Err(FerrumError::model(format!(
1062                    "linear_attention_prepare_varlen sequence {seq} has empty or non-monotonic range {}..{}",
1063                    cu[seq],
1064                    cu[seq + 1]
1065                )));
1066            }
1067            for token in cu[seq] as usize..cu[seq + 1] as usize {
1068                if token_rows[token] != seq as u32 {
1069                    return Err(FerrumError::model(format!(
1070                        "linear_attention_prepare_varlen token_seq_indices[{token}]={} != seq {seq}",
1071                        token_rows[token]
1072                    )));
1073                }
1074            }
1075        }
1076
1077        for seq in 0..batch {
1078            let token_start = cu[seq] as usize;
1079            let token_end = cu[seq + 1] as usize;
1080            let seq_tokens = token_end - token_start;
1081            let state_row_base = seq * conv_state_len;
1082
1083            for token in token_start..token_end {
1084                let local_token = token - token_start;
1085                for channel in 0..conv_channels {
1086                    let state_base = state_row_base + channel * state_len;
1087                    let mut acc = 0.0;
1088                    for kernel_idx in 0..conv_kernel {
1089                        let source =
1090                            local_token as isize + kernel_idx as isize - state_len as isize;
1091                        let x = if source >= 0 {
1092                            mixed_qkv_raw[(token_start + source as usize) * conv_channels + channel]
1093                        } else {
1094                            initial_conv_states[state_base + (state_len as isize + source) as usize]
1095                        };
1096                        acc += x * conv_weight[channel * conv_kernel + kernel_idx];
1097                    }
1098                    let conv = silu(acc);
1099                    if channel < qk_total {
1100                        query[token * qk_total + channel] = conv;
1101                    } else if channel < 2 * qk_total {
1102                        key[token * qk_total + (channel - qk_total)] = conv;
1103                    } else {
1104                        value[token * value_total + (channel - 2 * qk_total)] = conv;
1105                    }
1106                }
1107
1108                for value_head in 0..value_heads {
1109                    let gate_idx = token * value_heads + value_head;
1110                    g[gate_idx] =
1111                        -a_log[value_head].exp() * softplus(a_raw[gate_idx] + dt_bias[value_head]);
1112                    beta[gate_idx] = sigmoid(b_raw[gate_idx]);
1113                }
1114            }
1115
1116            for channel in 0..conv_channels {
1117                let state_base = state_row_base + channel * state_len;
1118                for pos in 0..state_len {
1119                    let source = seq_tokens as isize + pos as isize - state_len as isize;
1120                    final_conv_states[state_base + pos] = if source >= 0 {
1121                        mixed_qkv_raw[(token_start + source as usize) * conv_channels + channel]
1122                    } else {
1123                        initial_conv_states[state_base + (state_len as isize + source) as usize]
1124                    };
1125                }
1126            }
1127        }
1128
1129        if apply_qk_l2norm {
1130            for row in 0..total_tokens * key_heads {
1131                let base = row * key_dim;
1132                let mut q_sum = 0.0;
1133                let mut k_sum = 0.0;
1134                for d in 0..key_dim {
1135                    q_sum += query[base + d] * query[base + d];
1136                    k_sum += key[base + d] * key[base + d];
1137                }
1138                let q_inv = (q_sum + 1e-6).sqrt().recip();
1139                let k_inv = (k_sum + 1e-6).sqrt().recip();
1140                for d in 0..key_dim {
1141                    query[base + d] *= q_inv;
1142                    key[base + d] *= k_inv;
1143                }
1144            }
1145        }
1146        Ok(())
1147    }
1148
1149    #[allow(clippy::too_many_arguments)]
1150    fn linear_attention_prepare_varlen_packed_qkvz_ba_f32(
1151        _ctx: &mut Self::Context,
1152        mixed_qkvz_raw: &Self::Buffer,
1153        ba_raw: &Self::Buffer,
1154        conv_weight: &Self::Buffer,
1155        initial_conv_states: &Self::Buffer,
1156        a_log: &Self::Buffer,
1157        dt_bias: &Self::Buffer,
1158        cu_seqlens: &Self::Buffer,
1159        token_seq_indices: &Self::Buffer,
1160        query: &mut Self::Buffer,
1161        key: &mut Self::Buffer,
1162        value: &mut Self::Buffer,
1163        z: &mut Self::Buffer,
1164        g: &mut Self::Buffer,
1165        beta: &mut Self::Buffer,
1166        final_conv_states: &mut Self::Buffer,
1167        batch: usize,
1168        total_tokens: usize,
1169        key_heads: usize,
1170        value_heads: usize,
1171        key_dim: usize,
1172        value_dim: usize,
1173        conv_kernel: usize,
1174        apply_qk_l2norm: bool,
1175    ) -> Result<()> {
1176        if batch == 0
1177            || total_tokens == 0
1178            || key_heads == 0
1179            || value_heads == 0
1180            || key_dim == 0
1181            || value_dim == 0
1182            || conv_kernel == 0
1183        {
1184            return Err(FerrumError::model(format!(
1185                "linear_attention_prepare_varlen_packed shape must be positive, got batch={batch} total_tokens={total_tokens} key_heads={key_heads} value_heads={value_heads} key_dim={key_dim} value_dim={value_dim} conv_kernel={conv_kernel}"
1186            )));
1187        }
1188        let qk_total = key_heads * key_dim;
1189        let value_total = value_heads * value_dim;
1190        let conv_channels = 2 * qk_total + value_total;
1191        let qkvz_width = conv_channels + value_total;
1192        let ba_width = 2 * value_heads;
1193        let state_len = conv_kernel.saturating_sub(1);
1194        let conv_state_len = conv_channels * state_len;
1195        for (label, actual, expected) in [
1196            (
1197                "mixed_qkvz_raw",
1198                mixed_qkvz_raw.len(),
1199                total_tokens * qkvz_width,
1200            ),
1201            ("ba_raw", ba_raw.len(), total_tokens * ba_width),
1202            (
1203                "conv_weight",
1204                conv_weight.len(),
1205                conv_channels * conv_kernel,
1206            ),
1207            (
1208                "initial_conv_states",
1209                initial_conv_states.len(),
1210                batch * conv_state_len,
1211            ),
1212            ("a_log", a_log.len(), value_heads),
1213            ("dt_bias", dt_bias.len(), value_heads),
1214            ("query", query.len(), total_tokens * qk_total),
1215            ("key", key.len(), total_tokens * qk_total),
1216            ("value", value.len(), total_tokens * value_total),
1217            ("z", z.len(), total_tokens * value_total),
1218            ("g", g.len(), total_tokens * value_heads),
1219            ("beta", beta.len(), total_tokens * value_heads),
1220            (
1221                "final_conv_states",
1222                final_conv_states.len(),
1223                batch * conv_state_len,
1224            ),
1225        ] {
1226            if actual < expected {
1227                return Err(FerrumError::model(format!(
1228                    "linear_attention_prepare_varlen_packed {label} length {actual} < expected {expected}"
1229                )));
1230            }
1231        }
1232        let cu = cpu_read_u32_buffer(
1233            cu_seqlens,
1234            batch + 1,
1235            "linear_attention_prepare_varlen_packed cu_seqlens",
1236        )?;
1237        let token_rows = cpu_read_u32_buffer(
1238            token_seq_indices,
1239            total_tokens,
1240            "linear_attention_prepare_varlen_packed token_seq_indices",
1241        )?;
1242        if cu.first().copied() != Some(0) || cu.last().copied() != Some(total_tokens as u32) {
1243            return Err(FerrumError::model(format!(
1244                "linear_attention_prepare_varlen_packed cu_seqlens must start at 0 and end at total_tokens {total_tokens}, got first={:?} last={:?}",
1245                cu.first(),
1246                cu.last()
1247            )));
1248        }
1249        for seq in 0..batch {
1250            if cu[seq + 1] <= cu[seq] {
1251                return Err(FerrumError::model(format!(
1252                    "linear_attention_prepare_varlen_packed sequence {seq} has empty or non-monotonic range {}..{}",
1253                    cu[seq],
1254                    cu[seq + 1]
1255                )));
1256            }
1257            for token in cu[seq] as usize..cu[seq + 1] as usize {
1258                if token_rows[token] != seq as u32 {
1259                    return Err(FerrumError::model(format!(
1260                        "linear_attention_prepare_varlen_packed token_seq_indices[{token}]={} != seq {seq}",
1261                        token_rows[token]
1262                    )));
1263                }
1264            }
1265        }
1266
1267        for seq in 0..batch {
1268            let token_start = cu[seq] as usize;
1269            let token_end = cu[seq + 1] as usize;
1270            let seq_tokens = token_end - token_start;
1271            let state_row_base = seq * conv_state_len;
1272
1273            for token in token_start..token_end {
1274                let local_token = token - token_start;
1275                for channel in 0..conv_channels {
1276                    let state_base = state_row_base + channel * state_len;
1277                    let mut acc = 0.0;
1278                    for kernel_idx in 0..conv_kernel {
1279                        let source =
1280                            local_token as isize + kernel_idx as isize - state_len as isize;
1281                        let x = if source >= 0 {
1282                            mixed_qkvz_raw[(token_start + source as usize) * qkvz_width + channel]
1283                        } else {
1284                            initial_conv_states[state_base + (state_len as isize + source) as usize]
1285                        };
1286                        acc += x * conv_weight[channel * conv_kernel + kernel_idx];
1287                    }
1288                    let conv = silu(acc);
1289                    if channel < qk_total {
1290                        query[token * qk_total + channel] = conv;
1291                    } else if channel < 2 * qk_total {
1292                        key[token * qk_total + (channel - qk_total)] = conv;
1293                    } else {
1294                        value[token * value_total + (channel - 2 * qk_total)] = conv;
1295                    }
1296                }
1297
1298                let qkvz_base = token * qkvz_width;
1299                let z_base = token * value_total;
1300                z[z_base..z_base + value_total].copy_from_slice(
1301                    &mixed_qkvz_raw[qkvz_base + conv_channels..qkvz_base + qkvz_width],
1302                );
1303
1304                let ba_base = token * ba_width;
1305                for value_head in 0..value_heads {
1306                    let gate_idx = token * value_heads + value_head;
1307                    let b = ba_raw[ba_base + value_head];
1308                    let a = ba_raw[ba_base + value_heads + value_head];
1309                    g[gate_idx] = -a_log[value_head].exp() * softplus(a + dt_bias[value_head]);
1310                    beta[gate_idx] = sigmoid(b);
1311                }
1312            }
1313
1314            for channel in 0..conv_channels {
1315                let state_base = state_row_base + channel * state_len;
1316                for pos in 0..state_len {
1317                    let source = seq_tokens as isize + pos as isize - state_len as isize;
1318                    final_conv_states[state_base + pos] = if source >= 0 {
1319                        mixed_qkvz_raw[(token_start + source as usize) * qkvz_width + channel]
1320                    } else {
1321                        initial_conv_states[state_base + (state_len as isize + source) as usize]
1322                    };
1323                }
1324            }
1325        }
1326
1327        if apply_qk_l2norm {
1328            for row in 0..total_tokens * key_heads {
1329                let base = row * key_dim;
1330                let mut q_sum = 0.0;
1331                let mut k_sum = 0.0;
1332                for d in 0..key_dim {
1333                    q_sum += query[base + d] * query[base + d];
1334                    k_sum += key[base + d] * key[base + d];
1335                }
1336                let q_inv = (q_sum + 1e-6).sqrt().recip();
1337                let k_inv = (k_sum + 1e-6).sqrt().recip();
1338                for d in 0..key_dim {
1339                    query[base + d] *= q_inv;
1340                    key[base + d] *= k_inv;
1341                }
1342            }
1343        }
1344        Ok(())
1345    }
1346
1347    #[allow(clippy::too_many_arguments)]
1348    fn linear_attention_decode_prepare_f32(
1349        _ctx: &mut Self::Context,
1350        mixed_qkv_raw: &Self::Buffer,
1351        conv_weight: &Self::Buffer,
1352        conv_state: &Self::Buffer,
1353        a_raw: &Self::Buffer,
1354        b_raw: &Self::Buffer,
1355        a_log: &Self::Buffer,
1356        dt_bias: &Self::Buffer,
1357        query: &mut Self::Buffer,
1358        key: &mut Self::Buffer,
1359        value: &mut Self::Buffer,
1360        g: &mut Self::Buffer,
1361        beta: &mut Self::Buffer,
1362        next_conv_state: &mut Self::Buffer,
1363        key_heads: usize,
1364        value_heads: usize,
1365        key_dim: usize,
1366        value_dim: usize,
1367        conv_kernel: usize,
1368        apply_qk_l2norm: bool,
1369    ) -> Result<()> {
1370        validate_linear_attention_decode_prepare_shape(
1371            mixed_qkv_raw.len(),
1372            conv_weight.len(),
1373            conv_state.len(),
1374            a_raw.len(),
1375            b_raw.len(),
1376            a_log.len(),
1377            dt_bias.len(),
1378            query.len(),
1379            key.len(),
1380            value.len(),
1381            g.len(),
1382            beta.len(),
1383            next_conv_state.len(),
1384            key_heads,
1385            value_heads,
1386            key_dim,
1387            value_dim,
1388            conv_kernel,
1389        )?;
1390
1391        let qk_total = key_heads * key_dim;
1392        let value_total = value_heads * value_dim;
1393        let conv_channels = 2 * qk_total + value_total;
1394        let state_len = conv_kernel - 1;
1395        for channel in 0..conv_channels {
1396            let state_base = channel * state_len;
1397            let mut acc = 0.0;
1398            for kernel_idx in 0..conv_kernel {
1399                let x = if kernel_idx < state_len {
1400                    conv_state[state_base + kernel_idx]
1401                } else {
1402                    mixed_qkv_raw[channel]
1403                };
1404                acc += x * conv_weight[channel * conv_kernel + kernel_idx];
1405            }
1406
1407            if state_len > 0 {
1408                for pos in 0..state_len {
1409                    next_conv_state[state_base + pos] = if pos + 1 < state_len {
1410                        conv_state[state_base + pos + 1]
1411                    } else {
1412                        mixed_qkv_raw[channel]
1413                    };
1414                }
1415            }
1416
1417            let conv = silu(acc);
1418            if channel < qk_total {
1419                query[channel] = conv;
1420            } else if channel < 2 * qk_total {
1421                key[channel - qk_total] = conv;
1422            } else {
1423                value[channel - 2 * qk_total] = conv;
1424            }
1425        }
1426
1427        for value_head in 0..value_heads {
1428            g[value_head] =
1429                -a_log[value_head].exp() * softplus(a_raw[value_head] + dt_bias[value_head]);
1430            beta[value_head] = sigmoid(b_raw[value_head]);
1431        }
1432
1433        if apply_qk_l2norm {
1434            for row in 0..key_heads {
1435                let base = row * key_dim;
1436                let mut q_sum = 0.0;
1437                let mut k_sum = 0.0;
1438                for d in 0..key_dim {
1439                    q_sum += query[base + d] * query[base + d];
1440                    k_sum += key[base + d] * key[base + d];
1441                }
1442                let q_inv = (q_sum + 1e-6).sqrt().recip();
1443                let k_inv = (k_sum + 1e-6).sqrt().recip();
1444                for d in 0..key_dim {
1445                    query[base + d] *= q_inv;
1446                    key[base + d] *= k_inv;
1447                }
1448            }
1449        }
1450        Ok(())
1451    }
1452
1453    #[allow(clippy::too_many_arguments)]
1454    fn linear_attention_decode_prepare_batch_f32(
1455        _ctx: &mut Self::Context,
1456        mixed_qkv_raw: &Self::Buffer,
1457        conv_weight: &Self::Buffer,
1458        conv_states: &Self::Buffer,
1459        a_raw: &Self::Buffer,
1460        b_raw: &Self::Buffer,
1461        a_log: &Self::Buffer,
1462        dt_bias: &Self::Buffer,
1463        query: &mut Self::Buffer,
1464        key: &mut Self::Buffer,
1465        value: &mut Self::Buffer,
1466        g: &mut Self::Buffer,
1467        beta: &mut Self::Buffer,
1468        next_conv_states: &mut Self::Buffer,
1469        batch: usize,
1470        key_heads: usize,
1471        value_heads: usize,
1472        key_dim: usize,
1473        value_dim: usize,
1474        conv_kernel: usize,
1475        apply_qk_l2norm: bool,
1476    ) -> Result<()> {
1477        if batch == 0 {
1478            return Err(FerrumError::model(
1479                "linear_attention_decode_prepare_batch batch must be positive",
1480            ));
1481        }
1482        let qk_total = key_heads * key_dim;
1483        let value_total = value_heads * value_dim;
1484        let conv_channels = 2 * qk_total + value_total;
1485        let state_len = conv_kernel.saturating_sub(1);
1486        let conv_state_len = conv_channels * state_len;
1487        validate_linear_attention_prepare_shape(
1488            mixed_qkv_raw.len(),
1489            conv_weight.len(),
1490            a_raw.len(),
1491            b_raw.len(),
1492            a_log.len(),
1493            dt_bias.len(),
1494            query.len(),
1495            key.len(),
1496            value.len(),
1497            g.len(),
1498            beta.len(),
1499            batch,
1500            key_heads,
1501            value_heads,
1502            key_dim,
1503            value_dim,
1504            conv_kernel,
1505        )?;
1506        for (label, actual, expected) in [
1507            ("conv_states", conv_states.len(), batch * conv_state_len),
1508            (
1509                "next_conv_states",
1510                next_conv_states.len(),
1511                batch * conv_state_len,
1512            ),
1513        ] {
1514            if actual < expected {
1515                return Err(FerrumError::model(format!(
1516                    "linear_attention_decode_prepare_batch {label} length {actual} < expected {expected}"
1517                )));
1518            }
1519        }
1520
1521        for row in 0..batch {
1522            let conv_row_base = row * conv_channels;
1523            let state_row_base = row * conv_state_len;
1524            for channel in 0..conv_channels {
1525                let state_base = state_row_base + channel * state_len;
1526                let mut acc = 0.0;
1527                for kernel_idx in 0..conv_kernel {
1528                    let x = if kernel_idx < state_len {
1529                        conv_states[state_base + kernel_idx]
1530                    } else {
1531                        mixed_qkv_raw[conv_row_base + channel]
1532                    };
1533                    acc += x * conv_weight[channel * conv_kernel + kernel_idx];
1534                }
1535
1536                if state_len > 0 {
1537                    for pos in 0..state_len {
1538                        next_conv_states[state_base + pos] = if pos + 1 < state_len {
1539                            conv_states[state_base + pos + 1]
1540                        } else {
1541                            mixed_qkv_raw[conv_row_base + channel]
1542                        };
1543                    }
1544                }
1545
1546                let conv = silu(acc);
1547                if channel < qk_total {
1548                    query[row * qk_total + channel] = conv;
1549                } else if channel < 2 * qk_total {
1550                    key[row * qk_total + (channel - qk_total)] = conv;
1551                } else {
1552                    value[row * value_total + (channel - 2 * qk_total)] = conv;
1553                }
1554            }
1555
1556            for value_head in 0..value_heads {
1557                let gate_idx = row * value_heads + value_head;
1558                g[gate_idx] =
1559                    -a_log[value_head].exp() * softplus(a_raw[gate_idx] + dt_bias[value_head]);
1560                beta[gate_idx] = sigmoid(b_raw[gate_idx]);
1561            }
1562        }
1563
1564        if apply_qk_l2norm {
1565            for row in 0..batch * key_heads {
1566                let base = row * key_dim;
1567                let mut q_sum = 0.0;
1568                let mut k_sum = 0.0;
1569                for d in 0..key_dim {
1570                    q_sum += query[base + d] * query[base + d];
1571                    k_sum += key[base + d] * key[base + d];
1572                }
1573                let q_inv = (q_sum + 1e-6).sqrt().recip();
1574                let k_inv = (k_sum + 1e-6).sqrt().recip();
1575                for d in 0..key_dim {
1576                    query[base + d] *= q_inv;
1577                    key[base + d] *= k_inv;
1578                }
1579            }
1580        }
1581        Ok(())
1582    }
1583
1584    #[allow(clippy::too_many_arguments)]
1585    fn gated_rms_norm_f32(
1586        _ctx: &mut Self::Context,
1587        core: &Self::Buffer,
1588        z: &Self::Buffer,
1589        weight: &Self::Buffer,
1590        out: &mut Self::Buffer,
1591        tokens: usize,
1592        heads: usize,
1593        dim: usize,
1594        eps: f32,
1595    ) -> Result<()> {
1596        validate_gated_rms_norm_shape(
1597            core.len(),
1598            z.len(),
1599            weight.len(),
1600            out.len(),
1601            tokens,
1602            heads,
1603            dim,
1604        )?;
1605
1606        for row in 0..tokens * heads {
1607            let base = row * dim;
1608            let mut sum_sq = 0.0;
1609            for d in 0..dim {
1610                let x = core[base + d];
1611                sum_sq += x * x;
1612            }
1613            let inv = (sum_sq / dim as f32 + eps).sqrt().recip();
1614            for d in 0..dim {
1615                out[base + d] = core[base + d] * inv * weight[d] * silu(z[base + d]);
1616            }
1617        }
1618        Ok(())
1619    }
1620
1621    fn copy_slice(
1622        _ctx: &mut Self::Context,
1623        src: &Self::Buffer,
1624        src_offset: usize,
1625        dst: &mut Self::Buffer,
1626        dst_offset: usize,
1627        len: usize,
1628    ) {
1629        dst[dst_offset..dst_offset + len].copy_from_slice(&src[src_offset..src_offset + len]);
1630    }
1631
1632    fn embedding_lookup(
1633        _ctx: &mut Self::Context,
1634        table: &Self::Buffer,
1635        ids: &[u32],
1636        out: &mut Self::Buffer,
1637        dim: usize,
1638    ) {
1639        for (i, &id) in ids.iter().enumerate() {
1640            let src = id as usize * dim;
1641            out[i * dim..(i + 1) * dim].copy_from_slice(&table[src..src + dim]);
1642        }
1643    }
1644
1645    fn split_qkv(
1646        _ctx: &mut Self::Context,
1647        qkv: &Self::Buffer,
1648        q: &mut Self::Buffer,
1649        k: &mut Self::Buffer,
1650        v: &mut Self::Buffer,
1651        tokens: usize,
1652        q_dim: usize,
1653        kv_dim: usize,
1654    ) {
1655        let qkv_dim = q_dim + 2 * kv_dim;
1656        for t in 0..tokens {
1657            let base = t * qkv_dim;
1658            q[t * q_dim..(t + 1) * q_dim].copy_from_slice(&qkv[base..base + q_dim]);
1659            k[t * kv_dim..(t + 1) * kv_dim]
1660                .copy_from_slice(&qkv[base + q_dim..base + q_dim + kv_dim]);
1661            v[t * kv_dim..(t + 1) * kv_dim]
1662                .copy_from_slice(&qkv[base + q_dim + kv_dim..base + qkv_dim]);
1663        }
1664    }
1665
1666    fn fused_silu_mul_split(
1667        _ctx: &mut Self::Context,
1668        gate_up: &Self::Buffer,
1669        out: &mut Self::Buffer,
1670        tokens: usize,
1671        im: usize,
1672    ) {
1673        for t in 0..tokens {
1674            for i in 0..im {
1675                let g = gate_up[t * 2 * im + i];
1676                let u = gate_up[t * 2 * im + im + i];
1677                out[t * im + i] = (g / (1.0 + (-g).exp())) * u;
1678            }
1679        }
1680    }
1681
1682    fn fused_gelu_tanh_mul_split(
1683        _ctx: &mut Self::Context,
1684        gate_up: &Self::Buffer,
1685        out: &mut Self::Buffer,
1686        tokens: usize,
1687        im: usize,
1688    ) {
1689        const SQRT_2_OVER_PI: f32 = 0.797_884_56;
1690        for t in 0..tokens {
1691            for i in 0..im {
1692                let g = gate_up[t * 2 * im + i];
1693                let u = gate_up[t * 2 * im + im + i];
1694                let inner = SQRT_2_OVER_PI * (g + 0.044715 * g * g * g);
1695                out[t * im + i] = 0.5 * g * (1.0 + inner.tanh()) * u;
1696            }
1697        }
1698    }
1699
1700    fn scale_inplace(_ctx: &mut Self::Context, buf: &mut Self::Buffer, scale: f32, len: usize) {
1701        for x in buf[..len].iter_mut() {
1702            *x *= scale;
1703        }
1704    }
1705
1706    fn qk_norm_rope(
1707        _ctx: &mut Self::Context,
1708        input: &Self::Buffer,
1709        norm_w: &Self::Buffer,
1710        cos: &Self::Buffer,
1711        sin: &Self::Buffer,
1712        output: &mut Self::Buffer,
1713        tokens: usize,
1714        heads: usize,
1715        head_dim: usize,
1716        pos_offset: usize,
1717        eps: f32,
1718        mode: i32,
1719    ) {
1720        let half = head_dim / 2;
1721        let cos_len = cos.len();
1722        let sin_len = sin.len();
1723        debug_assert_eq!(cos_len, sin_len);
1724
1725        for t in 0..tokens {
1726            let pos = pos_offset + t;
1727            for h in 0..heads {
1728                // input row: [t, h, :]  stride = heads * head_dim
1729                let src_off = (t * heads + h) * head_dim;
1730                // output row: [h, t, :]  stride = tokens * head_dim
1731                let dst_off = (h * tokens + t) * head_dim;
1732
1733                // Mode 0: plain transpose.
1734                if mode == 0 {
1735                    for i in 0..head_dim {
1736                        output[dst_off + i] = input[src_off + i];
1737                    }
1738                    continue;
1739                }
1740
1741                // Optional RMS norm (mode 1 only).
1742                let scale = if mode == 1 {
1743                    let mut sum_sq = 0.0f32;
1744                    for i in 0..head_dim {
1745                        sum_sq += input[src_off + i] * input[src_off + i];
1746                    }
1747                    1.0f32 / (sum_sq / head_dim as f32 + eps).sqrt()
1748                } else {
1749                    1.0
1750                };
1751
1752                if mode == 3 {
1753                    // GGUF LLaMA / llama.cpp interleaved RoPE layout.
1754                    for i in 0..half {
1755                        let j = 2 * i;
1756                        let x0 = input[src_off + j];
1757                        let x1 = input[src_off + j + 1];
1758                        let c = cos[pos * half + i];
1759                        let s = sin[pos * half + i];
1760                        output[dst_off + j] = x0 * c - x1 * s;
1761                        output[dst_off + j + 1] = x1 * c + x0 * s;
1762                    }
1763                } else {
1764                    // Apply (norm?) + half-split RoPE to head-major output.
1765                    for i in 0..half {
1766                        let (x0_raw, x1_raw) = (input[src_off + i], input[src_off + i + half]);
1767                        let (x0, x1) = if mode == 1 {
1768                            (
1769                                x0_raw * scale * norm_w[i],
1770                                x1_raw * scale * norm_w[i + half],
1771                            )
1772                        } else {
1773                            (x0_raw, x1_raw)
1774                        };
1775                        let c = cos[pos * half + i];
1776                        let s = sin[pos * half + i];
1777                        output[dst_off + i] = x0 * c - x1 * s;
1778                        output[dst_off + i + half] = x1 * c + x0 * s;
1779                    }
1780                }
1781            }
1782        }
1783    }
1784
1785    fn qk_norm_rope_partial(
1786        _ctx: &mut Self::Context,
1787        input: &Self::Buffer,
1788        norm_w: &Self::Buffer,
1789        cos: &Self::Buffer,
1790        sin: &Self::Buffer,
1791        output: &mut Self::Buffer,
1792        tokens: usize,
1793        heads: usize,
1794        head_dim: usize,
1795        rope_dim: usize,
1796        input_stride: usize,
1797        input_offset: usize,
1798        input_head_stride: usize,
1799        pos_offset: usize,
1800        eps: f32,
1801        mode: i32,
1802    ) -> Result<()> {
1803        if tokens == 0 || heads == 0 || head_dim == 0 || rope_dim == 0 {
1804            return Err(FerrumError::model(format!(
1805                "qk_norm_rope_partial shape must be positive, got tokens={tokens} heads={heads} head_dim={head_dim} rope_dim={rope_dim}"
1806            )));
1807        }
1808        if rope_dim > head_dim || rope_dim % 2 != 0 {
1809            return Err(FerrumError::model(format!(
1810                "qk_norm_rope_partial rope_dim {rope_dim} must be even and <= head_dim {head_dim}"
1811            )));
1812        }
1813        if input_head_stride == 0 {
1814            return Err(FerrumError::model(
1815                "qk_norm_rope_partial input_head_stride must be positive",
1816            ));
1817        }
1818        let required_width = input_offset + (heads - 1) * input_head_stride + head_dim;
1819        if input_stride < required_width {
1820            return Err(FerrumError::model(format!(
1821                "qk_norm_rope_partial input_stride {input_stride} is too small for offset {input_offset}, heads {heads}, head_dim {head_dim}, input_head_stride {input_head_stride}"
1822            )));
1823        }
1824        let required_input = tokens * input_stride;
1825        let required_output = tokens * heads * head_dim;
1826        if input.len() < required_input || output.len() < required_output {
1827            return Err(FerrumError::model(format!(
1828                "qk_norm_rope_partial buffer too short: input {} need {required_input}, output {} need {required_output}",
1829                input.len(),
1830                output.len()
1831            )));
1832        }
1833        if mode != 0
1834            && (cos.len() < (pos_offset + tokens) * (rope_dim / 2)
1835                || sin.len() < (pos_offset + tokens) * (rope_dim / 2))
1836        {
1837            return Err(FerrumError::model(
1838                "qk_norm_rope_partial RoPE cache is too short",
1839            ));
1840        }
1841        if (mode == 1 || mode == 3) && norm_w.len() < head_dim {
1842            return Err(FerrumError::model(
1843                "qk_norm_rope_partial norm weight is too short",
1844            ));
1845        }
1846
1847        let rope_half = rope_dim / 2;
1848        for t in 0..tokens {
1849            let pos = pos_offset + t;
1850            for h in 0..heads {
1851                let src_off = t * input_stride + input_offset + h * input_head_stride;
1852                let dst_off = (h * tokens + t) * head_dim;
1853                let mut row = vec![0.0f32; head_dim];
1854
1855                let scale = if mode == 1 || mode == 3 {
1856                    let mut sum_sq = 0.0f32;
1857                    for i in 0..head_dim {
1858                        let value = input[src_off + i];
1859                        sum_sq += value * value;
1860                    }
1861                    (sum_sq / head_dim as f32 + eps).sqrt().recip()
1862                } else {
1863                    1.0
1864                };
1865                for i in 0..head_dim {
1866                    let mut value = input[src_off + i];
1867                    if mode == 1 || mode == 3 {
1868                        value *= scale * norm_w[i];
1869                    }
1870                    row[i] = value;
1871                }
1872
1873                if mode == 1 || mode == 2 {
1874                    for i in 0..rope_half {
1875                        let left = i;
1876                        let right = i + rope_half;
1877                        let x0 = row[left];
1878                        let x1 = row[right];
1879                        let c = cos[pos * rope_half + i];
1880                        let s = sin[pos * rope_half + i];
1881                        row[left] = x0 * c - x1 * s;
1882                        row[right] = x1 * c + x0 * s;
1883                    }
1884                } else if mode == 3 {
1885                    for i in 0..rope_half {
1886                        let left = 2 * i;
1887                        let right = left + 1;
1888                        let x0 = row[left];
1889                        let x1 = row[right];
1890                        let c = cos[pos * rope_half + i];
1891                        let s = sin[pos * rope_half + i];
1892                        row[left] = x0 * c - x1 * s;
1893                        row[right] = x1 * c + x0 * s;
1894                    }
1895                } else if mode != 0 {
1896                    return Err(FerrumError::model(format!(
1897                        "qk_norm_rope_partial unsupported mode {mode}"
1898                    )));
1899                }
1900
1901                output[dst_off..dst_off + head_dim].copy_from_slice(&row);
1902            }
1903        }
1904        Ok(())
1905    }
1906
1907    fn qwen35_apply_attention_gate(
1908        _ctx: &mut Self::Context,
1909        context: &mut Self::Buffer,
1910        query_raw: &Self::Buffer,
1911        tokens: usize,
1912        q_total: usize,
1913        q_proj_total: usize,
1914        head_dim: usize,
1915    ) -> Result<()> {
1916        if head_dim == 0 || q_total % head_dim != 0 {
1917            return Err(FerrumError::model(format!(
1918                "qwen35 attention gate requires q_total {q_total} to be divisible by head_dim {head_dim}"
1919            )));
1920        }
1921        let heads = q_total / head_dim;
1922        if q_proj_total < heads * 2 * head_dim {
1923            return Err(FerrumError::model(format!(
1924                "qwen35 attention gate requires q_proj_total >= heads*2*head_dim, got q_total={q_total} q_proj_total={q_proj_total} head_dim={head_dim}"
1925            )));
1926        }
1927        if context.len() < tokens * q_total || query_raw.len() < tokens * q_proj_total {
1928            return Err(FerrumError::model(
1929                "qwen35 attention gate buffer is too short",
1930            ));
1931        }
1932        for token in 0..tokens {
1933            let ctx_base = token * q_total;
1934            for dim in 0..q_total {
1935                let head = dim / head_dim;
1936                let head_dim_offset = dim % head_dim;
1937                let gate_idx =
1938                    token * q_proj_total + head * (2 * head_dim) + head_dim + head_dim_offset;
1939                context[ctx_base + dim] *= sigmoid(query_raw[gate_idx]);
1940            }
1941        }
1942        Ok(())
1943    }
1944
1945    fn qwen35_apply_token_gate(
1946        _ctx: &mut Self::Context,
1947        values: &mut Self::Buffer,
1948        gate: &Self::Buffer,
1949        tokens: usize,
1950        hidden_size: usize,
1951    ) -> Result<()> {
1952        if values.len() < tokens * hidden_size || gate.len() < tokens {
1953            return Err(FerrumError::model("qwen35 token gate buffer is too short"));
1954        }
1955        for token in 0..tokens {
1956            let scale = sigmoid(gate[token]);
1957            let base = token * hidden_size;
1958            for dim in 0..hidden_size {
1959                values[base + dim] *= scale;
1960            }
1961        }
1962        Ok(())
1963    }
1964
1965    fn qwen35_apply_token_gate_and_add_inplace(
1966        _ctx: &mut Self::Context,
1967        dst: &mut Self::Buffer,
1968        values: &mut Self::Buffer,
1969        gate: &Self::Buffer,
1970        tokens: usize,
1971        hidden_size: usize,
1972    ) -> Result<()> {
1973        if dst.len() < tokens * hidden_size
1974            || values.len() < tokens * hidden_size
1975            || gate.len() < tokens
1976        {
1977            return Err(FerrumError::model(
1978                "qwen35 token gate merge buffer is too short",
1979            ));
1980        }
1981        for token in 0..tokens {
1982            let scale = sigmoid(gate[token]);
1983            let base = token * hidden_size;
1984            for dim in 0..hidden_size {
1985                let idx = base + dim;
1986                values[idx] *= scale;
1987                dst[idx] += values[idx];
1988            }
1989        }
1990        Ok(())
1991    }
1992
1993    fn kv_cache_append_head_major(
1994        _ctx: &mut Self::Context,
1995        cache_k: &mut Self::Buffer,
1996        cache_v: &mut Self::Buffer,
1997        cache_len: usize,
1998        cache_capacity: usize,
1999        new_k_head_major: &Self::Buffer,
2000        new_v_head_major: &Self::Buffer,
2001        new_tokens: usize,
2002        nkv: usize,
2003        hd: usize,
2004    ) {
2005        debug_assert!(cache_len + new_tokens <= cache_capacity);
2006        debug_assert_eq!(cache_k.len(), nkv * cache_capacity * hd);
2007        debug_assert_eq!(cache_v.len(), nkv * cache_capacity * hd);
2008        // The source buffers may be sized for `max_tokens` (the prefill-
2009        // sized scratch) while only the first `nkv * new_tokens * hd`
2010        // entries are valid for this call. Allow >= so reusing scratch
2011        // across prefill and decode doesn't trip the assert.
2012        debug_assert!(new_k_head_major.len() >= nkv * new_tokens * hd);
2013        debug_assert!(new_v_head_major.len() >= nkv * new_tokens * hd);
2014
2015        for h in 0..nkv {
2016            let dst_base = h * cache_capacity * hd + cache_len * hd;
2017            let src_base = h * new_tokens * hd;
2018            cache_k[dst_base..dst_base + new_tokens * hd]
2019                .copy_from_slice(&new_k_head_major[src_base..src_base + new_tokens * hd]);
2020            cache_v[dst_base..dst_base + new_tokens * hd]
2021                .copy_from_slice(&new_v_head_major[src_base..src_base + new_tokens * hd]);
2022        }
2023    }
2024
2025    fn transpose_head_to_token(
2026        _ctx: &mut Self::Context,
2027        src: &Self::Buffer,
2028        dst: &mut Self::Buffer,
2029        tokens: usize,
2030        heads: usize,
2031        dim: usize,
2032    ) {
2033        for h in 0..heads {
2034            for t in 0..tokens {
2035                let s = (h * tokens + t) * dim;
2036                let d = (t * heads + h) * dim;
2037                dst[d..d + dim].copy_from_slice(&src[s..s + dim]);
2038            }
2039        }
2040    }
2041
2042    fn add_inplace(
2043        _ctx: &mut Self::Context,
2044        residual: &mut Self::Buffer,
2045        x: &Self::Buffer,
2046        len: usize,
2047    ) {
2048        for i in 0..len {
2049            residual[i] += x[i];
2050        }
2051    }
2052
2053    fn scaled_add_inplace(
2054        _ctx: &mut Self::Context,
2055        dst: &mut Self::Buffer,
2056        src: &Self::Buffer,
2057        scale: f32,
2058        len: usize,
2059    ) {
2060        for i in 0..len {
2061            dst[i] += scale * src[i];
2062        }
2063    }
2064
2065    fn add_bias(
2066        _ctx: &mut Self::Context,
2067        data: &mut Self::Buffer,
2068        bias: &Self::Buffer,
2069        rows: usize,
2070        cols: usize,
2071    ) {
2072        debug_assert_eq!(bias.len(), cols);
2073        for r in 0..rows {
2074            let off = r * cols;
2075            for c in 0..cols {
2076                data[off + c] += bias[c];
2077            }
2078        }
2079    }
2080
2081    fn layer_norm(
2082        _ctx: &mut Self::Context,
2083        x: &Self::Buffer,
2084        gamma: &Self::Buffer,
2085        beta: &Self::Buffer,
2086        eps: f32,
2087        out: &mut Self::Buffer,
2088        tokens: usize,
2089        dim: usize,
2090    ) {
2091        debug_assert_eq!(gamma.len(), dim);
2092        debug_assert_eq!(beta.len(), dim);
2093        for t in 0..tokens {
2094            let off = t * dim;
2095            // Compute mean + variance over `dim` in f64 for stability.
2096            let mut mean = 0.0f64;
2097            for i in 0..dim {
2098                mean += x[off + i] as f64;
2099            }
2100            mean /= dim as f64;
2101            let mut var = 0.0f64;
2102            for i in 0..dim {
2103                let d = x[off + i] as f64 - mean;
2104                var += d * d;
2105            }
2106            var /= dim as f64;
2107            let inv = 1.0f32 / ((var as f32) + eps).sqrt();
2108            let mean_f32 = mean as f32;
2109            for i in 0..dim {
2110                out[off + i] = (x[off + i] - mean_f32) * inv * gamma[i] + beta[i];
2111            }
2112        }
2113    }
2114
2115    fn gelu(_ctx: &mut Self::Context, x: &Self::Buffer, out: &mut Self::Buffer, len: usize) {
2116        // Exact GELU: 0.5 * x * (1 + erf(x / sqrt(2))).
2117        // Uses f64 for erf accuracy (matches torch.nn.functional.gelu default).
2118        for i in 0..len {
2119            let xi = x[i];
2120            out[i] = 0.5 * xi * (1.0 + libm_erf(xi / std::f32::consts::SQRT_2));
2121        }
2122    }
2123
2124    fn alloc(len: usize) -> Self::Buffer {
2125        vec![0.0f32; len]
2126    }
2127    fn to_vec(buf: &Self::Buffer, len: usize) -> Vec<f32> {
2128        buf[..len].to_vec()
2129    }
2130    fn from_slice(data: &[f32]) -> Self::Buffer {
2131        data.to_vec()
2132    }
2133}
2134
2135// ── Helpers ──────────────────────────────────────────────────────────────
2136
2137fn dot_product(a: &[f32], b: &[f32]) -> f32 {
2138    #[cfg(target_os = "macos")]
2139    {
2140        let mut result = 0.0f32;
2141        unsafe {
2142            vDSP_dotpr(a.as_ptr(), 1, b.as_ptr(), 1, &mut result, a.len() as u64);
2143        }
2144        result
2145    }
2146    #[cfg(not(target_os = "macos"))]
2147    {
2148        a.iter().zip(b).map(|(x, y)| x * y).sum()
2149    }
2150}
2151
2152#[allow(dead_code)]
2153fn apply_rope_impl(
2154    data: &mut [f32],
2155    tokens: usize,
2156    heads: usize,
2157    head_dim: usize,
2158    half: usize,
2159    cos: &[f32],
2160    sin: &[f32],
2161    positions: &[u32],
2162) {
2163    for t in 0..tokens {
2164        let pos = positions[t] as usize;
2165        for h in 0..heads {
2166            let base = t * heads * head_dim + h * head_dim;
2167            for i in 0..half {
2168                let c = cos[pos * half + i];
2169                let s = sin[pos * half + i];
2170                let x0 = data[base + i];
2171                let x1 = data[base + half + i];
2172                data[base + i] = x0 * c - x1 * s;
2173                data[base + half + i] = x1 * c + x0 * s;
2174            }
2175        }
2176    }
2177}
2178
2179fn cpu_attention(
2180    q: &[f32],
2181    k: &[f32],
2182    v: &[f32],
2183    out: &mut [f32],
2184    batch: usize,
2185    q_len: usize,
2186    kv_len: usize,
2187    causal: bool,
2188    pos_offset: usize,
2189    cfg: &AttnConfig,
2190) {
2191    let nh = cfg.num_heads;
2192    let nkv = cfg.num_kv_heads;
2193    let d = cfg.head_dim;
2194    let n_rep = nh / nkv;
2195    let scale = cfg.scale;
2196    // Per-head KV stride: 0 (the default) means contiguous (legacy
2197    // `kv_cache_append` path reallocates each layer). A non-zero value means
2198    // the cache is pre-allocated to `kv_seq_stride` rows per head but only
2199    // the first `kv_len` are valid — we skip the rest via `attend_len`.
2200    let kv_stride = if cfg.kv_seq_stride > 0 {
2201        cfg.kv_seq_stride
2202    } else {
2203        kv_len
2204    };
2205
2206    for b in 0..batch {
2207        for h in 0..nh {
2208            let kv_h = h / n_rep;
2209            let q_off = (b * nh + h) * q_len * d;
2210            let k_off = (b * nkv + kv_h) * kv_stride * d;
2211            let v_off = (b * nkv + kv_h) * kv_stride * d;
2212            let o_off = (b * nh + h) * q_len * d;
2213
2214            for qi in 0..q_len {
2215                let attend_end = if causal {
2216                    (pos_offset + qi + 1).min(kv_len)
2217                } else {
2218                    kv_len
2219                };
2220                let attend_start = if causal && cfg.sliding_window > 0 {
2221                    attend_end.saturating_sub(cfg.sliding_window)
2222                } else {
2223                    0
2224                };
2225                let mut max_score = f32::NEG_INFINITY;
2226                let mut sum_exp = 0.0f32;
2227                let mut acc = vec![0.0f32; d];
2228
2229                for ki in attend_start..attend_end {
2230                    let mut dot = 0.0f32;
2231                    for di in 0..d {
2232                        dot += q[q_off + qi * d + di] * k[k_off + ki * d + di];
2233                    }
2234                    let score = dot * scale;
2235                    if score > max_score {
2236                        let correction = (max_score - score).exp();
2237                        for di in 0..d {
2238                            acc[di] *= correction;
2239                        }
2240                        sum_exp *= correction;
2241                        max_score = score;
2242                    }
2243                    let w = (score - max_score).exp();
2244                    sum_exp += w;
2245                    for di in 0..d {
2246                        acc[di] += w * v[v_off + ki * d + di];
2247                    }
2248                }
2249
2250                if sum_exp > 0.0 {
2251                    let inv = 1.0 / sum_exp;
2252                    for di in 0..d {
2253                        out[o_off + qi * d + di] = acc[di] * inv;
2254                    }
2255                }
2256            }
2257        }
2258    }
2259}
2260
2261/// Minimal error-function approximation (Abramowitz & Stegun 7.1.26),
2262/// max error ~1.5e-7 which is comfortably below f32 round-off noise.
2263fn libm_erf(x: f32) -> f32 {
2264    let sign = if x < 0.0 { -1.0 } else { 1.0 };
2265    let x = x.abs();
2266    let t = 1.0 / (1.0 + 0.3275911 * x);
2267    let y = 1.0
2268        - (((((1.061_405_4 * t - 1.453_152_1) * t) + 1.421_413_8) * t - 0.284_496_72) * t
2269            + 0.254_829_6)
2270            * t
2271            * (-x * x).exp();
2272    sign * y
2273}
2274
2275// CPU has no graph-capture analogue; inherit BackendGraph defaults.
2276impl crate::backend::BackendGraph for CpuBackend {}
2277
2278// CPU has no multi-rank collectives; inherit BackendCollective defaults.
2279impl crate::backend::BackendCollective for CpuBackend {}
2280
2281/// Dequant raw GPTQ tensors → row-major `[n, k]` f32. Shared between
2282/// the per-tensor `load_gptq` and the MoE `load_gptq_stacked` impls.
2283fn cpu_dequant_gptq(
2284    qweight: &[i32],
2285    scales: &[f32],
2286    qzeros: &[i32],
2287    bits: u32,
2288    group_size: usize,
2289    k: usize,
2290    n: usize,
2291) -> Result<Vec<f32>> {
2292    if bits != 4 {
2293        return Err(FerrumError::unsupported(format!(
2294            "CPU GPTQ: only bits=4 supported (got {bits})"
2295        )));
2296    }
2297    let mut w = vec![0.0f32; n * k];
2298    let packed_rows = k / 8;
2299    for pr in 0..packed_rows {
2300        for col in 0..n {
2301            let packed = qweight[pr * n + col] as u32;
2302            for bi in 0..8 {
2303                let ki = pr * 8 + bi;
2304                let q = ((packed >> (bi * 4)) & 0xF) as i32;
2305                let grp = ki / group_size;
2306                let scale = scales[grp * n + col];
2307                let z_packed = qzeros[grp * (n / 8) + (col / 8)] as u32;
2308                let zero = (((z_packed >> ((col % 8) * 4)) & 0xF) as i32) + 1;
2309                let val = (q - zero) as f32 * scale;
2310                w[col * k + ki] = val;
2311            }
2312        }
2313    }
2314    Ok(w)
2315}
2316
2317impl crate::backend::BackendQuantMarlin for CpuBackend {
2318    fn load_gptq(
2319        qweight: &[i32],
2320        scales: &[f32],
2321        qzeros: &[i32],
2322        _g_idx: Option<&[i32]>,
2323        bias_host: Option<&[f32]>,
2324        bits: u32,
2325        group_size: usize,
2326        k: usize,
2327        n: usize,
2328    ) -> Result<Box<dyn crate::Linear<Self> + Send + Sync>> {
2329        let w = cpu_dequant_gptq(qweight, scales, qzeros, bits, group_size, k, n)?;
2330        // Phase 3e/2: dequantized weights become a CpuGptqLinear that
2331        // owns the (out_features, in_features) f32 matrix and runs
2332        // through the existing Self::gemm CPU path.
2333        Ok(Box::new(crate::quant_linear::cpu_dequant::CpuGptqLinear {
2334            weight_f32: w,
2335            bias: bias_host.map(|b| b.to_vec()),
2336            in_features: k,
2337            out_features: n,
2338        }))
2339    }
2340    fn load_gptq_stacked(
2341        qweights: &[&[i32]],
2342        scales: &[&[f32]],
2343        qzeros: &[&[i32]],
2344        _g_idx: Option<&[i32]>,
2345        bits: u32,
2346        group_size: usize,
2347        k: usize,
2348        n_per_expert: usize,
2349    ) -> Result<std::sync::Arc<dyn crate::MarlinExpertStack<Self>>> {
2350        // Phase 3e/2 addition: dequant each expert independently, concat
2351        // along N (rows in [n, k] layout). Used by MoE parity tests.
2352        let num_experts = qweights.len();
2353        if scales.len() != num_experts || qzeros.len() != num_experts {
2354            return Err(FerrumError::model(format!(
2355                "load_gptq_stacked: input slice lengths disagree (qw {num_experts}, sc {}, qz {})",
2356                scales.len(),
2357                qzeros.len()
2358            )));
2359        }
2360        let total_n = num_experts * n_per_expert;
2361        let mut all_w = Vec::with_capacity(total_n * k);
2362        for ((qw_e, sc_e), qz_e) in qweights.iter().zip(scales.iter()).zip(qzeros.iter()) {
2363            let w_e = cpu_dequant_gptq(qw_e, sc_e, qz_e, bits, group_size, k, n_per_expert)?;
2364            all_w.extend_from_slice(&w_e);
2365        }
2366        let store = std::sync::Arc::new(CpuGptqStore {
2367            weight_f32: all_w,
2368            k,
2369            n: total_n,
2370        });
2371        Ok(std::sync::Arc::new(
2372            crate::quant_linear::cpu_marlin_stack::CpuMarlinExpertStack::new(
2373                store,
2374                num_experts,
2375                n_per_expert,
2376                k,
2377            ),
2378        ))
2379    }
2380    // Phase C step 4b: make_stacked_expert_linear inlined into
2381    // CpuMarlinExpertStack::make_expert_linear.
2382    // Phase C step 4e: make_marlin_expert_stack subsumed by load_gptq_stacked.
2383    // gemm_gptq_with_offset_strided body moved to free function
2384    // cpu_gemm_gptq_with_offset_strided below — called by
2385    // CpuMarlinExpertStack::gemm_phase_batched.
2386}
2387
2388fn cpu_read_u32_buffer(buf: &[f32], n: usize, label: &str) -> Result<Vec<u32>> {
2389    let bytes = n
2390        .checked_mul(std::mem::size_of::<u32>())
2391        .ok_or_else(|| FerrumError::model(format!("{label}: byte length overflow")))?;
2392    if bytes > buf.len() * std::mem::size_of::<f32>() {
2393        return Err(FerrumError::model(format!(
2394            "{label}: buffer byte length {} < expected {bytes}",
2395            buf.len() * std::mem::size_of::<f32>()
2396        )));
2397    }
2398    let mut out = vec![0u32; n];
2399    unsafe {
2400        std::ptr::copy_nonoverlapping(
2401            buf.as_ptr() as *const u8,
2402            out.as_mut_ptr() as *mut u8,
2403            bytes,
2404        );
2405    }
2406    Ok(out)
2407}
2408
2409/// Free-function form of the deleted
2410/// `BackendQuantMarlin::gemm_gptq_with_offset_strided` (Phase C step 4e).
2411/// Single caller: `CpuMarlinExpertStack::gemm_phase_batched`.
2412#[allow(clippy::too_many_arguments)]
2413pub(crate) fn cpu_gemm_gptq_with_offset_strided(
2414    _ctx: &mut <CpuBackend as Backend>::Context,
2415    input: &<CpuBackend as Backend>::Buffer,
2416    in_row_offset: usize,
2417    weight: &CpuGptqStore,
2418    expert_offset: usize,
2419    expert_n: usize,
2420    output: &mut <CpuBackend as Backend>::Buffer,
2421    out_row_offset: usize,
2422    m: usize,
2423    k: usize,
2424) -> Result<()> {
2425    if expert_offset + expert_n > weight.n {
2426        return Err(FerrumError::model(format!(
2427            "cpu_gemm_gptq_with_offset_strided OOB: offset {expert_offset} + n {expert_n} > stacked_n {}",
2428            weight.n
2429        )));
2430    }
2431    if k != weight.k {
2432        return Err(FerrumError::model(format!(
2433            "cpu_gemm_gptq_with_offset_strided k mismatch: arg {k} vs weight.k {}",
2434            weight.k
2435        )));
2436    }
2437    let in_start = in_row_offset * k;
2438    let in_end = (in_row_offset + m) * k;
2439    let out_start = out_row_offset * expert_n;
2440    let out_end = (out_row_offset + m) * expert_n;
2441    let row_start = expert_offset * k;
2442    let row_end = (expert_offset + expert_n) * k;
2443    let weight_slice = weight.weight_f32[row_start..row_end].to_vec();
2444    let in_slice = input[in_start..in_end].to_vec();
2445    let mut out_slice = vec![0.0f32; m * expert_n];
2446    let mut ctx_local = ();
2447    CpuBackend::gemm(
2448        &mut ctx_local,
2449        &in_slice,
2450        &weight_slice,
2451        &mut out_slice,
2452        m,
2453        expert_n,
2454        k,
2455    );
2456    output[out_start..out_end].copy_from_slice(&out_slice);
2457    Ok(())
2458}
2459
2460impl crate::backend::BackendQuantGguf for CpuBackend {
2461    fn load_quant(
2462        kind: super::GgufQuantType,
2463        bytes: &[u8],
2464        n_rows: usize,
2465        n_cols: usize,
2466    ) -> Result<Box<dyn crate::Linear<Self> + Send + Sync>> {
2467        use super::GgufQuantType;
2468        let store = match kind {
2469            GgufQuantType::Q4K => {
2470                let total_elems = n_rows * n_cols;
2471                if total_elems % Q4_K_QK != 0 {
2472                    return Err(FerrumError::model(format!(
2473                        "load_quant Q4K: elements {total_elems} not a multiple of {Q4_K_QK}"
2474                    )));
2475                }
2476                let n_blocks = total_elems / Q4_K_QK;
2477                let expected = n_blocks * Q4_K_BLOCK_BYTES;
2478                if bytes.len() != expected {
2479                    return Err(FerrumError::model(format!(
2480                        "load_quant Q4K: bytes {} != expected {} ({n_blocks} × {Q4_K_BLOCK_BYTES})",
2481                        bytes.len(),
2482                        expected
2483                    )));
2484                }
2485                CpuQuantStore::Q4K {
2486                    weights: dequant_q4_k_cpu(bytes, n_blocks),
2487                    n_rows,
2488                    n_cols,
2489                }
2490            }
2491            other => {
2492                return Err(FerrumError::unsupported(format!(
2493                    "CPU load_quant: {other:?} not yet implemented"
2494                )));
2495            }
2496        };
2497        // Phase 3e/3: dispatch via CpuGgufLinear::forward instead of
2498        // a trait method.
2499        Ok(Box::new(crate::quant_linear::cpu_gguf::CpuGgufLinear {
2500            store,
2501            in_features: n_cols,
2502            out_features: n_rows,
2503        }))
2504    }
2505}
2506
2507// CPU has no paged-KV path; inherit unsupported defaults.
2508impl crate::backend::BackendPagedKv for CpuBackend {}
2509
2510// CPU has no native MoE dispatch; inherit unsupported defaults.
2511impl crate::backend::BackendMoeFused for CpuBackend {}
2512
2513// CPU: existing KV cache path treats fp16 buffer as f32 internally; mark as KvFp16 for compatibility.
2514impl crate::backend::BackendKvDtype<crate::backend::KvFp16> for CpuBackend {
2515    type KvBuffer = <Self as crate::backend::Backend>::Buffer;
2516    type KvScales = ();
2517}