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