1use half::f16;
5
6use super::{AttnConfig, Backend};
7use ferrum_types::{FerrumError, Result};
8
9pub mod vnext_ops;
10pub mod vnext_runtime;
11
12const 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; fn 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
32fn 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
316pub struct CpuGptqStore {
319 pub weight_f32: Vec<f32>, pub k: usize,
321 pub n: usize,
322}
323
324pub enum CpuQuantStore {
332 Q4K {
333 weights: Vec<f32>, n_rows: usize,
335 n_cols: usize,
336 },
337}
338
339impl Backend for CpuBackend {
340 type Buffer = Vec<f32>;
341 type Context = ();
342 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 fn alloc_typed(dtype: crate::backend::Dtype, n: usize) -> Self::Buffer {
373 let bytes = n * dtype.bytes_per_elem();
376 let f32_len = bytes.div_ceil(4);
377 vec![0.0f32; f32_len]
378 }
379
380 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 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 let src_off = (t * heads + h) * head_dim;
1733 let dst_off = (h * tokens + t) * head_dim;
1735
1736 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 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 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 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 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 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 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
2138fn 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 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
2264fn 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
2278impl crate::backend::BackendGraph for CpuBackend {}
2280
2281impl crate::backend::BackendCollective for CpuBackend {}
2283
2284fn 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 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 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 }
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#[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 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
2510impl crate::backend::BackendPagedKv for CpuBackend {}
2512
2513impl crate::backend::BackendMoeFused for CpuBackend {}
2515
2516impl crate::backend::BackendKvDtype<crate::backend::KvFp16> for CpuBackend {
2518 type KvBuffer = <Self as crate::backend::Backend>::Buffer;
2519 type KvScales = ();
2520}