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