1use crate::Engine;
6use crate::cache::Cache;
7use cudarc::driver::CudaSlice;
8use memra_gguf::config::ModelConfig;
9
10pub struct PrimeSlabs {
13 pub t_cap: usize,
14 pub h: CudaSlice<f32>,
15 pub x1: CudaSlice<f32>,
16 pub z: CudaSlice<f32>,
17 pub act: CudaSlice<f32>,
18 pub xa: CudaSlice<f32>,
19 pub xb: CudaSlice<f32>,
20 pub h16: CudaSlice<u8>,
21 pub z16: CudaSlice<u8>,
22 pub gate: CudaSlice<f32>, pub up: CudaSlice<f32>, pub ffn_out: CudaSlice<f32>, pub seg_glue: Vec<Option<cudarc::driver::CudaGraph>>,
32 pub mixed: CudaSlice<f32>,
36 pub seg_mid: Vec<Option<cudarc::driver::CudaGraph>>,
37 pub seg_t: usize,
38}
39
40unsafe impl Send for PrimeSlabs {}
43
44fn empty_cache_layers<T>(n: usize) -> Vec<Option<T>> {
45 std::iter::repeat_with(|| None).take(n).collect()
46}
47
48struct PrimeCacheStages<'a> {
53 parent: &'a mut Cache,
54 cut: usize,
55 stage0: Cache,
56 stage1: Cache,
57}
58
59impl<'a> PrimeCacheStages<'a> {
60 fn new(parent: &'a mut Cache, cut: usize) -> Self {
61 let n = parent.kv.len();
62 assert_eq!(parent.recur.len(), n, "cache layer vectors disagree");
63 assert!(cut <= n, "PP-2 cache cut {cut} exceeds {n} layers");
64 let mut kv0 = empty_cache_layers(n);
65 let mut kv1 = empty_cache_layers(n);
66 let mut recur0 = empty_cache_layers(n);
67 let mut recur1 = empty_cache_layers(n);
68 for i in 0..cut {
69 kv0[i] = parent.kv[i].take();
70 recur0[i] = parent.recur[i].take();
71 }
72 for i in cut..n {
73 kv1[i] = parent.kv[i].take();
74 recur1[i] = parent.recur[i].take();
75 }
76 let pos = parent.pos;
77 let max_ctx = parent.max_ctx;
78 Self {
79 parent,
80 cut,
81 stage0: Cache {
82 kv: kv0,
83 recur: recur0,
84 pos,
85 max_ctx,
86 last_logits_dev: None,
87 dflash_taps: None,
88 },
89 stage1: Cache {
90 kv: kv1,
91 recur: recur1,
92 pos,
93 max_ctx,
94 last_logits_dev: None,
95 dflash_taps: None,
96 },
97 }
98 }
99
100 fn parts(&mut self) -> (&mut Cache, &mut Cache) {
101 (&mut self.stage0, &mut self.stage1)
102 }
103}
104
105impl Drop for PrimeCacheStages<'_> {
106 fn drop(&mut self) {
107 let n = self.parent.kv.len();
108 for i in 0..n {
109 let source = if i < self.cut {
110 &mut self.stage0
111 } else {
112 &mut self.stage1
113 };
114 debug_assert!(self.parent.kv[i].is_none());
115 debug_assert!(self.parent.recur[i].is_none());
116 self.parent.kv[i] = source.kv[i].take();
117 self.parent.recur[i] = source.recur[i].take();
118 }
119 self.parent.pos = self.stage0.pos.min(self.stage1.pos);
120 }
121}
122
123pub(crate) struct AttnPre {
125 pub q: cudarc::driver::CudaSlice<f32>,
126 pub k: cudarc::driver::CudaSlice<f32>,
127 pub v: cudarc::driver::CudaSlice<f32>,
128 pub gate: Option<cudarc::driver::CudaSlice<f32>>,
129}
130
131pub(crate) struct GdnPrep {
133 pub hk: usize,
134 pub q_l2: cudarc::driver::CudaSlice<f32>,
135 pub k_l2: cudarc::driver::CudaSlice<f32>,
136 pub v_g: cudarc::driver::CudaSlice<f32>,
137 pub beta: cudarc::driver::CudaSlice<f32>,
138 pub g_log: cudarc::driver::CudaSlice<f32>,
139 pub kb16: Option<cudarc::driver::CudaSlice<u8>>,
140 pub qb16: Option<cudarc::driver::CudaSlice<u8>>,
141}
142
143pub(crate) struct VerifyStreamScratch {
145 pub pos_d: CudaSlice<i32>,
146 pub row_ctrs: Vec<CudaSlice<i32>>,
147}
148use crate::hybrid::{FullAttnLayer, HybridModel, LinearAttnLayer, Mixer, MoeWeights};
149
150struct MoeInputTraceWriter {
151 dir: std::path::PathBuf,
152 index: std::fs::File,
153 payloads: std::collections::HashMap<u16, (std::fs::File, u64)>,
154}
155
156static MOE_INPUT_TRACE_WRITER: std::sync::OnceLock<std::sync::Mutex<Option<MoeInputTraceWriter>>> =
157 std::sync::OnceLock::new();
158
159fn gdec_enabled() -> bool {
162 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
163 *E.get_or_init(|| {
164 std::env::var("MEMRA_MOE_GDEC")
165 .map(|v| v != "0")
166 .unwrap_or(true)
167 })
168}
169
170fn moe_slab_enabled() -> bool {
181 std::env::var("MEMRA_MOE_SLAB").as_deref() != Ok("0")
182}
183
184fn moe_grouped_enabled(_cfg: &ModelConfig, _prefill: bool) -> bool {
188 std::env::var("MEMRA_MOE_GROUPED")
189 .map(|value| value != "0")
190 .unwrap_or(false)
191}
192
193fn moe_prefetch_enabled() -> bool {
196 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
197 *E.get_or_init(|| {
198 std::env::var("MEMRA_MOE_PREFETCH").as_deref() == Ok("1")
199 || crate::spill_pread::worker_enabled()
200 })
201}
202
203fn moe_page_prefetch_window() -> usize {
208 static W: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
209 *W.get_or_init(|| {
210 page_prefetch_window_from_values(
211 std::env::var("MEMRA_MOE_PAGE_PREFETCH").as_deref() == Ok("1"),
212 std::env::var("MEMRA_MOE_PAGE_PREFETCH_WINDOW")
213 .ok()
214 .as_deref(),
215 )
216 })
217}
218
219fn page_prefetch_window_from_values(enabled: bool, raw_window: Option<&str>) -> usize {
220 if !enabled {
221 return 0;
222 }
223 raw_window.and_then(|value| value.parse().ok()).unwrap_or(1)
224}
225
226fn page_prefetch_positions(position: usize, len: usize, window: usize) -> std::ops::Range<usize> {
230 if window == 0 || position >= len {
231 return len..len;
232 }
233 let (start, count) = if position == 0 {
234 (1, window)
235 } else {
236 (position.saturating_add(window), 1)
237 };
238 let start = start.min(len);
239 start..start.saturating_add(count).min(len)
240}
241
242fn grouped_worker_prefetch_position(order_len: usize, current: Option<usize>) -> Option<usize> {
245 let position = current.map_or(0, |position| position.saturating_add(1));
246 (position < order_len).then_some(position)
247}
248
249fn worker_prefetch_window() -> usize {
254 static WINDOW: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
255 *WINDOW.get_or_init(|| {
256 let automatic = crate::spill_pread::configured_depth().saturating_sub(1) / 3;
257 std::env::var("MEMRA_SPILL_WORKER_EXPERT_WINDOW")
258 .ok()
259 .and_then(|value| value.parse::<usize>().ok())
260 .unwrap_or(automatic.max(1))
261 })
262}
263
264fn worker_prefetch_positions(position: usize, len: usize, window: usize) -> std::ops::Range<usize> {
268 if window == 0 || position >= len {
269 return len..len;
270 }
271 let (start, count) = if position == 0 {
272 (0, window)
273 } else {
274 (position.saturating_add(window).saturating_sub(1), 1)
275 };
276 let start = start.min(len);
277 start..start.saturating_add(count).min(len)
278}
279
280fn moe_dev_enabled() -> bool {
285 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
286 *E.get_or_init(|| {
287 std::env::var("MEMRA_MOE_DEV")
288 .map(|v| v != "0")
289 .unwrap_or(true)
290 && !matches!(std::env::var("MEMRA_FUSED_ROUTER").as_deref(), Ok("0"))
291 })
292}
293
294fn sigmoid_router_enabled() -> bool {
297 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
298 *E.get_or_init(|| {
299 std::env::var("MEMRA_SIG_ROUTER")
300 .map(|v| v != "0")
301 .unwrap_or(true)
302 })
303}
304
305fn moe_q8_enabled() -> bool {
310 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
311 *E.get_or_init(|| {
312 std::env::var("MEMRA_MOE_Q8")
313 .map(|v| v != "0")
314 .unwrap_or(true)
315 })
316}
317
318fn expert_dp4a_supported(qt: i32) -> bool {
321 qt == crate::QT_Q4_0
322 || qt == crate::QT_IQ3_S
323 || qt == crate::QT_IQ4_XS
324 || qt == crate::QT_Q3_K
325 || qt == crate::QT_Q4_K
326 || qt == crate::QT_Q6_K
327}
328
329fn q8_expert_supported(qt: i32) -> bool {
330 static KQ: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
336 let kq = *KQ.get_or_init(|| {
337 std::env::var("MEMRA_MOE_Q8_KQ")
338 .map(|v| v != "0")
339 .unwrap_or(true)
340 });
341 let nvfp4_q8 = std::env::var("MEMRA_MOE_Q8_NVFP4")
348 .map(|v| v != "0")
349 .unwrap_or(true);
350 qt == crate::QT_IQ3_S
351 || qt == crate::QT_IQ4_XS
352 || (nvfp4_q8 && qt == crate::QT_NVFP4)
353 || (kq && (qt == crate::QT_Q3_K || qt == crate::QT_Q4_K || qt == crate::QT_Q6_K))
354}
355
356fn q8_expert_dec_supported(qt: i32) -> bool {
359 qt == crate::QT_IQ3_S || qt == crate::QT_IQ4_XS || qt == crate::QT_Q4_0
360}
361
362fn f16g_proj_ok(qt: i32, in_f: usize) -> bool {
368 match qt {
369 crate::QT_Q4_0 => in_f % 32 == 0,
370 crate::QT_IQ4_XS | crate::QT_IQ3_S | crate::QT_Q3_K | crate::QT_Q4_K | crate::QT_Q6_K => {
371 in_f % 256 == 0
372 }
373 _ => false,
374 }
375}
376
377fn moe_prewarm_enabled() -> bool {
380 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
381 *E.get_or_init(|| {
382 std::env::var("MEMRA_MOE_PREWARM")
383 .map(|v| v != "0")
384 .unwrap_or(true)
385 })
386}
387
388fn cpu_expert_profile_admit_enabled() -> bool {
392 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
393 *E.get_or_init(|| std::env::var("MEMRA_CPU_EXPERT_FREEZE_PROFILE_ADMIT").as_deref() == Ok("1"))
394}
395
396pub const PRIME_MIN_T: usize = 16;
400const PRIME_PIPE_MICROBATCHES: usize = 8;
401const PRIME_PIPE_MIN_CHUNK: usize = 128;
402const PRIME_PIPE_EDGE_MIN_CHUNK: usize = 64;
403const PRIME_PIPE_LINEAR_WORK: usize = 8;
404
405fn prime_pp2_auto_geometry(n_layers: usize) -> bool {
406 crate::pp::prime_pp_on()
407 && !crate::pp::pp2_streams_off()
408 && crate::pp::pp_cuts(n_layers).is_some_and(|cuts| cuts.len() == 3)
409}
410
411pub fn prime_chunk_tokens(t: usize, n_layers: usize) -> usize {
415 if let Ok(value) = std::env::var("MEMRA_PRIME_CHUNK") {
416 let parsed = value
417 .parse::<usize>()
418 .unwrap_or(crate::cache::PRIME_CHUNK_MAX_TOKENS);
419 return if crate::cache::swa_ring_on() {
420 if parsed == 0 {
421 crate::cache::PRIME_CHUNK_MAX_TOKENS
422 } else {
423 parsed.min(crate::cache::PRIME_CHUNK_MAX_TOKENS)
424 }
425 } else {
426 parsed
427 };
428 }
429 let chunk = crate::cache::PRIME_CHUNK_MAX_TOKENS;
430 if prime_pp2_auto_geometry(n_layers) && t >= 2 * PRIME_PIPE_MIN_CHUNK {
431 chunk.min(
432 t.div_ceil(PRIME_PIPE_MICROBATCHES)
433 .max(PRIME_PIPE_MIN_CHUNK),
434 )
435 } else {
436 chunk
437 }
438}
439
440fn fixed_prime_chunk_ranges(t: usize, chunk: usize) -> Vec<(usize, usize)> {
441 fixed_prime_chunk_ranges_for_ring(t, chunk, crate::cache::swa_ring_on())
442}
443
444fn fixed_prime_chunk_ranges_for_ring(t: usize, chunk: usize, ring_on: bool) -> Vec<(usize, usize)> {
445 if chunk == 0 || t <= chunk {
446 return vec![(0, t)];
447 }
448 let mut ranges = Vec::with_capacity(t.div_ceil(chunk));
449 let mut start = 0usize;
450 while start < t {
451 let mut end = (start + chunk).min(t);
452 if t - end > 0 && t - end < PRIME_MIN_T {
453 if ring_on {
454 let shifted = t - PRIME_MIN_T;
455 end = if shifted > start { shifted } else { t };
456 } else {
457 end = t;
458 }
459 }
460 ranges.push((start, end));
461 start = end;
462 }
463 ranges
464}
465
466fn prime_chunk_work(prefix: usize, total: usize) -> u128 {
467 let prefix = prefix as u128;
468 prefix * (prefix + (PRIME_PIPE_LINEAR_WORK as u128) * (total as u128))
469}
470
471fn dynamic_prime_chunk_ranges(
472 t: usize,
473 fixed_chunk: usize,
474 fixed: &[(usize, usize)],
475) -> Vec<(usize, usize)> {
476 let n = fixed.len();
477 if n < 3 {
478 return fixed.to_vec();
479 }
480
481 let max_first = t - (n - 1) * PRIME_MIN_T;
482 let first = fixed_chunk
483 .div_ceil(2)
484 .max(PRIME_PIPE_EDGE_MIN_CHUNK)
485 .min(max_first);
486 let mut ranges = Vec::with_capacity(n);
487 ranges.push((0, first));
488
489 let first_work = prime_chunk_work(first, t);
490 let work_span = prime_chunk_work(t, t) - first_work;
491 let denominator = (n - 1) as u128;
492 let mut previous = first;
493 for boundary in 1..n - 1 {
494 let target = first_work * denominator + work_span * (boundary as u128);
495 let remaining = n - 1 - boundary;
496 let mut low = previous + PRIME_MIN_T;
497 let mut high = t - remaining * PRIME_MIN_T;
498 while low < high {
499 let mid = low + (high - low) / 2;
500 if prime_chunk_work(mid, t) * denominator >= target {
501 high = mid;
502 } else {
503 low = mid + 1;
504 }
505 }
506 ranges.push((previous, low));
507 previous = low;
508 }
509 ranges.push((previous, t));
510 ranges
511}
512
513pub fn prime_chunk_ranges(t: usize, n_layers: usize) -> Vec<(usize, usize)> {
517 let explicit_chunk = std::env::var_os("MEMRA_PRIME_CHUNK").is_some();
518 let chunk = prime_chunk_tokens(t, n_layers);
519 let fixed = fixed_prime_chunk_ranges(t, chunk);
520 let dynamic = match std::env::var("MEMRA_PRIME_CHUNK_SCHED") {
521 Ok(value) => value == "dynamic",
522 Err(_) => true,
523 };
524 if explicit_chunk || !dynamic || !prime_pp2_auto_geometry(n_layers) {
525 fixed
526 } else {
527 dynamic_prime_chunk_ranges(t, chunk, &fixed)
528 }
529}
530
531impl HybridModel {
532 fn prime_trace_path() -> Option<&'static str> {
537 static P: std::sync::OnceLock<Option<String>> = std::sync::OnceLock::new();
538 P.get_or_init(|| std::env::var("MEMRA_PRIME_TRACE").ok())
539 .as_deref()
540 }
541
542 pub fn forward(
544 &self,
545 e: &Engine,
546 tokens: &[u32],
547 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
548 if self.is_gemma4_e4b() {
549 return self.gemma4_e4b_forward(e, tokens, false);
550 }
551 if self.cfg.gemma4.is_some() {
552 return self.gemma4_forward(e, tokens, false);
553 }
554 let cfg = &self.cfg;
555 let n_embd = cfg.n_embd as usize;
556 let t = tokens.len();
557 let eps = cfg.rms_eps;
558 let pos: Vec<i32> = (0..t as i32).collect();
559 let pos_d = e.htod_i32(&pos)?;
560
561 let mut x = self.embed(e, tokens)?; for (il, layer) in self.layers.iter().enumerate() {
564 let mut h = e.uninit(t * n_embd)?;
566 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
567
568 let mixed = match &layer.mixer {
569 Mixer::Full(fa) => self.full_attn(e, fa, &h, &pos_d, t, il)?,
570 Mixer::Linear(la) => self.linear_attn(e, la, &h, t)?,
571 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
572 };
573
574 let mut x1 = e.uninit(t * n_embd)?;
576 e.add(&x, &mixed, &mut x1, t * n_embd)?;
577
578 let mut z = e.uninit(t * n_embd)?;
580 e.rms_norm(
581 &x1,
582 layer.post_attn_norm.float_data(),
583 &mut z,
584 n_embd,
585 t,
586 eps,
587 )?;
588 let ffn_out = match &layer.ffn {
589 crate::hybrid::Ffn::Dense {
590 ffn_gate,
591 ffn_up,
592 ffn_down,
593 } => {
594 let n_ff = ffn_gate.out_features();
595 let mut g2 = e.matmul_group(&[ffn_gate, ffn_up], &z, t)?;
596 let up = g2.pop().unwrap();
597 let gate = g2.pop().unwrap();
598 let mut act = e.uninit(t * n_ff)?;
599 Self::ffn_act_lim(
604 e,
605 &self.cfg,
606 &gate,
607 &up,
608 1.0,
609 1.0,
610 self.cfg.clamp_shexp_at(il as u32),
611 &mut act,
612 t * n_ff,
613 )?;
614 e.matmul(ffn_down, &act, t)?
615 }
616 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il_prefill(e, m, &z, t, il as u16)?,
617 };
618 let mut x2 = e.uninit(t * n_embd)?;
619 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
620 x = x2;
621 }
622
623 let mut hn = e.uninit(t * n_embd)?;
624 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
625 let logits = e.matmul(&self.output, &hn, t)?;
626 Ok(e.dtoh(&logits)?)
627 }
628
629 pub fn forward_last(
635 &self,
636 e: &Engine,
637 tokens: &[u32],
638 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
639 if self.cfg.gemma4.is_some() {
640 return self.gemma4_forward(e, tokens, true);
641 }
642 let cfg = &self.cfg;
643 let n_embd = cfg.n_embd as usize;
644 let t = tokens.len();
645 let eps = cfg.rms_eps;
646 let pos: Vec<i32> = (0..t as i32).collect();
647 let pos_d = e.htod_i32(&pos)?;
648
649 let mut x = self.embed(e, tokens)?; let probe = std::env::var("MEMRA_LAYER_PROBE").is_ok();
653 for (il, layer) in self.layers.iter().enumerate() {
654 let mut h = e.uninit(t * n_embd)?;
655 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
656 if probe {
657 e.stream().synchronize()?;
658 eprintln!("[probe] L{il} norm ok");
659 }
660 let mixed = match &layer.mixer {
661 Mixer::Full(fa) => self.full_attn(e, fa, &h, &pos_d, t, il)?,
662 Mixer::Linear(la) => self.linear_attn(e, la, &h, t)?,
663 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
664 };
665 if probe {
666 e.stream().synchronize()?;
667 eprintln!("[probe] L{il} mixer ok");
668 }
669 let mut x1 = e.uninit(t * n_embd)?;
670 e.add(&x, &mixed, &mut x1, t * n_embd)?;
671 let mut z = e.uninit(t * n_embd)?;
672 e.rms_norm(
673 &x1,
674 layer.post_attn_norm.float_data(),
675 &mut z,
676 n_embd,
677 t,
678 eps,
679 )?;
680 let ffn_out = match &layer.ffn {
681 crate::hybrid::Ffn::Dense {
682 ffn_gate,
683 ffn_up,
684 ffn_down,
685 } => {
686 let n_ff = ffn_gate.out_features();
687 let mut g2 = e.matmul_group(&[ffn_gate, ffn_up], &z, t)?;
688 let up = g2.pop().unwrap();
689 let gate = g2.pop().unwrap();
690 let mut act = e.uninit(t * n_ff)?;
691 Self::ffn_act_lim(
693 e,
694 &self.cfg,
695 &gate,
696 &up,
697 1.0,
698 1.0,
699 self.cfg.clamp_shexp_at(il as u32),
700 &mut act,
701 t * n_ff,
702 )?;
703 e.matmul(ffn_down, &act, t)?
704 }
705 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il_prefill(e, m, &z, t, il as u16)?,
706 };
707 if probe {
708 e.stream().synchronize()?;
709 eprintln!("[probe] L{il} ffn ok");
710 }
711 let mut x2 = e.uninit(t * n_embd)?;
712 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
713 x = x2;
714 }
715 let mut hn = e.uninit(t * n_embd)?;
717 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
718 let last = e.view(&hn, t * n_embd); let last_row = last.slice((t - 1) * n_embd..t * n_embd); let mut hlast = e.uninit(n_embd)?;
721 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
722 let logits = e.matmul(&self.output, &hlast, 1)?; Ok(e.dtoh(&logits)?)
724 }
725
726 pub fn prime_cache(
758 &self,
759 e: &Engine,
760 tokens: &[u32],
761 cache: &mut Cache,
762 queued_after: usize,
763 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
764 self.prime_cache_overlaid(e, tokens, cache, queued_after, None)
765 }
766
767 pub fn prime_cache_overlaid(
773 &self,
774 e: &Engine,
775 tokens: &[u32],
776 cache: &mut Cache,
777 queued_after: usize,
778 overlay: Option<&crate::vision::EmbedOverlay>,
779 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
780 let n_embd = self.cfg.n_embd as usize;
781 let t = tokens.len();
782 assert!(
786 t >= PRIME_MIN_T,
787 "prime_cache needs T >= {PRIME_MIN_T} (caller gates)"
788 );
789 assert!(
790 cache.pos + t <= cache.max_ctx,
791 "prime_cache: prompt exceeds cache max_ctx"
792 );
793
794 if self.is_gemma4_e4b() || self.cfg.gemma4.is_some() {
806 if self.is_gemma4_e4b() {
807 if overlay.is_some() {
808 return Err(
809 "vision embedding overlay is unsupported on gemma4 E4B (PLE prime)".into(),
810 );
811 }
812 return self.gemma4_e4b_prime(e, tokens, cache);
813 }
814 return self.gemma4_prime(e, tokens, cache, overlay);
819 }
820 let ranges = prime_chunk_ranges(t, self.layers.len());
821 let legacy_calllocal = std::env::var("MEMRA_PRIME_CALLLOCAL").as_deref() == Ok("1");
857 let seq_end = if legacy_calllocal {
858 cache.pos + t
859 } else {
860 cache.pos + t + queued_after
861 };
862 if ranges.len() == 1 {
863 return self.prime_chunk(e, tokens, cache, seq_end, 0, overlay);
864 }
865 if crate::pp::prime_pipe_on() && crate::pp::prime_pp_on() && !crate::pp::pp2_streams_off() {
870 if let Some(fence) = crate::pp::pp_cuts(self.layers.len()).filter(|f| f.len() == 3) {
871 if overlay.is_some() {
872 return Err(
873 "vision embedding overlay + pipelined PP prime unsupported (v1); \
874 run the serial prime (single device or MEMRA_PRIME_PIPE=0)"
875 .into(),
876 );
877 }
878 if crate::pp::pp_multi_stream_same_device() {
879 return Err(
880 "prime chunk pipeline refused with 2 stage streams on one device — \
881 that concurrent-stream placement remains quarantined by the deferred \
882 pp flake record. Use one device per stage or MEMRA_PRIME_PIPE=0 for \
883 the serial split."
884 .into(),
885 );
886 }
887 return self.prime_cache_pp2_pipelined(e, tokens, cache, seq_end, &ranges, &fence);
888 }
889 }
890 let mut hiddens = e.uninit(t * n_embd)?;
891 let mut last: Option<(Vec<f32>, CudaSlice<f32>)> = None;
892 for &(start, end) in &ranges {
893 if let Some(taps) = cache.dflash_taps.as_mut() {
895 taps.base = start;
896 }
897 let (l, hs, x) =
898 self.prime_chunk(e, &tokens[start..end], cache, seq_end, start, overlay)?;
899 e.copy_into(&mut hiddens, start * n_embd, &x, (end - start) * n_embd)?;
900 last = Some((l, hs));
901 }
902 let (logits, h_seed) = last.unwrap();
903 Ok((logits, h_seed, hiddens))
904 }
905
906 fn prime_cache_pp2_pipelined(
911 &self,
912 e: &Engine,
913 tokens: &[u32],
914 cache: &mut Cache,
915 seq_end: usize,
916 ranges: &[(usize, usize)],
917 fence: &[usize],
918 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
919 debug_assert_eq!(fence.len(), 3);
920 debug_assert!(ranges.len() >= 2);
921 let rt = crate::pp::PpNRt::get(e)?;
922 assert_eq!(
923 rt.n_stages(),
924 2,
925 "prime pipeline requires exactly two PP stages"
926 );
927 let n_embd = self.cfg.n_embd as usize;
928 let t = tokens.len();
929 let initial_base = cache.pos;
930 let caller_stream = e.stream();
931
932 rt.fence_stages_behind(&caller_stream)?;
937 let max_payload = ranges.iter().map(|(s, e)| (e - s) * n_embd).max().unwrap();
938 rt.prepare_overlap_slots(0, max_payload)?;
939
940 let mut hiddens = e.uninit(t * n_embd)?;
941 let mut last: Option<(Vec<f32>, CudaSlice<f32>)> = None;
942 let mut stage_caches = PrimeCacheStages::new(cache, fence[1]);
943 let (cache0, cache1) = stage_caches.parts();
944 let (first_start, first_end) = ranges[0];
945 let mut slot = self.prime_pp2_stage0_enqueue(
946 e,
947 rt,
948 &tokens[first_start..first_end],
949 cache0,
950 seq_end,
951 fence,
952 initial_base + first_start,
953 true,
954 )?;
955 cache0.pos = initial_base + first_end;
956
957 for (i, &(start, end)) in ranges.iter().enumerate() {
958 let base = initial_base + start;
959 debug_assert_eq!(
960 cache1.pos, base,
961 "stage 1 must drain chunks in original position order"
962 );
963 let (out, next_slot) = if let Some(&(next_start, next_end)) = ranges.get(i + 1) {
964 let next_base = initial_base + next_start;
965 debug_assert_eq!(
966 cache0.pos, next_base,
967 "stage 0 must issue chunks in original position order"
968 );
969 let cache0_stage = &mut *cache0;
970 std::thread::scope(|scope| -> Result<_, Box<dyn std::error::Error>> {
975 let stage0 = scope.spawn(move || -> Result<usize, String> {
976 let next = self
977 .prime_pp2_stage0_enqueue(
978 e,
979 rt,
980 &tokens[next_start..next_end],
981 cache0_stage,
982 seq_end,
983 fence,
984 next_base,
985 true,
986 )
987 .map_err(|err| err.to_string())?;
988 cache0_stage.pos = initial_base + next_end;
989 Ok(next)
990 });
991 let x = self.prime_pp2_stage1_enqueue(
992 e,
993 rt,
994 slot,
995 end - start,
996 cache1,
997 seq_end,
998 fence,
999 base,
1000 true,
1001 )?;
1002 let out = {
1003 rt.bind_stage(1)?;
1004 let _st1 = rt.enter(1);
1005 let e1 = rt.engine(1, e);
1006 self.prime_chunk_epilogue(e1, x, end - start, cache1)?
1007 };
1008 let next = stage0
1009 .join()
1010 .map_err(|_| "pipeprime stage-0 host walker panicked")?
1011 .map_err(|err| -> Box<dyn std::error::Error> { err.into() })?;
1012 Ok((out, Some(next)))
1013 })?
1014 } else {
1015 let x = self.prime_pp2_stage1_enqueue(
1016 e,
1017 rt,
1018 slot,
1019 end - start,
1020 cache1,
1021 seq_end,
1022 fence,
1023 base,
1024 true,
1025 )?;
1026 let out = {
1027 rt.bind_stage(1)?;
1028 let _st1 = rt.enter(1);
1029 let e1 = rt.engine(1, e);
1030 self.prime_chunk_epilogue(e1, x, end - start, cache1)?
1031 };
1032 (out, None)
1033 };
1034
1035 rt.publish_to(1, &caller_stream)?;
1036 e.copy_into(&mut hiddens, start * n_embd, &out.2, (end - start) * n_embd)?;
1037 last = Some((out.0, out.1));
1038 crate::pp::PRIME_SPLIT_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
1039
1040 if let Some(next) = next_slot {
1041 rt.fence_stages_behind(&caller_stream)?;
1046 slot = next;
1047 }
1048 }
1049
1050 debug_assert_eq!(cache0.pos, initial_base + t);
1051 debug_assert_eq!(cache1.pos, initial_base + t);
1052 let (logits, h_seed) = last.unwrap();
1053 Ok((logits, h_seed, hiddens))
1054 }
1055
1056 fn gdn_hk(e: &Engine, t: usize, num_v: usize, num_k: usize) -> usize {
1063 if Engine::gdn_db_on()
1064 && Engine::gdn_chunked_enabled()
1065 && t >= 16
1066 && e.gdn_mma_enabled(Engine::gdn_chunk_size())
1067 && num_k * 2 == num_v
1068 {
1069 num_k
1070 } else {
1071 num_v
1072 }
1073 }
1074
1075 fn f16out_on(e: &Engine, t: usize) -> bool {
1080 crate::f16_ffi::pp_f16_enabled()
1081 && t >= 16
1082 && !e.verify_exact_on()
1083 && std::env::var("MEMRA_F16OUT").as_deref() != Ok("0")
1084 }
1085
1086 pub fn prime_slabs_get(
1094 &self,
1095 e: &Engine,
1096 t: usize,
1097 n_embd: usize,
1098 n_ff_max: usize,
1099 ) -> Result<std::sync::Arc<std::sync::Mutex<PrimeSlabs>>, Box<dyn std::error::Error>> {
1100 let mut slabs = self.prime_slabs.lock().unwrap();
1101 let dev = e.ctx().ordinal();
1102 let need_new = match slabs.get(&dev) {
1103 None => true,
1104 Some(sl) => sl.lock().unwrap().t_cap < t,
1105 };
1106 if need_new {
1107 slabs.insert(
1108 dev,
1109 std::sync::Arc::new(std::sync::Mutex::new(PrimeSlabs {
1110 t_cap: t,
1111 h: e.uninit(t * n_embd)?,
1112 x1: e.uninit(t * n_embd)?,
1113 z: e.uninit(t * n_embd)?,
1114 act: e.uninit(t * n_ff_max)?,
1115 xa: e.uninit(t * n_embd)?,
1116 xb: e.uninit(t * n_embd)?,
1117 h16: e.alloc_u8_uninit(t * n_embd * 2)?,
1118 z16: e.alloc_u8_uninit(t * n_embd * 2)?,
1119 gate: e.uninit(t * n_ff_max)?,
1120 up: e.uninit(t * n_ff_max)?,
1121 ffn_out: e.uninit(t * n_embd)?,
1122 seg_glue: Vec::new(),
1123 mixed: e.uninit(t * n_embd)?,
1124 seg_mid: Vec::new(),
1125 seg_t: 0,
1126 })),
1127 );
1128 }
1129 Ok(slabs.get(&dev).expect("prime slab inserted").clone())
1130 }
1131
1132 fn prime_chunk(
1136 &self,
1137 e: &Engine,
1138 tokens: &[u32],
1139 cache: &mut Cache,
1140 seq_end: usize,
1141 chunk_off: usize,
1142 overlay: Option<&crate::vision::EmbedOverlay>,
1143 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
1144 if crate::pp::pp_host_bounce_active()
1145 && (self.cfg.gemma4.is_some() || !crate::pp::prime_pp_on())
1146 {
1147 return Err(
1148 "prime_chunk: refused with MEMRA_PP_HOST_BOUNCE=1 because this configuration \
1149 has no active prime stage split and would peer-read remote weights; keep \
1150 MEMRA_PRIME_PP enabled and use a PP-prime-supported model"
1151 .into(),
1152 );
1153 }
1154 if self.cfg.gemma4.is_none() && !crate::pp::pp2_streams_off() && crate::pp::prime_pp_on() {
1163 if let Some(fence) = crate::pp::pp_cuts(self.layers.len()) {
1164 if overlay.is_some() {
1165 return Err("vision embedding overlay + PP prime unsupported (v1); \
1166 run single-device or MEMRA_PRIME_PP=0"
1167 .into());
1168 }
1169 return self.prime_chunk_ppn(e, tokens, cache, seq_end, &fence);
1170 }
1171 }
1172 if crate::pp::pp_host_bounce_active() {
1173 return Err(
1174 "prime_chunk: MEMRA_PP_HOST_BOUNCE=1 found no valid prime stage split; \
1175 refusing an unsplit remote-weight walk"
1176 .into(),
1177 );
1178 }
1179 let t = tokens.len();
1180 let base = cache.pos;
1181 debug_assert!(
1182 seq_end >= base + t,
1183 "prime_chunk: seq_end must cover this chunk"
1184 );
1185 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
1186 let pos_d = e.htod_i32(&pos)?;
1187
1188 let mut x_embed = self.embed(e, tokens)?; if let Some(ov) = overlay {
1190 let n_embd = self.cfg.n_embd as usize;
1194 for &(pos, row_off, n_rows) in &ov.spans {
1195 let lo = pos.max(chunk_off);
1196 let hi = (pos + n_rows).min(chunk_off + t);
1197 if lo < hi {
1198 let src_row = row_off + (lo - pos);
1199 let view = ov
1200 .rows
1201 .slice(src_row * n_embd..(src_row + (hi - lo)) * n_embd);
1202 e.copy_view_into(
1203 &mut x_embed,
1204 (lo - chunk_off) * n_embd,
1205 &view,
1206 (hi - lo) * n_embd,
1207 )?;
1208 }
1209 }
1210 }
1211 let x = self.prime_layers(
1212 e,
1213 x_embed,
1214 0,
1215 self.layers.len(),
1216 &pos_d,
1217 t,
1218 base,
1219 cache,
1220 seq_end,
1221 )?;
1222 self.prime_chunk_epilogue(e, x, t, cache)
1223 }
1224
1225 #[allow(clippy::too_many_arguments)]
1241 fn prime_layers(
1242 &self,
1243 e: &Engine,
1244 x_in: CudaSlice<f32>,
1245 lo: usize,
1246 hi: usize,
1247 pos_d: &CudaSlice<i32>,
1248 t: usize,
1249 base: usize,
1250 cache: &mut Cache,
1251 seq_end: usize,
1252 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1253 let cfg = &self.cfg;
1254 let n_embd = cfg.n_embd as usize;
1255 let eps = cfg.rms_eps;
1256 let f16fuse = crate::f16_ffi::pp_f16_enabled() && t >= 16;
1260 let n_ff_max = self
1266 .layers
1267 .iter()
1268 .map(|l| match &l.ffn {
1269 crate::hybrid::Ffn::Dense { ffn_gate, .. } => ffn_gate.out_features(),
1270 _ => n_embd,
1271 })
1272 .max()
1273 .unwrap_or(n_embd)
1274 .max(n_embd);
1275 let use_slabs = std::env::var("MEMRA_PRIME_SLABS").as_deref() != Ok("0");
1276 let slab = if use_slabs {
1277 Some(self.prime_slabs_get(e, t, n_embd, n_ff_max)?)
1278 } else {
1279 None
1280 };
1281 let mut slab_guard = slab.as_ref().map(|sl| sl.lock().unwrap());
1282 let mut x_own; type SlabRefs<'a> = (
1284 &'a mut CudaSlice<f32>,
1285 &'a mut CudaSlice<f32>,
1286 &'a mut CudaSlice<f32>,
1287 &'a mut CudaSlice<f32>,
1288 &'a mut CudaSlice<u8>,
1289 &'a mut CudaSlice<u8>,
1290 &'a mut CudaSlice<f32>,
1291 &'a mut CudaSlice<f32>,
1292 &'a mut CudaSlice<f32>,
1293 );
1294 let (mut x_cur, mut x_nxt, sl): (
1295 &mut CudaSlice<f32>,
1296 &mut CudaSlice<f32>,
1297 Option<SlabRefs>,
1298 );
1299 let mut seg: Option<(
1300 &mut Vec<Option<cudarc::driver::CudaGraph>>,
1301 &mut Vec<Option<cudarc::driver::CudaGraph>>,
1302 &mut CudaSlice<f32>,
1303 &mut usize,
1304 )> = None;
1305 let mut x_own2;
1306 match slab_guard.as_mut() {
1307 Some(g) => {
1308 let slabs = &mut **g;
1309 e.copy_into(&mut slabs.xa, 0, &x_in, t * n_embd)?;
1310 let PrimeSlabs {
1311 xa,
1312 xb,
1313 h,
1314 x1,
1315 z,
1316 act,
1317 h16,
1318 z16,
1319 gate,
1320 up,
1321 ffn_out,
1322 seg_glue,
1323 mixed,
1324 seg_mid,
1325 seg_t,
1326 ..
1327 } = slabs;
1328 x_cur = xa;
1329 x_nxt = xb;
1330 seg = Some((seg_glue, seg_mid, mixed, seg_t));
1331 sl = Some((h, x1, z, act, h16, z16, gate, up, ffn_out));
1332 }
1333 None => {
1334 x_own = x_in;
1335 x_own2 = e.uninit(t * n_embd)?;
1336 x_cur = &mut x_own;
1337 x_nxt = &mut x_own2;
1338 sl = None;
1339 }
1340 }
1341 let mut alloc_h;
1342 let mut alloc_x1;
1343 let mut alloc_z;
1344 let mut alloc_act;
1345 let mut alloc_h16;
1346 let mut alloc_z16;
1347 let mut alloc_gate;
1348 let mut alloc_up;
1349 let mut alloc_fo;
1350 let (h, x1, z, act): (
1351 &mut CudaSlice<f32>,
1352 &mut CudaSlice<f32>,
1353 &mut CudaSlice<f32>,
1354 &mut CudaSlice<f32>,
1355 );
1356 let (h16, z16): (&mut CudaSlice<u8>, &mut CudaSlice<u8>);
1357 let (sl_gate, sl_up, sl_fo): (
1358 &mut CudaSlice<f32>,
1359 &mut CudaSlice<f32>,
1360 &mut CudaSlice<f32>,
1361 );
1362 match sl {
1363 Some((a, b, c, d, e16, f16b, g, u, fo)) => {
1364 h = a;
1365 x1 = b;
1366 z = c;
1367 act = d;
1368 h16 = e16;
1369 z16 = f16b;
1370 sl_gate = g;
1371 sl_up = u;
1372 sl_fo = fo;
1373 }
1374 None => {
1375 alloc_h = e.uninit(t * n_embd)?;
1376 alloc_x1 = e.uninit(t * n_embd)?;
1377 alloc_z = e.uninit(t * n_embd)?;
1378 alloc_act = e.uninit(t * n_ff_max)?;
1379 alloc_h16 = e.alloc_u8_uninit(t * n_embd * 2)?;
1380 alloc_z16 = e.alloc_u8_uninit(t * n_embd * 2)?;
1381 alloc_gate = e.uninit(t * n_ff_max)?;
1382 alloc_up = e.uninit(t * n_ff_max)?;
1383 alloc_fo = e.uninit(t * n_embd)?;
1384 h = &mut alloc_h;
1385 x1 = &mut alloc_x1;
1386 z = &mut alloc_z;
1387 act = &mut alloc_act;
1388 h16 = &mut alloc_h16;
1389 z16 = &mut alloc_z16;
1390 sl_gate = &mut alloc_gate;
1391 sl_up = &mut alloc_up;
1392 sl_fo = &mut alloc_fo;
1393 }
1394 }
1395 let n_layers = self.layers.len();
1400 let use_seg = f16fuse
1410 && seg.is_some()
1411 && self.cfg.step35.is_none()
1412 && lo == 0
1413 && hi == n_layers
1414 && std::env::var("MEMRA_PRIME_SEG").as_deref() == Ok("1");
1415 if let Some((sg, sm, _, st)) = seg.as_mut() {
1416 if **st != t {
1417 sg.clear();
1418 sg.extend((0..n_layers).map(|_| None));
1419 sm.clear();
1420 sm.extend((0..n_layers).map(|_| None));
1421 **st = t;
1422 }
1423 }
1424 {
1425 let layer_lo = &self.layers[lo];
1426 if f16fuse {
1427 e.rms_norm_f16out(
1428 x_cur,
1429 layer_lo.attn_norm.float_data(),
1430 h,
1431 h16,
1432 n_embd,
1433 t,
1434 eps,
1435 )?;
1436 } else {
1437 e.rms_norm(x_cur, layer_lo.attn_norm.float_data(), h, n_embd, t, eps)?;
1438 }
1439 }
1440 for il in lo..hi {
1441 let layer = &self.layers[il];
1442 let hx16 = if f16fuse { Some(&*h16) } else { None };
1443 if use_seg {
1444 let (pre, pre16, w_out) = match &layer.mixer {
1447 Mixer::Full(fa) => {
1448 let g3 = match hx16 {
1449 Some(xh) => e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], h, xh, t)?,
1450 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
1451 };
1452 let (pre, pre16) =
1453 self.full_attn_prime_core_inner(e, fa, g3, &pos_d, t, cache, il)?;
1454 (pre, pre16, &fa.wo)
1455 }
1456 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
1457 Mixer::Linear(la) => {
1458 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
1459 let g4 = match hx16 {
1460 Some(xh) => e.matmul_group_xh(&ws, h, xh, t)?,
1461 None => e.matmul_group(&ws, h, t)?,
1462 };
1463 let (pre, pre16) =
1464 self.linear_attn_prime_core_pad_inner(e, la, g4, t, cache, il, None)?;
1465 (pre, pre16, &la.ssm_out)
1466 }
1467 };
1468 {
1469 let (_, sm, mslab, _) = seg.as_mut().unwrap();
1470 let pre_n = pre.len() / t;
1471 let xh_pre = match pre16 {
1472 Some(x) => x,
1473 None => e.f16_act(&pre, t * pre_n, pre_n)?,
1474 };
1475 if !e.try_f16_gemm_pre_into(w_out, &xh_pre, t, mslab)? {
1476 let y = e.matmul(w_out, &pre, t)?;
1477 e.copy_into(mslab, 0, &y, t * n_embd)?;
1478 }
1479 if sm[il].is_none() {
1480 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
1481 let w_post = layer.post_attn_norm.float_data();
1482 e.stream().synchronize()?;
1483 e.stream()
1484 .begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
1485 let r = (|| -> Result<(), Box<dyn std::error::Error>> {
1486 e.add(x_cur, mslab, x1, t * n_embd)?;
1487 e.rms_norm_f16out(x1, w_post, z, z16, n_embd, t, eps)?;
1488 Ok(())
1489 })();
1490 let g = e.stream().end_capture(
1491 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH);
1492 r?;
1493 sm[il] = Some(g?.ok_or("S-mid capture produced no graph")?);
1494 }
1495 sm[il].as_ref().unwrap().launch()?;
1496 }
1497 } else {
1498 let mixed = match &layer.mixer {
1499 Mixer::Full(fa) => {
1500 self.full_attn_prime(e, fa, h, hx16, &pos_d, t, cache, il, seq_end)?
1501 }
1502 Mixer::Linear(la) => self.linear_attn_prime(e, la, h, hx16, t, cache, il)?,
1503 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
1504 };
1505 if f16fuse {
1506 e.add_rms_norm_f16out(
1509 x_cur,
1510 &mixed,
1511 layer.post_attn_norm.float_data(),
1512 x1,
1513 z,
1514 z16,
1515 n_embd,
1516 t,
1517 eps,
1518 )?;
1519 } else {
1520 e.add(x_cur, &mixed, x1, t * n_embd)?;
1521 e.rms_norm(x1, layer.post_attn_norm.float_data(), z, n_embd, t, eps)?;
1522 }
1523 }
1524 let zx16 = if f16fuse { Some(&*z16) } else { None };
1525 match &layer.ffn {
1526 crate::hybrid::Ffn::Dense {
1527 ffn_gate,
1528 ffn_up,
1529 ffn_down,
1530 } => {
1531 let n_ff = ffn_gate.out_features();
1532 let mut into_ok = false;
1535 if let Some(xh) = zx16 {
1536 into_ok = e.try_f16_gemm_pre_into(ffn_gate, xh, t, sl_gate)?
1537 && e.try_f16_gemm_pre_into(ffn_up, xh, t, sl_up)?;
1538 }
1539 if !into_ok {
1540 let mut g2 = match zx16 {
1541 Some(xh) => e.matmul_group_xh(&[ffn_gate, ffn_up], z, xh, t)?,
1542 None => e.matmul_group(&[ffn_gate, ffn_up], z, t)?,
1543 };
1544 let up_y = g2.pop().unwrap();
1545 let gate_y = g2.pop().unwrap();
1546 e.copy_into(sl_gate, 0, &gate_y, t * n_ff)?;
1547 e.copy_into(sl_up, 0, &up_y, t * n_ff)?;
1548 }
1549 let d_lim = self.cfg.clamp_shexp_at(il as u32);
1554 let act16 = if Self::f16out_on(e, t) && self.cfg.m3.is_none() && d_lim.is_none()
1555 {
1556 let mut a16 = e.alloc_u8_uninit(t * n_ff * 2)?;
1557 e.silu_mul_f16out(sl_gate, sl_up, act, &mut a16, t * n_ff)?;
1558 Some(a16)
1559 } else {
1560 Self::ffn_act_lim(
1561 e,
1562 &self.cfg,
1563 sl_gate,
1564 sl_up,
1565 1.0,
1566 1.0,
1567 d_lim,
1568 act,
1569 t * n_ff,
1570 )?;
1571 None
1572 };
1573 let xh_act = match act16 {
1575 Some(x) => x,
1576 None => e.f16_act(act, t * n_ff, n_ff)?,
1577 };
1578 if !e.try_f16_gemm_pre_into(ffn_down, &xh_act, t, sl_fo)? {
1579 let y = e.matmul(ffn_down, &*act, t)?;
1580 e.copy_into(sl_fo, 0, &y, t * n_embd)?;
1581 }
1582 }
1583 crate::hybrid::Ffn::Moe(m) => {
1584 let y = self.moe_ffn_il_prefill(e, m, z, t, il as u16)?;
1585 e.copy_into(sl_fo, 0, &y, t * n_embd)?;
1586 }
1587 }
1588 if use_seg && il + 1 < hi {
1589 let w_next = self.layers[il + 1].attn_norm.float_data();
1591 let (sg, _, _, _) = seg.as_mut().unwrap();
1592 if sg[il].is_none() {
1593 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
1594 e.stream().synchronize()?;
1595 e.stream()
1596 .begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
1597 let r = (|| -> Result<(), Box<dyn std::error::Error>> {
1598 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
1599 e.rms_norm_f16out(x_nxt, w_next, h, h16, n_embd, t, eps)?;
1600 Ok(())
1601 })();
1602 let g = e.stream().end_capture(
1603 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
1604 );
1605 r?;
1606 sg[il] = Some(g?.ok_or("S-glue capture produced no graph")?);
1607 }
1608 sg[il].as_ref().unwrap().launch()?;
1609 } else {
1610 if il + 1 < hi {
1611 let w_next = self.layers[il + 1].attn_norm.float_data();
1612 if f16fuse {
1613 e.add_rms_norm_f16out(x1, sl_fo, w_next, x_nxt, h, h16, n_embd, t, eps)?;
1614 } else {
1615 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
1616 e.rms_norm(x_nxt, w_next, h, n_embd, t, eps)?;
1617 }
1618 } else {
1619 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
1620 }
1621 }
1622 if let Some(path) = Self::prime_trace_path() {
1628 let row = (base + t - 1) as usize;
1629 let host = e.dtoh(x_nxt)?;
1630 let last = &host[(t - 1) * n_embd..t * n_embd];
1631 use std::io::Write as _;
1632 let mut f = std::fs::OpenOptions::new()
1633 .create(true)
1634 .append(true)
1635 .open(path)?;
1636 let mut h64: u64 = 0xcbf29ce484222325;
1637 for v in last {
1638 h64 ^= v.to_bits() as u64;
1639 h64 = h64.wrapping_mul(0x100000001b3);
1640 }
1641 writeln!(
1642 f,
1643 "{{\"pos\":{row},\"layer\":{il},\"t\":{t},\"base\":{base},\
1644 \"hash\":\"{h64:016x}\",\"v0\":{:.9e},\"v1\":{:.9e},\"v2\":{:.9e}}}",
1645 last[0], last[1], last[2]
1646 )?;
1647 }
1648 self.dflash_tap(e, cache, il, x_nxt, t)?;
1651 std::mem::swap(&mut x_cur, &mut x_nxt);
1652 }
1653 let mut x = e.uninit(t * n_embd)?;
1655 e.copy_into(&mut x, 0, x_cur, t * n_embd)?;
1656 drop(slab_guard);
1657 Ok(x)
1658 }
1659
1660 fn prime_chunk_epilogue(
1665 &self,
1666 e: &Engine,
1667 x: CudaSlice<f32>,
1668 t: usize,
1669 cache: &mut Cache,
1670 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
1671 let n_embd = self.cfg.n_embd as usize;
1672 let eps = self.cfg.rms_eps;
1673 let mut h_seed = e.uninit(n_embd)?;
1677 if !crate::spec::spec_hpost() {
1678 e.copy_view_into(
1679 &mut h_seed,
1680 0,
1681 &x.slice((t - 1) * n_embd..t * n_embd),
1682 n_embd,
1683 )?;
1684 }
1685 let mut hn = e.uninit(t * n_embd)?;
1687 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
1688 if crate::spec::spec_hpost() {
1689 e.copy_view_into(
1690 &mut h_seed,
1691 0,
1692 &hn.slice((t - 1) * n_embd..t * n_embd),
1693 n_embd,
1694 )?;
1695 }
1696 let last = e.view(&hn, t * n_embd);
1697 let last_row = last.slice((t - 1) * n_embd..t * n_embd);
1698 let mut hlast = e.uninit(n_embd)?;
1699 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
1700 let logits = e.matmul(&self.output, &hlast, 1)?;
1701 cache.pos += t;
1702 Ok((
1705 e.dtoh(&logits)?,
1706 h_seed,
1707 if crate::spec::spec_hpost() { hn } else { x },
1708 ))
1709 }
1710
1711 fn prime_chunk_ppn(
1735 &self,
1736 e: &Engine,
1737 tokens: &[u32],
1738 cache: &mut Cache,
1739 seq_end: usize,
1740 fence: &[usize],
1741 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
1742 let rt = crate::pp::PpNRt::get(e)?;
1743 let n_st = fence.len() - 1;
1744 assert_eq!(
1745 rt.n_stages(),
1746 n_st,
1747 "PpNRt stage count {} != fence stages {n_st}",
1748 rt.n_stages()
1749 );
1750 let n_embd = self.cfg.n_embd as usize;
1751 let t = tokens.len();
1752 let base = cache.pos;
1753 debug_assert!(
1754 seq_end >= base + t,
1755 "prime_chunk_ppn: seq_end must cover this chunk"
1756 );
1757 let payload = t * n_embd;
1758 let caller_stream = e.stream();
1762 rt.fence_stages_behind(&caller_stream)?;
1763
1764 if n_st == 2 {
1765 let slot =
1766 self.prime_pp2_stage0_enqueue(e, rt, tokens, cache, seq_end, fence, base, false)?;
1767 let x =
1768 self.prime_pp2_stage1_enqueue(e, rt, slot, t, cache, seq_end, fence, base, false)?;
1769 let out = {
1770 rt.bind_stage(1)?;
1771 let _st1 = rt.enter(1);
1772 let e1 = rt.engine(1, e);
1773 self.prime_chunk_epilogue(e1, x, t, cache)?
1774 };
1775 rt.publish_to(1, &caller_stream)?;
1776 crate::pp::PRIME_SPLIT_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
1777 return Ok(out);
1778 }
1779
1780 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
1781
1782 let mut slot = {
1784 let _st0 = rt.enter(0);
1785 let e0 = rt.engine(0, e);
1786 let pos_d = e0.htod_i32(&pos)?;
1787 let x = self.embed(e0, tokens)?;
1788 let x =
1789 self.prime_layers(e0, x, fence[0], fence[1], &pos_d, t, base, cache, seq_end)?;
1790 rt.tx(0, &x, payload)?
1791 };
1793
1794 for s in 1..n_st - 1 {
1796 let _st = rt.enter(s);
1797 let es = rt.engine(s, e);
1798 let pos_d = es.htod_i32(&pos)?;
1799 let x = rt.rx(s - 1, slot, payload)?;
1800 let x = self.prime_layers(
1801 es,
1802 x,
1803 fence[s],
1804 fence[s + 1],
1805 &pos_d,
1806 t,
1807 base,
1808 cache,
1809 seq_end,
1810 )?;
1811 slot = rt.tx(s, &x, payload)?;
1812 }
1813
1814 let _stl = rt.enter(n_st - 1);
1816 let el = rt.engine(n_st - 1, e);
1817 let pos_d = el.htod_i32(&pos)?;
1818 let x = rt.rx(n_st - 2, slot, payload)?;
1819 let x = self.prime_layers(
1820 el,
1821 x,
1822 fence[n_st - 1],
1823 fence[n_st],
1824 &pos_d,
1825 t,
1826 base,
1827 cache,
1828 seq_end,
1829 )?;
1830 let out = self.prime_chunk_epilogue(el, x, t, cache)?;
1831 rt.publish_to(n_st - 1, &caller_stream)?;
1837 crate::pp::PRIME_SPLIT_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
1838 Ok(out)
1839 }
1840
1841 fn prime_pp2_stage0_enqueue(
1842 &self,
1843 e: &Engine,
1844 rt: &crate::pp::PpNRt,
1845 tokens: &[u32],
1846 cache: &mut Cache,
1847 seq_end: usize,
1848 fence: &[usize],
1849 base: usize,
1850 pipelined: bool,
1851 ) -> Result<usize, Box<dyn std::error::Error>> {
1852 let t = tokens.len();
1853 let n_embd = self.cfg.n_embd as usize;
1854 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
1855 rt.bind_stage(0)?;
1856 let _st0 = rt.enter(0);
1857 let e0 = rt.engine(0, e);
1858 let pos_d = e0.htod_i32(&pos)?;
1859 let x = self.embed(e0, tokens)?;
1860 let _overlap = pipelined.then(crate::pp::enter_prime_pipe_stage);
1861 let x = self.prime_layers(e0, x, fence[0], fence[1], &pos_d, t, base, cache, seq_end)?;
1862 if pipelined {
1863 rt.tx_pipelined(0, &x, t * n_embd)
1864 } else {
1865 rt.tx(0, &x, t * n_embd)
1866 }
1867 }
1868
1869 fn prime_pp2_stage1_enqueue(
1870 &self,
1871 e: &Engine,
1872 rt: &crate::pp::PpNRt,
1873 slot: usize,
1874 t: usize,
1875 cache: &mut Cache,
1876 seq_end: usize,
1877 fence: &[usize],
1878 base: usize,
1879 pipelined: bool,
1880 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1881 let n_embd = self.cfg.n_embd as usize;
1882 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
1883 rt.bind_stage(1)?;
1884 let _st1 = rt.enter(1);
1885 let e1 = rt.engine(1, e);
1886 let pos_d = e1.htod_i32(&pos)?;
1887 let x = rt.rx(0, slot, t * n_embd)?;
1888 let _overlap = pipelined.then(crate::pp::enter_prime_pipe_stage);
1889 self.prime_layers(e1, x, fence[1], fence[2], &pos_d, t, base, cache, seq_end)
1890 }
1891
1892 pub fn prime_chunk_captured(
1908 &self,
1909 e: &Engine,
1910 x_in: &CudaSlice<f32>,
1911 pos_d: &CudaSlice<i32>,
1912 t: usize,
1913 cache: &mut Cache,
1914 len_d: &CudaSlice<i32>,
1915 logits_out: &mut CudaSlice<f32>,
1916 h_seed_out: &mut CudaSlice<f32>,
1917 ) -> Result<(), Box<dyn std::error::Error>> {
1918 let cfg = &self.cfg;
1919 let n_embd = cfg.n_embd as usize;
1920 let eps = cfg.rms_eps;
1921 let f16fuse = crate::f16_ffi::pp_f16_enabled() && t >= 16;
1922 let mut x = e.uninit(t * n_embd)?;
1923 e.copy_into(&mut x, 0, x_in, t * n_embd)?;
1924 for (il, layer) in self.layers.iter().enumerate() {
1925 let mut h = e.uninit(t * n_embd)?;
1926 let mut hx16: Option<CudaSlice<u8>> = None;
1927 if f16fuse {
1928 let mut b16 = e.alloc_u8_uninit(t * n_embd * 2)?;
1929 e.rms_norm_f16out(
1930 &x,
1931 layer.attn_norm.float_data(),
1932 &mut h,
1933 &mut b16,
1934 n_embd,
1935 t,
1936 eps,
1937 )?;
1938 hx16 = Some(b16);
1939 } else {
1940 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
1941 }
1942 let mixed = match &layer.mixer {
1943 Mixer::Full(fa) => {
1947 self.full_attn_prime(e, fa, &h, hx16.as_ref(), pos_d, t, cache, il, t)?
1948 }
1949 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
1950 Mixer::Linear(la) => {
1951 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
1952 let g4 = match hx16.as_ref() {
1953 Some(xh) => e.matmul_group_xh(&ws, &h, xh, t)?,
1954 None => e.matmul_group(&ws, &h, t)?,
1955 };
1956 self.linear_attn_prime_core_pad(e, la, g4, t, cache, il, Some(len_d))?
1957 }
1958 };
1959 let mut x1 = e.uninit(t * n_embd)?;
1960 e.add(&x, &mixed, &mut x1, t * n_embd)?;
1961 let mut z = e.uninit(t * n_embd)?;
1962 let mut zx16: Option<CudaSlice<u8>> = None;
1963 if f16fuse {
1964 let mut b16 = e.alloc_u8_uninit(t * n_embd * 2)?;
1965 e.rms_norm_f16out(
1966 &x1,
1967 layer.post_attn_norm.float_data(),
1968 &mut z,
1969 &mut b16,
1970 n_embd,
1971 t,
1972 eps,
1973 )?;
1974 zx16 = Some(b16);
1975 } else {
1976 e.rms_norm(
1977 &x1,
1978 layer.post_attn_norm.float_data(),
1979 &mut z,
1980 n_embd,
1981 t,
1982 eps,
1983 )?;
1984 }
1985 let ffn_out = match &layer.ffn {
1986 crate::hybrid::Ffn::Dense {
1987 ffn_gate,
1988 ffn_up,
1989 ffn_down,
1990 } => {
1991 let n_ff = ffn_gate.out_features();
1992 let mut g2 = match &zx16 {
1993 Some(xh) => e.matmul_group_xh(&[ffn_gate, ffn_up], &z, xh, t)?,
1994 None => e.matmul_group(&[ffn_gate, ffn_up], &z, t)?,
1995 };
1996 let up = g2.pop().unwrap();
1997 let gate = g2.pop().unwrap();
1998 let mut act = e.uninit(t * n_ff)?;
1999 Self::ffn_act_lim(
2001 e,
2002 &self.cfg,
2003 &gate,
2004 &up,
2005 1.0,
2006 1.0,
2007 self.cfg.clamp_shexp_at(il as u32),
2008 &mut act,
2009 t * n_ff,
2010 )?;
2011 e.matmul(ffn_down, &act, t)?
2012 }
2013 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il_prefill(e, m, &z, t, il as u16)?,
2014 };
2015 let mut x2 = e.uninit(t * n_embd)?;
2016 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
2017 x = x2;
2018 }
2019 if !crate::spec::spec_hpost() {
2021 e.row_gather_dev(&x, h_seed_out, len_d, n_embd)?;
2022 }
2023 let mut hn = e.uninit(t * n_embd)?;
2024 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
2025 if crate::spec::spec_hpost() {
2026 e.row_gather_dev(&hn, h_seed_out, len_d, n_embd)?;
2027 }
2028 let mut hlast = e.uninit(n_embd)?;
2029 e.row_gather_dev(&hn, &mut hlast, len_d, n_embd)?;
2030 let logits = e.matmul(&self.output, &hlast, 1)?;
2031 let nv = logits.len();
2032 e.copy_into(logits_out, 0, &logits, nv)?;
2033 Ok(())
2034 }
2035
2036 fn step35_prime_batch_on() -> bool {
2037 std::env::var("MEMRA_STEP35_PRIME_BATCH").as_deref() != Ok("0")
2038 }
2039
2040 #[allow(clippy::too_many_arguments)]
2043 fn step35_prime_batch_layers(
2044 &self,
2045 e: &Engine,
2046 mut x: CudaSlice<f32>,
2047 lo: usize,
2048 hi: usize,
2049 ts: &[usize],
2050 offs: &[usize],
2051 pos_ds: &[CudaSlice<i32>],
2052 caches: &mut [&mut Cache],
2053 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2054 let cfg = &self.cfg;
2055 let n_embd = cfg.n_embd as usize;
2056 let eps = cfg.rms_eps;
2057 let b = ts.len();
2058 let total: usize = ts.iter().sum();
2059 let f16fuse = crate::f16_ffi::pp_f16_enabled() && total >= 16;
2060
2061 let split = |e: &Engine,
2062 y: &CudaSlice<f32>,
2063 dim: usize|
2064 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
2065 let mut out = Vec::with_capacity(b);
2066 for s in 0..b {
2067 let mut ys = e.uninit(ts[s] * dim)?;
2068 e.copy_view_into(
2069 &mut ys,
2070 0,
2071 &y.slice(offs[s] * dim..(offs[s] + ts[s]) * dim),
2072 ts[s] * dim,
2073 )?;
2074 out.push(ys);
2075 }
2076 Ok(out)
2077 };
2078
2079 for il in lo..hi {
2080 let layer = &self.layers[il];
2081 let Mixer::Full(fa) = &layer.mixer else {
2082 return Err(format!("step35 layer {il} is not full-attn — corrupt config").into());
2083 };
2084
2085 let mut h = e.uninit(total * n_embd)?;
2086 let mut hx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
2087 if f16fuse {
2088 e.rms_norm_f16out(
2089 &x,
2090 layer.attn_norm.float_data(),
2091 &mut h,
2092 &mut hx16,
2093 n_embd,
2094 total,
2095 eps,
2096 )?;
2097 } else {
2098 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, total, eps)?;
2099 }
2100
2101 let gate_w = fa
2105 .attn_gate
2106 .as_ref()
2107 .ok_or("step35 layer is missing attn_gate.weight (head-wise attention gate)")?;
2108 let mut g4 = if f16fuse {
2109 e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv, gate_w], &h, &hx16, total)?
2110 } else {
2111 e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv, gate_w], &h, total)?
2112 };
2113 let gate = g4.pop().unwrap();
2114 let mut parts: Vec<Vec<CudaSlice<f32>>> =
2115 (0..b).map(|_| Vec::with_capacity(3)).collect();
2116 for (w, y) in [&fa.wq, &fa.wk, &fa.wv].iter().zip(g4) {
2117 for (s, ys) in split(e, &y, w.out_features())?.into_iter().enumerate() {
2118 parts[s].push(ys);
2119 }
2120 }
2121 let gates = split(e, &gate, gate_w.out_features())?;
2122 let geometry = self.step35_geom(il);
2123 let hd = geometry.head_dim_k as usize;
2124 let nh = geometry.n_head as usize;
2125 let mut ag_cat = e.uninit(total * nh * hd)?;
2126 for (s, (g3s, gate)) in parts.into_iter().zip(gates).enumerate() {
2127 let ag = self.step35_attn_pre_wo(
2128 e,
2129 fa,
2130 g3s,
2131 None,
2132 Some(&gate),
2133 &pos_ds[s],
2134 ts[s],
2135 Some(&mut *caches[s]),
2136 il,
2137 ts[s],
2138 )?;
2139 e.copy_into(&mut ag_cat, offs[s] * nh * hd, &ag, ts[s] * nh * hd)?;
2140 }
2141 let mixed = e.matmul(&fa.wo, &ag_cat, total)?;
2142
2143 let mut x1 = e.uninit(total * n_embd)?;
2144 let mut z = e.uninit(total * n_embd)?;
2145 let mut zx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
2146 if f16fuse {
2147 e.add_rms_norm_f16out(
2148 &x,
2149 &mixed,
2150 layer.post_attn_norm.float_data(),
2151 &mut x1,
2152 &mut z,
2153 &mut zx16,
2154 n_embd,
2155 total,
2156 eps,
2157 )?;
2158 } else {
2159 e.add(&x, &mixed, &mut x1, total * n_embd)?;
2160 e.rms_norm(
2161 &x1,
2162 layer.post_attn_norm.float_data(),
2163 &mut z,
2164 n_embd,
2165 total,
2166 eps,
2167 )?;
2168 }
2169
2170 let ffn_out = match &layer.ffn {
2171 crate::hybrid::Ffn::Dense {
2172 ffn_gate,
2173 ffn_up,
2174 ffn_down,
2175 } => {
2176 let n_ff = ffn_gate.out_features();
2177 let mut g2 = if f16fuse {
2178 e.matmul_group_xh(&[ffn_gate, ffn_up], &z, &zx16, total)?
2179 } else {
2180 e.matmul_group(&[ffn_gate, ffn_up], &z, total)?
2181 };
2182 let up = g2.pop().unwrap();
2183 let gate = g2.pop().unwrap();
2184 let mut act = e.uninit(total * n_ff)?;
2185 let d_lim = cfg.clamp_shexp_at(il as u32);
2186 if Self::f16out_on(e, total) && cfg.m3.is_none() && d_lim.is_none() {
2187 let mut a16 = e.alloc_u8_uninit(total * n_ff * 2)?;
2188 e.silu_mul_f16out(&gate, &up, &mut act, &mut a16, total * n_ff)?;
2189 match e.try_f16_gemm_pre(ffn_down, &a16, total)? {
2190 Some(y) => y,
2191 None => e.matmul(ffn_down, &act, total)?,
2192 }
2193 } else {
2194 Self::ffn_act_lim(
2195 e,
2196 cfg,
2197 &gate,
2198 &up,
2199 1.0,
2200 1.0,
2201 d_lim,
2202 &mut act,
2203 total * n_ff,
2204 )?;
2205 e.matmul(ffn_down, &act, total)?
2206 }
2207 }
2208 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il(e, m, &z, total, il as u16)?,
2209 };
2210 let mut x2 = e.uninit(total * n_embd)?;
2211 e.add(&x1, &ffn_out, &mut x2, total * n_embd)?;
2212 x = x2;
2213 }
2214 Ok(x)
2215 }
2216
2217 fn step35_prime_batch_epilogue(
2218 &self,
2219 e: &Engine,
2220 x: CudaSlice<f32>,
2221 ts: &[usize],
2222 offs: &[usize],
2223 caches: &mut [&mut Cache],
2224 ) -> Result<Vec<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
2225 let n_embd = self.cfg.n_embd as usize;
2226 let total: usize = ts.iter().sum();
2227 let mut hn = e.uninit(total * n_embd)?;
2228 e.rms_norm(
2229 &x,
2230 self.output_norm.float_data(),
2231 &mut hn,
2232 n_embd,
2233 total,
2234 self.cfg.rms_eps,
2235 )?;
2236
2237 let hidden_src = if crate::spec::spec_hpost() { &hn } else { &x };
2238 let mut out = Vec::with_capacity(ts.len());
2239 for s in 0..ts.len() {
2240 let mut hidden = e.uninit(ts[s] * n_embd)?;
2241 e.copy_view_into(
2242 &mut hidden,
2243 0,
2244 &hidden_src.slice(offs[s] * n_embd..(offs[s] + ts[s]) * n_embd),
2245 ts[s] * n_embd,
2246 )?;
2247 let last0 = (offs[s] + ts[s] - 1) * n_embd;
2248 let mut h_seed = e.uninit(n_embd)?;
2249 e.copy_view_into(
2250 &mut h_seed,
2251 0,
2252 &hidden_src.slice(last0..last0 + n_embd),
2253 n_embd,
2254 )?;
2255 let mut hlast = e.uninit(n_embd)?;
2257 e.copy_view_into(&mut hlast, 0, &hn.slice(last0..last0 + n_embd), n_embd)?;
2258 let logits = e.dtoh(&e.matmul(&self.output, &hlast, 1)?)?;
2259 caches[s].pos += ts[s];
2260 out.push((logits, h_seed, hidden));
2261 }
2262 Ok(out)
2263 }
2264
2265 fn step35_prime_cache_batch(
2266 &self,
2267 e: &Engine,
2268 prompts: &[&[u32]],
2269 caches: &mut [&mut Cache],
2270 ) -> Result<Vec<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
2271 if crate::pp::pp_host_bounce_active()
2272 && (!crate::pp::prime_pp_on() || crate::pp::pp_cuts(self.layers.len()).is_none())
2273 {
2274 return Err(
2275 "step35_prime_cache_batch: MEMRA_PP_HOST_BOUNCE=1 requires a valid prime \
2276 stage split; refusing an unsplit remote-weight walk"
2277 .into(),
2278 );
2279 }
2280 if !Self::step35_prime_batch_on() {
2281 return Err("step35 batched prime is disabled (MEMRA_STEP35_PRIME_BATCH=0)".into());
2282 }
2283 if caches.iter().any(|c| c.pos != 0) {
2284 return Err(
2285 "step35 batched prime currently supports complete fresh prompts only; \
2286 continuation/tick chunks require per-request queued_after"
2287 .into(),
2288 );
2289 }
2290
2291 let ts: Vec<usize> = prompts.iter().map(|p| p.len()).collect();
2292 for &t in &ts {
2293 assert!(
2294 t >= PRIME_MIN_T,
2295 "step35 batched prime needs T >= {PRIME_MIN_T}"
2296 );
2297 }
2298 for (s, c) in caches.iter().enumerate() {
2299 assert!(
2300 ts[s] <= c.max_ctx,
2301 "step35 batched prime exceeds cache max_ctx"
2302 );
2303 }
2304 let offs: Vec<usize> = ts
2305 .iter()
2306 .scan(0usize, |a, &t| {
2307 let o = *a;
2308 *a += t;
2309 Some(o)
2310 })
2311 .collect();
2312 let total: usize = ts.iter().sum();
2313 let payload = total * self.cfg.n_embd as usize;
2314 let cat_tokens: Vec<u32> = prompts.iter().flat_map(|p| p.iter().copied()).collect();
2315 let positions: Vec<Vec<i32>> = ts.iter().map(|&t| (0..t as i32).collect()).collect();
2316 let upload_positions =
2317 |e: &Engine| -> Result<Vec<CudaSlice<i32>>, Box<dyn std::error::Error>> {
2318 positions
2319 .iter()
2320 .map(|p| e.htod_i32(p))
2321 .collect::<Result<_, _>>()
2322 };
2323
2324 static ONCE: std::sync::Once = std::sync::Once::new();
2325 ONCE.call_once(|| {
2326 eprintln!(
2327 "[step35-prime-batch] first concat prime: B={} tokens={total}",
2328 prompts.len()
2329 );
2330 });
2331
2332 let out = if !crate::pp::pp2_streams_off() && crate::pp::prime_pp_on() {
2333 if let Some(fence) = crate::pp::pp_cuts(self.layers.len()) {
2334 let rt = crate::pp::PpNRt::get(e)?;
2335 let n_st = fence.len() - 1;
2336 assert_eq!(
2337 rt.n_stages(),
2338 n_st,
2339 "step35 prime batch stage count mismatch"
2340 );
2341 let caller_stream = e.stream();
2342 rt.fence_stages_behind(&caller_stream)?;
2343
2344 let mut slot = {
2345 let _st0 = rt.enter(0);
2346 let e0 = rt.engine(0, e);
2347 let pos_ds = upload_positions(e0)?;
2348 let x = self.embed(e0, &cat_tokens)?;
2349 let x = self.step35_prime_batch_layers(
2350 e0, x, fence[0], fence[1], &ts, &offs, &pos_ds, caches,
2351 )?;
2352 rt.tx(0, &x, payload)?
2353 };
2354 for s in 1..n_st - 1 {
2355 let _st = rt.enter(s);
2356 let es = rt.engine(s, e);
2357 let pos_ds = upload_positions(es)?;
2358 let x = rt.rx(s - 1, slot, payload)?;
2359 let x = self.step35_prime_batch_layers(
2360 es,
2361 x,
2362 fence[s],
2363 fence[s + 1],
2364 &ts,
2365 &offs,
2366 &pos_ds,
2367 caches,
2368 )?;
2369 slot = rt.tx(s, &x, payload)?;
2370 }
2371
2372 let _stl = rt.enter(n_st - 1);
2373 let el = rt.engine(n_st - 1, e);
2374 let pos_ds = upload_positions(el)?;
2375 let x = rt.rx(n_st - 2, slot, payload)?;
2376 let x = self.step35_prime_batch_layers(
2377 el,
2378 x,
2379 fence[n_st - 1],
2380 fence[n_st],
2381 &ts,
2382 &offs,
2383 &pos_ds,
2384 caches,
2385 )?;
2386 let out = self.step35_prime_batch_epilogue(el, x, &ts, &offs, caches)?;
2387 rt.publish_to(n_st - 1, &caller_stream)?;
2388 crate::pp::STEP35_PRIME_BATCH_SPLITS
2389 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2390 out
2391 } else {
2392 let pos_ds = upload_positions(e)?;
2393 let x = self.embed(e, &cat_tokens)?;
2394 let x = self.step35_prime_batch_layers(
2395 e,
2396 x,
2397 0,
2398 self.layers.len(),
2399 &ts,
2400 &offs,
2401 &pos_ds,
2402 caches,
2403 )?;
2404 self.step35_prime_batch_epilogue(e, x, &ts, &offs, caches)?
2405 }
2406 } else {
2407 let pos_ds = upload_positions(e)?;
2408 let x = self.embed(e, &cat_tokens)?;
2409 let x = self.step35_prime_batch_layers(
2410 e,
2411 x,
2412 0,
2413 self.layers.len(),
2414 &ts,
2415 &offs,
2416 &pos_ds,
2417 caches,
2418 )?;
2419 self.step35_prime_batch_epilogue(e, x, &ts, &offs, caches)?
2420 };
2421 crate::pp::STEP35_PRIME_BATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2422 Ok(out)
2423 }
2424
2425 pub fn prime_cache_batch(
2442 &self,
2443 e: &Engine,
2444 prompts: &[&[u32]],
2445 caches: &mut [&mut Cache],
2446 ) -> Result<Vec<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
2447 let cfg = &self.cfg;
2448 let n_embd = cfg.n_embd as usize;
2449 let eps = cfg.rms_eps;
2450 let b = prompts.len();
2451 assert!(b >= 1 && b == caches.len());
2452 let pos0s: Vec<usize> = caches.iter().map(|c| c.pos).collect();
2453 let carried = pos0s.iter().any(|&p| p > 0);
2454 if cfg.gemma4.is_some() {
2460 return Err(
2461 "prime_cache_batch: gemma4 has no batched prime core (per-layer \
2462 swa/global geometry, softcapped head) — use gemma4_prime per sequence"
2463 .into(),
2464 );
2465 }
2466 if cfg.step35.is_some() {
2469 return self.step35_prime_cache_batch(e, prompts, caches);
2470 }
2471 let ts: Vec<usize> = prompts.iter().map(|p| p.len()).collect();
2472 for &t in &ts {
2473 assert!(
2474 t >= PRIME_MIN_T,
2475 "prime_cache_batch needs T >= {PRIME_MIN_T}"
2476 );
2477 }
2478 for (s, c) in caches.iter().enumerate() {
2479 assert!(
2480 c.pos + ts[s] <= c.max_ctx,
2481 "prime_cache_batch: prompt exceeds cache max_ctx"
2482 );
2483 }
2484 let total: usize = ts.iter().sum();
2485 let offs: Vec<usize> = ts
2486 .iter()
2487 .scan(0usize, |a, &t| {
2488 let o = *a;
2489 *a += t;
2490 Some(o)
2491 })
2492 .collect();
2493 let pos_ds: Vec<CudaSlice<i32>> = ts
2495 .iter()
2496 .zip(&pos0s)
2497 .map(|(&t, &p0)| e.htod_i32(&(p0 as i32..(p0 + t) as i32).collect::<Vec<_>>()))
2498 .collect::<Result<_, _>>()?;
2499 let split = |e: &Engine,
2501 y: &CudaSlice<f32>,
2502 dim: usize|
2503 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
2504 let mut out = Vec::with_capacity(b);
2505 for s in 0..b {
2506 let mut ys = e.uninit(ts[s] * dim)?;
2507 e.copy_view_into(
2508 &mut ys,
2509 0,
2510 &y.slice(offs[s] * dim..(offs[s] + ts[s]) * dim),
2511 ts[s] * dim,
2512 )?;
2513 out.push(ys);
2514 }
2515 Ok(out)
2516 };
2517
2518 let cat_tokens: Vec<u32> = prompts.iter().flat_map(|p| p.iter().copied()).collect();
2519 let mut x = self.embed(e, &cat_tokens)?; for (il, layer) in self.layers.iter().enumerate() {
2521 let mut h = e.uninit(total * n_embd)?;
2522 let mut hx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
2523 e.rms_norm_f16out(
2524 &x,
2525 layer.attn_norm.float_data(),
2526 &mut h,
2527 &mut hx16,
2528 n_embd,
2529 total,
2530 eps,
2531 )?;
2532 let mut mixed = e.uninit(total * n_embd)?;
2534 match &layer.mixer {
2535 Mixer::Full(fa) => {
2536 let g3 = e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], &h, &hx16, total)?;
2537 let geometry = self.cfg.full_attention_geometry_at(il as u32);
2543 let (n_head, n_head_kv, head_dim) = (
2544 geometry.n_head as usize,
2545 geometry.n_head_kv as usize,
2546 geometry.head_dim_k as usize,
2547 );
2548 let fa_scale = geometry.attention_scale();
2549 let use_favl = !carried
2550 && (2..=8).contains(&b)
2551 && (head_dim == 256 || head_dim == 128)
2552 && geometry.attention_gate == memra_gguf::config::AttentionGateKind::FusedQ
2553 && std::env::var("MEMRA_NOFA").is_err()
2554 && std::env::var("MEMRA_FA_FLOOR").is_err()
2555 && std::env::var("MEMRA_FA_PP_W2").as_deref() != Ok("1")
2556 && std::env::var("MEMRA_FA_BF16KV").as_deref() != Ok("0")
2557 && std::env::var("MEMRA_FA_VL").as_deref() != Ok("0");
2558 if use_favl {
2559 let (qf_w, kf_w, vf_w) = (
2560 fa.wq.out_features(),
2561 fa.wk.out_features(),
2562 fa.wv.out_features(),
2563 );
2564 struct APre {
2565 q: CudaSlice<f32>,
2566 gate: Option<CudaSlice<f32>>,
2567 qn: CudaSlice<f32>,
2568 kn: CudaSlice<f32>,
2569 }
2570 let mut aps = Vec::with_capacity(b);
2571 for &t in ts.iter().take(b) {
2572 aps.push(APre {
2573 q: e.uninit(t * n_head * head_dim)?,
2574 gate: Some(e.uninit(t * n_head * head_dim)?),
2575 qn: e.uninit(t * n_head * head_dim)?,
2576 kn: e.uninit(t * n_head_kv * head_dim)?,
2577 });
2578 }
2579 let (kv_dim_k, kv_dim_v, ktb, vtb) = {
2580 let kvl = caches[0].kv[il].as_ref().unwrap();
2581 (kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes)
2582 };
2583 let pargs: Vec<crate::AttnPreVl> = (0..b)
2584 .map(|s| {
2585 let (o, t) = (offs[s], ts[s]);
2586 let kvl = caches[s].kv[il].as_ref().unwrap();
2587 assert!(
2588 kvl.len == 0 && kvl.len + t <= caches[s].max_ctx,
2589 "prime_cache_batch attn vl: fresh + capacity"
2590 );
2591 crate::AttnPreVl {
2592 qf: e.addr_f32v(&g3[0].slice(o * qf_w..(o + t) * qf_w)),
2593 kf: e.addr_f32v(&g3[1].slice(o * kf_w..(o + t) * kf_w)),
2594 vf: e.addr_f32v(&g3[2].slice(o * vf_w..(o + t) * vf_w)),
2595 q: e.addr_f32(&aps[s].q),
2596 gate: e.addr_f32(aps[s].gate.as_ref().unwrap()),
2597 qn: e.addr_f32(&aps[s].qn),
2598 kn: e.addr_f32(&aps[s].kn),
2599 kc: e.addr_u8(&kvl.k),
2600 vc: e.addr_u8(&kvl.v),
2601 t: t as i32,
2602 pad: 0,
2603 }
2604 })
2605 .collect();
2606 e.attn_pre_vl8(
2607 &pargs,
2608 fa.q_norm.float_data(),
2609 fa.k_norm.float_data(),
2610 head_dim,
2611 geometry.n_rot as usize,
2612 n_head,
2613 n_head_kv,
2614 self.cfg.rms_eps,
2615 geometry.rope_base,
2616 1.0,
2617 kv_dim_k,
2618 kv_dim_v,
2619 ktb,
2620 vtb,
2621 )?;
2622 for s in 0..b {
2623 let kvl = caches[s].kv[il].as_mut().unwrap();
2624 kvl.len += ts[s];
2625 let new_len = kvl.len as i32;
2626 e.set_i32_one(&mut kvl.len_d, new_len)?;
2627 }
2628 let mut attns = Vec::with_capacity(b);
2629 let mut mirrors = Vec::with_capacity(b);
2630 for &t in ts.iter().take(b) {
2631 attns.push(e.uninit(t * n_head * head_dim)?);
2632 let n = t * n_head_kv * head_dim;
2633 mirrors.push((e.alloc_u8_uninit(n * 2)?, e.alloc_u8_uninit(n * 2)?));
2634 }
2635 let fa3_on = match std::env::var("MEMRA_FA3").as_deref() {
2638 Ok("0") => false,
2639 Ok("1") => true,
2640 _ => cfg!(memra_hopper_mma),
2641 };
2642 if fa3_on {
2643 let mut q16s = Vec::with_capacity(b);
2644 let mut v16s = Vec::with_capacity(b);
2645 for s in 0..b {
2646 let t = ts[s];
2647 let mut q16 = e.alloc_u8_uninit(t * n_head * head_dim * 2)?;
2648 e.f32_to_bf16_into(&aps[s].qn, &mut q16, t * n_head * head_dim)?;
2649 let mut k16 = e.alloc_u8_uninit(t * n_head_kv * head_dim * 2)?;
2650 e.f32_to_bf16_into(&aps[s].kn, &mut k16, t * n_head_kv * head_dim)?;
2651 let mut v16 = e.alloc_u8_uninit(t * n_head_kv * head_dim * 2)?;
2652 e.f32_to_bf16_v(
2653 &g3[2].slice(offs[s] * vf_w..(offs[s] + t) * vf_w),
2654 &mut v16,
2655 t * n_head_kv * head_dim,
2656 )?;
2657 q16s.push(q16);
2658 v16s.push((k16, v16));
2659 }
2660 let mut qp = [core::ptr::null::<core::ffi::c_void>(); 8];
2661 let mut kp = qp;
2662 let mut vp = qp;
2663 let mut op = [core::ptr::null_mut::<f32>(); 8];
2664 let mut tsv = [0i32; 8];
2665 for s in 0..b {
2666 qp[s] = e.addr_u8(&q16s[s]) as *const core::ffi::c_void;
2667 kp[s] = e.addr_u8(&v16s[s].0) as *const core::ffi::c_void;
2668 vp[s] = e.addr_u8(&v16s[s].1) as *const core::ffi::c_void;
2669 op[s] = e.addr_f32(&attns[s]) as *mut f32;
2670 tsv[s] = ts[s] as i32;
2671 }
2672 let rc = unsafe {
2673 crate::fa3_vl_raw(
2674 qp.as_ptr(),
2675 kp.as_ptr(),
2676 vp.as_ptr(),
2677 op.as_ptr(),
2678 tsv.as_ptr(),
2679 b as i32,
2680 n_head as i32,
2681 n_head_kv as i32,
2682 head_dim as i32,
2683 fa_scale,
2684 e.stream().cu_stream() as *mut core::ffi::c_void,
2685 )
2686 };
2687 if rc != 0 {
2688 return Err(format!("memra_fa3_vl rc={rc}").into());
2689 }
2690 } else {
2691 let fargs: Vec<crate::FaSeqVl> = (0..b)
2692 .map(|s| crate::FaSeqVl {
2693 q: e.addr_f32(&aps[s].qn),
2694 k16: e.addr_u8(&mirrors[s].0),
2695 v16: e.addr_u8(&mirrors[s].1),
2696 o: e.addr_f32(&attns[s]),
2697 kf: e.addr_f32(&aps[s].kn),
2698 vf: e.addr_f32v(
2699 &g3[2].slice(offs[s] * vf_w..(offs[s] + ts[s]) * vf_w),
2700 ),
2701 t: ts[s] as i32,
2702 pad: 0,
2703 })
2704 .collect();
2705 e.fa_prefill_vl8(&fargs, head_dim, n_head, n_head_kv, fa_scale)?;
2706 }
2707 for (s, attn) in attns.into_iter().enumerate() {
2708 let (attn_g, ag16) = self.full_attn_prime_post_fa(
2709 e,
2710 attn,
2711 &aps[s].gate,
2712 ts[s],
2713 n_head,
2714 head_dim,
2715 )?;
2716 let mut done = false;
2717 if let Some(xh) = &ag16 {
2718 done = e.try_f16_gemm_pre_into_off(
2719 &fa.wo,
2720 xh,
2721 ts[s],
2722 &mut mixed,
2723 offs[s] * n_embd,
2724 )?;
2725 }
2726 if !done {
2727 let m = e.matmul(&fa.wo, &attn_g, ts[s])?;
2728 e.copy_into(&mut mixed, offs[s] * n_embd, &m, ts[s] * n_embd)?;
2729 }
2730 }
2731 } else {
2732 let mut parts: Vec<Vec<CudaSlice<f32>>> =
2733 (0..b).map(|_| Vec::new()).collect();
2734 for (w, y) in [&fa.wq, &fa.wk, &fa.wv].iter().zip(g3) {
2735 for (s, ys) in split(e, &y, w.out_features())?.into_iter().enumerate() {
2736 parts[s].push(ys);
2737 }
2738 }
2739 for (s, g3s) in parts.into_iter().enumerate() {
2740 let (attn_g, ag16) = self.full_attn_prime_core_inner(
2742 e, fa, g3s, &pos_ds[s], ts[s], caches[s], il,
2743 )?;
2744 let mut done = false;
2745 if let Some(xh) = &ag16 {
2746 done = e.try_f16_gemm_pre_into_off(
2747 &fa.wo,
2748 xh,
2749 ts[s],
2750 &mut mixed,
2751 offs[s] * n_embd,
2752 )?;
2753 }
2754 if !done {
2755 let m = e.matmul(&fa.wo, &attn_g, ts[s])?;
2756 e.copy_into(&mut mixed, offs[s] * n_embd, &m, ts[s] * n_embd)?;
2757 }
2758 }
2759 }
2760 }
2761 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
2762 Mixer::Linear(la) => {
2763 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
2768 let g4 = e.matmul_group_xh(&ws, &h, &hx16, total)?;
2769 let outs =
2770 self.linear_attn_prime_core_batch(e, la, &g4, &offs, &ts, caches, il)?;
2771 for (s, (gn, gn16)) in outs.into_iter().enumerate() {
2772 let (o, t) = (offs[s], ts[s]);
2773 let mut done = false;
2774 if let Some(xh) = &gn16 {
2775 done = e.try_f16_gemm_pre_into_off(
2776 &la.ssm_out,
2777 xh,
2778 t,
2779 &mut mixed,
2780 o * n_embd,
2781 )?;
2782 }
2783 if !done {
2784 let m = e.matmul(&la.ssm_out, &gn, t)?;
2785 e.copy_into(&mut mixed, o * n_embd, &m, t * n_embd)?;
2786 }
2787 }
2788 }
2789 }
2790 let mut x1 = e.uninit(total * n_embd)?;
2791 let mut z = e.uninit(total * n_embd)?;
2792 let mut zx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
2793 e.add_rms_norm_f16out(
2794 &x,
2795 &mixed,
2796 layer.post_attn_norm.float_data(),
2797 &mut x1,
2798 &mut z,
2799 &mut zx16,
2800 n_embd,
2801 total,
2802 eps,
2803 )?;
2804 let ffn_out = match &layer.ffn {
2805 crate::hybrid::Ffn::Dense {
2806 ffn_gate,
2807 ffn_up,
2808 ffn_down,
2809 } => {
2810 let n_ff = ffn_gate.out_features();
2811 let mut g2 = e.matmul_group_xh(&[ffn_gate, ffn_up], &z, &zx16, total)?;
2812 let up = g2.pop().unwrap();
2813 let gate = g2.pop().unwrap();
2814 let mut act = e.uninit(total * n_ff)?;
2815 let d_lim = self.cfg.clamp_shexp_at(il as u32);
2819 if Self::f16out_on(e, total) && self.cfg.m3.is_none() && d_lim.is_none() {
2820 let mut a16 = e.alloc_u8_uninit(total * n_ff * 2)?;
2821 e.silu_mul_f16out(&gate, &up, &mut act, &mut a16, total * n_ff)?;
2822 match e.try_f16_gemm_pre(ffn_down, &a16, total)? {
2823 Some(y) => y,
2824 None => e.matmul(ffn_down, &act, total)?,
2825 }
2826 } else {
2827 Self::ffn_act_lim(
2828 e,
2829 &self.cfg,
2830 &gate,
2831 &up,
2832 1.0,
2833 1.0,
2834 d_lim,
2835 &mut act,
2836 total * n_ff,
2837 )?;
2838 e.matmul(ffn_down, &act, total)?
2839 }
2840 }
2841 crate::hybrid::Ffn::Moe(m) => {
2842 self.moe_ffn_il_prefill(e, m, &z, total, il as u16)?
2843 }
2844 };
2845 let mut x2 = e.uninit(total * n_embd)?;
2846 e.add(&x1, &ffn_out, &mut x2, total * n_embd)?;
2847 x = x2;
2848 }
2849 let mut hn = e.uninit(total * n_embd)?;
2851 e.rms_norm(
2852 &x,
2853 self.output_norm.float_data(),
2854 &mut hn,
2855 n_embd,
2856 total,
2857 eps,
2858 )?;
2859 let mut hcat = e.uninit(b * n_embd)?;
2865 for s in 0..b {
2866 let last0 = (offs[s] + ts[s] - 1) * n_embd;
2867 e.copy_view_into(
2868 &mut hcat,
2869 s * n_embd,
2870 &hn.slice(last0..last0 + n_embd),
2871 n_embd,
2872 )?;
2873 }
2874 let logits_cat = if b >= 2 {
2875 e.try_f16_gemm(&self.output, &hcat, b)?
2876 } else {
2877 None
2878 };
2879 let logits_host: Option<Vec<f32>> = match &logits_cat {
2880 Some(lc) => Some(e.dtoh(lc)?),
2881 None => None,
2882 };
2883 let n_vocab = self.output.out_features();
2884 let mut hidden_all = if crate::spec::spec_hpost() {
2885 split(e, &hn, n_embd)?
2886 } else {
2887 split(e, &x, n_embd)?
2888 };
2889 let mut out = Vec::with_capacity(b);
2890 for s in 0..b {
2891 let last0 = (offs[s] + ts[s] - 1) * n_embd;
2892 let mut h_seed = e.uninit(n_embd)?;
2893 if !crate::spec::spec_hpost() {
2894 e.copy_view_into(&mut h_seed, 0, &x.slice(last0..last0 + n_embd), n_embd)?;
2895 } else {
2896 e.copy_view_into(&mut h_seed, 0, &hn.slice(last0..last0 + n_embd), n_embd)?;
2897 }
2898 let logits = match &logits_host {
2899 Some(lh) => lh[s * n_vocab..(s + 1) * n_vocab].to_vec(),
2900 None => {
2901 let mut hlast = e.uninit(n_embd)?;
2902 e.copy_view_into(&mut hlast, 0, &hn.slice(last0..last0 + n_embd), n_embd)?;
2903 e.dtoh(&e.matmul(&self.output, &hlast, 1)?)?
2904 }
2905 };
2906 caches[s].pos += ts[s];
2907 out.push((logits, h_seed, hidden_all.remove(0)));
2908 }
2909 Ok(out)
2910 }
2911
2912 #[allow(clippy::too_many_arguments)]
2923 fn full_attn_prime(
2924 &self,
2925 e: &Engine,
2926 fa: &FullAttnLayer,
2927 h: &CudaSlice<f32>,
2928 hx: Option<&CudaSlice<u8>>,
2929 pos_d: &CudaSlice<i32>,
2930 t: usize,
2931 cache: &mut Cache,
2932 il: usize,
2933 seq_end: usize,
2934 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2935 if self.cfg.step35.is_some() {
2936 return self.step35_attn_prime(e, fa, h, hx, pos_d, t, cache, il, seq_end);
2937 }
2938 let g3 = match hx {
2943 Some(xh) => e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], h, xh, t)?,
2944 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
2945 };
2946 self.full_attn_prime_core(e, fa, g3, pos_d, t, cache, il)
2947 }
2948
2949 fn full_attn_prime_core(
2953 &self,
2954 e: &Engine,
2955 fa: &FullAttnLayer,
2956 g3: Vec<CudaSlice<f32>>,
2957 pos_d: &CudaSlice<i32>,
2958 t: usize,
2959 cache: &mut Cache,
2960 il: usize,
2961 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2962 let (attn_g, ag16) = self.full_attn_prime_core_inner(e, fa, g3, pos_d, t, cache, il)?;
2963 if let Some(xh) = &ag16 {
2964 if let Some(y) = e.try_f16_gemm_pre(&fa.wo, xh, t)? {
2965 return Ok(y);
2966 }
2967 }
2968 Ok(e.matmul(&fa.wo, &attn_g, t)?)
2969 }
2970
2971 fn full_attn_prime_core_inner(
2972 &self,
2973 e: &Engine,
2974 fa: &FullAttnLayer,
2975 g3: Vec<CudaSlice<f32>>,
2976 pos_d: &CudaSlice<i32>,
2977 t: usize,
2978 cache: &mut Cache,
2979 il: usize,
2980 ) -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
2981 let cfg = &self.cfg;
2982 let geometry = cfg.full_attention_geometry_at(il as u32);
2983 let n_head = geometry.n_head as usize;
2984 let n_head_kv = geometry.n_head_kv as usize;
2985 let head_dim = geometry.head_dim_k as usize;
2986 let scale = geometry.attention_scale();
2987 let (pre, base_len) = self.full_attn_prime_pre_fa(e, fa, g3, pos_d, t, cache, il)?;
2988 let AttnPre { q, k, v, gate } = pre;
2989 let mut attn = e.uninit(t * n_head * head_dim)?;
2990 self.full_attn_prime_fa_dispatch(
2991 e, &q, &k, &v, &mut attn, base_len, t, cache, il, head_dim, n_head, n_head_kv, scale,
2992 )?;
2993 self.full_attn_prime_post_fa(e, attn, &gate, t, n_head, head_dim)
2994 }
2995
2996 #[allow(clippy::type_complexity)]
3000 fn full_attn_prime_pre_fa(
3001 &self,
3002 e: &Engine,
3003 fa: &FullAttnLayer,
3004 mut g3: Vec<CudaSlice<f32>>,
3005 pos_d: &CudaSlice<i32>,
3006 t: usize,
3007 cache: &mut Cache,
3008 il: usize,
3009 ) -> Result<(AttnPre, usize), Box<dyn std::error::Error>> {
3010 let cfg = &self.cfg;
3011 let geometry = cfg.full_attention_geometry_at(il as u32);
3012 let n_head = geometry.n_head as usize;
3013 let n_head_kv = geometry.n_head_kv as usize;
3014 let head_dim = geometry.head_dim_k as usize;
3015 let eps = cfg.rms_eps;
3016
3017 let gated = geometry.attention_gate == memra_gguf::config::AttentionGateKind::FusedQ;
3021 let v = g3.pop().unwrap();
3022 let mut k = g3.pop().unwrap();
3023 let qf = g3.pop().unwrap();
3024 let (mut q, gate) = if gated {
3025 let mut q = e.uninit(t * n_head * head_dim)?;
3026 let mut gate = e.uninit(t * n_head * head_dim)?;
3027 e.q_gate_split(&qf, &mut q, &mut gate, head_dim, n_head, t)?;
3028 (q, Some(gate))
3029 } else {
3030 (qf, None)
3031 };
3032
3033 let mut qn = e.uninit(t * n_head * head_dim)?;
3034 e.rms_norm(
3035 &q,
3036 fa.q_norm.float_data(),
3037 &mut qn,
3038 head_dim,
3039 n_head * t,
3040 eps,
3041 )?;
3042 q = qn;
3043 let mut kn = e.uninit(t * n_head_kv * head_dim)?;
3044 e.rms_norm(
3045 &k,
3046 fa.k_norm.float_data(),
3047 &mut kn,
3048 head_dim,
3049 n_head_kv * t,
3050 eps,
3051 )?;
3052 k = kn;
3053 let rope_dims = geometry.n_rot as usize;
3054 e.rope_neox(
3055 &mut q,
3056 pos_d,
3057 head_dim,
3058 rope_dims,
3059 n_head,
3060 t,
3061 geometry.rope_base,
3062 1.0,
3063 )?;
3064 e.rope_neox(
3065 &mut k,
3066 pos_d,
3067 head_dim,
3068 rope_dims,
3069 n_head_kv,
3070 t,
3071 geometry.rope_base,
3072 1.0,
3073 )?;
3074
3075 {
3078 let kvl = cache.kv[il].as_mut().unwrap();
3079 assert!(kvl.len + t <= cache.max_ctx, "prime_cache: KV overflow");
3080 e.append_kv_quantized_rows(
3081 &k,
3082 &v,
3083 &mut kvl.k,
3084 &mut kvl.v,
3085 kvl.len,
3086 t,
3087 kvl.kv_dim_k,
3088 kvl.kv_dim_v,
3089 kvl.k_tok_bytes,
3090 kvl.v_tok_bytes,
3091 crate::Engine::kv_fp8_on(),
3092 )?;
3093 kvl.len += t;
3094 let new_len = kvl.len as i32;
3095 e.set_i32_one(&mut kvl.len_d, new_len)?;
3096 }
3097
3098 let base_len = {
3099 let kvl = cache.kv[il].as_ref().unwrap();
3100 kvl.len - t };
3102 Ok((AttnPre { q, k, v, gate }, base_len))
3103 }
3104
3105 #[allow(clippy::too_many_arguments)]
3112 fn full_attn_prime_fa_dispatch(
3113 &self,
3114 e: &Engine,
3115 q: &CudaSlice<f32>,
3116 k: &CudaSlice<f32>,
3117 v: &CudaSlice<f32>,
3118 attn: &mut CudaSlice<f32>,
3119 base_len: usize,
3120 t: usize,
3121 cache: &mut Cache,
3122 il: usize,
3123 head_dim: usize,
3124 n_head: usize,
3125 n_head_kv: usize,
3126 scale: f32,
3127 ) -> Result<(), Box<dyn std::error::Error>> {
3128 if base_len == 0 && std::env::var("MEMRA_PRIME_F32CHUNK0").as_deref() == Ok("1") {
3141 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
3142 e.sdpa_naive(
3143 q, k, v, attn, head_dim, n_head, n_head_kv, t, t, scale, true,
3144 )?;
3145 } else {
3146 e.fa_prefill(
3147 q, k, v, attn, head_dim, n_head, n_head_kv, t, t, scale, true,
3148 )?;
3149 }
3150 return Ok(());
3151 }
3152 let kvl = cache.kv[il].as_ref().unwrap();
3153 let t_kv = base_len + t;
3154 let k_view = e.view_u8(&kvl.k, t_kv * kvl.k_tok_bytes);
3155 let v_view = e.view_u8(&kvl.v, t_kv * kvl.v_tok_bytes);
3156 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
3160 e.sdpa_naive_quantized_view(
3161 q,
3162 &k_view,
3163 &v_view,
3164 attn,
3165 head_dim,
3166 n_head,
3167 n_head_kv,
3168 t,
3169 t_kv,
3170 scale,
3171 true,
3172 kvl.k_tok_bytes,
3173 kvl.v_tok_bytes,
3174 )?;
3175 return Ok(());
3176 }
3177 let deqw = std::env::var("MEMRA_PRIME_DEQW")
3185 .map(|v| v != "0")
3186 .unwrap_or(true);
3187 if deqw {
3188 e.fa_prefill_view_ws(
3189 q,
3190 &k_view,
3191 &v_view,
3192 attn,
3193 head_dim,
3194 n_head,
3195 n_head_kv,
3196 t,
3197 t_kv,
3198 scale,
3199 true,
3200 kvl.k_tok_bytes,
3201 kvl.v_tok_bytes,
3202 crate::Engine::kv_fp8_on(),
3203 )?;
3204 } else {
3205 e.fa_prefill_view(
3206 q,
3207 &k_view,
3208 &v_view,
3209 attn,
3210 head_dim,
3211 n_head,
3212 n_head_kv,
3213 t,
3214 t_kv,
3215 scale,
3216 true,
3217 kvl.k_tok_bytes,
3218 kvl.v_tok_bytes,
3219 crate::Engine::kv_fp8_on(),
3220 )?;
3221 }
3222 Ok(())
3223 }
3224
3225 fn full_attn_prime_post_fa(
3228 &self,
3229 e: &Engine,
3230 attn: CudaSlice<f32>,
3231 gate: &Option<CudaSlice<f32>>,
3232 t: usize,
3233 n_head: usize,
3234 head_dim: usize,
3235 ) -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
3236 let (attn_g, ag16) = match gate {
3237 Some(gate) => {
3238 let n = t * n_head * head_dim;
3239 let mut ag = e.uninit(n)?;
3240 if Self::f16out_on(e, t) {
3241 let mut a16 = e.alloc_u8_uninit(n * 2)?;
3242 e.sig_mul_f16out(&attn, gate, &mut ag, &mut a16, n)?;
3243 (ag, Some(a16))
3244 } else {
3245 let mut gsig = e.uninit(n)?;
3246 e.sigmoid(gate, &mut gsig, n)?;
3247 e.mul(&attn, &gsig, &mut ag, n)?;
3248 (ag, None)
3249 }
3250 }
3251 None => (attn, None),
3252 };
3253 Ok((attn_g, ag16))
3254 }
3255
3256 fn linear_attn_prime(
3263 &self,
3264 e: &Engine,
3265 la: &LinearAttnLayer,
3266 h: &CudaSlice<f32>,
3267 hx: Option<&CudaSlice<u8>>,
3268 t: usize,
3269 cache: &mut Cache,
3270 il: usize,
3271 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3272 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
3274 let g4 = match hx {
3275 Some(xh) => e.matmul_group_xh(&ws, h, xh, t)?,
3276 None => e.matmul_group(&ws, h, t)?,
3277 };
3278 self.linear_attn_prime_core(e, la, g4, t, cache, il)
3279 }
3280
3281 fn linear_attn_prime_core(
3283 &self,
3284 e: &Engine,
3285 la: &LinearAttnLayer,
3286 mut g4: Vec<CudaSlice<f32>>,
3287 t: usize,
3288 cache: &mut Cache,
3289 il: usize,
3290 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3291 self.linear_attn_prime_core_pad(e, la, g4.drain(..).collect(), t, cache, il, None)
3292 }
3293
3294 #[allow(clippy::too_many_arguments)]
3298 fn linear_attn_prime_core_pad_inner(
3299 &self,
3300 e: &Engine,
3301 la: &LinearAttnLayer,
3302 mut g4: Vec<CudaSlice<f32>>,
3303 t: usize,
3304 cache: &mut Cache,
3305 il: usize,
3306 pad_len: Option<&CudaSlice<i32>>,
3307 ) -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
3308 let ssm = self.cfg.ssm.as_ref().unwrap();
3310 let d_state = ssm.state_size as usize;
3311 let num_k = ssm.group_count as usize;
3312 let num_v = ssm.time_step_rank as usize;
3313 let key_dim = d_state * num_k;
3314 let value_dim = d_state * num_v;
3315 let conv_dim = key_dim * 2 + value_dim;
3316 let alpha = g4.pop().unwrap(); let beta_raw = g4.pop().unwrap(); let z = g4.pop().unwrap(); let qkv_mixed = g4.pop().unwrap(); self.linear_attn_prime_core_pad_view(
3321 e,
3322 la,
3323 &qkv_mixed.slice(0..t * conv_dim),
3324 &z.slice(0..t * value_dim),
3325 &beta_raw.slice(0..t * num_v),
3326 &alpha.slice(0..t * num_v),
3327 t,
3328 cache,
3329 il,
3330 pad_len,
3331 )
3332 }
3333
3334 #[allow(clippy::too_many_arguments)]
3337 fn linear_attn_gdn_prep(
3338 &self,
3339 e: &Engine,
3340 la: &LinearAttnLayer,
3341 qkv_mixed: &cudarc::driver::CudaView<f32>,
3342 beta_raw: &cudarc::driver::CudaView<f32>,
3343 alpha: &cudarc::driver::CudaView<f32>,
3344 t: usize,
3345 cache: &mut Cache,
3346 il: usize,
3347 pad_len: Option<&CudaSlice<i32>>,
3348 ) -> Result<GdnPrep, Box<dyn std::error::Error>> {
3349 let cfg = &self.cfg;
3350 let ssm = cfg.ssm.as_ref().unwrap();
3351 let d_state = ssm.state_size as usize; let num_k = ssm.group_count as usize; let num_v = ssm.time_step_rank as usize; let d_conv = ssm.conv_kernel as usize; let key_dim = d_state * num_k; let value_dim = d_state * num_v; let conv_dim = key_dim * 2 + value_dim; let eps = cfg.rms_eps;
3359 debug_assert!(
3360 t >= d_conv - 1,
3361 "stateful conv needs T >= pad (PRIME_MIN_T gates)"
3362 );
3363
3364 let rl = cache.recur[il].as_mut().unwrap();
3369 let hk = Self::gdn_hk(e, t, num_v, num_k);
3370 let conv_fuse = std::env::var("MEMRA_CONV_FUSE").as_deref() != Ok("0");
3371 let hk = if conv_fuse { hk } else { num_v }; let mut q_g = e.uninit(d_state * hk * t)?;
3373 let mut k_g = e.uninit(d_state * hk * t)?;
3374 let mut v_g = e.uninit(d_state * num_v * t)?;
3375 if conv_fuse {
3376 e.ssm_conv1d_gdn_state_pad(
3377 qkv_mixed,
3378 &mut rl.conv_state,
3379 la.ssm_conv1d.float_data(),
3380 &mut q_g,
3381 &mut k_g,
3382 &mut v_g,
3383 conv_dim,
3384 t,
3385 d_conv,
3386 d_state,
3387 num_v,
3388 num_k,
3389 key_dim,
3390 hk,
3391 pad_len,
3392 )?;
3393 } else {
3394 let mut conv_out = e.uninit(conv_dim * t)?; e.ssm_conv1d_tm_state_pad_v(
3396 qkv_mixed,
3397 &mut rl.conv_state,
3398 la.ssm_conv1d.float_data(),
3399 &mut conv_out,
3400 conv_dim,
3401 t,
3402 d_conv,
3403 pad_len,
3404 )?;
3405 e.qkv_to_gdn_repack(
3406 &conv_out, &mut q_g, &mut k_g, &mut v_g, d_state, num_v, num_k, key_dim, t,
3407 )?;
3408 }
3409 let mut q_l2 = e.uninit(d_state * hk * t)?;
3410 let qb16 = if Engine::l2_v2_on(d_state) && e.gdn_wgmma_on(32) {
3414 let mut qb = e.alloc_u8_uninit(d_state * hk * t * 2)?;
3415 e.l2_norm_pp(&q_g, &mut q_l2, Some(&mut qb), d_state, hk * t, eps)?;
3416 Some(qb)
3417 } else {
3418 e.l2_norm_pp(&q_g, &mut q_l2, None, d_state, hk * t, eps)?;
3419 None
3420 };
3421 let mut k_l2 = e.uninit(d_state * hk * t)?;
3422 let kb16 = if Engine::l2_v2_on(d_state) {
3424 let mut kb = e.alloc_u8_uninit(d_state * hk * t * 2)?;
3425 e.l2_norm_pp(&k_g, &mut k_l2, Some(&mut kb), d_state, hk * t, eps)?;
3426 Some(kb)
3427 } else {
3428 e.l2_norm_pp(&k_g, &mut k_l2, None, d_state, hk * t, eps)?;
3429 None
3430 };
3431 let mut beta = e.uninit(t * num_v)?;
3432 e.sigmoid_v(beta_raw, &mut beta, t * num_v)?;
3433 let mut g_log = e.uninit(t * num_v)?;
3434 e.gdn_glog_v(
3435 alpha,
3436 la.ssm_dt.float_data(),
3437 la.ssm_a.float_data(),
3438 &mut g_log,
3439 num_v,
3440 t,
3441 )?;
3442 if let Some(len_d) = pad_len {
3443 e.gdn_pad_mask(&mut beta, &mut g_log, len_d, num_v, t)?;
3444 }
3445 Ok(GdnPrep {
3446 hk,
3447 q_l2,
3448 k_l2,
3449 v_g,
3450 beta,
3451 g_log,
3452 kb16,
3453 qb16,
3454 })
3455 }
3456
3457 #[allow(clippy::too_many_arguments)]
3462 fn linear_attn_prime_core_batch(
3463 &self,
3464 e: &Engine,
3465 la: &LinearAttnLayer,
3466 g4: &[CudaSlice<f32>],
3467 offs: &[usize],
3468 ts: &[usize],
3469 caches: &mut [&mut Cache],
3470 il: usize,
3471 ) -> Result<Vec<(CudaSlice<f32>, Option<CudaSlice<u8>>)>, Box<dyn std::error::Error>> {
3472 let ssm = self.cfg.ssm.as_ref().unwrap();
3473 let d_state = ssm.state_size as usize;
3474 let num_k = ssm.group_count as usize;
3475 let num_v = ssm.time_step_rank as usize;
3476 let key_dim = d_state * num_k;
3477 let value_dim = d_state * num_v;
3478 let conv_dim = key_dim * 2 + value_dim;
3479 let eps = self.cfg.rms_eps;
3480 let scale = 1.0 / (d_state as f32).sqrt();
3481 let b = ts.len();
3482 let c = Engine::gdn_chunk_size();
3483 let carried = caches.iter().any(|c| c.pos > 0);
3486 let use_vl = !carried
3487 && (2..=8).contains(&b)
3488 && Engine::gdn_chunked_enabled()
3489 && ts.iter().all(|&t| t >= 16)
3490 && e.gdn_mma_enabled(c)
3491 && std::env::var("MEMRA_GDN_VL").as_deref() != Ok("0");
3492 if !use_vl {
3493 return (0..b)
3494 .map(|s| {
3495 let (o, t) = (offs[s], ts[s]);
3496 self.linear_attn_prime_core_pad_view(
3497 e,
3498 la,
3499 &g4[0].slice(o * conv_dim..(o + t) * conv_dim),
3500 &g4[1].slice(o * value_dim..(o + t) * value_dim),
3501 &g4[2].slice(o * num_v..(o + t) * num_v),
3502 &g4[3].slice(o * num_v..(o + t) * num_v),
3503 t,
3504 caches[s],
3505 il,
3506 None,
3507 )
3508 })
3509 .collect();
3510 }
3511 struct SeqBufs {
3515 conv_out: CudaSlice<f32>,
3516 q_g: CudaSlice<f32>,
3517 k_g: CudaSlice<f32>,
3518 v_g: CudaSlice<f32>,
3519 q_l2: CudaSlice<f32>,
3520 k_l2: CudaSlice<f32>,
3521 beta: CudaSlice<f32>,
3522 g_log: CudaSlice<f32>,
3523 gn: CudaSlice<f32>,
3524 gn16: CudaSlice<u8>,
3525 }
3526 let d_conv = ssm.conv_kernel as usize;
3527 let f16o = Self::f16out_on(e, 16);
3528 let hk = Self::gdn_hk(e, 16, num_v, num_k); let mut sb = Vec::with_capacity(b);
3530 let mut pres = Vec::with_capacity(b);
3531 for &t in ts.iter().take(b) {
3532 sb.push(SeqBufs {
3533 conv_out: e.uninit(conv_dim * t)?,
3534 q_g: e.uninit(d_state * hk * t)?,
3535 k_g: e.uninit(d_state * hk * t)?,
3536 v_g: e.uninit(d_state * num_v * t)?,
3537 q_l2: e.uninit(d_state * hk * t)?,
3538 k_l2: e.uninit(d_state * hk * t)?,
3539 beta: e.uninit(t * num_v)?,
3540 g_log: e.uninit(t * num_v)?,
3541 gn: e.uninit(d_state * num_v * t)?,
3542 gn16: e.alloc_u8_uninit(d_state * num_v * t * 2)?,
3543 });
3544 pres.push(e.gdn_chunk_alloc(num_v, t, c, hk)?);
3545 }
3546 let prep_args: Vec<crate::GdnPrepVl> = (0..b)
3547 .map(|s| {
3548 let (o, t) = (offs[s], ts[s]);
3549 let rl = caches[s].recur[il].as_ref().unwrap();
3550 crate::GdnPrepVl {
3551 qkv: e.addr_f32v(&g4[0].slice(o * conv_dim..(o + t) * conv_dim)),
3552 conv_state: e.addr_f32(&rl.conv_state),
3553 conv_out: e.addr_f32(&sb[s].conv_out),
3554 q_g: e.addr_f32(&sb[s].q_g),
3555 k_g: e.addr_f32(&sb[s].k_g),
3556 v_g: e.addr_f32(&sb[s].v_g),
3557 q_l2: e.addr_f32(&sb[s].q_l2),
3558 k_l2: e.addr_f32(&sb[s].k_l2),
3559 beta_raw: e.addr_f32v(&g4[2].slice(o * num_v..(o + t) * num_v)),
3560 alpha: e.addr_f32v(&g4[3].slice(o * num_v..(o + t) * num_v)),
3561 beta: e.addr_f32(&sb[s].beta),
3562 g_log: e.addr_f32(&sb[s].g_log),
3563 o: e.addr_f32(&pres[s].o),
3564 z: e.addr_f32v(&g4[1].slice(o * value_dim..(o + t) * value_dim)),
3565 gn: e.addr_f32(&sb[s].gn),
3566 gn16: e.addr_u8(&sb[s].gn16),
3567 kb16: if Engine::l2_v2_on(d_state) {
3568 e.addr_u8(&pres[s].kb16)
3569 } else {
3570 0
3571 },
3572 qb16: if Engine::l2_v2_on(d_state) && e.gdn_wgmma_on(c) {
3573 e.addr_u8(&pres[s].qb16)
3574 } else {
3575 0
3576 },
3577 t: t as i32,
3578 pad: 0,
3579 }
3580 })
3581 .collect();
3582 let args: Vec<crate::GdnSeqVl> = (0..b)
3583 .map(|s| {
3584 let rl = caches[s].recur[il].as_ref().unwrap();
3585 crate::GdnSeqVl {
3586 kb16: e.addr_u8(&pres[s].kb16),
3587 gcum: e.addr_f32(&pres[s].gcum),
3588 beta: e.addr_f32(&sb[s].beta),
3589 u: e.addr_f32(&pres[s].u),
3590 wb16: e.addr_u8(&pres[s].wb16),
3591 y: e.addr_u8(&pres[s].y16),
3592 ssnap: e.addr_u8(&pres[s].ssnap16),
3593 state_in: e.addr_f32(&rl.ssm_state),
3594 state_out: e.addr_f32(&rl.ssm_state_alt),
3595 q: e.addr_f32(&sb[s].q_l2),
3596 p: e.addr_f32(&pres[s].p),
3597 o: e.addr_f32(&pres[s].o),
3598 k: e.addr_f32(&sb[s].k_l2),
3599 v: e.addr_f32(&sb[s].v_g),
3600 g: e.addr_f32(&sb[s].g_log),
3601 a: e.addr_f32(&pres[s].a),
3602 w: e.addr_f32(&pres[s].w),
3603 t: ts[s] as i32,
3604 nc: pres[s].nc as i32,
3605 }
3606 })
3607 .collect();
3608 e.gdn_prep_vl8(
3609 &prep_args,
3610 la.ssm_conv1d.float_data(),
3611 la.ssm_dt.float_data(),
3612 la.ssm_a.float_data(),
3613 conv_dim,
3614 d_conv,
3615 d_state,
3616 num_v,
3617 num_k,
3618 key_dim,
3619 hk,
3620 eps,
3621 )?;
3622 if !Engine::l2_v2_on(d_state) {
3625 e.gdn_mirror_vl8(&args, num_v, 0, hk)?;
3626 }
3627 let wq8: Option<crate::GdnWVl8> = if e.gdn_wgmma_on(c) {
3629 if !Engine::l2_v2_on(d_state) {
3631 for s in 0..b {
3632 e.f32_to_bf16_into(&sb[s].q_l2, &mut pres[s].qb16, d_state * hk * ts[s])?;
3633 }
3634 }
3635 let mut wa = [crate::GdnWVl::default(); 8];
3636 for s in 0..b {
3637 wa[s] = crate::GdnWVl {
3638 qb16: e.addr_u8(&pres[s].qb16),
3639 pb16: e.addr_u8(&pres[s].pb16),
3640 };
3641 }
3642 Some(crate::GdnWVl8(wa))
3643 } else {
3644 None
3645 };
3646 e.gdn_chunk_k123_vl8(&args, num_v, hk, wq8.as_ref())?;
3647 e.gdn_chunk_vl8(&args, num_v, scale, hk, wq8.as_ref())?;
3648 if f16o {
3649 e.gdn_tail_vl8(&prep_args, la.ssm_norm.float_data(), d_state, num_v, eps)?;
3650 }
3651 let mut out = Vec::with_capacity(b);
3653 for (s, bufs) in sb.into_iter().enumerate() {
3654 let rl = caches[s].recur[il].as_mut().unwrap();
3655 std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
3656 let (o, t) = (offs[s], ts[s]);
3657 let SeqBufs { mut gn, gn16, .. } = bufs;
3658 if f16o {
3659 out.push((gn, Some(gn16)));
3660 } else {
3661 let z_v = g4[1].slice(o * value_dim..(o + t) * value_dim);
3662 e.gated_rmsnorm_zv(
3663 &pres[s].o,
3664 la.ssm_norm.float_data(),
3665 &z_v,
3666 &mut gn,
3667 d_state,
3668 num_v * t,
3669 eps,
3670 )?;
3671 out.push((gn, None));
3672 }
3673 }
3674 Ok(out)
3675 }
3676
3677 #[allow(clippy::too_many_arguments)]
3681 fn linear_attn_prime_core_pad_view(
3682 &self,
3683 e: &Engine,
3684 la: &LinearAttnLayer,
3685 qkv_mixed: &cudarc::driver::CudaView<f32>,
3686 z: &cudarc::driver::CudaView<f32>,
3687 beta_raw: &cudarc::driver::CudaView<f32>,
3688 alpha: &cudarc::driver::CudaView<f32>,
3689 t: usize,
3690 cache: &mut Cache,
3691 il: usize,
3692 pad_len: Option<&CudaSlice<i32>>,
3693 ) -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
3694 let cfg = &self.cfg;
3695 let ssm = cfg.ssm.as_ref().unwrap();
3696 let d_state = ssm.state_size as usize; let num_v = ssm.time_step_rank as usize; let eps = cfg.rms_eps;
3699 let scale = 1.0 / (d_state as f32).sqrt();
3700
3701 let prep =
3702 self.linear_attn_gdn_prep(e, la, qkv_mixed, beta_raw, alpha, t, cache, il, pad_len)?;
3703
3704 let mut o = e.uninit(d_state * num_v * t)?;
3710 let rl = cache.recur[il].as_mut().unwrap();
3711 {
3712 let crate::cache::RecurLayer {
3713 ssm_state,
3714 ssm_state_alt,
3715 ..
3716 } = rl;
3717 e.gdn_scan_prefill(
3718 &prep.q_l2,
3719 &prep.k_l2,
3720 &prep.v_g,
3721 &prep.g_log,
3722 &prep.beta,
3723 prep.kb16.as_ref(),
3724 prep.qb16.as_ref(),
3725 ssm_state,
3726 ssm_state_alt,
3727 &mut o,
3728 num_v,
3729 t,
3730 scale,
3731 prep.hk,
3732 )?;
3733 }
3734 std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
3735
3736 let mut gn = e.uninit(d_state * num_v * t)?;
3739 let gn16 = if Self::f16out_on(e, t) {
3740 let mut g16 = e.alloc_u8_uninit(d_state * num_v * t * 2)?;
3741 e.gated_rmsnorm_f16out_zv(
3742 &o,
3743 la.ssm_norm.float_data(),
3744 z,
3745 &mut gn,
3746 &mut g16,
3747 d_state,
3748 num_v * t,
3749 eps,
3750 )?;
3751 Some(g16)
3752 } else {
3753 e.gated_rmsnorm_zv(
3754 &o,
3755 la.ssm_norm.float_data(),
3756 z,
3757 &mut gn,
3758 d_state,
3759 num_v * t,
3760 eps,
3761 )?;
3762 None
3763 };
3764 Ok((gn, gn16))
3765 }
3766
3767 #[allow(clippy::too_many_arguments)]
3769 fn linear_attn_prime_core_pad(
3770 &self,
3771 e: &Engine,
3772 la: &LinearAttnLayer,
3773 g4: Vec<CudaSlice<f32>>,
3774 t: usize,
3775 cache: &mut Cache,
3776 il: usize,
3777 pad_len: Option<&CudaSlice<i32>>,
3778 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3779 let (gn, gn16) = self.linear_attn_prime_core_pad_inner(e, la, g4, t, cache, il, pad_len)?;
3780 if let Some(xh) = &gn16 {
3781 if let Some(y) = e.try_f16_gemm_pre(&la.ssm_out, xh, t)? {
3782 return Ok(y);
3783 }
3784 }
3785 Ok(e.matmul(&la.ssm_out, &gn, t)?)
3786 }
3787
3788 pub fn full_attn(
3793 &self,
3794 e: &Engine,
3795 fa: &FullAttnLayer,
3796 h: &CudaSlice<f32>,
3797 pos_d: &CudaSlice<i32>,
3798 t: usize,
3799 il: usize,
3800 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3801 if self.cfg.step35.is_some() {
3802 return self.step35_attn(e, fa, h, pos_d, t, il);
3803 }
3804 let cfg = &self.cfg;
3805 let _n_embd = cfg.n_embd as usize;
3806 let geometry = cfg.full_attention_geometry_at(il as u32);
3807 let n_head = geometry.n_head as usize;
3808 let n_head_kv = geometry.n_head_kv as usize;
3809 let head_dim = geometry.head_dim_k as usize;
3810 let eps = cfg.rms_eps;
3811 let scale = geometry.attention_scale();
3812
3813 let gated = geometry.attention_gate == memra_gguf::config::AttentionGateKind::FusedQ;
3816 let mut g3 = e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?;
3818 let v = g3.pop().unwrap();
3819 let mut k = g3.pop().unwrap();
3820 let qf = g3.pop().unwrap();
3821 let (mut q, gate) = if gated {
3822 let mut q = e.uninit(t * n_head * head_dim)?;
3823 let mut gate = e.uninit(t * n_head * head_dim)?;
3824 e.q_gate_split(&qf, &mut q, &mut gate, head_dim, n_head, t)?;
3825 (q, Some(gate))
3826 } else {
3827 (qf, None)
3828 };
3829
3830 let mut qn = e.uninit(t * n_head * head_dim)?;
3832 e.rms_norm(
3833 &q,
3834 fa.q_norm.float_data(),
3835 &mut qn,
3836 head_dim,
3837 n_head * t,
3838 eps,
3839 )?;
3840 q = qn;
3841 let mut kn = e.uninit(t * n_head_kv * head_dim)?;
3842 e.rms_norm(
3843 &k,
3844 fa.k_norm.float_data(),
3845 &mut kn,
3846 head_dim,
3847 n_head_kv * t,
3848 eps,
3849 )?;
3850 k = kn;
3851 let rope_dims = geometry.n_rot as usize;
3852 e.rope_neox(
3853 &mut q,
3854 pos_d,
3855 head_dim,
3856 rope_dims,
3857 n_head,
3858 t,
3859 geometry.rope_base,
3860 1.0,
3861 )?;
3862 e.rope_neox(
3863 &mut k,
3864 pos_d,
3865 head_dim,
3866 rope_dims,
3867 n_head_kv,
3868 t,
3869 geometry.rope_base,
3870 1.0,
3871 )?;
3872
3873 let mut attn = e.uninit(t * n_head * head_dim)?;
3875 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
3878 e.sdpa_naive(
3880 &q, &k, &v, &mut attn, head_dim, n_head, n_head_kv, t, t, scale, true,
3881 )?;
3882 } else {
3883 e.fa_prefill(
3884 &q, &k, &v, &mut attn, head_dim, n_head, n_head_kv, t, t, scale, true,
3885 )?;
3886 }
3887
3888 let attn_g = match &gate {
3890 Some(gate) => {
3891 let mut gsig = e.uninit(t * n_head * head_dim)?;
3892 e.sigmoid(gate, &mut gsig, t * n_head * head_dim)?;
3893 let mut ag = e.uninit(t * n_head * head_dim)?;
3894 e.mul(&attn, &gsig, &mut ag, t * n_head * head_dim)?;
3895 ag
3896 }
3897 None => attn,
3898 };
3899
3900 let o = e.matmul(&fa.wo, &attn_g, t)?;
3902 Ok(o)
3903 }
3904
3905 pub fn linear_attn(
3907 &self,
3908 e: &Engine,
3909 la: &LinearAttnLayer,
3910 h: &CudaSlice<f32>,
3911 t: usize,
3912 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3913 let cfg = &self.cfg;
3914 let _n_embd = cfg.n_embd as usize;
3915 let ssm = cfg.ssm.as_ref().unwrap();
3916 let d_state = ssm.state_size as usize; let num_k = ssm.group_count as usize; let num_v = ssm.time_step_rank as usize; let d_conv = ssm.conv_kernel as usize; let head_k = d_state;
3921 let head_v = d_state;
3922 let key_dim = head_k * num_k; let value_dim = head_v * num_v; let conv_dim = key_dim * 2 + value_dim; let eps = cfg.rms_eps;
3926 let scale = 1.0 / (d_state as f32).sqrt();
3927
3928 let mut g4 = e.matmul_group(
3931 &[&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha],
3932 h,
3933 t,
3934 )?;
3935 let alpha = g4.pop().unwrap(); let beta_raw = g4.pop().unwrap(); let z = g4.pop().unwrap(); let qkv_mixed = g4.pop().unwrap(); let _ = (head_k, head_v);
3947 let mut q_g = e.uninit(d_state * num_v * t)?;
3948 let mut k_g = e.uninit(d_state * num_v * t)?;
3949 let mut v_g = e.uninit(d_state * num_v * t)?;
3950 e.ssm_conv1d_gdn(
3951 &qkv_mixed,
3952 la.ssm_conv1d.float_data(),
3953 &mut q_g,
3954 &mut k_g,
3955 &mut v_g,
3956 conv_dim,
3957 t,
3958 d_conv,
3959 d_state,
3960 num_v,
3961 num_k,
3962 key_dim,
3963 )?;
3964 let mut q_l2 = e.uninit(d_state * num_v * t)?;
3966 e.l2_norm(&q_g, &mut q_l2, d_state, num_v * t, eps)?;
3967 let mut k_l2 = e.uninit(d_state * num_v * t)?;
3968 e.l2_norm(&k_g, &mut k_l2, d_state, num_v * t, eps)?;
3969 let v_gd = v_g;
3970
3971 let mut beta = e.uninit(t * num_v)?;
3974 e.sigmoid(&beta_raw, &mut beta, t * num_v)?;
3975 let mut g_log = e.uninit(t * num_v)?;
3977 e.gdn_glog(
3978 &alpha,
3979 la.ssm_dt.float_data(),
3980 la.ssm_a.float_data(),
3981 &mut g_log,
3982 num_v,
3983 t,
3984 )?;
3985
3986 let state_in = e.zeros(d_state * d_state * num_v)?; let mut state_out = e.zeros(d_state * d_state * num_v)?;
3989 let mut o = e.uninit(d_state * num_v * t)?;
3990 e.gdn_scan_prefill(
3991 &q_l2,
3992 &k_l2,
3993 &v_gd,
3994 &g_log,
3995 &beta,
3996 None,
3997 None,
3998 &state_in,
3999 &mut state_out,
4000 &mut o,
4001 num_v,
4002 t,
4003 scale,
4004 num_v,
4005 )?;
4006
4007 let mut gn = e.uninit(d_state * num_v * t)?;
4012 e.gated_rmsnorm(
4013 &o,
4014 la.ssm_norm.float_data(),
4015 &z,
4016 &mut gn,
4017 d_state,
4018 num_v * t,
4019 eps,
4020 )?;
4021
4022 let out = e.matmul(&la.ssm_out, &gn, t)?;
4026 Ok(out)
4027 }
4028}
4029
4030impl HybridModel {
4031 pub fn moe_ffn_il(
4042 &self,
4043 e: &Engine,
4044 m: &MoeWeights,
4045 z: &CudaSlice<f32>,
4046 t: usize,
4047 il: u16,
4048 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4049 Self::moe_ffn_inner(e, m, z, None, t, &self.cfg, il, self.max_moe_block(), false)
4050 }
4051
4052 pub fn moe_ffn_il_prefill(
4055 &self,
4056 e: &Engine,
4057 m: &MoeWeights,
4058 z: &CudaSlice<f32>,
4059 t: usize,
4060 il: u16,
4061 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4062 Self::moe_ffn_inner(e, m, z, None, t, &self.cfg, il, self.max_moe_block(), true)
4063 }
4064
4065 pub fn moe_ffn_il_zq8(
4069 &self,
4070 e: &Engine,
4071 m: &MoeWeights,
4072 z: &CudaSlice<f32>,
4073 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
4074 t: usize,
4075 il: u16,
4076 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4077 Self::moe_ffn_inner(e, m, z, zq8, t, &self.cfg, il, self.max_moe_block(), false)
4078 }
4079
4080 pub(crate) fn moe_ffn(
4088 e: &Engine,
4089 m: &MoeWeights,
4090 z: &CudaSlice<f32>,
4091 t: usize,
4092 cfg: &ModelConfig,
4093 il: u16,
4094 max_block: usize,
4095 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4096 Self::moe_ffn_inner(e, m, z, None, t, cfg, il, max_block, false)
4097 }
4098
4099 #[allow(clippy::too_many_arguments)]
4100 pub(crate) fn moe_ffn_inner(
4101 e: &Engine,
4102 m: &MoeWeights,
4103 z: &CudaSlice<f32>,
4104 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
4105 t: usize,
4106 cfg: &ModelConfig,
4107 il: u16,
4108 max_block: usize,
4109 prefill: bool,
4110 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4111 let worker_io = crate::spill_pread::worker_enabled();
4112 let epoch_lfu = std::env::var_os("MEMRA_MOE_LFU_DECAY").is_some();
4113 if Engine::moe_cache_enabled() && (worker_io || epoch_lfu) {
4114 e.with_moe_cache(max_block, |cache, _| {
4115 cache.begin_forward_epoch(il, t);
4116 if worker_io {
4117 cache.begin_worker_scope();
4118 }
4119 Ok(())
4120 })?;
4121 }
4122 if Self::sigmoid_resident_dev_eligible(e, m, cfg) {
4123 let moe = cfg.moe.as_ref().unwrap();
4124 let n_expert = moe.expert_count as usize;
4125 let n_used = moe.expert_used_count as usize;
4126 let sigmoid = cfg.sigmoid_router().unwrap();
4127 let logits = Self::moe_router_logits(e, m, z, t, cfg)?;
4128 Self::trace_sigmoid_router_logits(e, il, t, n_expert, n_used, &logits, m, sigmoid)?;
4129 return Self::moe_ffn_sigmoid_dev(e, m, z, zq8, &logits, t, cfg, il, sigmoid);
4130 }
4131 if t > 1 && moe_grouped_enabled(cfg, prefill) {
4134 let grouped_out = Self::moe_ffn_grouped(e, m, z, t, cfg, il, max_block)?;
4135 if std::env::var("MEMRA_MOE_GATE").is_ok() {
4140 let seq_out = Self::moe_ffn_sequential(e, m, z, t, cfg, il, max_block)?;
4141 let g_host = e.dtoh(&grouped_out)?;
4142 let s_host = e.dtoh(&seq_out)?;
4143 let g_bytes: &[u8] = unsafe {
4144 std::slice::from_raw_parts(g_host.as_ptr() as *const u8, g_host.len() * 4)
4145 };
4146 let s_bytes: &[u8] = unsafe {
4147 std::slice::from_raw_parts(s_host.as_ptr() as *const u8, s_host.len() * 4)
4148 };
4149 if g_bytes == s_bytes {
4150 println!("moe-gate il={il} t={t} BYTE-IDENTICAL");
4151 } else {
4152 let diffs = g_host
4153 .iter()
4154 .zip(s_host.iter())
4155 .enumerate()
4156 .filter(|(_, (a, b))| a != b)
4157 .count();
4158 let maxdiff = g_host
4159 .iter()
4160 .zip(s_host.iter())
4161 .map(|(a, b)| (a - b).abs())
4162 .fold(0.0f32, f32::max);
4163 panic!(
4164 "moe-gate il={il} t={t} MISMATCH: {diffs}/{} elems differ, maxdiff={maxdiff:.6e}",
4165 g_host.len()
4166 );
4167 }
4168 }
4169 return Ok(grouped_out);
4170 }
4171 Self::moe_ffn_sequential_zq8(e, m, z, zq8, t, cfg, il, max_block)
4172 }
4173
4174 fn sigmoid_resident_dev_eligible(e: &Engine, m: &MoeWeights, cfg: &ModelConfig) -> bool {
4175 let Some(moe) = cfg.moe.as_ref() else {
4176 return false;
4177 };
4178 static OBSERVATION_MODE: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4181 let observation_mode = *OBSERVATION_MODE.get_or_init(|| {
4182 std::env::var("MEMRA_MOE_STATS").is_ok()
4183 || std::env::var("MEMRA_MOE_TRACE").is_ok()
4184 || std::env::var("MEMRA_MOE_WEIGHT_TRACE").is_ok()
4185 || std::env::var("MEMRA_MOE_INPUT_TRACE_DIR").is_ok()
4186 || std::env::var("MEMRA_MOE_GATE").is_ok()
4187 });
4188 cfg.step35.is_some()
4189 && sigmoid_router_enabled()
4190 && moe_dev_enabled()
4191 && moe_slab_enabled()
4192 && !observation_mode
4193 && moe.expert_used_count <= 8
4194 && m.has_uniform_expert_layout()
4195 && m.gate_exps.macros.is_none()
4196 && m.up_exps.macros.is_none()
4197 && m.down_exps.macros.is_none()
4198 && !m.has_macros
4199 && moe_q8_enabled()
4200 && q8_expert_supported(m.gate_exps.qtype)
4201 && q8_expert_supported(m.up_exps.qtype)
4202 && q8_expert_supported(m.down_exps.qtype)
4203 && m.dev_exps
4204 .as_ref()
4205 .is_some_and(|dev| dev.dev == e.ctx().ordinal())
4206 }
4207
4208 pub(crate) fn moe_ffn_sequential(
4210 e: &Engine,
4211 m: &MoeWeights,
4212 z: &CudaSlice<f32>,
4213 t: usize,
4214 cfg: &ModelConfig,
4215 il: u16,
4216 max_block: usize,
4217 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4218 Self::moe_ffn_sequential_zq8(e, m, z, None, t, cfg, il, max_block)
4219 }
4220
4221 fn moe_router_logits(
4225 e: &Engine,
4226 m: &MoeWeights,
4227 z: &CudaSlice<f32>,
4228 t: usize,
4229 cfg: &ModelConfig,
4230 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4231 if t < PRIME_MIN_T {
4232 if crate::router_kernel_on() {
4234 e.router_gemv(
4235 m.gate_inp.float_data(),
4236 z,
4237 cfg.n_embd as usize,
4238 m.gate_exps.n_expert,
4239 t,
4240 )
4241 } else {
4242 e.matmul_decode_exact(&m.gate_inp, z, t)
4243 }
4244 } else if crate::router_prefill_exact_on() && crate::router_kernel_on() {
4245 e.router_gemv(
4246 m.gate_inp.float_data(),
4247 z,
4248 cfg.n_embd as usize,
4249 m.gate_exps.n_expert,
4250 t,
4251 )
4252 } else {
4253 e.matmul(&m.gate_inp, z, t)
4254 }
4255 }
4256
4257 fn trace_moe_routes(
4261 il: u16,
4262 t: usize,
4263 sel_all: &[u32],
4264 weights: &[f32],
4265 ) -> Result<(), Box<dyn std::error::Error>> {
4266 use std::io::Write as _;
4267 if let Ok(path) = std::env::var("MEMRA_MOE_TRACE") {
4268 let mut f = std::fs::OpenOptions::new()
4269 .create(true)
4270 .append(true)
4271 .open(path)?;
4272 let ids: Vec<String> = sel_all.iter().map(|s| s.to_string()).collect();
4273 writeln!(f, "{} {} {}", il, t, ids.join(","))?;
4274 }
4275 if let Ok(path) = std::env::var("MEMRA_MOE_WEIGHT_TRACE") {
4276 let mut f = std::fs::OpenOptions::new()
4277 .create(true)
4278 .append(true)
4279 .open(path)?;
4280 let pairs: Vec<String> = sel_all
4281 .iter()
4282 .zip(weights)
4283 .map(|(expert, weight)| format!("{expert}:{weight:.9}"))
4284 .collect();
4285 writeln!(f, "{} {} {}", il, t, pairs.join(","))?;
4286 }
4287 Ok(())
4288 }
4289
4290 #[allow(clippy::too_many_arguments)]
4291 fn trace_sigmoid_router_logits(
4292 e: &Engine,
4293 il: u16,
4294 t: usize,
4295 n_expert: usize,
4296 n_used: usize,
4297 logits: &CudaSlice<f32>,
4298 m: &MoeWeights,
4299 (scaling_factor, route_norm): (f32, bool),
4300 ) -> Result<(), Box<dyn std::error::Error>> {
4301 if !crate::sigrouter_contract::served_logit_trace_enabled() || t != 1 {
4302 return Ok(());
4303 }
4304 let logits = e.dtoh(logits)?;
4305 let active: Vec<u8> = m
4306 .active_experts
4307 .as_ref()
4308 .map(|mask| mask.iter().map(|&enabled| u8::from(enabled)).collect())
4309 .unwrap_or_else(|| vec![1; n_expert]);
4310 let bias = m.exp_probs_b.clone().unwrap_or_else(|| vec![0.0; n_expert]);
4311 crate::sigrouter_contract::capture_served_logits(
4312 il as u32,
4313 t,
4314 n_expert,
4315 n_used,
4316 scaling_factor,
4317 route_norm,
4318 &active,
4319 &bias,
4320 &logits,
4321 )?;
4322 Ok(())
4323 }
4324
4325 fn trace_moe_input(
4330 e: &Engine,
4331 il: u16,
4332 t: usize,
4333 n_embd: usize,
4334 z: &CudaSlice<f32>,
4335 ) -> Result<(), Box<dyn std::error::Error>> {
4336 use std::io::Write as _;
4337 let Ok(dir) = std::env::var("MEMRA_MOE_INPUT_TRACE_DIR") else {
4338 return Ok(());
4339 };
4340 let host = e.dtoh(z)?;
4341 if host.len() != t * n_embd {
4342 return Err(format!(
4343 "MoE input trace shape mismatch at layer {il}: got {} values, expected {}x{}",
4344 host.len(),
4345 t,
4346 n_embd
4347 )
4348 .into());
4349 }
4350 let bytes = unsafe {
4351 std::slice::from_raw_parts(
4352 host.as_ptr().cast::<u8>(),
4353 host.len() * std::mem::size_of::<f32>(),
4354 )
4355 };
4356 let state = MOE_INPUT_TRACE_WRITER.get_or_init(|| std::sync::Mutex::new(None));
4357 let mut state = state
4358 .lock()
4359 .map_err(|_| "MoE input trace writer lock is poisoned")?;
4360 if state.is_none() {
4361 let dir = std::path::PathBuf::from(&dir);
4362 std::fs::create_dir_all(&dir)?;
4363 let index = std::fs::OpenOptions::new()
4364 .create(true)
4365 .append(true)
4366 .open(dir.join("index.jsonl"))?;
4367 *state = Some(MoeInputTraceWriter {
4368 dir,
4369 index,
4370 payloads: std::collections::HashMap::new(),
4371 });
4372 }
4373 let writer = state.as_mut().unwrap();
4374 if writer.dir != std::path::Path::new(&dir) {
4375 return Err("MEMRA_MOE_INPUT_TRACE_DIR changed after capture started".into());
4376 }
4377 let file_name = format!("layer-{il:03}.f32");
4378 if !writer.payloads.contains_key(&il) {
4379 let payload = std::fs::OpenOptions::new()
4380 .create(true)
4381 .append(true)
4382 .open(writer.dir.join(&file_name))?;
4383 let offset = payload.metadata()?.len();
4384 writer.payloads.insert(il, (payload, offset));
4385 }
4386 let (payload, offset) = writer.payloads.get_mut(&il).unwrap();
4387 let row_offset = *offset;
4388 payload.write_all(bytes)?;
4389 *offset += bytes.len() as u64;
4390 writeln!(
4391 writer.index,
4392 "{{\"format\":\"memra-moe-input-trace-v1\",\"layer\":{il},\"tokens\":{t},\
4393 \"hidden_size\":{n_embd},\"file\":\"{file_name}\",\"offset\":{row_offset},\
4394 \"payload_bytes\":{}}}",
4395 bytes.len()
4396 )?;
4397 Ok(())
4398 }
4399
4400 #[allow(clippy::too_many_arguments)]
4401 pub(crate) fn moe_ffn_sequential_zq8(
4402 e: &Engine,
4403 m: &MoeWeights,
4404 z: &CudaSlice<f32>,
4405 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
4406 t: usize,
4407 cfg: &ModelConfig,
4408 il: u16,
4409 max_block: usize,
4410 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4411 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
4412 let moe = cfg.moe.as_ref().unwrap();
4413 let n_embd = cfg.n_embd as usize; let n_expert = moe.expert_count as usize; let n_used = moe.expert_used_count as usize; let n_ff_exp = moe.expert_ff_length as usize; debug_assert_eq!(m.gate_exps.in_f, n_embd);
4420 debug_assert_eq!(m.gate_exps.out_f, n_ff_exp);
4421 debug_assert_eq!(m.down_exps.in_f, n_ff_exp); debug_assert_eq!(m.down_exps.out_f, n_embd); debug_assert_eq!(m.gate_exps.n_expert, n_expert);
4424
4425 let lim_exp = cfg.clamp_exp_at(il as u32);
4428 let lim_shexp = cfg.clamp_shexp_at(il as u32);
4429 let use_cache = Engine::moe_cache_enabled();
4430 let uniform_experts = m.has_uniform_expert_layout();
4431 let moe_q8 = uniform_experts
4432 && moe_q8_enabled()
4433 && q8_expert_supported(m.gate_exps.qtype)
4434 && q8_expert_supported(m.up_exps.qtype)
4435 && q8_expert_supported(m.down_exps.qtype);
4436 let cpu_expert_requested = crate::cpu_experts::configured();
4443 if cpu_expert_requested && (cfg.hy3.is_none() || cfg.m3.is_some()) {
4444 return Err(std::io::Error::other(
4445 "MEMRA_CPU_EXPERT_LIB is experimental and currently gated to Hy3",
4446 )
4447 .into());
4448 }
4449 let cpu_hybrid = cpu_expert_requested && t < PRIME_MIN_T && m.dev_exps.is_none();
4450 let freeze_cpu_residency = cpu_expert_requested
4456 && std::env::var("MEMRA_CPU_EXPERT_FREEZE_CACHE").as_deref() == Ok("1");
4457 let caller_warms_before_freeze = std::env::var("MEMRA_CPU_EXPERT_FREEZE_WARMUP_TOKENS")
4458 .ok()
4459 .and_then(|value| value.parse::<usize>().ok())
4460 .is_some_and(|tokens| tokens > 0);
4461 if cpu_hybrid && freeze_cpu_residency && !caller_warms_before_freeze {
4462 e.freeze_moe_cache();
4463 }
4464 let cache_frozen = use_cache && e.moe_cache_frozen();
4465 let cache_dispatch = use_cache && (!cache_frozen || cpu_hybrid);
4466
4467 let logits = Self::moe_router_logits(e, m, z, t, cfg)?;
4470 if let Some(sig) = cfg.sigmoid_router() {
4471 Self::trace_sigmoid_router_logits(e, il, t, n_expert, n_used, &logits, m, sig)?;
4472 }
4473
4474 let no_exp_macros = m.gate_exps.macros.is_none()
4513 && m.up_exps.macros.is_none()
4514 && m.down_exps.macros.is_none();
4515 if cfg.sigmoid_router().is_none()
4519 && cfg.m3.is_none()
4520 && cfg.hy3.is_none()
4521 && !cfg.swiglu_clamped_at(il as u32)
4522 && no_exp_macros
4523 && t >= PRIME_MIN_T
4524 && m.dev_exps.is_some()
4525 && moe_q8_enabled()
4526 && q8_expert_supported(m.gate_exps.qtype)
4527 && q8_expert_supported(m.up_exps.qtype)
4528 && q8_expert_supported(m.down_exps.qtype)
4529 && std::env::var("MEMRA_MOE_PAIRS")
4530 .map(|v| v != "0")
4531 .unwrap_or(true)
4532 && std::env::var("MEMRA_MOE_STATS").is_err()
4533 {
4534 return Self::moe_ffn_pairs(e, m, z, &logits, t, cfg);
4535 }
4536
4537 let dev_ok = uniform_experts
4555 && cfg.sigmoid_router().is_none()
4556 && cfg.m3.is_none()
4557 && cfg.hy3.is_none()
4558 && !cfg.swiglu_clamped_at(il as u32);
4559 let observe_routes = std::env::var("MEMRA_MOE_STATS").is_ok()
4563 || std::env::var("MEMRA_MOE_TRACE").is_ok()
4564 || std::env::var("MEMRA_MOE_WEIGHT_TRACE").is_ok()
4565 || std::env::var("MEMRA_MOE_INPUT_TRACE_DIR").is_ok();
4566 if dev_ok
4567 && t < PRIME_MIN_T
4568 && m.dev_exps.is_some()
4569 && n_used <= 8
4570 && moe_dev_enabled()
4571 && !observe_routes
4572 {
4573 return Self::moe_ffn_dev(e, m, z, zq8, &logits, t, cfg, il, max_block);
4574 }
4575 if dev_ok && use_cache && n_used <= 8 && moe_dev_enabled() && !observe_routes {
4576 let row_ok = e.with_moe_cache(max_block, |c, eng| {
4577 if moe_prewarm_enabled() {
4578 c.prewarm_layer(il, m, eng)?;
4579 }
4580 Ok(c.layer_dev_row(il, n_expert, eng)?.is_some())
4581 })?;
4582 if row_ok {
4583 return Self::moe_ffn_dev(e, m, z, zq8, &logits, t, cfg, il, max_block);
4584 }
4585 }
4586
4587 let (sel_all, w_all, routed_cpu_input) = if let Some(sig) = cfg.sigmoid_router() {
4589 if cpu_hybrid {
4590 let (sel, w, input) = Self::moe_route_sigmoid_with_input(
4591 e,
4592 &logits,
4593 z,
4594 t,
4595 n_expert,
4596 n_used,
4597 m.exp_probs_b.as_deref(),
4598 sig,
4599 m.active_experts.as_deref(),
4600 )?;
4601 (sel, w, Some(input))
4602 } else {
4603 let (sel, w) =
4604 Self::moe_route_sigmoid_cfg(e, &logits, t, n_expert, n_used, m, sig)?;
4605 (sel, w, None)
4606 }
4607 } else {
4608 let (sel, w) =
4609 Self::moe_route_cfg(e, &logits, t, n_expert, n_used, m.active_experts.as_deref())?;
4610 (sel, w, None)
4611 };
4612 crate::moesd::record_host_routes(il, n_expert, n_used, &sel_all)?;
4613
4614 Self::trace_moe_routes(il, t, &sel_all, &w_all)?;
4618 Self::trace_moe_input(e, il, t, n_embd, z)?;
4619
4620 let worker_disk_prefetch =
4632 cache_dispatch && crate::spill_pread::worker_enabled() && !cpu_hybrid;
4633 let promote_worker_h2d =
4634 t == 1 && worker_disk_prefetch && crate::spill_pread::copy_h2d_enabled();
4635 if promote_worker_h2d {
4636 let mut selected_blocks = Vec::with_capacity(n_used * 3);
4637 for &ex in sel_all.iter().take(n_used) {
4638 let ex = ex as u16;
4639 selected_blocks.extend([
4640 BlockId::new(il, PROJ_GATE, ex),
4641 BlockId::new(il, PROJ_UP, ex),
4642 BlockId::new(il, PROJ_DOWN, ex),
4643 ]);
4644 }
4645 for &ex in sel_all.iter().take(n_used) {
4646 Self::moe_prefetch_disk_expert(e, il, ex as usize, m, max_block, &selected_blocks)?;
4647 }
4648 e.with_moe_cache(max_block, |cache, eng| {
4649 cache.promote_worker_reads_at_safe_boundary(
4650 &selected_blocks,
4651 &selected_blocks,
4652 eng,
4653 )?;
4654 Ok(())
4655 })?;
4656 }
4657
4658 if t > 1 && std::env::var("MEMRA_MOE_STATS").is_ok() {
4661 let mut cnt = vec![0u32; n_expert];
4662 for &s in sel_all.iter() {
4663 cnt[s as usize] += 1;
4664 }
4665 let total = sel_all.len() as f64;
4666 let mut h = 0.0f64;
4667 let mut active = 0usize;
4668 for &c in &cnt {
4669 if c > 0 {
4670 active += 1;
4671 let p = c as f64 / total;
4672 h -= p * p.log2();
4673 }
4674 }
4675 let maxc = cnt.iter().copied().max().unwrap_or(0);
4676 println!(
4677 "moe-stats il={} t={} assignments={} active={}/{} entropy={:.3}b (max {:.3}b) mean_tok_per_active={:.2} max_tok_per_expert={}",
4678 il,
4679 t,
4680 sel_all.len(),
4681 active,
4682 n_expert,
4683 h,
4684 (n_expert as f64).log2(),
4685 total / active.max(1) as f64,
4686 maxc
4687 );
4688 }
4689
4690 let gdec_may_fire = uniform_experts
4703 && use_cache
4704 && n_used <= 8
4705 && gdec_enabled()
4706 && !cfg.swiglu_clamped_at(il as u32);
4707 let slab_local = m
4723 .dev_exps
4724 .as_ref()
4725 .filter(|d| !d.gu_il && moe_slab_enabled() && d.dev == e.ctx().ordinal());
4726 let slab_bases = slab_local.map(|d| {
4727 use cudarc::driver::DevicePtr;
4728 let s = e.stream();
4729 let (pg, _g0) = d.gate.device_ptr(&s);
4730 let (pu, _g1) = d.up.device_ptr(&s);
4731 let (pd, _g2) = d.down.device_ptr(&s);
4732 (pg as u64, pu as u64, pd as u64)
4733 });
4734 let slab_fused_may_fire = slab_bases.is_some()
4744 && n_used <= 8
4745 && gdec_enabled()
4746 && !cfg.swiglu_clamped_at(il as u32)
4747 && cfg.m3.is_none()
4748 && no_exp_macros
4749 && moe_q8;
4750 let mut moe_out = if gdec_may_fire || slab_fused_may_fire {
4753 e.uninit(t * n_embd)?
4754 } else {
4755 e.zeros(t * n_embd)?
4756 };
4757 let cpu_input = if cpu_hybrid {
4760 Some(routed_cpu_input.ok_or("CPU expert routing did not return the MoE input")?)
4761 } else {
4762 None
4763 };
4764
4765 let g_len = m.gate_exps.max_expert_bytes(); let u_len = m.up_exps.max_expert_bytes(); let d_len = m.down_exps.max_expert_bytes(); let mut scratch_g: Option<CudaSlice<u8>> = None;
4773 let mut scratch_u: Option<CudaSlice<u8>> = None;
4774 let mut scratch_d: Option<CudaSlice<u8>> = None;
4775 let page_window = moe_page_prefetch_window();
4783
4784 for tok in 0..t {
4787 let sel = &sel_all[tok * n_used..(tok + 1) * n_used];
4788 let w = &w_all[tok * n_used..(tok + 1) * n_used];
4789 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd); let mut tok_q8: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
4791
4792 let no_macros = m.gate_exps.macros.is_none()
4806 && m.up_exps.macros.is_none()
4807 && m.down_exps.macros.is_none();
4808 if slab_fused_may_fire {
4818 let (pg, pu, pd) = slab_bases.unwrap();
4819 let mut gp = [0u64; 8];
4820 let mut up = [0u64; 8];
4821 let mut dp = [0u64; 8];
4822 for (j, &ex) in sel.iter().enumerate() {
4823 let ex = ex as usize;
4824 gp[j] = pg + (ex * m.gate_exps.expert_stride) as u64;
4825 up[j] = pu + (ex * m.up_exps.expert_stride) as u64;
4826 dp[j] = pd + (ex * m.down_exps.expert_stride) as u64;
4827 }
4828 let mut wv = [0f32; 8];
4829 wv[..n_used].copy_from_slice(w);
4830 if tok_q8.is_none() {
4831 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
4832 }
4833 let (zq, zd) = tok_q8.as_ref().unwrap();
4834 let act = e.moe_gate_up_silu8_q8(
4835 crate::WPtr8(gp),
4836 crate::WPtr8(up),
4837 zq,
4838 zd,
4839 n_embd,
4840 n_ff_exp,
4841 n_used,
4842 m.gate_exps.qtype,
4843 m.up_exps.qtype,
4844 m.gate_exps.row_bytes,
4845 m.up_exps.row_bytes,
4846 )?;
4847 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
4848 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
4849 e.moe_down8_fma_q8(
4850 crate::WPtr8(dp),
4851 crate::F32x8(wv),
4852 &aq2,
4853 &ad2,
4854 &mut dst,
4855 n_ff_exp,
4856 n_embd,
4857 n_used,
4858 m.down_exps.qtype,
4859 m.down_exps.row_bytes,
4860 )?;
4861 continue;
4862 }
4863 if gdec_may_fire && moe_q8 && cfg.m3.is_none() && no_macros {
4864 if tok_q8.is_none() {
4865 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
4866 }
4867 let (zq, zd) = tok_q8.as_ref().unwrap();
4868 if Self::moe_gdec_token_q8(
4869 e,
4870 m,
4871 il,
4872 max_block,
4873 zq,
4874 zd,
4875 sel,
4876 w,
4877 &mut moe_out,
4878 tok,
4879 n_embd,
4880 n_ff_exp,
4881 n_used,
4882 )? {
4883 continue;
4884 }
4885 } else if gdec_may_fire
4886 && cfg.m3.is_none()
4887 && no_macros
4888 && Self::moe_gdec_token(
4889 e,
4890 m,
4891 il,
4892 max_block,
4893 &zt,
4894 sel,
4895 w,
4896 &mut moe_out,
4897 tok,
4898 n_embd,
4899 n_ff_exp,
4900 n_used,
4901 )?
4902 {
4903 continue;
4904 }
4905
4906 if gdec_may_fire || slab_fused_may_fire {
4912 let mut row = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
4913 e.memset_zeros_view(&mut row)?;
4914 }
4915
4916 let mut cpu_mask = vec![false; sel.len()];
4922 let cpu_worker = if let Some(host_input) = cpu_input.as_ref() {
4923 let gpu_resident = if use_cache {
4924 e.with_moe_cache(max_block, |cache, _| {
4925 Ok(sel
4926 .iter()
4927 .map(|&expert| {
4928 let expert = expert as u16;
4929 [PROJ_GATE, PROJ_UP, PROJ_DOWN]
4930 .into_iter()
4931 .filter(|&projection| {
4932 cache
4933 .resident(BlockId::new(il, projection, expert))
4934 .is_some()
4935 })
4936 .count()
4937 })
4938 .collect::<Vec<_>>())
4939 })?
4940 } else {
4941 vec![0; sel.len()]
4942 };
4943 let mut cpu_selected = Vec::new();
4944 for (index, (&expert, &route_weight)) in sel.iter().zip(w).enumerate() {
4945 if gpu_resident[index] != 3 {
4946 cpu_mask[index] = true;
4947 crate::cpu_experts::record_incomplete_gpu_residency(gpu_resident[index]);
4948 let expert = expert as usize;
4949 cpu_selected.push((expert, route_weight));
4950 }
4951 }
4952 if crate::cpu_experts::predictor_enabled() {
4953 let row = &host_input[tok * n_embd..(tok + 1) * n_embd];
4957 crate::cpu_experts::predictor_submit(il, row);
4958 }
4959 if cpu_selected.is_empty() {
4960 None
4961 } else {
4962 let row = &host_input[tok * n_embd..(tok + 1) * n_embd];
4963 let job = crate::cpu_experts::prepare_job(m, il, &cpu_selected, row)
4964 .map_err(std::io::Error::other)?;
4965 Some(crate::cpu_experts::submit(job).map_err(std::io::Error::other)?)
4966 }
4967 } else {
4968 None
4969 };
4970
4971 let worker_window = worker_disk_prefetch
4972 .then(worker_prefetch_window)
4973 .unwrap_or(0);
4974 for (j, &ex) in sel.iter().enumerate() {
4975 if cpu_mask[j] {
4976 continue;
4977 }
4978 let ex = ex as usize;
4979 if let Some(d) = slab_local {
4986 let gl = m.gate_exps.expert_layout(ex);
4987 let ul = m.up_exps.expert_layout(ex);
4988 let dl = m.down_exps.expert_layout(ex);
4989 let (g0, u0, d0) = (
4990 ex * m.gate_exps.expert_stride,
4991 ex * m.up_exps.expert_stride,
4992 ex * m.down_exps.expert_stride,
4993 );
4994 let (gate, up) = if moe_q8 {
4995 if tok_q8.is_none() {
4996 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
4997 }
4998 let (zq, zd) = tok_q8.as_ref().unwrap();
4999 (
5000 e.qmatvec_expert_q8(
5001 &d.gate,
5002 g0..g0 + gl.len,
5003 zq,
5004 zd,
5005 1,
5006 m.gate_exps.in_f,
5007 m.gate_exps.out_f,
5008 gl.qtype,
5009 gl.row_bytes,
5010 )?,
5011 e.qmatvec_expert_q8(
5012 &d.up,
5013 u0..u0 + ul.len,
5014 zq,
5015 zd,
5016 1,
5017 m.up_exps.in_f,
5018 m.up_exps.out_f,
5019 ul.qtype,
5020 ul.row_bytes,
5021 )?,
5022 )
5023 } else {
5024 (
5025 e.qmatvec_view(
5026 &d.gate,
5027 g0..g0 + gl.len,
5028 &zt,
5029 1,
5030 m.gate_exps.in_f,
5031 m.gate_exps.out_f,
5032 gl.qtype,
5033 gl.row_bytes,
5034 )?,
5035 e.qmatvec_view(
5036 &d.up,
5037 u0..u0 + ul.len,
5038 &zt,
5039 1,
5040 m.up_exps.in_f,
5041 m.up_exps.out_f,
5042 ul.qtype,
5043 ul.row_bytes,
5044 )?,
5045 )
5046 };
5047 let mut act = e.uninit(n_ff_exp)?;
5048 Self::ffn_act_lim(
5049 e,
5050 cfg,
5051 &gate,
5052 &up,
5053 m.gate_exps.macro_scale(ex),
5054 m.up_exps.macro_scale(ex),
5055 lim_exp,
5056 &mut act,
5057 n_ff_exp,
5058 )?;
5059 let y = if moe_q8 {
5060 let (aq2, ad2) = e.quantize_q8_1(&act, 1, n_ff_exp)?;
5061 e.qmatvec_expert_q8(
5062 &d.down,
5063 d0..d0 + dl.len,
5064 &aq2,
5065 &ad2,
5066 1,
5067 m.down_exps.in_f,
5068 m.down_exps.out_f,
5069 dl.qtype,
5070 dl.row_bytes,
5071 )?
5072 } else {
5073 let actv = act.slice(0..n_ff_exp);
5074 e.qmatvec_view(
5075 &d.down,
5076 d0..d0 + dl.len,
5077 &actv,
5078 1,
5079 m.down_exps.in_f,
5080 m.down_exps.out_f,
5081 dl.qtype,
5082 dl.row_bytes,
5083 )?
5084 };
5085 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
5086 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
5087 continue;
5088 }
5089 for next in page_prefetch_positions(j, sel.len(), page_window) {
5090 Self::moe_prefetch_host_expert(sel[next] as usize, m);
5091 }
5092 let keep = [
5093 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_GATE, ex as u16),
5094 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_UP, ex as u16),
5095 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_DOWN, ex as u16),
5096 ];
5097 if worker_disk_prefetch && worker_window > 0 {
5098 for next in worker_prefetch_positions(j, sel.len(), worker_window) {
5099 Self::moe_prefetch_disk_expert(
5100 e,
5101 il,
5102 sel[next] as usize,
5103 m,
5104 max_block,
5105 &keep,
5106 )?;
5107 }
5108 } else if cache_dispatch
5109 && !cpu_hybrid
5110 && moe_prefetch_enabled()
5111 && j + 1 < sel.len()
5112 {
5113 let next = sel[j + 1] as usize;
5114 Self::moe_prefetch_expert(e, il, next, m, max_block, &keep)?;
5115 }
5116 let [gate_q8, up_q8, down_q8] = [moe_q8; 3];
5117 if cache_dispatch && (gate_q8 || up_q8 || down_q8) {
5118 if (gate_q8 || up_q8) && tok_q8.is_none() {
5121 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
5122 }
5123 let gate = if gate_q8 {
5124 let (zq, zd) = tok_q8.as_ref().unwrap();
5125 Self::moe_cached_gemm_q8(e, il, PROJ_GATE, ex, m, max_block, zq, zd)?
5126 } else {
5127 Self::moe_cached_gemm(e, il, PROJ_GATE, ex, m, max_block, &zt)?
5128 };
5129 let up = if up_q8 {
5130 let (zq, zd) = tok_q8.as_ref().unwrap();
5131 Self::moe_cached_gemm_q8(e, il, PROJ_UP, ex, m, max_block, zq, zd)?
5132 } else {
5133 Self::moe_cached_gemm(e, il, PROJ_UP, ex, m, max_block, &zt)?
5134 };
5135 let mut act = e.uninit(n_ff_exp)?;
5136 Self::ffn_act_lim(
5137 e,
5138 cfg,
5139 &gate,
5140 &up,
5141 m.gate_exps.macro_scale(ex),
5142 m.up_exps.macro_scale(ex),
5143 lim_exp,
5144 &mut act,
5145 n_ff_exp,
5146 )?;
5147 let y = if down_q8 {
5148 let (aq2, ad2) = e.quantize_q8_1(&act, 1, n_ff_exp)?;
5149 Self::moe_cached_gemm_q8(e, il, PROJ_DOWN, ex, m, max_block, &aq2, &ad2)?
5150 } else {
5151 let actv = act.slice(0..n_ff_exp);
5152 Self::moe_cached_gemm(e, il, PROJ_DOWN, ex, m, max_block, &actv)?
5153 };
5154 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
5155 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
5157 } else if cache_dispatch {
5158 let gate = Self::moe_cached_gemm(e, il, PROJ_GATE, ex, m, max_block, &zt)?;
5163 let up = Self::moe_cached_gemm(e, il, PROJ_UP, ex, m, max_block, &zt)?;
5164 let mut act = e.uninit(n_ff_exp)?; Self::ffn_act_lim(
5166 e,
5167 cfg,
5168 &gate,
5169 &up,
5170 m.gate_exps.macro_scale(ex),
5171 m.up_exps.macro_scale(ex),
5172 lim_exp,
5173 &mut act,
5174 n_ff_exp,
5175 )?;
5176 let actv = act.slice(0..n_ff_exp);
5177 let y = Self::moe_cached_gemm(e, il, PROJ_DOWN, ex, m, max_block, &actv)?;
5178 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
5179 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
5181 } else if cache_frozen {
5182 let gate = Self::moe_frozen_gemm(
5187 e,
5188 il,
5189 PROJ_GATE,
5190 ex,
5191 m,
5192 max_block,
5193 &zt,
5194 &mut scratch_g,
5195 g_len,
5196 )?;
5197 let up = Self::moe_frozen_gemm(
5198 e,
5199 il,
5200 PROJ_UP,
5201 ex,
5202 m,
5203 max_block,
5204 &zt,
5205 &mut scratch_u,
5206 u_len,
5207 )?;
5208 let mut act = e.uninit(n_ff_exp)?;
5209 Self::ffn_act_lim(
5210 e,
5211 cfg,
5212 &gate,
5213 &up,
5214 m.gate_exps.macro_scale(ex),
5215 m.up_exps.macro_scale(ex),
5216 lim_exp,
5217 &mut act,
5218 n_ff_exp,
5219 )?;
5220 let actv = act.slice(0..n_ff_exp);
5221 let y = Self::moe_frozen_gemm(
5222 e,
5223 il,
5224 PROJ_DOWN,
5225 ex,
5226 m,
5227 max_block,
5228 &actv,
5229 &mut scratch_d,
5230 d_len,
5231 )?;
5232 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
5233 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
5234 } else {
5235 if scratch_g.is_none() {
5239 scratch_g = Some(e.alloc_u8_uninit(g_len)?);
5240 scratch_u = Some(e.alloc_u8_uninit(u_len)?);
5241 scratch_d = Some(e.alloc_u8_uninit(d_len)?);
5242 }
5243 let (sg, su, sd) = (
5244 scratch_g.as_mut().unwrap(),
5245 scratch_u.as_mut().unwrap(),
5246 scratch_d.as_mut().unwrap(),
5247 );
5248 let gl = m.gate_exps.expert_layout(ex);
5249 let ul = m.up_exps.expert_layout(ex);
5250 let dl = m.down_exps.expert_layout(ex);
5251 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
5252 let gate = e.qmatvec_view(
5253 sg,
5254 0..gl.len,
5255 &zt,
5256 1,
5257 m.gate_exps.in_f,
5258 m.gate_exps.out_f,
5259 gl.qtype,
5260 gl.row_bytes,
5261 )?;
5262
5263 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
5264 let up = e.qmatvec_view(
5265 su,
5266 0..ul.len,
5267 &zt,
5268 1,
5269 m.up_exps.in_f,
5270 m.up_exps.out_f,
5271 ul.qtype,
5272 ul.row_bytes,
5273 )?;
5274
5275 let mut act = e.uninit(n_ff_exp)?; Self::ffn_act_lim(
5277 e,
5278 cfg,
5279 &gate,
5280 &up,
5281 m.gate_exps.macro_scale(ex),
5282 m.up_exps.macro_scale(ex),
5283 lim_exp,
5284 &mut act,
5285 n_ff_exp,
5286 )?;
5287
5288 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
5289 let actv = act.slice(0..n_ff_exp);
5290 let y = e.qmatvec_view(
5291 sd,
5292 0..dl.len,
5293 &actv,
5294 1,
5295 m.down_exps.in_f,
5296 m.down_exps.out_f,
5297 dl.qtype,
5298 dl.row_bytes,
5299 )?;
5300
5301 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
5302 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
5303 }
5304 }
5305 if let Some(worker) = cpu_worker {
5306 let cpu_output = worker.wait().map_err(std::io::Error::other)?;
5307 let cpu_output = e.htod(&cpu_output)?;
5308 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
5309 e.axpy_into(&cpu_output, 1.0, &mut dst, n_embd)?;
5310 }
5311 if cpu_hybrid && !cache_frozen && cpu_expert_profile_admit_enabled() {
5312 for (j, &ex) in sel.iter().enumerate() {
5313 if cpu_mask[j] {
5314 Self::moe_profile_admit_expert(e, il, ex as usize, m, max_block)?;
5315 }
5316 }
5317 }
5318 }
5319
5320 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
5325 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
5326 {
5327 let n_ff_sh = gate_shexp.out_features(); let verify_t = t > 1 && t < PRIME_MIN_T;
5336 let (sg_gate, sg_up) = if t == 1 {
5337 match e.matmul_q8_fused2_x(gate_shexp, up_shexp, z)? {
5338 Some(pair) => pair,
5339 None => (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?),
5340 }
5341 } else if verify_t {
5342 (
5343 e.matmul_decode_exact(gate_shexp, z, t)?,
5344 e.matmul_decode_exact(up_shexp, z, t)?,
5345 )
5346 } else {
5347 (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?) };
5349 let mut sa = e.uninit(t * n_ff_sh)?; Self::ffn_act_lim(
5351 e,
5352 cfg,
5353 &sg_gate,
5354 &sg_up,
5355 1.0,
5356 1.0,
5357 lim_shexp,
5358 &mut sa,
5359 t * n_ff_sh,
5360 )?;
5361 let sh = if verify_t {
5362 e.matmul_decode_exact(down_shexp, &sa, t)?
5363 } else {
5364 e.matmul(down_shexp, &sa, t)?
5365 }; let g = match &m.gate_inp_shexp {
5379 Some(gate_inp_shexp) => {
5380 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
5381 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
5382 } else {
5383 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
5384 let mut g = e.uninit(t)?; e.sigmoid(&gs, &mut g, t)?;
5386 g
5387 }
5388 }
5389 None => e.htod(&vec![1.0f32; t])?,
5390 };
5391 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
5393 }
5394
5395 Ok(moe_out)
5396 }
5397
5398 pub fn stage1_h2d_per_token(&self) -> u64 {
5401 use crate::hybrid::Ffn;
5402 let n_used = self
5403 .cfg
5404 .moe
5405 .as_ref()
5406 .map(|m| m.expert_used_count as u64)
5407 .unwrap_or(0);
5408 let mut bytes = 0u64;
5409 for l in self.layers.iter() {
5410 if let Ffn::Moe(m) = &l.ffn {
5411 bytes += n_used
5412 * (m.gate_exps.max_expert_bytes()
5413 + m.up_exps.max_expert_bytes()
5414 + m.down_exps.max_expert_bytes()) as u64;
5415 }
5416 }
5417 bytes
5418 }
5419
5420 pub(crate) fn max_moe_block(&self) -> usize {
5424 use crate::hybrid::Ffn;
5425 let mut mx = 0usize;
5426 let mut scan = |ffn: &Ffn| {
5427 if let Ffn::Moe(m) = ffn {
5428 mx = mx
5429 .max(m.gate_exps.max_expert_bytes())
5430 .max(m.up_exps.max_expert_bytes())
5431 .max(m.down_exps.max_expert_bytes());
5432 }
5433 };
5434 for l in self.layers.iter() {
5435 scan(&l.ffn);
5436 }
5437 if let Some(mtp) = self.mtp.as_ref() {
5438 scan(&mtp.ffn);
5439 }
5440 mx
5441 }
5442
5443 pub(crate) fn moe_cache_block_sizes(&self) -> Vec<usize> {
5446 use crate::hybrid::Ffn;
5447 let mut sizes = Vec::new();
5448 let mut scan = |ffn: &Ffn| {
5449 let Ffn::Moe(m) = ffn else { return };
5450 for ex in 0..m.gate_exps.n_expert {
5451 if m.active_experts.as_ref().is_some_and(|active| !active[ex]) {
5452 continue;
5453 }
5454 for exps in [&m.gate_exps, &m.up_exps, &m.down_exps] {
5455 let len = exps.expert_layout(ex).len;
5456 if len > 0 {
5457 sizes.push(len);
5458 }
5459 }
5460 }
5461 };
5462 for layer in &self.layers {
5463 scan(&layer.ffn);
5464 }
5465 if let Some(mtp) = &self.mtp {
5466 scan(&mtp.ffn);
5467 }
5468 sizes
5469 }
5470
5471 pub fn save_cpu_expert_residency_profile(
5477 &self,
5478 e: &Engine,
5479 path: &std::path::Path,
5480 ) -> Result<(), Box<dyn std::error::Error>> {
5481 let Some(ids) = e.export_moe_residency() else {
5482 return Err("no MoE residency cache to persist".into());
5483 };
5484 let mut body = format!(
5485 "memra-freeze-profile v1 max_block={} blocks={}\n",
5486 self.max_moe_block(),
5487 ids.len()
5488 );
5489 for (layer, proj, ex) in &ids {
5490 body.push_str(&format!("{layer} {proj} {ex}\n"));
5491 }
5492 let tmp = path.with_extension("tmp");
5493 std::fs::write(&tmp, body)?;
5494 std::fs::rename(&tmp, path)?;
5495 println!(
5496 "[moe-cache] freeze profile saved: {} blocks -> {}",
5497 ids.len(),
5498 path.display()
5499 );
5500 Ok(())
5501 }
5502
5503 pub fn restore_cpu_expert_residency_profile(
5507 &self,
5508 e: &Engine,
5509 path: &std::path::Path,
5510 ) -> Result<bool, Box<dyn std::error::Error>> {
5511 use crate::hybrid::Ffn;
5512 use crate::moe_cache::BlockId;
5513 let Ok(content) = std::fs::read_to_string(path) else {
5514 return Ok(false);
5515 };
5516 let mut lines = content.lines();
5517 let Some(header) = lines.next() else {
5518 return Ok(false);
5519 };
5520 let expected = format!("memra-freeze-profile v1 max_block={}", self.max_moe_block());
5521 if !header.starts_with(&expected) {
5522 println!(
5523 "[moe-cache] freeze profile ignored (geometry mismatch): {}",
5524 path.display()
5525 );
5526 return Ok(false);
5527 }
5528 let mut by_layer: std::collections::HashMap<u16, Vec<BlockId>> =
5529 std::collections::HashMap::new();
5530 for line in lines {
5531 let mut fields = line.split_whitespace();
5532 let (Some(layer), Some(proj), Some(ex)) = (fields.next(), fields.next(), fields.next())
5533 else {
5534 continue;
5535 };
5536 let (Ok(layer), Ok(proj), Ok(ex)) =
5537 (layer.parse::<u16>(), proj.parse::<u8>(), ex.parse::<u16>())
5538 else {
5539 continue;
5540 };
5541 by_layer
5542 .entry(layer)
5543 .or_default()
5544 .push(BlockId::new(layer, proj, ex));
5545 }
5546 let requested: usize = by_layer.values().map(Vec::len).sum();
5547 if requested == 0 {
5548 return Ok(false);
5549 }
5550 let max_block = self.max_moe_block();
5551 let mut restaged = 0usize;
5552 let mut stage_layer =
5553 |layer_index: u16, ffn: &Ffn| -> Result<(), Box<dyn std::error::Error>> {
5554 let Ffn::Moe(m) = ffn else { return Ok(()) };
5555 let Some(ids) = by_layer.get(&layer_index) else {
5556 return Ok(());
5557 };
5558 e.with_moe_cache(max_block, |cache, eng| {
5559 for id in ids {
5560 if cache.restage_block(*id, m, eng)? {
5561 restaged += 1;
5562 }
5563 }
5564 Ok(())
5565 })
5566 };
5567 for (index, layer) in self.layers.iter().enumerate() {
5568 stage_layer(index as u16, &layer.ffn)?;
5569 }
5570 if let Some(mtp) = self.mtp.as_ref() {
5571 stage_layer(u16::MAX, &mtp.ffn)?;
5572 }
5573 e.freeze_moe_cache();
5574 println!(
5575 "[moe-cache] freeze profile restored: {restaged}/{requested} blocks restaged from {}",
5576 path.display()
5577 );
5578 Ok(true)
5579 }
5580
5581 pub fn freeze_cpu_expert_residency(
5583 &self,
5584 e: &Engine,
5585 ) -> Result<(), Box<dyn std::error::Error>> {
5586 e.freeze_moe_cache();
5587 Ok(())
5588 }
5589
5590 pub fn ffn_act(
5598 e: &Engine,
5599 cfg: &ModelConfig,
5600 gate: &CudaSlice<f32>,
5601 up: &CudaSlice<f32>,
5602 act: &mut CudaSlice<f32>,
5603 n: usize,
5604 ) -> Result<(), Box<dyn std::error::Error>> {
5605 Self::ffn_act_scaled(e, cfg, gate, up, 1.0, 1.0, act, n)
5606 }
5607
5608 #[allow(clippy::too_many_arguments)]
5612 pub(crate) fn ffn_act_scaled(
5613 e: &Engine,
5614 cfg: &ModelConfig,
5615 gate: &CudaSlice<f32>,
5616 up: &CudaSlice<f32>,
5617 gs: f32,
5618 us: f32,
5619 act: &mut CudaSlice<f32>,
5620 n: usize,
5621 ) -> Result<(), Box<dyn std::error::Error>> {
5622 Self::ffn_act_lim(e, cfg, gate, up, gs, us, None, act, n)
5623 }
5624
5625 #[allow(clippy::too_many_arguments)]
5634 pub(crate) fn ffn_act_lim(
5635 e: &Engine,
5636 cfg: &ModelConfig,
5637 gate: &CudaSlice<f32>,
5638 up: &CudaSlice<f32>,
5639 gs: f32,
5640 us: f32,
5641 limit: Option<f32>,
5642 act: &mut CudaSlice<f32>,
5643 n: usize,
5644 ) -> Result<(), Box<dyn std::error::Error>> {
5645 if let Some(m3) = cfg.m3.as_ref() {
5646 debug_assert!(
5647 limit.is_none(),
5648 "m3 swigluoai and step35 clamp are different archs"
5649 );
5650 return e.swigluoai_mul_scaled(
5651 gate,
5652 up,
5653 gs,
5654 us,
5655 m3.swiglu_alpha,
5656 m3.swiglu_limit,
5657 act,
5658 n,
5659 );
5660 }
5661 if let Some(l) = limit {
5662 return e.swiglu_clamped_mul_scaled(gate, up, gs, us, l, act, n);
5663 }
5664 if gs == 1.0 && us == 1.0 {
5665 return e.silu_mul(gate, up, act, n);
5666 }
5667 e.silu_mul_scaled(gate, up, gs, us, act, n)
5668 }
5669
5670 fn moe_route(
5676 e: &Engine,
5677 logits: &CudaSlice<f32>,
5678 t: usize,
5679 n_expert: usize,
5680 n_used: usize,
5681 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
5682 Self::moe_route_cfg(e, logits, t, n_expert, n_used, None)
5683 }
5684
5685 #[allow(clippy::too_many_arguments)]
5693 fn moe_route_sigmoid_cfg(
5694 e: &Engine,
5695 logits: &CudaSlice<f32>,
5696 t: usize,
5697 n_expert: usize,
5698 n_used: usize,
5699 m: &MoeWeights,
5700 (sf, route_norm): (f32, bool),
5701 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
5702 if sigmoid_router_enabled() {
5703 return e.moe_router_sigmoid_topk_host(
5704 logits,
5705 t,
5706 n_expert,
5707 n_used,
5708 m.active_count(),
5709 &m.exp_probs_b_dev,
5710 &m.active_experts_dev,
5711 sf,
5712 route_norm,
5713 );
5714 }
5715 let lg = e.dtoh(logits)?;
5716 Self::moe_route_sigmoid_host(
5717 &lg,
5718 t,
5719 n_expert,
5720 n_used,
5721 m.exp_probs_b.as_deref(),
5722 sf,
5723 route_norm,
5724 m.active_experts.as_deref(),
5725 )
5726 }
5727
5728 fn moe_route_cfg(
5731 e: &Engine,
5732 logits: &CudaSlice<f32>,
5733 t: usize,
5734 n_expert: usize,
5735 n_used: usize,
5736 active: Option<&[bool]>,
5737 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
5738 if active.is_none() && !matches!(std::env::var("MEMRA_FUSED_ROUTER").as_deref(), Ok("0")) {
5741 return e.moe_router_topk_host(logits, t, n_expert, n_used);
5742 }
5743 let lg = e.dtoh(logits)?; let mut sel = vec![0u32; t * n_used];
5746 let mut w_out = vec![0f32; t * n_used];
5747 for tok in 0..t {
5748 let row = &lg[tok * n_expert..(tok + 1) * n_expert];
5749 let maxl = row
5751 .iter()
5752 .enumerate()
5753 .filter(|(i, _)| active.is_none_or(|mask| mask[*i]))
5754 .map(|(_, &x)| x)
5755 .fold(f32::NEG_INFINITY, f32::max);
5756 let mut probs = vec![0f32; n_expert];
5757 let mut den = 0f32;
5758 for i in 0..n_expert {
5759 if active.is_some_and(|mask| !mask[i]) {
5760 continue;
5761 }
5762 let x = (row[i] - maxl).exp();
5763 probs[i] = x;
5764 den += x;
5765 }
5766 for p in probs.iter_mut() {
5767 *p /= den;
5768 }
5769 let mut idx: Vec<usize> = (0..n_expert)
5771 .filter(|&i| active.is_none_or(|mask| mask[i]))
5772 .collect();
5773 idx.sort_by(|&a, &b| probs[b].total_cmp(&probs[a]).then(a.cmp(&b)));
5774 let sl = &idx[..n_used];
5775 let mut wv: Vec<f32> = sl.iter().map(|&i| probs[i]).collect();
5776 let mut ws: f32 = wv.iter().sum();
5777 ws = ws.max(6.103515625e-5_f32); for x in wv.iter_mut() {
5779 *x /= ws;
5780 }
5781 for j in 0..n_used {
5782 sel[tok * n_used + j] = sl[j] as u32;
5783 w_out[tok * n_used + j] = wv[j];
5784 }
5785 }
5786 Ok((sel, w_out))
5787 }
5788
5789 #[allow(clippy::too_many_arguments)]
5790 fn moe_route_sigmoid_with_input(
5791 e: &Engine,
5792 logits: &CudaSlice<f32>,
5793 input: &CudaSlice<f32>,
5794 t: usize,
5795 n_expert: usize,
5796 n_used: usize,
5797 bias: Option<&[f32]>,
5798 (sf, route_norm): (f32, bool),
5799 active: Option<&[bool]>,
5800 ) -> Result<(Vec<u32>, Vec<f32>, Vec<f32>), Box<dyn std::error::Error>> {
5801 let (lg, input) = e.dtoh_pair(logits, input)?;
5802 let (sel, w) =
5803 Self::moe_route_sigmoid_host(&lg, t, n_expert, n_used, bias, sf, route_norm, active)?;
5804 Ok((sel, w, input))
5805 }
5806
5807 pub fn start_moe_prefetch_predictor(
5812 &self,
5813 e: &Engine,
5814 cfg: &ModelConfig,
5815 ) -> Result<(), Box<dyn std::error::Error>> {
5816 use crate::hybrid::Ffn;
5817 let Some(sig) = cfg.sigmoid_router() else {
5818 return Err("prefetch predictor requires a sigmoid-router arch".into());
5819 };
5820 let resident: std::collections::HashSet<(u16, u8, u16)> = e
5821 .export_moe_residency()
5822 .ok_or("prefetch predictor needs the frozen MoE residency cache")?
5823 .into_iter()
5824 .collect();
5825 let mut layers = Vec::new();
5826 for (index, layer) in self.layers.iter().enumerate() {
5827 let Ffn::Moe(m) = &layer.ffn else { continue };
5828 let crate::model::GpuTensor::Float { data, .. } = &m.gate_inp else {
5829 continue;
5830 };
5831 let router = e.dtoh(data)?;
5832 let n_expert = m.gate_exps.n_expert;
5833 let n_embd = m.gate_exps.in_f;
5834 if router.len() != n_embd * n_expert {
5835 continue;
5836 }
5837 let build = |exps: &crate::model::HostExps| {
5838 (0..n_expert)
5839 .map(|expert| crate::cpu_experts::predictor_projection(exps, expert))
5840 .collect::<Vec<_>>()
5841 };
5842 layers.push((
5843 index as u16,
5844 crate::cpu_experts::PredictLayerInit {
5845 router,
5846 bias: m.exp_probs_b.clone(),
5847 active: m.active_experts.clone(),
5848 n_embd,
5849 n_used: cfg
5850 .moe
5851 .as_ref()
5852 .map(|moe| moe.expert_used_count as usize)
5853 .ok_or("prefetch predictor requires MoE config")?,
5854 sig,
5855 weights_n_expert: n_expert,
5856 gate: build(&m.gate_exps),
5857 up: build(&m.up_exps),
5858 down: build(&m.down_exps),
5859 },
5860 ));
5861 }
5862 crate::cpu_experts::start_prefetch_predictor(layers, resident).map_err(|error| error.into())
5863 }
5864
5865 #[allow(clippy::too_many_arguments)]
5868 pub fn moe_route_sigmoid_host_public(
5869 logits: &[f32],
5870 t: usize,
5871 n_expert: usize,
5872 n_used: usize,
5873 bias: Option<&[f32]>,
5874 sf: f32,
5875 route_norm: bool,
5876 active: Option<&[bool]>,
5877 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
5878 Self::moe_route_sigmoid_host(logits, t, n_expert, n_used, bias, sf, route_norm, active)
5879 }
5880
5881 #[allow(clippy::too_many_arguments)]
5882 fn moe_route_sigmoid_host(
5883 lg: &[f32],
5884 t: usize,
5885 n_expert: usize,
5886 n_used: usize,
5887 bias: Option<&[f32]>,
5888 sf: f32,
5889 route_norm: bool,
5890 active: Option<&[bool]>,
5891 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
5892 let active_count = active
5893 .map(|mask| mask.iter().filter(|&&enabled| enabled).count())
5894 .unwrap_or(n_expert);
5895 crate::sigrouter_contract::validate_active_count(n_used, active_count)?;
5896 if lg.len() != t * n_expert {
5897 return Err(format!(
5898 "sigmoid router logits length mismatch: got {}, expected {}",
5899 lg.len(),
5900 t * n_expert,
5901 )
5902 .into());
5903 }
5904 let mut sel = vec![0u32; t * n_used];
5905 let mut w_out = vec![0f32; t * n_used];
5906 for tok in 0..t {
5907 let row = &lg[tok * n_expert..(tok + 1) * n_expert];
5908 let scores: Vec<f32> = row.iter().map(|&x| 1.0 / (1.0 + (-x).exp())).collect();
5909 let selsc: Vec<f32> = match bias {
5911 Some(b) => scores.iter().zip(b).map(|(s, bb)| s + bb).collect(),
5912 None => scores.clone(),
5913 };
5914 let mut idx: Vec<usize> = (0..n_expert)
5915 .filter(|&i| active.is_none_or(|mask| mask[i]))
5916 .collect();
5917 idx.sort_by(|&a, &b| selsc[b].total_cmp(&selsc[a]).then(a.cmp(&b)));
5918 let sl = &idx[..n_used];
5919 let mut wv: Vec<f32> = sl.iter().map(|&i| scores[i]).collect();
5920 if route_norm {
5921 let ws: f32 = wv.iter().sum::<f32>().max(1e-20);
5922 for x in wv.iter_mut() {
5923 *x = *x / ws * sf;
5924 }
5925 } else {
5926 for x in wv.iter_mut() {
5927 *x *= sf;
5928 }
5929 }
5930 for j in 0..n_used {
5931 sel[tok * n_used + j] = sl[j] as u32;
5932 w_out[tok * n_used + j] = wv[j];
5933 }
5934 }
5935 Ok((sel, w_out))
5936 }
5937
5938 #[allow(clippy::too_many_arguments)]
5942 fn moe_ffn_sigmoid_dev(
5943 e: &Engine,
5944 m: &MoeWeights,
5945 z: &CudaSlice<f32>,
5946 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
5947 logits: &CudaSlice<f32>,
5948 t: usize,
5949 cfg: &ModelConfig,
5950 il: u16,
5951 (scaling_factor, route_norm): (f32, bool),
5952 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5953 let moe = cfg.moe.as_ref().unwrap();
5954 let n_embd = cfg.n_embd as usize;
5955 let n_expert = moe.expert_count as usize;
5956 let n_used = moe.expert_used_count as usize;
5957 let n_ff_exp = moe.expert_ff_length as usize;
5958 let dev = m.dev_exps.as_ref().unwrap();
5959 debug_assert!(cfg.step35.is_some());
5960 debug_assert_eq!(dev.dev, e.ctx().ordinal());
5961 debug_assert!(m.has_uniform_expert_layout());
5962 debug_assert!(!m.has_macros);
5963
5964 let (sel_d, w_d) = e.moe_router_sigmoid_topk(
5965 logits,
5966 t,
5967 n_expert,
5968 n_used,
5969 m.active_count(),
5970 &m.exp_probs_b_dev,
5971 &m.active_experts_dev,
5972 scaling_factor,
5973 route_norm,
5974 )?;
5975 crate::moesd::record_device_routes(e, il, n_expert, n_used, &sel_d)?;
5976 let (gate_row_bytes, up_row_bytes) = if dev.gu_il {
5977 let combined = m.gate_exps.row_bytes + m.up_exps.row_bytes;
5978 (combined, combined)
5979 } else {
5980 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
5981 };
5982 let (zq, zd) = match (t, zq8) {
5983 (1, Some((q, d))) => (q.clone(), d.clone()),
5984 _ => e.quantize_q8_1(z, t, n_embd)?,
5985 };
5986 let n_pairs = t * n_used;
5987 let mut moe_out = if cfg.clamp_exp_at(il as u32).is_some() {
5988 let pair_tok: Vec<i32> = (0..n_pairs).map(|pair| (pair / n_used) as i32).collect();
5992 let pair_tok_d = e.htod_i32(&pair_tok)?;
5993 let gate = e.moe_pairs_matvec_q8(
5994 &dev.ptr_row,
5995 0,
5996 &pair_tok_d,
5997 &sel_d,
5998 &zq,
5999 &zd,
6000 n_embd,
6001 n_ff_exp,
6002 n_expert,
6003 n_pairs,
6004 m.gate_exps.qtype,
6005 gate_row_bytes,
6006 )?;
6007 let up = e.moe_pairs_matvec_q8(
6008 &dev.ptr_row,
6009 1,
6010 &pair_tok_d,
6011 &sel_d,
6012 &zq,
6013 &zd,
6014 n_embd,
6015 n_ff_exp,
6016 n_expert,
6017 n_pairs,
6018 m.up_exps.qtype,
6019 up_row_bytes,
6020 )?;
6021 let mut act = e.uninit(n_pairs * n_ff_exp)?;
6022 Self::ffn_act_lim(
6023 e,
6024 cfg,
6025 &gate,
6026 &up,
6027 1.0,
6028 1.0,
6029 cfg.clamp_exp_at(il as u32),
6030 &mut act,
6031 n_pairs * n_ff_exp,
6032 )?;
6033 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
6034 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
6035 let pair_self_d = e.htod_i32(&pair_self)?;
6036 let down = e.moe_pairs_matvec_q8(
6037 &dev.ptr_row,
6038 2,
6039 &pair_self_d,
6040 &sel_d,
6041 &aq2,
6042 &ad2,
6043 n_ff_exp,
6044 n_embd,
6045 n_expert,
6046 n_pairs,
6047 m.down_exps.qtype,
6048 m.down_exps.row_bytes,
6049 )?;
6050 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
6051 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
6052 let tok_off_d = e.htod_i32(&tok_off)?;
6053 let tok_ids_d = e.htod_i32(&tok_ids)?;
6054 let mut output = e.uninit(t * n_embd)?;
6055 e.moe_pairs_scatter(&down, &w_d, &tok_off_d, &tok_ids_d, &mut output, t, n_embd)?;
6056 output
6057 } else {
6058 let act = e.moe_gate_up_silu8_dev_q8_rows(
6059 &dev.ptr_row,
6060 &sel_d,
6061 &zq,
6062 &zd,
6063 t,
6064 n_embd,
6065 n_ff_exp,
6066 n_used,
6067 n_expert,
6068 m.gate_exps.qtype,
6069 m.up_exps.qtype,
6070 gate_row_bytes,
6071 up_row_bytes,
6072 &m.dev_macros,
6073 )?;
6074 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
6075 let mut output = e.uninit(t * n_embd)?;
6076 e.moe_down8_fma_dev_q8_rows_g(
6077 &dev.ptr_row,
6078 &sel_d,
6079 &w_d,
6080 &aq2,
6081 &ad2,
6082 &mut output,
6083 t,
6084 n_ff_exp,
6085 n_embd,
6086 n_used,
6087 n_expert,
6088 m.down_exps.qtype,
6089 m.down_exps.row_bytes,
6090 )?;
6091 output
6092 };
6093
6094 if std::env::var("MEMRA_SIG_ROUTER_DISPATCH_TRACE").as_deref() == Ok("1") {
6095 eprintln!(
6096 "[sigrouter-dev] layer={il} tokens={t} experts={n_expert} used={n_used} clamp={} gu_il={}",
6097 cfg.clamp_exp_at(il as u32).is_some(),
6098 dev.gu_il,
6099 );
6100 }
6101 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut moe_out)?;
6102 Ok(moe_out)
6103 }
6104
6105 fn moe_ffn_pairs(
6114 e: &Engine,
6115 m: &MoeWeights,
6116 z: &CudaSlice<f32>,
6117 logits: &CudaSlice<f32>,
6118 t: usize,
6119 cfg: &ModelConfig,
6120 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6121 let moe = cfg.moe.as_ref().unwrap();
6122 let n_embd = cfg.n_embd as usize;
6123 let n_expert = moe.expert_count as usize;
6124 let n_used = moe.expert_used_count as usize;
6125 let n_ff_exp = moe.expert_ff_length as usize;
6126 debug_assert!(
6131 !cfg.swiglu_clamped_anywhere(),
6132 "moe_ffn_pairs has no per-layer clamp: fused epilogues are plain SiLU"
6133 );
6134 let dev = m.dev_exps.as_ref().unwrap();
6135 let (rbg_d, rbu_d) = if dev.gu_il {
6137 let sxx = m.gate_exps.row_bytes + m.up_exps.row_bytes;
6138 (sxx, sxx)
6139 } else {
6140 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
6141 };
6142
6143 let (sel_all, w_all) = Self::moe_route(e, logits, t, n_expert, n_used)?;
6144 let n_pairs = t * n_used;
6145 let pair_tok: Vec<i32> = (0..n_pairs).map(|p| (p / n_used) as i32).collect();
6148 let pair_ex: Vec<i32> = sel_all.iter().map(|&x| x as i32).collect();
6149 let pair_w: Vec<f32> = w_all.clone();
6150 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
6151 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
6152 let pt = e.htod_i32(&pair_tok)?;
6153 let px = e.htod_i32(&pair_ex)?;
6154 let pw = e.htod(&pair_w)?;
6155 let toff = e.htod_i32(&tok_off)?;
6156 let tids = e.htod_i32(&tok_ids)?;
6157
6158 let mut by_ex: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
6162 for p in 0..n_pairs {
6163 by_ex[pair_ex[p] as usize].push(p as i32);
6164 }
6165 let mut ex_ids: Vec<i32> = Vec::new();
6166 let mut ex_off: Vec<i32> = vec![0];
6167 let mut ex_pairs: Vec<i32> = Vec::with_capacity(n_pairs);
6168 for (ex, list) in by_ex.iter().enumerate() {
6169 if list.is_empty() {
6170 continue;
6171 }
6172 ex_ids.push(ex as i32);
6173 ex_pairs.extend_from_slice(list);
6174 ex_off.push(ex_pairs.len() as i32);
6175 }
6176 let n_active = ex_ids.len();
6177 let exi = e.htod_i32(&ex_ids)?;
6178 let exo = e.htod_i32(&ex_off)?;
6179 let exp_d = e.htod_i32(&ex_pairs)?;
6180 let _ = &px; static MMA_T: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
6201 let mma_t = *MMA_T.get_or_init(|| {
6202 std::env::var("MEMRA_MOE_MMA_T")
6203 .ok()
6204 .and_then(|v| v.parse().ok())
6205 .unwrap_or(16)
6206 });
6207 let use_mma = std::env::var("MEMRA_MOE_MMA")
6208 .map(|v| v != "0")
6209 .unwrap_or(true)
6210 && t >= mma_t
6211 && q8_expert_dec_supported(m.gate_exps.qtype)
6212 && q8_expert_dec_supported(m.up_exps.qtype)
6213 && q8_expert_dec_supported(m.down_exps.qtype)
6214 && n_embd % 256 == 0
6215 && n_ff_exp % 256 == 0;
6216 let mma_capable = q8_expert_dec_supported(m.gate_exps.qtype)
6232 && q8_expert_dec_supported(m.up_exps.qtype)
6233 && q8_expert_dec_supported(m.down_exps.qtype)
6234 && n_embd % 256 == 0
6235 && n_ff_exp % 256 == 0;
6236 let f16g_mode = crate::moe_f16g_mode();
6237 let f16g = f16g_mode != 0
6238 && t >= mma_t
6239 && (f16g_mode != 3 || !mma_capable)
6240 && f16g_proj_ok(m.gate_exps.qtype, n_embd)
6241 && f16g_proj_ok(m.up_exps.qtype, n_embd)
6242 && f16g_proj_ok(m.down_exps.qtype, n_ff_exp);
6243 if use_mma || f16g {
6244 let y_down = if f16g {
6252 let csr_tok: Vec<i32> = ex_pairs.iter().map(|&p| p / n_used as i32).collect();
6256 let csr_tok_d = e.htod_i32(&csr_tok)?;
6257 let (z_f16, z_s) = e.moe_f16g_act(z, Some(&csr_tok_d), n_embd, n_pairs)?;
6258 let g_csr = e.moe_f16_grouped(
6259 &dev.ptr_row,
6260 0,
6261 n_expert,
6262 &exi,
6263 &ex_off,
6264 &exo,
6265 &z_f16,
6266 &z_s,
6267 n_embd,
6268 n_ff_exp,
6269 n_active,
6270 n_pairs,
6271 m.gate_exps.qtype,
6272 rbg_d,
6273 )?;
6274 let u_csr = e.moe_f16_grouped(
6275 &dev.ptr_row,
6276 1,
6277 n_expert,
6278 &exi,
6279 &ex_off,
6280 &exo,
6281 &z_f16,
6282 &z_s,
6283 n_embd,
6284 n_ff_exp,
6285 n_active,
6286 n_pairs,
6287 m.up_exps.qtype,
6288 rbu_d,
6289 )?;
6290 let act_csr = e.moe_pairs_silu_mul(&g_csr, &u_csr, n_pairs * n_ff_exp)?;
6291 let (a_f16, a_s) = e.moe_f16g_act(&act_csr, None, n_ff_exp, n_pairs)?;
6292 let d_csr = e.moe_f16_grouped(
6293 &dev.ptr_row,
6294 2,
6295 n_expert,
6296 &exi,
6297 &ex_off,
6298 &exo,
6299 &a_f16,
6300 &a_s,
6301 n_ff_exp,
6302 n_embd,
6303 n_active,
6304 n_pairs,
6305 m.down_exps.qtype,
6306 m.down_exps.row_bytes,
6307 )?;
6308 e.rows_permute(&d_csr, &exp_d, n_pairs, n_embd)?
6309 } else {
6310 let z_scr = e.mmq_iq_quantize_act(z, n_embd, t)?;
6312 let gate = e.mmq_iq_experts(
6313 &dev.ptr_row,
6314 0,
6315 n_expert,
6316 &exi,
6317 &exo,
6318 &exp_d,
6319 &pt,
6320 &z_scr,
6321 n_embd,
6322 n_ff_exp,
6323 n_active,
6324 n_pairs,
6325 t,
6326 m.gate_exps.qtype,
6327 rbg_d,
6328 )?;
6329 let up = e.mmq_iq_experts(
6330 &dev.ptr_row,
6331 1,
6332 n_expert,
6333 &exi,
6334 &exo,
6335 &exp_d,
6336 &pt,
6337 &z_scr,
6338 n_embd,
6339 n_ff_exp,
6340 n_active,
6341 n_pairs,
6342 t,
6343 m.up_exps.qtype,
6344 rbu_d,
6345 )?;
6346 let a_scr = if crate::moe_fuse_actq_on() {
6352 e.mmq_iq_fused_act_quant(&gate, &up, n_ff_exp, n_pairs, 0)?
6353 } else {
6354 let act = e.moe_pairs_silu_mul(&gate, &up, n_pairs * n_ff_exp)?;
6355 e.mmq_iq_quantize_act(&act, n_ff_exp, n_pairs)?
6356 };
6357 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
6358 let pself = e.htod_i32(&pair_self)?;
6359 e.mmq_iq_experts(
6360 &dev.ptr_row,
6361 2,
6362 n_expert,
6363 &exi,
6364 &exo,
6365 &exp_d,
6366 &pself,
6367 &a_scr,
6368 n_ff_exp,
6369 n_embd,
6370 n_active,
6371 n_pairs,
6372 n_pairs,
6373 m.down_exps.qtype,
6374 m.down_exps.row_bytes,
6375 )?
6376 };
6377 let mut moe_out = e.uninit(t * n_embd)?;
6378 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
6379 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
6380 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
6381 {
6382 let n_ff_sh = gate_shexp.out_features();
6383 let sg_gate = e.matmul(gate_shexp, z, t)?;
6384 let sg_up = e.matmul(up_shexp, z, t)?;
6385 let mut sa = e.uninit(t * n_ff_sh)?;
6386 Self::ffn_act(e, cfg, &sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
6387 let sh = e.matmul(down_shexp, &sa, t)?;
6388 let g = match &m.gate_inp_shexp {
6394 Some(gate_inp_shexp) if crate::router_prefill_exact_on() => {
6395 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
6396 }
6397 Some(gate_inp_shexp) => {
6398 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
6399 let mut g = e.uninit(t)?;
6400 e.sigmoid(&gs, &mut g, t)?;
6401 g
6402 }
6403 None => e.htod(&vec![1.0f32; t])?,
6404 };
6405 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
6406 }
6407 return Ok(moe_out);
6408 }
6409
6410 let dec = std::env::var("MEMRA_MOE_DEC")
6413 .map(|v| v != "0")
6414 .unwrap_or(true);
6415 let matvec = |proj,
6416 exi: &_,
6417 exo: &_,
6418 exp_d: &_,
6419 pt: &_,
6420 aq: &_,
6421 ad: &_,
6422 inf,
6423 outf,
6424 qtype,
6425 rb|
6426 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6427 let dec = dec && q8_expert_dec_supported(qtype);
6429 if dec {
6430 e.moe_pairs_matvec_q8_dec(
6431 &dev.ptr_row,
6432 proj,
6433 exi,
6434 exo,
6435 exp_d,
6436 pt,
6437 aq,
6438 ad,
6439 inf,
6440 outf,
6441 n_expert,
6442 n_active,
6443 n_pairs,
6444 qtype,
6445 rb,
6446 )
6447 } else {
6448 e.moe_pairs_matvec_q8_em(
6449 &dev.ptr_row,
6450 proj,
6451 exi,
6452 exo,
6453 exp_d,
6454 pt,
6455 aq,
6456 ad,
6457 inf,
6458 outf,
6459 n_expert,
6460 n_active,
6461 n_pairs,
6462 qtype,
6463 rb,
6464 )
6465 }
6466 };
6467 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
6468 let gate = matvec(
6469 0,
6470 &exi,
6471 &exo,
6472 &exp_d,
6473 &pt,
6474 &zq,
6475 &zd,
6476 n_embd,
6477 n_ff_exp,
6478 m.gate_exps.qtype,
6479 rbg_d,
6480 )?;
6481 let up = matvec(
6482 1,
6483 &exi,
6484 &exo,
6485 &exp_d,
6486 &pt,
6487 &zq,
6488 &zd,
6489 n_embd,
6490 n_ff_exp,
6491 m.up_exps.qtype,
6492 rbu_d,
6493 )?;
6494 let act = e.moe_pairs_silu_mul(&gate, &up, n_pairs * n_ff_exp)?;
6495 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
6496 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
6498 let pself = e.htod_i32(&pair_self)?;
6499 let y_down = matvec(
6500 2,
6501 &exi,
6502 &exo,
6503 &exp_d,
6504 &pself,
6505 &aq2,
6506 &ad2,
6507 n_ff_exp,
6508 n_embd,
6509 m.down_exps.qtype,
6510 m.down_exps.row_bytes,
6511 )?;
6512 let mut moe_out = e.uninit(t * n_embd)?; e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
6514
6515 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
6519 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
6520 {
6521 let n_ff_sh = gate_shexp.out_features();
6522 let step_exact = cfg.step35.is_some();
6526 let verify_t = step_exact && t > 1 && t < PRIME_MIN_T;
6527 let (sg_gate, sg_up) = if step_exact && t == 1 {
6528 match e.matmul_q8_fused2_x(gate_shexp, up_shexp, z)? {
6529 Some(pair) => pair,
6530 None => (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?),
6531 }
6532 } else if verify_t {
6533 let mut fused = None;
6534 if crate::spec::spec_fused_t()
6535 && (2..=4).contains(&t)
6536 && e.uses_q8_1_fast(gate_shexp)
6537 && e.uses_q8_1_fast(up_shexp)
6538 {
6539 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
6540 fused = e.matmul_q8_fused2_t(gate_shexp, up_shexp, &zq, &zd, t)?;
6541 }
6542 match fused {
6543 Some(pair) => pair,
6544 None => (
6545 e.matmul_decode_exact(gate_shexp, z, t)?,
6546 e.matmul_decode_exact(up_shexp, z, t)?,
6547 ),
6548 }
6549 } else {
6550 (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?)
6551 };
6552 let mut sa = e.uninit(t * n_ff_sh)?;
6553 e.silu_mul(&sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
6554 let sh = if verify_t {
6555 e.matmul_decode_exact(down_shexp, &sa, t)?
6556 } else {
6557 e.matmul(down_shexp, &sa, t)?
6558 };
6559 let g = match &m.gate_inp_shexp {
6564 Some(gate_inp_shexp) if crate::router_prefill_exact_on() => {
6565 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
6566 }
6567 Some(gate_inp_shexp) => {
6568 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
6569 let mut g = e.uninit(t)?;
6570 e.sigmoid(&gs, &mut g, t)?;
6571 g
6572 }
6573 None => e.htod(&vec![1.0f32; t])?,
6574 };
6575 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
6576 }
6577 Ok(moe_out)
6578 }
6579
6580 #[allow(clippy::too_many_arguments)]
6582 #[allow(clippy::too_many_arguments)]
6583 fn moe_ffn_dev(
6584 e: &Engine,
6585 m: &MoeWeights,
6586 z: &CudaSlice<f32>,
6587 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
6588 logits: &CudaSlice<f32>,
6589 t: usize,
6590 cfg: &ModelConfig,
6591 il: u16,
6592 max_block: usize,
6593 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6594 let moe = cfg.moe.as_ref().unwrap();
6595 let n_embd = cfg.n_embd as usize;
6596 let n_expert = moe.expert_count as usize;
6597 let n_used = moe.expert_used_count as usize;
6598 let n_ff_exp = moe.expert_ff_length as usize;
6599 debug_assert!(
6603 cfg.sigmoid_router().is_none(),
6604 "moe_ffn_dev routes SOFTMAX: a sigmoid-router arch would pick wrong experts"
6605 );
6606 debug_assert!(
6607 !cfg.swiglu_clamped_at(il as u32),
6608 "moe_ffn_dev's fused epilogue is plain SiLU: no clamped form"
6609 );
6610
6611 let (sel_d, mut w_d) = e.moe_router_topk(logits, t, n_expert, n_used)?;
6613 if m.has_macros {
6616 e.moe_w_scale_by_expert(&mut w_d, &sel_d, &m.dev_macros, n_expert, t * n_used)?;
6617 }
6618
6619 let mut moe_out = e.uninit(t * n_embd)?;
6621
6622 if let Some(dev) = m.dev_exps.as_ref() {
6625 let (rbg_d, rbu_d) = if dev.gu_il {
6628 let sxx = m.gate_exps.row_bytes + m.up_exps.row_bytes;
6629 (sxx, sxx)
6630 } else {
6631 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
6632 };
6633 let q8 = moe_q8_enabled()
6634 && q8_expert_supported(m.gate_exps.qtype)
6635 && q8_expert_supported(m.up_exps.qtype)
6636 && q8_expert_supported(m.down_exps.qtype);
6637 let rows_arm = q8
6646 && t > 1
6647 && crate::spec::spec_m2()
6648 && n_ff_exp == 512
6649 && n_used <= 8
6650 && std::env::var("MEMRA_MOE_DEVQ8_GU")
6651 .map(|v| v.is_empty() || v == "v")
6652 .unwrap_or(true)
6653 && std::env::var("MEMRA_MOE_DEVQ8_DOWN")
6654 .map(|v| v.is_empty() || v == "w8h2v")
6655 .unwrap_or(true);
6656 let csr_mode = std::env::var("MEMRA_MOE_CSR")
6665 .ok()
6666 .and_then(|v| v.parse::<i32>().ok())
6667 .unwrap_or(1);
6668 let csr_qt = |qt: i32| qt == crate::QT_IQ4_XS || qt == crate::QT_IQ3_S;
6669 let csr_arm = rows_arm
6670 && csr_mode > 0
6671 && t <= 10
6672 && csr_qt(m.gate_exps.qtype)
6673 && csr_qt(m.up_exps.qtype)
6674 && csr_qt(m.down_exps.qtype);
6675 if csr_arm {
6676 if csr_mode == 2 {
6677 static ENGAGED: std::sync::Once = std::sync::Once::new();
6678 ENGAGED.call_once(|| eprintln!("[memra] moe CSR byte-compare mode ON (t={t})"));
6679 }
6680 let n_pairs = t * n_used;
6681 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
6682 let act = e.moe_gate_up_silu8_dev_q8_csr(
6683 &dev.ptr_row,
6684 &sel_d,
6685 &zq,
6686 &zd,
6687 n_pairs,
6688 n_embd,
6689 n_ff_exp,
6690 n_used,
6691 n_expert,
6692 m.gate_exps.qtype,
6693 m.up_exps.qtype,
6694 rbg_d,
6695 rbu_d,
6696 )?;
6697 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
6698 e.moe_down8_fma_dev_q8_rows(
6702 &dev.ptr_row,
6703 &sel_d,
6704 &w_d,
6705 &aq2,
6706 &ad2,
6707 &mut moe_out,
6708 t,
6709 n_ff_exp,
6710 n_embd,
6711 n_used,
6712 n_expert,
6713 m.down_exps.qtype,
6714 m.down_exps.row_bytes,
6715 )?;
6716 if csr_mode == 2 {
6717 let act_r = e.moe_gate_up_silu8_dev_q8_rows(
6719 &dev.ptr_row,
6720 &sel_d,
6721 &zq,
6722 &zd,
6723 t,
6724 n_embd,
6725 n_ff_exp,
6726 n_used,
6727 n_expert,
6728 m.gate_exps.qtype,
6729 m.up_exps.qtype,
6730 rbg_d,
6731 rbu_d,
6732 &m.dev_macros,
6733 )?;
6734 let mut out_r = e.uninit(t * n_embd)?;
6735 let (aq2r, ad2r) = e.quantize_q8_1(&act_r, n_pairs, n_ff_exp)?;
6736 e.moe_down8_fma_dev_q8_rows(
6737 &dev.ptr_row,
6738 &sel_d,
6739 &w_d,
6740 &aq2r,
6741 &ad2r,
6742 &mut out_r,
6743 t,
6744 n_ff_exp,
6745 n_embd,
6746 n_used,
6747 n_expert,
6748 m.down_exps.qtype,
6749 m.down_exps.row_bytes,
6750 )?;
6751 let (a1, a2) = (e.dtoh(&act)?, e.dtoh(&act_r)?);
6752 let (o1, o2) = (e.dtoh(&moe_out)?, e.dtoh(&out_r)?);
6753 let ba = a1
6754 .iter()
6755 .zip(&a2)
6756 .filter(|(x, y)| x.to_bits() != y.to_bits())
6757 .count();
6758 let bo = o1
6759 .iter()
6760 .zip(&o2)
6761 .filter(|(x, y)| x.to_bits() != y.to_bits())
6762 .count();
6763 if ba + bo > 0 {
6764 eprintln!(
6765 "[csr-check] il={il} t={t} ACT diffs={ba}/{} OUT diffs={bo}/{}",
6766 a1.len(),
6767 o1.len()
6768 );
6769 let sel_h = e.dtoh_i32(&sel_d)?;
6771 let mut shown = 0;
6772 for (i, (x, y)) in a1.iter().zip(&a2).enumerate() {
6773 if x.to_bits() != y.to_bits() && shown < 4 {
6774 let (p, o) = (i / n_ff_exp, i % n_ff_exp);
6775 let ex = sel_h[p];
6776 let npx = sel_h.iter().filter(|&&v| v == ex).count();
6777 eprintln!(
6778 " ACT p={p} ex={ex} np={npx} o={o} csr={x:e} rows={y:e}"
6779 );
6780 shown += 1;
6781 }
6782 }
6783 std::process::exit(3);
6784 }
6785 }
6786 } else if rows_arm {
6787 if std::env::var("MEMRA_MOE_OVERLAP").as_deref() == Ok("1") {
6790 use std::sync::atomic::{AtomicU64, Ordering};
6791 static PAIRS: AtomicU64 = AtomicU64::new(0);
6792 static UNIQ: AtomicU64 = AtomicU64::new(0);
6793 static CALLS: AtomicU64 = AtomicU64::new(0);
6794 let sel_h = e.dtoh_i32(&sel_d)?;
6795 let mut u: Vec<i32> = sel_h.clone();
6796 u.sort_unstable();
6797 u.dedup();
6798 PAIRS.fetch_add(sel_h.len() as u64, Ordering::Relaxed);
6799 UNIQ.fetch_add(u.len() as u64, Ordering::Relaxed);
6800 let c = CALLS.fetch_add(1, Ordering::Relaxed) + 1;
6801 if c % 480 == 0 {
6802 let p = PAIRS.load(Ordering::Relaxed);
6803 let q = UNIQ.load(Ordering::Relaxed);
6804 eprintln!(
6805 "[overlap] calls={c} pairs={p} unique={q} ratio={:.3} (t={t})",
6806 q as f64 / p as f64
6807 );
6808 }
6809 }
6810 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
6811 let act = e.moe_gate_up_silu8_dev_q8_rows(
6812 &dev.ptr_row,
6813 &sel_d,
6814 &zq,
6815 &zd,
6816 t,
6817 n_embd,
6818 n_ff_exp,
6819 n_used,
6820 n_expert,
6821 m.gate_exps.qtype,
6822 m.up_exps.qtype,
6823 rbg_d,
6824 rbu_d,
6825 &m.dev_macros,
6826 )?;
6827 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
6828 e.moe_down8_fma_dev_q8_rows(
6829 &dev.ptr_row,
6830 &sel_d,
6831 &w_d,
6832 &aq2,
6833 &ad2,
6834 &mut moe_out,
6835 t,
6836 n_ff_exp,
6837 n_embd,
6838 n_used,
6839 n_expert,
6840 m.down_exps.qtype,
6841 m.down_exps.row_bytes,
6842 )?;
6843 } else {
6844 for tok in 0..t {
6845 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
6846 let selt = sel_d.slice(tok * n_used..(tok + 1) * n_used);
6847 let wt = w_d.slice(tok * n_used..(tok + 1) * n_used);
6848 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
6849 if q8 {
6850 let (zq, zd) = match (t, zq8) {
6851 (1, Some((q, d))) => (q.clone(), d.clone()),
6852 _ => e.quantize_q8_1_view(&zt, 1, n_embd)?,
6853 };
6854 let act = e.moe_gate_up_silu8_dev_q8(
6855 &dev.ptr_row,
6856 &selt,
6857 &zq,
6858 &zd,
6859 n_embd,
6860 n_ff_exp,
6861 n_used,
6862 n_expert,
6863 m.gate_exps.qtype,
6864 m.up_exps.qtype,
6865 rbg_d,
6866 rbu_d,
6867 &m.dev_macros,
6868 )?;
6869 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
6870 e.moe_down8_fma_dev_q8(
6871 &dev.ptr_row,
6872 &selt,
6873 &wt,
6874 &aq2,
6875 &ad2,
6876 &mut dst,
6877 n_ff_exp,
6878 n_embd,
6879 n_used,
6880 n_expert,
6881 m.down_exps.qtype,
6882 m.down_exps.row_bytes,
6883 )?;
6884 } else {
6885 let act = e.moe_gate_up_silu8_dev(
6886 &dev.ptr_row,
6887 &selt,
6888 &zt,
6889 n_embd,
6890 n_ff_exp,
6891 n_used,
6892 n_expert,
6893 m.gate_exps.qtype,
6894 m.up_exps.qtype,
6895 rbg_d,
6896 rbu_d,
6897 &m.dev_macros,
6898 )?;
6899 e.moe_down8_fma_dev(
6900 &dev.ptr_row,
6901 &selt,
6902 &wt,
6903 &act,
6904 &mut dst,
6905 n_ff_exp,
6906 n_embd,
6907 n_used,
6908 n_expert,
6909 m.down_exps.qtype,
6910 m.down_exps.row_bytes,
6911 )?;
6912 }
6913 }
6914 }
6915 } else {
6916 let q8 = moe_q8_enabled()
6923 && q8_expert_supported(m.gate_exps.qtype)
6924 && q8_expert_supported(m.up_exps.qtype)
6925 && q8_expert_supported(m.down_exps.qtype);
6926 e.with_moe_cache(max_block, |c, eng| {
6927 let row = c
6928 .layer_dev_row(il, n_expert, eng)?
6929 .ok_or("moe_ffn_dev: layer row vanished under the lock")?;
6930 for tok in 0..t {
6931 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
6932 let selt = sel_d.slice(tok * n_used..(tok + 1) * n_used);
6933 let wt = w_d.slice(tok * n_used..(tok + 1) * n_used);
6934 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
6935 if q8 {
6936 let (zq, zd) = match (t, zq8) {
6937 (1, Some((q, d))) => (q.clone(), d.clone()),
6938 _ => eng.quantize_q8_1_view(&zt, 1, n_embd)?,
6939 };
6940 let act = eng.moe_gate_up_silu8_dev_q8(
6941 row,
6942 &selt,
6943 &zq,
6944 &zd,
6945 n_embd,
6946 n_ff_exp,
6947 n_used,
6948 n_expert,
6949 m.gate_exps.qtype,
6950 m.up_exps.qtype,
6951 m.gate_exps.row_bytes,
6952 m.up_exps.row_bytes,
6953 &m.dev_macros,
6954 )?;
6955 let (aq2, ad2) = eng.quantize_q8_1(&act, n_used, n_ff_exp)?;
6956 eng.moe_down8_fma_dev_q8(
6957 row,
6958 &selt,
6959 &wt,
6960 &aq2,
6961 &ad2,
6962 &mut dst,
6963 n_ff_exp,
6964 n_embd,
6965 n_used,
6966 n_expert,
6967 m.down_exps.qtype,
6968 m.down_exps.row_bytes,
6969 )?;
6970 } else {
6971 let act = eng.moe_gate_up_silu8_dev(
6972 row,
6973 &selt,
6974 &zt,
6975 n_embd,
6976 n_ff_exp,
6977 n_used,
6978 n_expert,
6979 m.gate_exps.qtype,
6980 m.up_exps.qtype,
6981 m.gate_exps.row_bytes,
6982 m.up_exps.row_bytes,
6983 &m.dev_macros,
6984 )?;
6985 eng.moe_down8_fma_dev(
6986 row,
6987 &selt,
6988 &wt,
6989 &act,
6990 &mut dst,
6991 n_ff_exp,
6992 n_embd,
6993 n_used,
6994 n_expert,
6995 m.down_exps.qtype,
6996 m.down_exps.row_bytes,
6997 )?;
6998 }
6999 }
7000 c.hits += (t * 3 * n_used) as u64;
7002 Ok(())
7003 })?;
7004 }
7005
7006 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
7011 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
7012 {
7013 let n_ff_sh = gate_shexp.out_features();
7014 let verify_t = t > 1 && t < PRIME_MIN_T;
7017 let (sg_gate, sg_up) = if t == 1 {
7018 match e.matmul_q8_fused2_x(gate_shexp, up_shexp, z)? {
7019 Some(pair) => pair,
7020 None => (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?),
7021 }
7022 } else if verify_t {
7023 let mut fused = None;
7027 if crate::spec::spec_fused_t()
7028 && (2..=4).contains(&t)
7029 && e.uses_q8_1_fast(gate_shexp)
7030 && e.uses_q8_1_fast(up_shexp)
7031 {
7032 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
7033 fused = e.matmul_q8_fused2_t(gate_shexp, up_shexp, &zq, &zd, t)?;
7034 }
7035 match fused {
7036 Some(pair) => pair,
7037 None => (
7038 e.matmul_decode_exact(gate_shexp, z, t)?,
7039 e.matmul_decode_exact(up_shexp, z, t)?,
7040 ),
7041 }
7042 } else {
7043 (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?)
7044 };
7045 let mut sa = e.uninit(t * n_ff_sh)?; e.silu_mul(&sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
7047 let sh = if verify_t {
7048 e.matmul_decode_exact(down_shexp, &sa, t)?
7049 } else {
7050 e.matmul(down_shexp, &sa, t)?
7051 };
7052 let g = match &m.gate_inp_shexp {
7056 Some(gate_inp_shexp) => {
7057 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
7060 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
7061 } else {
7062 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
7063 let mut g = e.uninit(t)?;
7064 e.sigmoid(&gs, &mut g, t)?;
7065 g
7066 }
7067 }
7068 None => e.htod(&vec![1.0f32; t])?,
7069 };
7070 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
7071 }
7072
7073 Ok(moe_out)
7074 }
7075
7076 #[allow(clippy::too_many_arguments)]
7086 #[allow(clippy::too_many_arguments)]
7089 fn moe_gdec_token_q8(
7090 e: &Engine,
7091 m: &MoeWeights,
7092 il: u16,
7093 max_block: usize,
7094 zq: &CudaSlice<i8>,
7095 zd: &CudaSlice<f32>,
7096 sel: &[u32],
7097 w: &[f32],
7098 moe_out: &mut CudaSlice<f32>,
7099 tok: usize,
7100 n_embd: usize,
7101 n_ff_exp: usize,
7102 n_used: usize,
7103 ) -> Result<bool, Box<dyn std::error::Error>> {
7104 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
7105 use cudarc::driver::DevicePtr;
7106 let ptrs = e.with_moe_cache(max_block, |c, eng| {
7107 let mut g = [0u64; 8];
7108 let mut u = [0u64; 8];
7109 let mut d = [0u64; 8];
7110 for (j, &ex) in sel.iter().enumerate() {
7111 let ex = ex as u16;
7112 let (Some(sg), Some(su), Some(sd)) = (
7113 c.resident(BlockId::new(il, PROJ_GATE, ex)),
7114 c.resident(BlockId::new(il, PROJ_UP, ex)),
7115 c.resident(BlockId::new(il, PROJ_DOWN, ex)),
7116 ) else {
7117 return Ok(None);
7118 };
7119 let __s = eng.stream();
7120 let (pg, _e0) = c.slot(sg).device_ptr(&__s);
7121 let (pu, _e1) = c.slot(su).device_ptr(&__s);
7122 let (pd, _e2) = c.slot(sd).device_ptr(&__s);
7123 g[j] = pg as u64;
7124 u[j] = pu as u64;
7125 d[j] = pd as u64;
7126 }
7127 if cpu_expert_profile_admit_enabled() && !c.is_frozen() {
7128 for &ex in sel {
7129 let ex = ex as u16;
7130 for proj in [PROJ_GATE, PROJ_UP, PROJ_DOWN] {
7131 c.note_profile_hit(BlockId::new(il, proj, ex));
7132 }
7133 }
7134 }
7135 c.hits += (3 * n_used) as u64;
7136 Ok(Some((g, u, d)))
7137 })?;
7138 let Some((g, u, d)) = ptrs else {
7139 return Ok(false);
7140 };
7141 let mut wv = [0f32; 8];
7142 wv[..n_used].copy_from_slice(w);
7143 let act = e.moe_gate_up_silu8_q8(
7144 crate::WPtr8(g),
7145 crate::WPtr8(u),
7146 zq,
7147 zd,
7148 n_embd,
7149 n_ff_exp,
7150 n_used,
7151 m.gate_exps.qtype,
7152 m.up_exps.qtype,
7153 m.gate_exps.row_bytes,
7154 m.up_exps.row_bytes,
7155 )?;
7156 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
7158 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
7159 e.moe_down8_fma_q8(
7160 crate::WPtr8(d),
7161 crate::F32x8(wv),
7162 &aq2,
7163 &ad2,
7164 &mut dst,
7165 n_ff_exp,
7166 n_embd,
7167 n_used,
7168 m.down_exps.qtype,
7169 m.down_exps.row_bytes,
7170 )?;
7171 Ok(true)
7172 }
7173
7174 fn moe_gdec_token(
7175 e: &Engine,
7176 m: &MoeWeights,
7177 il: u16,
7178 max_block: usize,
7179 zt: &cudarc::driver::CudaView<f32>,
7180 sel: &[u32],
7181 w: &[f32],
7182 moe_out: &mut CudaSlice<f32>,
7183 tok: usize,
7184 n_embd: usize,
7185 n_ff_exp: usize,
7186 n_used: usize,
7187 ) -> Result<bool, Box<dyn std::error::Error>> {
7188 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
7189 use cudarc::driver::DevicePtr;
7190 let ptrs = e.with_moe_cache(max_block, |c, eng| {
7192 let mut g = [0u64; 8];
7193 let mut u = [0u64; 8];
7194 let mut d = [0u64; 8];
7195 for (j, &ex) in sel.iter().enumerate() {
7196 let ex = ex as u16;
7197 let (Some(sg), Some(su), Some(sd)) = (
7198 c.resident(BlockId::new(il, PROJ_GATE, ex)),
7199 c.resident(BlockId::new(il, PROJ_UP, ex)),
7200 c.resident(BlockId::new(il, PROJ_DOWN, ex)),
7201 ) else {
7202 return Ok(None);
7203 };
7204 let __s = eng.stream();
7205 let (pg, _e0) = c.slot(sg).device_ptr(&__s);
7206 let (pu, _e1) = c.slot(su).device_ptr(&__s);
7207 let (pd, _e2) = c.slot(sd).device_ptr(&__s);
7208 g[j] = pg as u64;
7209 u[j] = pu as u64;
7210 d[j] = pd as u64;
7211 }
7212 if cpu_expert_profile_admit_enabled() && !c.is_frozen() {
7213 for &ex in sel {
7214 let ex = ex as u16;
7215 for proj in [PROJ_GATE, PROJ_UP, PROJ_DOWN] {
7216 c.note_profile_hit(BlockId::new(il, proj, ex));
7217 }
7218 }
7219 }
7220 c.hits += (3 * n_used) as u64; Ok(Some((g, u, d)))
7222 })?;
7223 let Some((g, u, d)) = ptrs else {
7224 return Ok(false);
7225 };
7226 let mut wv = [0f32; 8];
7227 wv[..n_used].copy_from_slice(w);
7228 let act = e.moe_gate_up_silu8(
7230 crate::WPtr8(g),
7231 crate::WPtr8(u),
7232 zt,
7233 n_embd,
7234 n_ff_exp,
7235 n_used,
7236 m.gate_exps.qtype,
7237 m.up_exps.qtype,
7238 m.gate_exps.row_bytes,
7239 m.up_exps.row_bytes,
7240 )?;
7241 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
7242 e.moe_down8_fma_into(
7243 crate::WPtr8(d),
7244 crate::F32x8(wv),
7245 &act,
7246 &mut dst,
7247 n_ff_exp,
7248 n_embd,
7249 n_used,
7250 m.down_exps.qtype,
7251 m.down_exps.row_bytes,
7252 )?;
7253 Ok(true)
7254 }
7255
7256 fn moe_cached_gemm_q8(
7261 e: &Engine,
7262 il: u16,
7263 proj: u8,
7264 ex: usize,
7265 m: &MoeWeights,
7266 max_block: usize,
7267 aq: &CudaSlice<i8>,
7268 ad: &CudaSlice<f32>,
7269 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7270 use crate::moe_cache::{BlockId, DispatchSlot, PROJ_GATE, PROJ_UP};
7271 let exps = match proj {
7272 PROJ_GATE => &m.gate_exps,
7273 PROJ_UP => &m.up_exps,
7274 _ => &m.down_exps,
7275 };
7276 let layout = exps.expert_layout(ex);
7277 let id = BlockId::new(il, proj, ex as u16);
7278 let source = exps.expert_source(ex);
7279 e.with_moe_cache(max_block, |c, eng| {
7280 let slot = c.dispatch_source(id, source, eng)?;
7281 let DispatchSlot::Resident(sl) = slot;
7282 let buf = c.slot(sl);
7283 eng.qmatvec_expert_q8(
7284 buf,
7285 0..layout.len,
7286 aq,
7287 ad,
7288 1,
7289 exps.in_f,
7290 exps.out_f,
7291 layout.qtype,
7292 layout.row_bytes,
7293 )
7294 })
7295 }
7296
7297 fn moe_cached_gemm(
7298 e: &Engine,
7299 il: u16,
7300 proj: u8,
7301 ex: usize,
7302 m: &MoeWeights,
7303 max_block: usize,
7304 x: &cudarc::driver::CudaView<f32>,
7305 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7306 use crate::moe_cache::{BlockId, DispatchSlot, PROJ_GATE, PROJ_UP};
7307 let exps = match proj {
7308 PROJ_GATE => &m.gate_exps,
7309 PROJ_UP => &m.up_exps,
7310 _ => &m.down_exps,
7311 };
7312 let layout = exps.expert_layout(ex);
7313 let id = BlockId::new(il, proj, ex as u16);
7314 let source = exps.expert_source(ex);
7315 e.with_moe_cache(max_block, |c, eng| {
7317 let slot = c.dispatch_source(id, source, eng)?;
7318 let DispatchSlot::Resident(sl) = slot;
7321 let buf = c.slot(sl);
7322 eng.qmatvec_view(
7323 buf,
7324 0..layout.len,
7325 x,
7326 1,
7327 exps.in_f,
7328 exps.out_f,
7329 layout.qtype,
7330 layout.row_bytes,
7331 )
7332 })
7333 }
7334
7335 fn moe_profile_admit_expert(
7339 e: &Engine,
7340 il: u16,
7341 ex: usize,
7342 m: &MoeWeights,
7343 max_block: usize,
7344 ) -> Result<(), Box<dyn std::error::Error>> {
7345 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
7346 e.with_moe_cache(max_block, |cache, eng| {
7347 for (proj, exps) in [
7348 (PROJ_GATE, &m.gate_exps),
7349 (PROJ_UP, &m.up_exps),
7350 (PROJ_DOWN, &m.down_exps),
7351 ] {
7352 let id = BlockId::new(il, proj, ex as u16);
7353 let _ = cache.dispatch_source(id, exps.expert_source(ex), eng)?;
7354 }
7355 Ok(())
7356 })
7357 }
7358
7359 #[allow(clippy::too_many_arguments)]
7362 fn moe_frozen_gemm(
7363 e: &Engine,
7364 il: u16,
7365 proj: u8,
7366 ex: usize,
7367 m: &MoeWeights,
7368 max_block: usize,
7369 x: &cudarc::driver::CudaView<f32>,
7370 scratch: &mut Option<CudaSlice<u8>>,
7371 scratch_len: usize,
7372 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7373 use crate::moe_cache::{BlockId, PROJ_GATE, PROJ_UP};
7374 let exps = match proj {
7375 PROJ_GATE => &m.gate_exps,
7376 PROJ_UP => &m.up_exps,
7377 _ => &m.down_exps,
7378 };
7379 let layout = exps.expert_layout(ex);
7380 let id = BlockId::new(il, proj, ex as u16);
7381 if let Some(output) = e.with_moe_cache(max_block, |cache, eng| {
7382 let Some(slot) = cache.resident(id) else {
7383 return Ok(None);
7384 };
7385 let buf = cache.slot(slot);
7386 Ok(Some(eng.qmatvec_view(
7387 buf,
7388 0..layout.len,
7389 x,
7390 1,
7391 exps.in_f,
7392 exps.out_f,
7393 layout.qtype,
7394 layout.row_bytes,
7395 )?))
7396 })? {
7397 return Ok(output);
7398 }
7399 if scratch.is_none() {
7400 *scratch = Some(e.alloc_u8_uninit(scratch_len)?);
7401 }
7402 let scratch = scratch.as_mut().unwrap();
7403 e.stage_expert(exps.expert_bytes(ex), scratch, 0)?;
7404 e.qmatvec_view(
7405 scratch,
7406 0..layout.len,
7407 x,
7408 1,
7409 exps.in_f,
7410 exps.out_f,
7411 layout.qtype,
7412 layout.row_bytes,
7413 )
7414 }
7415
7416 fn moe_prefetch_expert(
7417 e: &Engine,
7418 il: u16,
7419 ex: usize,
7420 m: &MoeWeights,
7421 max_block: usize,
7422 keep: &[crate::moe_cache::BlockId],
7423 ) -> Result<(), Box<dyn std::error::Error>> {
7424 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
7425 e.with_moe_cache(max_block, |c, eng| {
7426 for (proj, exps) in [
7427 (PROJ_GATE, &m.gate_exps),
7428 (PROJ_UP, &m.up_exps),
7429 (PROJ_DOWN, &m.down_exps),
7430 ] {
7431 let id = BlockId::new(il, proj, ex as u16);
7432 let _ = c.prefetch_source(id, exps.expert_source(ex), keep, eng)?;
7433 }
7434 Ok(())
7435 })
7436 }
7437
7438 fn moe_prefetch_disk_expert(
7441 e: &Engine,
7442 il: u16,
7443 ex: usize,
7444 m: &MoeWeights,
7445 max_block: usize,
7446 keep: &[crate::moe_cache::BlockId],
7447 ) -> Result<(), Box<dyn std::error::Error>> {
7448 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
7449 e.with_moe_cache(max_block, |c, eng| {
7450 for (proj, exps) in [
7451 (PROJ_GATE, &m.gate_exps),
7452 (PROJ_UP, &m.up_exps),
7453 (PROJ_DOWN, &m.down_exps),
7454 ] {
7455 let source = exps.expert_source(ex);
7456 if let crate::model::ExpertSource::Disk { .. } = &source {
7457 let id = BlockId::new(il, proj, ex as u16);
7458 let _ = c.prefetch_source(id, source, keep, eng)?;
7459 }
7460 }
7461 Ok(())
7462 })
7463 }
7464
7465 #[inline]
7466 fn moe_prefetch_host_expert(ex: usize, m: &MoeWeights) {
7467 let _ = m.gate_exps.prefetch_expert_pages(ex);
7468 let _ = m.up_exps.prefetch_expert_pages(ex);
7469 let _ = m.down_exps.prefetch_expert_pages(ex);
7470 }
7471}
7472
7473impl HybridModel {
7490 #[allow(clippy::too_many_arguments)]
7494 fn moe_ffn_grouped_resident_q8(
7495 e: &Engine,
7496 m: &MoeWeights,
7497 z: &CudaSlice<f32>,
7498 t: usize,
7499 cfg: &ModelConfig,
7500 il: u16,
7501 sel_all: &[u32],
7502 w_all: &[f32],
7503 table: &CudaSlice<u64>,
7504 gu_il: bool,
7505 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7506 let moe = cfg.moe.as_ref().unwrap();
7507 let n_embd = cfg.n_embd as usize;
7508 let n_expert = moe.expert_count as usize;
7509 let n_used = moe.expert_used_count as usize;
7510 let n_ff_exp = moe.expert_ff_length as usize;
7511 let n_pairs = t * n_used;
7512 debug_assert_eq!(sel_all.len(), n_pairs);
7513 debug_assert_eq!(w_all.len(), n_pairs);
7514 debug_assert!(
7515 m.gate_exps.macros.is_none()
7516 && m.up_exps.macros.is_none()
7517 && m.down_exps.macros.is_none(),
7518 "resident grouped q8 does not fold per-expert macro scales",
7519 );
7520
7521 if !cfg.swiglu_clamped_at(il as u32) {
7527 let sel: Vec<i32> = sel_all.iter().map(|&expert| expert as i32).collect();
7528 let sel_d = e.htod_i32(&sel)?;
7529 let w_d = e.htod(w_all)?;
7530 let (gate_row_bytes, up_row_bytes) = if gu_il {
7531 let combined = m.gate_exps.row_bytes + m.up_exps.row_bytes;
7532 (combined, combined)
7533 } else {
7534 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
7535 };
7536 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
7537 let act = e.moe_gate_up_silu8_dev_q8_rows(
7538 table,
7539 &sel_d,
7540 &zq,
7541 &zd,
7542 t,
7543 n_embd,
7544 n_ff_exp,
7545 n_used,
7546 n_expert,
7547 m.gate_exps.qtype,
7548 m.up_exps.qtype,
7549 gate_row_bytes,
7550 up_row_bytes,
7551 &m.dev_macros,
7552 )?;
7553 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
7554 let mut moe_out = e.uninit(t * n_embd)?;
7555 e.moe_down8_fma_dev_q8_rows_g(
7556 table,
7557 &sel_d,
7558 &w_d,
7559 &aq2,
7560 &ad2,
7561 &mut moe_out,
7562 t,
7563 n_ff_exp,
7564 n_embd,
7565 n_used,
7566 n_expert,
7567 m.down_exps.qtype,
7568 m.down_exps.row_bytes,
7569 )?;
7570
7571 if std::env::var("MEMRA_MOE_STATS").is_ok() {
7572 let mut counts = vec![0usize; n_expert];
7573 for &expert in sel_all {
7574 counts[expert as usize] += 1;
7575 }
7576 let mut sizes: Vec<usize> =
7577 counts.into_iter().filter(|&count| count != 0).collect();
7578 sizes.sort_unstable();
7579 let mean = sizes.iter().sum::<usize>() as f64 / sizes.len().max(1) as f64;
7580 println!(
7581 "moe-grouped il={il} t={t} dispatch=resident-q8-rows active={}/{} \
7582 m_e: min={} median={} mean={mean:.1} max={}",
7583 sizes.len(),
7584 n_expert,
7585 sizes.first().copied().unwrap_or(0),
7586 sizes.get(sizes.len() / 2).copied().unwrap_or(0),
7587 sizes.last().copied().unwrap_or(0),
7588 );
7589 }
7590 return Ok(moe_out);
7591 }
7592
7593 let pair_tok: Vec<i32> = (0..n_pairs).map(|pair| (pair / n_used) as i32).collect();
7597 let pair_ex: Vec<i32> = sel_all.iter().map(|&expert| expert as i32).collect();
7598 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
7599 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
7600
7601 let mut by_expert: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
7602 for (pair, &expert) in pair_ex.iter().enumerate() {
7603 by_expert[expert as usize].push(pair as i32);
7604 }
7605
7606 let pair_tok_d = e.htod_i32(&pair_tok)?;
7607 let pair_ex_d = e.htod_i32(&pair_ex)?;
7608 let pair_w_d = e.htod(w_all)?;
7609 let tok_off_d = e.htod_i32(&tok_off)?;
7610 let tok_ids_d = e.htod_i32(&tok_ids)?;
7611
7612 let matvec = |proj: i32,
7613 pair_rows: &CudaSlice<i32>,
7614 aq: &CudaSlice<i8>,
7615 ad: &CudaSlice<f32>,
7616 in_f: usize,
7617 out_f: usize,
7618 qtype: i32,
7619 row_bytes: usize|
7620 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7621 e.moe_pairs_matvec_q8(
7622 table, proj, pair_rows, &pair_ex_d, aq, ad, in_f, out_f, n_expert, n_pairs, qtype,
7623 row_bytes,
7624 )
7625 };
7626
7627 let (gate_row_bytes, up_row_bytes) = if gu_il {
7628 let combined = m.gate_exps.row_bytes + m.up_exps.row_bytes;
7629 (combined, combined)
7630 } else {
7631 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
7632 };
7633 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
7634 let gate = matvec(
7635 0,
7636 &pair_tok_d,
7637 &zq,
7638 &zd,
7639 n_embd,
7640 n_ff_exp,
7641 m.gate_exps.qtype,
7642 gate_row_bytes,
7643 )?;
7644 let up = matvec(
7645 1,
7646 &pair_tok_d,
7647 &zq,
7648 &zd,
7649 n_embd,
7650 n_ff_exp,
7651 m.up_exps.qtype,
7652 up_row_bytes,
7653 )?;
7654 let mut act = e.uninit(n_pairs * n_ff_exp)?;
7655 Self::ffn_act_lim(
7656 e,
7657 cfg,
7658 &gate,
7659 &up,
7660 1.0,
7661 1.0,
7662 cfg.clamp_exp_at(il as u32),
7663 &mut act,
7664 n_pairs * n_ff_exp,
7665 )?;
7666 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
7667 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
7668 let pair_self_d = e.htod_i32(&pair_self)?;
7669 let down = matvec(
7670 2,
7671 &pair_self_d,
7672 &aq2,
7673 &ad2,
7674 n_ff_exp,
7675 n_embd,
7676 m.down_exps.qtype,
7677 m.down_exps.row_bytes,
7678 )?;
7679 let mut moe_out = e.uninit(t * n_embd)?;
7680 e.moe_pairs_scatter(
7681 &down,
7682 &pair_w_d,
7683 &tok_off_d,
7684 &tok_ids_d,
7685 &mut moe_out,
7686 t,
7687 n_embd,
7688 )?;
7689
7690 if std::env::var("MEMRA_MOE_STATS").is_ok() {
7691 let mut sizes: Vec<usize> = by_expert
7692 .iter()
7693 .filter_map(|pairs| (!pairs.is_empty()).then_some(pairs.len()))
7694 .collect();
7695 sizes.sort_unstable();
7696 let mean = sizes.iter().sum::<usize>() as f64 / sizes.len().max(1) as f64;
7697 println!(
7698 "moe-grouped il={il} t={t} dispatch=resident-q8-clamped-pairs active={}/{} \
7699 m_e: min={} median={} mean={mean:.1} max={}",
7700 sizes.len(),
7701 n_expert,
7702 sizes.first().copied().unwrap_or(0),
7703 sizes.get(sizes.len() / 2).copied().unwrap_or(0),
7704 sizes.last().copied().unwrap_or(0),
7705 );
7706 }
7707 Ok(moe_out)
7708 }
7709
7710 fn moe_ffn_grouped_add_shared(
7711 e: &Engine,
7712 m: &MoeWeights,
7713 z: &CudaSlice<f32>,
7714 t: usize,
7715 cfg: &ModelConfig,
7716 il: u16,
7717 moe_out: &mut CudaSlice<f32>,
7718 ) -> Result<(), Box<dyn std::error::Error>> {
7719 let n_embd = cfg.n_embd as usize;
7720 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
7721 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
7722 {
7723 let n_ff_sh = gate_shexp.out_features();
7724 let sg_gate = e.matmul(gate_shexp, z, t)?;
7725 let sg_up = e.matmul(up_shexp, z, t)?;
7726 let mut sa = e.uninit(t * n_ff_sh)?;
7727 Self::ffn_act_lim(
7728 e,
7729 cfg,
7730 &sg_gate,
7731 &sg_up,
7732 1.0,
7733 1.0,
7734 cfg.clamp_shexp_at(il as u32),
7735 &mut sa,
7736 t * n_ff_sh,
7737 )?;
7738 let sh = e.matmul(down_shexp, &sa, t)?;
7739 let gate = match &m.gate_inp_shexp {
7740 Some(gate_inp_shexp) => {
7741 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
7742 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
7743 } else {
7744 let raw = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
7745 let mut gate = e.uninit(t)?;
7746 e.sigmoid(&raw, &mut gate, t)?;
7747 gate
7748 }
7749 }
7750 None => e.htod(&vec![1.0f32; t])?,
7751 };
7752 e.add_scaled_rows(&sh, &gate, moe_out, n_embd, t)?;
7753 }
7754 Ok(())
7755 }
7756
7757 pub(crate) fn moe_ffn_grouped(
7760 e: &Engine,
7761 m: &MoeWeights,
7762 z: &CudaSlice<f32>,
7763 t: usize,
7764 cfg: &ModelConfig,
7765 il: u16,
7766 max_block: usize,
7767 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7768 let moe = cfg.moe.as_ref().unwrap();
7769 let n_embd = cfg.n_embd as usize;
7770 let n_expert = moe.expert_count as usize;
7771 let n_used = moe.expert_used_count as usize;
7772 let n_ff_exp = moe.expert_ff_length as usize;
7773 let lim_exp = cfg.clamp_exp_at(il as u32);
7775
7776 let logits = Self::moe_router_logits(e, m, z, t, cfg)?;
7780 if let Some(sig) = cfg.sigmoid_router() {
7781 Self::trace_sigmoid_router_logits(e, il, t, n_expert, n_used, &logits, m, sig)?;
7782 }
7783 let (sel_all, w_all) = if let Some(sig) = cfg.sigmoid_router() {
7784 Self::moe_route_sigmoid_cfg(e, &logits, t, n_expert, n_used, m, sig)?
7785 } else {
7786 Self::moe_route_cfg(e, &logits, t, n_expert, n_used, m.active_experts.as_deref())?
7787 };
7788 crate::moesd::record_host_routes(il, n_expert, n_used, &sel_all)?;
7789 Self::trace_moe_routes(il, t, &sel_all, &w_all)?;
7790 Self::trace_moe_input(e, il, t, n_embd, z)?;
7791
7792 let no_exp_macros = m.gate_exps.macros.is_none()
7797 && m.up_exps.macros.is_none()
7798 && m.down_exps.macros.is_none();
7799 let resident_q8 = m.dev_exps.as_ref().filter(|dev| {
7800 m.has_uniform_expert_layout()
7801 && no_exp_macros
7802 && moe_q8_enabled()
7803 && q8_expert_supported(m.gate_exps.qtype)
7804 && q8_expert_supported(m.up_exps.qtype)
7805 && q8_expert_supported(m.down_exps.qtype)
7806 && moe_slab_enabled()
7807 && dev.dev == e.ctx().ordinal()
7808 });
7809 if let Some(dev) = resident_q8 {
7810 let mut moe_out = Self::moe_ffn_grouped_resident_q8(
7811 e,
7812 m,
7813 z,
7814 t,
7815 cfg,
7816 il,
7817 &sel_all,
7818 &w_all,
7819 &dev.ptr_row,
7820 dev.gu_il,
7821 )?;
7822 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut moe_out)?;
7823 return Ok(moe_out);
7824 }
7825
7826 struct ExpertGroup {
7830 tok_indices: Vec<i32>, slot_indices: Vec<i32>, weights: Vec<f32>, }
7834 let mut groups: Vec<ExpertGroup> = (0..n_expert)
7835 .map(|_| ExpertGroup {
7836 tok_indices: Vec::new(),
7837 slot_indices: Vec::new(),
7838 weights: Vec::new(),
7839 })
7840 .collect();
7841
7842 for tok in 0..t {
7843 for j in 0..n_used {
7844 let ex = sel_all[tok * n_used + j] as usize;
7845 let w = w_all[tok * n_used + j];
7846 groups[ex].tok_indices.push(tok as i32);
7847 groups[ex].slot_indices.push(j as i32);
7848 groups[ex].weights.push(w);
7849 }
7850 }
7851
7852 let mut slot_buf = e.zeros(t * n_used * n_embd)?;
7855 let mut wbuf = e.zeros(t * n_used)?; let g_len = m.gate_exps.max_expert_bytes();
7859 let u_len = m.up_exps.max_expert_bytes();
7860 let d_len = m.down_exps.max_expert_bytes();
7861 let moe_q8 = m.has_uniform_expert_layout()
7862 && moe_q8_enabled()
7863 && q8_expert_supported(m.gate_exps.qtype)
7864 && q8_expert_supported(m.up_exps.qtype)
7865 && q8_expert_supported(m.down_exps.qtype);
7866 let slab_local = m
7869 .dev_exps
7870 .as_ref()
7871 .filter(|dev| !dev.gu_il && moe_slab_enabled() && dev.dev == e.ctx().ordinal());
7872 let use_cache =
7873 slab_local.is_none() && Engine::moe_cache_enabled() && !e.moe_cache_frozen();
7874 let grouped_q8 = moe_q8 && (slab_local.is_some() || use_cache);
7877
7878 let (mut scratch_g, mut scratch_u, mut scratch_d) = if slab_local.is_none() && !use_cache {
7880 (
7881 Some(e.alloc_u8(g_len)?),
7882 Some(e.alloc_u8(u_len)?),
7883 Some(e.alloc_u8(d_len)?),
7884 )
7885 } else {
7886 (None, None, None)
7887 };
7888
7889 let mut order: Vec<usize> = (0..n_expert)
7900 .filter(|&ex| !groups[ex].tok_indices.is_empty())
7901 .collect();
7902 order.sort_by(|&a, &b| {
7903 groups[b]
7904 .tok_indices
7905 .len()
7906 .cmp(&groups[a].tok_indices.len())
7907 .then(a.cmp(&b))
7908 });
7909 let mut m_dist: Vec<usize> = Vec::new(); let page_window = moe_page_prefetch_window();
7911 let worker_disk_prefetch = use_cache && crate::spill_pread::worker_enabled();
7912 if worker_disk_prefetch {
7913 if let Some(first) = grouped_worker_prefetch_position(order.len(), None) {
7914 Self::moe_prefetch_disk_expert(e, il, order[first], m, max_block, &[])?;
7915 }
7916 }
7917 for (order_pos, &ex) in order.iter().enumerate() {
7918 for next in page_prefetch_positions(order_pos, order.len(), page_window) {
7919 Self::moe_prefetch_host_expert(order[next], m);
7920 }
7921 if worker_disk_prefetch {
7922 if let Some(next) = grouped_worker_prefetch_position(order.len(), Some(order_pos)) {
7923 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
7924 let keep = [
7925 BlockId::new(il, PROJ_GATE, ex as u16),
7926 BlockId::new(il, PROJ_UP, ex as u16),
7927 BlockId::new(il, PROJ_DOWN, ex as u16),
7928 ];
7929 Self::moe_prefetch_disk_expert(e, il, order[next], m, max_block, &keep)?;
7930 }
7931 }
7932 let grp = &groups[ex];
7933 let m_e = grp.tok_indices.len();
7934 m_dist.push(m_e);
7935 let gl = m.gate_exps.expert_layout(ex);
7936 let ul = m.up_exps.expert_layout(ex);
7937 let dl = m.down_exps.expert_layout(ex);
7938
7939 let tok_idx_d = e.htod_i32(&grp.tok_indices)?;
7943 let slot_idx_d = e.htod_i32(&grp.slot_indices)?;
7944 let dmac = m.down_exps.macro_scale(ex);
7945 let weight_d = if dmac == 1.0 {
7946 e.htod(&grp.weights)?
7947 } else {
7948 let scaled: Vec<f32> = grp.weights.iter().map(|&w| w * dmac).collect();
7949 e.htod(&scaled)?
7950 };
7951
7952 let mut gathered = e.zeros(m_e * n_embd)?;
7954 e.gather_rows(z, &tok_idx_d, &mut gathered, n_embd, m_e)?;
7955 let gv = gathered.slice(0..m_e * n_embd);
7956
7957 let y = if let Some(dev) = slab_local {
7960 let gate_start = ex * m.gate_exps.expert_stride;
7961 let up_start = ex * m.up_exps.expert_stride;
7962 let down_start = ex * m.down_exps.expert_stride;
7963 if grouped_q8 {
7964 let (zq, zd) = e.quantize_q8_1(&gathered, m_e, n_embd)?;
7965 let gate = e.qmatvec_expert_q8(
7966 &dev.gate,
7967 gate_start..gate_start + gl.len,
7968 &zq,
7969 &zd,
7970 m_e,
7971 m.gate_exps.in_f,
7972 m.gate_exps.out_f,
7973 gl.qtype,
7974 gl.row_bytes,
7975 )?;
7976 let up = e.qmatvec_expert_q8(
7977 &dev.up,
7978 up_start..up_start + ul.len,
7979 &zq,
7980 &zd,
7981 m_e,
7982 m.up_exps.in_f,
7983 m.up_exps.out_f,
7984 ul.qtype,
7985 ul.row_bytes,
7986 )?;
7987 let mut act = e.uninit(m_e * n_ff_exp)?;
7988 Self::ffn_act_lim(
7989 e,
7990 cfg,
7991 &gate,
7992 &up,
7993 m.gate_exps.macro_scale(ex),
7994 m.up_exps.macro_scale(ex),
7995 lim_exp,
7996 &mut act,
7997 m_e * n_ff_exp,
7998 )?;
7999 let (aq2, ad2) = e.quantize_q8_1(&act, m_e, n_ff_exp)?;
8000 e.qmatvec_expert_q8(
8001 &dev.down,
8002 down_start..down_start + dl.len,
8003 &aq2,
8004 &ad2,
8005 m_e,
8006 m.down_exps.in_f,
8007 m.down_exps.out_f,
8008 dl.qtype,
8009 dl.row_bytes,
8010 )?
8011 } else {
8012 let gate = e.qmatvec_view(
8013 &dev.gate,
8014 gate_start..gate_start + gl.len,
8015 &gv,
8016 m_e,
8017 m.gate_exps.in_f,
8018 m.gate_exps.out_f,
8019 gl.qtype,
8020 gl.row_bytes,
8021 )?;
8022 let up = e.qmatvec_view(
8023 &dev.up,
8024 up_start..up_start + ul.len,
8025 &gv,
8026 m_e,
8027 m.up_exps.in_f,
8028 m.up_exps.out_f,
8029 ul.qtype,
8030 ul.row_bytes,
8031 )?;
8032 let mut act = e.uninit(m_e * n_ff_exp)?;
8033 Self::ffn_act_lim(
8034 e,
8035 cfg,
8036 &gate,
8037 &up,
8038 m.gate_exps.macro_scale(ex),
8039 m.up_exps.macro_scale(ex),
8040 lim_exp,
8041 &mut act,
8042 m_e * n_ff_exp,
8043 )?;
8044 let actv = act.slice(0..m_e * n_ff_exp);
8045 e.qmatvec_view(
8046 &dev.down,
8047 down_start..down_start + dl.len,
8048 &actv,
8049 m_e,
8050 m.down_exps.in_f,
8051 m.down_exps.out_f,
8052 dl.qtype,
8053 dl.row_bytes,
8054 )?
8055 }
8056 } else if use_cache {
8057 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
8058 if grouped_q8 {
8059 let (zq, zd) = e.quantize_q8_1(&gathered, m_e, n_embd)?;
8060 let gate = e.with_moe_cache(max_block, |cache, eng| {
8061 let id = BlockId::new(il, PROJ_GATE, ex as u16);
8062 let slot = cache.dispatch_source(id, m.gate_exps.expert_source(ex), eng)?;
8063 eng.qmatvec_expert_q8(
8064 cache.buf(slot),
8065 0..gl.len,
8066 &zq,
8067 &zd,
8068 m_e,
8069 m.gate_exps.in_f,
8070 m.gate_exps.out_f,
8071 gl.qtype,
8072 gl.row_bytes,
8073 )
8074 })?;
8075 let up = e.with_moe_cache(max_block, |cache, eng| {
8076 let id = BlockId::new(il, PROJ_UP, ex as u16);
8077 let slot = cache.dispatch_source(id, m.up_exps.expert_source(ex), eng)?;
8078 eng.qmatvec_expert_q8(
8079 cache.buf(slot),
8080 0..ul.len,
8081 &zq,
8082 &zd,
8083 m_e,
8084 m.up_exps.in_f,
8085 m.up_exps.out_f,
8086 ul.qtype,
8087 ul.row_bytes,
8088 )
8089 })?;
8090 let mut act = e.uninit(m_e * n_ff_exp)?;
8091 Self::ffn_act_lim(
8092 e,
8093 cfg,
8094 &gate,
8095 &up,
8096 m.gate_exps.macro_scale(ex),
8097 m.up_exps.macro_scale(ex),
8098 lim_exp,
8099 &mut act,
8100 m_e * n_ff_exp,
8101 )?;
8102 let (aq2, ad2) = e.quantize_q8_1(&act, m_e, n_ff_exp)?;
8103 e.with_moe_cache(max_block, |cache, eng| {
8104 let id = BlockId::new(il, PROJ_DOWN, ex as u16);
8105 let slot = cache.dispatch_source(id, m.down_exps.expert_source(ex), eng)?;
8106 eng.qmatvec_expert_q8(
8107 cache.buf(slot),
8108 0..dl.len,
8109 &aq2,
8110 &ad2,
8111 m_e,
8112 m.down_exps.in_f,
8113 m.down_exps.out_f,
8114 dl.qtype,
8115 dl.row_bytes,
8116 )
8117 })?
8118 } else {
8119 let gate = e.with_moe_cache(max_block, |cache, eng| {
8120 let id = BlockId::new(il, PROJ_GATE, ex as u16);
8121 let slot = cache.dispatch_source(id, m.gate_exps.expert_source(ex), eng)?;
8122 eng.qmatvec_view(
8123 cache.buf(slot),
8124 0..gl.len,
8125 &gv,
8126 m_e,
8127 m.gate_exps.in_f,
8128 m.gate_exps.out_f,
8129 gl.qtype,
8130 gl.row_bytes,
8131 )
8132 })?;
8133 let up = e.with_moe_cache(max_block, |cache, eng| {
8134 let id = BlockId::new(il, PROJ_UP, ex as u16);
8135 let slot = cache.dispatch_source(id, m.up_exps.expert_source(ex), eng)?;
8136 eng.qmatvec_view(
8137 cache.buf(slot),
8138 0..ul.len,
8139 &gv,
8140 m_e,
8141 m.up_exps.in_f,
8142 m.up_exps.out_f,
8143 ul.qtype,
8144 ul.row_bytes,
8145 )
8146 })?;
8147 let mut act = e.uninit(m_e * n_ff_exp)?;
8148 Self::ffn_act_lim(
8149 e,
8150 cfg,
8151 &gate,
8152 &up,
8153 m.gate_exps.macro_scale(ex),
8154 m.up_exps.macro_scale(ex),
8155 lim_exp,
8156 &mut act,
8157 m_e * n_ff_exp,
8158 )?;
8159 let actv = act.slice(0..m_e * n_ff_exp);
8160 e.with_moe_cache(max_block, |cache, eng| {
8161 let id = BlockId::new(il, PROJ_DOWN, ex as u16);
8162 let slot = cache.dispatch_source(id, m.down_exps.expert_source(ex), eng)?;
8163 eng.qmatvec_view(
8164 cache.buf(slot),
8165 0..dl.len,
8166 &actv,
8167 m_e,
8168 m.down_exps.in_f,
8169 m.down_exps.out_f,
8170 dl.qtype,
8171 dl.row_bytes,
8172 )
8173 })?
8174 }
8175 } else {
8176 let sg = scratch_g.as_mut().unwrap();
8177 let su = scratch_u.as_mut().unwrap();
8178 let sd = scratch_d.as_mut().unwrap();
8179 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
8180 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
8181 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
8182 if grouped_q8 {
8183 let (zq, zd) = e.quantize_q8_1(&gathered, m_e, n_embd)?;
8184 let gate = e.qmatvec_expert_q8(
8185 sg,
8186 0..gl.len,
8187 &zq,
8188 &zd,
8189 m_e,
8190 m.gate_exps.in_f,
8191 m.gate_exps.out_f,
8192 gl.qtype,
8193 gl.row_bytes,
8194 )?;
8195 let up = e.qmatvec_expert_q8(
8196 su,
8197 0..ul.len,
8198 &zq,
8199 &zd,
8200 m_e,
8201 m.up_exps.in_f,
8202 m.up_exps.out_f,
8203 ul.qtype,
8204 ul.row_bytes,
8205 )?;
8206 let mut act = e.uninit(m_e * n_ff_exp)?;
8207 Self::ffn_act_lim(
8208 e,
8209 cfg,
8210 &gate,
8211 &up,
8212 m.gate_exps.macro_scale(ex),
8213 m.up_exps.macro_scale(ex),
8214 lim_exp,
8215 &mut act,
8216 m_e * n_ff_exp,
8217 )?;
8218 let (aq2, ad2) = e.quantize_q8_1(&act, m_e, n_ff_exp)?;
8219 e.qmatvec_expert_q8(
8220 sd,
8221 0..dl.len,
8222 &aq2,
8223 &ad2,
8224 m_e,
8225 m.down_exps.in_f,
8226 m.down_exps.out_f,
8227 dl.qtype,
8228 dl.row_bytes,
8229 )?
8230 } else {
8231 let gate = e.qmatvec_view(
8232 sg,
8233 0..gl.len,
8234 &gv,
8235 m_e,
8236 m.gate_exps.in_f,
8237 m.gate_exps.out_f,
8238 gl.qtype,
8239 gl.row_bytes,
8240 )?;
8241 let up = e.qmatvec_view(
8242 su,
8243 0..ul.len,
8244 &gv,
8245 m_e,
8246 m.up_exps.in_f,
8247 m.up_exps.out_f,
8248 ul.qtype,
8249 ul.row_bytes,
8250 )?;
8251 let mut act = e.uninit(m_e * n_ff_exp)?;
8252 Self::ffn_act_lim(
8253 e,
8254 cfg,
8255 &gate,
8256 &up,
8257 m.gate_exps.macro_scale(ex),
8258 m.up_exps.macro_scale(ex),
8259 lim_exp,
8260 &mut act,
8261 m_e * n_ff_exp,
8262 )?;
8263 let actv = act.slice(0..m_e * n_ff_exp);
8264 e.qmatvec_view(
8265 sd,
8266 0..dl.len,
8267 &actv,
8268 m_e,
8269 m.down_exps.in_f,
8270 m.down_exps.out_f,
8271 dl.qtype,
8272 dl.row_bytes,
8273 )?
8274 }
8275 };
8276
8277 e.scatter_slot(
8279 &y,
8280 &tok_idx_d,
8281 &slot_idx_d,
8282 &weight_d,
8283 &mut slot_buf,
8284 &mut wbuf,
8285 n_embd,
8286 n_used,
8287 m_e,
8288 )?;
8289 }
8290
8291 let mut moe_out = e.zeros(t * n_embd)?;
8293 e.reduce_slots(&slot_buf, &wbuf, &mut moe_out, n_embd, n_used, t)?;
8294
8295 if std::env::var("MEMRA_MOE_STATS").is_ok() && !m_dist.is_empty() {
8297 m_dist.sort_unstable();
8298 let active = m_dist.len();
8299 let mean = m_dist.iter().sum::<usize>() as f64 / active as f64;
8300 let median = m_dist[active / 2];
8301 let max_m = *m_dist.last().unwrap();
8302 let min_m = m_dist[0];
8303 let above16 = m_dist.iter().filter(|&&x| x >= 16).count();
8304 println!(
8305 "moe-grouped il={il} t={t} active={active}/{n_expert} \
8306 m_e: min={min_m} median={median} mean={mean:.1} max={max_m} \
8307 above_gemm_threshold(>=16)={above16}/{active}"
8308 );
8309 }
8310
8311 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut moe_out)?;
8312 Ok(moe_out)
8313 }
8314
8315 pub(crate) fn moe_ffn_lockstep(
8322 &self,
8323 e: &Engine,
8324 m: &MoeWeights,
8325 zbatch: &CudaSlice<f32>,
8326 mrows: usize,
8327 il: u16,
8328 max_block: usize,
8329 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8330 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
8331 let cfg = &self.cfg;
8332 let moe = cfg.moe.as_ref().unwrap();
8333 let n_embd = cfg.n_embd as usize;
8334 let n_expert = moe.expert_count as usize;
8335 let n_used = moe.expert_used_count as usize;
8336 let n_ff_exp = moe.expert_ff_length as usize;
8337 let lim_exp = cfg.clamp_exp_at(il as u32);
8339 let lim_shexp = cfg.clamp_shexp_at(il as u32);
8340
8341 let logits = e.matmul(&m.gate_inp, zbatch, mrows)?;
8342 if let Some(sig) = cfg.sigmoid_router() {
8343 Self::trace_sigmoid_router_logits(e, il, mrows, n_expert, n_used, &logits, m, sig)?;
8344 }
8345 let (sel_all, w_all) = if let Some(sig) = cfg.sigmoid_router() {
8346 Self::moe_route_sigmoid_cfg(e, &logits, mrows, n_expert, n_used, m, sig)?
8347 } else {
8348 Self::moe_route_cfg(
8349 e,
8350 &logits,
8351 mrows,
8352 n_expert,
8353 n_used,
8354 m.active_experts.as_deref(),
8355 )?
8356 };
8357 Self::trace_moe_routes(il, mrows, &sel_all, &w_all)?;
8358
8359 let resident_expert: Vec<bool> = e.with_moe_cache(max_block, |c, _| {
8361 Ok((0..n_expert)
8362 .map(|ex| {
8363 [PROJ_GATE, PROJ_UP, PROJ_DOWN]
8364 .into_iter()
8365 .all(|p| c.resident(BlockId::new(il, p, ex as u16)).is_some())
8366 })
8367 .collect())
8368 })?;
8369
8370 struct Group {
8371 rows: Vec<i32>,
8372 slots: Vec<i32>,
8373 weights: Vec<f32>,
8374 }
8375 let mut groups: std::collections::HashMap<usize, Group> = Default::default();
8376 let mut cpu_rows: Vec<Vec<(usize, f32)>> = vec![Vec::new(); mrows];
8377 let mut cpu_by_expert: std::collections::HashMap<usize, Vec<(usize, f32)>> =
8378 Default::default();
8379 for row in 0..mrows {
8380 for j in 0..n_used {
8381 let ex = sel_all[row * n_used + j] as usize;
8382 let w = w_all[row * n_used + j];
8383 if resident_expert[ex] {
8384 let group = groups.entry(ex).or_insert_with(|| Group {
8385 rows: Vec::new(),
8386 slots: Vec::new(),
8387 weights: Vec::new(),
8388 });
8389 group.rows.push(row as i32);
8390 group.slots.push(j as i32);
8391 group.weights.push(w);
8392 } else {
8393 crate::cpu_experts::record_incomplete_gpu_residency(0);
8394 cpu_rows[row].push((ex, w));
8395 cpu_by_expert.entry(ex).or_default().push((row, w));
8396 }
8397 }
8398 }
8399
8400 let host_rows = e.dtoh(zbatch)?;
8406 let rows_ok = crate::cpu_experts::rows_supported();
8407 enum CpuPart {
8408 Single { row: usize },
8409 Rows { rows: Vec<usize> },
8410 }
8411 let mut tickets: Vec<(CpuPart, crate::cpu_experts::CpuExpertTicket)> = Vec::new();
8412 let mut rows_served: std::collections::HashSet<(usize, usize)> = Default::default();
8413 if rows_ok {
8414 let mut shared: Vec<(usize, Vec<(usize, f32)>)> = cpu_by_expert
8415 .into_iter()
8416 .filter(|(_, rows)| rows.len() >= 2)
8417 .collect();
8418 shared.sort_by_key(|(ex, _)| *ex);
8419 for (ex, mut row_weights) in shared {
8420 row_weights.sort_by_key(|(row, _)| *row);
8421 let inputs: Vec<(&[f32], f32)> = row_weights
8422 .iter()
8423 .map(|&(row, w)| (&host_rows[row * n_embd..(row + 1) * n_embd], w))
8424 .collect();
8425 let job = crate::cpu_experts::prepare_rows_job(m, ex, &inputs)
8426 .map_err(std::io::Error::other)?;
8427 for &(row, _) in &row_weights {
8428 rows_served.insert((row, ex));
8429 }
8430 tickets.push((
8431 CpuPart::Rows {
8432 rows: row_weights.iter().map(|&(row, _)| row).collect(),
8433 },
8434 crate::cpu_experts::submit_rows(job).map_err(std::io::Error::other)?,
8435 ));
8436 }
8437 }
8438 for (row, selected) in cpu_rows.iter().enumerate() {
8439 let leftover: Vec<(usize, f32)> = selected
8440 .iter()
8441 .copied()
8442 .filter(|&(ex, _)| !rows_served.contains(&(row, ex)))
8443 .collect();
8444 if leftover.is_empty() {
8445 continue;
8446 }
8447 let host_row = &host_rows[row * n_embd..(row + 1) * n_embd];
8448 let job = crate::cpu_experts::prepare_job(m, il, &leftover, host_row)
8449 .map_err(std::io::Error::other)?;
8450 tickets.push((
8451 CpuPart::Single { row },
8452 crate::cpu_experts::submit(job).map_err(std::io::Error::other)?,
8453 ));
8454 }
8455
8456 let mut slot_buf = e.zeros(mrows * n_used * n_embd)?;
8457 let mut wbuf = e.zeros(mrows * n_used)?;
8458 let mut order: Vec<usize> = groups.keys().copied().collect();
8459 order.sort_by(|&a, &b| {
8460 groups[&b]
8461 .rows
8462 .len()
8463 .cmp(&groups[&a].rows.len())
8464 .then(a.cmp(&b))
8465 });
8466 for &ex in &order {
8467 let group = &groups[&ex];
8468 let m_e = group.rows.len();
8469 let gl = m.gate_exps.expert_layout(ex);
8470 let ul = m.up_exps.expert_layout(ex);
8471 let dl = m.down_exps.expert_layout(ex);
8472 let row_idx_d = e.htod_i32(&group.rows)?;
8473 let slot_idx_d = e.htod_i32(&group.slots)?;
8474 let dmac = m.down_exps.macro_scale(ex);
8475 let weight_d = if dmac == 1.0 {
8476 e.htod(&group.weights)?
8477 } else {
8478 let scaled: Vec<f32> = group.weights.iter().map(|&w| w * dmac).collect();
8479 e.htod(&scaled)?
8480 };
8481 let mut gathered = e.zeros(m_e * n_embd)?;
8482 e.gather_rows(zbatch, &row_idx_d, &mut gathered, n_embd, m_e)?;
8483 let gv = gathered.slice(0..m_e * n_embd);
8484 let gate = e.with_moe_cache(max_block, |c, eng| {
8485 let slot = c
8486 .resident(BlockId::new(il, PROJ_GATE, ex as u16))
8487 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
8488 eng.qmatvec_view(
8489 c.buf(crate::moe_cache::DispatchSlot::Resident(slot)),
8490 0..gl.len,
8491 &gv,
8492 m_e,
8493 m.gate_exps.in_f,
8494 m.gate_exps.out_f,
8495 gl.qtype,
8496 gl.row_bytes,
8497 )
8498 })?;
8499 let up = e.with_moe_cache(max_block, |c, eng| {
8500 let slot = c
8501 .resident(BlockId::new(il, PROJ_UP, ex as u16))
8502 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
8503 eng.qmatvec_view(
8504 c.buf(crate::moe_cache::DispatchSlot::Resident(slot)),
8505 0..ul.len,
8506 &gv,
8507 m_e,
8508 m.up_exps.in_f,
8509 m.up_exps.out_f,
8510 ul.qtype,
8511 ul.row_bytes,
8512 )
8513 })?;
8514 let mut act = e.zeros(m_e * n_ff_exp)?;
8515 Self::ffn_act_lim(
8516 e,
8517 cfg,
8518 &gate,
8519 &up,
8520 m.gate_exps.macro_scale(ex),
8521 m.up_exps.macro_scale(ex),
8522 lim_exp,
8523 &mut act,
8524 m_e * n_ff_exp,
8525 )?;
8526 let actv = act.slice(0..m_e * n_ff_exp);
8527 let y = e.with_moe_cache(max_block, |c, eng| {
8528 let slot = c
8529 .resident(BlockId::new(il, PROJ_DOWN, ex as u16))
8530 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
8531 eng.qmatvec_view(
8532 c.buf(crate::moe_cache::DispatchSlot::Resident(slot)),
8533 0..dl.len,
8534 &actv,
8535 m_e,
8536 m.down_exps.in_f,
8537 m.down_exps.out_f,
8538 dl.qtype,
8539 dl.row_bytes,
8540 )
8541 })?;
8542 e.scatter_slot(
8543 &y,
8544 &row_idx_d,
8545 &slot_idx_d,
8546 &weight_d,
8547 &mut slot_buf,
8548 &mut wbuf,
8549 n_embd,
8550 n_used,
8551 m_e,
8552 )?;
8553 }
8554 let mut moe_out = e.zeros(mrows * n_embd)?;
8555 e.reduce_slots(&slot_buf, &wbuf, &mut moe_out, n_embd, n_used, mrows)?;
8556
8557 let mut row_sums: Vec<Option<Vec<f32>>> = vec![None; mrows];
8559 for (part, ticket) in tickets {
8560 let cpu_output = ticket.wait().map_err(std::io::Error::other)?;
8561 let mut add_row = |row: usize, chunk: &[f32]| {
8562 let sum = row_sums[row].get_or_insert_with(|| vec![0.0f32; n_embd]);
8563 for (accumulator, value) in sum.iter_mut().zip(chunk) {
8564 *accumulator += value;
8565 }
8566 };
8567 match part {
8568 CpuPart::Single { row } => add_row(row, &cpu_output),
8569 CpuPart::Rows { rows } => {
8570 for (slot, row) in rows.into_iter().enumerate() {
8571 add_row(row, &cpu_output[slot * n_embd..(slot + 1) * n_embd]);
8572 }
8573 }
8574 }
8575 }
8576 for (row, sum) in row_sums.into_iter().enumerate() {
8577 let Some(sum) = sum else { continue };
8578 let cpu_output = e.htod(&sum)?;
8579 let mut dst = moe_out.slice_mut(row * n_embd..(row + 1) * n_embd);
8580 e.axpy_into(&cpu_output, 1.0, &mut dst, n_embd)?;
8581 }
8582
8583 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
8584 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
8585 {
8586 let n_ff_sh = gate_shexp.out_features();
8587 let sg_gate = e.matmul(gate_shexp, zbatch, mrows)?;
8588 let sg_up = e.matmul(up_shexp, zbatch, mrows)?;
8589 let mut sa = e.zeros(mrows * n_ff_sh)?;
8590 Self::ffn_act_lim(
8591 e,
8592 cfg,
8593 &sg_gate,
8594 &sg_up,
8595 1.0,
8596 1.0,
8597 lim_shexp,
8598 &mut sa,
8599 mrows * n_ff_sh,
8600 )?;
8601 let sh = e.matmul(down_shexp, &sa, mrows)?;
8602 let g = match &m.gate_inp_shexp {
8605 Some(gate_inp_shexp) => {
8606 e.sigmoid_dot_rows(zbatch, gate_inp_shexp.float_data(), n_embd, mrows)?
8607 }
8608 None => e.htod(&vec![1.0f32; mrows])?,
8609 };
8610 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, mrows)?;
8611 }
8612
8613 Ok(moe_out)
8614 }
8615}
8616
8617impl HybridModel {
8623 pub(crate) fn gemma4_geom(&self, il: usize) -> (usize, usize, usize, f32, f32, bool) {
8625 let g = self.cfg.gemma4.as_ref().unwrap();
8626 let swa = g.swa_pattern[il];
8627 let hd = if swa {
8628 g.key_length_swa
8629 } else {
8630 g.key_length_global
8631 } as usize;
8632 (
8636 hd,
8637 g.head_count_kv[il] as usize,
8638 self.cfg.n_head as usize,
8639 if swa {
8640 g.rope_base_swa
8641 } else {
8642 g.rope_base_global
8643 },
8644 1.0,
8645 swa,
8646 )
8647 }
8648
8649 pub(crate) fn gemma4_suppress(
8653 &self,
8654 e: &Engine,
8655 ld: &mut CudaSlice<f32>,
8656 t: usize,
8657 ) -> Result<(), Box<dyn std::error::Error>> {
8658 if let Some((ids, n)) = self.gemma4_aux.as_ref().and_then(|a| a.suppress_d.as_ref()) {
8659 #[cfg(debug_assertions)]
8664 crate::debug_assert_tensor_stream_device(
8665 ids,
8666 &e.stream(),
8667 "gemma4_suppress.suppress_d",
8668 );
8669 e.mask_ids_rows(ld, ids, *n, self.output.out_features(), t)?;
8670 }
8671 Ok(())
8672 }
8673
8674 #[allow(clippy::too_many_arguments)]
8679 fn gemma4_attn_prime(
8680 &self,
8681 e: &Engine,
8682 fa: &crate::hybrid::FullAttnLayer,
8683 il: usize,
8684 h: &CudaSlice<f32>,
8685 pos_d: &CudaSlice<i32>,
8686 t: usize,
8687 cache: Option<&mut Cache>,
8688 island: Option<&CudaSlice<i32>>,
8689 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8690 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
8691 let eps = self.cfg.rms_eps;
8692 let aux = self.gemma4_aux.as_ref().unwrap();
8693 let ones = aux.ones(e);
8694 #[cfg(debug_assertions)]
8695 crate::debug_assert_tensor_stream_device(ones, &e.stream(), "gemma4_attn_prime.ones");
8696
8697 e.mmq_act_begin();
8700 let q0 = e.matmul(&fa.wq, h, t)?; if il == 0 && std::env::var("MEMRA_G4_PRIME_TRACE").as_deref() == Ok("1") {
8702 let v = e.dtoh(&q0)?;
8703 let nan = v.iter().filter(|x| x.is_nan()).count();
8704 let amax = v.iter().fold(0f32, |a, x| a.max(x.abs()));
8705 eprintln!(
8706 "[g4-prime-trace] L0 q0: nan={nan}/{} amax={amax:.3}",
8707 v.len()
8708 );
8709 }
8710 let k0 = e.matmul(&fa.wk, h, t)?; let v0 = if swa {
8714 e.matmul(&fa.wv, h, t)?
8715 } else {
8716 e.clone_dtod(&k0)?
8717 };
8718 if il == 0 && std::env::var("MEMRA_G4_PRIME_TRACE").as_deref() == Ok("1") {
8719 for (tag, buf) in [("k0", &k0), ("v0", &v0)] {
8720 let v = e.dtoh(buf)?;
8721 let nan = v.iter().filter(|x| x.is_nan()).count();
8722 let amax = v.iter().fold(0f32, |a, x| a.max(x.abs()));
8723 eprintln!(
8724 "[g4-prime-trace] L0 {tag}: nan={nan}/{} amax={amax:.3}",
8725 v.len()
8726 );
8727 }
8728 }
8729
8730 let mut q = e.uninit(t * nh * hd)?;
8731 let mut k = e.uninit(t * nkv * hd)?;
8732 let mut v = e.uninit(t * nkv * hd)?;
8734 static EMIT: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8738 let emit = island.is_none()
8741 && t >= 16
8742 && crate::Engine::qkvnorm_w_on_prefill(nh * t + 2 * nkv * t, hd)
8743 && *EMIT.get_or_init(|| {
8744 std::env::var("MEMRA_FA_EMIT")
8745 .map(|s| s != "0")
8746 .unwrap_or(true)
8747 });
8748 let mut qb = e.alloc_uninit::<u8>(if emit { t * nh * hd * 2 } else { 1 })?;
8749 let mut kb = e.alloc_uninit::<u8>(if emit { t * nkv * hd * 2 } else { 1 })?;
8750 let mut vb = e.alloc_uninit::<u8>(if emit { t * nkv * hd * 2 } else { 1 })?;
8751 let v_f16 = emit
8754 && crate::fa_f16pv_on()
8755 && match hd {
8756 512 => true,
8757 256 => swa && crate::faw_hp_on() && nh % 2 == 0 && (nh / nkv) % 2 == 0,
8758 _ => false,
8759 };
8760 if emit {
8761 e.rms_norm_qkv_w4b(
8762 &q0,
8763 &k0,
8764 &v0,
8765 fa.q_norm.float_data(),
8766 fa.k_norm.float_data(),
8767 ones,
8768 &mut q,
8769 &mut k,
8770 &mut v,
8771 &mut vb,
8772 hd,
8773 nh * t,
8774 nkv * t,
8775 eps,
8776 v_f16,
8777 )?;
8778 } else {
8779 e.rms_norm_qkv(
8780 &q0,
8781 &k0,
8782 &v0,
8783 fa.q_norm.float_data(),
8784 fa.k_norm.float_data(),
8785 ones,
8786 &mut q,
8787 &mut k,
8788 &mut v,
8789 hd,
8790 nh * t,
8791 nkv * t,
8792 eps,
8793 )?;
8794 }
8795
8796 let ff = if swa {
8797 None
8798 } else {
8799 Some(
8800 aux.rope_freqs(e)
8801 .expect("gemma4 global rope needs rope_freqs.weight"),
8802 )
8803 };
8804 #[cfg(debug_assertions)]
8805 if let Some(ff) = ff {
8806 crate::debug_assert_tensor_stream_device(
8807 ff,
8808 &e.stream(),
8809 "gemma4_attn_prime.rope_freqs",
8810 );
8811 }
8812 if emit {
8813 e.rope_neox2_bf16e(
8814 &mut q, &mut k, &mut qb, &mut kb, pos_d, hd, hd, nh, nkv, t, base, 1.0, ff,
8815 )?;
8816 } else {
8817 e.rope_neox2(&mut q, &mut k, pos_d, hd, hd, nh, nkv, t, base, 1.0, ff)?;
8818 }
8819
8820 if let Some(cache) = cache {
8821 let kvl = cache.kv[il].as_mut().unwrap();
8822 assert_eq!(kvl.len, 0, "gemma4 prime is fresh-prompt only (v0)");
8823 e.append_kv_quantized_rows(
8824 &k,
8825 &v,
8826 &mut kvl.k,
8827 &mut kvl.v,
8828 kvl.len,
8829 t,
8830 kvl.kv_dim_k,
8831 kvl.kv_dim_v,
8832 kvl.k_tok_bytes,
8833 kvl.v_tok_bytes,
8834 (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on()),
8835 )?;
8836 kvl.len += t;
8837 }
8838 let mut attn = e.zeros(t * nh * hd)?;
8839 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
8843 if let Some(span) = island {
8844 let w = if swa && t > win { win } else { 0 };
8849 e.sdpa_naive_island(&q, &k, &v, &mut attn, span, hd, nh, nkv, t, t, scale, w)?;
8850 } else if swa && t > win {
8851 if hd == 256 && std::env::var("MEMRA_NOFA").is_err() {
8852 if emit {
8853 e.fa_prefill_w_pre(
8854 &qb, &kb, &vb, &mut attn, hd, nh, nkv, t, t, scale, true, win, v_f16,
8855 )?;
8856 } else {
8857 e.fa_prefill_w(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
8858 }
8859 } else {
8860 e.sdpa_naive_w(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
8861 }
8862 } else if hd == 256 && std::env::var("MEMRA_NOFA").is_err() {
8863 e.fa_prefill(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
8864 } else if hd == 512 && std::env::var("MEMRA_NOFA").is_err() {
8865 if emit {
8866 e.fa_prefill_hd512_pre(
8867 &qb, &kb, &vb, &mut attn, hd, nh, nkv, t, t, scale, true, v_f16,
8868 )?;
8869 } else {
8870 e.fa_prefill_hd512(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
8871 }
8872 } else {
8873 e.sdpa_naive(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
8874 }
8875 Ok(e.matmul(&fa.wo, &attn, t)?)
8876 }
8877
8878 fn gemma4_attn(
8880 &self,
8881 e: &Engine,
8882 fa: &crate::hybrid::FullAttnLayer,
8883 il: usize,
8884 h: &CudaSlice<f32>,
8885 pos_d: &CudaSlice<i32>,
8886 t: usize,
8887 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8888 self.gemma4_attn_prime(e, fa, il, h, pos_d, t, None, None)
8889 }
8890
8891 fn gemma4_moe_q8(
8896 &self,
8897 e: &Engine,
8898 m: &crate::hybrid::MoeWeights,
8899 bits: &crate::hybrid::Gemma4MoeBits,
8900 mq: &(CudaSlice<i8>, CudaSlice<f32>),
8901 router_in: &CudaSlice<f32>,
8902 t: usize,
8903 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8904 let cfg = &self.cfg;
8905 let moe = cfg.moe.as_ref().unwrap();
8906 let n_embd = cfg.n_embd as usize;
8907 let n_expert = moe.expert_count as usize;
8908 let n_used = moe.expert_used_count as usize;
8909 let n_ff_exp = moe.expert_ff_length as usize;
8910 let logits = if crate::router_kernel_on() {
8914 e.router_gemv(m.gate_inp.float_data(), router_in, n_embd, n_expert, t)?
8915 } else {
8916 e.matmul(&m.gate_inp, router_in, t)?
8917 };
8918 let dev = m.dev_exps.as_ref().unwrap();
8919 let (sel_d, w_d) =
8920 e.moe_router_topk_scaled(&logits, t, n_expert, n_used, &bits.per_expert_scale_d)?;
8921 let (zq, zd) = mq;
8922 if t == 1 {
8923 let selv = sel_d.slice(0..n_used);
8924 let wv = w_d.slice(0..n_used);
8925 let act = e.moe_gate_up_gelu8_dev_q8(
8926 &dev.ptr_row,
8927 &selv,
8928 zq,
8929 zd,
8930 n_embd,
8931 n_ff_exp,
8932 n_used,
8933 n_expert,
8934 m.gate_exps.qtype,
8935 m.up_exps.qtype,
8936 m.gate_exps.row_bytes,
8937 m.up_exps.row_bytes,
8938 )?;
8939 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
8940 let mut moe_out = e.uninit(n_embd)?;
8941 e.moe_down8_fma_dev_q8(
8942 &dev.ptr_row,
8943 &selv,
8944 &wv,
8945 &aq2,
8946 &ad2,
8947 &mut moe_out.slice_mut(0..n_embd),
8948 n_ff_exp,
8949 n_embd,
8950 n_used,
8951 n_expert,
8952 m.down_exps.qtype,
8953 m.down_exps.row_bytes,
8954 )?;
8955 return Ok(moe_out);
8956 }
8957 let csr = t <= 10 && std::env::var("MEMRA_GEMMA_CSR").as_deref() != Ok("0");
8958 let act = if csr {
8959 e.moe_gate_up_gelu8_dev_q8_csr(
8960 &dev.ptr_row,
8961 &sel_d,
8962 zq,
8963 zd,
8964 t * n_used,
8965 n_embd,
8966 n_ff_exp,
8967 n_used,
8968 n_expert,
8969 m.gate_exps.qtype,
8970 m.up_exps.qtype,
8971 m.gate_exps.row_bytes,
8972 m.up_exps.row_bytes,
8973 )?
8974 } else {
8975 e.moe_gate_up_gelu8_dev_q8_rows(
8976 &dev.ptr_row,
8977 &sel_d,
8978 zq,
8979 zd,
8980 t,
8981 n_embd,
8982 n_ff_exp,
8983 n_used,
8984 n_expert,
8985 m.gate_exps.qtype,
8986 m.up_exps.qtype,
8987 m.gate_exps.row_bytes,
8988 m.up_exps.row_bytes,
8989 )?
8990 };
8991 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
8992 let mut moe_out = e.uninit(t * n_embd)?;
8993 e.moe_down8_fma_dev_q8_rows_g(
8996 &dev.ptr_row,
8997 &sel_d,
8998 &w_d,
8999 &aq2,
9000 &ad2,
9001 &mut moe_out,
9002 t,
9003 n_ff_exp,
9004 n_embd,
9005 n_used,
9006 n_expert,
9007 m.down_exps.qtype,
9008 m.down_exps.row_bytes,
9009 )?;
9010 Ok(moe_out)
9011 }
9012
9013 fn gemma4_moe(
9017 &self,
9018 e: &Engine,
9019 m: &crate::hybrid::MoeWeights,
9020 bits: &crate::hybrid::Gemma4MoeBits,
9021 moe_in: &CudaSlice<f32>,
9022 router_in: &CudaSlice<f32>,
9023 t: usize,
9024 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9025 let cfg = &self.cfg;
9026 let moe = cfg.moe.as_ref().unwrap();
9027 let n_embd = cfg.n_embd as usize;
9028 let n_expert = moe.expert_count as usize;
9029 let n_used = moe.expert_used_count as usize;
9030 let n_ff_exp = moe.expert_ff_length as usize;
9031
9032 let logits = if t < PRIME_MIN_T && crate::router_kernel_on() {
9036 e.router_gemv(m.gate_inp.float_data(), router_in, n_embd, n_expert, t)?
9037 } else {
9038 e.matmul(&m.gate_inp, router_in, t)?
9039 };
9040
9041 if t < PRIME_MIN_T
9046 && m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
9047 && expert_dp4a_supported(m.gate_exps.qtype)
9048 && expert_dp4a_supported(m.up_exps.qtype)
9049 && expert_dp4a_supported(m.down_exps.qtype)
9050 && std::env::var("MEMRA_GEMMA_MOE_FAST").as_deref() != Ok("0")
9051 {
9052 let dev = m.dev_exps.as_ref().unwrap();
9053 let (sel_d, w_d) =
9054 e.moe_router_topk_scaled(&logits, t, n_expert, n_used, &bits.per_expert_scale_d)?;
9055 if t == 1 {
9056 let (zq, zd) = e.quantize_q8_1(moe_in, 1, n_embd)?;
9057 let selv = sel_d.slice(0..n_used);
9058 let wv = w_d.slice(0..n_used);
9059 let act = e.moe_gate_up_gelu8_dev_q8(
9060 &dev.ptr_row,
9061 &selv,
9062 &zq,
9063 &zd,
9064 n_embd,
9065 n_ff_exp,
9066 n_used,
9067 n_expert,
9068 m.gate_exps.qtype,
9069 m.up_exps.qtype,
9070 m.gate_exps.row_bytes,
9071 m.up_exps.row_bytes,
9072 )?;
9073 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
9074 let mut moe_out = e.uninit(n_embd)?;
9075 e.moe_down8_fma_dev_q8(
9076 &dev.ptr_row,
9077 &selv,
9078 &wv,
9079 &aq2,
9080 &ad2,
9081 &mut moe_out.slice_mut(0..n_embd),
9082 n_ff_exp,
9083 n_embd,
9084 n_used,
9085 n_expert,
9086 m.down_exps.qtype,
9087 m.down_exps.row_bytes,
9088 )?;
9089 return Ok(moe_out);
9090 }
9091 let (zq, zd) = e.quantize_q8_1(moe_in, t, n_embd)?;
9096 let csr = t <= 10 && std::env::var("MEMRA_GEMMA_CSR").as_deref() != Ok("0");
9097 let act = if csr {
9098 e.moe_gate_up_gelu8_dev_q8_csr(
9099 &dev.ptr_row,
9100 &sel_d,
9101 &zq,
9102 &zd,
9103 t * n_used,
9104 n_embd,
9105 n_ff_exp,
9106 n_used,
9107 n_expert,
9108 m.gate_exps.qtype,
9109 m.up_exps.qtype,
9110 m.gate_exps.row_bytes,
9111 m.up_exps.row_bytes,
9112 )?
9113 } else {
9114 e.moe_gate_up_gelu8_dev_q8_rows(
9115 &dev.ptr_row,
9116 &sel_d,
9117 &zq,
9118 &zd,
9119 t,
9120 n_embd,
9121 n_ff_exp,
9122 n_used,
9123 n_expert,
9124 m.gate_exps.qtype,
9125 m.up_exps.qtype,
9126 m.gate_exps.row_bytes,
9127 m.up_exps.row_bytes,
9128 )?
9129 };
9130 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
9131 let mut moe_out = e.uninit(t * n_embd)?;
9132 e.moe_down8_fma_dev_q8_rows_g(
9133 &dev.ptr_row,
9134 &sel_d,
9135 &w_d,
9136 &aq2,
9137 &ad2,
9138 &mut moe_out,
9139 t,
9140 n_ff_exp,
9141 n_embd,
9142 n_used,
9143 n_expert,
9144 m.down_exps.qtype,
9145 m.down_exps.row_bytes,
9146 )?;
9147 return Ok(moe_out);
9148 }
9149
9150 let (sel_all, mut w_all) = Self::moe_route(e, &logits, t, n_expert, n_used)?;
9151 for (i, &sx) in sel_all.iter().enumerate() {
9152 w_all[i] *= bits.per_expert_scale[sx as usize];
9153 }
9154
9155 if t >= PRIME_MIN_T
9159 && m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
9160 && expert_dp4a_supported(m.gate_exps.qtype)
9161 && expert_dp4a_supported(m.up_exps.qtype)
9162 && expert_dp4a_supported(m.down_exps.qtype)
9163 && std::env::var("MEMRA_GEMMA_MOE_PAIRS").as_deref() != Ok("0")
9164 {
9165 let dev = m.dev_exps.as_ref().unwrap();
9166 let n_pairs = t * n_used;
9167 let pair_ex: Vec<i32> = sel_all.iter().map(|&x| x as i32).collect();
9168 let pair_tok: Vec<i32> = (0..n_pairs).map(|p| (p / n_used) as i32).collect();
9169 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
9170 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
9171 let pt = e.htod_i32(&pair_tok)?;
9172 let pw = e.htod(&w_all)?;
9173 let toff = e.htod_i32(&tok_off)?;
9174 let tids = e.htod_i32(&tok_ids)?;
9175 let mut by_ex: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
9176 for p in 0..n_pairs {
9177 by_ex[pair_ex[p] as usize].push(p as i32);
9178 }
9179 let mut ex_ids: Vec<i32> = Vec::new();
9180 let mut ex_off: Vec<i32> = vec![0];
9181 let mut ex_pairs: Vec<i32> = Vec::with_capacity(n_pairs);
9182 for (ex, list) in by_ex.iter().enumerate() {
9183 if list.is_empty() {
9184 continue;
9185 }
9186 ex_ids.push(ex as i32);
9187 ex_pairs.extend_from_slice(list);
9188 ex_off.push(ex_pairs.len() as i32);
9189 }
9190 let n_active = ex_ids.len();
9191 let exi = e.htod_i32(&ex_ids)?;
9192 let exo = e.htod_i32(&ex_off)?;
9193 let exp_d = e.htod_i32(&ex_pairs)?;
9194 if crate::moe_f16g_gemma_on()
9202 && f16g_proj_ok(m.gate_exps.qtype, n_embd)
9203 && f16g_proj_ok(m.up_exps.qtype, n_embd)
9204 && f16g_proj_ok(m.down_exps.qtype, n_ff_exp)
9205 {
9206 let csr_tok: Vec<i32> = ex_pairs.iter().map(|&p| p / n_used as i32).collect();
9207 let csr_tok_d = e.htod_i32(&csr_tok)?;
9208 let (z_f16, z_s) = e.moe_f16g_act(moe_in, Some(&csr_tok_d), n_embd, n_pairs)?;
9209 let g_csr = e.moe_f16_grouped(
9210 &dev.ptr_row,
9211 0,
9212 n_expert,
9213 &exi,
9214 &ex_off,
9215 &exo,
9216 &z_f16,
9217 &z_s,
9218 n_embd,
9219 n_ff_exp,
9220 n_active,
9221 n_pairs,
9222 m.gate_exps.qtype,
9223 m.gate_exps.row_bytes,
9224 )?;
9225 let u_csr = e.moe_f16_grouped(
9226 &dev.ptr_row,
9227 1,
9228 n_expert,
9229 &exi,
9230 &ex_off,
9231 &exo,
9232 &z_f16,
9233 &z_s,
9234 n_embd,
9235 n_ff_exp,
9236 n_active,
9237 n_pairs,
9238 m.up_exps.qtype,
9239 m.up_exps.row_bytes,
9240 )?;
9241 let act_csr = e.moe_pairs_gelu_mul(&g_csr, &u_csr, n_pairs * n_ff_exp)?;
9242 let (a_f16, a_s) = e.moe_f16g_act(&act_csr, None, n_ff_exp, n_pairs)?;
9243 let d_csr = e.moe_f16_grouped(
9244 &dev.ptr_row,
9245 2,
9246 n_expert,
9247 &exi,
9248 &ex_off,
9249 &exo,
9250 &a_f16,
9251 &a_s,
9252 n_ff_exp,
9253 n_embd,
9254 n_active,
9255 n_pairs,
9256 m.down_exps.qtype,
9257 m.down_exps.row_bytes,
9258 )?;
9259 let y_down = e.rows_permute(&d_csr, &exp_d, n_pairs, n_embd)?;
9260 let mut moe_out = e.uninit(t * n_embd)?;
9261 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
9262 if std::env::var("MEMRA_F16G_DEBUG").is_ok() {
9263 let scan = |v: &[f32]| v.iter().filter(|x| !x.is_finite()).count();
9264 let (yd, mo) = (e.dtoh(&y_down)?, e.dtoh(&moe_out)?);
9265 eprintln!(
9266 "[f16g-debug] post-permute bad={} post-scatter bad={}",
9267 scan(&yd),
9268 scan(&mo)
9269 );
9270 }
9271 return Ok(moe_out);
9272 }
9273 let mma =
9276 n_embd % 256 == 0 && std::env::var("MEMRA_GEMMA_MOE_MMA").as_deref() != Ok("0");
9277 let (gate, up) = if mma {
9278 let z_scr = e.mmq_iq_quantize_act(moe_in, n_embd, t)?;
9279 (
9280 e.mmq_iq_experts(
9281 &dev.ptr_row,
9282 0,
9283 n_expert,
9284 &exi,
9285 &exo,
9286 &exp_d,
9287 &pt,
9288 &z_scr,
9289 n_embd,
9290 n_ff_exp,
9291 n_active,
9292 n_pairs,
9293 t,
9294 m.gate_exps.qtype,
9295 m.gate_exps.row_bytes,
9296 )?,
9297 e.mmq_iq_experts(
9298 &dev.ptr_row,
9299 1,
9300 n_expert,
9301 &exi,
9302 &exo,
9303 &exp_d,
9304 &pt,
9305 &z_scr,
9306 n_embd,
9307 n_ff_exp,
9308 n_active,
9309 n_pairs,
9310 t,
9311 m.up_exps.qtype,
9312 m.up_exps.row_bytes,
9313 )?,
9314 )
9315 } else {
9316 let (zq, zd) = e.quantize_q8_1(moe_in, t, n_embd)?;
9317 (
9318 e.moe_pairs_matvec_q8_dec(
9319 &dev.ptr_row,
9320 0,
9321 &exi,
9322 &exo,
9323 &exp_d,
9324 &pt,
9325 &zq,
9326 &zd,
9327 n_embd,
9328 n_ff_exp,
9329 n_expert,
9330 n_active,
9331 n_pairs,
9332 m.gate_exps.qtype,
9333 m.gate_exps.row_bytes,
9334 )?,
9335 e.moe_pairs_matvec_q8_dec(
9336 &dev.ptr_row,
9337 1,
9338 &exi,
9339 &exo,
9340 &exp_d,
9341 &pt,
9342 &zq,
9343 &zd,
9344 n_embd,
9345 n_ff_exp,
9346 n_expert,
9347 n_active,
9348 n_pairs,
9349 m.up_exps.qtype,
9350 m.up_exps.row_bytes,
9351 )?,
9352 )
9353 };
9354 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
9355 let pself = e.htod_i32(&pair_self)?;
9356 let y_down = if mma {
9368 let in_pad = n_ff_exp.div_ceil(256) * 256;
9369 let a_scr = if crate::moe_fuse_actq_on() {
9370 e.mmq_iq_fused_act_quant(&gate, &up, n_ff_exp, n_pairs, 1)?
9371 } else {
9372 let act = e.moe_pairs_gelu_mul(&gate, &up, n_pairs * n_ff_exp)?;
9373 e.mmq_iq_quantize_act(&act, n_ff_exp, n_pairs)?
9374 };
9375 e.mmq_iq_experts(
9376 &dev.ptr_row,
9377 2,
9378 n_expert,
9379 &exi,
9380 &exo,
9381 &exp_d,
9382 &pself,
9383 &a_scr,
9384 in_pad,
9385 n_embd,
9386 n_active,
9387 n_pairs,
9388 n_pairs,
9389 m.down_exps.qtype,
9390 m.down_exps.row_bytes,
9391 )?
9392 } else {
9393 let act = e.moe_pairs_gelu_mul(&gate, &up, n_pairs * n_ff_exp)?;
9394 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
9395 e.moe_pairs_matvec_q8_dec(
9396 &dev.ptr_row,
9397 2,
9398 &exi,
9399 &exo,
9400 &exp_d,
9401 &pself,
9402 &aq2,
9403 &ad2,
9404 n_ff_exp,
9405 n_embd,
9406 n_expert,
9407 n_active,
9408 n_pairs,
9409 m.down_exps.qtype,
9410 m.down_exps.row_bytes,
9411 )?
9412 };
9413 let mut moe_out = e.uninit(t * n_embd)?;
9414 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
9415 return Ok(moe_out);
9416 }
9417
9418 let g_len = m.gate_exps.expert_stride;
9419 let u_len = m.up_exps.expert_stride;
9420 let d_len = m.down_exps.expert_stride;
9421 let dev = m.dev_exps.as_ref().filter(|d| !d.gu_il);
9425 let (mut sg, mut su, mut sd) = if dev.is_some() {
9426 (None, None, None)
9427 } else {
9428 (
9429 Some(e.alloc_u8_uninit(g_len)?),
9430 Some(e.alloc_u8_uninit(u_len)?),
9431 Some(e.alloc_u8_uninit(d_len)?),
9432 )
9433 };
9434 let mut moe_out = e.zeros(t * n_embd)?;
9435 for tok in 0..t {
9436 let sel = &sel_all[tok * n_used..(tok + 1) * n_used];
9437 let w = &w_all[tok * n_used..(tok + 1) * n_used];
9438 let zt = moe_in.slice(tok * n_embd..(tok + 1) * n_embd);
9439 for (j, &ex) in sel.iter().enumerate() {
9440 let ex = ex as usize;
9441 let gate = match dev {
9442 Some(d) => e.qmatvec_view(
9443 &d.gate,
9444 ex * g_len..(ex + 1) * g_len,
9445 &zt,
9446 1,
9447 m.gate_exps.in_f,
9448 m.gate_exps.out_f,
9449 m.gate_exps.qtype,
9450 m.gate_exps.row_bytes,
9451 )?,
9452 None => {
9453 let sg = sg.as_mut().unwrap();
9454 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
9455 e.qmatvec_view(
9456 sg,
9457 0..g_len,
9458 &zt,
9459 1,
9460 m.gate_exps.in_f,
9461 m.gate_exps.out_f,
9462 m.gate_exps.qtype,
9463 m.gate_exps.row_bytes,
9464 )?
9465 }
9466 };
9467 let up = match dev {
9468 Some(d) => e.qmatvec_view(
9469 &d.up,
9470 ex * u_len..(ex + 1) * u_len,
9471 &zt,
9472 1,
9473 m.up_exps.in_f,
9474 m.up_exps.out_f,
9475 m.up_exps.qtype,
9476 m.up_exps.row_bytes,
9477 )?,
9478 None => {
9479 let su = su.as_mut().unwrap();
9480 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
9481 e.qmatvec_view(
9482 su,
9483 0..u_len,
9484 &zt,
9485 1,
9486 m.up_exps.in_f,
9487 m.up_exps.out_f,
9488 m.up_exps.qtype,
9489 m.up_exps.row_bytes,
9490 )?
9491 }
9492 };
9493 let mut act = e.uninit(n_ff_exp)?;
9494 e.gelu_tanh_mul(&gate, &up, &mut act, n_ff_exp)?;
9495 let actv = act.slice(0..n_ff_exp);
9496 let y = match dev {
9497 Some(d) => e.qmatvec_view(
9498 &d.down,
9499 ex * d_len..(ex + 1) * d_len,
9500 &actv,
9501 1,
9502 m.down_exps.in_f,
9503 m.down_exps.out_f,
9504 m.down_exps.qtype,
9505 m.down_exps.row_bytes,
9506 )?,
9507 None => {
9508 let sd = sd.as_mut().unwrap();
9509 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
9510 e.qmatvec_view(
9511 sd,
9512 0..d_len,
9513 &actv,
9514 1,
9515 m.down_exps.in_f,
9516 m.down_exps.out_f,
9517 m.down_exps.qtype,
9518 m.down_exps.row_bytes,
9519 )?
9520 }
9521 };
9522 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
9523 e.axpy_into(&y, w[j], &mut dst, n_embd)?;
9524 }
9525 }
9526 Ok(moe_out)
9527 }
9528
9529 fn gemma4_layer(
9531 &self,
9532 e: &Engine,
9533 il: usize,
9534 layer: &crate::hybrid::HybridLayer,
9535 x: &CudaSlice<f32>,
9536 pos_d: &CudaSlice<i32>,
9537 t: usize,
9538 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9539 let n_embd = self.cfg.n_embd as usize;
9540 let eps = self.cfg.rms_eps;
9541
9542 let mut h = e.zeros(t * n_embd)?;
9543 e.rms_norm(x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
9544 let Mixer::Full(fa) = &layer.mixer else {
9545 panic!("gemma4 layer {il} not full-attn")
9546 };
9547 let o = self.gemma4_attn(e, fa, il, &h, pos_d, t)?;
9548 let mut cur = e.zeros(t * n_embd)?;
9550 e.rms_norm(
9551 &o,
9552 layer.post_attn_norm.float_data(),
9553 &mut cur,
9554 n_embd,
9555 t,
9556 eps,
9557 )?;
9558 self.gemma4_layer_tail_add(e, layer, &cur, x, t)
9559 }
9560
9561 fn gemma4_layer_tail_add(
9565 &self,
9566 e: &Engine,
9567 layer: &crate::hybrid::HybridLayer,
9568 cur: &CudaSlice<f32>,
9569 x: &CudaSlice<f32>,
9570 t: usize,
9571 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9572 Ok(self.gemma4_layer_tail_add_n(e, layer, cur, x, t, None)?.0)
9573 }
9574
9575 fn gemma4_layer_tail_add_n(
9578 &self,
9579 e: &Engine,
9580 layer: &crate::hybrid::HybridLayer,
9581 cur: &CudaSlice<f32>,
9582 x: &CudaSlice<f32>,
9583 t: usize,
9584 next_norm: Option<&CudaSlice<f32>>,
9585 ) -> Result<(CudaSlice<f32>, Option<CudaSlice<f32>>), Box<dyn std::error::Error>> {
9586 let n_embd = self.cfg.n_embd as usize;
9587 let bits = layer.gemma4.as_ref().unwrap();
9588 let (sn, attn_out) = self.gemma4_layer_tail_core(e, layer, cur, x, t)?;
9589 let mut xn = e.uninit(t * n_embd)?;
9590 match next_norm {
9591 Some(w) => {
9592 let mut hn = e.uninit(t * n_embd)?;
9593 e.add_scale_rms_norm(
9594 &sn,
9595 &attn_out,
9596 bits.layer_scale,
9597 w,
9598 &mut xn,
9599 &mut hn,
9600 n_embd,
9601 t,
9602 self.cfg.rms_eps,
9603 )?;
9604 Ok((xn, Some(hn)))
9605 }
9606 None => {
9607 e.add_scale(&sn, &attn_out, bits.layer_scale, &mut xn, t * n_embd)?;
9608 Ok((xn, None))
9609 }
9610 }
9611 }
9612
9613 fn gemma4_layer_tail_core(
9616 &self,
9617 e: &Engine,
9618 layer: &crate::hybrid::HybridLayer,
9619 cur: &CudaSlice<f32>,
9620 x: &CudaSlice<f32>,
9621 t: usize,
9622 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
9623 self.gemma4_layer_tail_core_pn(e, layer, cur, x, t, None, false)
9624 }
9625
9626 fn gemma4_layer_tail_core_pn(
9633 &self,
9634 e: &Engine,
9635 layer: &crate::hybrid::HybridLayer,
9636 cur: &CudaSlice<f32>,
9637 x: &CudaSlice<f32>,
9638 t: usize,
9639 pre_norm: Option<&CudaSlice<f32>>,
9640 defer_post_norm: bool,
9641 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
9642 let n_embd = self.cfg.n_embd as usize;
9643 let eps = self.cfg.rms_eps;
9644 let bits = layer.gemma4.as_ref().unwrap();
9645
9646 let Some(mbits) = bits.moe_bits.as_ref() else {
9649 let crate::hybrid::Ffn::Dense {
9650 ffn_gate,
9651 ffn_up,
9652 ffn_down,
9653 } = &layer.ffn
9654 else {
9655 panic!("gemma4 dense layer without Dense ffn")
9656 };
9657 let mut attn_out = e.uninit(t * n_embd)?;
9658 let mut zsh = e.uninit(t * n_embd)?;
9659 let mut zpair: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
9662 match pre_norm {
9663 Some(wa) if t == 1 => {
9664 zpair = Some(e.rms_pre_add_rms_norm_q8z(
9665 cur,
9666 wa,
9667 x,
9668 bits.ffn_norm.float_data(),
9669 &mut attn_out,
9670 &mut zsh,
9671 n_embd,
9672 t,
9673 eps,
9674 )?);
9675 }
9676 Some(wa) => e.rms_pre_add_rms_norm(
9677 cur,
9678 wa,
9679 x,
9680 bits.ffn_norm.float_data(),
9681 &mut attn_out,
9682 &mut zsh,
9683 n_embd,
9684 t,
9685 eps,
9686 )?,
9687 None => e.add_rms_norm(
9688 cur,
9689 x,
9690 bits.ffn_norm.float_data(),
9691 &mut attn_out,
9692 &mut zsh,
9693 n_embd,
9694 t,
9695 eps,
9696 )?,
9697 }
9698 let n_ff = ffn_gate.out_features();
9699 let (gate, up) = if t == 1 {
9705 let (zq, zd) = match zpair {
9706 Some(p) => p,
9707 None => e.quantize_q8_1(&zsh, 1, n_embd)?,
9708 };
9709 match e.matmul_q4_fused2(ffn_gate, ffn_up, &zq, &zd)? {
9710 Some(p) => p,
9711 None => match e.matmul_nvfp4_fused2(ffn_gate, ffn_up, &zq, &zd, 1)? {
9713 Some(p) => p,
9714 None => (
9715 e.matmul_pre(ffn_gate, &zq, &zd, &zsh, 1)?,
9716 e.matmul_pre(ffn_up, &zq, &zd, &zsh, 1)?,
9717 ),
9718 },
9719 }
9720 } else {
9721 static F2B: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9726 let f2b = *F2B.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
9727 let fused = if f2b {
9728 let (zq, zd) = e.quantize_q8_1(&zsh, t, n_embd)?;
9729 e.matmul_q4_fused2_batched(ffn_gate, ffn_up, &zq, &zd, t)?
9730 } else {
9731 None
9732 };
9733 match fused {
9734 Some(p) => p,
9735 None => {
9736 e.mmq_act_begin();
9738 (e.matmul(ffn_gate, &zsh, t)?, e.matmul(ffn_up, &zsh, t)?)
9739 }
9740 }
9741 };
9742 let mut act = e.uninit(t * n_ff)?;
9743 let f0 = if e.uses_q8_1_fast(ffn_down) {
9746 let upv = e.view(&up, t * n_ff);
9747 let up_all = upv.slice(0..t * n_ff);
9748 let (aq, ad) = e.gelu_tanh_mul_q8_1(&gate, &up_all, &mut act, n_ff, t)?;
9749 e.matmul_pre(ffn_down, &aq, &ad, &act, t)?
9750 } else {
9751 e.gelu_tanh_mul(&gate, &up, &mut act, t * n_ff)?;
9752 e.matmul(ffn_down, &act, t)?
9753 };
9754 if defer_post_norm {
9755 return Ok((f0, attn_out));
9756 }
9757 let mut sn = e.uninit(t * n_embd)?;
9758 e.rms_norm(
9759 &f0,
9760 bits.post_ffw_norm.float_data(),
9761 &mut sn,
9762 n_embd,
9763 t,
9764 eps,
9765 )?;
9766 return Ok((sn, attn_out));
9767 };
9768
9769 assert!(pre_norm.is_none(), "pre-norm fold is dense-entry only");
9770 let mut attn_out = e.uninit(t * n_embd)?;
9775 let mut router_in = e.uninit(t * n_embd)?;
9776 let fast_moe = match &layer.ffn {
9777 crate::hybrid::Ffn::Moe(m) => {
9778 m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
9779 && expert_dp4a_supported(m.gate_exps.qtype)
9780 && expert_dp4a_supported(m.up_exps.qtype)
9781 && expert_dp4a_supported(m.down_exps.qtype)
9782 && std::env::var("MEMRA_GEMMA_MOE_FAST").as_deref() != Ok("0")
9783 }
9784 _ => false,
9785 };
9786 let q8z = t < PRIME_MIN_T && fast_moe;
9787 let (zsh_f32, zsh_q8, moe_q8) = if q8z {
9788 let (z0, m2) = e.add_rms_norm3_q8z(
9789 cur,
9790 x,
9791 bits.ffn_norm.float_data(),
9792 &mbits.router_scale_pre,
9793 mbits.pre_ffw_norm_2.float_data(),
9794 &mut attn_out,
9795 &mut router_in,
9796 n_embd,
9797 t,
9798 eps,
9799 )?;
9800 (None, Some(z0), Some(m2))
9801 } else {
9802 let mut zsh = e.uninit(t * n_embd)?;
9803 let mut moe_in = e.uninit(t * n_embd)?;
9804 e.add_rms_norm3(
9805 cur,
9806 x,
9807 bits.ffn_norm.float_data(),
9808 &mbits.router_scale_pre,
9809 mbits.pre_ffw_norm_2.float_data(),
9810 &mut attn_out,
9811 &mut zsh,
9812 &mut router_in,
9813 &mut moe_in,
9814 n_embd,
9815 t,
9816 eps,
9817 )?;
9818 (Some((zsh, moe_in)), None, None)
9819 };
9820 let attn_out2 = attn_out;
9821 #[allow(unused_variables)]
9822 let attn_out = &attn_out2;
9823 let n_ff = mbits.shared_gate.out_features();
9824 let (gate, up) = if let Some((zq, zd)) = zsh_q8.as_ref() {
9825 if t == 1 {
9826 match e.matmul_q4_fused2(&mbits.shared_gate, &mbits.shared_up, zq, zd)? {
9827 Some(p) => p,
9828 None => match e.matmul_nvfp4_fused2(
9829 &mbits.shared_gate,
9830 &mbits.shared_up,
9831 zq,
9832 zd,
9833 1,
9834 )? {
9835 Some(p) => p,
9836 None => {
9837 let h0 = e.zeros(0)?;
9838 (
9839 e.matmul_pre(&mbits.shared_gate, zq, zd, &h0, 1)?,
9840 e.matmul_pre(&mbits.shared_up, zq, zd, &h0, 1)?,
9841 )
9842 }
9843 },
9844 }
9845 } else {
9846 let h0 = e.zeros(0)?;
9848 (
9849 e.matmul_pre(&mbits.shared_gate, zq, zd, &h0, t)?,
9850 e.matmul_pre(&mbits.shared_up, zq, zd, &h0, t)?,
9851 )
9852 }
9853 } else {
9854 let (zsh, _) = zsh_f32.as_ref().unwrap();
9855 (
9856 e.matmul(&mbits.shared_gate, zsh, t)?,
9857 e.matmul(&mbits.shared_up, zsh, t)?,
9858 )
9859 };
9860 let mut act = e.uninit(t * n_ff)?;
9861 e.gelu_tanh_mul(&gate, &up, &mut act, t * n_ff)?;
9862 let mlp0 = e.matmul(&mbits.shared_down, &act, t)?;
9863 let crate::hybrid::Ffn::Moe(m) = &layer.ffn else {
9864 panic!("gemma4 layer not MoE")
9865 };
9866 let moe0 = match (&moe_q8, &zsh_f32) {
9867 (Some(mq), _) => self.gemma4_moe_q8(e, m, mbits, mq, &router_in, t)?,
9868 (None, Some((_, moe_in))) => self.gemma4_moe(e, m, mbits, moe_in, &router_in, t)?,
9869 _ => unreachable!(),
9870 };
9871 let mut mlp = e.uninit(t * n_embd)?;
9873 let mut moe = e.uninit(t * n_embd)?;
9874 e.rms_norm2x(
9875 &mlp0,
9876 &moe0,
9877 mbits.post_ffw_norm_1.float_data(),
9878 mbits.post_ffw_norm_2.float_data(),
9879 &mut mlp,
9880 &mut moe,
9881 n_embd,
9882 t,
9883 eps,
9884 )?;
9885
9886 let mut sum = e.uninit(t * n_embd)?;
9889 let mut sn = e.uninit(t * n_embd)?;
9890 e.add_rms_norm(
9891 &mlp,
9892 &moe,
9893 bits.post_ffw_norm.float_data(),
9894 &mut sum,
9895 &mut sn,
9896 n_embd,
9897 t,
9898 eps,
9899 )?;
9900 Ok((sn, attn_out2))
9901 }
9902
9903 pub(crate) fn gemma4_layer_tail_add_nq_pn(
9913 &self,
9914 e: &Engine,
9915 layer: &crate::hybrid::HybridLayer,
9916 o: &CudaSlice<f32>,
9917 x: &CudaSlice<f32>,
9918 t: usize,
9919 next_norm: Option<&CudaSlice<f32>>,
9920 ) -> Result<(CudaSlice<f32>, Option<(CudaSlice<i8>, CudaSlice<f32>)>), Box<dyn std::error::Error>>
9921 {
9922 let n_embd = self.cfg.n_embd as usize;
9923 let eps = self.cfg.rms_eps;
9924 let bits = layer.gemma4.as_ref().unwrap();
9925 if Engine::g4_pnfold_on() && matches!(layer.ffn, crate::hybrid::Ffn::Dense { .. }) {
9926 let (f0, attn_out) = self.gemma4_layer_tail_core_pn(
9927 e,
9928 layer,
9929 o,
9930 x,
9931 t,
9932 Some(layer.post_attn_norm.float_data()),
9933 true,
9934 )?;
9935 let mut xn = e.uninit(t * n_embd)?;
9936 return match next_norm {
9937 Some(w) => {
9938 let pair = e.rms_pre_add_scale_rms_norm_q8_1(
9939 &f0,
9940 bits.post_ffw_norm.float_data(),
9941 &attn_out,
9942 bits.layer_scale,
9943 w,
9944 &mut xn,
9945 n_embd,
9946 t,
9947 eps,
9948 )?;
9949 Ok((xn, Some(pair)))
9950 }
9951 None => {
9952 let mut sn = e.uninit(t * n_embd)?;
9953 e.rms_norm(
9954 &f0,
9955 bits.post_ffw_norm.float_data(),
9956 &mut sn,
9957 n_embd,
9958 t,
9959 eps,
9960 )?;
9961 e.add_scale(&sn, &attn_out, bits.layer_scale, &mut xn, t * n_embd)?;
9962 Ok((xn, None))
9963 }
9964 };
9965 }
9966 let mut cur = e.uninit(t * n_embd)?;
9967 e.rms_norm(
9968 o,
9969 layer.post_attn_norm.float_data(),
9970 &mut cur,
9971 n_embd,
9972 t,
9973 eps,
9974 )?;
9975 self.gemma4_layer_tail_add_nq(e, layer, &cur, x, t, next_norm)
9976 }
9977
9978 pub(crate) fn gemma4_layer_tail_add_nq(
9979 &self,
9980 e: &Engine,
9981 layer: &crate::hybrid::HybridLayer,
9982 cur: &CudaSlice<f32>,
9983 x: &CudaSlice<f32>,
9984 t: usize,
9985 next_norm: Option<&CudaSlice<f32>>,
9986 ) -> Result<(CudaSlice<f32>, Option<(CudaSlice<i8>, CudaSlice<f32>)>), Box<dyn std::error::Error>>
9987 {
9988 let n_embd = self.cfg.n_embd as usize;
9989 let bits = layer.gemma4.as_ref().unwrap();
9990 let (sn, attn_out) = self.gemma4_layer_tail_core(e, layer, cur, x, t)?;
9991 let mut xn = e.uninit(t * n_embd)?;
9992 match next_norm {
9993 Some(w) => {
9994 let pair = e.add_scale_rms_norm_q8_1(
9995 &sn,
9996 &attn_out,
9997 bits.layer_scale,
9998 w,
9999 &mut xn,
10000 n_embd,
10001 t,
10002 self.cfg.rms_eps,
10003 )?;
10004 Ok((xn, Some(pair)))
10005 }
10006 None => {
10007 e.add_scale(&sn, &attn_out, bits.layer_scale, &mut xn, t * n_embd)?;
10008 Ok((xn, None))
10009 }
10010 }
10011 }
10012
10013 fn gemma4_forward(
10016 &self,
10017 e: &Engine,
10018 tokens: &[u32],
10019 last_only: bool,
10020 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
10021 if self.is_gemma4_e4b() {
10024 return self.gemma4_e4b_forward(e, tokens, last_only);
10025 }
10026 let n_embd = self.cfg.n_embd as usize;
10027 let t = tokens.len();
10028 let pos: Vec<i32> = (0..t as i32).collect();
10029 let pos_d = e.htod_i32(&pos)?;
10030
10031 let mut x = self.embed(e, tokens)?;
10032 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
10033 let probe = std::env::var("MEMRA_GEMMA_PROBE").is_ok();
10036 let stat =
10037 |e: &Engine, x: &CudaSlice<f32>, tag: &str| -> Result<(), Box<dyn std::error::Error>> {
10038 let h = e.dtoh(x)?;
10039 let bad = h.iter().filter(|v| !v.is_finite()).count();
10040 let mx = h
10041 .iter()
10042 .filter(|v| v.is_finite())
10043 .fold(0.0f32, |m, v| m.max(v.abs()));
10044 eprintln!(
10045 "[gemma-probe] {tag}: tok0_first3={:?} bad={bad} max={mx:.3e}",
10046 &h[..3]
10047 );
10048 Ok(())
10049 };
10050 if probe {
10051 stat(e, &x, "embed")?;
10052 }
10053 for (il, layer) in self.layers.iter().enumerate() {
10054 x = self.gemma4_layer(e, il, layer, &x, &pos_d, t)?;
10055 if probe {
10056 stat(e, &x, &format!("L{il}"))?;
10057 }
10058 }
10059 let mut hn = e.zeros(t * n_embd)?;
10060 e.rms_norm(
10061 &x,
10062 self.output_norm.float_data(),
10063 &mut hn,
10064 n_embd,
10065 t,
10066 self.cfg.rms_eps,
10067 )?;
10068 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
10069 let n_vocab = self.output.out_features();
10070 let logits = if last_only {
10071 let hv = e.view(&hn, t * n_embd);
10072 let last_row = hv.slice((t - 1) * n_embd..t * n_embd);
10073 let mut hlast = e.zeros(n_embd)?;
10074 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
10075 let mut ld = e.matmul(&self.output, &hlast, 1)?;
10076 e.softcap(&mut ld, cap, n_vocab)?;
10077 self.gemma4_suppress(e, &mut ld, 1)?;
10078 e.dtoh(&ld)?
10079 } else {
10080 let mut ld = e.matmul(&self.output, &hn, t)?;
10081 e.softcap(&mut ld, cap, t * n_vocab)?;
10082 self.gemma4_suppress(e, &mut ld, t)?;
10083 e.dtoh(&ld)?
10084 };
10085 Ok(logits)
10086 }
10087
10088 pub(crate) fn gemma4_prime(
10093 &self,
10094 e: &Engine,
10095 tokens: &[u32],
10096 cache: &mut Cache,
10097 overlay: Option<&crate::vision::EmbedOverlay>,
10098 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
10099 if cache.pos != 0 {
10104 return Err(
10105 "gemma4 prime v0 is fresh-prompt only (no continuation/chunked prime) \
10106 — prime the full prompt in one call or decode tokenwise"
10107 .into(),
10108 );
10109 }
10110 let n_embd = self.cfg.n_embd as usize;
10111 let eps = self.cfg.rms_eps;
10112 let t = tokens.len();
10113 let pos: Vec<i32> = (0..t as i32).collect();
10114 let pos_d = e.htod_i32(&pos)?;
10115 let mut x = self.embed(e, tokens)?;
10116 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
10117 let island: Option<CudaSlice<i32>> = match overlay {
10124 Some(ov) => {
10125 let mut span_id = vec![-1i32; t];
10126 for (i, &(pos, row_off, n_rows)) in ov.spans.iter().enumerate() {
10127 if pos + n_rows > t {
10128 return Err(format!(
10129 "gemma4 overlay span {i} [{pos}, {}) exceeds the prompt ({t})",
10130 pos + n_rows
10131 )
10132 .into());
10133 }
10134 let view = ov.rows.slice(row_off * n_embd..(row_off + n_rows) * n_embd);
10135 e.copy_view_into(&mut x, pos * n_embd, &view, n_rows * n_embd)?;
10136 for s in span_id.iter_mut().skip(pos).take(n_rows) {
10137 *s = i as i32;
10138 }
10139 }
10140 if std::env::var("MEMRA_GV_FORCE_CAUSAL").as_deref() == Ok("1") {
10144 None
10145 } else {
10146 Some(e.htod_i32(&span_id)?)
10147 }
10148 }
10149 None => None,
10150 };
10151 for (il, layer) in self.layers.iter().enumerate() {
10152 let mut h = e.zeros(t * n_embd)?;
10153 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
10154 let Mixer::Full(fa) = &layer.mixer else {
10155 panic!("gemma4 layer not full-attn")
10156 };
10157 let trace = il == 0 && std::env::var("MEMRA_G4_PRIME_TRACE").as_deref() == Ok("1");
10158 if trace {
10159 let v = e.dtoh(&h)?;
10160 let nan = v.iter().filter(|x| x.is_nan()).count();
10161 eprintln!("[g4-prime-trace] L0 post-attn_norm: nan={nan}/{}", v.len());
10162 }
10163 let o =
10164 self.gemma4_attn_prime(e, fa, il, &h, &pos_d, t, Some(cache), island.as_ref())?;
10165 if trace {
10166 let v = e.dtoh(&o)?;
10167 let nan = v.iter().filter(|x| x.is_nan()).count();
10168 eprintln!("[g4-prime-trace] L0 post-attn: nan={nan}/{}", v.len());
10169 }
10170 let mut cur = e.zeros(t * n_embd)?;
10171 e.rms_norm(
10172 &o,
10173 layer.post_attn_norm.float_data(),
10174 &mut cur,
10175 n_embd,
10176 t,
10177 eps,
10178 )?;
10179 x = self.gemma4_layer_tail_add(e, layer, &cur, &x, t)?;
10180 self.dflash_tap(e, cache, il, &x, t)?;
10181 if std::env::var("MEMRA_G4_PRIME_TRACE").as_deref() == Ok("1") {
10183 let h = e.dtoh(&x)?;
10184 let nan = h.iter().filter(|v| v.is_nan()).count();
10185 let amax = h.iter().fold(0f32, |a, v| a.max(v.abs()));
10186 eprintln!(
10187 "[g4-prime-trace] layer {il}: nan={nan}/{} amax={amax:.3}",
10188 h.len()
10189 );
10190 if nan > 0 {
10191 return Err(format!("g4-prime-trace: first NaN at layer {il}").into());
10192 }
10193 }
10194 }
10195 cache.pos += t;
10196 let hiddens = e.clone_dtod(&x)?;
10197 let xv = e.view(&x, t * n_embd);
10198 let last_row = xv.slice((t - 1) * n_embd..t * n_embd);
10199 let mut h_seed = e.zeros(n_embd)?;
10200 e.copy_view_into(&mut h_seed, 0, &last_row, n_embd)?;
10201 let mut hn = e.uninit(n_embd)?;
10202 e.rms_norm(
10203 &h_seed,
10204 self.output_norm.float_data(),
10205 &mut hn,
10206 n_embd,
10207 1,
10208 eps,
10209 )?;
10210 let mut ld = e.matmul(&self.output, &hn, 1)?;
10211 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
10212 e.softcap(&mut ld, cap, self.output.out_features())?;
10213 self.gemma4_suppress(e, &mut ld, 1)?;
10214 let logits = e.dtoh(&ld)?;
10215 Ok((logits, h_seed, hiddens))
10216 }
10217
10218 fn gemma4_decode_attn(
10223 &self,
10224 e: &Engine,
10225 fa: &crate::hybrid::FullAttnLayer,
10226 il: usize,
10227 hq: &CudaSlice<i8>,
10228 hdq: &CudaSlice<f32>,
10229 pos_d: &CudaSlice<i32>,
10230 cache: &mut Cache,
10231 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
10232 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
10233 let eps = self.cfg.rms_eps;
10234 let aux = self.gemma4_aux.as_ref().unwrap();
10235 let ones = aux.ones(e);
10236 #[cfg(debug_assertions)]
10237 crate::debug_assert_tensor_stream_device(ones, &e.stream(), "gemma4_decode_attn.ones");
10238 let (hq, hdq) = (hq, hdq);
10239 let h0 = e.zeros(0)?;
10240 let h = &h0;
10241 let (q0, k0, v0) = if swa {
10242 match e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, &hq, &hdq)? {
10243 Some(t3) => t3,
10244 None => match e.matmul_nvfp4_fused2(&fa.wq, &fa.wk, &hq, &hdq, 1)? {
10247 Some((q0, k0)) => {
10248 let v0 = e.matmul_pre(&fa.wv, &hq, &hdq, h, 1)?;
10249 (q0, k0, v0)
10250 }
10251 None => (
10252 e.matmul_pre(&fa.wq, &hq, &hdq, h, 1)?,
10253 e.matmul_pre(&fa.wk, &hq, &hdq, h, 1)?,
10254 e.matmul_pre(&fa.wv, &hq, &hdq, h, 1)?,
10255 ),
10256 },
10257 }
10258 } else {
10259 let (q0, k0) = match e.matmul_q4_fused2(&fa.wq, &fa.wk, &hq, &hdq)? {
10260 Some(p) => p,
10261 None => match e.matmul_nvfp4_fused2(&fa.wq, &fa.wk, &hq, &hdq, 1)? {
10262 Some(p) => p,
10263 None => (
10264 e.matmul_pre(&fa.wq, &hq, &hdq, h, 1)?,
10265 e.matmul_pre(&fa.wk, &hq, &hdq, h, 1)?,
10266 ),
10267 },
10268 };
10269 let v0 = e.clone_dtod(&k0)?;
10270 (q0, k0, v0)
10271 };
10272 let mut q = e.uninit(nh * hd)?;
10273 let mut k = e.uninit(nkv * hd)?;
10274 let mut v = e.uninit(nkv * hd)?;
10275 let ff = if swa {
10278 None
10279 } else {
10280 Some(
10281 aux.rope_freqs(e)
10282 .expect("gemma4 global rope needs rope_freqs.weight"),
10283 )
10284 };
10285 #[cfg(debug_assertions)]
10286 if let Some(ff) = ff {
10287 crate::debug_assert_tensor_stream_device(
10288 ff,
10289 &e.stream(),
10290 "gemma4_decode_attn.rope_freqs",
10291 );
10292 }
10293 let kvl = cache.kv[il].as_mut().unwrap();
10294 let kv_fp8 = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
10295 if crate::Engine::qkv_append_on() {
10296 e.rms_norm_qkv_rope_append(
10300 &q0,
10301 &k0,
10302 &v0,
10303 fa.q_norm.float_data(),
10304 fa.k_norm.float_data(),
10305 ones,
10306 &mut q,
10307 &mut k,
10308 &mut v,
10309 hd,
10310 nh,
10311 nkv,
10312 pos_d,
10313 nh,
10314 nkv,
10315 base,
10316 1.0,
10317 ff,
10318 eps,
10319 &mut kvl.k,
10320 &mut kvl.v,
10321 kvl.len,
10322 kvl.k_tok_bytes,
10323 kvl.v_tok_bytes,
10324 kv_fp8,
10325 )?;
10326 } else {
10327 e.rms_norm_qkv_rope(
10328 &q0,
10329 &k0,
10330 &v0,
10331 fa.q_norm.float_data(),
10332 fa.k_norm.float_data(),
10333 ones,
10334 &mut q,
10335 &mut k,
10336 &mut v,
10337 hd,
10338 nh,
10339 nkv,
10340 pos_d,
10341 nh,
10342 nkv,
10343 base,
10344 1.0,
10345 ff,
10346 eps,
10347 )?;
10348 e.append_kv_quantized(
10349 &k,
10350 &v,
10351 &mut kvl.k,
10352 &mut kvl.v,
10353 kvl.len,
10354 kvl.kv_dim_k,
10355 kvl.kv_dim_v,
10356 kvl.k_tok_bytes,
10357 kvl.v_tok_bytes,
10358 kv_fp8,
10359 )?;
10360 }
10361 kvl.len += 1;
10362 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
10366 let mut attn = e.uninit(nh * hd)?;
10367 if !swa
10369 && hd == 512
10370 && kvl.len >= crate::fa512_min_tkv()
10371 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
10372 {
10373 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
10374 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
10375 let base = kvl.len as i32;
10377 e.i32_set_k(&mut kvl.len_d, base)?;
10378 e.fa_decode_rows(
10379 &q,
10380 &kp,
10381 &vp,
10382 &mut attn,
10383 hd,
10384 nh,
10385 nkv,
10386 kvl.len - 1,
10387 1,
10388 scale,
10389 kvl.k_tok_bytes,
10390 kvl.v_tok_bytes,
10391 Some((&kvl.len_d, -1)),
10392 false,
10393 false,
10394 None,
10395 )?;
10396 return Ok(e.matmul(&fa.wo, &attn, 1)?);
10397 }
10398 if swa
10400 && kvl.len > win
10401 && hd == 256
10402 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
10403 {
10404 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
10405 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
10406 let base = kvl.len as i32;
10407 e.i32_set_k(&mut kvl.len_d, base)?;
10408 e.fa_decode_rows_w(
10409 &q,
10410 &kp,
10411 &vp,
10412 &mut attn,
10413 hd,
10414 nh,
10415 nkv,
10416 &kvl.len_d,
10417 -1,
10418 1,
10419 scale,
10420 win,
10421 kvl.k_tok_bytes,
10422 kvl.v_tok_bytes,
10423 None,
10424 )?;
10425 return Ok(e.matmul(&fa.wo, &attn, 1)?);
10426 }
10427 let (off_tok, t_kv) = if swa && kvl.len > win {
10428 (kvl.len - win, win)
10429 } else {
10430 (0, kvl.len)
10431 };
10432 let k_view = e.view_u8_range(
10433 &kvl.k,
10434 off_tok * kvl.k_tok_bytes,
10435 (off_tok + t_kv) * kvl.k_tok_bytes,
10436 );
10437 let v_view = e.view_u8_range(
10438 &kvl.v,
10439 off_tok * kvl.v_tok_bytes,
10440 (off_tok + t_kv) * kvl.v_tok_bytes,
10441 );
10442 e.fa_decode_kvmod(
10443 &q,
10444 &k_view,
10445 &v_view,
10446 &mut attn,
10447 hd,
10448 nh,
10449 nkv,
10450 t_kv,
10451 scale,
10452 kvl.k_tok_bytes,
10453 kvl.v_tok_bytes,
10454 swa && crate::Engine::wkv_on(),
10455 )?;
10456 Ok(e.matmul(&fa.wo, &attn, 1)?)
10457 }
10458
10459 #[allow(clippy::too_many_arguments)]
10466 pub fn gemma4_decode_step_dc(
10467 &self,
10468 e: &Engine,
10469 token_d: &CudaSlice<u32>,
10470 pos_d: &mut CudaSlice<i32>,
10471 embd_gpu: &CudaSlice<u8>,
10472 embd_qt: i32,
10473 embd_rb: usize,
10474 cache: &mut Cache,
10475 n_vocab: usize,
10476 cap_bucket_max: Option<(usize, usize)>,
10477 ) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
10478 let mut tok_out = e.stream().alloc_zeros::<u32>(1)?;
10479 self.gemma4_decode_step_dc_into(
10480 e,
10481 token_d,
10482 pos_d,
10483 embd_gpu,
10484 embd_qt,
10485 embd_rb,
10486 cache,
10487 n_vocab,
10488 cap_bucket_max,
10489 &mut tok_out,
10490 )?;
10491 Ok(tok_out)
10492 }
10493
10494 #[allow(clippy::too_many_arguments)]
10497 pub fn gemma4_decode_step_dc_into(
10498 &self,
10499 e: &Engine,
10500 token_d: &CudaSlice<u32>,
10501 pos_d: &mut CudaSlice<i32>,
10502 embd_gpu: &CudaSlice<u8>,
10503 embd_qt: i32,
10504 embd_rb: usize,
10505 cache: &mut Cache,
10506 n_vocab: usize,
10507 cap_bucket_max: Option<(usize, usize)>,
10508 tok_out: &mut CudaSlice<u32>,
10509 ) -> Result<(), Box<dyn std::error::Error>> {
10510 let n_embd = self.cfg.n_embd as usize;
10511 let eps = self.cfg.rms_eps;
10512 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
10513 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
10514 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
10515 let n_layers = self.layers.len();
10516 for (il, layer) in self.layers.iter().enumerate() {
10517 let (hq, hdq) = match h_carry.take() {
10518 Some(p) => p,
10519 None => {
10520 e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, 1, eps)?
10521 }
10522 };
10523 let Mixer::Full(fa) = &layer.mixer else {
10524 panic!("gemma4 layer {il} not full-attn")
10525 };
10526 let o =
10527 self.gemma4_decode_attn_dc(e, fa, il, &hq, &hdq, pos_d, cache, cap_bucket_max)?;
10528 let next_norm = if il + 1 < n_layers {
10529 Some(self.layers[il + 1].attn_norm.float_data())
10530 } else {
10531 None
10532 };
10533 let (xn, hn) = self.gemma4_layer_tail_add_nq_pn(e, layer, &o, &x, 1, next_norm)?;
10534 x = xn;
10535 h_carry = hn;
10536 }
10537 let mut hn = e.uninit(n_embd)?;
10538 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
10539 let mut logits = e.matmul(&self.output, &hn, 1)?;
10540 self.gemma4_suppress(e, &mut logits, 1)?; e.argmax_token_device_into(&logits, tok_out, n_vocab)?;
10542 e.inc_seqlen(pos_d)?;
10543 if cap_bucket_max.is_none() {
10544 cache.pos += 1;
10545 }
10546 Ok(())
10547 }
10548
10549 pub fn g4_dc_slots(&self, e: &Engine) -> Result<G4DcSlots, Box<dyn std::error::Error>> {
10556 let n_embd = self.cfg.n_embd as usize;
10557 let n_vocab = self.output.out_features();
10558 let n_layers = self.layers.len();
10559 let (mut qmax, mut kvmax, mut ffmax) = (0usize, 0usize, 0usize);
10560 for il in 0..n_layers {
10561 let (hd, nkv, nh, _b, _s, _w) = self.gemma4_geom(il);
10562 qmax = qmax.max(nh * hd);
10563 kvmax = kvmax.max(nkv * hd);
10564 if let crate::hybrid::Ffn::Dense { ffn_gate, .. } = &self.layers[il].ffn {
10565 ffmax = ffmax.max(ffn_gate.out_features());
10566 }
10567 }
10568 Ok(G4DcSlots {
10569 x: e.uninit(n_embd)?,
10570 xn: e.uninit(n_embd)?,
10571 cur: e.uninit(n_embd)?,
10572 hq: e.alloc_i8_uninit(n_embd)?,
10573 hd_: e.uninit(n_embd / 32)?,
10574 q0: e.uninit(qmax)?,
10575 k0: e.uninit(kvmax)?,
10576 v0: e.uninit(kvmax)?,
10577 q: e.uninit(qmax)?,
10578 k: e.uninit(kvmax)?,
10579 v: e.uninit(kvmax)?,
10580 attn: e.uninit(qmax)?,
10581 o: e.uninit(n_embd)?,
10582 attn_out: e.uninit(n_embd)?,
10583 zsh: e.uninit(n_embd)?,
10584 zq: e.alloc_i8_uninit(n_embd.max(qmax))?,
10587 zd: e.uninit(n_embd.max(qmax) / 32)?,
10588 gate: e.uninit(ffmax)?,
10589 up: e.uninit(ffmax)?,
10590 act: e.uninit(ffmax)?,
10591 actq: e.alloc_i8_uninit(ffmax)?,
10592 actd: e.uninit(ffmax / 32)?,
10593 f0: e.uninit(n_embd)?,
10594 sn: e.uninit(n_embd)?,
10595 hn: e.uninit(n_embd)?,
10596 logits: e.uninit(n_vocab)?,
10597 })
10598 }
10599
10600 fn g4_matvec_m1_into(
10603 &self,
10604 e: &Engine,
10605 w: &crate::model::GpuTensor,
10606 aq: &CudaSlice<i8>,
10607 ad: &CudaSlice<f32>,
10608 y: &mut CudaSlice<f32>,
10609 ) -> Result<(), Box<dyn std::error::Error>> {
10610 use crate::model::GpuTensor;
10611 let (bytes, qtype, row_bytes, scale, rp) = match w {
10612 GpuTensor::Quant {
10613 bytes,
10614 qtype,
10615 row_bytes,
10616 scale,
10617 rp,
10618 ..
10619 } => (bytes, *qtype, *row_bytes, *scale, *rp),
10620 _ => return Err("g4_matvec_m1_into: non-quant tensor".into()),
10621 };
10622 let (mbytes, mrp) = match w {
10623 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
10624 _ => (bytes, rp),
10625 };
10626 e.qmatvec_mmvq_into(
10627 mbytes,
10628 aq,
10629 ad,
10630 1,
10631 w.in_features(),
10632 w.out_features(),
10633 qtype,
10634 row_bytes,
10635 scale,
10636 mrp,
10637 y,
10638 )
10639 }
10640
10641 #[allow(clippy::too_many_arguments)]
10645 pub fn gemma4_decode_step_dc_slotted(
10646 &self,
10647 e: &Engine,
10648 token_d: &CudaSlice<u32>,
10649 pos_d: &mut CudaSlice<i32>,
10650 embd_gpu: &CudaSlice<u8>,
10651 embd_qt: i32,
10652 embd_rb: usize,
10653 cache: &mut Cache,
10654 n_vocab: usize,
10655 cap_bucket_max: Option<(usize, usize)>,
10656 sl: &mut G4DcSlots,
10657 tok_out: &mut CudaSlice<u32>,
10658 ring: Option<(&mut CudaSlice<u32>, usize)>,
10659 ) -> Result<(), Box<dyn std::error::Error>> {
10660 let n_embd = self.cfg.n_embd as usize;
10661 let eps = self.cfg.rms_eps;
10662 e.embed_gather_device_into(embd_gpu, token_d, &mut sl.x, n_embd, embd_qt, embd_rb)?;
10663 e.scale_inplace(&mut sl.x, (n_embd as f32).sqrt(), n_embd)?;
10664 let n_layers = self.layers.len();
10665 let mut has_carry = false;
10666 for il in 0..n_layers {
10667 if !has_carry {
10668 e.rms_norm_q8_1_into(
10669 &sl.x,
10670 self.layers[il].attn_norm.float_data(),
10671 n_embd,
10672 1,
10673 eps,
10674 &mut sl.hq,
10675 &mut sl.hd_,
10676 )?;
10677 }
10678 has_carry = true;
10679 let layer = &self.layers[il];
10680 let Mixer::Full(fa) = &layer.mixer else {
10681 panic!("gemma4 layer {il} not full-attn")
10682 };
10683 self.gemma4_decode_attn_dc_slotted(e, fa, il, pos_d, cache, cap_bucket_max, sl)?;
10684 if !Engine::g4_pnfold_on() {
10687 e.rms_norm(
10688 &sl.o,
10689 layer.post_attn_norm.float_data(),
10690 &mut sl.cur,
10691 n_embd,
10692 1,
10693 eps,
10694 )?;
10695 }
10696 let next_norm = if il + 1 < n_layers {
10697 Some(self.layers[il + 1].attn_norm.float_data())
10698 } else {
10699 None
10700 };
10701 self.gemma4_layer_tail_slotted(e, layer, next_norm, sl)?;
10702 std::mem::swap(&mut sl.x, &mut sl.xn);
10703 }
10704 e.rms_norm(
10705 &sl.x,
10706 self.output_norm.float_data(),
10707 &mut sl.hn,
10708 n_embd,
10709 1,
10710 eps,
10711 )?;
10712 e.quantize_q8_1_into(&sl.hn, 1, n_embd, &mut sl.zq, &mut sl.zd)?;
10713 {
10715 let (zq, zd) = (&sl.zq, &sl.zd);
10716 let zq = unsafe { &*(zq as *const CudaSlice<i8>) };
10717 let zd = unsafe { &*(zd as *const CudaSlice<f32>) };
10718 self.g4_matvec_m1_into(e, &self.output, zq, zd, &mut sl.logits)?;
10719 }
10720 self.gemma4_suppress(e, &mut sl.logits, 1)?;
10721 e.argmax_token_device_into(&sl.logits, tok_out, n_vocab)?;
10722 if let Some((ring, base)) = ring {
10723 e.plain_tok_ring(tok_out, pos_d, base, ring)?;
10727 }
10728 e.inc_seqlen(pos_d)?;
10729 if cap_bucket_max.is_none() {
10730 cache.pos += 1;
10731 }
10732 Ok(())
10733 }
10734
10735 #[allow(clippy::too_many_arguments)]
10737 fn gemma4_decode_attn_dc_slotted(
10738 &self,
10739 e: &Engine,
10740 fa: &crate::hybrid::FullAttnLayer,
10741 il: usize,
10742 pos_d: &CudaSlice<i32>,
10743 cache: &mut Cache,
10744 cap_bucket_max: Option<(usize, usize)>,
10745 sl: &mut G4DcSlots,
10746 ) -> Result<(), Box<dyn std::error::Error>> {
10747 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
10748 let eps = self.cfg.rms_eps;
10749 let aux = self.gemma4_aux.as_ref().unwrap();
10750 let ones = aux.ones(e);
10751 #[cfg(debug_assertions)]
10752 crate::debug_assert_tensor_stream_device(
10753 ones,
10754 &e.stream(),
10755 "gemma4_decode_attn_dc_slotted.ones",
10756 );
10757 {
10758 let hq = unsafe { &*(&sl.hq as *const CudaSlice<i8>) };
10759 let hdq = unsafe { &*(&sl.hd_ as *const CudaSlice<f32>) };
10760 if swa {
10761 if !e.matmul_q4_fused3_into(
10762 &fa.wq, &fa.wk, &fa.wv, hq, hdq, &mut sl.q0, &mut sl.k0, &mut sl.v0,
10763 )? {
10764 if e.matmul_nvfp4_fused2_into(&fa.wq, &fa.wk, hq, hdq, &mut sl.q0, &mut sl.k0)?
10768 {
10769 self.g4_matvec_m1_into(e, &fa.wv, hq, hdq, &mut sl.v0)?;
10770 } else {
10771 return Err("slotted step: fused3 unavailable (non-uniform trunk)".into());
10772 }
10773 }
10774 } else {
10775 if !e.matmul_q4_fused2_into(&fa.wq, &fa.wk, hq, hdq, &mut sl.q0, &mut sl.k0)?
10776 && !e
10777 .matmul_nvfp4_fused2_into(&fa.wq, &fa.wk, hq, hdq, &mut sl.q0, &mut sl.k0)?
10778 {
10779 return Err("slotted step: fused2 unavailable".into());
10780 }
10781 let k0r = unsafe { &*(&sl.k0 as *const CudaSlice<f32>) };
10782 e.copy_into(&mut sl.v0, 0, k0r, nkv * hd)?;
10783 }
10784 }
10785 let ff = if swa {
10788 None
10789 } else {
10790 Some(
10791 aux.rope_freqs(e)
10792 .expect("gemma4 global rope needs rope_freqs.weight"),
10793 )
10794 };
10795 #[cfg(debug_assertions)]
10796 if let Some(ff) = ff {
10797 crate::debug_assert_tensor_stream_device(
10798 ff,
10799 &e.stream(),
10800 "gemma4_decode_attn_dc_slotted.rope_freqs",
10801 );
10802 }
10803 let kvl = cache.kv[il].as_mut().unwrap();
10804 let kv_fp8 = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
10805 if crate::Engine::qkv_append_on() {
10806 e.rms_norm_qkv_rope_append_dc(
10808 &sl.q0,
10809 &sl.k0,
10810 &sl.v0,
10811 fa.q_norm.float_data(),
10812 fa.k_norm.float_data(),
10813 ones,
10814 &mut sl.q,
10815 &mut sl.k,
10816 &mut sl.v,
10817 hd,
10818 nh,
10819 nkv,
10820 pos_d,
10821 nh,
10822 nkv,
10823 base,
10824 1.0,
10825 ff,
10826 eps,
10827 &mut kvl.k,
10828 &mut kvl.v,
10829 &kvl.len_d,
10830 kvl.k_tok_bytes,
10831 kvl.v_tok_bytes,
10832 kv_fp8,
10833 )?;
10834 } else {
10835 e.rms_norm_qkv_rope(
10836 &sl.q0,
10837 &sl.k0,
10838 &sl.v0,
10839 fa.q_norm.float_data(),
10840 fa.k_norm.float_data(),
10841 ones,
10842 &mut sl.q,
10843 &mut sl.k,
10844 &mut sl.v,
10845 hd,
10846 nh,
10847 nkv,
10848 pos_d,
10849 nh,
10850 nkv,
10851 base,
10852 1.0,
10853 ff,
10854 eps,
10855 )?;
10856 e.append_kv_quantized_dc(
10857 &sl.k,
10858 &sl.v,
10859 &mut kvl.k,
10860 &mut kvl.v,
10861 &kvl.len_d,
10862 kvl.kv_dim_k,
10863 kvl.kv_dim_v,
10864 kvl.k_tok_bytes,
10865 kvl.v_tok_bytes,
10866 kv_fp8,
10867 )?;
10868 }
10869 e.inc_seqlen(&mut kvl.len_d)?;
10870 let (b_swa, b_glob) = cap_bucket_max.expect("slotted step is capture-only");
10871 let k_view = e.view_u8(&kvl.k, kvl.k.len());
10872 let v_view = e.view_u8(&kvl.v, kvl.v.len());
10873 let rows_on = std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0");
10874 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
10875 let mut fa_q8 = false;
10879 if !swa && hd == 512 && b_glob >= crate::fa512_min_tkv() && rows_on {
10880 e.fa_decode_rows(
10881 &sl.q,
10882 &k_view,
10883 &v_view,
10884 &mut sl.attn,
10885 hd,
10886 nh,
10887 nkv,
10888 b_glob - 1,
10889 1,
10890 scale,
10891 kvl.k_tok_bytes,
10892 kvl.v_tok_bytes,
10893 Some((&kvl.len_d, -1)),
10894 false,
10895 false,
10896 Some((&mut sl.zq, &mut sl.zd)),
10897 )?;
10898 fa_q8 = true;
10899 } else if swa && b_swa > win && hd == 256 && rows_on {
10900 e.fa_decode_rows_w(
10901 &sl.q,
10902 &k_view,
10903 &v_view,
10904 &mut sl.attn,
10905 hd,
10906 nh,
10907 nkv,
10908 &kvl.len_d,
10909 -1,
10910 1,
10911 scale,
10912 win,
10913 kvl.k_tok_bytes,
10914 kvl.v_tok_bytes,
10915 Some((&mut sl.zq, &mut sl.zd)),
10916 )?;
10917 fa_q8 = true;
10918 } else {
10919 let b = if swa { b_swa } else { b_glob };
10920 e.fa_decode_dc(
10921 &sl.q,
10922 &k_view,
10923 &v_view,
10924 &mut sl.attn,
10925 hd,
10926 nh,
10927 nkv,
10928 &kvl.len_d,
10929 b,
10930 scale,
10931 kvl.k_tok_bytes,
10932 kvl.v_tok_bytes,
10933 swa && crate::Engine::wkv_on(),
10934 )?;
10935 }
10936 if !fa_q8 {
10937 let aq = unsafe { &*(&sl.attn as *const CudaSlice<f32>) };
10938 e.quantize_q8_1_into(aq, 1, nh * hd, &mut sl.zq, &mut sl.zd)?;
10939 }
10940 {
10941 let zq = unsafe { &*(&sl.zq as *const CudaSlice<i8>) };
10942 let zd = unsafe { &*(&sl.zd as *const CudaSlice<f32>) };
10943 self.g4_matvec_m1_into(e, &fa.wo, zq, zd, &mut sl.o)?;
10944 }
10945 Ok(())
10946 }
10947
10948 fn gemma4_layer_tail_slotted(
10951 &self,
10952 e: &Engine,
10953 layer: &crate::hybrid::HybridLayer,
10954 next_norm: Option<&CudaSlice<f32>>,
10955 sl: &mut G4DcSlots,
10956 ) -> Result<(), Box<dyn std::error::Error>> {
10957 let n_embd = self.cfg.n_embd as usize;
10958 let eps = self.cfg.rms_eps;
10959 let bits = layer.gemma4.as_ref().unwrap();
10960 let crate::hybrid::Ffn::Dense {
10961 ffn_gate,
10962 ffn_up,
10963 ffn_down,
10964 } = &layer.ffn
10965 else {
10966 return Err("slotted tail: dense ffn only".into());
10967 };
10968 let pnfold = Engine::g4_pnfold_on();
10969 if pnfold {
10970 let or = unsafe { &*(&sl.o as *const CudaSlice<f32>) };
10973 let xr = unsafe { &*(&sl.x as *const CudaSlice<f32>) };
10974 e.rms_pre_add_rms_norm_q8z_into(
10975 or,
10976 layer.post_attn_norm.float_data(),
10977 xr,
10978 bits.ffn_norm.float_data(),
10979 &mut sl.attn_out,
10980 &mut sl.zsh,
10981 n_embd,
10982 1,
10983 eps,
10984 &mut sl.zq,
10985 &mut sl.zd,
10986 )?;
10987 } else {
10988 e.add_rms_norm(
10989 &sl.cur,
10990 &sl.x,
10991 bits.ffn_norm.float_data(),
10992 &mut sl.attn_out,
10993 &mut sl.zsh,
10994 n_embd,
10995 1,
10996 eps,
10997 )?;
10998 }
10999 let n_ff = ffn_gate.out_features();
11000 if !pnfold {
11001 let zshr = unsafe { &*(&sl.zsh as *const CudaSlice<f32>) };
11002 e.quantize_q8_1_into(zshr, 1, n_embd, &mut sl.zq, &mut sl.zd)?;
11003 }
11004 {
11005 let zq = unsafe { &*(&sl.zq as *const CudaSlice<i8>) };
11006 let zd = unsafe { &*(&sl.zd as *const CudaSlice<f32>) };
11007 if !e.matmul_q4_fused2_into(ffn_gate, ffn_up, zq, zd, &mut sl.gate, &mut sl.up)?
11008 && !e.matmul_nvfp4_fused2_into(
11009 ffn_gate,
11010 ffn_up,
11011 zq,
11012 zd,
11013 &mut sl.gate,
11014 &mut sl.up,
11015 )?
11016 {
11017 return Err("slotted tail: ffn fused2 unavailable".into());
11018 }
11019 }
11020 debug_assert!(e.uses_q8_1_fast(ffn_down));
11021 {
11022 let upr = unsafe { &*(&sl.up as *const CudaSlice<f32>) };
11023 let upv = e.view(upr, n_ff);
11024 let up_all = upv.slice(0..n_ff);
11025 let gr = unsafe { &*(&sl.gate as *const CudaSlice<f32>) };
11026 e.gelu_tanh_mul_q8_1_into(
11027 gr,
11028 &up_all,
11029 &mut sl.act,
11030 n_ff,
11031 1,
11032 &mut sl.actq,
11033 &mut sl.actd,
11034 )?;
11035 }
11036 {
11037 let aq = unsafe { &*(&sl.actq as *const CudaSlice<i8>) };
11038 let ad = unsafe { &*(&sl.actd as *const CudaSlice<f32>) };
11039 self.g4_matvec_m1_into(e, ffn_down, aq, ad, &mut sl.f0)?;
11040 }
11041 if pnfold {
11042 if let Some(w) = next_norm {
11045 let f0r = unsafe { &*(&sl.f0 as *const CudaSlice<f32>) };
11046 let aor = unsafe { &*(&sl.attn_out as *const CudaSlice<f32>) };
11047 e.rms_pre_add_scale_rms_norm_q8_1_into(
11048 f0r,
11049 bits.post_ffw_norm.float_data(),
11050 aor,
11051 bits.layer_scale,
11052 w,
11053 &mut sl.xn,
11054 n_embd,
11055 1,
11056 eps,
11057 &mut sl.hq,
11058 &mut sl.hd_,
11059 )?;
11060 return Ok(());
11061 }
11062 }
11063 e.rms_norm(
11064 &sl.f0,
11065 bits.post_ffw_norm.float_data(),
11066 &mut sl.sn,
11067 n_embd,
11068 1,
11069 eps,
11070 )?;
11071 match next_norm {
11072 Some(w) => {
11073 e.add_scale_rms_norm_q8_1_into(
11074 &sl.sn,
11075 &sl.attn_out,
11076 bits.layer_scale,
11077 w,
11078 &mut sl.xn,
11079 n_embd,
11080 1,
11081 eps,
11082 &mut sl.hq,
11083 &mut sl.hd_,
11084 )?;
11085 }
11086 None => {
11087 e.add_scale(&sl.sn, &sl.attn_out, bits.layer_scale, &mut sl.xn, n_embd)?;
11088 }
11089 }
11090 Ok(())
11091 }
11092
11093 #[allow(clippy::too_many_arguments)]
11095 fn gemma4_decode_attn_dc(
11096 &self,
11097 e: &Engine,
11098 fa: &crate::hybrid::FullAttnLayer,
11099 il: usize,
11100 hq: &CudaSlice<i8>,
11101 hdq: &CudaSlice<f32>,
11102 pos_d: &CudaSlice<i32>,
11103 cache: &mut Cache,
11104 cap_bucket_max: Option<(usize, usize)>,
11105 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
11106 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
11107 let eps = self.cfg.rms_eps;
11108 let aux = self.gemma4_aux.as_ref().unwrap();
11109 let ones = aux.ones(e);
11110 #[cfg(debug_assertions)]
11111 crate::debug_assert_tensor_stream_device(ones, &e.stream(), "gemma4_decode_attn_dc.ones");
11112 let (q0, k0, v0) = if swa {
11113 match e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)? {
11114 Some(t3) => t3,
11115 None => match e.matmul_nvfp4_fused2(&fa.wq, &fa.wk, hq, hdq, 1)? {
11117 Some((q0, k0)) => {
11118 let h0 = e.zeros(0)?;
11119 let v0 = e.matmul_pre(&fa.wv, hq, hdq, &h0, 1)?;
11120 (q0, k0, v0)
11121 }
11122 None => {
11123 let h0 = e.zeros(0)?;
11124 (
11125 e.matmul_pre(&fa.wq, hq, hdq, &h0, 1)?,
11126 e.matmul_pre(&fa.wk, hq, hdq, &h0, 1)?,
11127 e.matmul_pre(&fa.wv, hq, hdq, &h0, 1)?,
11128 )
11129 }
11130 },
11131 }
11132 } else {
11133 let (q0, k0) = match e.matmul_q4_fused2(&fa.wq, &fa.wk, hq, hdq)? {
11134 Some(p) => p,
11135 None => match e.matmul_nvfp4_fused2(&fa.wq, &fa.wk, hq, hdq, 1)? {
11136 Some(p) => p,
11137 None => {
11138 let h0 = e.zeros(0)?;
11139 (
11140 e.matmul_pre(&fa.wq, hq, hdq, &h0, 1)?,
11141 e.matmul_pre(&fa.wk, hq, hdq, &h0, 1)?,
11142 )
11143 }
11144 },
11145 };
11146 let v0 = e.clone_dtod(&k0)?;
11147 (q0, k0, v0)
11148 };
11149 let mut q = e.uninit(nh * hd)?;
11150 let mut k = e.uninit(nkv * hd)?;
11151 let mut v = e.uninit(nkv * hd)?;
11152 let ff = if swa {
11154 None
11155 } else {
11156 Some(
11157 aux.rope_freqs(e)
11158 .expect("gemma4 global rope needs rope_freqs.weight"),
11159 )
11160 };
11161 #[cfg(debug_assertions)]
11162 if let Some(ff) = ff {
11163 crate::debug_assert_tensor_stream_device(
11164 ff,
11165 &e.stream(),
11166 "gemma4_decode_attn_dc.rope_freqs",
11167 );
11168 }
11169 let kvl = cache.kv[il].as_mut().unwrap();
11170 let kv_fp8 = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
11171 if crate::Engine::qkv_append_on() {
11172 e.rms_norm_qkv_rope_append_dc(
11174 &q0,
11175 &k0,
11176 &v0,
11177 fa.q_norm.float_data(),
11178 fa.k_norm.float_data(),
11179 ones,
11180 &mut q,
11181 &mut k,
11182 &mut v,
11183 hd,
11184 nh,
11185 nkv,
11186 pos_d,
11187 nh,
11188 nkv,
11189 base,
11190 1.0,
11191 ff,
11192 eps,
11193 &mut kvl.k,
11194 &mut kvl.v,
11195 &kvl.len_d,
11196 kvl.k_tok_bytes,
11197 kvl.v_tok_bytes,
11198 kv_fp8,
11199 )?;
11200 } else {
11201 e.rms_norm_qkv_rope(
11202 &q0,
11203 &k0,
11204 &v0,
11205 fa.q_norm.float_data(),
11206 fa.k_norm.float_data(),
11207 ones,
11208 &mut q,
11209 &mut k,
11210 &mut v,
11211 hd,
11212 nh,
11213 nkv,
11214 pos_d,
11215 nh,
11216 nkv,
11217 base,
11218 1.0,
11219 ff,
11220 eps,
11221 )?;
11222 e.append_kv_quantized_dc(
11223 &k,
11224 &v,
11225 &mut kvl.k,
11226 &mut kvl.v,
11227 &kvl.len_d,
11228 kvl.kv_dim_k,
11229 kvl.kv_dim_v,
11230 kvl.k_tok_bytes,
11231 kvl.v_tok_bytes,
11232 kv_fp8,
11233 )?;
11234 }
11235 e.inc_seqlen(&mut kvl.len_d)?;
11236 let mut attn = e.uninit(nh * hd)?;
11237 let mut fa_q8: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
11240 match cap_bucket_max {
11245 None => {
11246 kvl.len += 1;
11250 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
11251 if !swa
11252 && hd == 512
11253 && kvl.len >= crate::fa512_min_tkv()
11254 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
11255 {
11256 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
11259 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
11260 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
11261 e.fa_decode_rows(
11262 &q,
11263 &kp,
11264 &vp,
11265 &mut attn,
11266 hd,
11267 nh,
11268 nkv,
11269 kvl.len - 1,
11270 1,
11271 scale,
11272 kvl.k_tok_bytes,
11273 kvl.v_tok_bytes,
11274 Some((&kvl.len_d, -1)),
11275 false,
11276 false,
11277 Some((&mut aq8, &mut ad8)),
11278 )?;
11279 fa_q8 = Some((aq8, ad8));
11280 } else if swa
11281 && kvl.len > win
11282 && hd == 256
11283 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
11284 {
11285 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
11287 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
11288 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
11289 e.fa_decode_rows_w(
11290 &q,
11291 &kp,
11292 &vp,
11293 &mut attn,
11294 hd,
11295 nh,
11296 nkv,
11297 &kvl.len_d,
11298 -1,
11299 1,
11300 scale,
11301 win,
11302 kvl.k_tok_bytes,
11303 kvl.v_tok_bytes,
11304 Some((&mut aq8, &mut ad8)),
11305 )?;
11306 fa_q8 = Some((aq8, ad8));
11307 } else {
11308 let (off_tok, t_kv) = if swa && kvl.len > win {
11309 (kvl.len - win, win)
11310 } else {
11311 (0, kvl.len)
11312 };
11313 let k_view = e.view_u8_range(
11314 &kvl.k,
11315 off_tok * kvl.k_tok_bytes,
11316 (off_tok + t_kv) * kvl.k_tok_bytes,
11317 );
11318 let v_view = e.view_u8_range(
11319 &kvl.v,
11320 off_tok * kvl.v_tok_bytes,
11321 (off_tok + t_kv) * kvl.v_tok_bytes,
11322 );
11323 e.fa_decode_kvmod(
11324 &q,
11325 &k_view,
11326 &v_view,
11327 &mut attn,
11328 hd,
11329 nh,
11330 nkv,
11331 t_kv,
11332 scale,
11333 kvl.k_tok_bytes,
11334 kvl.v_tok_bytes,
11335 swa && crate::Engine::wkv_on(),
11336 )?;
11337 }
11338 }
11339 Some((b_swa, b_glob)) => {
11340 let k_view = e.view_u8(&kvl.k, kvl.k.len());
11346 let v_view = e.view_u8(&kvl.v, kvl.v.len());
11347 let rows_on = std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0");
11348 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
11349 if !swa && hd == 512 && b_glob >= crate::fa512_min_tkv() && rows_on {
11350 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
11351 e.fa_decode_rows(
11352 &q,
11353 &k_view,
11354 &v_view,
11355 &mut attn,
11356 hd,
11357 nh,
11358 nkv,
11359 b_glob - 1,
11360 1,
11361 scale,
11362 kvl.k_tok_bytes,
11363 kvl.v_tok_bytes,
11364 Some((&kvl.len_d, -1)),
11365 false,
11366 false,
11367 Some((&mut aq8, &mut ad8)),
11368 )?;
11369 fa_q8 = Some((aq8, ad8));
11370 } else if swa && b_swa > win && hd == 256 && rows_on {
11371 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
11372 e.fa_decode_rows_w(
11373 &q,
11374 &k_view,
11375 &v_view,
11376 &mut attn,
11377 hd,
11378 nh,
11379 nkv,
11380 &kvl.len_d,
11381 -1,
11382 1,
11383 scale,
11384 win,
11385 kvl.k_tok_bytes,
11386 kvl.v_tok_bytes,
11387 Some((&mut aq8, &mut ad8)),
11388 )?;
11389 fa_q8 = Some((aq8, ad8));
11390 } else {
11391 let b = if swa { b_swa } else { b_glob };
11392 e.fa_decode_dc(
11393 &q,
11394 &k_view,
11395 &v_view,
11396 &mut attn,
11397 hd,
11398 nh,
11399 nkv,
11400 &kvl.len_d,
11401 b,
11402 scale,
11403 kvl.k_tok_bytes,
11404 kvl.v_tok_bytes,
11405 swa && crate::Engine::wkv_on(),
11406 )?;
11407 }
11408 }
11409 }
11410 if let Some((aq8, ad8)) = fa_q8 {
11413 let mut y = e.uninit(fa.wo.out_features())?;
11414 self.g4_matvec_m1_into(e, &fa.wo, &aq8, &ad8, &mut y)?;
11415 return Ok(y);
11416 }
11417 Ok(e.matmul(&fa.wo, &attn, 1)?)
11418 }
11419
11420 pub fn gemma4_generate_graph(
11425 &self,
11426 e: &Engine,
11427 prompt_pos: usize,
11428 first_token: u32,
11429 cache: &mut Cache,
11430 max_new: usize,
11431 eos: &[u32],
11432 mut on_token: impl FnMut(u32) -> bool,
11433 ) -> Result<(Vec<u32>, crate::decode::StopReason), Box<dyn std::error::Error>> {
11434 if self.is_gemma4_e4b() {
11435 return Err(
11436 "E4B graph serving is unwired (HANDOVER-E4B.md) — dc-eager is the serving arm"
11437 .into(),
11438 );
11439 }
11440 use crate::decode::StopReason;
11441 let n_vocab = self.output.out_features();
11442 let n_embd = self.cfg.n_embd as usize;
11443 let embd_gpu = self
11444 .embd_gpu
11445 .get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload"));
11446 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
11447 for kvl in cache.kv.iter_mut().flatten() {
11448 e.set_i32_one(&mut kvl.len_d, kvl.len as i32)?;
11449 }
11450 let mut token_d = e.stream().clone_htod(&[first_token])?;
11451 let mut pos_d = e.htod_i32(&[prompt_pos as i32])?;
11452 let g4 = self.cfg.gemma4.as_ref().unwrap();
11453 let (hd_s, hd_g) = (g4.key_length_swa as usize, g4.key_length_global as usize);
11454 let nkv_s = g4
11456 .head_count_kv
11457 .iter()
11458 .zip(g4.swa_pattern.iter())
11459 .find(|p| *p.1)
11460 .map(|p| *p.0 as usize)
11461 .unwrap_or(8);
11462 let nkv_g = g4
11463 .head_count_kv
11464 .iter()
11465 .zip(g4.swa_pattern.iter())
11466 .find(|p| !*p.1)
11467 .map(|p| *p.0 as usize)
11468 .unwrap_or(2);
11469 let mut graphs: std::collections::HashMap<
11470 ((bool, usize), (bool, usize), bool, bool),
11471 (
11472 cudarc::driver::CudaGraph,
11473 Vec<Box<dyn std::any::Any + Send>>,
11474 ),
11475 > = Default::default();
11476 let mut slots = self.g4_dc_slots(e)?;
11479 const RING: usize = 64;
11482 const DRAIN: usize = 1;
11488 let mut ring = e.stream().alloc_zeros::<u32>(RING)?;
11489 let ring_base = prompt_pos;
11490 let mut out = Vec::with_capacity(max_new);
11491 let mut reason = StopReason::MaxNew;
11492 let mut next = first_token;
11493 let mut captures = 0usize;
11494 for _ in 0..max_new {
11495 out.push(next);
11496 if eos.contains(&next) {
11497 reason = StopReason::Eos;
11498 break;
11499 }
11500 if !on_token(next) {
11501 reason = StopReason::Callback;
11502 break;
11503 }
11504 let t_kv = cache.pos + 1;
11505 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
11513 let f512 = crate::fa512_min_tkv();
11514 let key_s = if t_kv > win {
11515 (true, usize::MAX)
11516 } else {
11517 e.fa_bucket_key(t_kv, hd_s, nkv_s, crate::Engine::wkv_on())
11518 };
11519 let (key_g, rung_end) = if t_kv >= f512 {
11520 let end = (t_kv + 1).next_power_of_two().max(f512 * 2);
11523 ((true, end), end)
11524 } else {
11525 (e.fa_bucket_key(t_kv, hd_g, nkv_g, false), t_kv)
11526 };
11527 let key = (key_s, key_g, t_kv >= f512, t_kv > win);
11528 if !graphs.contains_key(&key) {
11529 let bucket_max = (t_kv, rung_end);
11530 let snap = cache.snapshot(e)?;
11532 let pos_save = e.dtoh_i32_one(&pos_d)?;
11533 let len_save: Vec<Option<i32>> = cache
11534 .kv
11535 .iter()
11536 .map(|k| k.as_ref().map(|kvl| e.dtoh_i32_one(&kvl.len_d).unwrap()))
11537 .collect();
11538 let tok_save = e.dtoh_u32_one(&token_d)?;
11539 let graph = {
11544 let tok_ref = &mut token_d;
11545 let pos_ref = &mut pos_d;
11546 let cache_ref = &mut *cache;
11547 let slots_ref = &mut slots;
11548 let ring_ref = &mut ring;
11549 e.capture_graph_retained_flags(
11550 cudarc::driver::sys::CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_USE_NODE_PRIORITY,
11551 |e| {
11552 let tok_in = unsafe { &*(tok_ref as *const CudaSlice<u32>) };
11554 let sl = unsafe { &mut *(slots_ref as *mut G4DcSlots) };
11555 let rg = unsafe { &mut *(ring_ref as *mut CudaSlice<u32>) };
11556 self.gemma4_decode_step_dc_slotted(e, tok_in, pos_ref, embd_gpu, qt, rb,
11557 cache_ref, n_vocab, Some(bucket_max),
11558 sl, tok_ref, Some((rg, ring_base)))
11559 })?
11560 };
11561 cache.rollback(e, &snap, 0)?;
11562 e.set_i32_one(&mut pos_d, pos_save)?;
11563 for (il, ls) in len_save.iter().enumerate() {
11564 if let (Some(kvl), Some(v)) = (cache.kv[il].as_mut(), ls) {
11565 e.set_i32_one(&mut kvl.len_d, *v)?;
11566 }
11567 }
11568 e.set_u32_one(&mut token_d, tok_save)?;
11569 if std::env::var("MEMRA_GRAPH_CENSUS").as_deref() == Ok("1") {
11570 if let Ok(c) = crate::graph_update::node_census(&graph.0) {
11571 eprintln!("[graph-census] {c:?}");
11572 }
11573 }
11574 graphs.insert(key, graph);
11575 captures += 1;
11576 }
11577 let mut chunk = 1usize;
11582 let drain_cap: usize = std::env::var("MEMRA_GRAPH_DRAIN")
11583 .ok()
11584 .and_then(|v| v.parse().ok())
11585 .unwrap_or(DRAIN);
11586 while chunk < drain_cap && out.len() + chunk < max_new {
11587 let t_next = cache.pos + 1 + chunk;
11588 let key_s2 = if t_next > win {
11589 (true, usize::MAX)
11590 } else {
11591 e.fa_bucket_key(t_next, hd_s, nkv_s, crate::Engine::wkv_on())
11592 };
11593 let key_g2 = if t_next >= f512 {
11594 (true, (t_next + 1).next_power_of_two().max(f512 * 2))
11595 } else {
11596 e.fa_bucket_key(t_next, hd_g, nkv_g, false)
11597 };
11598 if (key_s2, key_g2, t_next >= f512, t_next > win) != key {
11599 break;
11600 }
11601 chunk += 1;
11602 }
11603 let g = &graphs.get(&key).unwrap().0;
11604 for _ in 0..chunk {
11605 g.launch()?;
11606 }
11607 e.stream().synchronize()?;
11608 let ringh = e.dtoh_u32(&ring)?;
11609 for j in 0..chunk {
11610 let pos_j = cache.pos + j;
11611 let tok_j = ringh[(pos_j - ring_base) % RING];
11612 cache.pos += 0; if j + 1 == chunk {
11614 next = tok_j;
11615 } else {
11616 out.push(tok_j);
11617 if eos.contains(&tok_j) || !on_token(tok_j) {
11618 reason = if eos.contains(&tok_j) {
11619 StopReason::Eos
11620 } else {
11621 StopReason::Callback
11622 };
11623 let keep = cache.pos + j + 1;
11625 e.set_i32_one(&mut pos_d, keep as i32)?;
11626 for kvl in cache.kv.iter_mut().filter_map(|k| k.as_mut()) {
11627 e.set_i32_one(&mut kvl.len_d, keep as i32)?;
11628 kvl.len = keep;
11629 }
11630 cache.pos = keep;
11631 if std::env::var("MEMRA_GRAPH_STATS").is_ok() {
11632 eprintln!("[gemma-graph] captures={captures} buckets={}", graphs.len());
11633 }
11634 return Ok((out, reason));
11635 }
11636 }
11637 }
11638 cache.pos += chunk;
11639 for kvl in cache.kv.iter_mut().filter_map(|k| k.as_mut()) {
11640 kvl.len += chunk;
11641 }
11642 }
11643 if std::env::var("MEMRA_GRAPH_STATS").is_ok() {
11644 eprintln!("[gemma-graph] captures={captures} buckets={}", graphs.len());
11645 }
11646 Ok((out, reason))
11647 }
11648
11649 pub(crate) fn gemma4_decode_step_t(
11655 &self,
11656 e: &Engine,
11657 tokens: &[u32],
11658 pos0: usize,
11659 cache: &mut Cache,
11660 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
11661 Ok(self.gemma4_decode_step_t_h(e, tokens, pos0, cache)?.0)
11662 }
11663
11664 pub(crate) fn gemma4_decode_step_t_am(
11668 &self,
11669 e: &Engine,
11670 tokens: &[u32],
11671 pos0: usize,
11672 cache: &mut Cache,
11673 ) -> Result<(Vec<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
11674 let (ld, hn) = self.gemma4_verify_trunk(e, tokens, pos0, cache, None)?;
11675 let t = tokens.len();
11676 let n_vocab = self.output.out_features();
11677 let mut toks = e.stream().alloc_zeros::<u32>(t)?;
11678 for i in 0..t {
11679 e.argmax_token_device_col(&ld, i, n_vocab, &mut toks, i)?;
11680 }
11681 Ok((e.dtoh_u32(&toks)?, hn))
11682 }
11683
11684 pub(crate) fn gemma4_decode_step_t_am_dev(
11687 &self,
11688 e: &Engine,
11689 tok_d: &CudaSlice<u32>,
11690 t: usize,
11691 pos0: usize,
11692 cache: &mut Cache,
11693 ) -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
11694 let (ld, hn) = self.gemma4_verify_trunk(e, &vec![0u32; t], pos0, cache, Some(tok_d))?;
11695 let n_vocab = self.output.out_features();
11696 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
11697 for i in 0..t {
11698 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
11699 }
11700 Ok((vam, hn))
11701 }
11702
11703 pub(crate) fn gemma4_decode_step_t_h(
11706 &self,
11707 e: &Engine,
11708 tokens: &[u32],
11709 pos0: usize,
11710 cache: &mut Cache,
11711 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
11712 let (mut ld, hn) = self.gemma4_verify_trunk(e, tokens, pos0, cache, None)?;
11713 let t = tokens.len();
11714 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
11715 e.softcap(&mut ld, cap, t * self.output.out_features())?;
11716 Ok((e.dtoh(&ld)?, hn))
11717 }
11718
11719 pub(crate) fn verify_stream_scratch(
11722 &self,
11723 e: &Engine,
11724 cap: usize,
11725 ) -> Result<VerifyStreamScratch, Box<dyn std::error::Error>> {
11726 Ok(VerifyStreamScratch {
11727 pos_d: e.htod_i32(&vec![0i32; cap])?,
11728 row_ctrs: (0..cap)
11729 .map(|_| e.htod_i32(&[0]))
11730 .collect::<Result<_, _>>()?,
11731 })
11732 }
11733
11734 pub(crate) fn gemma4_verify_t_am_stream(
11742 &self,
11743 e: &Engine,
11744 tok_d: &CudaSlice<u32>,
11745 t: usize,
11746 ctr: &CudaSlice<i32>,
11747 hint: usize,
11748 cache: &mut Cache,
11749 scr: &mut VerifyStreamScratch,
11750 ) -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
11751 let n_embd = self.cfg.n_embd as usize;
11752 let eps = self.cfg.rms_eps;
11753 assert!(t <= scr.row_ctrs.len() && t <= 64);
11754 e.i32_iota_from(ctr, &mut scr.pos_d, t)?;
11755 for i in 0..t {
11756 e.i32_copy_add(ctr, &mut scr.row_ctrs[i], (i + 1) as i32)?;
11757 }
11758 let (pos_d, row_ctrs) = (&scr.pos_d, &scr.row_ctrs);
11759 let embd_gpu = self
11760 .embd_gpu
11761 .get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload"));
11762 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
11763 let mut x = e.embed_gather_device_td(embd_gpu, tok_d, t, n_embd, qt, rb)?;
11764 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
11765 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
11766 let n_layers = self.layers.len();
11767 for (il, layer) in self.layers.iter().enumerate() {
11768 let (hq, hdq) = match h_carry.take() {
11769 Some(p) => p,
11770 None => {
11771 e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, t, eps)?
11772 }
11773 };
11774 let Mixer::Full(fa) = &layer.mixer else {
11775 panic!("gemma4 layer {il} not full-attn")
11776 };
11777 let o = self
11778 .gemma4_verify_attn_stream(e, fa, il, &hq, &hdq, pos_d, t, cache, hint, row_ctrs)?;
11779 let next_norm = if il + 1 < n_layers {
11780 Some(self.layers[il + 1].attn_norm.float_data())
11781 } else {
11782 None
11783 };
11784 let (xn, hn) = self.gemma4_layer_tail_add_nq_pn(e, layer, &o, &x, t, next_norm)?;
11785 x = xn;
11786 h_carry = hn;
11787 self.dflash_tap(e, cache, il, &x, t)?;
11788 }
11789 let mut hn = e.uninit(t * n_embd)?;
11790 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
11791 let ld = e.matmul(&self.output, &hn, t)?;
11792 let n_vocab = self.output.out_features();
11793 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
11794 for i in 0..t {
11795 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
11796 }
11797 Ok((vam, hn))
11798 }
11799
11800 pub(crate) fn dflash_tap(
11807 &self,
11808 e: &Engine,
11809 cache: &mut Cache,
11810 il: usize,
11811 x: &CudaSlice<f32>,
11812 t: usize,
11813 ) -> Result<(), Box<dyn std::error::Error>> {
11814 let Some(taps) = cache.dflash_taps.as_mut() else {
11815 return Ok(());
11816 };
11817 let Some(slot) = taps.layer_ids.iter().position(|&l| l == il) else {
11818 return Ok(());
11819 };
11820 let h = taps.hidden;
11821 let n_taps = taps.layer_ids.len();
11822 let base = taps.base;
11823 debug_assert!(
11824 base + t <= taps.t,
11825 "tap window {base}+{t} exceeds sink {}",
11826 taps.t
11827 );
11828 let xv = e.view(x, t * h);
11829 for r in 0..t {
11830 let row = xv.slice(r * h..(r + 1) * h);
11831 e.copy_view_into(&mut taps.buf, (base + r) * n_taps * h + slot * h, &row, h)?;
11832 }
11833 Ok(())
11834 }
11835
11836 fn gemma4_verify_trunk(
11837 &self,
11838 e: &Engine,
11839 tokens: &[u32],
11840 pos0: usize,
11841 cache: &mut Cache,
11842 tok_dev: Option<&CudaSlice<u32>>,
11843 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
11844 let n_embd = self.cfg.n_embd as usize;
11845 let eps = self.cfg.rms_eps;
11846 let t = tokens.len();
11847 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
11848 let pos_d = e.htod_i32(&pos)?;
11849 let mut x = match tok_dev {
11850 Some(td) => {
11851 let embd_gpu = self
11852 .embd_gpu
11853 .get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload"));
11854 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
11855 e.embed_gather_device_td(embd_gpu, td, t, n_embd, qt, rb)?
11856 }
11857 None => e.htod(&self.embd.gather(n_embd, tokens))?,
11858 };
11859 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
11860 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
11861 let n_layers = self.layers.len();
11862 for (il, layer) in self.layers.iter().enumerate() {
11863 let (hq, hdq) = match h_carry.take() {
11864 Some(p) => p,
11865 None => {
11866 e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, t, eps)?
11867 }
11868 };
11869 let Mixer::Full(fa) = &layer.mixer else {
11870 panic!("gemma4 layer {il} not full-attn")
11871 };
11872 let o = self.gemma4_verify_attn(e, fa, il, &hq, &hdq, &pos_d, t, cache)?;
11873 let next_norm = if il + 1 < n_layers {
11874 Some(self.layers[il + 1].attn_norm.float_data())
11875 } else {
11876 None
11877 };
11878 let (xn, hn) = self.gemma4_layer_tail_add_nq_pn(e, layer, &o, &x, t, next_norm)?;
11879 x = xn;
11880 h_carry = hn;
11881 self.dflash_tap(e, cache, il, &x, t)?;
11882 }
11883 let mut hn = e.uninit(t * n_embd)?;
11884 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
11885 let mut ld = e.matmul(&self.output, &hn, t)?;
11886 self.gemma4_suppress(e, &mut ld, t)?; cache.pos += t;
11888 Ok((ld, hn))
11889 }
11890
11891 #[allow(clippy::too_many_arguments)]
11899 fn gemma4_verify_attn_stream(
11900 &self,
11901 e: &Engine,
11902 fa: &crate::hybrid::FullAttnLayer,
11903 il: usize,
11904 hq: &CudaSlice<i8>,
11905 hdq: &CudaSlice<f32>,
11906 pos_d: &CudaSlice<i32>,
11907 t: usize,
11908 cache: &mut Cache,
11909 hint: usize,
11910 row_ctrs: &[CudaSlice<i32>],
11911 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
11912 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
11913 let eps = self.cfg.rms_eps;
11914 let aux = self.gemma4_aux.as_ref().unwrap();
11915 let ones = aux.ones(e);
11916 #[cfg(debug_assertions)]
11917 crate::debug_assert_tensor_stream_device(
11918 ones,
11919 &e.stream(),
11920 "gemma4_verify_attn_stream.ones",
11921 );
11922 let h0 = e.zeros(0)?;
11923 let h = &h0;
11924 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
11927 let f2b = *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
11928 let fused_qkv = if f2b {
11929 if swa {
11930 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
11931 .map(|(a, b, c)| (a, b, Some(c)))
11932 } else {
11933 e.matmul_q4_fused2_batched(&fa.wq, &fa.wk, hq, hdq, t)?
11934 .map(|(a, b)| (a, b, None))
11935 }
11936 } else {
11937 None
11938 };
11939 let (q0, k0, v0) = match fused_qkv {
11940 Some((a, b, cv)) => {
11941 let v = match cv {
11942 Some(c) => c,
11943 None => e.clone_dtod(&b)?,
11944 };
11945 (a, b, v)
11946 }
11947 None => {
11948 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
11949 let k0 = e.matmul_pre(&fa.wk, hq, hdq, h, t)?;
11950 let v0 = if swa {
11951 e.matmul_pre(&fa.wv, hq, hdq, h, t)?
11952 } else {
11953 e.clone_dtod(&k0)?
11954 };
11955 (q0, k0, v0)
11956 }
11957 };
11958 let mut q = e.uninit(t * nh * hd)?;
11959 let mut k = e.uninit(t * nkv * hd)?;
11960 let mut v = e.uninit(t * nkv * hd)?;
11961 let ff = if swa {
11964 None
11965 } else {
11966 Some(
11967 aux.rope_freqs(e)
11968 .expect("gemma4 global rope needs rope_freqs.weight"),
11969 )
11970 };
11971 #[cfg(debug_assertions)]
11972 if let Some(ff) = ff {
11973 crate::debug_assert_tensor_stream_device(
11974 ff,
11975 &e.stream(),
11976 "gemma4_verify_attn_stream.rope_freqs",
11977 );
11978 }
11979 e.rms_norm_qkv_rope(
11980 &q0,
11981 &k0,
11982 &v0,
11983 fa.q_norm.float_data(),
11984 fa.k_norm.float_data(),
11985 ones,
11986 &mut q,
11987 &mut k,
11988 &mut v,
11989 hd,
11990 nh * t,
11991 nkv * t,
11992 pos_d,
11993 nh,
11994 nkv,
11995 base,
11996 1.0,
11997 ff,
11998 eps,
11999 )?;
12000 let kvl = cache.kv[il].as_mut().unwrap();
12001 e.append_kv_quantized_rows_dc(
12003 &k,
12004 &v,
12005 &mut kvl.k,
12006 &mut kvl.v,
12007 &kvl.len_d,
12008 t,
12009 kvl.kv_dim_k,
12010 kvl.kv_dim_v,
12011 kvl.k_tok_bytes,
12012 kvl.v_tok_bytes,
12013 (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on()),
12014 )?;
12015 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
12018 let mut attn = e.uninit(t * nh * hd)?;
12019 let k_view = e.view_u8(&kvl.k, kvl.k.len());
12020 let v_view = e.view_u8(&kvl.v, kvl.v.len());
12021 if swa && hint + 1 >= win {
12024 e.fa_decode_rows_w(
12027 &q,
12028 &k_view,
12029 &v_view,
12030 &mut attn,
12031 hd,
12032 nh,
12033 nkv,
12034 &kvl.len_d,
12035 0,
12036 t,
12037 scale,
12038 win,
12039 kvl.k_tok_bytes,
12040 kvl.v_tok_bytes,
12041 None,
12042 )?;
12043 } else if hd == 512 && hint + t < crate::fa512_min_tkv() {
12044 let bucket = (hint + t + 2)
12057 .next_power_of_two()
12058 .min(crate::fa512_min_tkv().saturating_sub(1));
12059 let qv = e.view(&q, t * nh * hd);
12060 for i in 0..t {
12061 let q_row = qv.slice(i * nh * hd..(i + 1) * nh * hd);
12062 let mut q_one = e.uninit(nh * hd)?;
12063 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
12064 let mut a_one = e.uninit(nh * hd)?;
12065 e.fa_decode_dc(
12066 &q_one,
12067 &k_view,
12068 &v_view,
12069 &mut a_one,
12070 hd,
12071 nh,
12072 nkv,
12073 &row_ctrs[i],
12074 bucket,
12075 scale,
12076 kvl.k_tok_bytes,
12077 kvl.v_tok_bytes,
12078 false,
12079 )?;
12080 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
12081 }
12082 } else if hd == 512 {
12083 e.fa_decode_rows(
12086 &q,
12087 &k_view,
12088 &v_view,
12089 &mut attn,
12090 hd,
12091 nh,
12092 nkv,
12093 hint,
12094 t,
12095 scale,
12096 kvl.k_tok_bytes,
12097 kvl.v_tok_bytes,
12098 Some((&kvl.len_d, 0)),
12099 false,
12100 false,
12101 None,
12102 )?;
12103 } else {
12104 e.fa_decode_rows_dc(
12106 &q,
12107 &k_view,
12108 &v_view,
12109 &mut attn,
12110 hd,
12111 nh,
12112 nkv,
12113 &kvl.len_d,
12114 hint + t,
12115 t,
12116 scale,
12117 kvl.k_tok_bytes,
12118 kvl.v_tok_bytes,
12119 0,
12120 swa && crate::Engine::wkv_on(),
12121 )?;
12122 }
12123 Ok(e.matmul(&fa.wo, &attn, t)?)
12124 }
12125
12126 fn gemma4_verify_attn(
12127 &self,
12128 e: &Engine,
12129 fa: &crate::hybrid::FullAttnLayer,
12130 il: usize,
12131 hq: &CudaSlice<i8>,
12132 hdq: &CudaSlice<f32>,
12133 pos_d: &CudaSlice<i32>,
12134 t: usize,
12135 cache: &mut Cache,
12136 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12137 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
12138 let eps = self.cfg.rms_eps;
12139 let aux = self.gemma4_aux.as_ref().unwrap();
12140 let ones = aux.ones(e);
12141 #[cfg(debug_assertions)]
12142 crate::debug_assert_tensor_stream_device(ones, &e.stream(), "gemma4_verify_attn.ones");
12143 let n_embd = self.cfg.n_embd as usize;
12144 let _ = n_embd;
12145
12146 let h0 = e.zeros(0)?;
12147 let h = &h0;
12148 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12151 let f2b = *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
12152 let fused_qkv = if f2b {
12153 if swa {
12154 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
12155 .map(|(a, b, c)| (a, b, Some(c)))
12156 } else {
12157 e.matmul_q4_fused2_batched(&fa.wq, &fa.wk, hq, hdq, t)?
12158 .map(|(a, b)| (a, b, None))
12159 }
12160 } else {
12161 None
12162 };
12163 let (q0, k0, v0) = match fused_qkv {
12164 Some((a, b, cv)) => {
12165 let v = match cv {
12166 Some(c) => c,
12167 None => e.clone_dtod(&b)?,
12168 };
12169 (a, b, v)
12170 }
12171 None => {
12172 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
12173 let k0 = e.matmul_pre(&fa.wk, hq, hdq, h, t)?;
12174 let v0 = if swa {
12175 e.matmul_pre(&fa.wv, hq, hdq, h, t)?
12176 } else {
12177 e.clone_dtod(&k0)?
12178 };
12179 (q0, k0, v0)
12180 }
12181 };
12182 let mut q = e.uninit(t * nh * hd)?;
12183 let mut k = e.uninit(t * nkv * hd)?;
12184 let mut v = e.uninit(t * nkv * hd)?;
12185 let ff = if swa {
12188 None
12189 } else {
12190 Some(
12191 aux.rope_freqs(e)
12192 .expect("gemma4 global rope needs rope_freqs.weight"),
12193 )
12194 };
12195 #[cfg(debug_assertions)]
12196 if let Some(ff) = ff {
12197 crate::debug_assert_tensor_stream_device(
12198 ff,
12199 &e.stream(),
12200 "gemma4_verify_attn.rope_freqs",
12201 );
12202 }
12203 e.rms_norm_qkv_rope(
12204 &q0,
12205 &k0,
12206 &v0,
12207 fa.q_norm.float_data(),
12208 fa.k_norm.float_data(),
12209 ones,
12210 &mut q,
12211 &mut k,
12212 &mut v,
12213 hd,
12214 nh * t,
12215 nkv * t,
12216 pos_d,
12217 nh,
12218 nkv,
12219 base,
12220 1.0,
12221 ff,
12222 eps,
12223 )?;
12224 let kvl = cache.kv[il].as_mut().unwrap();
12225 let base_len = kvl.len;
12226 e.append_kv_quantized_rows(
12227 &k,
12228 &v,
12229 &mut kvl.k,
12230 &mut kvl.v,
12231 base_len,
12232 t,
12233 kvl.kv_dim_k,
12234 kvl.kv_dim_v,
12235 kvl.k_tok_bytes,
12236 kvl.v_tok_bytes,
12237 (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on()),
12238 )?;
12239 kvl.len += t;
12240 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
12241 let mut attn = e.uninit(t * nh * hd)?;
12242 let rows_ok = (hd == 256 && base_len + 1 >= crate::fa_vec_min_tkv())
12245 || (hd == 512 && !swa && base_len + 1 >= crate::fa512_min_tkv());
12248 if rows_ok && (!swa || base_len + t <= win) {
12249 let k_view = e.view_u8(&kvl.k, (base_len + t) * kvl.k_tok_bytes);
12250 let v_view = e.view_u8(&kvl.v, (base_len + t) * kvl.v_tok_bytes);
12251 if hd == 512 {
12252 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
12254 e.fa_decode_rows(
12255 &q,
12256 &k_view,
12257 &v_view,
12258 &mut attn,
12259 hd,
12260 nh,
12261 nkv,
12262 base_len,
12263 t,
12264 scale,
12265 kvl.k_tok_bytes,
12266 kvl.v_tok_bytes,
12267 Some((&kvl.len_d, 0)),
12268 false,
12269 swa && crate::Engine::wkv_on(),
12270 None,
12271 )?;
12272 } else {
12273 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
12277 e.fa_decode_rows_dc(
12278 &q,
12279 &k_view,
12280 &v_view,
12281 &mut attn,
12282 hd,
12283 nh,
12284 nkv,
12285 &kvl.len_d,
12286 base_len + t,
12287 t,
12288 scale,
12289 kvl.k_tok_bytes,
12290 kvl.v_tok_bytes,
12291 0,
12292 swa && crate::Engine::wkv_on(),
12293 )?;
12294 }
12295 return Ok(e.matmul(&fa.wo, &attn, t)?);
12296 }
12297 if hd == 256
12305 && swa
12306 && base_len + 1 >= win
12307 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
12308 {
12309 let k_view = e.view_u8(&kvl.k, (base_len + t) * kvl.k_tok_bytes);
12310 let v_view = e.view_u8(&kvl.v, (base_len + t) * kvl.v_tok_bytes);
12311 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
12312 e.fa_decode_rows_w(
12313 &q,
12314 &k_view,
12315 &v_view,
12316 &mut attn,
12317 hd,
12318 nh,
12319 nkv,
12320 &kvl.len_d,
12321 0,
12322 t,
12323 scale,
12324 win,
12325 kvl.k_tok_bytes,
12326 kvl.v_tok_bytes,
12327 None,
12328 )?;
12329 return Ok(e.matmul(&fa.wo, &attn, t)?);
12330 }
12331 for i in 0..t {
12332 let avail = base_len + i + 1;
12333 let (off_tok, t_kv) = if swa && avail > win {
12334 (avail - win, win)
12335 } else {
12336 (0, avail)
12337 };
12338 let k_view = e.view_u8_range(
12339 &kvl.k,
12340 off_tok * kvl.k_tok_bytes,
12341 (off_tok + t_kv) * kvl.k_tok_bytes,
12342 );
12343 let v_view = e.view_u8_range(
12344 &kvl.v,
12345 off_tok * kvl.v_tok_bytes,
12346 (off_tok + t_kv) * kvl.v_tok_bytes,
12347 );
12348 let qi = e.view(&q, t * nh * hd);
12349 let q_row = qi.slice(i * nh * hd..(i + 1) * nh * hd);
12350 let mut q_one = e.uninit(nh * hd)?;
12351 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
12352 let mut a_one = e.uninit(nh * hd)?;
12353 if swa
12357 && avail > win
12358 && hd == 256
12359 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
12360 {
12361 let kp = e.view_u8(&kvl.k, avail * kvl.k_tok_bytes);
12362 let vp = e.view_u8(&kvl.v, avail * kvl.v_tok_bytes);
12363 e.i32_set_k(&mut kvl.len_d, (avail - 1) as i32)?;
12364 e.fa_decode_rows_w(
12365 &q_one,
12366 &kp,
12367 &vp,
12368 &mut a_one,
12369 hd,
12370 nh,
12371 nkv,
12372 &kvl.len_d,
12373 0,
12374 1,
12375 scale,
12376 win,
12377 kvl.k_tok_bytes,
12378 kvl.v_tok_bytes,
12379 None,
12380 )?;
12381 } else if !swa
12382 && hd == 512
12383 && avail >= crate::fa512_min_tkv()
12384 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
12385 {
12386 let kp = e.view_u8(&kvl.k, avail * kvl.k_tok_bytes);
12387 let vp = e.view_u8(&kvl.v, avail * kvl.v_tok_bytes);
12388 e.i32_set_k(&mut kvl.len_d, (avail - 1) as i32)?;
12389 e.fa_decode_rows(
12390 &q_one,
12391 &kp,
12392 &vp,
12393 &mut a_one,
12394 hd,
12395 nh,
12396 nkv,
12397 avail - 1,
12398 1,
12399 scale,
12400 kvl.k_tok_bytes,
12401 kvl.v_tok_bytes,
12402 Some((&kvl.len_d, 0)),
12403 false,
12404 false,
12405 None,
12406 )?;
12407 } else {
12408 e.fa_decode_kvmod(
12409 &q_one,
12410 &k_view,
12411 &v_view,
12412 &mut a_one,
12413 hd,
12414 nh,
12415 nkv,
12416 t_kv,
12417 scale,
12418 kvl.k_tok_bytes,
12419 kvl.v_tok_bytes,
12420 swa && crate::Engine::wkv_on(),
12421 )?;
12422 }
12423 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
12424 }
12425 Ok(e.matmul(&fa.wo, &attn, t)?)
12426 }
12427
12428 pub(crate) fn gemma4_decode_step_h(
12431 &self,
12432 e: &Engine,
12433 token: u32,
12434 cache: &mut Cache,
12435 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
12436 if let Some(split) = crate::pp::pp2_split(self.layers.len()) {
12441 return self.gemma4_decode_step_h_pp2(e, token, cache, split);
12442 }
12443 if crate::pp::pp_cuts(self.layers.len()).is_some() {
12444 crate::pp::warn_unwired_once("gemma4 eager decode (N>2)");
12445 }
12446 let n_embd = self.cfg.n_embd as usize;
12447 let eps = self.cfg.rms_eps;
12448 let pos_d = e.htod_i32(&[cache.pos as i32])?;
12449 let mut x = e.htod(&self.embd.gather(n_embd, &[token]))?;
12450 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
12451 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
12454 let n_layers = self.layers.len();
12455 for (il, layer) in self.layers.iter().enumerate() {
12456 let (hq, hdq) = match h_carry.take() {
12457 Some(p) => p,
12458 None => {
12459 e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, 1, eps)?
12460 }
12461 };
12462 let Mixer::Full(fa) = &layer.mixer else {
12463 panic!("gemma4 layer {il} not full-attn")
12464 };
12465 let o = self.gemma4_decode_attn(e, fa, il, &hq, &hdq, &pos_d, cache)?;
12466 let next_norm = if il + 1 < n_layers {
12467 Some(self.layers[il + 1].attn_norm.float_data())
12468 } else {
12469 None
12470 };
12471 let (xn, hn) = self.gemma4_layer_tail_add_nq_pn(e, layer, &o, &x, 1, next_norm)?;
12472 x = xn;
12473 h_carry = hn;
12474 }
12475 let mut hn = e.uninit(n_embd)?;
12476 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
12477 let h_seed = e.clone_dtod(&x)?;
12478 let mut ld = e.matmul(&self.output, &hn, 1)?;
12479 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
12480 e.softcap(&mut ld, cap, self.output.out_features())?; self.gemma4_suppress(e, &mut ld, 1)?;
12482 let logits = e.dtoh(&ld)?;
12483 cache.pos += 1;
12484 Ok((logits, h_seed))
12485 }
12486
12487 fn gemma4_decode_layers(
12495 &self,
12496 e: &Engine,
12497 mut x: CudaSlice<f32>,
12498 lo: usize,
12499 hi: usize,
12500 pos_d: &CudaSlice<i32>,
12501 cache: &mut Cache,
12502 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12503 let n_embd = self.cfg.n_embd as usize;
12504 let eps = self.cfg.rms_eps;
12505 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
12506 for il in lo..hi {
12507 let layer = &self.layers[il];
12508 let (hq, hdq) = match h_carry.take() {
12509 Some(p) => p,
12510 None => {
12512 e.rms_norm_q8_1(&x, self.layers[il].attn_norm.float_data(), n_embd, 1, eps)?
12513 }
12514 };
12515 let Mixer::Full(fa) = &layer.mixer else {
12516 panic!("gemma4 layer {il} not full-attn")
12517 };
12518 let o = self.gemma4_decode_attn(e, fa, il, &hq, &hdq, pos_d, cache)?;
12519 let next_norm = if il + 1 < hi {
12520 Some(self.layers[il + 1].attn_norm.float_data())
12521 } else {
12522 None
12523 };
12524 let (xn, hn) = self.gemma4_layer_tail_add_nq_pn(e, layer, &o, &x, 1, next_norm)?;
12525 x = xn;
12526 h_carry = hn;
12527 }
12528 Ok(x)
12529 }
12530
12531 fn gemma4_decode_step_h_pp2(
12539 &self,
12540 e: &Engine,
12541 token: u32,
12542 cache: &mut Cache,
12543 split: usize,
12544 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
12545 if crate::pp::pp2_streams_off() {
12546 return self.gemma4_decode_step_h_pp2_samestream(e, token, cache, split);
12547 }
12548 let rt = crate::pp::Pp2Rt::get(e)?;
12549 let e0 = rt.engine(0, e);
12550 let e1 = rt.engine(1, e);
12551 let n_embd = self.cfg.n_embd as usize;
12552 let eps = self.cfg.rms_eps;
12553 let pos = cache.pos as i32;
12554
12555 let slot = {
12557 let _st0 = rt.enter(0);
12558 let pos_d = e0.htod_i32(&[pos])?;
12559 #[cfg(debug_assertions)]
12560 crate::debug_assert_tensor_stream_device(
12561 &pos_d,
12562 &e0.stream(),
12563 "gemma4_decode_step_h_pp2.stage0.pos_d",
12564 );
12565 let mut x = e0.htod(&self.embd.gather(n_embd, &[token]))?;
12566 e0.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
12567 let x = self.gemma4_decode_layers(e0, x, 0, split, &pos_d, cache)?;
12568 rt.tx(0, &x, n_embd)?
12569 };
12570
12571 let _st1 = rt.enter(1);
12573 let pos_d = e1.htod_i32(&[pos])?;
12574 #[cfg(debug_assertions)]
12575 crate::debug_assert_tensor_stream_device(
12576 &pos_d,
12577 &e1.stream(),
12578 "gemma4_decode_step_h_pp2.stage1.pos_d",
12579 );
12580 let x = rt.rx(0, slot, n_embd)?;
12581 let x = self.gemma4_decode_layers(e1, x, split, self.layers.len(), &pos_d, cache)?;
12582
12583 let mut hn = e1.uninit(n_embd)?;
12584 e1.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
12585 let h_seed = e1.clone_dtod(&x)?;
12586 let mut ld = e1.matmul(&self.output, &hn, 1)?;
12587 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
12588 e1.softcap(&mut ld, cap, self.output.out_features())?;
12589 self.gemma4_suppress(e1, &mut ld, 1)?;
12590 let logits = e1.dtoh(&ld)?;
12591 cache.pos += 1;
12592 Ok((logits, h_seed))
12593 }
12594
12595 fn gemma4_decode_step_h_pp2_samestream(
12598 &self,
12599 e: &Engine,
12600 token: u32,
12601 cache: &mut Cache,
12602 split: usize,
12603 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
12604 let n_embd = self.cfg.n_embd as usize;
12605 let eps = self.cfg.rms_eps;
12606 let pos_d = e.htod_i32(&[cache.pos as i32])?;
12607
12608 let mut x = e.htod(&self.embd.gather(n_embd, &[token]))?;
12610 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
12611 let x = self.gemma4_decode_layers(e, x, 0, split, &pos_d, cache)?;
12612
12613 let boundary_tx = e.clone_dtod(&x)?;
12615 let boundary_rx = e.clone_dtod(&boundary_tx)?;
12616
12617 let x =
12619 self.gemma4_decode_layers(e, boundary_rx, split, self.layers.len(), &pos_d, cache)?;
12620
12621 let mut hn = e.uninit(n_embd)?;
12622 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
12623 let h_seed = e.clone_dtod(&x)?;
12624 let mut ld = e.matmul(&self.output, &hn, 1)?;
12625 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
12626 e.softcap(&mut ld, cap, self.output.out_features())?;
12627 self.gemma4_suppress(e, &mut ld, 1)?;
12628 let logits = e.dtoh(&ld)?;
12629 cache.pos += 1;
12630 Ok((logits, h_seed))
12631 }
12632}
12633
12634impl HybridModel {
12653 pub(crate) fn step35_geom(&self, il: usize) -> memra_gguf::config::LayerGeometry {
12656 let geometry = self
12657 .cfg
12658 .layer_geometry(il as u32)
12659 .unwrap_or_else(|| panic!("step35 layer {il} has no geometry-table row"));
12660 debug_assert_eq!(
12661 geometry.attention_gate,
12662 memra_gguf::config::AttentionGateKind::SeparateHead
12663 );
12664 geometry
12665 }
12666
12667 #[allow(clippy::too_many_arguments)]
12727 fn step35_attn_pre_wo(
12728 &self,
12729 e: &Engine,
12730 fa: &FullAttnLayer,
12731 mut g3: Vec<CudaSlice<f32>>,
12732 hg: Option<&CudaSlice<f32>>,
12733 gt_pre: Option<&CudaSlice<f32>>,
12734 pos_d: &CudaSlice<i32>,
12735 t: usize,
12736 cache: Option<&mut Cache>,
12737 il: usize,
12738 seq_end: usize,
12739 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12740 let geometry = self.step35_geom(il);
12741 let hd = geometry.head_dim_k as usize;
12742 let nkv = geometry.n_head_kv as usize;
12743 let nh = geometry.n_head as usize;
12744 let rbase = geometry.rope_base;
12745 let scale = geometry.attention_scale();
12746 let swa = geometry.window.is_some();
12747 let eps = self.cfg.rms_eps;
12748 let win = geometry.window.unwrap_or(0) as usize;
12749 let n_rot = geometry.n_rot as usize;
12750
12751 let v = g3.pop().unwrap();
12752 let k0 = g3.pop().unwrap();
12753 let q0 = g3.pop().unwrap();
12754
12755 let mut q = e.uninit(t * nh * hd)?;
12759 e.rms_norm(&q0, fa.q_norm.float_data(), &mut q, hd, nh * t, eps)?;
12760 let mut k = e.uninit(t * nkv * hd)?;
12761 e.rms_norm(&k0, fa.k_norm.float_data(), &mut k, hd, nkv * t, eps)?;
12762 let ff = if geometry.rope_factors {
12763 self.step35_aux.as_ref().and_then(|a| a.rope_freqs(e))
12764 } else {
12765 None
12766 };
12767 #[cfg(debug_assertions)]
12768 if let Some(ff) = ff {
12769 crate::debug_assert_tensor_stream_device(
12770 ff,
12771 &e.stream(),
12772 "step35_attn_pre_wo.rope_freqs",
12773 );
12774 }
12775 e.rope_neox2(&mut q, &mut k, pos_d, hd, n_rot, nh, nkv, t, rbase, 1.0, ff)?;
12776
12777 let mut attn = e.uninit(t * nh * hd)?;
12778 match cache {
12779 Some(cache) => {
12780 let base_len = cache.kv[il].as_ref().unwrap().len;
12781 let legacy_tkv = std::env::var("MEMRA_STEP35_SWA_TKV").as_deref() == Ok("1");
12783 let legacy_calllocal = std::env::var("MEMRA_PRIME_CALLLOCAL").as_deref() == Ok("1");
12784 let off = if swa {
12785 let raw = base_len.saturating_sub(win - 1);
12786 if legacy_tkv || legacy_calllocal {
12787 raw
12788 } else {
12789 raw & !31usize
12790 }
12791 } else {
12792 0
12793 };
12794 {
12795 let kvl = cache.kv[il].as_mut().unwrap();
12796 assert!(kvl.len + t <= cache.max_ctx, "step35 prime: KV overflow");
12797 let write_row = e.prepare_kv_append(kvl, off, t)?;
12798 e.append_kv_quantized_rows(
12799 &k,
12800 &v,
12801 &mut kvl.k,
12802 &mut kvl.v,
12803 write_row,
12804 t,
12805 kvl.kv_dim_k,
12806 kvl.kv_dim_v,
12807 kvl.k_tok_bytes,
12808 kvl.v_tok_bytes,
12809 crate::Engine::kv_fp8_on(),
12810 )?;
12811 kvl.len += t;
12812 let new_len = kvl.len as i32;
12813 e.set_i32_one(&mut kvl.len_d, new_len)?;
12814 }
12815 let kvl = cache.kv[il].as_ref().unwrap();
12816 let t_kv = base_len + t - off;
12839 let physical = kvl.physical_rows(off, off + t_kv)?;
12840 let k_view = e.view_u8_range(
12841 &kvl.k,
12842 physical.start * kvl.k_tok_bytes,
12843 physical.end * kvl.k_tok_bytes,
12844 );
12845 let v_view = e.view_u8_range(
12846 &kvl.v,
12847 physical.start * kvl.v_tok_bytes,
12848 physical.end * kvl.v_tok_bytes,
12849 );
12850 let swa_naive = if legacy_tkv {
12862 t_kv > win
12863 } else {
12864 seq_end > win
12865 };
12866 if swa && swa_naive {
12867 if std::env::var("MEMRA_STEP35_SWA_FA").as_deref() == Ok("0") {
12880 e.sdpa_naive_w_quantized_view(
12881 &q,
12882 &k_view,
12883 &v_view,
12884 &mut attn,
12885 hd,
12886 nh,
12887 nkv,
12888 t,
12889 t_kv,
12890 scale,
12891 true,
12892 win,
12893 kvl.k_tok_bytes,
12894 kvl.v_tok_bytes,
12895 )?;
12896 } else {
12897 e.fa_prefill_view_ws_w_hd128(
12898 &q,
12899 &k_view,
12900 &v_view,
12901 &mut attn,
12902 hd,
12903 nh,
12904 nkv,
12905 t,
12906 t_kv,
12907 scale,
12908 true,
12909 win,
12910 kvl.k_tok_bytes,
12911 kvl.v_tok_bytes,
12912 )?;
12913 }
12914 } else if std::env::var("MEMRA_NOFA").is_ok() {
12915 e.sdpa_naive_quantized_view(
12916 &q,
12917 &k_view,
12918 &v_view,
12919 &mut attn,
12920 hd,
12921 nh,
12922 nkv,
12923 t,
12924 t_kv,
12925 scale,
12926 true,
12927 kvl.k_tok_bytes,
12928 kvl.v_tok_bytes,
12929 )?;
12930 } else {
12931 e.fa_prefill_view_ws(
12936 &q,
12937 &k_view,
12938 &v_view,
12939 &mut attn,
12940 hd,
12941 nh,
12942 nkv,
12943 t,
12944 t_kv,
12945 scale,
12946 true,
12947 kvl.k_tok_bytes,
12948 kvl.v_tok_bytes,
12949 crate::Engine::kv_fp8_on(),
12950 )?;
12951 }
12952 }
12953 None => {
12954 debug_assert_eq!(
12959 seq_end, t,
12960 "step35 cacheless prefill is monolithic (seq_end == t)"
12961 );
12962 if swa && seq_end > win {
12963 e.sdpa_naive_w(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
12964 } else if std::env::var("MEMRA_NOFA").is_ok() {
12965 e.sdpa_naive(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
12966 } else {
12967 e.fa_prefill(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
12968 }
12969 }
12970 }
12971
12972 let gw = fa
12975 .attn_gate
12976 .as_ref()
12977 .ok_or("step35 layer is missing attn_gate.weight (head-wise attention gate)")?;
12978 let gt_owned = if gt_pre.is_none() {
12979 Some(e.matmul(
12980 gw,
12981 hg.ok_or("step35 attention needs hg when gt_pre is absent")?,
12982 t,
12983 )?)
12984 } else {
12985 None
12986 };
12987 let gt = gt_pre.or(gt_owned.as_ref()).unwrap();
12988 let mut ag = e.uninit(t * nh * hd)?;
12989 e.attn_head_gate(&attn, gt, &mut ag, None, hd, nh, t)?;
12990 Ok(ag)
12991 }
12992
12993 pub(crate) fn step35_attn(
12996 &self,
12997 e: &Engine,
12998 fa: &FullAttnLayer,
12999 h: &CudaSlice<f32>,
13000 pos_d: &CudaSlice<i32>,
13001 t: usize,
13002 il: usize,
13003 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
13004 let g3 = e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?;
13005 let ag = self.step35_attn_pre_wo(e, fa, g3, Some(h), None, pos_d, t, None, il, t)?;
13007 Ok(e.matmul(&fa.wo, &ag, t)?)
13008 }
13009
13010 #[allow(clippy::too_many_arguments)]
13017 pub(crate) fn step35_attn_prime(
13018 &self,
13019 e: &Engine,
13020 fa: &FullAttnLayer,
13021 h: &CudaSlice<f32>,
13022 hx: Option<&CudaSlice<u8>>,
13023 pos_d: &CudaSlice<i32>,
13024 t: usize,
13025 cache: &mut Cache,
13026 il: usize,
13027 seq_end: usize,
13028 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
13029 let g3 = match hx {
13030 Some(xh) => e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], h, xh, t)?,
13031 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
13032 };
13033 let ag =
13034 self.step35_attn_pre_wo(e, fa, g3, Some(h), None, pos_d, t, Some(cache), il, seq_end)?;
13035 Ok(e.matmul(&fa.wo, &ag, t)?)
13036 }
13037
13038 #[allow(clippy::too_many_arguments)]
13048 pub(crate) fn step35_decode_attn(
13049 &self,
13050 e: &Engine,
13051 fa: &FullAttnLayer,
13052 il: usize,
13053 h: &CudaSlice<f32>,
13054 pre_q: Option<(&CudaSlice<i8>, &CudaSlice<f32>)>,
13055 pos_d: &CudaSlice<i32>,
13056 cache: &mut Cache,
13057 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
13058 let geometry = self.step35_geom(il);
13059 let hd = geometry.head_dim_k as usize;
13060 let nkv = geometry.n_head_kv as usize;
13061 let nh = geometry.n_head as usize;
13062 let rbase = geometry.rope_base;
13063 let scale = geometry.attention_scale();
13064 let swa = geometry.window.is_some();
13065 let eps = self.cfg.rms_eps;
13066 let win = geometry.window.unwrap_or(0) as usize;
13067 let n_rot = geometry.n_rot as usize;
13068 let n_embd = self.cfg.n_embd as usize;
13069 let gw = fa
13070 .attn_gate
13071 .as_ref()
13072 .ok_or("step35 layer is missing attn_gate.weight (head-wise attention gate)")?;
13073
13074 let (q0, k0, v0, gt) = match pre_q {
13075 Some((hq, hdq)) => {
13076 debug_assert!(
13077 e.uses_q8_1_fast(gw),
13078 "step35 pre-quantized decode requires attn_gate on the q8_1 fast path \
13079 (h is a zero-length placeholder here) — see mixer_in_q8_1_fast"
13080 );
13081 let (a, b, c) = match e.matmul_q8_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)? {
13082 Some(t3) => t3,
13083 None => (
13084 e.matmul_pre(&fa.wq, hq, hdq, h, 1)?,
13085 e.matmul_pre(&fa.wk, hq, hdq, h, 1)?,
13086 e.matmul_pre(&fa.wv, hq, hdq, h, 1)?,
13087 ),
13088 };
13089 let gt = e.matmul_pre(gw, hq, hdq, h, 1)?;
13090 (a, b, c, gt)
13091 }
13092 None => {
13093 if e.uses_q8_1_fast(&fa.wq)
13094 && e.uses_q8_1_fast(&fa.wk)
13095 && e.uses_q8_1_fast(&fa.wv)
13096 && e.uses_q8_1_fast(gw)
13097 {
13098 let (hq, hdq) = e.quantize_q8_1(h, 1, n_embd)?;
13099 let (a, b, c) = match e.matmul_q8_fused3(&fa.wq, &fa.wk, &fa.wv, &hq, &hdq)? {
13100 Some(t3) => t3,
13101 None => (
13102 e.matmul_pre(&fa.wq, &hq, &hdq, h, 1)?,
13103 e.matmul_pre(&fa.wk, &hq, &hdq, h, 1)?,
13104 e.matmul_pre(&fa.wv, &hq, &hdq, h, 1)?,
13105 ),
13106 };
13107 let gt = e.matmul_pre(gw, &hq, &hdq, h, 1)?;
13108 (a, b, c, gt)
13109 } else {
13110 (
13111 e.matmul(&fa.wq, h, 1)?,
13112 e.matmul(&fa.wk, h, 1)?,
13113 e.matmul(&fa.wv, h, 1)?,
13114 e.matmul(gw, h, 1)?,
13115 )
13116 }
13117 }
13118 };
13119
13120 let mut q = e.uninit(nh * hd)?;
13121 e.rms_norm(&q0, fa.q_norm.float_data(), &mut q, hd, nh, eps)?;
13122 let mut k = e.uninit(nkv * hd)?;
13123 e.rms_norm(&k0, fa.k_norm.float_data(), &mut k, hd, nkv, eps)?;
13124 let ff = if swa {
13125 None
13126 } else {
13127 self.step35_aux.as_ref().and_then(|a| a.rope_freqs(e))
13128 };
13129 #[cfg(debug_assertions)]
13130 if let Some(ff) = ff {
13131 crate::debug_assert_tensor_stream_device(
13132 ff,
13133 &e.stream(),
13134 "step35_decode_attn.rope_freqs",
13135 );
13136 }
13137 e.rope_neox2(&mut q, &mut k, pos_d, hd, n_rot, nh, nkv, 1, rbase, 1.0, ff)?;
13138
13139 if std::env::var("MEMRA_NOFA").is_ok() {
13140 return Err(
13141 "MEMRA_NOFA (naive f32 SDPA) is incompatible with the quantized KV \
13142 cache; unset MEMRA_NOFA to use fa_decode"
13143 .into(),
13144 );
13145 }
13146 let kvl = cache.kv[il].as_mut().unwrap();
13147 let next_len = kvl.len + 1;
13148 let (off, t_kv) = if swa && next_len > win {
13149 (next_len - win, win)
13150 } else {
13151 (0, next_len)
13152 };
13153 let write_row = e.prepare_kv_append(kvl, off & !31usize, 1)?;
13154 e.append_kv_quantized(
13155 &k,
13156 &v0,
13157 &mut kvl.k,
13158 &mut kvl.v,
13159 write_row,
13160 kvl.kv_dim_k,
13161 kvl.kv_dim_v,
13162 kvl.k_tok_bytes,
13163 kvl.v_tok_bytes,
13164 crate::Engine::kv_fp8_on(),
13165 )?;
13166 kvl.len = next_len;
13167 let physical = kvl.physical_rows(off, off + t_kv)?;
13168 let k_view = e.view_u8_range(
13169 &kvl.k,
13170 physical.start * kvl.k_tok_bytes,
13171 physical.end * kvl.k_tok_bytes,
13172 );
13173 let v_view = e.view_u8_range(
13174 &kvl.v,
13175 physical.start * kvl.v_tok_bytes,
13176 physical.end * kvl.v_tok_bytes,
13177 );
13178 let mut attn = e.uninit(nh * hd)?;
13179 e.fa_decode_kvmod(
13180 &q,
13181 &k_view,
13182 &v_view,
13183 &mut attn,
13184 hd,
13185 nh,
13186 nkv,
13187 t_kv,
13188 scale,
13189 kvl.k_tok_bytes,
13190 kvl.v_tok_bytes,
13191 crate::Engine::kv_fp8_on(),
13192 )?;
13193
13194 let mut ag = e.uninit(nh * hd)?;
13195 e.attn_head_gate(&attn, >, &mut ag, None, hd, nh, 1)?;
13196 Ok(e.matmul(&fa.wo, &ag, 1)?)
13197 }
13198}
13199
13200impl HybridModel {
13209 pub fn is_gemma4_e4b(&self) -> bool {
13210 self.gemma4_aux.as_ref().is_some_and(|a| a.e4b.is_some())
13211 }
13212
13213 fn gemma4_e4b_geom(&self, il: usize) -> (usize, usize, usize, f32, f32, bool) {
13217 let g = self.cfg.gemma4.as_ref().unwrap();
13218 let swa = g.swa_pattern[il];
13219 let hd = if swa {
13220 g.key_length_swa
13221 } else {
13222 g.key_length_global
13223 } as usize;
13224 let Mixer::Full(fa) = &self.layers[il].mixer else {
13225 panic!("e4b layer {il} not full-attn")
13226 };
13227 let nh = fa.wq.out_features() / hd;
13228 let nkv = fa.wk.out_features() / hd;
13229 (
13230 hd,
13231 nkv,
13232 nh,
13233 if swa {
13234 g.rope_base_swa
13235 } else {
13236 g.rope_base_global
13237 },
13238 1.0,
13239 swa,
13240 )
13241 }
13242
13243 fn gemma4_e4b_kv_target(&self, il: usize) -> Option<usize> {
13245 self.layers[il]
13246 .gemma4
13247 .as_ref()
13248 .and_then(|b| b.e4b.as_ref())
13249 .and_then(|e4| e4.kv_share.map(|t| t as usize))
13250 }
13251
13252 fn gemma4_e4b_inp_pl(
13257 &self,
13258 e: &Engine,
13259 tokens: &[u32],
13260 x_scaled: &CudaSlice<f32>,
13261 t: usize,
13262 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
13263 let tok_d = e.stream().clone_htod(&tokens.to_vec())?;
13264 self.gemma4_e4b_inp_pl_dev(e, &tok_d, x_scaled, t)
13265 }
13266
13267 fn gemma4_e4b_inp_pl_dev(
13269 &self,
13270 e: &Engine,
13271 tok_d: &CudaSlice<u32>,
13272 x_scaled: &CudaSlice<f32>,
13273 t: usize,
13274 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
13275 let aux = self.gemma4_aux.as_ref().unwrap();
13276 let m = aux.e4b.as_ref().unwrap();
13277 let n_embd = self.cfg.n_embd as usize;
13278 let n_layer = self.layers.len();
13279 let width = m.n_epl * n_layer;
13280 let tbl = m.tok_tbl_gpu.get_or_init(|| {
13281 e.upload_u8(&m.tok_embd_bytes)
13282 .expect("e4b per-layer token table upload")
13283 });
13284 let mut a =
13285 e.embed_gather_device_td(tbl, tok_d, t, width, m.tok_embd_qt, m.tok_embd_row_bytes)?;
13286 e.scale_inplace(&mut a, (m.n_epl as f32).sqrt(), t * width)?;
13287 let mut p = e.matmul(&m.model_proj, x_scaled, t)?;
13288 e.scale_inplace(&mut p, 1.0 / (n_embd as f32).sqrt(), t * width)?;
13289 let mut pn = e.uninit(t * width)?;
13290 e.rms_norm(
13291 &p,
13292 m.proj_norm.float_data(),
13293 &mut pn,
13294 m.n_epl,
13295 t * n_layer,
13296 self.cfg.rms_eps,
13297 )?;
13298 let mut out = e.uninit(t * width)?;
13299 e.add_scale(&a, &pn, 1.0 / 2f32.sqrt(), &mut out, t * width)?;
13300 Ok(out)
13301 }
13302
13303 #[allow(clippy::too_many_arguments)]
13308 fn gemma4_e4b_attn(
13309 &self,
13310 e: &Engine,
13311 il: usize,
13312 hq: &CudaSlice<i8>,
13313 hdq: &CudaSlice<f32>,
13314 pos_d: &CudaSlice<i32>,
13315 t: usize,
13316 cache: &mut Cache,
13317 dc_bucket: Option<usize>,
13318 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
13319 let (hd, nkv, nh, base, scale, swa) = self.gemma4_e4b_geom(il);
13320 let eps = self.cfg.rms_eps;
13321 let aux = self.gemma4_aux.as_ref().unwrap();
13322 let ones = aux.ones(e);
13323 #[cfg(debug_assertions)]
13324 crate::debug_assert_tensor_stream_device(ones, &e.stream(), "gemma4_e4b_attn.ones");
13325 let Mixer::Full(fa) = &self.layers[il].mixer else {
13326 unreachable!()
13327 };
13328 let h0 = e.zeros(0)?;
13332 let h = &h0;
13333
13334 let ff = if swa {
13335 None
13336 } else {
13337 Some(
13338 aux.rope_freqs(e)
13339 .expect("e4b global rope needs rope_freqs.weight"),
13340 )
13341 };
13342 #[cfg(debug_assertions)]
13343 if let Some(ff) = ff {
13344 crate::debug_assert_tensor_stream_device(ff, &e.stream(), "gemma4_e4b_attn.rope_freqs");
13345 }
13346 let share = self.gemma4_e4b_kv_target(il);
13347 let mut kv_f32: Option<(CudaSlice<f32>, CudaSlice<f32>)> = None;
13349 let mut q;
13350 if let Some(_tgt) = share {
13351 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
13352 q = e.uninit(t * nh * hd)?;
13353 let mut kdummy = e.uninit(1)?;
13356 let mut vdummy = e.uninit(1)?;
13357 e.rms_norm_qkv_rope(
13358 &q0,
13359 &q0,
13360 &q0,
13361 fa.q_norm.float_data(),
13362 fa.q_norm.float_data(),
13363 ones,
13364 &mut q,
13365 &mut kdummy,
13366 &mut vdummy,
13367 hd,
13368 nh * t,
13369 0,
13370 pos_d,
13371 nh,
13372 1,
13373 base,
13374 1.0,
13375 ff,
13376 eps,
13377 )?;
13378 } else {
13379 let e4bits = self.layers[il].gemma4.as_ref().and_then(|g| g.e4b.as_ref());
13383 let cat = e4bits.and_then(|e4| e4.qkv_cat.as_ref());
13384 q = e.uninit(t * nh * hd)?;
13385 let mut k = e.uninit(t * nkv * hd)?;
13386 let mut v = e.uninit(t * nkv * hd)?;
13387 if t == 1 && cat.is_some() {
13388 let qkv0 = e.matmul_pre(cat.unwrap(), hq, hdq, h, 1)?;
13389 e.rms_norm_qkv_rope_cat(
13390 &qkv0,
13391 fa.q_norm.float_data(),
13392 fa.k_norm.float_data(),
13393 ones,
13394 &mut q,
13395 &mut k,
13396 &mut v,
13397 hd,
13398 nh,
13399 nkv,
13400 pos_d,
13401 nh,
13402 nkv,
13403 base,
13404 1.0,
13405 ff,
13406 eps,
13407 )?;
13408 } else {
13409 let (q0, k0, v0) = match if t == 1 {
13410 e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)?
13411 } else {
13412 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
13415 if *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0")) {
13416 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
13417 } else {
13418 None
13419 }
13420 } {
13421 Some(triple) => triple,
13422 None => (
13423 e.matmul_pre(&fa.wq, hq, hdq, h, t)?,
13424 e.matmul_pre(&fa.wk, hq, hdq, h, t)?,
13425 e.matmul_pre(&fa.wv, hq, hdq, h, t)?,
13426 ), };
13428 e.rms_norm_qkv_rope(
13431 &q0,
13432 &k0,
13433 &v0,
13434 fa.q_norm.float_data(),
13435 fa.k_norm.float_data(),
13436 ones,
13437 &mut q,
13438 &mut k,
13439 &mut v,
13440 hd,
13441 nh * t,
13442 nkv * t,
13443 pos_d,
13444 nh,
13445 nkv,
13446 base,
13447 1.0,
13448 ff,
13449 eps,
13450 )?;
13451 }
13452 let kvl = cache.kv[il].as_mut().unwrap();
13453 let cls = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
13457 if dc_bucket.is_some() {
13458 debug_assert!(t == 1);
13463 e.append_kv_quantized_row_dc_inc(
13465 &k,
13466 &v,
13467 &mut kvl.k,
13468 &mut kvl.v,
13469 &mut kvl.len_d,
13470 kvl.kv_dim_k,
13471 kvl.kv_dim_v,
13472 kvl.k_tok_bytes,
13473 kvl.v_tok_bytes,
13474 cls,
13475 )?;
13476 } else {
13477 e.append_kv_quantized_rows(
13478 &k,
13479 &v,
13480 &mut kvl.k,
13481 &mut kvl.v,
13482 kvl.len,
13483 t,
13484 kvl.kv_dim_k,
13485 kvl.kv_dim_v,
13486 kvl.k_tok_bytes,
13487 kvl.v_tok_bytes,
13488 cls,
13489 )?;
13490 kvl.len += t;
13491 }
13492 kv_f32 = Some((k, v));
13493 }
13494 let kvl_idx = share.unwrap_or(il);
13497 let kvl = cache.kv[kvl_idx].as_ref().unwrap();
13498 let base_len = kvl.len - t; let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
13500 let mut attn = e.uninit(t * nh * hd)?;
13501 if t > 1 && base_len == 0 && std::env::var("MEMRA_NOFA").is_err() {
13513 if let Some((kf, vf)) = &kv_f32 {
13514 if hd == 256 && t <= win {
13515 e.fa_prefill(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale, true)?;
13516 return Ok(e.matmul(&fa.wo, &attn, t)?);
13517 }
13518 if hd == 256 && swa && t > win {
13519 e.fa_prefill_w(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
13520 return Ok(e.matmul(&fa.wo, &attn, t)?);
13521 }
13522 if hd == 512 && !swa {
13523 e.fa_prefill_hd512(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale, true)?;
13524 return Ok(e.matmul(&fa.wo, &attn, t)?);
13525 }
13526 } else if share.is_some() {
13527 let g = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
13528 let k_view = e.view_u8(&kvl.k, kvl.k.len());
13529 let v_view = e.view_u8(&kvl.v, kvl.v.len());
13530 if hd == 256 && (!swa || t <= win) {
13531 e.fa_prefill_view(
13533 &q,
13534 &k_view,
13535 &v_view,
13536 &mut attn,
13537 hd,
13538 nh,
13539 nkv,
13540 t,
13541 t,
13542 scale,
13543 true,
13544 kvl.k_tok_bytes,
13545 kvl.v_tok_bytes,
13546 g,
13547 )?;
13548 return Ok(e.matmul(&fa.wo, &attn, t)?);
13549 }
13550 let kv_dim = nkv * hd;
13553 let mut kf = e.uninit(t * kv_dim)?;
13554 let mut vf = e.uninit(t * kv_dim)?;
13555 e.fa_dequant_kv_view_f32(
13556 &k_view,
13557 &v_view,
13558 &mut kf,
13559 &mut vf,
13560 kv_dim,
13561 kv_dim,
13562 t,
13563 kvl.k_tok_bytes,
13564 kvl.v_tok_bytes,
13565 g,
13566 )?;
13567 if hd == 512 {
13568 e.fa_prefill_hd512(&q, &kf, &vf, &mut attn, hd, nh, nkv, t, t, scale, true)?;
13569 } else {
13570 e.fa_prefill_w(&q, &kf, &vf, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
13571 }
13572 return Ok(e.matmul(&fa.wo, &attn, t)?);
13573 }
13574 }
13575 if let Some(bucket) = dc_bucket {
13576 assert!(t == 1);
13581 let bucket = if hd == 512 && win <= crate::fa512_min_tkv() {
13587 bucket.min(crate::fa512_min_tkv().saturating_sub(1))
13588 } else {
13589 bucket
13590 };
13591 let k_view = e.view_u8(&kvl.k, kvl.k.len());
13592 let v_view = e.view_u8(&kvl.v, kvl.v.len());
13593 let g = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
13594 if crate::Engine::wpf_level() >= 1 {
13602 e.prefetch_weight_l2(&fa.wo)?;
13603 }
13604 if e.uses_q8_1_fast(&fa.wo) {
13607 let mut oq = e.alloc_i8_uninit(nh * hd)?;
13608 let mut od = e.zeros(nh * hd / 32)?;
13609 e.fa_decode_dc_q8(
13610 &q,
13611 &k_view,
13612 &v_view,
13613 &mut attn,
13614 hd,
13615 nh,
13616 nkv,
13617 &kvl.len_d,
13618 bucket,
13619 scale,
13620 kvl.k_tok_bytes,
13621 kvl.v_tok_bytes,
13622 g,
13623 Some((&mut oq, &mut od)),
13624 )?;
13625 return Ok(e.matmul_pre(&fa.wo, &oq, &od, &attn, t)?);
13626 }
13627 e.fa_decode_dc(
13628 &q,
13629 &k_view,
13630 &v_view,
13631 &mut attn,
13632 hd,
13633 nh,
13634 nkv,
13635 &kvl.len_d,
13636 bucket,
13637 scale,
13638 kvl.k_tok_bytes,
13639 kvl.v_tok_bytes,
13640 g,
13641 )?;
13642 return Ok(e.matmul(&fa.wo, &attn, t)?);
13643 }
13644 for i in 0..t {
13645 let avail = base_len + i + 1;
13646 let (off_tok, t_kv) = if swa && avail > win {
13647 (avail - win, win)
13648 } else {
13649 (0, avail)
13650 };
13651 let k_view = e.view_u8_range(
13652 &kvl.k,
13653 off_tok * kvl.k_tok_bytes,
13654 (off_tok + t_kv) * kvl.k_tok_bytes,
13655 );
13656 let v_view = e.view_u8_range(
13657 &kvl.v,
13658 off_tok * kvl.v_tok_bytes,
13659 (off_tok + t_kv) * kvl.v_tok_bytes,
13660 );
13661 let qv = e.view(&q, t * nh * hd);
13662 let q_row = qv.slice(i * nh * hd..(i + 1) * nh * hd);
13663 let mut q_one = e.uninit(nh * hd)?;
13664 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
13665 let mut a_one = e.uninit(nh * hd)?;
13666 e.fa_decode_kvmod(
13670 &q_one,
13671 &k_view,
13672 &v_view,
13673 &mut a_one,
13674 hd,
13675 nh,
13676 nkv,
13677 t_kv,
13678 scale,
13679 kvl.k_tok_bytes,
13680 kvl.v_tok_bytes,
13681 (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on()),
13682 )?;
13683 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
13684 }
13685 Ok(e.matmul(&fa.wo, &attn, t)?)
13686 }
13687
13688 fn gemma4_e4b_trunk(
13693 &self,
13694 e: &Engine,
13695 tokens: &[u32],
13696 pos0: usize,
13697 cache: &mut Cache,
13698 head_last: bool,
13699 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
13700 let n_embd = self.cfg.n_embd as usize;
13701 let t = tokens.len();
13702 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
13703 let pos_d = e.htod_i32(&pos)?;
13704 let mut x = e.htod(&self.embd.gather(n_embd, tokens))?;
13705 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
13706 let inp_pl = self.gemma4_e4b_inp_pl(e, tokens, &x, t)?;
13707 self.gemma4_e4b_trunk_core(e, x, inp_pl, &pos_d, t, cache, None, true, head_last)
13708 }
13709
13710 fn gemma4_e4b_trunk_core(
13714 &self,
13715 e: &Engine,
13716 x_in: CudaSlice<f32>,
13717 inp_pl: CudaSlice<f32>,
13718 pos_d: &CudaSlice<i32>,
13719 t: usize,
13720 cache: &mut Cache,
13721 dc_bucket: Option<usize>,
13722 cap_logits: bool,
13723 head_last: bool,
13724 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
13725 let n_embd = self.cfg.n_embd as usize;
13726 let eps = self.cfg.rms_eps;
13727 let n_layer = self.layers.len();
13728 let mut x = x_in;
13729 let aux_e4b = self.gemma4_aux.as_ref().unwrap().e4b.as_ref().unwrap();
13730 let n_epl = aux_e4b.n_epl;
13731
13732 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
13738 for il in 0..n_layer {
13739 let layer = &self.layers[il];
13740 let (hq, hdq) = match h_carry.take() {
13741 Some(p) => p,
13742 None => e.rms_norm_q8_1(&x, layer.attn_norm.float_data(), n_embd, t, eps)?,
13743 };
13744 let o = self.gemma4_e4b_attn(e, il, &hq, &hdq, pos_d, t, cache, dc_bucket)?;
13745 let bits = layer.gemma4.as_ref().unwrap();
13748 let e4b = bits.e4b.as_ref().expect("e4b layer bits");
13749 let fuse_exit = e.uses_q8_1_fast(&e4b.inp_gate);
13760 let (sn, attn_out) = self.gemma4_layer_tail_core_pn(
13761 e,
13762 layer,
13763 &o,
13764 &x,
13765 t,
13766 Some(layer.post_attn_norm.float_data()),
13767 fuse_exit,
13768 )?;
13769 let mut resid = e.uninit(t * n_embd)?;
13770 let g = if fuse_exit {
13776 let (rq, rd) = e.rms_pre_add_q8_1(
13778 &sn,
13779 bits.post_ffw_norm.float_data(),
13780 &attn_out,
13781 &mut resid,
13782 n_embd,
13783 t,
13784 self.cfg.rms_eps,
13785 )?;
13786 e.matmul_pre(&e4b.inp_gate, &rq, &rd, &resid, t)?
13787 } else {
13788 e.add(&sn, &attn_out, &mut resid, t * n_embd)?;
13789 e.matmul(&e4b.inp_gate, &resid, t)?
13790 };
13791 let mut act = e.uninit(t * n_epl)?;
13792 let y = if t == 1 && e.uses_q8_1_fast(&e4b.proj) {
13793 let ipv = e.view(&inp_pl, n_epl * n_layer);
13794 let row = ipv.slice(il * n_epl..(il + 1) * n_epl);
13795 let (aq, ad) = e.gelu_tanh_mul_q8_1(&g, &row, &mut act, n_epl, 1)?;
13796 e.matmul_pre(&e4b.proj, &aq, &ad, &act, t)?
13797 } else {
13798 let mut inp_this = e.uninit(t * n_epl)?;
13799 e.copy_rows_strided(
13800 &inp_pl,
13801 &mut inp_this,
13802 n_epl,
13803 t,
13804 n_epl * n_layer,
13805 il * n_epl,
13806 )?;
13807 e.gelu_tanh_mul(&g, &inp_this, &mut act, t * n_epl)?;
13808 e.matmul(&e4b.proj, &act, t)?
13809 };
13810 let next_norm = if il + 1 < n_layer {
13813 self.layers[il + 1].attn_norm.float_data()
13814 } else {
13815 self.output_norm.float_data()
13816 };
13817 let mut xn = e.uninit(t * n_embd)?;
13818 let pair = e.rms_pre_add_scale_rms_norm_q8_1(
13819 &y,
13820 e4b.post_norm.float_data(),
13821 &resid,
13822 bits.layer_scale,
13823 next_norm,
13824 &mut xn,
13825 n_embd,
13826 t,
13827 eps,
13828 )?;
13829 h_carry = Some(pair);
13830 x = xn;
13831 }
13832 let (oq, odq) = h_carry.take().unwrap();
13836 let h0 = e.zeros(0)?;
13837 let hm = if head_last { 1 } else { t };
13838 let (hq, hd) = if head_last && t > 1 {
13839 let mut q1 = e.uninit_i8(n_embd)?;
13840 e.dtod_copy_view_i8(&oq.slice((t - 1) * n_embd..t * n_embd), &mut q1)?;
13841 let nb = n_embd / 32;
13842 let mut d1 = e.uninit(nb)?;
13843 e.dtod_copy_view(&odq.slice((t - 1) * nb..t * nb), &mut d1)?;
13844 (q1, d1)
13845 } else {
13846 (oq, odq)
13847 };
13848 let mut ld = e.matmul_pre(&self.output, &hq, &hd, &h0, hm)?;
13849 if cap_logits {
13853 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
13854 e.softcap(&mut ld, cap, hm * self.output.out_features())?;
13855 }
13856 self.gemma4_suppress(e, &mut ld, hm)?; Ok((ld, x))
13858 }
13859
13860 pub fn gemma4_e4b_decode_step_t_am_dev(
13867 &self,
13868 e: &Engine,
13869 tok_d: &CudaSlice<u32>,
13870 t: usize,
13871 pos0: usize,
13872 cache: &mut Cache,
13873 ) -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
13874 let n_embd = self.cfg.n_embd as usize;
13875 let eps = self.cfg.rms_eps;
13876 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
13877 let pos_d = e.htod_i32(&pos)?;
13878 let embd_gpu = self
13879 .embd_gpu
13880 .get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload"));
13881 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
13882 let mut x = e.embed_gather_device_td(embd_gpu, tok_d, t, n_embd, qt, rb)?;
13883 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
13884 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, tok_d, &x, t)?;
13885 let (ld, xp) =
13886 self.gemma4_e4b_trunk_core(e, x, inp_pl, &pos_d, t, cache, None, true, false)?;
13887 let n_vocab = self.output.out_features();
13890 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
13891 for i in 0..t {
13892 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
13893 }
13894 let mut hn = e.uninit(t * n_embd)?;
13895 e.rms_norm(&xp, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
13896 cache.pos += t;
13897 Ok((vam, hn))
13898 }
13899
13900 pub(crate) fn gemma4_e4b_decode_step_t_h(
13903 &self,
13904 e: &Engine,
13905 tokens: &[u32],
13906 pos0: usize,
13907 cache: &mut Cache,
13908 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
13909 let n_embd = self.cfg.n_embd as usize;
13910 let eps = self.cfg.rms_eps;
13911 let t = tokens.len();
13912 let (ld, xp) = self.gemma4_e4b_trunk(e, tokens, pos0, cache, false)?;
13913 let mut hn = e.uninit(t * n_embd)?;
13914 e.rms_norm(&xp, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
13915 cache.pos += t;
13916 Ok((e.dtoh(&ld)?, hn))
13917 }
13918
13919 pub fn gemma4_e4b_decode_step_dcg(
13925 &self,
13926 e: &Engine,
13927 token_d: &mut CudaSlice<u32>,
13928 pos_d: &mut CudaSlice<i32>,
13929 embd_gpu: &CudaSlice<u8>,
13930 embd_qt: i32,
13931 embd_rb: usize,
13932 cache: &mut Cache,
13933 n_vocab: usize,
13934 bucket: usize,
13935 ) -> Result<(), Box<dyn std::error::Error>> {
13936 let n_embd = self.cfg.n_embd as usize;
13937 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
13938 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
13939 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, token_d, &x, 1)?;
13940 let (ld, _x) =
13941 self.gemma4_e4b_trunk_core(e, x, inp_pl, pos_d, 1, cache, Some(bucket), false, false)?;
13942 e.argmax_token_device_into(&ld, token_d, n_vocab)?;
13943 e.inc_seqlen(pos_d)?;
13944 Ok(())
13945 }
13946
13947 #[allow(clippy::too_many_arguments)]
13955 pub fn gemma4_e4b_decode_step_dc(
13956 &self,
13957 e: &Engine,
13958 token_d: &CudaSlice<u32>,
13959 pos_d: &mut CudaSlice<i32>,
13960 embd_gpu: &CudaSlice<u8>,
13961 embd_qt: i32,
13962 embd_rb: usize,
13963 cache: &mut Cache,
13964 n_vocab: usize,
13965 ) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
13966 let n_embd = self.cfg.n_embd as usize;
13967 let eps = self.cfg.rms_eps;
13968 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
13969 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
13970 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, token_d, &x, 1)?;
13971 let (ld, _x) =
13972 self.gemma4_e4b_trunk_core(e, x, inp_pl, pos_d, 1, cache, None, false, false)?;
13973 let mut tok_out = e.stream().alloc_zeros::<u32>(1)?;
13974 e.argmax_token_device_into(&ld, &mut tok_out, n_vocab)?;
13975 e.inc_seqlen(pos_d)?;
13976 cache.pos += 1;
13977 let _ = eps;
13978 Ok(tok_out)
13979 }
13980
13981 pub(crate) fn gemma4_e4b_decode_step_h(
13984 &self,
13985 e: &Engine,
13986 token: u32,
13987 cache: &mut Cache,
13988 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
13989 let (ld, x) = self.gemma4_e4b_trunk(e, &[token], cache.pos, cache, false)?;
13990 let logits = e.dtoh(&ld)?;
13991 cache.pos += 1;
13992 Ok((logits, x))
13993 }
13994
13995 pub(crate) fn gemma4_e4b_prime(
13999 &self,
14000 e: &Engine,
14001 tokens: &[u32],
14002 cache: &mut Cache,
14003 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
14004 if cache.pos != 0 {
14007 return Err(
14008 "e4b prime is fresh-prompt only (v0) — prime the full prompt in one \
14009 call or decode tokenwise"
14010 .into(),
14011 );
14012 }
14013 let n_embd = self.cfg.n_embd as usize;
14014 let t = tokens.len();
14015 let (ld, x) = self.gemma4_e4b_trunk(e, tokens, 0, cache, true)?;
14016 cache.pos += t;
14017 let last = e.dtoh(&ld)?; let xv = e.view(&x, t * n_embd);
14019 let row = xv.slice((t - 1) * n_embd..t * n_embd);
14020 let mut h_seed = e.uninit(n_embd)?;
14021 e.copy_view_into(&mut h_seed, 0, &row, n_embd)?;
14022 Ok((last, h_seed, x))
14023 }
14024
14025 pub(crate) fn gemma4_e4b_forward(
14027 &self,
14028 e: &Engine,
14029 tokens: &[u32],
14030 last_only: bool,
14031 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
14032 let mut cache = Cache::new(e, &self.cfg, tokens.len() + 8)?;
14033 let (ld, _x) = self.gemma4_e4b_trunk(e, tokens, 0, &mut cache, last_only)?;
14034 Ok(e.dtoh(&ld)?) }
14036}
14037
14038#[cfg(test)]
14039mod prime_chunk_schedule_tests {
14040 use super::{
14041 PRIME_MIN_T, PRIME_PIPE_MIN_CHUNK, dynamic_prime_chunk_ranges, fixed_prime_chunk_ranges,
14042 fixed_prime_chunk_ranges_for_ring,
14043 };
14044
14045 fn sizes(ranges: &[(usize, usize)]) -> Vec<usize> {
14046 ranges.iter().map(|(start, end)| end - start).collect()
14047 }
14048
14049 fn auto_chunk(t: usize) -> usize {
14050 t.div_ceil(8).max(PRIME_PIPE_MIN_CHUNK).min(4096)
14051 }
14052
14053 #[test]
14054 fn fixed_schedule_retains_measured_geometry() {
14055 assert_eq!(
14056 sizes(&fixed_prime_chunk_ranges(461, 128)),
14057 vec![128, 128, 128, 77]
14058 );
14059 assert_eq!(
14060 sizes(&fixed_prime_chunk_ranges(1833, 230)),
14061 vec![230, 230, 230, 230, 230, 230, 230, 223]
14062 );
14063 assert_eq!(sizes(&fixed_prime_chunk_ranges(4096, 512)), vec![512; 8]);
14064 let capped = sizes(&fixed_prime_chunk_ranges_for_ring(8200, 4096, true));
14065 assert_eq!(capped, vec![4096, 4088, 16]);
14066 assert!(capped.iter().all(|&rows| rows <= 4096));
14067 assert_eq!(
14068 sizes(&fixed_prime_chunk_ranges_for_ring(4100, 4096, false)),
14069 vec![4100],
14070 "flag-off schedule remains byte-for-byte the legacy monolithic tail",
14071 );
14072 }
14073
14074 #[test]
14075 fn dynamic_schedule_matches_registered_shapes() {
14076 let cases = [
14077 (461, vec![64, 141, 132, 124]),
14078 (1833, vec![115, 269, 260, 252, 244, 237, 231, 225]),
14079 (4096, vec![256, 602, 580, 563, 545, 531, 516, 503]),
14080 ];
14081 for (t, expected) in cases {
14082 let chunk = auto_chunk(t);
14083 let fixed = fixed_prime_chunk_ranges(t, chunk);
14084 assert_eq!(
14085 sizes(&dynamic_prime_chunk_ranges(t, chunk, &fixed)),
14086 expected
14087 );
14088 }
14089 }
14090
14091 #[test]
14092 fn dynamic_schedule_covers_exactly_and_shrinks_after_fill() {
14093 for t in 256..=8192 {
14094 let chunk = auto_chunk(t);
14095 let fixed = fixed_prime_chunk_ranges(t, chunk);
14096 let dynamic = dynamic_prime_chunk_ranges(t, chunk, &fixed);
14097 assert_eq!(dynamic.len(), fixed.len(), "T={t}");
14098 assert_eq!(dynamic.first().unwrap().0, 0, "T={t}");
14099 assert_eq!(dynamic.last().unwrap().1, t, "T={t}");
14100 for pair in dynamic.windows(2) {
14101 assert_eq!(pair[0].1, pair[1].0, "T={t}");
14102 }
14103 assert!(
14104 dynamic
14105 .iter()
14106 .all(|(start, end)| end - start >= PRIME_MIN_T),
14107 "T={t} sizes={:?}",
14108 sizes(&dynamic)
14109 );
14110 if dynamic.len() >= 3 {
14111 let chunk_sizes = sizes(&dynamic);
14112 assert!(
14113 chunk_sizes[0] < chunk_sizes[1],
14114 "T={t} sizes={chunk_sizes:?}"
14115 );
14116 assert!(
14117 chunk_sizes[1..].windows(2).all(|pair| pair[0] >= pair[1]),
14118 "T={t} sizes={chunk_sizes:?}"
14119 );
14120 }
14121 }
14122 }
14123}
14124
14125#[cfg(test)]
14126mod page_prefetch_tests {
14127 use super::{
14128 grouped_worker_prefetch_position, page_prefetch_positions,
14129 page_prefetch_window_from_values, worker_prefetch_positions,
14130 };
14131
14132 #[test]
14133 fn page_prefetch_window_keeps_existing_opt_in_default() {
14134 assert_eq!(page_prefetch_window_from_values(false, None), 0);
14135 assert_eq!(page_prefetch_window_from_values(false, Some("8")), 0);
14136 assert_eq!(page_prefetch_window_from_values(true, None), 1);
14137 assert_eq!(page_prefetch_window_from_values(true, Some("bad")), 1);
14138 assert_eq!(page_prefetch_window_from_values(true, Some("0")), 0);
14139 assert_eq!(page_prefetch_window_from_values(true, Some("8")), 8);
14140 }
14141
14142 #[test]
14143 fn rolling_page_prefetch_advises_each_future_expert_once() {
14144 let advised: Vec<_> = (0..7)
14145 .flat_map(|position| page_prefetch_positions(position, 7, 3))
14146 .collect();
14147 assert_eq!(advised, vec![1, 2, 3, 4, 5, 6]);
14148
14149 let one_ahead: Vec<_> = (0..4)
14150 .flat_map(|position| page_prefetch_positions(position, 4, 1))
14151 .collect();
14152 assert_eq!(one_ahead, vec![1, 2, 3]);
14153 assert!(page_prefetch_positions(0, 4, 0).is_empty());
14154 }
14155
14156 #[test]
14157 fn grouped_worker_prefetch_primes_first_then_each_known_next_once() {
14158 assert_eq!(grouped_worker_prefetch_position(0, None), None);
14159 let positions: Vec<_> = std::iter::once(grouped_worker_prefetch_position(4, None).unwrap())
14160 .chain(
14161 (0..4).filter_map(|position| grouped_worker_prefetch_position(4, Some(position))),
14162 )
14163 .collect();
14164 assert_eq!(positions, vec![0, 1, 2, 3]);
14165 assert_eq!(grouped_worker_prefetch_position(1, Some(0)), None);
14166 }
14167
14168 #[test]
14169 fn rolling_worker_prefetch_primes_current_and_each_future_expert_once() {
14170 let queued: Vec<_> = (0..8)
14171 .flat_map(|position| worker_prefetch_positions(position, 8, 5))
14172 .collect();
14173 assert_eq!(queued, (0..8).collect::<Vec<_>>());
14174
14175 let one_at_a_time: Vec<_> = (0..4)
14176 .flat_map(|position| worker_prefetch_positions(position, 4, 1))
14177 .collect();
14178 assert_eq!(one_at_a_time, vec![0, 1, 2, 3]);
14179 assert!(worker_prefetch_positions(0, 4, 0).is_empty());
14180 }
14181}
14182
14183pub struct G4DcSlots {
14184 x: CudaSlice<f32>,
14185 xn: CudaSlice<f32>,
14186 cur: CudaSlice<f32>,
14187 hq: CudaSlice<i8>,
14188 hd_: CudaSlice<f32>,
14189 q0: CudaSlice<f32>,
14190 k0: CudaSlice<f32>,
14191 v0: CudaSlice<f32>,
14192 q: CudaSlice<f32>,
14193 k: CudaSlice<f32>,
14194 v: CudaSlice<f32>,
14195 attn: CudaSlice<f32>,
14196 o: CudaSlice<f32>,
14197 attn_out: CudaSlice<f32>,
14198 zsh: CudaSlice<f32>,
14199 zq: CudaSlice<i8>,
14200 zd: CudaSlice<f32>,
14201 gate: CudaSlice<f32>,
14202 up: CudaSlice<f32>,
14203 act: CudaSlice<f32>,
14204 actq: CudaSlice<i8>,
14205 actd: CudaSlice<f32>,
14206 f0: CudaSlice<f32>,
14207 sn: CudaSlice<f32>,
14208 hn: CudaSlice<f32>,
14209 logits: CudaSlice<f32>,
14210}