1use crate::Engine;
6use crate::cache::Cache;
7use cudarc::driver::CudaSlice;
8use memra_gguf::config::ModelConfig;
9
10pub struct PrimeSlabs {
14 pub t_cap: usize,
15 pub h: CudaSlice<f32>,
16 pub x1: CudaSlice<f32>,
17 pub z: CudaSlice<f32>,
18 pub act: CudaSlice<f32>,
19 pub xa: CudaSlice<f32>,
20 pub xb: CudaSlice<f32>,
21 pub h16: CudaSlice<u8>,
22 pub z16: CudaSlice<u8>,
23 pub gate: CudaSlice<f32>, pub up: CudaSlice<f32>, pub ffn_out: CudaSlice<f32>, pub seg_glue: Vec<Option<cudarc::driver::CudaGraph>>,
33 pub mixed: CudaSlice<f32>,
37 pub seg_mid: Vec<Option<cudarc::driver::CudaGraph>>,
38 pub seg_t: usize,
39}
40
41unsafe impl Send for PrimeSlabs {}
44
45fn shexp_gate_up_t1(
54 e: &Engine,
55 gate_shexp: &crate::model::GpuTensor,
56 up_shexp: &crate::model::GpuTensor,
57 z: &CudaSlice<f32>,
58 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
59) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
60 let is_nvfp4 = |w: &crate::model::GpuTensor| matches!(w, crate::model::GpuTensor::Quant { qtype, .. } if *qtype == crate::QT_NVFP4);
61 if is_nvfp4(gate_shexp) && is_nvfp4(up_shexp) {
62 let pair = match zq8 {
66 Some((zq, zd)) => e.matmul_nvfp4_fused2(gate_shexp, up_shexp, zq, zd, 1)?,
67 None => {
68 let (zq, zd) = e.quantize_q8_1(z, 1, gate_shexp.in_features())?;
69 e.matmul_nvfp4_fused2(gate_shexp, up_shexp, &zq, &zd, 1)?
70 }
71 };
72 if let Some(pair) = pair {
73 return Ok(pair);
74 }
75 }
76 match e.matmul_q8_fused2_x(gate_shexp, up_shexp, z)? {
77 Some(pair) => Ok(pair),
78 None => Ok((e.matmul(gate_shexp, z, 1)?, e.matmul(up_shexp, z, 1)?)),
79 }
80}
81
82fn active_matrix_values(
83 available: usize,
84 rows: usize,
85 columns: usize,
86 label: &str,
87) -> Result<usize, String> {
88 let required = rows
89 .checked_mul(columns)
90 .ok_or_else(|| format!("{label} shape overflows: {rows}x{columns}"))?;
91 if available < required {
92 return Err(format!(
93 "{label} has {available} values, fewer than the active {rows}x{columns} ({required})"
94 ));
95 }
96 Ok(required)
97}
98
99fn step_grouped_decode_shape(prefill: bool, tokens: usize) -> bool {
100 !prefill && tokens == 1
101}
102
103fn parse_step_ep_grouped_prefill(value: Option<&str>) -> Result<bool, String> {
104 match value {
105 None | Some("") | Some("0") => Ok(false),
106 Some("1") => Ok(true),
107 Some(value) => Err(format!(
108 "MEMRA_STEP_EP_GROUPED_PREFILL={value:?} is invalid; expected 0 or 1"
109 )),
110 }
111}
112
113fn step_ep_grouped_prefill_enabled() -> Result<bool, String> {
114 parse_step_ep_grouped_prefill(
115 std::env::var("MEMRA_STEP_EP_GROUPED_PREFILL")
116 .ok()
117 .as_deref(),
118 )
119}
120
121fn step_grouped_prefill_shape(enabled: bool, prefill: bool, tokens: usize) -> bool {
122 enabled && prefill && (PRIME_MIN_T..=crate::cache::PRIME_CHUNK_MAX_TOKENS).contains(&tokens)
123}
124
125fn parse_step_tp_prefill(value: Option<&str>) -> Result<bool, String> {
126 match value {
127 None | Some("") | Some("0") => Ok(false),
128 Some("1") => Ok(true),
129 Some(value) => Err(format!(
130 "MEMRA_STEP_TP_PREFILL={value:?} is invalid; expected 0 or 1"
131 )),
132 }
133}
134
135fn step_tp_prefill_enabled() -> Result<bool, String> {
136 parse_step_tp_prefill(std::env::var("MEMRA_STEP_TP_PREFILL").ok().as_deref())
137}
138
139fn validate_step_prime_batch_modes(tp_prefill: bool, grouped_prefill: bool) -> Result<(), String> {
140 if grouped_prefill && !tp_prefill {
141 return Err("MEMRA_STEP_EP_GROUPED_PREFILL=1 requires MEMRA_STEP_TP_PREFILL=1".into());
142 }
143 if tp_prefill {
144 return Err(
145 "Step TP4 cross-request prime batching did not clear the live-server performance \
146 gate; use per-session grouped prefill"
147 .into(),
148 );
149 }
150 Ok(())
151}
152
153fn step_tp_prefill_shape(
154 enabled: bool,
155 tokens: usize,
156 ranks: usize,
157 native_p2p: bool,
158 has_rank_local_attention: bool,
159 fp8_kv: bool,
160) -> bool {
161 enabled
169 && tokens >= PRIME_MIN_T
170 && matches!(ranks, 2 | 4)
171 && native_p2p
172 && has_rank_local_attention
173 && !fp8_kv
174}
175
176fn empty_cache_layers<T>(n: usize) -> Vec<Option<T>> {
177 std::iter::repeat_with(|| None).take(n).collect()
178}
179
180struct PrimeCacheStages<'a> {
185 parent: &'a mut Cache,
186 cut: usize,
187 stage0: Cache,
188 stage1: Cache,
189}
190
191impl<'a> PrimeCacheStages<'a> {
192 fn new(parent: &'a mut Cache, cut: usize) -> Self {
193 let n = parent.kv.len();
194 assert_eq!(parent.recur.len(), n, "cache layer vectors disagree");
195 assert!(cut <= n, "PP-2 cache cut {cut} exceeds {n} layers");
196 let mut kv0 = empty_cache_layers(n);
197 let mut kv1 = empty_cache_layers(n);
198 let mut tp_kv0 = empty_cache_layers(n);
199 let mut tp_kv1 = empty_cache_layers(n);
200 let mut recur0 = empty_cache_layers(n);
201 let mut recur1 = empty_cache_layers(n);
202 for i in 0..cut {
203 kv0[i] = parent.kv[i].take();
204 tp_kv0[i] = parent.tp_kv[i].take();
205 recur0[i] = parent.recur[i].take();
206 }
207 for i in cut..n {
208 kv1[i] = parent.kv[i].take();
209 tp_kv1[i] = parent.tp_kv[i].take();
210 recur1[i] = parent.recur[i].take();
211 }
212 let pos = parent.pos;
213 let max_ctx = parent.max_ctx;
214 Self {
215 parent,
216 cut,
217 stage0: Cache {
218 kv: kv0,
219 tp_kv: tp_kv0,
220 recur: recur0,
221 pos,
222 max_ctx,
223 last_logits_dev: None,
224 dflash_taps: None,
225 },
226 stage1: Cache {
227 kv: kv1,
228 tp_kv: tp_kv1,
229 recur: recur1,
230 pos,
231 max_ctx,
232 last_logits_dev: None,
233 dflash_taps: None,
234 },
235 }
236 }
237
238 fn parts(&mut self) -> (&mut Cache, &mut Cache) {
239 (&mut self.stage0, &mut self.stage1)
240 }
241}
242
243impl Drop for PrimeCacheStages<'_> {
244 fn drop(&mut self) {
245 let n = self.parent.kv.len();
246 for i in 0..n {
247 let source = if i < self.cut {
248 &mut self.stage0
249 } else {
250 &mut self.stage1
251 };
252 debug_assert!(self.parent.kv[i].is_none());
253 debug_assert!(self.parent.tp_kv[i].is_none());
254 debug_assert!(self.parent.recur[i].is_none());
255 self.parent.kv[i] = source.kv[i].take();
256 self.parent.tp_kv[i] = source.tp_kv[i].take();
257 self.parent.recur[i] = source.recur[i].take();
258 }
259 self.parent.pos = self.stage0.pos.min(self.stage1.pos);
260 }
261}
262
263pub(crate) struct AttnPre {
265 pub q: cudarc::driver::CudaSlice<f32>,
266 pub k: cudarc::driver::CudaSlice<f32>,
267 pub v: cudarc::driver::CudaSlice<f32>,
268 pub gate: Option<cudarc::driver::CudaSlice<f32>>,
269}
270
271pub(crate) struct GdnPrep {
273 pub hk: usize,
274 pub q_l2: cudarc::driver::CudaSlice<f32>,
275 pub k_l2: cudarc::driver::CudaSlice<f32>,
276 pub v_g: cudarc::driver::CudaSlice<f32>,
277 pub beta: cudarc::driver::CudaSlice<f32>,
278 pub g_log: cudarc::driver::CudaSlice<f32>,
279 pub kb16: Option<cudarc::driver::CudaSlice<u8>>,
280 pub qb16: Option<cudarc::driver::CudaSlice<u8>>,
281}
282
283pub(crate) struct VerifyStreamScratch {
285 pub pos_d: CudaSlice<i32>,
286 pub row_ctrs: Vec<CudaSlice<i32>>,
287}
288use crate::hybrid::{FullAttnLayer, HybridModel, LinearAttnLayer, Mixer, MoeWeights};
289
290struct MoeInputTraceWriter {
291 dir: std::path::PathBuf,
292 index: std::fs::File,
293 payloads: std::collections::HashMap<u16, (std::fs::File, u64)>,
294}
295
296static MOE_INPUT_TRACE_WRITER: std::sync::OnceLock<std::sync::Mutex<Option<MoeInputTraceWriter>>> =
297 std::sync::OnceLock::new();
298
299fn gdec_enabled() -> bool {
302 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
303 *E.get_or_init(|| {
304 std::env::var("MEMRA_MOE_GDEC")
305 .map(|v| v != "0")
306 .unwrap_or(true)
307 })
308}
309
310fn moe_slab_enabled() -> bool {
321 std::env::var("MEMRA_MOE_SLAB").as_deref() != Ok("0")
322}
323
324fn moe_grouped_enabled(_cfg: &ModelConfig, _prefill: bool) -> bool {
328 std::env::var("MEMRA_MOE_GROUPED")
329 .map(|value| value != "0")
330 .unwrap_or(false)
331}
332
333fn moe_prefetch_enabled() -> bool {
336 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
337 *E.get_or_init(|| {
338 std::env::var("MEMRA_MOE_PREFETCH").as_deref() == Ok("1")
339 || crate::spill_pread::worker_enabled()
340 })
341}
342
343fn moe_page_prefetch_window() -> usize {
348 static W: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
349 *W.get_or_init(|| {
350 page_prefetch_window_from_values(
351 std::env::var("MEMRA_MOE_PAGE_PREFETCH").as_deref() == Ok("1"),
352 std::env::var("MEMRA_MOE_PAGE_PREFETCH_WINDOW")
353 .ok()
354 .as_deref(),
355 )
356 })
357}
358
359fn page_prefetch_window_from_values(enabled: bool, raw_window: Option<&str>) -> usize {
360 if !enabled {
361 return 0;
362 }
363 raw_window.and_then(|value| value.parse().ok()).unwrap_or(1)
364}
365
366fn page_prefetch_positions(position: usize, len: usize, window: usize) -> std::ops::Range<usize> {
370 if window == 0 || position >= len {
371 return len..len;
372 }
373 let (start, count) = if position == 0 {
374 (1, window)
375 } else {
376 (position.saturating_add(window), 1)
377 };
378 let start = start.min(len);
379 start..start.saturating_add(count).min(len)
380}
381
382fn grouped_worker_prefetch_position(order_len: usize, current: Option<usize>) -> Option<usize> {
385 let position = current.map_or(0, |position| position.saturating_add(1));
386 (position < order_len).then_some(position)
387}
388
389fn worker_prefetch_window() -> usize {
394 static WINDOW: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
395 *WINDOW.get_or_init(|| {
396 let automatic = crate::spill_pread::configured_depth().saturating_sub(1) / 3;
397 std::env::var("MEMRA_SPILL_WORKER_EXPERT_WINDOW")
398 .ok()
399 .and_then(|value| value.parse::<usize>().ok())
400 .unwrap_or(automatic.max(1))
401 })
402}
403
404fn worker_prefetch_positions(position: usize, len: usize, window: usize) -> std::ops::Range<usize> {
408 if window == 0 || position >= len {
409 return len..len;
410 }
411 let (start, count) = if position == 0 {
412 (0, window)
413 } else {
414 (position.saturating_add(window).saturating_sub(1), 1)
415 };
416 let start = start.min(len);
417 start..start.saturating_add(count).min(len)
418}
419
420fn moe_dev_enabled() -> bool {
425 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
426 *E.get_or_init(|| {
427 std::env::var("MEMRA_MOE_DEV")
428 .map(|v| v != "0")
429 .unwrap_or(true)
430 && !matches!(std::env::var("MEMRA_FUSED_ROUTER").as_deref(), Ok("0"))
431 })
432}
433
434fn sigmoid_router_enabled() -> bool {
437 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
438 *E.get_or_init(|| {
439 std::env::var("MEMRA_SIG_ROUTER")
440 .map(|v| v != "0")
441 .unwrap_or(true)
442 })
443}
444
445fn moe_q8_enabled() -> bool {
450 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
451 *E.get_or_init(|| {
452 std::env::var("MEMRA_MOE_Q8")
453 .map(|v| v != "0")
454 .unwrap_or(true)
455 })
456}
457
458fn expert_dp4a_supported(qt: i32) -> bool {
461 qt == crate::QT_Q4_0
462 || qt == crate::QT_IQ3_S
463 || qt == crate::QT_IQ4_XS
464 || qt == crate::QT_Q3_K
465 || qt == crate::QT_Q4_K
466 || qt == crate::QT_Q6_K
467}
468
469fn q8_expert_supported(qt: i32) -> bool {
470 static KQ: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
476 let kq = *KQ.get_or_init(|| {
477 std::env::var("MEMRA_MOE_Q8_KQ")
478 .map(|v| v != "0")
479 .unwrap_or(true)
480 });
481 let nvfp4_q8 = std::env::var("MEMRA_MOE_Q8_NVFP4")
488 .map(|v| v != "0")
489 .unwrap_or(true);
490 qt == crate::QT_IQ3_S
491 || qt == crate::QT_IQ4_XS
492 || (nvfp4_q8 && qt == crate::QT_NVFP4)
493 || (kq && (qt == crate::QT_Q3_K || qt == crate::QT_Q4_K || qt == crate::QT_Q6_K))
494}
495
496fn q8_expert_dec_supported(qt: i32) -> bool {
499 qt == crate::QT_IQ3_S || qt == crate::QT_IQ4_XS || qt == crate::QT_Q4_0
500}
501
502fn f16g_proj_ok(qt: i32, in_f: usize) -> bool {
508 match qt {
509 crate::QT_Q4_0 => in_f % 32 == 0,
510 crate::QT_IQ4_XS | crate::QT_IQ3_S | crate::QT_Q3_K | crate::QT_Q4_K | crate::QT_Q6_K => {
511 in_f % 256 == 0
512 }
513 crate::QT_NVFP4 => in_f % 64 == 0,
518 _ => false,
519 }
520}
521
522fn moe_prewarm_enabled() -> bool {
525 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
526 *E.get_or_init(|| {
527 std::env::var("MEMRA_MOE_PREWARM")
528 .map(|v| v != "0")
529 .unwrap_or(true)
530 })
531}
532
533fn cpu_expert_profile_admit_enabled() -> bool {
537 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
538 *E.get_or_init(|| std::env::var("MEMRA_CPU_EXPERT_FREEZE_PROFILE_ADMIT").as_deref() == Ok("1"))
539}
540
541pub const PRIME_MIN_T: usize = 16;
545
546fn step_gemm_prime_suffix_on() -> bool {
575 std::env::var("MEMRA_STEP_GEMM_PRIME_SUFFIX").as_deref() != Ok("0")
576}
577
578const MOE_DEV_MAX_T: usize = 16;
587const PRIME_PIPE_MICROBATCHES: usize = 8;
588const PRIME_PIPE_MIN_CHUNK: usize = 128;
589const PRIME_PIPE_EDGE_MIN_CHUNK: usize = 64;
590const PRIME_PIPE_LINEAR_WORK: usize = 8;
591
592fn prime_pp2_auto_geometry(n_layers: usize) -> bool {
593 crate::pp::prime_pp_on()
594 && !crate::pp::pp2_streams_off()
595 && crate::pp::pp_cuts(n_layers).is_some_and(|cuts| cuts.len() == 3)
596}
597
598pub fn prime_chunk_tokens(t: usize, n_layers: usize) -> usize {
602 if let Ok(value) = std::env::var("MEMRA_PRIME_CHUNK") {
603 let parsed = value
604 .parse::<usize>()
605 .unwrap_or(crate::cache::PRIME_CHUNK_MAX_TOKENS);
606 return if crate::cache::swa_ring_on() {
607 if parsed == 0 {
608 crate::cache::PRIME_CHUNK_MAX_TOKENS
609 } else {
610 parsed.min(crate::cache::PRIME_CHUNK_MAX_TOKENS)
611 }
612 } else {
613 parsed
614 };
615 }
616 let chunk = crate::cache::PRIME_CHUNK_MAX_TOKENS;
617 if prime_pp2_auto_geometry(n_layers) && t >= 2 * PRIME_PIPE_MIN_CHUNK {
618 chunk.min(
619 t.div_ceil(PRIME_PIPE_MICROBATCHES)
620 .max(PRIME_PIPE_MIN_CHUNK),
621 )
622 } else {
623 chunk
624 }
625}
626
627fn fixed_prime_chunk_ranges(t: usize, chunk: usize) -> Vec<(usize, usize)> {
628 fixed_prime_chunk_ranges_for_ring(t, chunk, crate::cache::swa_ring_on())
629}
630
631fn fixed_prime_chunk_ranges_for_ring(t: usize, chunk: usize, ring_on: bool) -> Vec<(usize, usize)> {
632 if chunk == 0 || t <= chunk {
633 return vec![(0, t)];
634 }
635 let mut ranges = Vec::with_capacity(t.div_ceil(chunk));
636 let mut start = 0usize;
637 while start < t {
638 let mut end = (start + chunk).min(t);
639 if t - end > 0 && t - end < PRIME_MIN_T {
640 if ring_on {
641 let shifted = t - PRIME_MIN_T;
642 end = if shifted > start { shifted } else { t };
643 } else {
644 end = t;
645 }
646 }
647 ranges.push((start, end));
648 start = end;
649 }
650 ranges
651}
652
653fn prime_chunk_work(prefix: usize, total: usize) -> u128 {
654 let prefix = prefix as u128;
655 prefix * (prefix + (PRIME_PIPE_LINEAR_WORK as u128) * (total as u128))
656}
657
658fn dynamic_prime_chunk_ranges(
659 t: usize,
660 fixed_chunk: usize,
661 fixed: &[(usize, usize)],
662) -> Vec<(usize, usize)> {
663 let n = fixed.len();
664 if n < 3 {
665 return fixed.to_vec();
666 }
667
668 let max_first = t - (n - 1) * PRIME_MIN_T;
669 let first = fixed_chunk
670 .div_ceil(2)
671 .max(PRIME_PIPE_EDGE_MIN_CHUNK)
672 .min(max_first);
673 let mut ranges = Vec::with_capacity(n);
674 ranges.push((0, first));
675
676 let first_work = prime_chunk_work(first, t);
677 let work_span = prime_chunk_work(t, t) - first_work;
678 let denominator = (n - 1) as u128;
679 let mut previous = first;
680 for boundary in 1..n - 1 {
681 let target = first_work * denominator + work_span * (boundary as u128);
682 let remaining = n - 1 - boundary;
683 let mut low = previous + PRIME_MIN_T;
684 let mut high = t - remaining * PRIME_MIN_T;
685 while low < high {
686 let mid = low + (high - low) / 2;
687 if prime_chunk_work(mid, t) * denominator >= target {
688 high = mid;
689 } else {
690 low = mid + 1;
691 }
692 }
693 ranges.push((previous, low));
694 previous = low;
695 }
696 ranges.push((previous, t));
697 ranges
698}
699
700pub fn prime_chunk_ranges(t: usize, n_layers: usize, gdn_grid: bool) -> Vec<(usize, usize)> {
710 let explicit_chunk = std::env::var_os("MEMRA_PRIME_CHUNK").is_some();
711 let chunk = prime_chunk_tokens(t, n_layers);
712 let fixed = fixed_prime_chunk_ranges(t, chunk);
713 let dynamic = match std::env::var("MEMRA_PRIME_CHUNK_SCHED") {
714 Ok(value) => value == "dynamic",
715 Err(_) => true,
716 };
717 if explicit_chunk {
718 return fixed;
719 }
720 let ranges = if !dynamic || !prime_pp2_auto_geometry(n_layers) {
721 fixed
722 } else {
723 dynamic_prime_chunk_ranges(t, chunk, &fixed)
724 };
725 if gdn_grid && std::env::var("MEMRA_PRIME_GRID_ALIGN").as_deref() != Ok("0") {
729 align_prime_ranges_to_gdn(&ranges, t, Engine::gdn_chunk_size())
730 } else {
731 ranges
732 }
733}
734
735pub fn align_prime_ranges_to_gdn(
755 ranges: &[(usize, usize)],
756 t: usize,
757 c: usize,
758) -> Vec<(usize, usize)> {
759 if c == 0 || ranges.len() < 2 {
760 return ranges.to_vec();
761 }
762 let mut out: Vec<(usize, usize)> = Vec::with_capacity(ranges.len());
763 let mut start = 0usize;
764 for (i, &(_, end)) in ranges.iter().enumerate() {
765 let e = if i + 1 == ranges.len() {
766 t
767 } else {
768 end / c * c
769 };
770 if e > start {
771 out.push((start, e));
772 start = e;
773 } }
775 debug_assert_eq!(out.last().map(|&(_, e)| e), Some(t));
776 out
777}
778
779struct HeadSplit {
780 pin: u64,
781 w1: CudaSlice<u8>,
782 hn1: CudaSlice<f32>,
783 y1: CudaSlice<f32>,
784 logits_e: CudaSlice<f32>,
785 ev_hn: cudarc::driver::CudaEvent,
786 ev_done: cudarc::driver::CudaEvent,
787 raw_hn1: u64,
788 raw_y1: u64,
789 raw_logits_hi: u64,
790 samp: Option<SampScratch>,
795}
796
797struct SampScratch {
798 pb: CudaSlice<f32>,
799 th: CudaSlice<f32>,
800 z: CudaSlice<f32>,
801 mx: CudaSlice<f32>,
802 rows: CudaSlice<i32>,
803}
804static HEAD_SPLIT_WS: std::sync::Mutex<Option<HeadSplit>> = std::sync::Mutex::new(None);
806
807#[allow(clippy::type_complexity)]
811static DEV1_ROUTER_REPS: std::sync::Mutex<
812 Option<(
813 std::collections::HashMap<u16, (CudaSlice<f32>, CudaSlice<f32>, CudaSlice<u8>)>,
814 Option<CudaSlice<f32>>,
815 )>,
816> = std::sync::Mutex::new(None);
817
818#[allow(clippy::type_complexity)]
822static SHEXP_D1_REPS: std::sync::Mutex<
823 Option<std::collections::HashMap<u16, (CudaSlice<u8>, CudaSlice<u8>, CudaSlice<u8>)>>,
824> = std::sync::Mutex::new(None);
825#[allow(clippy::type_complexity)]
826static SHEXP_D1_WS: std::sync::Mutex<
827 Option<(
828 (usize, usize),
829 CudaSlice<f32>,
830 CudaSlice<f32>,
831 CudaSlice<f32>,
832 cudarc::driver::CudaEvent,
833 cudarc::driver::CudaEvent,
834 )>,
835> = std::sync::Mutex::new(None);
836
837static SHEXP_OV_WS: std::sync::Mutex<
839 Option<(usize, usize, usize, CudaSlice<f32>, CudaSlice<f32>)>,
840> = std::sync::Mutex::new(None);
841
842impl HybridModel {
843 pub fn gdn_prime_grid_on(&self) -> bool {
849 Engine::gdn_chunked_enabled()
850 && self
851 .layers
852 .iter()
853 .any(|l| matches!(l.mixer, crate::hybrid::Mixer::Linear(_)))
854 }
855
856 fn step35_tp_device_resident(e: &Engine, tp: &crate::hybrid::StepTpQkv) -> bool {
861 tp.runtime.native_p2p() && tp.runtime.root_shares_ctx(e)
862 }
863
864 fn step35_tp_qkv(
865 &self,
866 e: &Engine,
867 fa: &FullAttnLayer,
868 h: &CudaSlice<f32>,
869 t: usize,
870 ) -> Result<Option<Vec<CudaSlice<f32>>>, Box<dyn std::error::Error>> {
871 let Some(tp) = fa.step_tp_qkv.as_ref() else {
872 return Ok(None);
873 };
874 let values = active_matrix_values(
875 h.len(),
876 t,
877 self.cfg.n_embd as usize,
878 "Step TP QKV activation",
879 )?;
880 if Self::step35_tp_device_resident(e, tp) {
889 e.stream().synchronize()?;
892 let q = tp
893 .runtime
894 .bf16_column_parallel_resident_native_device(&tp.q, h, t)?;
895 let k = tp
896 .runtime
897 .bf16_column_parallel_resident_native_device(&tp.k, h, t)?;
898 let v = tp
899 .runtime
900 .bf16_column_parallel_resident_native_device(&tp.v, h, t)?;
901 Self::step35_tp_log_once(tp, "qkv", "device-resident");
902 return Ok(Some(vec![q, k, v]));
903 }
904 let host = e.dtoh_view(&h.slice(0..values))?;
905 let q = if tp.runtime.native_p2p() {
906 tp.runtime
907 .bf16_column_parallel_resident_native(&tp.q, &host, t)?
908 } else {
909 tp.runtime
910 .bf16_column_parallel_resident(&tp.q, &host, t)?
911 .gathered
912 };
913 let k = if tp.runtime.native_p2p() {
914 tp.runtime
915 .bf16_column_parallel_resident_native(&tp.k, &host, t)?
916 } else {
917 tp.runtime
918 .bf16_column_parallel_resident(&tp.k, &host, t)?
919 .gathered
920 };
921 let v = if tp.runtime.native_p2p() {
922 tp.runtime
923 .bf16_column_parallel_resident_native(&tp.v, &host, t)?
924 } else {
925 tp.runtime
926 .bf16_column_parallel_resident(&tp.v, &host, t)?
927 .gathered
928 };
929 Self::step35_tp_log_once(tp, "qkv", "host-canonical");
930 Ok(Some(vec![e.htod(&q)?, e.htod(&k)?, e.htod(&v)?]))
931 }
932
933 fn step35_tp_log_once(tp: &crate::hybrid::StepTpQkv, proj: &str, activation: &'static str) {
937 use std::sync::atomic::{AtomicBool, Ordering};
938 static LOGGED: [AtomicBool; 4] = [
939 AtomicBool::new(false),
940 AtomicBool::new(false),
941 AtomicBool::new(false),
942 AtomicBool::new(false),
943 ];
944 let idx = 2 * usize::from(proj == "o") + usize::from(activation == "device-resident");
945 if LOGGED[idx].swap(true, Ordering::Relaxed) {
946 return;
947 }
948 eprintln!(
949 "[step-tp-{proj}] execute layer={} devices={:?} projections={proj} \
950 tensor_parallel=true attention_local=true kv_local=true transport={} \
951 native_p2p={} bulk_p2p={} activation={activation} \
952 output={} performance_claim=false (logged once per transport)",
953 tp.layer,
954 tp.devices,
955 tp.runtime.transport_label(),
956 tp.runtime.native_p2p(),
957 tp.runtime.bulk_p2p(),
958 if activation == "device-resident" {
959 "root-resident"
960 } else {
961 "root-readback"
962 },
963 );
964 }
965
966 fn step35_tp_o(
967 &self,
968 e: &Engine,
969 fa: &FullAttnLayer,
970 activation: &CudaSlice<f32>,
971 tokens: usize,
972 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
973 let Some(tp) = fa.step_tp_qkv.as_ref() else {
974 return Ok(None);
975 };
976 if Self::step35_tp_device_resident(e, tp) {
980 e.stream().synchronize()?; let output = tp
982 .runtime
983 .step_bf16_row_parallel_resident_native_device(&tp.o, activation, tokens)?;
984 Self::step35_tp_log_once(tp, "o", "device-resident");
985 return Ok(Some(output));
986 }
987 let host = e.dtoh(activation)?;
988 let output = if tp.runtime.native_p2p() {
989 tp.runtime
990 .step_bf16_row_parallel_resident_native(&tp.o, &host, tokens)?
991 } else {
992 tp.runtime
993 .step_bf16_row_parallel_resident(&tp.o, &host, tokens)?
994 };
995 Self::step35_tp_log_once(tp, "o", "host-canonical");
996 Ok(Some(e.htod(&output)?))
997 }
998
999 fn step35_o(
1000 &self,
1001 e: &Engine,
1002 fa: &FullAttnLayer,
1003 activation: &CudaSlice<f32>,
1004 tokens: usize,
1005 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1006 match self.step35_tp_o(e, fa, activation, tokens)? {
1007 Some(output) => Ok(output),
1008 None => e.matmul(&fa.wo, activation, tokens),
1009 }
1010 }
1011
1012 fn prime_trace_path() -> Option<&'static str> {
1017 static P: std::sync::OnceLock<Option<String>> = std::sync::OnceLock::new();
1018 P.get_or_init(|| std::env::var("MEMRA_PRIME_TRACE").ok())
1019 .as_deref()
1020 }
1021
1022 fn prime_anatomy_on() -> bool {
1028 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1029 *E.get_or_init(|| std::env::var("MEMRA_PRIME_ANATOMY").as_deref() == Ok("1"))
1030 }
1031
1032 fn prime_anatomy_slots() -> &'static [std::sync::atomic::AtomicU64; 5] {
1033 static S: [std::sync::atomic::AtomicU64; 5] = [
1034 std::sync::atomic::AtomicU64::new(0), std::sync::atomic::AtomicU64::new(0), std::sync::atomic::AtomicU64::new(0), std::sync::atomic::AtomicU64::new(0), std::sync::atomic::AtomicU64::new(0), ];
1040 &S
1041 }
1042
1043 pub fn forward(
1045 &self,
1046 e: &Engine,
1047 tokens: &[u32],
1048 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
1049 if self.is_gemma4_e4b() {
1050 return self.gemma4_e4b_forward(e, tokens, false);
1051 }
1052 if self.uses_gemma_program() {
1053 return self.gemma4_forward(e, tokens, false);
1054 }
1055 let cfg = &self.cfg;
1056 let n_embd = cfg.n_embd as usize;
1057 let t = tokens.len();
1058 let eps = cfg.rms_eps;
1059 let pos: Vec<i32> = (0..t as i32).collect();
1060 let pos_d = e.htod_i32(&pos)?;
1061
1062 let mut x = self.embed(e, tokens)?; for (il, layer) in self.layers.iter().enumerate() {
1065 let mut h = e.uninit(t * n_embd)?;
1067 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
1068
1069 let mixed = match &layer.mixer {
1070 Mixer::Full(fa) => self.full_attn(e, fa, &h, &pos_d, t, il)?,
1071 Mixer::Linear(la) => self.linear_attn(e, la, &h, t)?,
1072 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
1073 };
1074
1075 let mut x1 = e.uninit(t * n_embd)?;
1077 e.add(&x, &mixed, &mut x1, t * n_embd)?;
1078
1079 let mut z = e.uninit(t * n_embd)?;
1081 e.rms_norm(
1082 &x1,
1083 layer.post_attn_norm.float_data(),
1084 &mut z,
1085 n_embd,
1086 t,
1087 eps,
1088 )?;
1089 let ffn_out = match &layer.ffn {
1090 crate::hybrid::Ffn::Dense {
1091 ffn_gate,
1092 ffn_up,
1093 ffn_down,
1094 } => {
1095 let n_ff = ffn_gate.out_features();
1096 let mut g2 = e.matmul_group(&[ffn_gate, ffn_up], &z, t)?;
1097 let up = g2.pop().unwrap();
1098 let gate = g2.pop().unwrap();
1099 let mut act = e.uninit(t * n_ff)?;
1100 Self::ffn_act_lim(
1105 e,
1106 &self.cfg,
1107 &gate,
1108 &up,
1109 1.0,
1110 1.0,
1111 self.cfg.clamp_shexp_at(il as u32),
1112 &mut act,
1113 t * n_ff,
1114 )?;
1115 e.matmul(ffn_down, &act, t)?
1116 }
1117 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il_prefill(e, m, &z, t, il as u16)?,
1118 };
1119 let mut x2 = e.uninit(t * n_embd)?;
1120 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
1121 x = x2;
1122 }
1123
1124 let mut hn = e.uninit(t * n_embd)?;
1125 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
1126 let logits = e.matmul(&self.output, &hn, t)?;
1127 Ok(e.dtoh(&logits)?)
1128 }
1129
1130 pub fn forward_last(
1136 &self,
1137 e: &Engine,
1138 tokens: &[u32],
1139 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
1140 if self.uses_gemma_program() {
1141 return self.gemma4_forward(e, tokens, true);
1142 }
1143 let cfg = &self.cfg;
1144 let n_embd = cfg.n_embd as usize;
1145 let t = tokens.len();
1146 let eps = cfg.rms_eps;
1147 let pos: Vec<i32> = (0..t as i32).collect();
1148 let pos_d = e.htod_i32(&pos)?;
1149
1150 let mut x = self.embed(e, tokens)?; let probe = std::env::var("MEMRA_LAYER_PROBE").is_ok();
1154 let anat = Self::prime_anatomy_on();
1155 let mut anat_last = if anat {
1156 e.stream().synchronize()?;
1157 Some(std::time::Instant::now())
1158 } else {
1159 None
1160 };
1161 macro_rules! anat_mark {
1162 ($slot:expr) => {
1163 if let Some(ts) = anat_last.as_mut() {
1164 e.stream().synchronize()?;
1165 Self::prime_anatomy_slots()[$slot].fetch_add(
1166 ts.elapsed().as_nanos() as u64,
1167 std::sync::atomic::Ordering::Relaxed,
1168 );
1169 *ts = std::time::Instant::now();
1170 }
1171 };
1172 }
1173 for (il, layer) in self.layers.iter().enumerate() {
1174 let mut h = e.uninit(t * n_embd)?;
1175 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
1176 if probe {
1177 e.stream().synchronize()?;
1178 eprintln!("[probe] L{il} norm ok");
1179 }
1180 anat_mark!(4);
1181 let mixed = match &layer.mixer {
1182 Mixer::Full(fa) => {
1183 let y = self.full_attn(e, fa, &h, &pos_d, t, il)?;
1184 anat_mark!(0);
1185 y
1186 }
1187 Mixer::Linear(la) => {
1188 let y = self.linear_attn(e, la, &h, t)?;
1189 anat_mark!(1);
1190 y
1191 }
1192 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
1193 };
1194 if probe {
1195 e.stream().synchronize()?;
1196 eprintln!("[probe] L{il} mixer ok");
1197 }
1198 let mut x1 = e.uninit(t * n_embd)?;
1199 e.add(&x, &mixed, &mut x1, t * n_embd)?;
1200 let mut z = e.uninit(t * n_embd)?;
1201 e.rms_norm(
1202 &x1,
1203 layer.post_attn_norm.float_data(),
1204 &mut z,
1205 n_embd,
1206 t,
1207 eps,
1208 )?;
1209 anat_mark!(4);
1210 let ffn_out = match &layer.ffn {
1211 crate::hybrid::Ffn::Dense {
1212 ffn_gate,
1213 ffn_up,
1214 ffn_down,
1215 } => {
1216 let n_ff = ffn_gate.out_features();
1217 let mut g2 = e.matmul_group(&[ffn_gate, ffn_up], &z, t)?;
1218 let up = g2.pop().unwrap();
1219 let gate = g2.pop().unwrap();
1220 let mut act = e.uninit(t * n_ff)?;
1221 Self::ffn_act_lim(
1223 e,
1224 &self.cfg,
1225 &gate,
1226 &up,
1227 1.0,
1228 1.0,
1229 self.cfg.clamp_shexp_at(il as u32),
1230 &mut act,
1231 t * n_ff,
1232 )?;
1233 let y = e.matmul(ffn_down, &act, t)?;
1234 anat_mark!(3);
1235 y
1236 }
1237 crate::hybrid::Ffn::Moe(m) => {
1238 let y = self.moe_ffn_il_prefill(e, m, &z, t, il as u16)?;
1239 anat_mark!(2);
1240 y
1241 }
1242 };
1243 if probe {
1244 e.stream().synchronize()?;
1245 eprintln!("[probe] L{il} ffn ok");
1246 }
1247 let mut x2 = e.uninit(t * n_embd)?;
1248 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
1249 x = x2;
1250 }
1251 if anat {
1252 let s = Self::prime_anatomy_slots();
1253 let ms = |i: usize| s[i].load(std::sync::atomic::Ordering::Relaxed) as f64 / 1.0e6;
1254 eprintln!(
1255 "[prime-anatomy] cumulative ms: attn_full={:.1} gdn_linear={:.1} moe={:.1} \
1256 dense={:.1} norms_adds={:.1} (t={t}, forward_last)",
1257 ms(0),
1258 ms(1),
1259 ms(2),
1260 ms(3),
1261 ms(4)
1262 );
1263 }
1264 let mut hn = e.uninit(t * n_embd)?;
1266 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
1267 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)?;
1270 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
1271 let logits = e.matmul(&self.output, &hlast, 1)?; Ok(e.dtoh(&logits)?)
1273 }
1274
1275 pub fn prime_cache(
1307 &self,
1308 e: &Engine,
1309 tokens: &[u32],
1310 cache: &mut Cache,
1311 queued_after: usize,
1312 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
1313 self.prime_cache_overlaid(e, tokens, cache, queued_after, None)
1314 }
1315
1316 pub fn prime_cache_overlaid(
1322 &self,
1323 e: &Engine,
1324 tokens: &[u32],
1325 cache: &mut Cache,
1326 queued_after: usize,
1327 overlay: Option<&crate::vision::EmbedOverlay>,
1328 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
1329 let n_embd = self.cfg.n_embd as usize;
1330 let t = tokens.len();
1331 let legacy_calllocal = std::env::var("MEMRA_PRIME_CALLLOCAL").as_deref() == Ok("1");
1360 let seq_end = if legacy_calllocal {
1361 cache.pos + t
1362 } else {
1363 cache.pos + t + queued_after
1364 };
1365 if overlay.is_none()
1370 && (cache.pos == 0 || step_gemm_prime_suffix_on())
1371 && t >= PRIME_MIN_T
1372 && crate::step_gemm_prime_on()
1373 && self.uses_sliding_gated_moe_program()
1374 {
1375 let n_embd = self.cfg.n_embd as usize;
1376 let base = cache.pos;
1377 let width = crate::cache::PRIME_CHUNK_MAX_TOKENS;
1378 let mut hiddens = e.uninit(t * n_embd)?;
1379 let mut last: Option<(Vec<f32>, CudaSlice<f32>)> = None;
1380 let mut start = 0usize;
1381 while start < t {
1382 let mut end = (start + width).min(t);
1385 if t - end > 0 && t - end < PRIME_MIN_T {
1386 end = t;
1387 }
1388 let mut out = self.step35_prime_cache_batch(
1389 e,
1390 &[&tokens[start..end]],
1391 &mut [cache],
1392 &[seq_end],
1393 )?;
1394 if out.len() != 1 {
1395 return Err("B=1 batched prime returned a non-singleton".into());
1396 }
1397 let (logits, h_seed, hidden) = out.remove(0);
1398 e.copy_into(
1399 &mut hiddens,
1400 start * n_embd,
1401 &hidden,
1402 (end - start) * n_embd,
1403 )?;
1404 last = Some((logits, h_seed));
1405 start = end;
1406 }
1407 let (logits, h_seed) = last.expect("prime produced no chunk");
1408 eprintln!(
1413 "[gemm-prime] ENGAGED t={t} base={base} seq_end={seq_end} chunks<={width} (GEMM trunk + grouped MoE)"
1414 );
1415 return Ok((logits, h_seed, hiddens));
1416 }
1417 if self.uses_sliding_gated_moe_program() {
1418 eprintln!(
1419 "[gemm-prime] WALK t={t} base={} seq_end={seq_end} (batched prime declined)",
1420 cache.pos
1421 );
1422 }
1423 if overlay.is_none() {
1424 if let Some(out) = self.step35_prime_trows(e, tokens, cache)? {
1425 return Ok(out);
1426 }
1427 }
1428 assert!(
1432 t >= PRIME_MIN_T,
1433 "prime_cache needs T >= {PRIME_MIN_T} (caller gates)"
1434 );
1435 assert!(
1436 cache.pos + t <= cache.max_ctx,
1437 "prime_cache: prompt exceeds cache max_ctx"
1438 );
1439
1440 if self.is_gemma4_e4b() || self.uses_gemma_program() {
1452 if self.is_gemma4_e4b() {
1453 if overlay.is_some() {
1454 return Err(
1455 "vision embedding overlay is unsupported on gemma4 E4B (PLE prime)".into(),
1456 );
1457 }
1458 return self.gemma4_e4b_prime(e, tokens, cache);
1459 }
1460 return self.gemma4_prime(e, tokens, cache, overlay);
1465 }
1466 let ranges = prime_chunk_ranges(t, self.layers.len(), self.gdn_prime_grid_on());
1467 if ranges.len() == 1 {
1505 return self.prime_chunk(e, tokens, cache, seq_end, 0, overlay);
1506 }
1507 if crate::pp::prime_pipe_on() && crate::pp::prime_pp_on() && !crate::pp::pp2_streams_off() {
1512 if let Some(fence) = crate::pp::pp_cuts(self.layers.len()).filter(|f| f.len() == 3) {
1513 if overlay.is_some() {
1514 return Err(
1515 "vision embedding overlay + pipelined PP prime unsupported (v1); \
1516 run the serial prime (single device or MEMRA_PRIME_PIPE=0)"
1517 .into(),
1518 );
1519 }
1520 if crate::pp::pp_multi_stream_same_device() {
1521 return Err(
1522 "prime chunk pipeline refused with 2 stage streams on one device — \
1523 that concurrent-stream placement remains quarantined by the deferred \
1524 pp flake record. Use one device per stage or MEMRA_PRIME_PIPE=0 for \
1525 the serial split."
1526 .into(),
1527 );
1528 }
1529 return self.prime_cache_pp2_pipelined(e, tokens, cache, seq_end, &ranges, &fence);
1530 }
1531 }
1532 let mut hiddens = e.uninit(t * n_embd)?;
1533 let mut last: Option<(Vec<f32>, CudaSlice<f32>)> = None;
1534 for &(start, end) in &ranges {
1535 if let Some(taps) = cache.dflash_taps.as_mut() {
1537 taps.base = start;
1538 }
1539 let (l, hs, x) =
1540 self.prime_chunk(e, &tokens[start..end], cache, seq_end, start, overlay)?;
1541 e.copy_into(&mut hiddens, start * n_embd, &x, (end - start) * n_embd)?;
1542 last = Some((l, hs));
1543 }
1544 let (logits, h_seed) = last.unwrap();
1545 Ok((logits, h_seed, hiddens))
1546 }
1547
1548 fn prime_cache_pp2_pipelined(
1553 &self,
1554 e: &Engine,
1555 tokens: &[u32],
1556 cache: &mut Cache,
1557 seq_end: usize,
1558 ranges: &[(usize, usize)],
1559 fence: &[usize],
1560 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
1561 debug_assert_eq!(fence.len(), 3);
1562 debug_assert!(ranges.len() >= 2);
1563 let rt = crate::pp::PpNRt::get(e)?;
1564 assert_eq!(
1565 rt.n_stages(),
1566 2,
1567 "prime pipeline requires exactly two PP stages"
1568 );
1569 let n_embd = self.cfg.n_embd as usize;
1570 let t = tokens.len();
1571 let initial_base = cache.pos;
1572 let caller_stream = e.stream();
1573
1574 rt.fence_stages_behind(&caller_stream)?;
1579 let max_payload = ranges.iter().map(|(s, e)| (e - s) * n_embd).max().unwrap();
1580 rt.prepare_overlap_slots(0, max_payload)?;
1581
1582 let mut hiddens = e.uninit(t * n_embd)?;
1583 let mut last: Option<(Vec<f32>, CudaSlice<f32>)> = None;
1584 let mut stage_caches = PrimeCacheStages::new(cache, fence[1]);
1585 let (cache0, cache1) = stage_caches.parts();
1586 let (first_start, first_end) = ranges[0];
1587 let mut slot = self.prime_pp2_stage0_enqueue(
1588 e,
1589 rt,
1590 &tokens[first_start..first_end],
1591 cache0,
1592 seq_end,
1593 fence,
1594 initial_base + first_start,
1595 true,
1596 )?;
1597 cache0.pos = initial_base + first_end;
1598
1599 for (i, &(start, end)) in ranges.iter().enumerate() {
1600 let base = initial_base + start;
1601 debug_assert_eq!(
1602 cache1.pos, base,
1603 "stage 1 must drain chunks in original position order"
1604 );
1605 let (out, next_slot) = if let Some(&(next_start, next_end)) = ranges.get(i + 1) {
1606 let next_base = initial_base + next_start;
1607 debug_assert_eq!(
1608 cache0.pos, next_base,
1609 "stage 0 must issue chunks in original position order"
1610 );
1611 let cache0_stage = &mut *cache0;
1612 std::thread::scope(|scope| -> Result<_, Box<dyn std::error::Error>> {
1617 let stage0 = scope.spawn(move || -> Result<usize, String> {
1618 let next = self
1619 .prime_pp2_stage0_enqueue(
1620 e,
1621 rt,
1622 &tokens[next_start..next_end],
1623 cache0_stage,
1624 seq_end,
1625 fence,
1626 next_base,
1627 true,
1628 )
1629 .map_err(|err| err.to_string())?;
1630 cache0_stage.pos = initial_base + next_end;
1631 Ok(next)
1632 });
1633 let x = self.prime_pp2_stage1_enqueue(
1634 e,
1635 rt,
1636 slot,
1637 end - start,
1638 cache1,
1639 seq_end,
1640 fence,
1641 base,
1642 true,
1643 )?;
1644 let out = {
1645 rt.bind_stage(1)?;
1646 let _st1 = rt.enter(1);
1647 let e1 = rt.engine(1, e);
1648 self.prime_chunk_epilogue(e1, x, end - start, cache1)?
1649 };
1650 let next = stage0
1651 .join()
1652 .map_err(|_| "pipeprime stage-0 host walker panicked")?
1653 .map_err(|err| -> Box<dyn std::error::Error> { err.into() })?;
1654 Ok((out, Some(next)))
1655 })?
1656 } else {
1657 let x = self.prime_pp2_stage1_enqueue(
1658 e,
1659 rt,
1660 slot,
1661 end - start,
1662 cache1,
1663 seq_end,
1664 fence,
1665 base,
1666 true,
1667 )?;
1668 let out = {
1669 rt.bind_stage(1)?;
1670 let _st1 = rt.enter(1);
1671 let e1 = rt.engine(1, e);
1672 self.prime_chunk_epilogue(e1, x, end - start, cache1)?
1673 };
1674 (out, None)
1675 };
1676
1677 rt.publish_to(1, &caller_stream)?;
1678 e.copy_into(&mut hiddens, start * n_embd, &out.2, (end - start) * n_embd)?;
1679 last = Some((out.0, out.1));
1680 crate::pp::PRIME_SPLIT_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
1681
1682 if let Some(next) = next_slot {
1683 rt.fence_stages_behind(&caller_stream)?;
1688 slot = next;
1689 }
1690 }
1691
1692 debug_assert_eq!(cache0.pos, initial_base + t);
1693 debug_assert_eq!(cache1.pos, initial_base + t);
1694 let (logits, h_seed) = last.unwrap();
1695 Ok((logits, h_seed, hiddens))
1696 }
1697
1698 fn gdn_hk(e: &Engine, t: usize, num_v: usize, num_k: usize) -> usize {
1705 if Engine::gdn_db_on()
1706 && Engine::gdn_chunked_enabled()
1707 && t >= 16
1708 && e.gdn_mma_enabled(Engine::gdn_chunk_size())
1709 && num_k * 2 == num_v
1710 {
1711 num_k
1712 } else {
1713 num_v
1714 }
1715 }
1716
1717 fn f16out_on(e: &Engine, t: usize) -> bool {
1722 crate::f16_ffi::pp_f16_enabled()
1723 && t >= 16
1724 && !e.verify_exact_on()
1725 && std::env::var("MEMRA_F16OUT").as_deref() != Ok("0")
1726 }
1727
1728 pub fn prime_slabs_get(
1736 &self,
1737 e: &Engine,
1738 t: usize,
1739 n_embd: usize,
1740 n_ff_max: usize,
1741 ) -> Result<std::sync::Arc<std::sync::Mutex<PrimeSlabs>>, Box<dyn std::error::Error>> {
1742 let mut slabs = self.prime_slabs.lock().unwrap();
1743 let dev = e.ctx().ordinal();
1744 let need_new = match slabs.get(&dev) {
1745 None => true,
1746 Some(sl) => sl.lock().unwrap().t_cap < t,
1747 };
1748 if need_new {
1749 slabs.insert(
1750 dev,
1751 std::sync::Arc::new(std::sync::Mutex::new(PrimeSlabs {
1752 t_cap: t,
1753 h: e.uninit(t * n_embd)?,
1754 x1: e.uninit(t * n_embd)?,
1755 z: e.uninit(t * n_embd)?,
1756 act: e.uninit(t * n_ff_max)?,
1757 xa: e.uninit(t * n_embd)?,
1758 xb: e.uninit(t * n_embd)?,
1759 h16: e.alloc_u8_uninit(t * n_embd * 2)?,
1760 z16: e.alloc_u8_uninit(t * n_embd * 2)?,
1761 gate: e.uninit(t * n_ff_max)?,
1762 up: e.uninit(t * n_ff_max)?,
1763 ffn_out: e.uninit(t * n_embd)?,
1764 seg_glue: Vec::new(),
1765 mixed: e.uninit(t * n_embd)?,
1766 seg_mid: Vec::new(),
1767 seg_t: 0,
1768 })),
1769 );
1770 }
1771 Ok(slabs.get(&dev).expect("prime slab inserted").clone())
1772 }
1773
1774 fn prime_chunk(
1778 &self,
1779 e: &Engine,
1780 tokens: &[u32],
1781 cache: &mut Cache,
1782 seq_end: usize,
1783 chunk_off: usize,
1784 overlay: Option<&crate::vision::EmbedOverlay>,
1785 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
1786 if crate::pp::pp_host_bounce_active()
1787 && (self.uses_gemma_program() || !crate::pp::prime_pp_on())
1788 {
1789 return Err(
1790 "prime_chunk: refused with MEMRA_PP_HOST_BOUNCE=1 because this configuration \
1791 has no active prime stage split and would peer-read remote weights; keep \
1792 MEMRA_PRIME_PP enabled and use a PP-prime-supported model"
1793 .into(),
1794 );
1795 }
1796 if !self.uses_gemma_program() && !crate::pp::pp2_streams_off() && crate::pp::prime_pp_on() {
1805 if let Some(fence) = crate::pp::pp_cuts(self.layers.len()) {
1806 if overlay.is_some() {
1807 return Err("vision embedding overlay + PP prime unsupported (v1); \
1808 run single-device or MEMRA_PRIME_PP=0"
1809 .into());
1810 }
1811 return self.prime_chunk_ppn(e, tokens, cache, seq_end, &fence);
1812 }
1813 }
1814 if crate::pp::pp_host_bounce_active() {
1815 return Err(
1816 "prime_chunk: MEMRA_PP_HOST_BOUNCE=1 found no valid prime stage split; \
1817 refusing an unsplit remote-weight walk"
1818 .into(),
1819 );
1820 }
1821 let t = tokens.len();
1822 let base = cache.pos;
1823 debug_assert!(
1824 seq_end >= base + t,
1825 "prime_chunk: seq_end must cover this chunk"
1826 );
1827 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
1828 let pos_d = e.htod_i32(&pos)?;
1829
1830 let mut x_embed = self.embed(e, tokens)?; if let Some(ov) = overlay {
1832 let n_embd = self.cfg.n_embd as usize;
1836 for &(pos, row_off, n_rows) in &ov.spans {
1837 let lo = pos.max(chunk_off);
1838 let hi = (pos + n_rows).min(chunk_off + t);
1839 if lo < hi {
1840 let src_row = row_off + (lo - pos);
1841 let view = ov
1842 .rows
1843 .slice(src_row * n_embd..(src_row + (hi - lo)) * n_embd);
1844 e.copy_view_into(
1845 &mut x_embed,
1846 (lo - chunk_off) * n_embd,
1847 &view,
1848 (hi - lo) * n_embd,
1849 )?;
1850 }
1851 }
1852 }
1853 let x = self.prime_layers(
1854 e,
1855 x_embed,
1856 0,
1857 self.layers.len(),
1858 &pos_d,
1859 t,
1860 base,
1861 cache,
1862 seq_end,
1863 )?;
1864 self.prime_chunk_epilogue(e, x, t, cache)
1865 }
1866
1867 #[allow(clippy::too_many_arguments)]
1883 fn prime_layers(
1884 &self,
1885 e: &Engine,
1886 x_in: CudaSlice<f32>,
1887 lo: usize,
1888 hi: usize,
1889 pos_d: &CudaSlice<i32>,
1890 t: usize,
1891 base: usize,
1892 cache: &mut Cache,
1893 seq_end: usize,
1894 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1895 let cfg = &self.cfg;
1896 let n_embd = cfg.n_embd as usize;
1897 let eps = cfg.rms_eps;
1898 let f16fuse = crate::f16_ffi::pp_f16_enabled() && t >= 16;
1902 let n_ff_max = self
1909 .layers
1910 .iter()
1911 .map(|l| match &l.ffn {
1912 crate::hybrid::Ffn::Dense { ffn_gate, .. } => ffn_gate.out_features(),
1913 _ => n_embd,
1914 })
1915 .max()
1916 .unwrap_or(n_embd)
1917 .max(n_embd);
1918 let use_slabs = std::env::var("MEMRA_PRIME_SLABS").as_deref() != Ok("0");
1919 let slab = if use_slabs {
1920 Some(self.prime_slabs_get(e, t, n_embd, n_ff_max)?)
1921 } else {
1922 None
1923 };
1924 let mut slab_guard = slab.as_ref().map(|sl| sl.lock().unwrap());
1925 let mut x_own; type SlabRefs<'a> = (
1927 &'a mut CudaSlice<f32>,
1928 &'a mut CudaSlice<f32>,
1929 &'a mut CudaSlice<f32>,
1930 &'a mut CudaSlice<f32>,
1931 &'a mut CudaSlice<u8>,
1932 &'a mut CudaSlice<u8>,
1933 &'a mut CudaSlice<f32>,
1934 &'a mut CudaSlice<f32>,
1935 &'a mut CudaSlice<f32>,
1936 );
1937 let (mut x_cur, mut x_nxt, sl): (
1938 &mut CudaSlice<f32>,
1939 &mut CudaSlice<f32>,
1940 Option<SlabRefs>,
1941 );
1942 let mut seg: Option<(
1943 &mut Vec<Option<cudarc::driver::CudaGraph>>,
1944 &mut Vec<Option<cudarc::driver::CudaGraph>>,
1945 &mut CudaSlice<f32>,
1946 &mut usize,
1947 )> = None;
1948 let mut x_own2;
1949 match slab_guard.as_mut() {
1950 Some(g) => {
1951 let slabs = &mut **g;
1952 e.copy_into(&mut slabs.xa, 0, &x_in, t * n_embd)?;
1953 let PrimeSlabs {
1954 xa,
1955 xb,
1956 h,
1957 x1,
1958 z,
1959 act,
1960 h16,
1961 z16,
1962 gate,
1963 up,
1964 ffn_out,
1965 seg_glue,
1966 mixed,
1967 seg_mid,
1968 seg_t,
1969 ..
1970 } = slabs;
1971 x_cur = xa;
1972 x_nxt = xb;
1973 seg = Some((seg_glue, seg_mid, mixed, seg_t));
1974 sl = Some((h, x1, z, act, h16, z16, gate, up, ffn_out));
1975 }
1976 None => {
1977 x_own = x_in;
1978 x_own2 = e.uninit(t * n_embd)?;
1979 x_cur = &mut x_own;
1980 x_nxt = &mut x_own2;
1981 sl = None;
1982 }
1983 }
1984 let mut alloc_h;
1985 let mut alloc_x1;
1986 let mut alloc_z;
1987 let mut alloc_act;
1988 let mut alloc_h16;
1989 let mut alloc_z16;
1990 let mut alloc_gate;
1991 let mut alloc_up;
1992 let mut alloc_fo;
1993 let (h, x1, z, act): (
1994 &mut CudaSlice<f32>,
1995 &mut CudaSlice<f32>,
1996 &mut CudaSlice<f32>,
1997 &mut CudaSlice<f32>,
1998 );
1999 let (h16, z16): (&mut CudaSlice<u8>, &mut CudaSlice<u8>);
2000 let (sl_gate, sl_up, sl_fo): (
2001 &mut CudaSlice<f32>,
2002 &mut CudaSlice<f32>,
2003 &mut CudaSlice<f32>,
2004 );
2005 match sl {
2006 Some((a, b, c, d, e16, f16b, g, u, fo)) => {
2007 h = a;
2008 x1 = b;
2009 z = c;
2010 act = d;
2011 h16 = e16;
2012 z16 = f16b;
2013 sl_gate = g;
2014 sl_up = u;
2015 sl_fo = fo;
2016 }
2017 None => {
2018 alloc_h = e.uninit(t * n_embd)?;
2019 alloc_x1 = e.uninit(t * n_embd)?;
2020 alloc_z = e.uninit(t * n_embd)?;
2021 alloc_act = e.uninit(t * n_ff_max)?;
2022 alloc_h16 = e.alloc_u8_uninit(t * n_embd * 2)?;
2023 alloc_z16 = e.alloc_u8_uninit(t * n_embd * 2)?;
2024 alloc_gate = e.uninit(t * n_ff_max)?;
2025 alloc_up = e.uninit(t * n_ff_max)?;
2026 alloc_fo = e.uninit(t * n_embd)?;
2027 h = &mut alloc_h;
2028 x1 = &mut alloc_x1;
2029 z = &mut alloc_z;
2030 act = &mut alloc_act;
2031 h16 = &mut alloc_h16;
2032 z16 = &mut alloc_z16;
2033 sl_gate = &mut alloc_gate;
2034 sl_up = &mut alloc_up;
2035 sl_fo = &mut alloc_fo;
2036 }
2037 }
2038 let n_layers = self.layers.len();
2043 let use_seg = f16fuse
2053 && seg.is_some()
2054 && !self.uses_sliding_gated_moe_program()
2055 && lo == 0
2056 && hi == n_layers
2057 && std::env::var("MEMRA_PRIME_SEG").as_deref() == Ok("1")
2058 && {
2064 let ok = crate::spec::graph_launch_headroom_ok(e);
2065 if !ok {
2066 static NOTED: std::sync::Once = std::sync::Once::new();
2067 NOTED.call_once(|| crate::spec::graph_replay_suspended_note("prime-seg"));
2068 }
2069 ok
2070 };
2071 if let Some((sg, sm, _, st)) = seg.as_mut() {
2072 if **st != t {
2073 sg.clear();
2074 sg.extend((0..n_layers).map(|_| None));
2075 sm.clear();
2076 sm.extend((0..n_layers).map(|_| None));
2077 **st = t;
2078 }
2079 }
2080 {
2081 let layer_lo = &self.layers[lo];
2082 if f16fuse {
2083 e.rms_norm_f16out(
2084 x_cur,
2085 layer_lo.attn_norm.float_data(),
2086 h,
2087 h16,
2088 n_embd,
2089 t,
2090 eps,
2091 )?;
2092 } else {
2093 e.rms_norm(x_cur, layer_lo.attn_norm.float_data(), h, n_embd, t, eps)?;
2094 }
2095 }
2096 let anat = Self::prime_anatomy_on();
2097 let mut anat_last = if anat {
2098 e.stream().synchronize()?;
2099 Some(std::time::Instant::now())
2100 } else {
2101 None
2102 };
2103 macro_rules! anat_mark {
2105 ($slot:expr) => {
2106 if let Some(ts) = anat_last.as_mut() {
2107 e.stream().synchronize()?;
2108 Self::prime_anatomy_slots()[$slot].fetch_add(
2109 ts.elapsed().as_nanos() as u64,
2110 std::sync::atomic::Ordering::Relaxed,
2111 );
2112 *ts = std::time::Instant::now();
2113 }
2114 };
2115 }
2116 for il in lo..hi {
2117 let layer = &self.layers[il];
2118 let hx16 = if f16fuse { Some(&*h16) } else { None };
2119 if use_seg {
2120 let (pre, pre16, w_out) = match &layer.mixer {
2123 Mixer::Full(fa) => {
2124 let g3 = match hx16 {
2125 Some(xh) => e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], h, xh, t)?,
2126 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
2127 };
2128 let (pre, pre16) =
2129 self.full_attn_prime_core_inner(e, fa, g3, &pos_d, t, cache, il)?;
2130 (pre, pre16, &fa.wo)
2131 }
2132 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
2133 Mixer::Linear(la) => {
2134 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
2135 let g4 = match hx16 {
2136 Some(xh) => e.matmul_group_xh(&ws, h, xh, t)?,
2137 None => e.matmul_group(&ws, h, t)?,
2138 };
2139 let (pre, pre16) =
2140 self.linear_attn_prime_core_pad_inner(e, la, g4, t, cache, il, None)?;
2141 (pre, pre16, &la.ssm_out)
2142 }
2143 };
2144 {
2145 let (_, sm, mslab, _) = seg.as_mut().unwrap();
2146 let pre_n = pre.len() / t;
2147 let xh_pre = match pre16 {
2148 Some(x) => x,
2149 None => e.f16_act(&pre, t * pre_n, pre_n)?,
2150 };
2151 if !e.try_f16_gemm_pre_into(w_out, &xh_pre, t, mslab)? {
2152 let y = e.matmul(w_out, &pre, t)?;
2153 e.copy_into(mslab, 0, &y, t * n_embd)?;
2154 }
2155 if sm[il].is_none() {
2156 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
2157 let w_post = layer.post_attn_norm.float_data();
2158 e.stream().synchronize()?;
2159 e.stream()
2160 .begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
2161 let r = (|| -> Result<(), Box<dyn std::error::Error>> {
2162 e.add(x_cur, mslab, x1, t * n_embd)?;
2163 e.rms_norm_f16out(x1, w_post, z, z16, n_embd, t, eps)?;
2164 Ok(())
2165 })();
2166 let g = e.stream().end_capture(
2167 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH);
2168 r?;
2169 sm[il] = Some(g?.ok_or("S-mid capture produced no graph")?);
2170 }
2171 sm[il].as_ref().unwrap().launch()?;
2172 }
2173 } else {
2174 let mixed = match &layer.mixer {
2175 Mixer::Full(fa) => {
2176 let y =
2177 self.full_attn_prime(e, fa, h, hx16, &pos_d, t, cache, il, seq_end)?;
2178 anat_mark!(0);
2179 y
2180 }
2181 Mixer::Linear(la) => {
2182 let y = self.linear_attn_prime(e, la, h, hx16, t, cache, il)?;
2183 anat_mark!(1);
2184 y
2185 }
2186 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
2187 };
2188 if f16fuse {
2189 e.add_rms_norm_f16out(
2192 x_cur,
2193 &mixed,
2194 layer.post_attn_norm.float_data(),
2195 x1,
2196 z,
2197 z16,
2198 n_embd,
2199 t,
2200 eps,
2201 )?;
2202 } else {
2203 e.add(x_cur, &mixed, x1, t * n_embd)?;
2204 e.rms_norm(x1, layer.post_attn_norm.float_data(), z, n_embd, t, eps)?;
2205 }
2206 anat_mark!(4);
2207 }
2208 let zx16 = if f16fuse { Some(&*z16) } else { None };
2209 match &layer.ffn {
2210 crate::hybrid::Ffn::Dense {
2211 ffn_gate,
2212 ffn_up,
2213 ffn_down,
2214 } => {
2215 let n_ff = ffn_gate.out_features();
2216 let mut into_ok = false;
2219 if let Some(xh) = zx16 {
2220 into_ok = e.try_f16_gemm_pre_into(ffn_gate, xh, t, sl_gate)?
2221 && e.try_f16_gemm_pre_into(ffn_up, xh, t, sl_up)?;
2222 }
2223 if !into_ok {
2224 let mut g2 = match zx16 {
2225 Some(xh) => e.matmul_group_xh(&[ffn_gate, ffn_up], z, xh, t)?,
2226 None => e.matmul_group(&[ffn_gate, ffn_up], z, t)?,
2227 };
2228 let up_y = g2.pop().unwrap();
2229 let gate_y = g2.pop().unwrap();
2230 e.copy_into(sl_gate, 0, &gate_y, t * n_ff)?;
2231 e.copy_into(sl_up, 0, &up_y, t * n_ff)?;
2232 }
2233 let d_lim = self.cfg.clamp_shexp_at(il as u32);
2238 let act16 = if Self::f16out_on(e, t) && self.cfg.m3.is_none() && d_lim.is_none()
2239 {
2240 let mut a16 = e.alloc_u8_uninit(t * n_ff * 2)?;
2241 e.silu_mul_f16out(sl_gate, sl_up, act, &mut a16, t * n_ff)?;
2242 Some(a16)
2243 } else {
2244 Self::ffn_act_lim(
2245 e,
2246 &self.cfg,
2247 sl_gate,
2248 sl_up,
2249 1.0,
2250 1.0,
2251 d_lim,
2252 act,
2253 t * n_ff,
2254 )?;
2255 None
2256 };
2257 let xh_act = match act16 {
2259 Some(x) => x,
2260 None => e.f16_act(act, t * n_ff, n_ff)?,
2261 };
2262 if !e.try_f16_gemm_pre_into(ffn_down, &xh_act, t, sl_fo)? {
2263 let y = e.matmul(ffn_down, &*act, t)?;
2264 e.copy_into(sl_fo, 0, &y, t * n_embd)?;
2265 }
2266 }
2267 crate::hybrid::Ffn::Moe(m) => {
2268 let y = self.moe_ffn_il_prefill(e, m, z, t, il as u16)?;
2269 e.copy_into(sl_fo, 0, &y, t * n_embd)?;
2270 anat_mark!(2);
2271 }
2272 }
2273 if let (crate::hybrid::Ffn::Dense { .. }, true) = (&layer.ffn, anat) {
2274 anat_mark!(3);
2275 }
2276 if use_seg && il + 1 < hi {
2277 let w_next = self.layers[il + 1].attn_norm.float_data();
2279 let (sg, _, _, _) = seg.as_mut().unwrap();
2280 if sg[il].is_none() {
2281 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
2282 e.stream().synchronize()?;
2283 e.stream()
2284 .begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
2285 let r = (|| -> Result<(), Box<dyn std::error::Error>> {
2286 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
2287 e.rms_norm_f16out(x_nxt, w_next, h, h16, n_embd, t, eps)?;
2288 Ok(())
2289 })();
2290 let g = e.stream().end_capture(
2291 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
2292 );
2293 r?;
2294 sg[il] = Some(g?.ok_or("S-glue capture produced no graph")?);
2295 }
2296 sg[il].as_ref().unwrap().launch()?;
2297 } else {
2298 if il + 1 < hi {
2299 let w_next = self.layers[il + 1].attn_norm.float_data();
2300 if f16fuse {
2301 e.add_rms_norm_f16out(x1, sl_fo, w_next, x_nxt, h, h16, n_embd, t, eps)?;
2302 } else {
2303 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
2304 e.rms_norm(x_nxt, w_next, h, n_embd, t, eps)?;
2305 }
2306 } else {
2307 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
2308 }
2309 }
2310 anat_mark!(4);
2311 if let Some(path) = Self::prime_trace_path() {
2317 let row = (base + t - 1) as usize;
2318 let host = e.dtoh(x_nxt)?;
2319 let last = &host[(t - 1) * n_embd..t * n_embd];
2320 use std::io::Write as _;
2321 let mut f = std::fs::OpenOptions::new()
2322 .create(true)
2323 .append(true)
2324 .open(path)?;
2325 let mut h64: u64 = 0xcbf29ce484222325;
2326 for v in last {
2327 h64 ^= v.to_bits() as u64;
2328 h64 = h64.wrapping_mul(0x100000001b3);
2329 }
2330 writeln!(
2331 f,
2332 "{{\"pos\":{row},\"layer\":{il},\"t\":{t},\"base\":{base},\
2333 \"hash\":\"{h64:016x}\",\"v0\":{:.9e},\"v1\":{:.9e},\"v2\":{:.9e}}}",
2334 last[0], last[1], last[2]
2335 )?;
2336 }
2337 self.dflash_tap(e, cache, il, x_nxt, t)?;
2340 std::mem::swap(&mut x_cur, &mut x_nxt);
2341 }
2342 if anat {
2343 let s = Self::prime_anatomy_slots();
2344 let ms = |i: usize| s[i].load(std::sync::atomic::Ordering::Relaxed) as f64 / 1.0e6;
2345 eprintln!(
2346 "[prime-anatomy] cumulative ms: attn_full={:.1} gdn_linear={:.1} moe={:.1} \
2347 dense={:.1} norms_adds={:.1} (t={t}, layers {lo}..{hi})",
2348 ms(0),
2349 ms(1),
2350 ms(2),
2351 ms(3),
2352 ms(4)
2353 );
2354 }
2355 let mut x = e.uninit(t * n_embd)?;
2357 e.copy_into(&mut x, 0, x_cur, t * n_embd)?;
2358 drop(slab_guard);
2359 Ok(x)
2360 }
2361
2362 fn prime_chunk_epilogue(
2367 &self,
2368 e: &Engine,
2369 x: CudaSlice<f32>,
2370 t: usize,
2371 cache: &mut Cache,
2372 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
2373 let n_embd = self.cfg.n_embd as usize;
2374 let eps = self.cfg.rms_eps;
2375 let mut h_seed = e.uninit(n_embd)?;
2379 if !crate::spec::spec_hpost() {
2380 e.copy_view_into(
2381 &mut h_seed,
2382 0,
2383 &x.slice((t - 1) * n_embd..t * n_embd),
2384 n_embd,
2385 )?;
2386 }
2387 let mut hn = e.uninit(t * n_embd)?;
2389 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
2390 if crate::spec::spec_hpost() {
2391 e.copy_view_into(
2392 &mut h_seed,
2393 0,
2394 &hn.slice((t - 1) * n_embd..t * n_embd),
2395 n_embd,
2396 )?;
2397 }
2398 let last = e.view(&hn, t * n_embd);
2399 let last_row = last.slice((t - 1) * n_embd..t * n_embd);
2400 let mut hlast = e.uninit(n_embd)?;
2401 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
2402 let logits = e.matmul(&self.output, &hlast, 1)?;
2403 cache.pos += t;
2404 Ok((
2407 e.dtoh(&logits)?,
2408 h_seed,
2409 if crate::spec::spec_hpost() { hn } else { x },
2410 ))
2411 }
2412
2413 pub fn hidden_postnorm_row(
2419 &self,
2420 e: &Engine,
2421 hiddens: &CudaSlice<f32>,
2422 row: usize,
2423 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
2424 let n_embd = self.cfg.n_embd as usize;
2425 let mut x1 = e.uninit(n_embd)?;
2426 e.copy_view_into(
2427 &mut x1,
2428 0,
2429 &hiddens.slice(row * n_embd..(row + 1) * n_embd),
2430 n_embd,
2431 )?;
2432 if crate::spec::spec_hpost() {
2433 return Ok(e.dtoh(&x1)?);
2434 }
2435 let mut hn = e.uninit(n_embd)?;
2436 e.rms_norm(
2437 &x1,
2438 self.output_norm.float_data(),
2439 &mut hn,
2440 n_embd,
2441 1,
2442 self.cfg.rms_eps,
2443 )?;
2444 Ok(e.dtoh(&hn)?)
2445 }
2446
2447 fn prime_chunk_ppn(
2471 &self,
2472 e: &Engine,
2473 tokens: &[u32],
2474 cache: &mut Cache,
2475 seq_end: usize,
2476 fence: &[usize],
2477 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
2478 let rt = crate::pp::PpNRt::get(e)?;
2479 let n_st = fence.len() - 1;
2480 assert_eq!(
2481 rt.n_stages(),
2482 n_st,
2483 "PpNRt stage count {} != fence stages {n_st}",
2484 rt.n_stages()
2485 );
2486 let n_embd = self.cfg.n_embd as usize;
2487 let t = tokens.len();
2488 let base = cache.pos;
2489 debug_assert!(
2490 seq_end >= base + t,
2491 "prime_chunk_ppn: seq_end must cover this chunk"
2492 );
2493 let payload = t * n_embd;
2494 let caller_stream = e.stream();
2498 rt.fence_stages_behind(&caller_stream)?;
2499
2500 if n_st == 2 {
2501 let slot =
2502 self.prime_pp2_stage0_enqueue(e, rt, tokens, cache, seq_end, fence, base, false)?;
2503 let x =
2504 self.prime_pp2_stage1_enqueue(e, rt, slot, t, cache, seq_end, fence, base, false)?;
2505 let out = {
2506 rt.bind_stage(1)?;
2507 let _st1 = rt.enter(1);
2508 let e1 = rt.engine(1, e);
2509 self.prime_chunk_epilogue(e1, x, t, cache)?
2510 };
2511 rt.publish_to(1, &caller_stream)?;
2512 crate::pp::PRIME_SPLIT_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2513 return Ok(out);
2514 }
2515
2516 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
2517
2518 let mut slot = {
2520 let _st0 = rt.enter(0);
2521 let e0 = rt.engine(0, e);
2522 let pos_d = e0.htod_i32(&pos)?;
2523 let x = self.embed(e0, tokens)?;
2524 let x =
2525 self.prime_layers(e0, x, fence[0], fence[1], &pos_d, t, base, cache, seq_end)?;
2526 rt.tx(0, &x, payload)?
2527 };
2529
2530 for s in 1..n_st - 1 {
2532 let _st = rt.enter(s);
2533 let es = rt.engine(s, e);
2534 let pos_d = es.htod_i32(&pos)?;
2535 let x = rt.rx(s - 1, slot, payload)?;
2536 let x = self.prime_layers(
2537 es,
2538 x,
2539 fence[s],
2540 fence[s + 1],
2541 &pos_d,
2542 t,
2543 base,
2544 cache,
2545 seq_end,
2546 )?;
2547 slot = rt.tx(s, &x, payload)?;
2548 }
2549
2550 let _stl = rt.enter(n_st - 1);
2552 let el = rt.engine(n_st - 1, e);
2553 let pos_d = el.htod_i32(&pos)?;
2554 let x = rt.rx(n_st - 2, slot, payload)?;
2555 let x = self.prime_layers(
2556 el,
2557 x,
2558 fence[n_st - 1],
2559 fence[n_st],
2560 &pos_d,
2561 t,
2562 base,
2563 cache,
2564 seq_end,
2565 )?;
2566 let out = self.prime_chunk_epilogue(el, x, t, cache)?;
2567 rt.publish_to(n_st - 1, &caller_stream)?;
2573 crate::pp::PRIME_SPLIT_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2574 Ok(out)
2575 }
2576
2577 fn prime_pp2_stage0_enqueue(
2578 &self,
2579 e: &Engine,
2580 rt: &crate::pp::PpNRt,
2581 tokens: &[u32],
2582 cache: &mut Cache,
2583 seq_end: usize,
2584 fence: &[usize],
2585 base: usize,
2586 pipelined: bool,
2587 ) -> Result<usize, Box<dyn std::error::Error>> {
2588 let t = tokens.len();
2589 let n_embd = self.cfg.n_embd as usize;
2590 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
2591 rt.bind_stage(0)?;
2592 let _st0 = rt.enter(0);
2593 let e0 = rt.engine(0, e);
2594 let pos_d = e0.htod_i32(&pos)?;
2595 let x = self.embed(e0, tokens)?;
2596 let _overlap = pipelined.then(crate::pp::enter_prime_pipe_stage);
2597 let x = self.prime_layers(e0, x, fence[0], fence[1], &pos_d, t, base, cache, seq_end)?;
2598 if pipelined {
2599 rt.tx_pipelined(0, &x, t * n_embd)
2600 } else {
2601 rt.tx(0, &x, t * n_embd)
2602 }
2603 }
2604
2605 fn prime_pp2_stage1_enqueue(
2606 &self,
2607 e: &Engine,
2608 rt: &crate::pp::PpNRt,
2609 slot: usize,
2610 t: usize,
2611 cache: &mut Cache,
2612 seq_end: usize,
2613 fence: &[usize],
2614 base: usize,
2615 pipelined: bool,
2616 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2617 let n_embd = self.cfg.n_embd as usize;
2618 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
2619 rt.bind_stage(1)?;
2620 let _st1 = rt.enter(1);
2621 let e1 = rt.engine(1, e);
2622 let pos_d = e1.htod_i32(&pos)?;
2623 let x = rt.rx(0, slot, t * n_embd)?;
2624 let _overlap = pipelined.then(crate::pp::enter_prime_pipe_stage);
2625 self.prime_layers(e1, x, fence[1], fence[2], &pos_d, t, base, cache, seq_end)
2626 }
2627
2628 pub fn prime_chunk_captured(
2644 &self,
2645 e: &Engine,
2646 x_in: &CudaSlice<f32>,
2647 pos_d: &CudaSlice<i32>,
2648 t: usize,
2649 cache: &mut Cache,
2650 len_d: &CudaSlice<i32>,
2651 logits_out: &mut CudaSlice<f32>,
2652 h_seed_out: &mut CudaSlice<f32>,
2653 ) -> Result<(), Box<dyn std::error::Error>> {
2654 let cfg = &self.cfg;
2655 let n_embd = cfg.n_embd as usize;
2656 let eps = cfg.rms_eps;
2657 let f16fuse = crate::f16_ffi::pp_f16_enabled() && t >= 16;
2658 let mut x = e.uninit(t * n_embd)?;
2659 e.copy_into(&mut x, 0, x_in, t * n_embd)?;
2660 for (il, layer) in self.layers.iter().enumerate() {
2661 let mut h = e.uninit(t * n_embd)?;
2662 let mut hx16: Option<CudaSlice<u8>> = None;
2663 if f16fuse {
2664 let mut b16 = e.alloc_u8_uninit(t * n_embd * 2)?;
2665 e.rms_norm_f16out(
2666 &x,
2667 layer.attn_norm.float_data(),
2668 &mut h,
2669 &mut b16,
2670 n_embd,
2671 t,
2672 eps,
2673 )?;
2674 hx16 = Some(b16);
2675 } else {
2676 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
2677 }
2678 let mixed = match &layer.mixer {
2679 Mixer::Full(fa) => {
2683 self.full_attn_prime(e, fa, &h, hx16.as_ref(), pos_d, t, cache, il, t)?
2684 }
2685 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
2686 Mixer::Linear(la) => {
2687 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
2688 let g4 = match hx16.as_ref() {
2689 Some(xh) => e.matmul_group_xh(&ws, &h, xh, t)?,
2690 None => e.matmul_group(&ws, &h, t)?,
2691 };
2692 self.linear_attn_prime_core_pad(e, la, g4, t, cache, il, Some(len_d))?
2693 }
2694 };
2695 let mut x1 = e.uninit(t * n_embd)?;
2696 e.add(&x, &mixed, &mut x1, t * n_embd)?;
2697 let mut z = e.uninit(t * n_embd)?;
2698 let mut zx16: Option<CudaSlice<u8>> = None;
2699 if f16fuse {
2700 let mut b16 = e.alloc_u8_uninit(t * n_embd * 2)?;
2701 e.rms_norm_f16out(
2702 &x1,
2703 layer.post_attn_norm.float_data(),
2704 &mut z,
2705 &mut b16,
2706 n_embd,
2707 t,
2708 eps,
2709 )?;
2710 zx16 = Some(b16);
2711 } else {
2712 e.rms_norm(
2713 &x1,
2714 layer.post_attn_norm.float_data(),
2715 &mut z,
2716 n_embd,
2717 t,
2718 eps,
2719 )?;
2720 }
2721 let ffn_out = match &layer.ffn {
2722 crate::hybrid::Ffn::Dense {
2723 ffn_gate,
2724 ffn_up,
2725 ffn_down,
2726 } => {
2727 let n_ff = ffn_gate.out_features();
2728 let mut g2 = match &zx16 {
2729 Some(xh) => e.matmul_group_xh(&[ffn_gate, ffn_up], &z, xh, t)?,
2730 None => e.matmul_group(&[ffn_gate, ffn_up], &z, t)?,
2731 };
2732 let up = g2.pop().unwrap();
2733 let gate = g2.pop().unwrap();
2734 let mut act = e.uninit(t * n_ff)?;
2735 Self::ffn_act_lim(
2737 e,
2738 &self.cfg,
2739 &gate,
2740 &up,
2741 1.0,
2742 1.0,
2743 self.cfg.clamp_shexp_at(il as u32),
2744 &mut act,
2745 t * n_ff,
2746 )?;
2747 e.matmul(ffn_down, &act, t)?
2748 }
2749 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il_prefill(e, m, &z, t, il as u16)?,
2750 };
2751 let mut x2 = e.uninit(t * n_embd)?;
2752 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
2753 x = x2;
2754 }
2755 if !crate::spec::spec_hpost() {
2757 e.row_gather_dev(&x, h_seed_out, len_d, n_embd)?;
2758 }
2759 let mut hn = e.uninit(t * n_embd)?;
2760 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
2761 if crate::spec::spec_hpost() {
2762 e.row_gather_dev(&hn, h_seed_out, len_d, n_embd)?;
2763 }
2764 let mut hlast = e.uninit(n_embd)?;
2765 e.row_gather_dev(&hn, &mut hlast, len_d, n_embd)?;
2766 let logits = e.matmul(&self.output, &hlast, 1)?;
2767 let nv = logits.len();
2768 e.copy_into(logits_out, 0, &logits, nv)?;
2769 Ok(())
2770 }
2771
2772 fn step35_prime_batch_on() -> bool {
2773 std::env::var("MEMRA_STEP35_PRIME_BATCH").as_deref() != Ok("0")
2774 }
2775
2776 #[allow(clippy::too_many_arguments)]
2779 #[allow(clippy::too_many_arguments)]
2784 fn step35_prime_batch_layers(
2785 &self,
2786 e: &Engine,
2787 mut x: CudaSlice<f32>,
2788 lo: usize,
2789 hi: usize,
2790 ts: &[usize],
2791 offs: &[usize],
2792 seq_ends: &[usize],
2793 pos_ds: &[CudaSlice<i32>],
2794 caches: &mut [&mut Cache],
2795 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2796 let cfg = &self.cfg;
2797 let n_embd = cfg.n_embd as usize;
2798 let eps = cfg.rms_eps;
2799 let b = ts.len();
2800 let total: usize = ts.iter().sum();
2801 let f16fuse = crate::f16_ffi::pp_f16_enabled() && total >= 16;
2802
2803 let split = |e: &Engine,
2804 y: &CudaSlice<f32>,
2805 dim: usize|
2806 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
2807 let mut out = Vec::with_capacity(b);
2808 for s in 0..b {
2809 let mut ys = e.uninit(ts[s] * dim)?;
2810 e.copy_view_into(
2811 &mut ys,
2812 0,
2813 &y.slice(offs[s] * dim..(offs[s] + ts[s]) * dim),
2814 ts[s] * dim,
2815 )?;
2816 out.push(ys);
2817 }
2818 Ok(out)
2819 };
2820
2821 let prof = std::env::var("MEMRA_PRIME_PROF").as_deref() == Ok("1");
2826 let mut ph = [0f64; 4]; let mut mark = |e: &Engine, acc: usize, t0: &mut std::time::Instant, ph: &mut [f64; 4]| {
2828 if prof {
2829 let _ = e.stream().synchronize();
2830 ph[acc] += t0.elapsed().as_secs_f64() * 1e3;
2831 *t0 = std::time::Instant::now();
2832 }
2833 };
2834 let mut pt = std::time::Instant::now();
2835 for il in lo..hi {
2836 let layer = &self.layers[il];
2837 let Mixer::Full(fa) = &layer.mixer else {
2838 return Err(format!("step35 layer {il} is not full-attn — corrupt config").into());
2839 };
2840
2841 let mut h = e.uninit(total * n_embd)?;
2842 let mut hx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
2843 if f16fuse {
2844 e.rms_norm_f16out(
2845 &x,
2846 layer.attn_norm.float_data(),
2847 &mut h,
2848 &mut hx16,
2849 n_embd,
2850 total,
2851 eps,
2852 )?;
2853 } else {
2854 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, total, eps)?;
2855 }
2856
2857 let gate_w = fa
2861 .attn_gate
2862 .as_ref()
2863 .ok_or("step35 layer is missing attn_gate.weight (head-wise attention gate)")?;
2864 let mut g4 = if f16fuse {
2865 e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv, gate_w], &h, &hx16, total)?
2866 } else {
2867 e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv, gate_w], &h, total)?
2868 };
2869 let gate = g4.pop().unwrap();
2870 let mut parts: Vec<Vec<CudaSlice<f32>>> =
2871 (0..b).map(|_| Vec::with_capacity(3)).collect();
2872 for (w, y) in [&fa.wq, &fa.wk, &fa.wv].iter().zip(g4) {
2873 for (s, ys) in split(e, &y, w.out_features())?.into_iter().enumerate() {
2874 parts[s].push(ys);
2875 }
2876 }
2877 let gates = split(e, &gate, gate_w.out_features())?;
2878 let geometry = self.step35_geom(il);
2879 let hd = geometry.head_dim_k as usize;
2880 let nh = geometry.n_head as usize;
2881 let mut ag_cat = e.uninit(total * nh * hd)?;
2882 for (s, (g3s, gate)) in parts.into_iter().zip(gates).enumerate() {
2883 mark(e, 0, &mut pt, &mut ph);
2884 let ag = self.step35_attn_pre_wo(
2885 e,
2886 fa,
2887 g3s,
2888 None,
2889 Some(&gate),
2890 &pos_ds[s],
2891 ts[s],
2892 Some(&mut *caches[s]),
2893 il,
2894 seq_ends[s],
2895 )?;
2896 e.copy_into(&mut ag_cat, offs[s] * nh * hd, &ag, ts[s] * nh * hd)?;
2897 }
2898 mark(e, 1, &mut pt, &mut ph);
2899 let mixed = e.matmul(&fa.wo, &ag_cat, total)?;
2900
2901 let mut x1 = e.uninit(total * n_embd)?;
2902 let mut z = e.uninit(total * n_embd)?;
2903 let mut zx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
2904 if f16fuse {
2905 e.add_rms_norm_f16out(
2906 &x,
2907 &mixed,
2908 layer.post_attn_norm.float_data(),
2909 &mut x1,
2910 &mut z,
2911 &mut zx16,
2912 n_embd,
2913 total,
2914 eps,
2915 )?;
2916 } else {
2917 e.add(&x, &mixed, &mut x1, total * n_embd)?;
2918 e.rms_norm(
2919 &x1,
2920 layer.post_attn_norm.float_data(),
2921 &mut z,
2922 n_embd,
2923 total,
2924 eps,
2925 )?;
2926 }
2927
2928 mark(e, 2, &mut pt, &mut ph);
2929 let ffn_out = match &layer.ffn {
2930 crate::hybrid::Ffn::Dense {
2931 ffn_gate,
2932 ffn_up,
2933 ffn_down,
2934 } => {
2935 let n_ff = ffn_gate.out_features();
2936 let mut g2 = if f16fuse {
2937 e.matmul_group_xh(&[ffn_gate, ffn_up], &z, &zx16, total)?
2938 } else {
2939 e.matmul_group(&[ffn_gate, ffn_up], &z, total)?
2940 };
2941 let up = g2.pop().unwrap();
2942 let gate = g2.pop().unwrap();
2943 let mut act = e.uninit(total * n_ff)?;
2944 let d_lim = cfg.clamp_shexp_at(il as u32);
2945 if Self::f16out_on(e, total) && cfg.m3.is_none() && d_lim.is_none() {
2946 let mut a16 = e.alloc_u8_uninit(total * n_ff * 2)?;
2947 e.silu_mul_f16out(&gate, &up, &mut act, &mut a16, total * n_ff)?;
2948 match e.try_f16_gemm_pre(ffn_down, &a16, total)? {
2949 Some(y) => y,
2950 None => e.matmul(ffn_down, &act, total)?,
2951 }
2952 } else {
2953 Self::ffn_act_lim(
2954 e,
2955 cfg,
2956 &gate,
2957 &up,
2958 1.0,
2959 1.0,
2960 d_lim,
2961 &mut act,
2962 total * n_ff,
2963 )?;
2964 e.matmul(ffn_down, &act, total)?
2965 }
2966 }
2967 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il(e, m, &z, total, il as u16)?,
2968 };
2969 let mut x2 = e.uninit(total * n_embd)?;
2970 e.add(&x1, &ffn_out, &mut x2, total * n_embd)?;
2971 x = x2;
2972 mark(e, 3, &mut pt, &mut ph);
2973 }
2974 if prof {
2975 eprintln!(
2976 "[prime-prof] t={total} layers={} norm+qkv={:.0}ms attn={:.0}ms o_proj={:.0}ms moe={:.0}ms",
2977 hi - lo,
2978 ph[0],
2979 ph[1],
2980 ph[2],
2981 ph[3]
2982 );
2983 }
2984 Ok(x)
2985 }
2986
2987 fn step35_prime_batch_epilogue(
2988 &self,
2989 e: &Engine,
2990 x: CudaSlice<f32>,
2991 ts: &[usize],
2992 offs: &[usize],
2993 caches: &mut [&mut Cache],
2994 ) -> Result<Vec<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
2995 let n_embd = self.cfg.n_embd as usize;
2996 let total: usize = ts.iter().sum();
2997 let mut hn = e.uninit(total * n_embd)?;
2998 e.rms_norm(
2999 &x,
3000 self.output_norm.float_data(),
3001 &mut hn,
3002 n_embd,
3003 total,
3004 self.cfg.rms_eps,
3005 )?;
3006
3007 let hidden_src = if crate::spec::spec_hpost() { &hn } else { &x };
3008 let mut out = Vec::with_capacity(ts.len());
3009 for s in 0..ts.len() {
3010 let mut hidden = e.uninit(ts[s] * n_embd)?;
3011 e.copy_view_into(
3012 &mut hidden,
3013 0,
3014 &hidden_src.slice(offs[s] * n_embd..(offs[s] + ts[s]) * n_embd),
3015 ts[s] * n_embd,
3016 )?;
3017 let last0 = (offs[s] + ts[s] - 1) * n_embd;
3018 let mut h_seed = e.uninit(n_embd)?;
3019 e.copy_view_into(
3020 &mut h_seed,
3021 0,
3022 &hidden_src.slice(last0..last0 + n_embd),
3023 n_embd,
3024 )?;
3025 let mut hlast = e.uninit(n_embd)?;
3027 e.copy_view_into(&mut hlast, 0, &hn.slice(last0..last0 + n_embd), n_embd)?;
3028 let logits = e.dtoh(&e.matmul(&self.output, &hlast, 1)?)?;
3029 caches[s].pos += ts[s];
3030 out.push((logits, h_seed, hidden));
3031 }
3032 Ok(out)
3033 }
3034
3035 fn step35_prime_cache_batch(
3040 &self,
3041 e: &Engine,
3042 prompts: &[&[u32]],
3043 caches: &mut [&mut Cache],
3044 seq_ends: &[usize],
3045 ) -> Result<Vec<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
3046 assert_eq!(
3047 seq_ends.len(),
3048 prompts.len(),
3049 "step35 batched prime: one seq_end per sequence"
3050 );
3051 validate_step_prime_batch_modes(
3052 step_tp_prefill_enabled()?,
3053 step_ep_grouped_prefill_enabled()?,
3054 )?;
3055 if crate::pp::pp_host_bounce_active()
3056 && (!crate::pp::prime_pp_on() || crate::pp::pp_cuts(self.layers.len()).is_none())
3057 {
3058 return Err(
3059 "step35_prime_cache_batch: MEMRA_PP_HOST_BOUNCE=1 requires a valid prime \
3060 stage split; refusing an unsplit remote-weight walk"
3061 .into(),
3062 );
3063 }
3064 if !Self::step35_prime_batch_on() {
3065 return Err("step35 batched prime is disabled (MEMRA_STEP35_PRIME_BATCH=0)".into());
3066 }
3067 if prompts.len() > 1 && caches.iter().any(|c| c.pos != 0) {
3072 return Err(
3073 "step35 batched prime supports continuation only at B=1; a cross-request batch \
3074 at mixed positions requires per-request queued_after"
3075 .into(),
3076 );
3077 }
3078
3079 let ts: Vec<usize> = prompts.iter().map(|p| p.len()).collect();
3080 for &t in &ts {
3081 assert!(
3082 t >= PRIME_MIN_T,
3083 "step35 batched prime needs T >= {PRIME_MIN_T}"
3084 );
3085 }
3086 for (s, c) in caches.iter().enumerate() {
3087 assert!(
3090 c.pos + ts[s] <= c.max_ctx,
3091 "step35 batched prime exceeds cache max_ctx"
3092 );
3093 assert!(
3094 seq_ends[s] >= c.pos + ts[s],
3095 "step35 batched prime: seq_end must cover this chunk"
3096 );
3097 }
3098 let legacy_tsend = std::env::var("MEMRA_STEP35_PRIME_BATCH_TSEND").as_deref() == Ok("1");
3105 let seq_ends_eff: Vec<usize> = if legacy_tsend {
3106 ts.clone()
3107 } else {
3108 seq_ends.to_vec()
3109 };
3110 let offs: Vec<usize> = ts
3111 .iter()
3112 .scan(0usize, |a, &t| {
3113 let o = *a;
3114 *a += t;
3115 Some(o)
3116 })
3117 .collect();
3118 let total: usize = ts.iter().sum();
3119 let payload = total * self.cfg.n_embd as usize;
3120 let cat_tokens: Vec<u32> = prompts.iter().flat_map(|p| p.iter().copied()).collect();
3121 let positions: Vec<Vec<i32>> = ts
3127 .iter()
3128 .zip(caches.iter())
3129 .map(|(&t, c)| {
3130 let base = c.pos as i32;
3131 (0..t as i32).map(|i| base + i).collect()
3132 })
3133 .collect();
3134 let upload_positions =
3135 |e: &Engine| -> Result<Vec<CudaSlice<i32>>, Box<dyn std::error::Error>> {
3136 positions
3137 .iter()
3138 .map(|p| e.htod_i32(p))
3139 .collect::<Result<_, _>>()
3140 };
3141
3142 static ONCE: std::sync::Once = std::sync::Once::new();
3143 ONCE.call_once(|| {
3144 eprintln!(
3145 "[step35-prime-batch] first concat prime: B={} tokens={total}",
3146 prompts.len()
3147 );
3148 });
3149
3150 let out = if !crate::pp::pp2_streams_off() && crate::pp::prime_pp_on() {
3151 if let Some(fence) = crate::pp::pp_cuts(self.layers.len()) {
3152 let rt = crate::pp::PpNRt::get(e)?;
3153 let n_st = fence.len() - 1;
3154 assert_eq!(
3155 rt.n_stages(),
3156 n_st,
3157 "step35 prime batch stage count mismatch"
3158 );
3159 let caller_stream = e.stream();
3160 rt.fence_stages_behind(&caller_stream)?;
3161
3162 let mut slot = {
3163 let _st0 = rt.enter(0);
3164 let e0 = rt.engine(0, e);
3165 let pos_ds = upload_positions(e0)?;
3166 let x = self.embed(e0, &cat_tokens)?;
3167 let x = self.step35_prime_batch_layers(
3168 e0,
3169 x,
3170 fence[0],
3171 fence[1],
3172 &ts,
3173 &offs,
3174 &seq_ends_eff,
3175 &pos_ds,
3176 caches,
3177 )?;
3178 rt.tx(0, &x, payload)?
3179 };
3180 for s in 1..n_st - 1 {
3181 let _st = rt.enter(s);
3182 let es = rt.engine(s, e);
3183 let pos_ds = upload_positions(es)?;
3184 let x = rt.rx(s - 1, slot, payload)?;
3185 let x = self.step35_prime_batch_layers(
3186 es,
3187 x,
3188 fence[s],
3189 fence[s + 1],
3190 &ts,
3191 &offs,
3192 &seq_ends_eff,
3193 &pos_ds,
3194 caches,
3195 )?;
3196 slot = rt.tx(s, &x, payload)?;
3197 }
3198
3199 let _stl = rt.enter(n_st - 1);
3200 let el = rt.engine(n_st - 1, e);
3201 let pos_ds = upload_positions(el)?;
3202 let x = rt.rx(n_st - 2, slot, payload)?;
3203 let x = self.step35_prime_batch_layers(
3204 el,
3205 x,
3206 fence[n_st - 1],
3207 fence[n_st],
3208 &ts,
3209 &offs,
3210 &seq_ends_eff,
3211 &pos_ds,
3212 caches,
3213 )?;
3214 let out = self.step35_prime_batch_epilogue(el, x, &ts, &offs, caches)?;
3215 rt.publish_to(n_st - 1, &caller_stream)?;
3216 crate::pp::STEP35_PRIME_BATCH_SPLITS
3217 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
3218 out
3219 } else {
3220 let pos_ds = upload_positions(e)?;
3221 let x = self.embed(e, &cat_tokens)?;
3222 let x = self.step35_prime_batch_layers(
3223 e,
3224 x,
3225 0,
3226 self.layers.len(),
3227 &ts,
3228 &offs,
3229 &seq_ends_eff,
3230 &pos_ds,
3231 caches,
3232 )?;
3233 self.step35_prime_batch_epilogue(e, x, &ts, &offs, caches)?
3234 }
3235 } else {
3236 let pos_ds = upload_positions(e)?;
3237 let x = self.embed(e, &cat_tokens)?;
3238 let x = self.step35_prime_batch_layers(
3239 e,
3240 x,
3241 0,
3242 self.layers.len(),
3243 &ts,
3244 &offs,
3245 &seq_ends_eff,
3246 &pos_ds,
3247 caches,
3248 )?;
3249 self.step35_prime_batch_epilogue(e, x, &ts, &offs, caches)?
3250 };
3251 crate::pp::STEP35_PRIME_BATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
3252 Ok(out)
3253 }
3254
3255 pub fn prime_cache_batch(
3272 &self,
3273 e: &Engine,
3274 prompts: &[&[u32]],
3275 caches: &mut [&mut Cache],
3276 ) -> Result<Vec<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
3277 if crate::pp::pp_cuts(self.layers.len()).is_some()
3278 && !self.rewrite_allowed(memra_gguf::execution_manifest::RewriteSurface::Pipeline)
3279 {
3280 return Err("pipeline rewrite is not qualified for batched prime".into());
3281 }
3282 if !self.rewrite_allowed(memra_gguf::execution_manifest::RewriteSurface::CarriedPrime) {
3283 if !self.rewrite_allowed(memra_gguf::execution_manifest::RewriteSurface::DecodeEager) {
3284 return Err("neither batched-prime nor eager rewrite is qualified".into());
3285 }
3286 if prompts.len() != caches.len() {
3287 return Err("prime fallback prompt/cache shape mismatch".into());
3288 }
3289 static ONCE: std::sync::Once = std::sync::Once::new();
3290 ONCE.call_once(|| {
3291 eprintln!(
3292 "[rewrite] carried-prime.v1 unqualified; using individual native eager primes"
3293 );
3294 });
3295 return prompts
3296 .iter()
3297 .copied()
3298 .zip(caches.iter_mut())
3299 .map(|(prompt, cache)| self.prime_cache(e, prompt, cache, 0))
3300 .collect();
3301 }
3302 let cfg = &self.cfg;
3303 let n_embd = cfg.n_embd as usize;
3304 let eps = cfg.rms_eps;
3305 let b = prompts.len();
3306 assert!(b >= 1 && b == caches.len());
3307 let pos0s: Vec<usize> = caches.iter().map(|c| c.pos).collect();
3308 let carried = pos0s.iter().any(|&p| p > 0);
3309 if self.uses_gemma_program() {
3315 return Err(
3316 "prime_cache_batch: gemma4 has no batched prime core (per-layer \
3317 swa/global geometry, softcapped head) — use gemma4_prime per sequence"
3318 .into(),
3319 );
3320 }
3321 if self.uses_sliding_gated_moe_program() {
3324 let seq_ends: Vec<usize> = caches
3329 .iter()
3330 .zip(prompts.iter())
3331 .map(|(c, p)| c.pos + p.len())
3332 .collect();
3333 return self.step35_prime_cache_batch(e, prompts, caches, &seq_ends);
3334 }
3335 let ts: Vec<usize> = prompts.iter().map(|p| p.len()).collect();
3336 for &t in &ts {
3337 assert!(
3338 t >= PRIME_MIN_T,
3339 "prime_cache_batch needs T >= {PRIME_MIN_T}"
3340 );
3341 }
3342 for (s, c) in caches.iter().enumerate() {
3343 assert!(
3344 c.pos + ts[s] <= c.max_ctx,
3345 "prime_cache_batch: prompt exceeds cache max_ctx"
3346 );
3347 }
3348 let total: usize = ts.iter().sum();
3349 let offs: Vec<usize> = ts
3350 .iter()
3351 .scan(0usize, |a, &t| {
3352 let o = *a;
3353 *a += t;
3354 Some(o)
3355 })
3356 .collect();
3357 let pos_ds: Vec<CudaSlice<i32>> = ts
3359 .iter()
3360 .zip(&pos0s)
3361 .map(|(&t, &p0)| e.htod_i32(&(p0 as i32..(p0 + t) as i32).collect::<Vec<_>>()))
3362 .collect::<Result<_, _>>()?;
3363 let split = |e: &Engine,
3365 y: &CudaSlice<f32>,
3366 dim: usize|
3367 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
3368 let mut out = Vec::with_capacity(b);
3369 for s in 0..b {
3370 let mut ys = e.uninit(ts[s] * dim)?;
3371 e.copy_view_into(
3372 &mut ys,
3373 0,
3374 &y.slice(offs[s] * dim..(offs[s] + ts[s]) * dim),
3375 ts[s] * dim,
3376 )?;
3377 out.push(ys);
3378 }
3379 Ok(out)
3380 };
3381
3382 let cat_tokens: Vec<u32> = prompts.iter().flat_map(|p| p.iter().copied()).collect();
3383 let mut x = self.embed(e, &cat_tokens)?; for (il, layer) in self.layers.iter().enumerate() {
3385 let mut h = e.uninit(total * n_embd)?;
3386 let mut hx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
3387 e.rms_norm_f16out(
3388 &x,
3389 layer.attn_norm.float_data(),
3390 &mut h,
3391 &mut hx16,
3392 n_embd,
3393 total,
3394 eps,
3395 )?;
3396 let mut mixed = e.uninit(total * n_embd)?;
3398 match &layer.mixer {
3399 Mixer::Full(fa) => {
3400 let g3 = e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], &h, &hx16, total)?;
3401 let geometry = self.cfg.full_attention_geometry_at(il as u32);
3407 let (n_head, n_head_kv, head_dim) = (
3408 geometry.n_head as usize,
3409 geometry.n_head_kv as usize,
3410 geometry.head_dim_k as usize,
3411 );
3412 let fa_scale = geometry.attention_scale();
3413 let use_favl = !carried
3414 && (2..=8).contains(&b)
3415 && (head_dim == 256 || head_dim == 128)
3416 && geometry.attention_gate == memra_gguf::config::AttentionGateKind::FusedQ
3417 && std::env::var("MEMRA_NOFA").is_err()
3418 && std::env::var("MEMRA_FA_FLOOR").is_err()
3419 && std::env::var("MEMRA_FA_PP_W2").as_deref() != Ok("1")
3420 && std::env::var("MEMRA_FA_BF16KV").as_deref() != Ok("0")
3421 && std::env::var("MEMRA_FA_VL").as_deref() != Ok("0");
3422 if use_favl {
3423 let (qf_w, kf_w, vf_w) = (
3424 fa.wq.out_features(),
3425 fa.wk.out_features(),
3426 fa.wv.out_features(),
3427 );
3428 memra_gguf::config::check_fused_q_gate_extent(qf_w, head_dim, n_head, 1)?;
3433 struct APre {
3434 q: CudaSlice<f32>,
3435 gate: Option<CudaSlice<f32>>,
3436 qn: CudaSlice<f32>,
3437 kn: CudaSlice<f32>,
3438 }
3439 let mut aps = Vec::with_capacity(b);
3440 for &t in ts.iter().take(b) {
3441 aps.push(APre {
3442 q: e.uninit(t * n_head * head_dim)?,
3443 gate: Some(e.uninit(t * n_head * head_dim)?),
3444 qn: e.uninit(t * n_head * head_dim)?,
3445 kn: e.uninit(t * n_head_kv * head_dim)?,
3446 });
3447 }
3448 let (kv_dim_k, kv_dim_v, ktb, vtb) = {
3449 let kvl = caches[0].kv[il].as_ref().unwrap();
3450 (kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes)
3451 };
3452 let pargs: Vec<crate::AttnPreVl> = (0..b)
3453 .map(|s| {
3454 let (o, t) = (offs[s], ts[s]);
3455 let kvl = caches[s].kv[il].as_ref().unwrap();
3456 assert!(
3457 kvl.len == 0 && kvl.len + t <= caches[s].max_ctx,
3458 "prime_cache_batch attn vl: fresh + capacity"
3459 );
3460 crate::AttnPreVl {
3461 qf: e.addr_f32v(&g3[0].slice(o * qf_w..(o + t) * qf_w)),
3462 kf: e.addr_f32v(&g3[1].slice(o * kf_w..(o + t) * kf_w)),
3463 vf: e.addr_f32v(&g3[2].slice(o * vf_w..(o + t) * vf_w)),
3464 q: e.addr_f32(&aps[s].q),
3465 gate: e.addr_f32(aps[s].gate.as_ref().unwrap()),
3466 qn: e.addr_f32(&aps[s].qn),
3467 kn: e.addr_f32(&aps[s].kn),
3468 kc: e.addr_u8(&kvl.k),
3469 vc: e.addr_u8(&kvl.v),
3470 t: t as i32,
3471 pad: 0,
3472 }
3473 })
3474 .collect();
3475 e.attn_pre_vl8(
3476 &pargs,
3477 fa.q_norm.float_data(),
3478 fa.k_norm.float_data(),
3479 head_dim,
3480 geometry.n_rot as usize,
3481 n_head,
3482 n_head_kv,
3483 self.cfg.rms_eps,
3484 geometry.rope_base,
3485 1.0,
3486 kv_dim_k,
3487 kv_dim_v,
3488 ktb,
3489 vtb,
3490 )?;
3491 for s in 0..b {
3492 let kvl = caches[s].kv[il].as_mut().unwrap();
3493 kvl.len += ts[s];
3494 let new_len = kvl.len as i32;
3495 e.set_i32_one(&mut kvl.len_d, new_len)?;
3496 }
3497 let mut attns = Vec::with_capacity(b);
3498 let mut mirrors = Vec::with_capacity(b);
3499 for &t in ts.iter().take(b) {
3500 attns.push(e.uninit(t * n_head * head_dim)?);
3501 let n = t * n_head_kv * head_dim;
3502 mirrors.push((e.alloc_u8_uninit(n * 2)?, e.alloc_u8_uninit(n * 2)?));
3503 }
3504 let fa3_on = match std::env::var("MEMRA_FA3").as_deref() {
3507 Ok("0") => false,
3508 Ok("1") => {
3512 crate::refuse_portable_force(
3513 "MEMRA_FA3=1",
3514 "the sm_90a fa3/bf16 kernels",
3515 );
3516 true
3517 }
3518 _ => cfg!(memra_hopper_mma),
3519 };
3520 if fa3_on {
3521 let mut q16s = Vec::with_capacity(b);
3522 let mut v16s = Vec::with_capacity(b);
3523 for s in 0..b {
3524 let t = ts[s];
3525 let mut q16 = e.alloc_u8_uninit(t * n_head * head_dim * 2)?;
3526 e.f32_to_bf16_into(&aps[s].qn, &mut q16, t * n_head * head_dim)?;
3527 let mut k16 = e.alloc_u8_uninit(t * n_head_kv * head_dim * 2)?;
3528 e.f32_to_bf16_into(&aps[s].kn, &mut k16, t * n_head_kv * head_dim)?;
3529 let mut v16 = e.alloc_u8_uninit(t * n_head_kv * head_dim * 2)?;
3530 e.f32_to_bf16_v(
3531 &g3[2].slice(offs[s] * vf_w..(offs[s] + t) * vf_w),
3532 &mut v16,
3533 t * n_head_kv * head_dim,
3534 )?;
3535 q16s.push(q16);
3536 v16s.push((k16, v16));
3537 }
3538 let mut qp = [core::ptr::null::<core::ffi::c_void>(); 8];
3539 let mut kp = qp;
3540 let mut vp = qp;
3541 let mut op = [core::ptr::null_mut::<f32>(); 8];
3542 let mut tsv = [0i32; 8];
3543 for s in 0..b {
3544 qp[s] = e.addr_u8(&q16s[s]) as *const core::ffi::c_void;
3545 kp[s] = e.addr_u8(&v16s[s].0) as *const core::ffi::c_void;
3546 vp[s] = e.addr_u8(&v16s[s].1) as *const core::ffi::c_void;
3547 op[s] = e.addr_f32(&attns[s]) as *mut f32;
3548 tsv[s] = ts[s] as i32;
3549 }
3550 let rc = unsafe {
3551 crate::fa3_vl_raw(
3552 qp.as_ptr(),
3553 kp.as_ptr(),
3554 vp.as_ptr(),
3555 op.as_ptr(),
3556 tsv.as_ptr(),
3557 b as i32,
3558 n_head as i32,
3559 n_head_kv as i32,
3560 head_dim as i32,
3561 fa_scale,
3562 e.stream().cu_stream() as *mut core::ffi::c_void,
3563 )
3564 };
3565 if rc != 0 {
3566 return Err(format!("memra_fa3_vl rc={rc}").into());
3567 }
3568 } else {
3569 let fargs: Vec<crate::FaSeqVl> = (0..b)
3570 .map(|s| crate::FaSeqVl {
3571 q: e.addr_f32(&aps[s].qn),
3572 k16: e.addr_u8(&mirrors[s].0),
3573 v16: e.addr_u8(&mirrors[s].1),
3574 o: e.addr_f32(&attns[s]),
3575 kf: e.addr_f32(&aps[s].kn),
3576 vf: e.addr_f32v(
3577 &g3[2].slice(offs[s] * vf_w..(offs[s] + ts[s]) * vf_w),
3578 ),
3579 t: ts[s] as i32,
3580 pad: 0,
3581 })
3582 .collect();
3583 e.fa_prefill_vl8(&fargs, head_dim, n_head, n_head_kv, fa_scale)?;
3584 }
3585 for (s, attn) in attns.into_iter().enumerate() {
3586 let (attn_g, ag16) = self.full_attn_prime_post_fa(
3587 e,
3588 attn,
3589 &aps[s].gate,
3590 ts[s],
3591 n_head,
3592 head_dim,
3593 )?;
3594 let mut done = false;
3595 if let Some(xh) = &ag16 {
3596 done = e.try_f16_gemm_pre_into_off(
3597 &fa.wo,
3598 xh,
3599 ts[s],
3600 &mut mixed,
3601 offs[s] * n_embd,
3602 )?;
3603 }
3604 if !done {
3605 let m = e.matmul(&fa.wo, &attn_g, ts[s])?;
3606 e.copy_into(&mut mixed, offs[s] * n_embd, &m, ts[s] * n_embd)?;
3607 }
3608 }
3609 } else {
3610 let mut parts: Vec<Vec<CudaSlice<f32>>> =
3611 (0..b).map(|_| Vec::new()).collect();
3612 for (w, y) in [&fa.wq, &fa.wk, &fa.wv].iter().zip(g3) {
3613 for (s, ys) in split(e, &y, w.out_features())?.into_iter().enumerate() {
3614 parts[s].push(ys);
3615 }
3616 }
3617 for (s, g3s) in parts.into_iter().enumerate() {
3618 let (attn_g, ag16) = self.full_attn_prime_core_inner(
3620 e, fa, g3s, &pos_ds[s], ts[s], caches[s], il,
3621 )?;
3622 let mut done = false;
3623 if let Some(xh) = &ag16 {
3624 done = e.try_f16_gemm_pre_into_off(
3625 &fa.wo,
3626 xh,
3627 ts[s],
3628 &mut mixed,
3629 offs[s] * n_embd,
3630 )?;
3631 }
3632 if !done {
3633 let m = e.matmul(&fa.wo, &attn_g, ts[s])?;
3634 e.copy_into(&mut mixed, offs[s] * n_embd, &m, ts[s] * n_embd)?;
3635 }
3636 }
3637 }
3638 }
3639 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
3640 Mixer::Linear(la) => {
3641 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
3646 let g4 = e.matmul_group_xh(&ws, &h, &hx16, total)?;
3647 let outs =
3648 self.linear_attn_prime_core_batch(e, la, &g4, &offs, &ts, caches, il)?;
3649 for (s, (gn, gn16)) in outs.into_iter().enumerate() {
3650 let (o, t) = (offs[s], ts[s]);
3651 let mut done = false;
3652 if let Some(xh) = &gn16 {
3653 done = e.try_f16_gemm_pre_into_off(
3654 &la.ssm_out,
3655 xh,
3656 t,
3657 &mut mixed,
3658 o * n_embd,
3659 )?;
3660 }
3661 if !done {
3662 let m = e.matmul(&la.ssm_out, &gn, t)?;
3663 e.copy_into(&mut mixed, o * n_embd, &m, t * n_embd)?;
3664 }
3665 }
3666 }
3667 }
3668 let mut x1 = e.uninit(total * n_embd)?;
3669 let mut z = e.uninit(total * n_embd)?;
3670 let mut zx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
3671 e.add_rms_norm_f16out(
3672 &x,
3673 &mixed,
3674 layer.post_attn_norm.float_data(),
3675 &mut x1,
3676 &mut z,
3677 &mut zx16,
3678 n_embd,
3679 total,
3680 eps,
3681 )?;
3682 let ffn_out = match &layer.ffn {
3683 crate::hybrid::Ffn::Dense {
3684 ffn_gate,
3685 ffn_up,
3686 ffn_down,
3687 } => {
3688 let n_ff = ffn_gate.out_features();
3689 let mut g2 = e.matmul_group_xh(&[ffn_gate, ffn_up], &z, &zx16, total)?;
3690 let up = g2.pop().unwrap();
3691 let gate = g2.pop().unwrap();
3692 let mut act = e.uninit(total * n_ff)?;
3693 let d_lim = self.cfg.clamp_shexp_at(il as u32);
3697 if Self::f16out_on(e, total) && self.cfg.m3.is_none() && d_lim.is_none() {
3698 let mut a16 = e.alloc_u8_uninit(total * n_ff * 2)?;
3699 e.silu_mul_f16out(&gate, &up, &mut act, &mut a16, total * n_ff)?;
3700 match e.try_f16_gemm_pre(ffn_down, &a16, total)? {
3701 Some(y) => y,
3702 None => e.matmul(ffn_down, &act, total)?,
3703 }
3704 } else {
3705 Self::ffn_act_lim(
3706 e,
3707 &self.cfg,
3708 &gate,
3709 &up,
3710 1.0,
3711 1.0,
3712 d_lim,
3713 &mut act,
3714 total * n_ff,
3715 )?;
3716 e.matmul(ffn_down, &act, total)?
3717 }
3718 }
3719 crate::hybrid::Ffn::Moe(m) => {
3720 self.moe_ffn_il_prefill(e, m, &z, total, il as u16)?
3721 }
3722 };
3723 let mut x2 = e.uninit(total * n_embd)?;
3724 e.add(&x1, &ffn_out, &mut x2, total * n_embd)?;
3725 x = x2;
3726 }
3727 let mut hn = e.uninit(total * n_embd)?;
3729 e.rms_norm(
3730 &x,
3731 self.output_norm.float_data(),
3732 &mut hn,
3733 n_embd,
3734 total,
3735 eps,
3736 )?;
3737 let mut hcat = e.uninit(b * n_embd)?;
3743 for s in 0..b {
3744 let last0 = (offs[s] + ts[s] - 1) * n_embd;
3745 e.copy_view_into(
3746 &mut hcat,
3747 s * n_embd,
3748 &hn.slice(last0..last0 + n_embd),
3749 n_embd,
3750 )?;
3751 }
3752 let logits_cat = if b >= 2 {
3753 e.try_f16_gemm(&self.output, &hcat, b)?
3754 } else {
3755 None
3756 };
3757 let logits_host: Option<Vec<f32>> = match &logits_cat {
3758 Some(lc) => Some(e.dtoh(lc)?),
3759 None => None,
3760 };
3761 let n_vocab = self.output.out_features();
3762 let mut hidden_all = if crate::spec::spec_hpost() {
3763 split(e, &hn, n_embd)?
3764 } else {
3765 split(e, &x, n_embd)?
3766 };
3767 let mut out = Vec::with_capacity(b);
3768 for s in 0..b {
3769 let last0 = (offs[s] + ts[s] - 1) * n_embd;
3770 let mut h_seed = e.uninit(n_embd)?;
3771 if !crate::spec::spec_hpost() {
3772 e.copy_view_into(&mut h_seed, 0, &x.slice(last0..last0 + n_embd), n_embd)?;
3773 } else {
3774 e.copy_view_into(&mut h_seed, 0, &hn.slice(last0..last0 + n_embd), n_embd)?;
3775 }
3776 let logits = match &logits_host {
3777 Some(lh) => lh[s * n_vocab..(s + 1) * n_vocab].to_vec(),
3778 None => {
3779 let mut hlast = e.uninit(n_embd)?;
3780 e.copy_view_into(&mut hlast, 0, &hn.slice(last0..last0 + n_embd), n_embd)?;
3781 e.dtoh(&e.matmul(&self.output, &hlast, 1)?)?
3782 }
3783 };
3784 caches[s].pos += ts[s];
3785 out.push((logits, h_seed, hidden_all.remove(0)));
3786 }
3787 Ok(out)
3788 }
3789
3790 #[allow(clippy::too_many_arguments)]
3801 fn full_attn_prime(
3802 &self,
3803 e: &Engine,
3804 fa: &FullAttnLayer,
3805 h: &CudaSlice<f32>,
3806 hx: Option<&CudaSlice<u8>>,
3807 pos_d: &CudaSlice<i32>,
3808 t: usize,
3809 cache: &mut Cache,
3810 il: usize,
3811 seq_end: usize,
3812 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3813 if self.uses_sliding_gated_moe_program() {
3814 return self.step35_attn_prime(e, fa, h, hx, pos_d, t, cache, il, seq_end);
3815 }
3816 let g3 = match hx {
3821 Some(xh) => e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], h, xh, t)?,
3822 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
3823 };
3824 self.full_attn_prime_core(e, fa, g3, pos_d, t, cache, il)
3825 }
3826
3827 fn full_attn_prime_core(
3831 &self,
3832 e: &Engine,
3833 fa: &FullAttnLayer,
3834 g3: Vec<CudaSlice<f32>>,
3835 pos_d: &CudaSlice<i32>,
3836 t: usize,
3837 cache: &mut Cache,
3838 il: usize,
3839 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3840 let (attn_g, ag16) = self.full_attn_prime_core_inner(e, fa, g3, pos_d, t, cache, il)?;
3841 if let Some(xh) = &ag16 {
3842 if let Some(y) = e.try_f16_gemm_pre(&fa.wo, xh, t)? {
3843 return Ok(y);
3844 }
3845 }
3846 Ok(e.matmul(&fa.wo, &attn_g, t)?)
3847 }
3848
3849 fn full_attn_prime_core_inner(
3850 &self,
3851 e: &Engine,
3852 fa: &FullAttnLayer,
3853 g3: Vec<CudaSlice<f32>>,
3854 pos_d: &CudaSlice<i32>,
3855 t: usize,
3856 cache: &mut Cache,
3857 il: usize,
3858 ) -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
3859 let cfg = &self.cfg;
3860 let geometry = cfg.full_attention_geometry_at(il as u32);
3861 let n_head = geometry.n_head as usize;
3862 let n_head_kv = geometry.n_head_kv as usize;
3863 let head_dim = geometry.head_dim_k as usize;
3864 let scale = geometry.attention_scale();
3865 let (pre, base_len) = self.full_attn_prime_pre_fa(e, fa, g3, pos_d, t, cache, il)?;
3866 let AttnPre { q, k, v, gate } = pre;
3867 let mut attn = e.uninit(t * n_head * head_dim)?;
3868 self.full_attn_prime_fa_dispatch(
3869 e, &q, &k, &v, &mut attn, base_len, t, cache, il, head_dim, n_head, n_head_kv, scale,
3870 )?;
3871 self.full_attn_prime_post_fa(e, attn, &gate, t, n_head, head_dim)
3872 }
3873
3874 #[allow(clippy::type_complexity)]
3878 fn full_attn_prime_pre_fa(
3879 &self,
3880 e: &Engine,
3881 fa: &FullAttnLayer,
3882 mut g3: Vec<CudaSlice<f32>>,
3883 pos_d: &CudaSlice<i32>,
3884 t: usize,
3885 cache: &mut Cache,
3886 il: usize,
3887 ) -> Result<(AttnPre, usize), Box<dyn std::error::Error>> {
3888 let cfg = &self.cfg;
3889 let geometry = cfg.full_attention_geometry_at(il as u32);
3890 let n_head = geometry.n_head as usize;
3891 let n_head_kv = geometry.n_head_kv as usize;
3892 let head_dim = geometry.head_dim_k as usize;
3893 let eps = cfg.rms_eps;
3894
3895 let gated = geometry.attention_gate == memra_gguf::config::AttentionGateKind::FusedQ;
3899 let v = g3.pop().unwrap();
3900 let mut k = g3.pop().unwrap();
3901 let qf = g3.pop().unwrap();
3902 let (mut q, gate) = if gated {
3903 let mut q = e.uninit(t * n_head * head_dim)?;
3904 let mut gate = e.uninit(t * n_head * head_dim)?;
3905 e.q_gate_split(&qf, &mut q, &mut gate, head_dim, n_head, t)?;
3906 (q, Some(gate))
3907 } else {
3908 (qf, None)
3909 };
3910
3911 let mut qn = e.uninit(t * n_head * head_dim)?;
3912 e.rms_norm(
3913 &q,
3914 fa.q_norm.float_data(),
3915 &mut qn,
3916 head_dim,
3917 n_head * t,
3918 eps,
3919 )?;
3920 q = qn;
3921 let mut kn = e.uninit(t * n_head_kv * head_dim)?;
3922 e.rms_norm(
3923 &k,
3924 fa.k_norm.float_data(),
3925 &mut kn,
3926 head_dim,
3927 n_head_kv * t,
3928 eps,
3929 )?;
3930 k = kn;
3931 let rope_dims = geometry.n_rot as usize;
3932 e.rope_neox(
3933 &mut q,
3934 pos_d,
3935 head_dim,
3936 rope_dims,
3937 n_head,
3938 t,
3939 geometry.rope_base,
3940 1.0,
3941 )?;
3942 e.rope_neox(
3943 &mut k,
3944 pos_d,
3945 head_dim,
3946 rope_dims,
3947 n_head_kv,
3948 t,
3949 geometry.rope_base,
3950 1.0,
3951 )?;
3952
3953 {
3956 let kvl = cache.kv[il].as_mut().unwrap();
3957 assert!(kvl.len + t <= cache.max_ctx, "prime_cache: KV overflow");
3958 e.append_kv_quantized_rows(
3959 &k,
3960 &v,
3961 &mut kvl.k,
3962 &mut kvl.v,
3963 kvl.len,
3964 t,
3965 kvl.kv_dim_k,
3966 kvl.kv_dim_v,
3967 kvl.k_tok_bytes,
3968 kvl.v_tok_bytes,
3969 crate::Engine::kv_fp8_on(),
3970 )?;
3971 kvl.len += t;
3972 let new_len = kvl.len as i32;
3973 e.set_i32_one(&mut kvl.len_d, new_len)?;
3974 }
3975
3976 let base_len = {
3977 let kvl = cache.kv[il].as_ref().unwrap();
3978 kvl.len - t };
3980 Ok((AttnPre { q, k, v, gate }, base_len))
3981 }
3982
3983 #[allow(clippy::too_many_arguments)]
3990 fn full_attn_prime_fa_dispatch(
3991 &self,
3992 e: &Engine,
3993 q: &CudaSlice<f32>,
3994 k: &CudaSlice<f32>,
3995 v: &CudaSlice<f32>,
3996 attn: &mut CudaSlice<f32>,
3997 base_len: usize,
3998 t: usize,
3999 cache: &mut Cache,
4000 il: usize,
4001 head_dim: usize,
4002 n_head: usize,
4003 n_head_kv: usize,
4004 scale: f32,
4005 ) -> Result<(), Box<dyn std::error::Error>> {
4006 if base_len == 0 && std::env::var("MEMRA_PRIME_F32CHUNK0").as_deref() == Ok("1") {
4019 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
4020 e.sdpa_naive(
4021 q, k, v, attn, head_dim, n_head, n_head_kv, t, t, scale, true,
4022 )?;
4023 } else {
4024 e.fa_prefill(
4025 q, k, v, attn, head_dim, n_head, n_head_kv, t, t, scale, true,
4026 )?;
4027 }
4028 return Ok(());
4029 }
4030 let kvl = cache.kv[il].as_ref().unwrap();
4031 let t_kv = base_len + t;
4032 let k_view = e.view_u8(&kvl.k, t_kv * kvl.k_tok_bytes);
4033 let v_view = e.view_u8(&kvl.v, t_kv * kvl.v_tok_bytes);
4034 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
4038 e.sdpa_naive_quantized_view(
4039 q,
4040 &k_view,
4041 &v_view,
4042 attn,
4043 head_dim,
4044 n_head,
4045 n_head_kv,
4046 t,
4047 t_kv,
4048 scale,
4049 true,
4050 kvl.k_tok_bytes,
4051 kvl.v_tok_bytes,
4052 )?;
4053 return Ok(());
4054 }
4055 let deqw = std::env::var("MEMRA_PRIME_DEQW")
4063 .map(|v| v != "0")
4064 .unwrap_or(true);
4065 if deqw {
4066 e.fa_prefill_view_ws(
4067 q,
4068 &k_view,
4069 &v_view,
4070 attn,
4071 head_dim,
4072 n_head,
4073 n_head_kv,
4074 t,
4075 t_kv,
4076 scale,
4077 true,
4078 kvl.k_tok_bytes,
4079 kvl.v_tok_bytes,
4080 crate::Engine::kv_fp8_on(),
4081 )?;
4082 } else {
4083 e.fa_prefill_view(
4084 q,
4085 &k_view,
4086 &v_view,
4087 attn,
4088 head_dim,
4089 n_head,
4090 n_head_kv,
4091 t,
4092 t_kv,
4093 scale,
4094 true,
4095 kvl.k_tok_bytes,
4096 kvl.v_tok_bytes,
4097 crate::Engine::kv_fp8_on(),
4098 )?;
4099 }
4100 Ok(())
4101 }
4102
4103 fn full_attn_prime_post_fa(
4106 &self,
4107 e: &Engine,
4108 attn: CudaSlice<f32>,
4109 gate: &Option<CudaSlice<f32>>,
4110 t: usize,
4111 n_head: usize,
4112 head_dim: usize,
4113 ) -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
4114 let (attn_g, ag16) = match gate {
4115 Some(gate) => {
4116 let n = t * n_head * head_dim;
4117 let mut ag = e.uninit(n)?;
4118 if Self::f16out_on(e, t) {
4119 let mut a16 = e.alloc_u8_uninit(n * 2)?;
4120 e.sig_mul_f16out(&attn, gate, &mut ag, &mut a16, n)?;
4121 (ag, Some(a16))
4122 } else {
4123 let mut gsig = e.uninit(n)?;
4124 e.sigmoid(gate, &mut gsig, n)?;
4125 e.mul(&attn, &gsig, &mut ag, n)?;
4126 (ag, None)
4127 }
4128 }
4129 None => (attn, None),
4130 };
4131 Ok((attn_g, ag16))
4132 }
4133
4134 fn linear_attn_prime(
4141 &self,
4142 e: &Engine,
4143 la: &LinearAttnLayer,
4144 h: &CudaSlice<f32>,
4145 hx: Option<&CudaSlice<u8>>,
4146 t: usize,
4147 cache: &mut Cache,
4148 il: usize,
4149 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4150 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
4152 let g4 = match hx {
4153 Some(xh) => e.matmul_group_xh(&ws, h, xh, t)?,
4154 None => e.matmul_group(&ws, h, t)?,
4155 };
4156 self.linear_attn_prime_core(e, la, g4, t, cache, il)
4157 }
4158
4159 fn linear_attn_prime_core(
4161 &self,
4162 e: &Engine,
4163 la: &LinearAttnLayer,
4164 mut g4: Vec<CudaSlice<f32>>,
4165 t: usize,
4166 cache: &mut Cache,
4167 il: usize,
4168 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4169 self.linear_attn_prime_core_pad(e, la, g4.drain(..).collect(), t, cache, il, None)
4170 }
4171
4172 #[allow(clippy::too_many_arguments)]
4176 fn linear_attn_prime_core_pad_inner(
4177 &self,
4178 e: &Engine,
4179 la: &LinearAttnLayer,
4180 mut g4: Vec<CudaSlice<f32>>,
4181 t: usize,
4182 cache: &mut Cache,
4183 il: usize,
4184 pad_len: Option<&CudaSlice<i32>>,
4185 ) -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
4186 let geometry = la.geometry;
4188 let d_state = geometry.key_head_dim as usize;
4189 let num_k = geometry.key_heads as usize;
4190 let num_v = geometry.value_heads as usize;
4191 let key_dim = d_state * num_k;
4192 let value_dim = geometry.value_head_dim as usize * num_v;
4193 let conv_dim = key_dim * 2 + value_dim;
4194 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(
4199 e,
4200 la,
4201 &qkv_mixed.slice(0..t * conv_dim),
4202 &z.slice(0..t * value_dim),
4203 &beta_raw.slice(0..t * num_v),
4204 &alpha.slice(0..t * num_v),
4205 t,
4206 cache,
4207 il,
4208 pad_len,
4209 )
4210 }
4211
4212 #[allow(clippy::too_many_arguments)]
4215 fn linear_attn_gdn_prep(
4216 &self,
4217 e: &Engine,
4218 la: &LinearAttnLayer,
4219 qkv_mixed: &cudarc::driver::CudaView<f32>,
4220 beta_raw: &cudarc::driver::CudaView<f32>,
4221 alpha: &cudarc::driver::CudaView<f32>,
4222 t: usize,
4223 cache: &mut Cache,
4224 il: usize,
4225 pad_len: Option<&CudaSlice<i32>>,
4226 ) -> Result<GdnPrep, Box<dyn std::error::Error>> {
4227 let cfg = &self.cfg;
4228 let geometry = la.geometry;
4229 let d_state = geometry.key_head_dim as usize;
4230 let num_k = geometry.key_heads as usize;
4231 let num_v = geometry.value_heads as usize;
4232 let d_conv = geometry.conv_kernel as usize;
4233 let key_dim = d_state * num_k; let value_dim = geometry.value_head_dim as usize * num_v;
4235 let conv_dim = key_dim * 2 + value_dim; let eps = cfg.rms_eps;
4237 debug_assert!(
4238 t >= d_conv - 1,
4239 "stateful conv needs T >= pad (PRIME_MIN_T gates)"
4240 );
4241
4242 let rl = cache.recur[il].as_mut().unwrap();
4247 let hk = Self::gdn_hk(e, t, num_v, num_k);
4248 let conv_fuse = std::env::var("MEMRA_CONV_FUSE").as_deref() != Ok("0");
4249 let hk = if conv_fuse { hk } else { num_v }; let mut q_g = e.uninit(d_state * hk * t)?;
4251 let mut k_g = e.uninit(d_state * hk * t)?;
4252 let mut v_g = e.uninit(d_state * num_v * t)?;
4253 if conv_fuse {
4254 e.ssm_conv1d_gdn_state_pad(
4255 qkv_mixed,
4256 &mut rl.conv_state,
4257 la.ssm_conv1d.float_data(),
4258 &mut q_g,
4259 &mut k_g,
4260 &mut v_g,
4261 conv_dim,
4262 t,
4263 d_conv,
4264 d_state,
4265 num_v,
4266 num_k,
4267 key_dim,
4268 hk,
4269 pad_len,
4270 )?;
4271 } else {
4272 let mut conv_out = e.uninit(conv_dim * t)?; e.ssm_conv1d_tm_state_pad_v(
4274 qkv_mixed,
4275 &mut rl.conv_state,
4276 la.ssm_conv1d.float_data(),
4277 &mut conv_out,
4278 conv_dim,
4279 t,
4280 d_conv,
4281 pad_len,
4282 )?;
4283 e.qkv_to_gdn_repack(
4284 &conv_out, &mut q_g, &mut k_g, &mut v_g, d_state, num_v, num_k, key_dim, t,
4285 )?;
4286 }
4287 let mut q_l2 = e.uninit(d_state * hk * t)?;
4288 let qb16 = if Engine::l2_v2_on(d_state) && e.gdn_wgmma_on(32) {
4292 let mut qb = e.alloc_u8_uninit(d_state * hk * t * 2)?;
4293 e.l2_norm_pp(&q_g, &mut q_l2, Some(&mut qb), d_state, hk * t, eps)?;
4294 Some(qb)
4295 } else {
4296 e.l2_norm_pp(&q_g, &mut q_l2, None, d_state, hk * t, eps)?;
4297 None
4298 };
4299 let mut k_l2 = e.uninit(d_state * hk * t)?;
4300 let kb16 = if Engine::l2_v2_on(d_state) {
4302 let mut kb = e.alloc_u8_uninit(d_state * hk * t * 2)?;
4303 e.l2_norm_pp(&k_g, &mut k_l2, Some(&mut kb), d_state, hk * t, eps)?;
4304 Some(kb)
4305 } else {
4306 e.l2_norm_pp(&k_g, &mut k_l2, None, d_state, hk * t, eps)?;
4307 None
4308 };
4309 let mut beta = e.uninit(t * num_v)?;
4310 e.sigmoid_v(beta_raw, &mut beta, t * num_v)?;
4311 let mut g_log = e.uninit(t * num_v)?;
4312 e.gdn_glog_v(
4313 alpha,
4314 la.ssm_dt.float_data(),
4315 la.ssm_a.float_data(),
4316 &mut g_log,
4317 num_v,
4318 t,
4319 )?;
4320 if let Some(len_d) = pad_len {
4321 e.gdn_pad_mask(&mut beta, &mut g_log, len_d, num_v, t)?;
4322 }
4323 Ok(GdnPrep {
4324 hk,
4325 q_l2,
4326 k_l2,
4327 v_g,
4328 beta,
4329 g_log,
4330 kb16,
4331 qb16,
4332 })
4333 }
4334
4335 #[allow(clippy::too_many_arguments)]
4340 fn linear_attn_prime_core_batch(
4341 &self,
4342 e: &Engine,
4343 la: &LinearAttnLayer,
4344 g4: &[CudaSlice<f32>],
4345 offs: &[usize],
4346 ts: &[usize],
4347 caches: &mut [&mut Cache],
4348 il: usize,
4349 ) -> Result<Vec<(CudaSlice<f32>, Option<CudaSlice<u8>>)>, Box<dyn std::error::Error>> {
4350 let geometry = la.geometry;
4351 let d_state = geometry.key_head_dim as usize;
4352 let num_k = geometry.key_heads as usize;
4353 let num_v = geometry.value_heads as usize;
4354 let d_conv = geometry.conv_kernel as usize;
4355 let key_dim = d_state * num_k;
4356 let value_dim = geometry.value_head_dim as usize * num_v;
4357 let conv_dim = key_dim * 2 + value_dim;
4358 let eps = self.cfg.rms_eps;
4359 let scale = 1.0 / (d_state as f32).sqrt();
4360 let b = ts.len();
4361 let c = Engine::gdn_chunk_size();
4362 let carried = caches.iter().any(|c| c.pos > 0);
4365 let use_vl = !carried
4366 && (2..=8).contains(&b)
4367 && Engine::gdn_chunked_enabled()
4368 && ts.iter().all(|&t| t >= 16)
4369 && e.gdn_mma_enabled(c)
4370 && std::env::var("MEMRA_GDN_VL").as_deref() != Ok("0");
4371 if !use_vl {
4372 return (0..b)
4373 .map(|s| {
4374 let (o, t) = (offs[s], ts[s]);
4375 self.linear_attn_prime_core_pad_view(
4376 e,
4377 la,
4378 &g4[0].slice(o * conv_dim..(o + t) * conv_dim),
4379 &g4[1].slice(o * value_dim..(o + t) * value_dim),
4380 &g4[2].slice(o * num_v..(o + t) * num_v),
4381 &g4[3].slice(o * num_v..(o + t) * num_v),
4382 t,
4383 caches[s],
4384 il,
4385 None,
4386 )
4387 })
4388 .collect();
4389 }
4390 struct SeqBufs {
4394 conv_out: CudaSlice<f32>,
4395 q_g: CudaSlice<f32>,
4396 k_g: CudaSlice<f32>,
4397 v_g: CudaSlice<f32>,
4398 q_l2: CudaSlice<f32>,
4399 k_l2: CudaSlice<f32>,
4400 beta: CudaSlice<f32>,
4401 g_log: CudaSlice<f32>,
4402 gn: CudaSlice<f32>,
4403 gn16: CudaSlice<u8>,
4404 }
4405 let f16o = Self::f16out_on(e, 16);
4406 let hk = Self::gdn_hk(e, 16, num_v, num_k); let mut sb = Vec::with_capacity(b);
4408 let mut pres = Vec::with_capacity(b);
4409 for &t in ts.iter().take(b) {
4410 sb.push(SeqBufs {
4411 conv_out: e.uninit(conv_dim * t)?,
4412 q_g: e.uninit(d_state * hk * t)?,
4413 k_g: e.uninit(d_state * hk * t)?,
4414 v_g: e.uninit(d_state * num_v * t)?,
4415 q_l2: e.uninit(d_state * hk * t)?,
4416 k_l2: e.uninit(d_state * hk * t)?,
4417 beta: e.uninit(t * num_v)?,
4418 g_log: e.uninit(t * num_v)?,
4419 gn: e.uninit(d_state * num_v * t)?,
4420 gn16: e.alloc_u8_uninit(d_state * num_v * t * 2)?,
4421 });
4422 pres.push(e.gdn_chunk_alloc(num_v, t, c, hk)?);
4423 }
4424 let prep_args: Vec<crate::GdnPrepVl> = (0..b)
4425 .map(|s| {
4426 let (o, t) = (offs[s], ts[s]);
4427 let rl = caches[s].recur[il].as_ref().unwrap();
4428 crate::GdnPrepVl {
4429 qkv: e.addr_f32v(&g4[0].slice(o * conv_dim..(o + t) * conv_dim)),
4430 conv_state: e.addr_f32(&rl.conv_state),
4431 conv_out: e.addr_f32(&sb[s].conv_out),
4432 q_g: e.addr_f32(&sb[s].q_g),
4433 k_g: e.addr_f32(&sb[s].k_g),
4434 v_g: e.addr_f32(&sb[s].v_g),
4435 q_l2: e.addr_f32(&sb[s].q_l2),
4436 k_l2: e.addr_f32(&sb[s].k_l2),
4437 beta_raw: e.addr_f32v(&g4[2].slice(o * num_v..(o + t) * num_v)),
4438 alpha: e.addr_f32v(&g4[3].slice(o * num_v..(o + t) * num_v)),
4439 beta: e.addr_f32(&sb[s].beta),
4440 g_log: e.addr_f32(&sb[s].g_log),
4441 o: e.addr_f32(&pres[s].o),
4442 z: e.addr_f32v(&g4[1].slice(o * value_dim..(o + t) * value_dim)),
4443 gn: e.addr_f32(&sb[s].gn),
4444 gn16: e.addr_u8(&sb[s].gn16),
4445 kb16: if Engine::l2_v2_on(d_state) {
4446 e.addr_u8(&pres[s].kb16)
4447 } else {
4448 0
4449 },
4450 qb16: if Engine::l2_v2_on(d_state) && e.gdn_wgmma_on(c) {
4451 e.addr_u8(&pres[s].qb16)
4452 } else {
4453 0
4454 },
4455 t: t as i32,
4456 pad: 0,
4457 }
4458 })
4459 .collect();
4460 let args: Vec<crate::GdnSeqVl> = (0..b)
4461 .map(|s| {
4462 let rl = caches[s].recur[il].as_ref().unwrap();
4463 crate::GdnSeqVl {
4464 kb16: e.addr_u8(&pres[s].kb16),
4465 gcum: e.addr_f32(&pres[s].gcum),
4466 beta: e.addr_f32(&sb[s].beta),
4467 u: e.addr_f32(&pres[s].u),
4468 wb16: e.addr_u8(&pres[s].wb16),
4469 y: e.addr_u8(&pres[s].y16),
4470 ssnap: e.addr_u8(&pres[s].ssnap16),
4471 state_in: e.addr_f32(&rl.ssm_state),
4472 state_out: e.addr_f32(&rl.ssm_state_alt),
4473 q: e.addr_f32(&sb[s].q_l2),
4474 p: e.addr_f32(&pres[s].p),
4475 o: e.addr_f32(&pres[s].o),
4476 k: e.addr_f32(&sb[s].k_l2),
4477 v: e.addr_f32(&sb[s].v_g),
4478 g: e.addr_f32(&sb[s].g_log),
4479 a: e.addr_f32(&pres[s].a),
4480 w: e.addr_f32(&pres[s].w),
4481 t: ts[s] as i32,
4482 nc: pres[s].nc as i32,
4483 }
4484 })
4485 .collect();
4486 e.gdn_prep_vl8(
4487 &prep_args,
4488 la.ssm_conv1d.float_data(),
4489 la.ssm_dt.float_data(),
4490 la.ssm_a.float_data(),
4491 conv_dim,
4492 d_conv,
4493 d_state,
4494 num_v,
4495 num_k,
4496 key_dim,
4497 hk,
4498 eps,
4499 )?;
4500 if !Engine::l2_v2_on(d_state) {
4503 e.gdn_mirror_vl8(&args, num_v, 0, hk)?;
4504 }
4505 let wq8: Option<crate::GdnWVl8> = if e.gdn_wgmma_on(c) {
4507 if !Engine::l2_v2_on(d_state) {
4509 for s in 0..b {
4510 e.f32_to_bf16_into(&sb[s].q_l2, &mut pres[s].qb16, d_state * hk * ts[s])?;
4511 }
4512 }
4513 let mut wa = [crate::GdnWVl::default(); 8];
4514 for s in 0..b {
4515 wa[s] = crate::GdnWVl {
4516 qb16: e.addr_u8(&pres[s].qb16),
4517 pb16: e.addr_u8(&pres[s].pb16),
4518 };
4519 }
4520 Some(crate::GdnWVl8(wa))
4521 } else {
4522 None
4523 };
4524 e.gdn_chunk_k123_vl8(&args, num_v, hk, wq8.as_ref())?;
4525 e.gdn_chunk_vl8(&args, num_v, scale, hk, wq8.as_ref())?;
4526 if f16o {
4527 e.gdn_tail_vl8(&prep_args, la.ssm_norm.float_data(), d_state, num_v, eps)?;
4528 }
4529 let mut out = Vec::with_capacity(b);
4531 for (s, bufs) in sb.into_iter().enumerate() {
4532 let rl = caches[s].recur[il].as_mut().unwrap();
4533 std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
4534 let (o, t) = (offs[s], ts[s]);
4535 let SeqBufs { mut gn, gn16, .. } = bufs;
4536 if f16o {
4537 out.push((gn, Some(gn16)));
4538 } else {
4539 let z_v = g4[1].slice(o * value_dim..(o + t) * value_dim);
4540 e.gated_rmsnorm_zv(
4541 &pres[s].o,
4542 la.ssm_norm.float_data(),
4543 &z_v,
4544 &mut gn,
4545 d_state,
4546 num_v * t,
4547 eps,
4548 )?;
4549 out.push((gn, None));
4550 }
4551 }
4552 Ok(out)
4553 }
4554
4555 #[allow(clippy::too_many_arguments)]
4559 fn linear_attn_prime_core_pad_view(
4560 &self,
4561 e: &Engine,
4562 la: &LinearAttnLayer,
4563 qkv_mixed: &cudarc::driver::CudaView<f32>,
4564 z: &cudarc::driver::CudaView<f32>,
4565 beta_raw: &cudarc::driver::CudaView<f32>,
4566 alpha: &cudarc::driver::CudaView<f32>,
4567 t: usize,
4568 cache: &mut Cache,
4569 il: usize,
4570 pad_len: Option<&CudaSlice<i32>>,
4571 ) -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
4572 let cfg = &self.cfg;
4573 let geometry = la.geometry;
4574 let d_state = geometry.key_head_dim as usize;
4575 let num_v = geometry.value_heads as usize;
4576 let eps = cfg.rms_eps;
4577 let scale = 1.0 / (d_state as f32).sqrt();
4578
4579 let prep =
4580 self.linear_attn_gdn_prep(e, la, qkv_mixed, beta_raw, alpha, t, cache, il, pad_len)?;
4581
4582 let mut o = e.uninit(d_state * num_v * t)?;
4588 let rl = cache.recur[il].as_mut().unwrap();
4589 {
4590 let crate::cache::RecurLayer {
4591 ssm_state,
4592 ssm_state_alt,
4593 ..
4594 } = rl;
4595 e.gdn_scan_prefill(
4596 &prep.q_l2,
4597 &prep.k_l2,
4598 &prep.v_g,
4599 &prep.g_log,
4600 &prep.beta,
4601 prep.kb16.as_ref(),
4602 prep.qb16.as_ref(),
4603 ssm_state,
4604 ssm_state_alt,
4605 &mut o,
4606 num_v,
4607 t,
4608 scale,
4609 prep.hk,
4610 )?;
4611 }
4612 std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
4613
4614 let mut gn = e.uninit(d_state * num_v * t)?;
4617 let gn16 = if Self::f16out_on(e, t) {
4618 let mut g16 = e.alloc_u8_uninit(d_state * num_v * t * 2)?;
4619 e.gated_rmsnorm_f16out_zv(
4620 &o,
4621 la.ssm_norm.float_data(),
4622 z,
4623 &mut gn,
4624 &mut g16,
4625 d_state,
4626 num_v * t,
4627 eps,
4628 )?;
4629 Some(g16)
4630 } else {
4631 e.gated_rmsnorm_zv(
4632 &o,
4633 la.ssm_norm.float_data(),
4634 z,
4635 &mut gn,
4636 d_state,
4637 num_v * t,
4638 eps,
4639 )?;
4640 None
4641 };
4642 Ok((gn, gn16))
4643 }
4644
4645 #[allow(clippy::too_many_arguments)]
4647 fn linear_attn_prime_core_pad(
4648 &self,
4649 e: &Engine,
4650 la: &LinearAttnLayer,
4651 g4: Vec<CudaSlice<f32>>,
4652 t: usize,
4653 cache: &mut Cache,
4654 il: usize,
4655 pad_len: Option<&CudaSlice<i32>>,
4656 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4657 let (gn, gn16) = self.linear_attn_prime_core_pad_inner(e, la, g4, t, cache, il, pad_len)?;
4658 if let Some(xh) = &gn16 {
4659 if let Some(y) = e.try_f16_gemm_pre(&la.ssm_out, xh, t)? {
4660 return Ok(y);
4661 }
4662 }
4663 Ok(e.matmul(&la.ssm_out, &gn, t)?)
4664 }
4665
4666 pub fn full_attn(
4671 &self,
4672 e: &Engine,
4673 fa: &FullAttnLayer,
4674 h: &CudaSlice<f32>,
4675 pos_d: &CudaSlice<i32>,
4676 t: usize,
4677 il: usize,
4678 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4679 if self.uses_sliding_gated_moe_program() {
4680 return self.step35_attn(e, fa, h, pos_d, t, il);
4681 }
4682 let cfg = &self.cfg;
4683 let _n_embd = cfg.n_embd as usize;
4684 let geometry = cfg.full_attention_geometry_at(il as u32);
4685 let n_head = geometry.n_head as usize;
4686 let n_head_kv = geometry.n_head_kv as usize;
4687 let head_dim = geometry.head_dim_k as usize;
4688 let eps = cfg.rms_eps;
4689 let scale = geometry.attention_scale();
4690
4691 let gated = geometry.attention_gate == memra_gguf::config::AttentionGateKind::FusedQ;
4694 let mut g3 = e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?;
4696 let v = g3.pop().unwrap();
4697 let mut k = g3.pop().unwrap();
4698 let qf = g3.pop().unwrap();
4699 let (mut q, gate) = if gated {
4700 let mut q = e.uninit(t * n_head * head_dim)?;
4701 let mut gate = e.uninit(t * n_head * head_dim)?;
4702 e.q_gate_split(&qf, &mut q, &mut gate, head_dim, n_head, t)?;
4703 (q, Some(gate))
4704 } else {
4705 (qf, None)
4706 };
4707
4708 let mut qn = e.uninit(t * n_head * head_dim)?;
4710 e.rms_norm(
4711 &q,
4712 fa.q_norm.float_data(),
4713 &mut qn,
4714 head_dim,
4715 n_head * t,
4716 eps,
4717 )?;
4718 q = qn;
4719 let mut kn = e.uninit(t * n_head_kv * head_dim)?;
4720 e.rms_norm(
4721 &k,
4722 fa.k_norm.float_data(),
4723 &mut kn,
4724 head_dim,
4725 n_head_kv * t,
4726 eps,
4727 )?;
4728 k = kn;
4729 let rope_dims = geometry.n_rot as usize;
4730 e.rope_neox(
4731 &mut q,
4732 pos_d,
4733 head_dim,
4734 rope_dims,
4735 n_head,
4736 t,
4737 geometry.rope_base,
4738 1.0,
4739 )?;
4740 e.rope_neox(
4741 &mut k,
4742 pos_d,
4743 head_dim,
4744 rope_dims,
4745 n_head_kv,
4746 t,
4747 geometry.rope_base,
4748 1.0,
4749 )?;
4750
4751 let mut attn = e.uninit(t * n_head * head_dim)?;
4753 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
4756 e.sdpa_naive(
4758 &q, &k, &v, &mut attn, head_dim, n_head, n_head_kv, t, t, scale, true,
4759 )?;
4760 } else {
4761 e.fa_prefill(
4762 &q, &k, &v, &mut attn, head_dim, n_head, n_head_kv, t, t, scale, true,
4763 )?;
4764 }
4765
4766 let attn_g = match &gate {
4768 Some(gate) => {
4769 let mut gsig = e.uninit(t * n_head * head_dim)?;
4770 e.sigmoid(gate, &mut gsig, t * n_head * head_dim)?;
4771 let mut ag = e.uninit(t * n_head * head_dim)?;
4772 e.mul(&attn, &gsig, &mut ag, t * n_head * head_dim)?;
4773 ag
4774 }
4775 None => attn,
4776 };
4777
4778 let o = e.matmul(&fa.wo, &attn_g, t)?;
4780 Ok(o)
4781 }
4782
4783 pub fn linear_attn(
4785 &self,
4786 e: &Engine,
4787 la: &LinearAttnLayer,
4788 h: &CudaSlice<f32>,
4789 t: usize,
4790 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4791 let cfg = &self.cfg;
4792 let _n_embd = cfg.n_embd as usize;
4793 let geometry = la.geometry;
4794 let d_state = geometry.key_head_dim as usize;
4795 let num_k = geometry.key_heads as usize;
4796 let num_v = geometry.value_heads as usize;
4797 let d_conv = geometry.conv_kernel as usize;
4798 let head_k = d_state;
4799 let head_v = geometry.value_head_dim as usize;
4800 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;
4804 let scale = 1.0 / (d_state as f32).sqrt();
4805
4806 let mut g4 = e.matmul_group(
4809 &[&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha],
4810 h,
4811 t,
4812 )?;
4813 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);
4825 let mut q_g = e.uninit(d_state * num_v * t)?;
4826 let mut k_g = e.uninit(d_state * num_v * t)?;
4827 let mut v_g = e.uninit(d_state * num_v * t)?;
4828 e.ssm_conv1d_gdn(
4829 &qkv_mixed,
4830 la.ssm_conv1d.float_data(),
4831 &mut q_g,
4832 &mut k_g,
4833 &mut v_g,
4834 conv_dim,
4835 t,
4836 d_conv,
4837 d_state,
4838 num_v,
4839 num_k,
4840 key_dim,
4841 )?;
4842 let mut q_l2 = e.uninit(d_state * num_v * t)?;
4844 e.l2_norm(&q_g, &mut q_l2, d_state, num_v * t, eps)?;
4845 let mut k_l2 = e.uninit(d_state * num_v * t)?;
4846 e.l2_norm(&k_g, &mut k_l2, d_state, num_v * t, eps)?;
4847 let v_gd = v_g;
4848
4849 let mut beta = e.uninit(t * num_v)?;
4852 e.sigmoid(&beta_raw, &mut beta, t * num_v)?;
4853 let mut g_log = e.uninit(t * num_v)?;
4855 e.gdn_glog(
4856 &alpha,
4857 la.ssm_dt.float_data(),
4858 la.ssm_a.float_data(),
4859 &mut g_log,
4860 num_v,
4861 t,
4862 )?;
4863
4864 let state_in = e.zeros(d_state * d_state * num_v)?; let mut state_out = e.zeros(d_state * d_state * num_v)?;
4867 let mut o = e.uninit(d_state * num_v * t)?;
4868 e.gdn_scan_prefill(
4869 &q_l2,
4870 &k_l2,
4871 &v_gd,
4872 &g_log,
4873 &beta,
4874 None,
4875 None,
4876 &state_in,
4877 &mut state_out,
4878 &mut o,
4879 num_v,
4880 t,
4881 scale,
4882 num_v,
4883 )?;
4884
4885 let mut gn = e.uninit(d_state * num_v * t)?;
4890 e.gated_rmsnorm(
4891 &o,
4892 la.ssm_norm.float_data(),
4893 &z,
4894 &mut gn,
4895 d_state,
4896 num_v * t,
4897 eps,
4898 )?;
4899
4900 let out = e.matmul(&la.ssm_out, &gn, t)?;
4904 Ok(out)
4905 }
4906}
4907
4908impl HybridModel {
4909 pub fn moe_ffn_il(
4920 &self,
4921 e: &Engine,
4922 m: &MoeWeights,
4923 z: &CudaSlice<f32>,
4924 t: usize,
4925 il: u16,
4926 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4927 Self::moe_ffn_inner(
4928 e,
4929 m,
4930 z,
4931 None,
4932 t,
4933 &self.cfg,
4934 il,
4935 self.max_moe_block(),
4936 false,
4937 None,
4938 self.uses_sliding_gated_moe_program(),
4939 )
4940 }
4941
4942 pub fn moe_ffn_il_prefill(
4945 &self,
4946 e: &Engine,
4947 m: &MoeWeights,
4948 z: &CudaSlice<f32>,
4949 t: usize,
4950 il: u16,
4951 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4952 Self::moe_ffn_inner(
4953 e,
4954 m,
4955 z,
4956 None,
4957 t,
4958 &self.cfg,
4959 il,
4960 self.max_moe_block(),
4961 true,
4962 Some(&self.step_grouped_prefill),
4963 self.uses_sliding_gated_moe_program(),
4964 )
4965 }
4966
4967 pub fn moe_ffn_il_zq8(
4971 &self,
4972 e: &Engine,
4973 m: &MoeWeights,
4974 z: &CudaSlice<f32>,
4975 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
4976 t: usize,
4977 il: u16,
4978 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4979 Self::moe_ffn_inner(
4980 e,
4981 m,
4982 z,
4983 zq8,
4984 t,
4985 &self.cfg,
4986 il,
4987 self.max_moe_block(),
4988 false,
4989 None,
4990 self.uses_sliding_gated_moe_program(),
4991 )
4992 }
4993
4994 pub(crate) fn moe_ffn(
5002 e: &Engine,
5003 m: &MoeWeights,
5004 z: &CudaSlice<f32>,
5005 t: usize,
5006 cfg: &ModelConfig,
5007 il: u16,
5008 max_block: usize,
5009 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5010 Self::moe_ffn_inner(e, m, z, None, t, cfg, il, max_block, false, None, false)
5011 }
5012
5013 #[allow(clippy::too_many_arguments)]
5014 pub(crate) fn moe_ffn_inner(
5015 e: &Engine,
5016 m: &MoeWeights,
5017 z: &CudaSlice<f32>,
5018 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
5019 t: usize,
5020 cfg: &ModelConfig,
5021 il: u16,
5022 max_block: usize,
5023 prefill: bool,
5024 grouped_prefill: Option<&std::sync::Mutex<crate::hybrid::StepEpGroupedPrefill>>,
5025 sliding_gated_moe: bool,
5026 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5027 let worker_io = crate::spill_pread::worker_enabled();
5028 let epoch_lfu = std::env::var_os("MEMRA_MOE_LFU_DECAY").is_some();
5029 if Engine::moe_cache_enabled() && (worker_io || epoch_lfu) {
5030 e.with_moe_cache(max_block, |cache, _| {
5031 cache.begin_forward_epoch(il, t);
5032 if worker_io {
5033 cache.begin_worker_scope();
5034 }
5035 Ok(())
5036 })?;
5037 }
5038 if m.step_ep.is_some() || m.step_tp.is_some() {
5039 let moe = cfg
5040 .moe
5041 .as_ref()
5042 .ok_or("Step distributed execution requires MoE model metadata")?;
5043 let n_embd = cfg.n_embd as usize;
5044 let n_expert = moe.expert_count as usize;
5045 let n_used = moe.expert_used_count as usize;
5046 let sigmoid = cfg
5047 .sigmoid_router()
5048 .ok_or("Step distributed execution requires the Step sigmoid router")?;
5049 let logits = Self::moe_router_logits(e, m, z, t, cfg)?;
5050 Self::trace_sigmoid_router_logits(e, il, t, n_expert, n_used, &logits, m, sigmoid)?;
5051 let grouped_prefill_requested = prefill && step_ep_grouped_prefill_enabled()?;
5052 if grouped_prefill_requested && !step_tp_prefill_enabled()? {
5053 return Err(
5054 "MEMRA_STEP_EP_GROUPED_PREFILL=1 requires MEMRA_STEP_TP_PREFILL=1".into(),
5055 );
5056 }
5057 if grouped_prefill_requested && !step_grouped_prefill_shape(true, prefill, t) {
5058 return Err(format!(
5059 "Step grouped prefill tokens {t} are outside the qualified {}..={} range",
5060 PRIME_MIN_T,
5061 crate::cache::PRIME_CHUNK_MAX_TOKENS,
5062 )
5063 .into());
5064 }
5065 let grouped_decode_shape = step_grouped_decode_shape(prefill, t);
5066 let grouped_prefill_shape =
5067 step_grouped_prefill_shape(grouped_prefill_requested, prefill, t);
5068 if let Some(ep) = m.step_ep.as_ref().filter(|ep| {
5069 ep.grouped_decode.is_some() && (grouped_decode_shape || grouped_prefill_shape)
5070 }) {
5071 let (selected, route_weights) =
5072 Self::moe_route_sigmoid_cfg(e, &logits, t, n_expert, n_used, m, sigmoid)?;
5073 crate::moesd::record_host_routes(il, n_expert, n_used, &selected)?;
5074 Self::trace_moe_routes(il, t, &selected, &route_weights)?;
5075 Self::trace_moe_input(e, il, t, n_embd, z)?;
5076 let selected = selected
5077 .iter()
5078 .map(|&expert| expert as usize)
5079 .collect::<Vec<_>>();
5080
5081 e.stream().synchronize()?;
5084 let execute = |state: &mut crate::hybrid::StepEpGroupedDecode| {
5085 state.projection.set_activation_limit(ep.activation_limit)?;
5086 ep.runtime
5087 .refresh_step_grouped_expert_parallel_gate_from_root_device(
5088 ep.experts.e4m3()?,
5089 &mut state.projection,
5090 z,
5091 t,
5092 &selected,
5093 )?;
5094 ep.runtime.refresh_step_grouped_expert_parallel_combine(
5095 &state.projection,
5096 &mut state.combine,
5097 &route_weights,
5098 )?;
5099 ep.runtime.execute_step_grouped_expert_parallel_gate(
5100 ep.experts.e4m3()?,
5101 &mut state.projection,
5102 )?;
5103 ep.runtime.execute_step_grouped_expert_parallel_combine(
5104 &state.projection,
5105 &mut state.combine,
5106 )?;
5107 let mut output = ep.runtime.copy_step_grouped_expert_parallel_combine_root(
5108 &state.projection,
5109 &state.combine,
5110 e,
5111 )?;
5112 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut output)?;
5113 if prefill {
5114 e.stream().synchronize()?;
5117 }
5118 eprintln!(
5119 "[step-tp-ep-grouped] execute layer={il} tokens={t} devices={:?} \
5120 attention_layout=tensor-parallel expert_layout=expert-parallel \
5121 expert_transport={} native_p2p=true route_control=host-narrow \
5122 input=root-device projection_workspaces=persistent \
5123 combine=root-device output=owning-stage-device \
5124 prefill={prefill} batched_decode=false capacity={} \
5125 performance_claim=false",
5126 ep.devices,
5127 ep.runtime.transport_label(),
5128 state.projection.max_tokens(),
5129 );
5130 Ok::<_, Box<dyn std::error::Error>>(output)
5131 };
5132
5133 if grouped_prefill_shape {
5134 let grouped_prefill = grouped_prefill
5135 .ok_or("Step grouped prefill has no model-scoped executor")?;
5136 let mut shared = grouped_prefill
5137 .lock()
5138 .map_err(|_| "Step grouped prefill state lock is poisoned")?;
5139 let needs_prepare = shared.state.as_ref().is_none_or(|state| {
5140 state.devices != ep.devices
5141 || state.grouped.projection.max_tokens() < t
5142 || state.grouped.projection.input_width() != n_embd
5143 || state.grouped.projection.expert_width()
5144 != moe.expert_ff_length as usize
5145 });
5146 if needs_prepare {
5147 let seed_input = vec![0.0f32; n_embd];
5148 let seed_selected = &selected[..n_used];
5149 let seed_weights = &route_weights[..n_used];
5150 let projection = ep
5151 .runtime
5152 .prepare_step_grouped_expert_parallel_gate_with_capacity(
5153 ep.experts.e4m3()?,
5154 &seed_input,
5155 1,
5156 seed_selected,
5157 ep.activation_limit,
5158 t,
5159 )?;
5160 let combine = ep.runtime.prepare_step_grouped_expert_parallel_combine(
5161 &projection,
5162 seed_weights,
5163 )?;
5164 shared.state = Some(crate::hybrid::StepEpGroupedPrefillState {
5165 devices: ep.devices.clone(),
5166 grouped: crate::hybrid::StepEpGroupedDecode {
5167 projection,
5168 combine,
5169 },
5170 });
5171 eprintln!(
5172 "[step-tp-ep-grouped-prefill] prepare capacity={t} devices={:?} \
5173 shared_across_layers=true performance_claim=false",
5174 ep.devices,
5175 );
5176 }
5177 return execute(
5178 &mut shared
5179 .state
5180 .as_mut()
5181 .expect("Step grouped prefill state prepared above")
5182 .grouped,
5183 );
5184 }
5185
5186 let mut grouped = ep
5187 .grouped_decode
5188 .as_ref()
5189 .expect("grouped decode presence checked above")
5190 .lock()
5191 .map_err(|_| "Step grouped decode state lock is poisoned")?;
5192 return execute(&mut grouped);
5193 }
5194 if grouped_prefill_shape {
5195 return Err(
5196 "Step grouped prefill requires native-P2P expert-owner device arithmetic"
5197 .into(),
5198 );
5199 }
5200 if t >= 16 && crate::step_gemm_prime_on() {
5215 if let Some(tp) = &m.step_tp {
5216 if let crate::hybrid::StepTpExpertBank::Nvfp4(bank) = &tp.experts {
5217 let mprof =
5225 std::env::var("MEMRA_PRIME_PROF").as_deref() == Ok("1") && t >= 16;
5226 let mut mt = std::time::Instant::now();
5227 let (selected, route_weights) = Self::moe_route_sigmoid_cfg(
5228 e, &logits, t, n_expert, n_used, m, sigmoid,
5229 )?;
5230 let sel_i32: Vec<i32> = selected.iter().map(|&x| x as i32).collect();
5231 let d_router = if mprof {
5232 let _ = e.stream().synchronize();
5233 let v = mt.elapsed().as_secs_f64() * 1e3;
5234 mt = std::time::Instant::now();
5235 v
5236 } else {
5237 0.0
5238 };
5239 let mdet = std::env::var("MEMRA_MOE_DETERM").as_deref() == Ok("1")
5253 && t >= 16
5254 && il < 4;
5255 if mdet {
5256 let a = tp.runtime.run_tensor_parallel_routes_nvfp4_prime_grouped(
5257 bank,
5258 e,
5259 z,
5260 t,
5261 &sel_i32,
5262 &route_weights,
5263 n_used,
5264 tp.activation_limit,
5265 )?;
5266 let b = tp.runtime.run_tensor_parallel_routes_nvfp4_prime_grouped(
5267 bank,
5268 e,
5269 z,
5270 t,
5271 &sel_i32,
5272 &route_weights,
5273 n_used,
5274 tp.activation_limit,
5275 )?;
5276 let (ha, hb) = (e.dtoh(&a)?, e.dtoh(&b)?);
5277 let mut md = 0.0f32;
5278 let mut ndiff = 0usize;
5279 for (x, y) in ha.iter().zip(hb.iter()) {
5280 let d = (x - y).abs();
5281 if d > 0.0 {
5282 ndiff += 1;
5283 }
5284 if d > md {
5285 md = d;
5286 }
5287 }
5288 eprintln!(
5289 "[moe-determ] il={il} t={t} maxdiff={md:.3e} \
5290 differing={ndiff}/{} -> {}",
5291 ha.len(),
5292 if ndiff == 0 {
5293 "IDENTICAL"
5294 } else {
5295 "NONDETERMINISTIC"
5296 }
5297 );
5298 }
5299 let mut output =
5300 tp.runtime.run_tensor_parallel_routes_nvfp4_prime_grouped(
5301 bank,
5302 e,
5303 z,
5304 t,
5305 &sel_i32,
5306 &route_weights,
5307 n_used,
5308 tp.activation_limit,
5309 )?;
5310 let d_gemm = if mprof {
5311 let _ = e.stream().synchronize();
5312 let v = mt.elapsed().as_secs_f64() * 1e3;
5313 mt = std::time::Instant::now();
5314 v
5315 } else {
5316 0.0
5317 };
5318 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut output)?;
5319 if mprof {
5320 let _ = e.stream().synchronize();
5321 let d_shared = mt.elapsed().as_secs_f64() * 1e3;
5322 eprintln!(
5326 "[moe-prof] il={il} t={t} router={d_router:.1}ms \
5327 gemm={d_gemm:.1}ms shared={d_shared:.1}ms"
5328 );
5329 }
5330 return Ok(output);
5331 }
5332 }
5333 }
5334 if t == 1
5335 && crate::tp::step_nvfp4_dev_routes_enabled()?
5336 && crate::tp::step_tp_dev_router_enabled()?
5337 {
5338 if let Some(tp) = &m.step_tp {
5339 if let crate::hybrid::StepTpExpertBank::Nvfp4(bank) = &tp.experts {
5340 let (sf, route_norm) = sigmoid;
5341 static D1_ROUTER: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
5348 let d1_router = *D1_ROUTER.get_or_init(|| {
5349 std::env::var("MEMRA_DEV1_ROUTER").as_deref() == Ok("1")
5350 });
5351 if d1_router {
5352 let (sf_h, rn_h) = sigmoid;
5353 let n_ex = m.gate_exps.n_expert;
5354 let act_ct = m.active_count();
5355 let _ = tp.runtime.nvfp4_routes_prestage_with(
5356 bank,
5357 e,
5358 z,
5359 |rank1, in1, sel1, w1| {
5360 let mut guard = DEV1_ROUTER_REPS
5361 .lock()
5362 .map_err(|_| "dev1 router replica lock")?;
5363 let (reps, scratch) =
5364 guard.get_or_insert_with(|| (Default::default(), None));
5365 if !reps.contains_key(&il) {
5366 use cudarc::driver::DevicePtr;
5367 let (g1, p1, a1) = (
5368 rank1.htod(&vec![0.0f32; n_ex * n_embd])?,
5369 rank1.htod(&vec![0.0f32; n_ex])?,
5370 rank1.alloc_u8_uninit(n_ex)?,
5371 );
5372 for (src, dst_len, dst) in [
5373 (
5374 {
5375 let s = e.stream();
5376 let (p, _g) =
5377 m.gate_inp.float_data().device_ptr(&s);
5378 p as u64
5379 },
5380 n_ex * n_embd * 4,
5381 {
5382 let s = rank1.stream();
5383 let (p, _g) = g1.device_ptr(&s);
5384 p as u64
5385 },
5386 ),
5387 (
5388 {
5389 let s = e.stream();
5390 let (p, _g) = m.exp_probs_b_dev.device_ptr(&s);
5391 p as u64
5392 },
5393 n_ex * 4,
5394 {
5395 let s = rank1.stream();
5396 let (p, _g) = p1.device_ptr(&s);
5397 p as u64
5398 },
5399 ),
5400 (
5401 {
5402 let s = e.stream();
5403 let (p, _g) =
5404 m.active_experts_dev.device_ptr(&s);
5405 p as u64
5406 },
5407 n_ex,
5408 {
5409 let s = rank1.stream();
5410 let (p, _g) = a1.device_ptr(&s);
5411 p as u64
5412 },
5413 ),
5414 ] {
5415 crate::tp::raw_copy_bytes(dst, src, dst_len, rank1)?;
5416 }
5417 rank1.stream().synchronize()?;
5418 reps.insert(il, (g1, p1, a1));
5419 }
5420 if scratch.is_none() {
5421 *scratch = Some(rank1.htod(&vec![0.0f32; n_ex])?);
5422 }
5423 let (g1, p1, a1) = reps.get(&il).expect("armed above");
5424 let logits1 = scratch.as_mut().expect("armed above");
5425 rank1.router_gemv_into(g1, in1, logits1, n_embd, n_ex, 1)?;
5426 rank1.moe_router_sigmoid_topk_into(
5427 logits1, 1, n_ex, n_used, act_ct, p1, a1, sf_h, rn_h, sel1,
5428 w1,
5429 )?;
5430 Ok(true)
5431 },
5432 )?;
5433 } else {
5434 let _ = tp.runtime.nvfp4_routes_prestage(bank, e, z)?;
5435 }
5436 static SELW: std::sync::Mutex<
5440 Option<(usize, CudaSlice<i32>, CudaSlice<f32>)>,
5441 > = std::sync::Mutex::new(None);
5442 let mut selw = SELW.lock().map_err(|_| "selw lock poisoned")?;
5443 if selw.as_ref().is_none_or(|(d, ..)| *d != e.ctx().ordinal()) {
5444 *selw = Some((
5445 e.ctx().ordinal(),
5446 e.htod_i32(&vec![0i32; n_used])?,
5447 e.htod(&vec![0.0f32; n_used])?,
5448 ));
5449 }
5450 let (_, sel_d, w_d) = selw.as_mut().expect("armed above");
5451 e.moe_router_sigmoid_topk_into(
5452 &logits,
5453 t,
5454 n_expert,
5455 n_used,
5456 m.active_count(),
5457 &m.exp_probs_b_dev,
5458 &m.active_experts_dev,
5459 sf,
5460 route_norm,
5461 sel_d,
5462 w_d,
5463 )?;
5464 crate::moesd::record_device_routes(e, il, n_expert, n_used, &sel_d)?;
5465 static SHEXP_OV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
5469 let shexp_ov = *SHEXP_OV.get_or_init(|| {
5470 std::env::var("MEMRA_SHEXP_OVERLAP").as_deref() == Ok("1")
5471 });
5472 static SHEXP_D1: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
5476 let shexp_d1 = *SHEXP_D1.get_or_init(|| {
5477 std::env::var("MEMRA_SHEXP_DEV1").as_deref() == Ok("1")
5478 }) && tp.runtime.rank_engine(1).is_some();
5479 static TAIL3: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
5483 let tail3 = *TAIL3
5484 .get_or_init(|| std::env::var("MEMRA_TAIL_ADD3").as_deref() != Ok("0"));
5485 let mut ov_issued = false;
5486 let mut d1_issued = false;
5487 let mut tail_folded = false;
5488 let mut output = if shexp_d1 {
5489 let rank1 = tp.runtime.rank_engine(1).expect("checked above");
5490 tp.runtime
5491 .run_tensor_parallel_routes_nvfp4_device_routed_prejoin(
5492 bank,
5493 e,
5494 z,
5495 &sel_d,
5496 &w_d,
5497 n_used,
5498 tp.activation_limit,
5499 || {
5500 d1_issued = Self::shexp_dev1_issue(
5501 e, rank1, m, z, cfg, il, n_embd,
5502 )?;
5503 Ok(())
5504 },
5505 )?
5506 } else if shexp_ov {
5507 let post_add = if tail3 {
5512 Self::shexp_overlap_tail_ptrs(e, m, cfg, n_embd)?
5513 } else {
5514 None
5515 };
5516 let used_post = post_add.is_some();
5517 let out = tp
5518 .runtime
5519 .run_tensor_parallel_routes_nvfp4_device_routed_prejoin_add3(
5520 bank,
5521 e,
5522 z,
5523 &sel_d,
5524 &w_d,
5525 n_used,
5526 tp.activation_limit,
5527 || {
5528 ov_issued =
5529 Self::shexp_overlap_issue(e, m, z, cfg, il, n_embd)?;
5530 Ok(())
5531 },
5532 post_add,
5533 )?;
5534 if used_post && ov_issued {
5539 tail_folded = true; }
5541 out
5542 } else {
5543 tp.runtime.run_tensor_parallel_routes_nvfp4_device_routed(
5544 bank,
5545 e,
5546 z,
5547 &sel_d,
5548 &w_d,
5549 n_used,
5550 tp.activation_limit,
5551 )?
5552 };
5553 if output.len() != t * n_embd {
5554 return Err(format!(
5555 "Step tp routed output has {} values, expected {t}x{n_embd}",
5556 output.len()
5557 )
5558 .into());
5559 }
5560 if tail_folded {
5561 } else if d1_issued {
5563 Self::shexp_dev1_apply(e, &mut output, n_embd)?;
5564 } else if ov_issued {
5565 Self::shexp_overlap_apply(e, &mut output, n_embd)?;
5566 } else {
5567 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut output)?;
5568 }
5569 static DR_LOGGED: std::sync::atomic::AtomicU64 =
5570 std::sync::atomic::AtomicU64::new(0);
5571 let layer_bit = 1u64 << (il as u64 % 64);
5572 if DR_LOGGED.fetch_or(layer_bit, std::sync::atomic::Ordering::Relaxed)
5573 & layer_bit
5574 == 0
5575 {
5576 eprintln!(
5577 "[step-tp] execute layer={il} tokens={t} devices={:?} \
5578 expert_transport={} native_p2p={} router=device \
5579 activation=host-canonical accumulation=host-canonical \
5580 output=e-device io=device performance_claim=false \
5581 (logged once per layer)",
5582 tp.devices,
5583 tp.runtime.transport_label(),
5584 tp.runtime.native_p2p(),
5585 );
5586 }
5587 return Ok(output);
5588 }
5589 }
5590 }
5591 static ROUTE_NS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
5595 static ROUTE_CALLS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
5596 let route_timing = std::env::var("MEMRA_STEP_TP_TIMING").as_deref() == Ok("1");
5597 let route_started = route_timing.then(std::time::Instant::now);
5598 let (selected, route_weights, input) = Self::moe_route_sigmoid_with_input(
5599 e,
5600 &logits,
5601 z,
5602 t,
5603 n_embd,
5604 n_expert,
5605 n_used,
5606 m.exp_probs_b.as_deref(),
5607 sigmoid,
5608 m.active_experts.as_deref(),
5609 )?;
5610 if let Some(started) = route_started {
5611 use std::sync::atomic::Ordering;
5612 let ns = ROUTE_NS.fetch_add(started.elapsed().as_nanos() as u64, Ordering::Relaxed)
5613 + started.elapsed().as_nanos() as u64;
5614 let calls = ROUTE_CALLS.fetch_add(1, Ordering::Relaxed) + 1;
5615 if calls % 430 == 0 {
5616 eprintln!(
5617 "[moe-route-timing] calls={calls} total_ms={:.1} avg_us={:.1}",
5618 ns as f64 / 1.0e6,
5619 ns as f64 / calls as f64 / 1.0e3,
5620 );
5621 }
5622 }
5623 crate::moesd::record_host_routes(il, n_expert, n_used, &selected)?;
5624 Self::trace_moe_routes(il, t, &selected, &route_weights)?;
5625 Self::trace_moe_input(e, il, t, n_embd, z)?;
5626 let selected = selected
5627 .iter()
5628 .map(|&expert| expert as usize)
5629 .collect::<Vec<_>>();
5630 if t == 1 && crate::tp::step_nvfp4_dev_routes_enabled()? {
5635 if let Some(tp) = &m.step_tp {
5636 if let crate::hybrid::StepTpExpertBank::Nvfp4(bank) = &tp.experts {
5637 let mut output = tp.runtime.run_tensor_parallel_routes_nvfp4_device_io(
5638 bank,
5639 e,
5640 z,
5641 &selected,
5642 &route_weights,
5643 n_used,
5644 tp.activation_limit,
5645 )?;
5646 if output.len() != t * n_embd {
5647 return Err(format!(
5648 "Step tp routed output has {} values, expected {t}x{n_embd}",
5649 output.len()
5650 )
5651 .into());
5652 }
5653 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut output)?;
5654 static IO_LOGGED: std::sync::atomic::AtomicU64 =
5655 std::sync::atomic::AtomicU64::new(0);
5656 let layer_bit = 1u64 << (il as u64 % 64);
5657 if IO_LOGGED.fetch_or(layer_bit, std::sync::atomic::Ordering::Relaxed)
5658 & layer_bit
5659 == 0
5660 {
5661 eprintln!(
5662 "[step-tp] execute layer={il} tokens={t} devices={:?} \
5663 expert_transport={} native_p2p={} activation=host-canonical \
5664 accumulation=host-canonical output=e-device io=device \
5665 performance_claim=false (logged once per layer)",
5666 tp.devices,
5667 tp.runtime.transport_label(),
5668 tp.runtime.native_p2p(),
5669 );
5670 }
5671 return Ok(output);
5672 }
5673 }
5674 }
5675 let (routed, mode, devices, transport, native_p2p) = if let Some(tp) = &m.step_tp {
5676 (
5677 match &tp.experts {
5678 crate::hybrid::StepTpExpertBank::E4m3(bank) => {
5679 tp.runtime.run_tensor_parallel_routes(
5680 bank,
5681 &input,
5682 t,
5683 &selected,
5684 &route_weights,
5685 n_used,
5686 )?
5687 }
5688 crate::hybrid::StepTpExpertBank::Nvfp4(bank) => {
5689 if t == 1 && crate::tp::step_nvfp4_dev_routes_enabled()? {
5690 tp.runtime.run_tensor_parallel_routes_nvfp4_device(
5691 bank,
5692 &input,
5693 &selected,
5694 &route_weights,
5695 n_used,
5696 tp.activation_limit,
5697 )?
5698 } else {
5699 tp.runtime.run_tensor_parallel_routes_nvfp4(
5700 bank,
5701 &input,
5702 t,
5703 &selected,
5704 &route_weights,
5705 n_used,
5706 tp.activation_limit,
5707 )?
5708 }
5709 }
5710 },
5711 "tp",
5712 &tp.devices,
5713 tp.runtime.transport_label(),
5714 tp.runtime.native_p2p(),
5715 )
5716 } else {
5717 let ep = m
5718 .step_ep
5719 .as_ref()
5720 .ok_or("Step distributed runtime has no EP or TP state")?;
5721 (
5722 match &ep.experts {
5723 crate::hybrid::StepEpExpertBank::E4m3(bank) => {
5724 ep.runtime.run_routed_experts(
5725 bank,
5726 &input,
5727 t,
5728 &selected,
5729 &route_weights,
5730 n_used,
5731 ep.activation_limit,
5732 )?
5733 }
5734 crate::hybrid::StepEpExpertBank::Nvfp4(bank) => {
5735 ep.runtime.run_routed_experts_nvfp4(
5736 bank,
5737 &input,
5738 t,
5739 &selected,
5740 &route_weights,
5741 n_used,
5742 ep.activation_limit,
5743 )?
5744 }
5745 },
5746 if ep.configured_by_tp { "tp-ep" } else { "ep" },
5747 &ep.devices,
5748 ep.runtime.transport_label(),
5749 ep.runtime.native_p2p(),
5750 )
5751 };
5752 if routed.len() != t * n_embd {
5753 return Err(format!(
5754 "Step {mode} routed output has {} values, expected {t}x{n_embd}",
5755 routed.len()
5756 )
5757 .into());
5758 }
5759 let mut output = e.htod(&routed)?;
5760 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut output)?;
5761 static STEP_LOGGED: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
5764 let layer_bit = 1u64 << (il as u64 % 64);
5765 if STEP_LOGGED.fetch_or(layer_bit, std::sync::atomic::Ordering::Relaxed) & layer_bit
5766 == 0
5767 {
5768 eprintln!(
5769 "[step-{mode}] execute layer={il} tokens={t} devices={devices:?} \
5770 expert_transport={transport} native_p2p={native_p2p} \
5771 activation={} accumulation={} output={} \
5772 performance_claim=false (logged once per layer)",
5773 if let Some(ep) = &m.step_ep {
5774 ep.runtime.expert_activation_label()
5775 } else {
5776 "host-canonical"
5777 },
5778 if let Some(ep) = &m.step_ep {
5779 ep.runtime.expert_accumulation_label()
5780 } else {
5781 "host-canonical"
5782 },
5783 if let Some(ep) = &m.step_ep {
5784 ep.runtime.expert_output_label()
5785 } else {
5786 "host-accumulated"
5787 },
5788 );
5789 if let Some(ep) = &m.step_ep {
5790 if let Some(limit) = ep.activation_limit {
5791 eprintln!(
5792 "[step-ep-clamp] execute layer={il} tokens={t} routed_clamp={limit} \
5793 formula=min-silu-times-clamped-up performance_claim=false"
5794 );
5795 }
5796 }
5797 }
5798 return Ok(output);
5799 }
5800 if Self::sigmoid_resident_dev_eligible(e, m, cfg, sliding_gated_moe) {
5801 let moe = cfg.moe.as_ref().unwrap();
5802 let n_expert = moe.expert_count as usize;
5803 let n_used = moe.expert_used_count as usize;
5804 let sigmoid = cfg.sigmoid_router().unwrap();
5805 let logits = Self::moe_router_logits(e, m, z, t, cfg)?;
5806 Self::trace_sigmoid_router_logits(e, il, t, n_expert, n_used, &logits, m, sigmoid)?;
5807 return Self::moe_ffn_sigmoid_dev(e, m, z, zq8, &logits, t, cfg, il, sigmoid);
5808 }
5809 if t > 1 && moe_grouped_enabled(cfg, prefill) {
5812 let grouped_out = Self::moe_ffn_grouped(e, m, z, t, cfg, il, max_block)?;
5813 if std::env::var("MEMRA_MOE_GATE").is_ok() {
5818 let seq_out = Self::moe_ffn_sequential(e, m, z, t, cfg, il, max_block)?;
5819 let g_host = e.dtoh(&grouped_out)?;
5820 let s_host = e.dtoh(&seq_out)?;
5821 let g_bytes: &[u8] = unsafe {
5822 std::slice::from_raw_parts(g_host.as_ptr() as *const u8, g_host.len() * 4)
5823 };
5824 let s_bytes: &[u8] = unsafe {
5825 std::slice::from_raw_parts(s_host.as_ptr() as *const u8, s_host.len() * 4)
5826 };
5827 if g_bytes == s_bytes {
5828 println!("moe-gate il={il} t={t} BYTE-IDENTICAL");
5829 } else {
5830 let diffs = g_host
5831 .iter()
5832 .zip(s_host.iter())
5833 .enumerate()
5834 .filter(|(_, (a, b))| a != b)
5835 .count();
5836 let maxdiff = g_host
5837 .iter()
5838 .zip(s_host.iter())
5839 .map(|(a, b)| (a - b).abs())
5840 .fold(0.0f32, f32::max);
5841 panic!(
5842 "moe-gate il={il} t={t} MISMATCH: {diffs}/{} elems differ, maxdiff={maxdiff:.6e}",
5843 g_host.len()
5844 );
5845 }
5846 }
5847 return Ok(grouped_out);
5848 }
5849 Self::moe_ffn_sequential_zq8(e, m, z, zq8, t, cfg, il, max_block)
5850 }
5851
5852 fn sigmoid_resident_dev_eligible(
5853 e: &Engine,
5854 m: &MoeWeights,
5855 cfg: &ModelConfig,
5856 sliding_gated_moe: bool,
5857 ) -> bool {
5858 let Some(moe) = cfg.moe.as_ref() else {
5859 return false;
5860 };
5861 static OBSERVATION_MODE: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
5864 let observation_mode = *OBSERVATION_MODE.get_or_init(|| {
5865 std::env::var("MEMRA_MOE_STATS").is_ok()
5866 || std::env::var("MEMRA_MOE_TRACE").is_ok()
5867 || std::env::var("MEMRA_MOE_WEIGHT_TRACE").is_ok()
5868 || std::env::var("MEMRA_MOE_INPUT_TRACE_DIR").is_ok()
5869 || std::env::var("MEMRA_MOE_GATE").is_ok()
5870 });
5871 let resident_layout_supported = m.dev_exps.as_ref().is_some_and(|dev| {
5872 if dev.dev != e.ctx().ordinal() {
5873 return false;
5874 }
5875 let q8 = moe_q8_enabled()
5876 && q8_expert_supported(m.gate_exps.qtype)
5877 && q8_expert_supported(m.up_exps.qtype)
5878 && q8_expert_supported(m.down_exps.qtype);
5879 let fp8 = dev.fp8_blk.is_some()
5880 && m.gate_exps.qtype == crate::QT_F8_E4M3_BLK
5881 && m.up_exps.qtype == crate::QT_F8_E4M3_BLK
5882 && m.down_exps.qtype == crate::QT_F8_E4M3_BLK;
5883 q8 || fp8
5884 });
5885 sliding_gated_moe
5886 && sigmoid_router_enabled()
5887 && moe_dev_enabled()
5888 && moe_slab_enabled()
5889 && !observation_mode
5890 && moe.expert_used_count <= 8
5891 && m.has_uniform_expert_layout()
5892 && m.gate_exps.macros.is_none()
5893 && m.up_exps.macros.is_none()
5894 && m.down_exps.macros.is_none()
5895 && !m.has_macros
5896 && resident_layout_supported
5897 }
5898
5899 pub(crate) fn moe_ffn_sequential(
5901 e: &Engine,
5902 m: &MoeWeights,
5903 z: &CudaSlice<f32>,
5904 t: usize,
5905 cfg: &ModelConfig,
5906 il: u16,
5907 max_block: usize,
5908 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5909 Self::moe_ffn_sequential_zq8(e, m, z, None, t, cfg, il, max_block)
5910 }
5911
5912 fn moe_router_logits(
5916 e: &Engine,
5917 m: &MoeWeights,
5918 z: &CudaSlice<f32>,
5919 t: usize,
5920 cfg: &ModelConfig,
5921 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5922 if t < PRIME_MIN_T {
5923 if crate::router_kernel_on() {
5925 e.router_gemv(
5926 m.gate_inp.float_data(),
5927 z,
5928 cfg.n_embd as usize,
5929 m.gate_exps.n_expert,
5930 t,
5931 )
5932 } else {
5933 e.matmul_decode_exact(&m.gate_inp, z, t)
5934 }
5935 } else if crate::router_prefill_exact_on() && crate::router_kernel_on() {
5936 e.router_gemv(
5937 m.gate_inp.float_data(),
5938 z,
5939 cfg.n_embd as usize,
5940 m.gate_exps.n_expert,
5941 t,
5942 )
5943 } else {
5944 e.matmul(&m.gate_inp, z, t)
5945 }
5946 }
5947
5948 fn trace_moe_routes(
5952 il: u16,
5953 t: usize,
5954 sel_all: &[u32],
5955 weights: &[f32],
5956 ) -> Result<(), Box<dyn std::error::Error>> {
5957 use std::io::Write as _;
5958 if let Ok(path) = std::env::var("MEMRA_MOE_TRACE") {
5959 let mut f = std::fs::OpenOptions::new()
5960 .create(true)
5961 .append(true)
5962 .open(path)?;
5963 let ids: Vec<String> = sel_all.iter().map(|s| s.to_string()).collect();
5964 writeln!(f, "{} {} {}", il, t, ids.join(","))?;
5965 }
5966 if let Ok(path) = std::env::var("MEMRA_MOE_WEIGHT_TRACE") {
5967 let mut f = std::fs::OpenOptions::new()
5968 .create(true)
5969 .append(true)
5970 .open(path)?;
5971 let pairs: Vec<String> = sel_all
5972 .iter()
5973 .zip(weights)
5974 .map(|(expert, weight)| format!("{expert}:{weight:.9}"))
5975 .collect();
5976 writeln!(f, "{} {} {}", il, t, pairs.join(","))?;
5977 }
5978 Ok(())
5979 }
5980
5981 #[allow(clippy::too_many_arguments)]
5982 fn trace_sigmoid_router_logits(
5983 e: &Engine,
5984 il: u16,
5985 t: usize,
5986 n_expert: usize,
5987 n_used: usize,
5988 logits: &CudaSlice<f32>,
5989 m: &MoeWeights,
5990 (scaling_factor, route_norm): (f32, bool),
5991 ) -> Result<(), Box<dyn std::error::Error>> {
5992 if !crate::sigrouter_contract::served_logit_trace_enabled() || t != 1 {
5993 return Ok(());
5994 }
5995 let logits = e.dtoh(logits)?;
5996 let active: Vec<u8> = m
5997 .active_experts
5998 .as_ref()
5999 .map(|mask| mask.iter().map(|&enabled| u8::from(enabled)).collect())
6000 .unwrap_or_else(|| vec![1; n_expert]);
6001 let bias = m.exp_probs_b.clone().unwrap_or_else(|| vec![0.0; n_expert]);
6002 crate::sigrouter_contract::capture_served_logits(
6003 il as u32,
6004 t,
6005 n_expert,
6006 n_used,
6007 scaling_factor,
6008 route_norm,
6009 &active,
6010 &bias,
6011 &logits,
6012 )?;
6013 Ok(())
6014 }
6015
6016 fn trace_moe_input(
6021 e: &Engine,
6022 il: u16,
6023 t: usize,
6024 n_embd: usize,
6025 z: &CudaSlice<f32>,
6026 ) -> Result<(), Box<dyn std::error::Error>> {
6027 use std::io::Write as _;
6028 let Ok(dir) = std::env::var("MEMRA_MOE_INPUT_TRACE_DIR") else {
6029 return Ok(());
6030 };
6031 let values = active_matrix_values(z.len(), t, n_embd, "MoE input trace activation")?;
6032 let host = e.dtoh_view(&z.slice(0..values))?;
6033 let bytes = unsafe {
6034 std::slice::from_raw_parts(
6035 host.as_ptr().cast::<u8>(),
6036 host.len() * std::mem::size_of::<f32>(),
6037 )
6038 };
6039 let state = MOE_INPUT_TRACE_WRITER.get_or_init(|| std::sync::Mutex::new(None));
6040 let mut state = state
6041 .lock()
6042 .map_err(|_| "MoE input trace writer lock is poisoned")?;
6043 if state.is_none() {
6044 let dir = std::path::PathBuf::from(&dir);
6045 std::fs::create_dir_all(&dir)?;
6046 let index = std::fs::OpenOptions::new()
6047 .create(true)
6048 .append(true)
6049 .open(dir.join("index.jsonl"))?;
6050 *state = Some(MoeInputTraceWriter {
6051 dir,
6052 index,
6053 payloads: std::collections::HashMap::new(),
6054 });
6055 }
6056 let writer = state.as_mut().unwrap();
6057 if writer.dir != std::path::Path::new(&dir) {
6058 return Err("MEMRA_MOE_INPUT_TRACE_DIR changed after capture started".into());
6059 }
6060 let file_name = format!("layer-{il:03}.f32");
6061 if !writer.payloads.contains_key(&il) {
6062 let payload = std::fs::OpenOptions::new()
6063 .create(true)
6064 .append(true)
6065 .open(writer.dir.join(&file_name))?;
6066 let offset = payload.metadata()?.len();
6067 writer.payloads.insert(il, (payload, offset));
6068 }
6069 let (payload, offset) = writer.payloads.get_mut(&il).unwrap();
6070 let row_offset = *offset;
6071 payload.write_all(bytes)?;
6072 *offset += bytes.len() as u64;
6073 writeln!(
6074 writer.index,
6075 "{{\"format\":\"memra-moe-input-trace-v1\",\"layer\":{il},\"tokens\":{t},\
6076 \"hidden_size\":{n_embd},\"file\":\"{file_name}\",\"offset\":{row_offset},\
6077 \"payload_bytes\":{}}}",
6078 bytes.len()
6079 )?;
6080 Ok(())
6081 }
6082
6083 #[allow(clippy::too_many_arguments)]
6084 pub(crate) fn moe_ffn_sequential_zq8(
6085 e: &Engine,
6086 m: &MoeWeights,
6087 z: &CudaSlice<f32>,
6088 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
6089 t: usize,
6090 cfg: &ModelConfig,
6091 il: u16,
6092 max_block: usize,
6093 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6094 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
6095 let moe = cfg.moe.as_ref().unwrap();
6096 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);
6103 debug_assert_eq!(m.gate_exps.out_f, n_ff_exp);
6104 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);
6107
6108 let lim_exp = cfg.clamp_exp_at(il as u32);
6111 let lim_shexp = cfg.clamp_shexp_at(il as u32);
6112 let use_cache = Engine::moe_cache_enabled();
6113 let uniform_experts = m.has_uniform_expert_layout();
6114 let moe_q8 = uniform_experts
6115 && moe_q8_enabled()
6116 && q8_expert_supported(m.gate_exps.qtype)
6117 && q8_expert_supported(m.up_exps.qtype)
6118 && q8_expert_supported(m.down_exps.qtype);
6119 let cpu_expert_requested = crate::cpu_experts::configured();
6126 if cpu_expert_requested && (cfg.hy3.is_none() || cfg.m3.is_some()) {
6127 return Err(std::io::Error::other(
6128 "MEMRA_CPU_EXPERT_LIB is experimental and currently gated to Hy3",
6129 )
6130 .into());
6131 }
6132 let cpu_hybrid = cpu_expert_requested && t < PRIME_MIN_T && m.dev_exps.is_none();
6133 let freeze_cpu_residency = cpu_expert_requested
6139 && std::env::var("MEMRA_CPU_EXPERT_FREEZE_CACHE").as_deref() == Ok("1");
6140 let caller_warms_before_freeze = std::env::var("MEMRA_CPU_EXPERT_FREEZE_WARMUP_TOKENS")
6141 .ok()
6142 .and_then(|value| value.parse::<usize>().ok())
6143 .is_some_and(|tokens| tokens > 0);
6144 if cpu_hybrid && freeze_cpu_residency && !caller_warms_before_freeze {
6145 e.freeze_moe_cache();
6146 }
6147 let cache_frozen = use_cache && e.moe_cache_frozen();
6148 let cache_dispatch = use_cache && (!cache_frozen || cpu_hybrid);
6149
6150 let logits = Self::moe_router_logits(e, m, z, t, cfg)?;
6153 if let Some(sig) = cfg.sigmoid_router() {
6154 Self::trace_sigmoid_router_logits(e, il, t, n_expert, n_used, &logits, m, sig)?;
6155 }
6156
6157 let no_exp_macros = m.gate_exps.macros.is_none()
6196 && m.up_exps.macros.is_none()
6197 && m.down_exps.macros.is_none();
6198 if cfg.sigmoid_router().is_none()
6202 && cfg.m3.is_none()
6203 && cfg.hy3.is_none()
6204 && !cfg.swiglu_clamped_at(il as u32)
6205 && no_exp_macros
6206 && t > MOE_DEV_MAX_T
6210 && m.dev_exps.is_some()
6211 && moe_q8_enabled()
6212 && q8_expert_supported(m.gate_exps.qtype)
6213 && q8_expert_supported(m.up_exps.qtype)
6214 && q8_expert_supported(m.down_exps.qtype)
6215 && std::env::var("MEMRA_MOE_PAIRS")
6216 .map(|v| v != "0")
6217 .unwrap_or(true)
6218 && std::env::var("MEMRA_MOE_STATS").is_err()
6219 {
6220 return Self::moe_ffn_pairs(e, m, z, &logits, t, cfg);
6221 }
6222
6223 let dev_ok = uniform_experts
6241 && cfg.sigmoid_router().is_none()
6242 && cfg.m3.is_none()
6243 && cfg.hy3.is_none()
6244 && !cfg.swiglu_clamped_at(il as u32);
6245 let observe_routes = std::env::var("MEMRA_MOE_STATS").is_ok()
6249 || std::env::var("MEMRA_MOE_TRACE").is_ok()
6250 || std::env::var("MEMRA_MOE_WEIGHT_TRACE").is_ok()
6251 || std::env::var("MEMRA_MOE_INPUT_TRACE_DIR").is_ok();
6252 if dev_ok
6253 && t <= MOE_DEV_MAX_T
6254 && m.dev_exps.is_some()
6255 && n_used <= 8
6256 && moe_dev_enabled()
6257 && !observe_routes
6258 {
6259 return Self::moe_ffn_dev(e, m, z, zq8, &logits, t, cfg, il, max_block);
6260 }
6261 if dev_ok && use_cache && n_used <= 8 && moe_dev_enabled() && !observe_routes {
6262 let row_ok = e.with_moe_cache(max_block, |c, eng| {
6263 if moe_prewarm_enabled() {
6264 c.prewarm_layer(il, m, eng)?;
6265 }
6266 Ok(c.layer_dev_row(il, n_expert, eng)?.is_some())
6267 })?;
6268 if row_ok {
6269 return Self::moe_ffn_dev(e, m, z, zq8, &logits, t, cfg, il, max_block);
6270 }
6271 }
6272
6273 let (sel_all, w_all, routed_cpu_input) = if let Some(sig) = cfg.sigmoid_router() {
6275 if cpu_hybrid {
6276 let (sel, w, input) = Self::moe_route_sigmoid_with_input(
6277 e,
6278 &logits,
6279 z,
6280 t,
6281 n_embd,
6282 n_expert,
6283 n_used,
6284 m.exp_probs_b.as_deref(),
6285 sig,
6286 m.active_experts.as_deref(),
6287 )?;
6288 (sel, w, Some(input))
6289 } else {
6290 let (sel, w) =
6291 Self::moe_route_sigmoid_cfg(e, &logits, t, n_expert, n_used, m, sig)?;
6292 (sel, w, None)
6293 }
6294 } else {
6295 let (sel, w) =
6296 Self::moe_route_cfg(e, &logits, t, n_expert, n_used, m.active_experts.as_deref())?;
6297 (sel, w, None)
6298 };
6299 crate::moesd::record_host_routes(il, n_expert, n_used, &sel_all)?;
6300
6301 Self::trace_moe_routes(il, t, &sel_all, &w_all)?;
6305 Self::trace_moe_input(e, il, t, n_embd, z)?;
6306
6307 let worker_disk_prefetch =
6319 cache_dispatch && crate::spill_pread::worker_enabled() && !cpu_hybrid;
6320 let promote_worker_h2d =
6321 t == 1 && worker_disk_prefetch && crate::spill_pread::copy_h2d_enabled();
6322 if promote_worker_h2d {
6323 let mut selected_blocks = Vec::with_capacity(n_used * 3);
6324 for &ex in sel_all.iter().take(n_used) {
6325 let ex = ex as u16;
6326 selected_blocks.extend([
6327 BlockId::new(il, PROJ_GATE, ex),
6328 BlockId::new(il, PROJ_UP, ex),
6329 BlockId::new(il, PROJ_DOWN, ex),
6330 ]);
6331 }
6332 for &ex in sel_all.iter().take(n_used) {
6333 Self::moe_prefetch_disk_expert(e, il, ex as usize, m, max_block, &selected_blocks)?;
6334 }
6335 e.with_moe_cache(max_block, |cache, eng| {
6336 cache.promote_worker_reads_at_safe_boundary(
6337 &selected_blocks,
6338 &selected_blocks,
6339 eng,
6340 )?;
6341 Ok(())
6342 })?;
6343 }
6344
6345 if t > 1 && std::env::var("MEMRA_MOE_STATS").is_ok() {
6348 let mut cnt = vec![0u32; n_expert];
6349 for &s in sel_all.iter() {
6350 cnt[s as usize] += 1;
6351 }
6352 let total = sel_all.len() as f64;
6353 let mut h = 0.0f64;
6354 let mut active = 0usize;
6355 for &c in &cnt {
6356 if c > 0 {
6357 active += 1;
6358 let p = c as f64 / total;
6359 h -= p * p.log2();
6360 }
6361 }
6362 let maxc = cnt.iter().copied().max().unwrap_or(0);
6363 println!(
6364 "moe-stats il={} t={} assignments={} active={}/{} entropy={:.3}b (max {:.3}b) mean_tok_per_active={:.2} max_tok_per_expert={}",
6365 il,
6366 t,
6367 sel_all.len(),
6368 active,
6369 n_expert,
6370 h,
6371 (n_expert as f64).log2(),
6372 total / active.max(1) as f64,
6373 maxc
6374 );
6375 }
6376
6377 let gdec_may_fire = uniform_experts
6390 && use_cache
6391 && n_used <= 8
6392 && gdec_enabled()
6393 && !cfg.swiglu_clamped_at(il as u32);
6394 let slab_local = m
6410 .dev_exps
6411 .as_ref()
6412 .filter(|d| !d.gu_il && moe_slab_enabled() && d.dev == e.ctx().ordinal());
6413 let slab_bases = slab_local.map(|d| {
6414 use cudarc::driver::DevicePtr;
6415 let s = e.stream();
6416 let (pg, _g0) = d.gate.device_ptr(&s);
6417 let (pu, _g1) = d.up.device_ptr(&s);
6418 let (pd, _g2) = d.down.device_ptr(&s);
6419 (pg as u64, pu as u64, pd as u64)
6420 });
6421 let slab_fused_may_fire = slab_bases.is_some()
6431 && n_used <= 8
6432 && gdec_enabled()
6433 && !cfg.swiglu_clamped_at(il as u32)
6434 && cfg.m3.is_none()
6435 && no_exp_macros
6436 && moe_q8;
6437 let mut moe_out = if gdec_may_fire || slab_fused_may_fire {
6440 e.uninit(t * n_embd)?
6441 } else {
6442 e.zeros(t * n_embd)?
6443 };
6444 let cpu_input = if cpu_hybrid {
6447 Some(routed_cpu_input.ok_or("CPU expert routing did not return the MoE input")?)
6448 } else {
6449 None
6450 };
6451
6452 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;
6460 let mut scratch_u: Option<CudaSlice<u8>> = None;
6461 let mut scratch_d: Option<CudaSlice<u8>> = None;
6462 let page_window = moe_page_prefetch_window();
6470
6471 for tok in 0..t {
6474 let sel = &sel_all[tok * n_used..(tok + 1) * n_used];
6475 let w = &w_all[tok * n_used..(tok + 1) * n_used];
6476 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd); let mut tok_q8: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
6478
6479 let no_macros = m.gate_exps.macros.is_none()
6493 && m.up_exps.macros.is_none()
6494 && m.down_exps.macros.is_none();
6495 if slab_fused_may_fire {
6505 let (pg, pu, pd) = slab_bases.unwrap();
6506 let mut gp = [0u64; 8];
6507 let mut up = [0u64; 8];
6508 let mut dp = [0u64; 8];
6509 for (j, &ex) in sel.iter().enumerate() {
6510 let ex = ex as usize;
6511 gp[j] = pg + (ex * m.gate_exps.expert_stride) as u64;
6512 up[j] = pu + (ex * m.up_exps.expert_stride) as u64;
6513 dp[j] = pd + (ex * m.down_exps.expert_stride) as u64;
6514 }
6515 let mut wv = [0f32; 8];
6516 wv[..n_used].copy_from_slice(w);
6517 if tok_q8.is_none() {
6518 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
6519 }
6520 let (zq, zd) = tok_q8.as_ref().unwrap();
6521 let act = e.moe_gate_up_silu8_q8(
6522 crate::WPtr8(gp),
6523 crate::WPtr8(up),
6524 zq,
6525 zd,
6526 n_embd,
6527 n_ff_exp,
6528 n_used,
6529 m.gate_exps.qtype,
6530 m.up_exps.qtype,
6531 m.gate_exps.row_bytes,
6532 m.up_exps.row_bytes,
6533 )?;
6534 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
6535 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
6536 e.moe_down8_fma_q8(
6537 crate::WPtr8(dp),
6538 crate::F32x8(wv),
6539 &aq2,
6540 &ad2,
6541 &mut dst,
6542 n_ff_exp,
6543 n_embd,
6544 n_used,
6545 m.down_exps.qtype,
6546 m.down_exps.row_bytes,
6547 )?;
6548 continue;
6549 }
6550 if gdec_may_fire && moe_q8 && cfg.m3.is_none() && no_macros {
6551 if tok_q8.is_none() {
6552 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
6553 }
6554 let (zq, zd) = tok_q8.as_ref().unwrap();
6555 if Self::moe_gdec_token_q8(
6556 e,
6557 m,
6558 il,
6559 max_block,
6560 zq,
6561 zd,
6562 sel,
6563 w,
6564 &mut moe_out,
6565 tok,
6566 n_embd,
6567 n_ff_exp,
6568 n_used,
6569 )? {
6570 continue;
6571 }
6572 } else if gdec_may_fire
6573 && cfg.m3.is_none()
6574 && no_macros
6575 && Self::moe_gdec_token(
6576 e,
6577 m,
6578 il,
6579 max_block,
6580 &zt,
6581 sel,
6582 w,
6583 &mut moe_out,
6584 tok,
6585 n_embd,
6586 n_ff_exp,
6587 n_used,
6588 )?
6589 {
6590 continue;
6591 }
6592
6593 if gdec_may_fire || slab_fused_may_fire {
6599 let mut row = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
6600 e.memset_zeros_view(&mut row)?;
6601 }
6602
6603 let mut cpu_mask = vec![false; sel.len()];
6609 let cpu_worker = if let Some(host_input) = cpu_input.as_ref() {
6610 let gpu_resident = if use_cache {
6611 e.with_moe_cache(max_block, |cache, _| {
6612 Ok(sel
6613 .iter()
6614 .map(|&expert| {
6615 let expert = expert as u16;
6616 [PROJ_GATE, PROJ_UP, PROJ_DOWN]
6617 .into_iter()
6618 .filter(|&projection| {
6619 cache
6620 .resident(BlockId::new(il, projection, expert))
6621 .is_some()
6622 })
6623 .count()
6624 })
6625 .collect::<Vec<_>>())
6626 })?
6627 } else {
6628 vec![0; sel.len()]
6629 };
6630 let mut cpu_selected = Vec::new();
6631 for (index, (&expert, &route_weight)) in sel.iter().zip(w).enumerate() {
6632 if gpu_resident[index] != 3 {
6633 cpu_mask[index] = true;
6634 crate::cpu_experts::record_incomplete_gpu_residency(gpu_resident[index]);
6635 let expert = expert as usize;
6636 cpu_selected.push((expert, route_weight));
6637 }
6638 }
6639 if crate::cpu_experts::predictor_enabled() {
6640 let row = &host_input[tok * n_embd..(tok + 1) * n_embd];
6644 crate::cpu_experts::predictor_submit(il, row);
6645 }
6646 if cpu_selected.is_empty() {
6647 None
6648 } else {
6649 let row = &host_input[tok * n_embd..(tok + 1) * n_embd];
6650 let job = crate::cpu_experts::prepare_job(m, il, &cpu_selected, row)
6651 .map_err(std::io::Error::other)?;
6652 Some(crate::cpu_experts::submit(job).map_err(std::io::Error::other)?)
6653 }
6654 } else {
6655 None
6656 };
6657
6658 let worker_window = worker_disk_prefetch
6659 .then(worker_prefetch_window)
6660 .unwrap_or(0);
6661 for (j, &ex) in sel.iter().enumerate() {
6662 if cpu_mask[j] {
6663 continue;
6664 }
6665 let ex = ex as usize;
6666 if let Some(d) = slab_local {
6673 let gl = m.gate_exps.expert_layout(ex);
6674 let ul = m.up_exps.expert_layout(ex);
6675 let dl = m.down_exps.expert_layout(ex);
6676 let (g0, u0, d0) = (
6677 ex * m.gate_exps.expert_stride,
6678 ex * m.up_exps.expert_stride,
6679 ex * m.down_exps.expert_stride,
6680 );
6681 let (gate, up) = if moe_q8 {
6682 if tok_q8.is_none() {
6683 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
6684 }
6685 let (zq, zd) = tok_q8.as_ref().unwrap();
6686 (
6687 e.qmatvec_expert_q8(
6688 &d.gate,
6689 g0..g0 + gl.len,
6690 zq,
6691 zd,
6692 1,
6693 m.gate_exps.in_f,
6694 m.gate_exps.out_f,
6695 gl.qtype,
6696 gl.row_bytes,
6697 )?,
6698 e.qmatvec_expert_q8(
6699 &d.up,
6700 u0..u0 + ul.len,
6701 zq,
6702 zd,
6703 1,
6704 m.up_exps.in_f,
6705 m.up_exps.out_f,
6706 ul.qtype,
6707 ul.row_bytes,
6708 )?,
6709 )
6710 } else {
6711 (
6712 e.qmatvec_view(
6713 &d.gate,
6714 g0..g0 + gl.len,
6715 &zt,
6716 1,
6717 m.gate_exps.in_f,
6718 m.gate_exps.out_f,
6719 gl.qtype,
6720 gl.row_bytes,
6721 )?,
6722 e.qmatvec_view(
6723 &d.up,
6724 u0..u0 + ul.len,
6725 &zt,
6726 1,
6727 m.up_exps.in_f,
6728 m.up_exps.out_f,
6729 ul.qtype,
6730 ul.row_bytes,
6731 )?,
6732 )
6733 };
6734 let mut act = e.uninit(n_ff_exp)?;
6735 Self::ffn_act_lim(
6736 e,
6737 cfg,
6738 &gate,
6739 &up,
6740 m.gate_exps.macro_scale(ex),
6741 m.up_exps.macro_scale(ex),
6742 lim_exp,
6743 &mut act,
6744 n_ff_exp,
6745 )?;
6746 let y = if moe_q8 {
6747 let (aq2, ad2) = e.quantize_q8_1(&act, 1, n_ff_exp)?;
6748 e.qmatvec_expert_q8(
6749 &d.down,
6750 d0..d0 + dl.len,
6751 &aq2,
6752 &ad2,
6753 1,
6754 m.down_exps.in_f,
6755 m.down_exps.out_f,
6756 dl.qtype,
6757 dl.row_bytes,
6758 )?
6759 } else {
6760 let actv = act.slice(0..n_ff_exp);
6761 e.qmatvec_view(
6762 &d.down,
6763 d0..d0 + dl.len,
6764 &actv,
6765 1,
6766 m.down_exps.in_f,
6767 m.down_exps.out_f,
6768 dl.qtype,
6769 dl.row_bytes,
6770 )?
6771 };
6772 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
6773 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
6774 continue;
6775 }
6776 for next in page_prefetch_positions(j, sel.len(), page_window) {
6777 Self::moe_prefetch_host_expert(sel[next] as usize, m);
6778 }
6779 let keep = [
6780 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_GATE, ex as u16),
6781 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_UP, ex as u16),
6782 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_DOWN, ex as u16),
6783 ];
6784 if worker_disk_prefetch && worker_window > 0 {
6785 for next in worker_prefetch_positions(j, sel.len(), worker_window) {
6786 Self::moe_prefetch_disk_expert(
6787 e,
6788 il,
6789 sel[next] as usize,
6790 m,
6791 max_block,
6792 &keep,
6793 )?;
6794 }
6795 } else if cache_dispatch
6796 && !cpu_hybrid
6797 && moe_prefetch_enabled()
6798 && j + 1 < sel.len()
6799 {
6800 let next = sel[j + 1] as usize;
6801 Self::moe_prefetch_expert(e, il, next, m, max_block, &keep)?;
6802 }
6803 let [gate_q8, up_q8, down_q8] = [moe_q8; 3];
6804 if cache_dispatch && (gate_q8 || up_q8 || down_q8) {
6805 if (gate_q8 || up_q8) && tok_q8.is_none() {
6808 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
6809 }
6810 let gate = if gate_q8 {
6811 let (zq, zd) = tok_q8.as_ref().unwrap();
6812 Self::moe_cached_gemm_q8(e, il, PROJ_GATE, ex, m, max_block, zq, zd)?
6813 } else {
6814 Self::moe_cached_gemm(e, il, PROJ_GATE, ex, m, max_block, &zt)?
6815 };
6816 let up = if up_q8 {
6817 let (zq, zd) = tok_q8.as_ref().unwrap();
6818 Self::moe_cached_gemm_q8(e, il, PROJ_UP, ex, m, max_block, zq, zd)?
6819 } else {
6820 Self::moe_cached_gemm(e, il, PROJ_UP, ex, m, max_block, &zt)?
6821 };
6822 let mut act = e.uninit(n_ff_exp)?;
6823 Self::ffn_act_lim(
6824 e,
6825 cfg,
6826 &gate,
6827 &up,
6828 m.gate_exps.macro_scale(ex),
6829 m.up_exps.macro_scale(ex),
6830 lim_exp,
6831 &mut act,
6832 n_ff_exp,
6833 )?;
6834 let y = if down_q8 {
6835 let (aq2, ad2) = e.quantize_q8_1(&act, 1, n_ff_exp)?;
6836 Self::moe_cached_gemm_q8(e, il, PROJ_DOWN, ex, m, max_block, &aq2, &ad2)?
6837 } else {
6838 let actv = act.slice(0..n_ff_exp);
6839 Self::moe_cached_gemm(e, il, PROJ_DOWN, ex, m, max_block, &actv)?
6840 };
6841 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
6842 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
6844 } else if cache_dispatch {
6845 let gate = Self::moe_cached_gemm(e, il, PROJ_GATE, ex, m, max_block, &zt)?;
6850 let up = Self::moe_cached_gemm(e, il, PROJ_UP, ex, m, max_block, &zt)?;
6851 let mut act = e.uninit(n_ff_exp)?; Self::ffn_act_lim(
6853 e,
6854 cfg,
6855 &gate,
6856 &up,
6857 m.gate_exps.macro_scale(ex),
6858 m.up_exps.macro_scale(ex),
6859 lim_exp,
6860 &mut act,
6861 n_ff_exp,
6862 )?;
6863 let actv = act.slice(0..n_ff_exp);
6864 let y = Self::moe_cached_gemm(e, il, PROJ_DOWN, ex, m, max_block, &actv)?;
6865 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
6866 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
6868 } else if cache_frozen {
6869 let gate = Self::moe_frozen_gemm(
6874 e,
6875 il,
6876 PROJ_GATE,
6877 ex,
6878 m,
6879 max_block,
6880 &zt,
6881 &mut scratch_g,
6882 g_len,
6883 )?;
6884 let up = Self::moe_frozen_gemm(
6885 e,
6886 il,
6887 PROJ_UP,
6888 ex,
6889 m,
6890 max_block,
6891 &zt,
6892 &mut scratch_u,
6893 u_len,
6894 )?;
6895 let mut act = e.uninit(n_ff_exp)?;
6896 Self::ffn_act_lim(
6897 e,
6898 cfg,
6899 &gate,
6900 &up,
6901 m.gate_exps.macro_scale(ex),
6902 m.up_exps.macro_scale(ex),
6903 lim_exp,
6904 &mut act,
6905 n_ff_exp,
6906 )?;
6907 let actv = act.slice(0..n_ff_exp);
6908 let y = Self::moe_frozen_gemm(
6909 e,
6910 il,
6911 PROJ_DOWN,
6912 ex,
6913 m,
6914 max_block,
6915 &actv,
6916 &mut scratch_d,
6917 d_len,
6918 )?;
6919 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
6920 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
6921 } else {
6922 if scratch_g.is_none() {
6926 scratch_g = Some(e.alloc_u8_uninit(g_len)?);
6927 scratch_u = Some(e.alloc_u8_uninit(u_len)?);
6928 scratch_d = Some(e.alloc_u8_uninit(d_len)?);
6929 }
6930 let (sg, su, sd) = (
6931 scratch_g.as_mut().unwrap(),
6932 scratch_u.as_mut().unwrap(),
6933 scratch_d.as_mut().unwrap(),
6934 );
6935 let gl = m.gate_exps.expert_layout(ex);
6936 let ul = m.up_exps.expert_layout(ex);
6937 let dl = m.down_exps.expert_layout(ex);
6938 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
6939 let gate = e.qmatvec_view(
6940 sg,
6941 0..gl.len,
6942 &zt,
6943 1,
6944 m.gate_exps.in_f,
6945 m.gate_exps.out_f,
6946 gl.qtype,
6947 gl.row_bytes,
6948 )?;
6949
6950 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
6951 let up = e.qmatvec_view(
6952 su,
6953 0..ul.len,
6954 &zt,
6955 1,
6956 m.up_exps.in_f,
6957 m.up_exps.out_f,
6958 ul.qtype,
6959 ul.row_bytes,
6960 )?;
6961
6962 let mut act = e.uninit(n_ff_exp)?; Self::ffn_act_lim(
6964 e,
6965 cfg,
6966 &gate,
6967 &up,
6968 m.gate_exps.macro_scale(ex),
6969 m.up_exps.macro_scale(ex),
6970 lim_exp,
6971 &mut act,
6972 n_ff_exp,
6973 )?;
6974
6975 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
6976 let actv = act.slice(0..n_ff_exp);
6977 let y = e.qmatvec_view(
6978 sd,
6979 0..dl.len,
6980 &actv,
6981 1,
6982 m.down_exps.in_f,
6983 m.down_exps.out_f,
6984 dl.qtype,
6985 dl.row_bytes,
6986 )?;
6987
6988 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
6989 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
6990 }
6991 }
6992 if let Some(worker) = cpu_worker {
6993 let cpu_output = worker.wait().map_err(std::io::Error::other)?;
6994 let cpu_output = e.htod(&cpu_output)?;
6995 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
6996 e.axpy_into(&cpu_output, 1.0, &mut dst, n_embd)?;
6997 }
6998 if cpu_hybrid && !cache_frozen && cpu_expert_profile_admit_enabled() {
6999 for (j, &ex) in sel.iter().enumerate() {
7000 if cpu_mask[j] {
7001 Self::moe_profile_admit_expert(e, il, ex as usize, m, max_block)?;
7002 }
7003 }
7004 }
7005 }
7006
7007 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
7012 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
7013 {
7014 let n_ff_sh = gate_shexp.out_features(); let verify_t = t > 1 && t < PRIME_MIN_T;
7023 let (sg_gate, sg_up) = if t == 1 {
7024 shexp_gate_up_t1(e, gate_shexp, up_shexp, z, zq8)?
7025 } else if verify_t {
7026 (
7027 e.matmul_decode_exact(gate_shexp, z, t)?,
7028 e.matmul_decode_exact(up_shexp, z, t)?,
7029 )
7030 } else {
7031 (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?) };
7033 let mut sa = e.uninit(t * n_ff_sh)?; Self::ffn_act_lim(
7035 e,
7036 cfg,
7037 &sg_gate,
7038 &sg_up,
7039 1.0,
7040 1.0,
7041 lim_shexp,
7042 &mut sa,
7043 t * n_ff_sh,
7044 )?;
7045 let sh = if verify_t {
7046 e.matmul_decode_exact(down_shexp, &sa, t)?
7047 } else {
7048 e.matmul(down_shexp, &sa, t)?
7049 }; let g = match &m.gate_inp_shexp {
7063 Some(gate_inp_shexp) => {
7064 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
7065 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
7066 } else {
7067 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
7068 let mut g = e.uninit(t)?; e.sigmoid(&gs, &mut g, t)?;
7070 g
7071 }
7072 }
7073 None => e.htod(&vec![1.0f32; t])?,
7074 };
7075 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
7077 }
7078
7079 Ok(moe_out)
7080 }
7081
7082 pub fn stage1_h2d_per_token(&self) -> u64 {
7085 use crate::hybrid::Ffn;
7086 let n_used = self
7087 .cfg
7088 .moe
7089 .as_ref()
7090 .map(|m| m.expert_used_count as u64)
7091 .unwrap_or(0);
7092 let mut bytes = 0u64;
7093 for l in self.layers.iter() {
7094 if let Ffn::Moe(m) = &l.ffn {
7095 bytes += n_used
7096 * (m.gate_exps.max_expert_bytes()
7097 + m.up_exps.max_expert_bytes()
7098 + m.down_exps.max_expert_bytes()) as u64;
7099 }
7100 }
7101 bytes
7102 }
7103
7104 pub(crate) fn max_moe_block(&self) -> usize {
7108 use crate::hybrid::Ffn;
7109 let mut mx = 0usize;
7110 let mut scan = |ffn: &Ffn| {
7111 if let Ffn::Moe(m) = ffn {
7112 mx = mx
7113 .max(m.gate_exps.max_expert_bytes())
7114 .max(m.up_exps.max_expert_bytes())
7115 .max(m.down_exps.max_expert_bytes());
7116 }
7117 };
7118 for l in self.layers.iter() {
7119 scan(&l.ffn);
7120 }
7121 if let Some(mtp) = self.mtp.as_ref() {
7122 scan(&mtp.ffn);
7123 }
7124 mx
7125 }
7126
7127 pub(crate) fn moe_cache_block_sizes(&self) -> Vec<usize> {
7130 use crate::hybrid::Ffn;
7131 let mut sizes = Vec::new();
7132 let mut scan = |ffn: &Ffn| {
7133 let Ffn::Moe(m) = ffn else { return };
7134 for ex in 0..m.gate_exps.n_expert {
7135 if m.active_experts.as_ref().is_some_and(|active| !active[ex]) {
7136 continue;
7137 }
7138 for exps in [&m.gate_exps, &m.up_exps, &m.down_exps] {
7139 let len = exps.expert_layout(ex).len;
7140 if len > 0 {
7141 sizes.push(len);
7142 }
7143 }
7144 }
7145 };
7146 for layer in &self.layers {
7147 scan(&layer.ffn);
7148 }
7149 if let Some(mtp) = &self.mtp {
7150 scan(&mtp.ffn);
7151 }
7152 sizes
7153 }
7154
7155 pub fn save_cpu_expert_residency_profile(
7161 &self,
7162 e: &Engine,
7163 path: &std::path::Path,
7164 ) -> Result<(), Box<dyn std::error::Error>> {
7165 let Some(ids) = e.export_moe_residency() else {
7166 return Err("no MoE residency cache to persist".into());
7167 };
7168 let mut body = format!(
7169 "memra-freeze-profile v1 max_block={} blocks={}\n",
7170 self.max_moe_block(),
7171 ids.len()
7172 );
7173 for (layer, proj, ex) in &ids {
7174 body.push_str(&format!("{layer} {proj} {ex}\n"));
7175 }
7176 let tmp = path.with_extension("tmp");
7177 std::fs::write(&tmp, body)?;
7178 std::fs::rename(&tmp, path)?;
7179 println!(
7180 "[moe-cache] freeze profile saved: {} blocks -> {}",
7181 ids.len(),
7182 path.display()
7183 );
7184 Ok(())
7185 }
7186
7187 pub fn restore_cpu_expert_residency_profile(
7191 &self,
7192 e: &Engine,
7193 path: &std::path::Path,
7194 ) -> Result<bool, Box<dyn std::error::Error>> {
7195 use crate::hybrid::Ffn;
7196 use crate::moe_cache::BlockId;
7197 let Ok(content) = std::fs::read_to_string(path) else {
7198 return Ok(false);
7199 };
7200 let mut lines = content.lines();
7201 let Some(header) = lines.next() else {
7202 return Ok(false);
7203 };
7204 let expected = format!("memra-freeze-profile v1 max_block={}", self.max_moe_block());
7205 if !header.starts_with(&expected) {
7206 println!(
7207 "[moe-cache] freeze profile ignored (geometry mismatch): {}",
7208 path.display()
7209 );
7210 return Ok(false);
7211 }
7212 let mut by_layer: std::collections::HashMap<u16, Vec<BlockId>> =
7213 std::collections::HashMap::new();
7214 for line in lines {
7215 let mut fields = line.split_whitespace();
7216 let (Some(layer), Some(proj), Some(ex)) = (fields.next(), fields.next(), fields.next())
7217 else {
7218 continue;
7219 };
7220 let (Ok(layer), Ok(proj), Ok(ex)) =
7221 (layer.parse::<u16>(), proj.parse::<u8>(), ex.parse::<u16>())
7222 else {
7223 continue;
7224 };
7225 by_layer
7226 .entry(layer)
7227 .or_default()
7228 .push(BlockId::new(layer, proj, ex));
7229 }
7230 let requested: usize = by_layer.values().map(Vec::len).sum();
7231 if requested == 0 {
7232 return Ok(false);
7233 }
7234 let max_block = self.max_moe_block();
7235 let mut restaged = 0usize;
7236 let mut stage_layer =
7237 |layer_index: u16, ffn: &Ffn| -> Result<(), Box<dyn std::error::Error>> {
7238 let Ffn::Moe(m) = ffn else { return Ok(()) };
7239 let Some(ids) = by_layer.get(&layer_index) else {
7240 return Ok(());
7241 };
7242 e.with_moe_cache(max_block, |cache, eng| {
7243 for id in ids {
7244 if cache.restage_block(*id, m, eng)? {
7245 restaged += 1;
7246 }
7247 }
7248 Ok(())
7249 })
7250 };
7251 for (index, layer) in self.layers.iter().enumerate() {
7252 stage_layer(index as u16, &layer.ffn)?;
7253 }
7254 if let Some(mtp) = self.mtp.as_ref() {
7255 stage_layer(u16::MAX, &mtp.ffn)?;
7256 }
7257 e.freeze_moe_cache();
7258 println!(
7259 "[moe-cache] freeze profile restored: {restaged}/{requested} blocks restaged from {}",
7260 path.display()
7261 );
7262 Ok(true)
7263 }
7264
7265 pub fn freeze_cpu_expert_residency(
7267 &self,
7268 e: &Engine,
7269 ) -> Result<(), Box<dyn std::error::Error>> {
7270 e.freeze_moe_cache();
7271 Ok(())
7272 }
7273
7274 pub fn ffn_act(
7282 e: &Engine,
7283 cfg: &ModelConfig,
7284 gate: &CudaSlice<f32>,
7285 up: &CudaSlice<f32>,
7286 act: &mut CudaSlice<f32>,
7287 n: usize,
7288 ) -> Result<(), Box<dyn std::error::Error>> {
7289 Self::ffn_act_scaled(e, cfg, gate, up, 1.0, 1.0, act, n)
7290 }
7291
7292 #[allow(clippy::too_many_arguments)]
7296 pub(crate) fn ffn_act_scaled(
7297 e: &Engine,
7298 cfg: &ModelConfig,
7299 gate: &CudaSlice<f32>,
7300 up: &CudaSlice<f32>,
7301 gs: f32,
7302 us: f32,
7303 act: &mut CudaSlice<f32>,
7304 n: usize,
7305 ) -> Result<(), Box<dyn std::error::Error>> {
7306 Self::ffn_act_lim(e, cfg, gate, up, gs, us, None, act, n)
7307 }
7308
7309 #[allow(clippy::too_many_arguments)]
7318 pub(crate) fn ffn_act_lim(
7319 e: &Engine,
7320 cfg: &ModelConfig,
7321 gate: &CudaSlice<f32>,
7322 up: &CudaSlice<f32>,
7323 gs: f32,
7324 us: f32,
7325 limit: Option<f32>,
7326 act: &mut CudaSlice<f32>,
7327 n: usize,
7328 ) -> Result<(), Box<dyn std::error::Error>> {
7329 if let Some(m3) = cfg.m3.as_ref() {
7330 debug_assert!(
7331 limit.is_none(),
7332 "m3 swigluoai and step35 clamp are different archs"
7333 );
7334 return e.swigluoai_mul_scaled(
7335 gate,
7336 up,
7337 gs,
7338 us,
7339 m3.swiglu_alpha,
7340 m3.swiglu_limit,
7341 act,
7342 n,
7343 );
7344 }
7345 if let Some(l) = limit {
7346 return e.swiglu_clamped_mul_scaled(gate, up, gs, us, l, act, n);
7347 }
7348 if gs == 1.0 && us == 1.0 {
7349 return e.silu_mul(gate, up, act, n);
7350 }
7351 e.silu_mul_scaled(gate, up, gs, us, act, n)
7352 }
7353
7354 fn moe_route(
7360 e: &Engine,
7361 logits: &CudaSlice<f32>,
7362 t: usize,
7363 n_expert: usize,
7364 n_used: usize,
7365 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
7366 Self::moe_route_cfg(e, logits, t, n_expert, n_used, None)
7367 }
7368
7369 #[allow(clippy::too_many_arguments)]
7377 fn moe_route_sigmoid_cfg(
7378 e: &Engine,
7379 logits: &CudaSlice<f32>,
7380 t: usize,
7381 n_expert: usize,
7382 n_used: usize,
7383 m: &MoeWeights,
7384 (sf, route_norm): (f32, bool),
7385 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
7386 if sigmoid_router_enabled() {
7387 return e.moe_router_sigmoid_topk_host(
7388 logits,
7389 t,
7390 n_expert,
7391 n_used,
7392 m.active_count(),
7393 &m.exp_probs_b_dev,
7394 &m.active_experts_dev,
7395 sf,
7396 route_norm,
7397 );
7398 }
7399 let lg = e.dtoh(logits)?;
7400 Self::moe_route_sigmoid_host(
7401 &lg,
7402 t,
7403 n_expert,
7404 n_used,
7405 m.exp_probs_b.as_deref(),
7406 sf,
7407 route_norm,
7408 m.active_experts.as_deref(),
7409 )
7410 }
7411
7412 fn moe_route_cfg(
7415 e: &Engine,
7416 logits: &CudaSlice<f32>,
7417 t: usize,
7418 n_expert: usize,
7419 n_used: usize,
7420 active: Option<&[bool]>,
7421 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
7422 if active.is_none() && !matches!(std::env::var("MEMRA_FUSED_ROUTER").as_deref(), Ok("0")) {
7425 return e.moe_router_topk_host(logits, t, n_expert, n_used);
7426 }
7427 let lg = e.dtoh(logits)?; let mut sel = vec![0u32; t * n_used];
7430 let mut w_out = vec![0f32; t * n_used];
7431 for tok in 0..t {
7432 let row = &lg[tok * n_expert..(tok + 1) * n_expert];
7433 let maxl = row
7435 .iter()
7436 .enumerate()
7437 .filter(|(i, _)| active.is_none_or(|mask| mask[*i]))
7438 .map(|(_, &x)| x)
7439 .fold(f32::NEG_INFINITY, f32::max);
7440 let mut probs = vec![0f32; n_expert];
7441 let mut den = 0f32;
7442 for i in 0..n_expert {
7443 if active.is_some_and(|mask| !mask[i]) {
7444 continue;
7445 }
7446 let x = (row[i] - maxl).exp();
7447 probs[i] = x;
7448 den += x;
7449 }
7450 for p in probs.iter_mut() {
7451 *p /= den;
7452 }
7453 let mut idx: Vec<usize> = (0..n_expert)
7455 .filter(|&i| active.is_none_or(|mask| mask[i]))
7456 .collect();
7457 idx.sort_by(|&a, &b| probs[b].total_cmp(&probs[a]).then(a.cmp(&b)));
7458 let sl = &idx[..n_used];
7459 let mut wv: Vec<f32> = sl.iter().map(|&i| probs[i]).collect();
7460 let mut ws: f32 = wv.iter().sum();
7461 ws = ws.max(6.103515625e-5_f32); for x in wv.iter_mut() {
7463 *x /= ws;
7464 }
7465 for j in 0..n_used {
7466 sel[tok * n_used + j] = sl[j] as u32;
7467 w_out[tok * n_used + j] = wv[j];
7468 }
7469 }
7470 Ok((sel, w_out))
7471 }
7472
7473 #[allow(clippy::too_many_arguments)]
7474 fn moe_route_sigmoid_with_input(
7475 e: &Engine,
7476 logits: &CudaSlice<f32>,
7477 input: &CudaSlice<f32>,
7478 t: usize,
7479 in_features: usize,
7480 n_expert: usize,
7481 n_used: usize,
7482 bias: Option<&[f32]>,
7483 (sf, route_norm): (f32, bool),
7484 active: Option<&[bool]>,
7485 ) -> Result<(Vec<u32>, Vec<f32>, Vec<f32>), Box<dyn std::error::Error>> {
7486 let logit_values =
7487 active_matrix_values(logits.len(), t, n_expert, "sigmoid router logits")?;
7488 let input_values =
7489 active_matrix_values(input.len(), t, in_features, "sigmoid router input")?;
7490 let (lg, input) = e.dtoh_pair_views(
7491 &logits.slice(0..logit_values),
7492 &input.slice(0..input_values),
7493 )?;
7494 let (sel, w) =
7495 Self::moe_route_sigmoid_host(&lg, t, n_expert, n_used, bias, sf, route_norm, active)?;
7496 Ok((sel, w, input))
7497 }
7498
7499 pub fn start_moe_prefetch_predictor(
7504 &self,
7505 e: &Engine,
7506 cfg: &ModelConfig,
7507 ) -> Result<(), Box<dyn std::error::Error>> {
7508 use crate::hybrid::Ffn;
7509 let Some(sig) = cfg.sigmoid_router() else {
7510 return Err("prefetch predictor requires a sigmoid-router arch".into());
7511 };
7512 let resident: std::collections::HashSet<(u16, u8, u16)> = e
7513 .export_moe_residency()
7514 .ok_or("prefetch predictor needs the frozen MoE residency cache")?
7515 .into_iter()
7516 .collect();
7517 let mut layers = Vec::new();
7518 for (index, layer) in self.layers.iter().enumerate() {
7519 let Ffn::Moe(m) = &layer.ffn else { continue };
7520 let crate::model::GpuTensor::Float { data, .. } = &m.gate_inp else {
7521 continue;
7522 };
7523 let router = e.dtoh(data)?;
7524 let n_expert = m.gate_exps.n_expert;
7525 let n_embd = m.gate_exps.in_f;
7526 if router.len() != n_embd * n_expert {
7527 continue;
7528 }
7529 let build = |exps: &crate::model::HostExps| {
7530 (0..n_expert)
7531 .map(|expert| crate::cpu_experts::predictor_projection(exps, expert))
7532 .collect::<Vec<_>>()
7533 };
7534 layers.push((
7535 index as u16,
7536 crate::cpu_experts::PredictLayerInit {
7537 router,
7538 bias: m.exp_probs_b.clone(),
7539 active: m.active_experts.clone(),
7540 n_embd,
7541 n_used: cfg
7542 .moe
7543 .as_ref()
7544 .map(|moe| moe.expert_used_count as usize)
7545 .ok_or("prefetch predictor requires MoE config")?,
7546 sig,
7547 weights_n_expert: n_expert,
7548 gate: build(&m.gate_exps),
7549 up: build(&m.up_exps),
7550 down: build(&m.down_exps),
7551 },
7552 ));
7553 }
7554 crate::cpu_experts::start_prefetch_predictor(layers, resident).map_err(|error| error.into())
7555 }
7556
7557 #[allow(clippy::too_many_arguments)]
7560 pub fn moe_route_sigmoid_host_public(
7561 logits: &[f32],
7562 t: usize,
7563 n_expert: usize,
7564 n_used: usize,
7565 bias: Option<&[f32]>,
7566 sf: f32,
7567 route_norm: bool,
7568 active: Option<&[bool]>,
7569 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
7570 Self::moe_route_sigmoid_host(logits, t, n_expert, n_used, bias, sf, route_norm, active)
7571 }
7572
7573 #[allow(clippy::too_many_arguments)]
7574 fn moe_route_sigmoid_host(
7575 lg: &[f32],
7576 t: usize,
7577 n_expert: usize,
7578 n_used: usize,
7579 bias: Option<&[f32]>,
7580 sf: f32,
7581 route_norm: bool,
7582 active: Option<&[bool]>,
7583 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
7584 let active_count = active
7585 .map(|mask| mask.iter().filter(|&&enabled| enabled).count())
7586 .unwrap_or(n_expert);
7587 crate::sigrouter_contract::validate_active_count(n_used, active_count)?;
7588 if lg.len() != t * n_expert {
7589 return Err(format!(
7590 "sigmoid router logits length mismatch: got {}, expected {}",
7591 lg.len(),
7592 t * n_expert,
7593 )
7594 .into());
7595 }
7596 let mut sel = vec![0u32; t * n_used];
7597 let mut w_out = vec![0f32; t * n_used];
7598 for tok in 0..t {
7599 let row = &lg[tok * n_expert..(tok + 1) * n_expert];
7600 let scores: Vec<f32> = row.iter().map(|&x| 1.0 / (1.0 + (-x).exp())).collect();
7601 let selsc: Vec<f32> = match bias {
7603 Some(b) => scores.iter().zip(b).map(|(s, bb)| s + bb).collect(),
7604 None => scores.clone(),
7605 };
7606 let mut idx: Vec<usize> = (0..n_expert)
7607 .filter(|&i| active.is_none_or(|mask| mask[i]))
7608 .collect();
7609 idx.sort_by(|&a, &b| selsc[b].total_cmp(&selsc[a]).then(a.cmp(&b)));
7610 let sl = &idx[..n_used];
7611 let mut wv: Vec<f32> = sl.iter().map(|&i| scores[i]).collect();
7612 if route_norm {
7613 let ws: f32 = wv.iter().sum::<f32>().max(1e-20);
7614 for x in wv.iter_mut() {
7615 *x = *x / ws * sf;
7616 }
7617 } else {
7618 for x in wv.iter_mut() {
7619 *x *= sf;
7620 }
7621 }
7622 for j in 0..n_used {
7623 sel[tok * n_used + j] = sl[j] as u32;
7624 w_out[tok * n_used + j] = wv[j];
7625 }
7626 }
7627 Ok((sel, w_out))
7628 }
7629
7630 #[allow(clippy::too_many_arguments)]
7634 fn moe_ffn_sigmoid_dev(
7635 e: &Engine,
7636 m: &MoeWeights,
7637 z: &CudaSlice<f32>,
7638 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
7639 logits: &CudaSlice<f32>,
7640 t: usize,
7641 cfg: &ModelConfig,
7642 il: u16,
7643 (scaling_factor, route_norm): (f32, bool),
7644 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7645 let moe = cfg.moe.as_ref().unwrap();
7646 let n_embd = cfg.n_embd as usize;
7647 let n_expert = moe.expert_count as usize;
7648 let n_used = moe.expert_used_count as usize;
7649 let n_ff_exp = moe.expert_ff_length as usize;
7650 let dev = m.dev_exps.as_ref().unwrap();
7651 debug_assert_eq!(dev.dev, e.ctx().ordinal());
7652 debug_assert!(m.has_uniform_expert_layout());
7653 debug_assert!(!m.has_macros);
7654
7655 let (sel_d, w_d) = e.moe_router_sigmoid_topk(
7656 logits,
7657 t,
7658 n_expert,
7659 n_used,
7660 m.active_count(),
7661 &m.exp_probs_b_dev,
7662 &m.active_experts_dev,
7663 scaling_factor,
7664 route_norm,
7665 )?;
7666 crate::moesd::record_device_routes(e, il, n_expert, n_used, &sel_d)?;
7667 if let Some(fp8) = dev.fp8_blk.as_ref() {
7668 debug_assert_eq!(m.gate_exps.qtype, crate::QT_F8_E4M3_BLK);
7669 debug_assert_eq!(m.up_exps.qtype, crate::QT_F8_E4M3_BLK);
7670 debug_assert_eq!(m.down_exps.qtype, crate::QT_F8_E4M3_BLK);
7671 debug_assert_eq!(fp8.gate.rows, m.gate_exps.out_f.div_ceil(128));
7672 debug_assert_eq!(fp8.up.rows, m.up_exps.out_f.div_ceil(128));
7673 debug_assert_eq!(fp8.down.rows, m.down_exps.out_f.div_ceil(128));
7674
7675 let selected = e.dtoh_i32(&sel_d)?;
7682 let route_weights = e.dtoh(&w_d)?;
7683 let mut moe_out = e.zeros(t * n_embd)?;
7684 for tok in 0..t {
7685 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
7686 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
7687 for j in 0..n_used {
7688 let pair = tok * n_used + j;
7689 let expert = selected[pair] as usize;
7690 let gate = Self::moe_resident_fp8_e4m3(
7691 e,
7692 &m.gate_exps,
7693 &dev.gate,
7694 &fp8.gate,
7695 expert,
7696 &zt,
7697 1,
7698 )?;
7699 let up = Self::moe_resident_fp8_e4m3(
7700 e, &m.up_exps, &dev.up, &fp8.up, expert, &zt, 1,
7701 )?;
7702 let mut act = e.uninit(n_ff_exp)?;
7703 Self::ffn_act_lim(
7704 e,
7705 cfg,
7706 &gate,
7707 &up,
7708 1.0,
7709 1.0,
7710 cfg.clamp_exp_at(il as u32),
7711 &mut act,
7712 n_ff_exp,
7713 )?;
7714 let act = act.slice(0..n_ff_exp);
7715 let down = Self::moe_resident_fp8_e4m3(
7716 e,
7717 &m.down_exps,
7718 &dev.down,
7719 &fp8.down,
7720 expert,
7721 &act,
7722 1,
7723 )?;
7724 e.axpy_into(&down, route_weights[pair], &mut dst, n_embd)?;
7725 }
7726 }
7727 if std::env::var("MEMRA_SIG_ROUTER_DISPATCH_TRACE").as_deref() == Ok("1") {
7728 eprintln!(
7729 "[sigrouter-dev] layer={il} tokens={t} experts={n_expert} used={n_used} \
7730 native=fp8blk-w8a8-e4m3-reference clamp={}",
7731 cfg.clamp_exp_at(il as u32).is_some(),
7732 );
7733 }
7734 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut moe_out)?;
7735 return Ok(moe_out);
7736 }
7737 let (gate_row_bytes, up_row_bytes) = if dev.gu_il {
7738 let combined = m.gate_exps.row_bytes + m.up_exps.row_bytes;
7739 (combined, combined)
7740 } else {
7741 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
7742 };
7743 let (zq, zd) = match (t, zq8) {
7744 (1, Some((q, d))) => (q.clone(), d.clone()),
7745 _ => e.quantize_q8_1(z, t, n_embd)?,
7746 };
7747 let n_pairs = t * n_used;
7748 let mut moe_out = if cfg.clamp_exp_at(il as u32).is_some() {
7749 let pair_tok: Vec<i32> = (0..n_pairs).map(|pair| (pair / n_used) as i32).collect();
7753 let pair_tok_d = e.htod_i32(&pair_tok)?;
7754 let gate = e.moe_pairs_matvec_q8(
7755 &dev.ptr_row,
7756 0,
7757 &pair_tok_d,
7758 &sel_d,
7759 &zq,
7760 &zd,
7761 n_embd,
7762 n_ff_exp,
7763 n_expert,
7764 n_pairs,
7765 m.gate_exps.qtype,
7766 gate_row_bytes,
7767 )?;
7768 let up = e.moe_pairs_matvec_q8(
7769 &dev.ptr_row,
7770 1,
7771 &pair_tok_d,
7772 &sel_d,
7773 &zq,
7774 &zd,
7775 n_embd,
7776 n_ff_exp,
7777 n_expert,
7778 n_pairs,
7779 m.up_exps.qtype,
7780 up_row_bytes,
7781 )?;
7782 let mut act = e.uninit(n_pairs * n_ff_exp)?;
7783 Self::ffn_act_lim(
7784 e,
7785 cfg,
7786 &gate,
7787 &up,
7788 1.0,
7789 1.0,
7790 cfg.clamp_exp_at(il as u32),
7791 &mut act,
7792 n_pairs * n_ff_exp,
7793 )?;
7794 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
7795 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
7796 let pair_self_d = e.htod_i32(&pair_self)?;
7797 let down = e.moe_pairs_matvec_q8(
7798 &dev.ptr_row,
7799 2,
7800 &pair_self_d,
7801 &sel_d,
7802 &aq2,
7803 &ad2,
7804 n_ff_exp,
7805 n_embd,
7806 n_expert,
7807 n_pairs,
7808 m.down_exps.qtype,
7809 m.down_exps.row_bytes,
7810 )?;
7811 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
7812 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
7813 let tok_off_d = e.htod_i32(&tok_off)?;
7814 let tok_ids_d = e.htod_i32(&tok_ids)?;
7815 let mut output = e.uninit(t * n_embd)?;
7816 e.moe_pairs_scatter(&down, &w_d, &tok_off_d, &tok_ids_d, &mut output, t, n_embd)?;
7817 output
7818 } else {
7819 let act = e.moe_gate_up_silu8_dev_q8_rows(
7820 &dev.ptr_row,
7821 &sel_d,
7822 &zq,
7823 &zd,
7824 t,
7825 n_embd,
7826 n_ff_exp,
7827 n_used,
7828 n_expert,
7829 m.gate_exps.qtype,
7830 m.up_exps.qtype,
7831 gate_row_bytes,
7832 up_row_bytes,
7833 &m.dev_macros,
7834 )?;
7835 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
7836 let mut output = e.uninit(t * n_embd)?;
7837 e.moe_down8_fma_dev_q8_rows_g(
7838 &dev.ptr_row,
7839 &sel_d,
7840 &w_d,
7841 &aq2,
7842 &ad2,
7843 &mut output,
7844 t,
7845 n_ff_exp,
7846 n_embd,
7847 n_used,
7848 n_expert,
7849 m.down_exps.qtype,
7850 m.down_exps.row_bytes,
7851 )?;
7852 output
7853 };
7854
7855 if std::env::var("MEMRA_SIG_ROUTER_DISPATCH_TRACE").as_deref() == Ok("1") {
7856 eprintln!(
7857 "[sigrouter-dev] layer={il} tokens={t} experts={n_expert} used={n_used} clamp={} gu_il={}",
7858 cfg.clamp_exp_at(il as u32).is_some(),
7859 dev.gu_il,
7860 );
7861 }
7862 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut moe_out)?;
7863 Ok(moe_out)
7864 }
7865
7866 #[allow(clippy::too_many_arguments)]
7867 fn moe_resident_fp8_e4m3(
7868 e: &Engine,
7869 exps: &crate::model::HostExps,
7870 bytes: &CudaSlice<u8>,
7871 scales: &crate::hybrid::DevExpertFp8ProjectionScales,
7872 expert: usize,
7873 x: &cudarc::driver::CudaView<f32>,
7874 m: usize,
7875 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7876 let layout = exps.expert_layout(expert);
7877 debug_assert_eq!(layout.qtype, crate::QT_F8_E4M3_BLK);
7878 debug_assert_eq!(scales.rows * scales.cols, scales.expert_stride);
7879 let byte_start = expert * exps.expert_stride;
7880 let scale_start = expert * scales.expert_stride;
7881 let weight = bytes.slice(byte_start..byte_start + layout.len);
7882 let scale = scales
7883 .scales
7884 .slice(scale_start..scale_start + scales.expert_stride);
7885 e.qmatvec_mmq_fp8_blk_view(&weight, &scale, x, m, exps.in_f, exps.out_f)
7886 }
7887
7888 fn moe_ffn_pairs(
7897 e: &Engine,
7898 m: &MoeWeights,
7899 z: &CudaSlice<f32>,
7900 logits: &CudaSlice<f32>,
7901 t: usize,
7902 cfg: &ModelConfig,
7903 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7904 let moe = cfg.moe.as_ref().unwrap();
7905 let n_embd = cfg.n_embd as usize;
7906 let n_expert = moe.expert_count as usize;
7907 let n_used = moe.expert_used_count as usize;
7908 let n_ff_exp = moe.expert_ff_length as usize;
7909 debug_assert!(
7914 !cfg.swiglu_clamped_anywhere(),
7915 "moe_ffn_pairs has no per-layer clamp: fused epilogues are plain SiLU"
7916 );
7917 let dev = m.dev_exps.as_ref().unwrap();
7918 let (rbg_d, rbu_d) = if dev.gu_il {
7920 let sxx = m.gate_exps.row_bytes + m.up_exps.row_bytes;
7921 (sxx, sxx)
7922 } else {
7923 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
7924 };
7925
7926 let (sel_all, w_all) = Self::moe_route(e, logits, t, n_expert, n_used)?;
7927 let n_pairs = t * n_used;
7928 let pair_tok: Vec<i32> = (0..n_pairs).map(|p| (p / n_used) as i32).collect();
7931 let pair_ex: Vec<i32> = sel_all.iter().map(|&x| x as i32).collect();
7932 let pair_w: Vec<f32> = w_all.clone();
7933 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
7934 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
7935 let pt = e.htod_i32(&pair_tok)?;
7936 let px = e.htod_i32(&pair_ex)?;
7937 let pw = e.htod(&pair_w)?;
7938 let toff = e.htod_i32(&tok_off)?;
7939 let tids = e.htod_i32(&tok_ids)?;
7940
7941 let mut by_ex: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
7945 for p in 0..n_pairs {
7946 by_ex[pair_ex[p] as usize].push(p as i32);
7947 }
7948 let mut ex_ids: Vec<i32> = Vec::new();
7949 let mut ex_off: Vec<i32> = vec![0];
7950 let mut ex_pairs: Vec<i32> = Vec::with_capacity(n_pairs);
7951 for (ex, list) in by_ex.iter().enumerate() {
7952 if list.is_empty() {
7953 continue;
7954 }
7955 ex_ids.push(ex as i32);
7956 ex_pairs.extend_from_slice(list);
7957 ex_off.push(ex_pairs.len() as i32);
7958 }
7959 let n_active = ex_ids.len();
7960 let exi = e.htod_i32(&ex_ids)?;
7961 let exo = e.htod_i32(&ex_off)?;
7962 let exp_d = e.htod_i32(&ex_pairs)?;
7963 let _ = &px; static MMA_T: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
7984 let mma_t = *MMA_T.get_or_init(|| {
7985 std::env::var("MEMRA_MOE_MMA_T")
7986 .ok()
7987 .and_then(|v| v.parse().ok())
7988 .unwrap_or(16)
7989 });
7990 let use_mma = std::env::var("MEMRA_MOE_MMA")
7991 .map(|v| v != "0")
7992 .unwrap_or(true)
7993 && t >= mma_t
7994 && q8_expert_dec_supported(m.gate_exps.qtype)
7995 && q8_expert_dec_supported(m.up_exps.qtype)
7996 && q8_expert_dec_supported(m.down_exps.qtype)
7997 && n_embd % 256 == 0
7998 && n_ff_exp % 256 == 0;
7999 let mma_capable = q8_expert_dec_supported(m.gate_exps.qtype)
8015 && q8_expert_dec_supported(m.up_exps.qtype)
8016 && q8_expert_dec_supported(m.down_exps.qtype)
8017 && n_embd % 256 == 0
8018 && n_ff_exp % 256 == 0;
8019 let f16g_mode = crate::moe_f16g_mode();
8020 let f16g = f16g_mode != 0
8021 && t >= mma_t
8022 && (f16g_mode != 3 || !mma_capable)
8023 && f16g_proj_ok(m.gate_exps.qtype, n_embd)
8024 && f16g_proj_ok(m.up_exps.qtype, n_embd)
8025 && f16g_proj_ok(m.down_exps.qtype, n_ff_exp);
8026 if use_mma || f16g {
8027 let y_down = if f16g {
8035 let csr_tok: Vec<i32> = ex_pairs.iter().map(|&p| p / n_used as i32).collect();
8039 let csr_tok_d = e.htod_i32(&csr_tok)?;
8040 let (z_f16, z_s) = e.moe_f16g_act(z, Some(&csr_tok_d), n_embd, n_pairs)?;
8041 let g_csr = e.moe_f16_grouped(
8042 &dev.ptr_row,
8043 0,
8044 n_expert,
8045 &exi,
8046 &ex_off,
8047 &exo,
8048 &z_f16,
8049 &z_s,
8050 n_embd,
8051 n_ff_exp,
8052 n_active,
8053 n_pairs,
8054 m.gate_exps.qtype,
8055 rbg_d,
8056 )?;
8057 let u_csr = e.moe_f16_grouped(
8058 &dev.ptr_row,
8059 1,
8060 n_expert,
8061 &exi,
8062 &ex_off,
8063 &exo,
8064 &z_f16,
8065 &z_s,
8066 n_embd,
8067 n_ff_exp,
8068 n_active,
8069 n_pairs,
8070 m.up_exps.qtype,
8071 rbu_d,
8072 )?;
8073 let act_csr = e.moe_pairs_silu_mul(&g_csr, &u_csr, n_pairs * n_ff_exp)?;
8074 let (a_f16, a_s) = e.moe_f16g_act(&act_csr, None, n_ff_exp, n_pairs)?;
8075 let d_csr = e.moe_f16_grouped(
8076 &dev.ptr_row,
8077 2,
8078 n_expert,
8079 &exi,
8080 &ex_off,
8081 &exo,
8082 &a_f16,
8083 &a_s,
8084 n_ff_exp,
8085 n_embd,
8086 n_active,
8087 n_pairs,
8088 m.down_exps.qtype,
8089 m.down_exps.row_bytes,
8090 )?;
8091 e.rows_permute(&d_csr, &exp_d, n_pairs, n_embd)?
8092 } else {
8093 let z_scr = e.mmq_iq_quantize_act(z, n_embd, t)?;
8095 let gate = e.mmq_iq_experts(
8096 &dev.ptr_row,
8097 0,
8098 n_expert,
8099 &exi,
8100 &exo,
8101 &exp_d,
8102 &pt,
8103 &z_scr,
8104 n_embd,
8105 n_ff_exp,
8106 n_active,
8107 n_pairs,
8108 t,
8109 m.gate_exps.qtype,
8110 rbg_d,
8111 )?;
8112 let up = e.mmq_iq_experts(
8113 &dev.ptr_row,
8114 1,
8115 n_expert,
8116 &exi,
8117 &exo,
8118 &exp_d,
8119 &pt,
8120 &z_scr,
8121 n_embd,
8122 n_ff_exp,
8123 n_active,
8124 n_pairs,
8125 t,
8126 m.up_exps.qtype,
8127 rbu_d,
8128 )?;
8129 let a_scr = if crate::moe_fuse_actq_on() {
8135 e.mmq_iq_fused_act_quant(&gate, &up, n_ff_exp, n_pairs, 0)?
8136 } else {
8137 let act = e.moe_pairs_silu_mul(&gate, &up, n_pairs * n_ff_exp)?;
8138 e.mmq_iq_quantize_act(&act, n_ff_exp, n_pairs)?
8139 };
8140 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
8141 let pself = e.htod_i32(&pair_self)?;
8142 e.mmq_iq_experts(
8143 &dev.ptr_row,
8144 2,
8145 n_expert,
8146 &exi,
8147 &exo,
8148 &exp_d,
8149 &pself,
8150 &a_scr,
8151 n_ff_exp,
8152 n_embd,
8153 n_active,
8154 n_pairs,
8155 n_pairs,
8156 m.down_exps.qtype,
8157 m.down_exps.row_bytes,
8158 )?
8159 };
8160 let mut moe_out = e.uninit(t * n_embd)?;
8161 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
8162 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
8163 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
8164 {
8165 let n_ff_sh = gate_shexp.out_features();
8166 let sg_gate = e.matmul(gate_shexp, z, t)?;
8167 let sg_up = e.matmul(up_shexp, z, t)?;
8168 let mut sa = e.uninit(t * n_ff_sh)?;
8169 Self::ffn_act(e, cfg, &sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
8170 let sh = e.matmul(down_shexp, &sa, t)?;
8171 let g = match &m.gate_inp_shexp {
8177 Some(gate_inp_shexp) if crate::router_prefill_exact_on() => {
8178 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
8179 }
8180 Some(gate_inp_shexp) => {
8181 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
8182 let mut g = e.uninit(t)?;
8183 e.sigmoid(&gs, &mut g, t)?;
8184 g
8185 }
8186 None => e.htod(&vec![1.0f32; t])?,
8187 };
8188 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
8189 }
8190 return Ok(moe_out);
8191 }
8192
8193 let dec = std::env::var("MEMRA_MOE_DEC")
8196 .map(|v| v != "0")
8197 .unwrap_or(true);
8198 let matvec = |proj,
8199 exi: &_,
8200 exo: &_,
8201 exp_d: &_,
8202 pt: &_,
8203 aq: &_,
8204 ad: &_,
8205 inf,
8206 outf,
8207 qtype,
8208 rb|
8209 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8210 let dec = dec && q8_expert_dec_supported(qtype);
8212 if dec {
8213 e.moe_pairs_matvec_q8_dec(
8214 &dev.ptr_row,
8215 proj,
8216 exi,
8217 exo,
8218 exp_d,
8219 pt,
8220 aq,
8221 ad,
8222 inf,
8223 outf,
8224 n_expert,
8225 n_active,
8226 n_pairs,
8227 qtype,
8228 rb,
8229 )
8230 } else {
8231 e.moe_pairs_matvec_q8_em(
8232 &dev.ptr_row,
8233 proj,
8234 exi,
8235 exo,
8236 exp_d,
8237 pt,
8238 aq,
8239 ad,
8240 inf,
8241 outf,
8242 n_expert,
8243 n_active,
8244 n_pairs,
8245 qtype,
8246 rb,
8247 )
8248 }
8249 };
8250 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
8251 let gate = matvec(
8252 0,
8253 &exi,
8254 &exo,
8255 &exp_d,
8256 &pt,
8257 &zq,
8258 &zd,
8259 n_embd,
8260 n_ff_exp,
8261 m.gate_exps.qtype,
8262 rbg_d,
8263 )?;
8264 let up = matvec(
8265 1,
8266 &exi,
8267 &exo,
8268 &exp_d,
8269 &pt,
8270 &zq,
8271 &zd,
8272 n_embd,
8273 n_ff_exp,
8274 m.up_exps.qtype,
8275 rbu_d,
8276 )?;
8277 let act = e.moe_pairs_silu_mul(&gate, &up, n_pairs * n_ff_exp)?;
8278 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
8279 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
8281 let pself = e.htod_i32(&pair_self)?;
8282 let y_down = matvec(
8283 2,
8284 &exi,
8285 &exo,
8286 &exp_d,
8287 &pself,
8288 &aq2,
8289 &ad2,
8290 n_ff_exp,
8291 n_embd,
8292 m.down_exps.qtype,
8293 m.down_exps.row_bytes,
8294 )?;
8295 let mut moe_out = e.uninit(t * n_embd)?; e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
8297
8298 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
8302 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
8303 {
8304 let n_ff_sh = gate_shexp.out_features();
8305 let step_exact = true;
8309 let verify_t = step_exact && t > 1 && t < PRIME_MIN_T;
8310 let (sg_gate, sg_up) = if step_exact && t == 1 {
8311 shexp_gate_up_t1(e, gate_shexp, up_shexp, z, None)?
8312 } else if verify_t {
8313 let mut fused = None;
8314 if crate::spec::spec_fused_t()
8315 && (2..=4).contains(&t)
8316 && e.uses_q8_1_fast(gate_shexp)
8317 && e.uses_q8_1_fast(up_shexp)
8318 {
8319 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
8320 fused = e.matmul_q8_fused2_t(gate_shexp, up_shexp, &zq, &zd, t)?;
8321 }
8322 match fused {
8323 Some(pair) => pair,
8324 None => (
8325 e.matmul_decode_exact(gate_shexp, z, t)?,
8326 e.matmul_decode_exact(up_shexp, z, t)?,
8327 ),
8328 }
8329 } else {
8330 (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?)
8331 };
8332 let mut sa = e.uninit(t * n_ff_sh)?;
8333 e.silu_mul(&sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
8334 let sh = if verify_t {
8335 e.matmul_decode_exact(down_shexp, &sa, t)?
8336 } else {
8337 e.matmul(down_shexp, &sa, t)?
8338 };
8339 let g = match &m.gate_inp_shexp {
8344 Some(gate_inp_shexp) if crate::router_prefill_exact_on() => {
8345 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
8346 }
8347 Some(gate_inp_shexp) => {
8348 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
8349 let mut g = e.uninit(t)?;
8350 e.sigmoid(&gs, &mut g, t)?;
8351 g
8352 }
8353 None => e.htod(&vec![1.0f32; t])?,
8354 };
8355 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
8356 }
8357 Ok(moe_out)
8358 }
8359
8360 #[allow(clippy::too_many_arguments)]
8362 #[allow(clippy::too_many_arguments)]
8363 fn moe_ffn_dev(
8364 e: &Engine,
8365 m: &MoeWeights,
8366 z: &CudaSlice<f32>,
8367 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
8368 logits: &CudaSlice<f32>,
8369 t: usize,
8370 cfg: &ModelConfig,
8371 il: u16,
8372 max_block: usize,
8373 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8374 let moe = cfg.moe.as_ref().unwrap();
8375 let n_embd = cfg.n_embd as usize;
8376 let n_expert = moe.expert_count as usize;
8377 let n_used = moe.expert_used_count as usize;
8378 let n_ff_exp = moe.expert_ff_length as usize;
8379 debug_assert!(
8383 cfg.sigmoid_router().is_none(),
8384 "moe_ffn_dev routes SOFTMAX: a sigmoid-router arch would pick wrong experts"
8385 );
8386 debug_assert!(
8387 !cfg.swiglu_clamped_at(il as u32),
8388 "moe_ffn_dev's fused epilogue is plain SiLU: no clamped form"
8389 );
8390
8391 let (sel_d, mut w_d) = e.moe_router_topk(logits, t, n_expert, n_used)?;
8393 if m.has_macros {
8396 e.moe_w_scale_by_expert(&mut w_d, &sel_d, &m.dev_macros, n_expert, t * n_used)?;
8397 }
8398
8399 let mut moe_out = e.uninit(t * n_embd)?;
8401
8402 if let Some(dev) = m.dev_exps.as_ref() {
8405 let (rbg_d, rbu_d) = if dev.gu_il {
8408 let sxx = m.gate_exps.row_bytes + m.up_exps.row_bytes;
8409 (sxx, sxx)
8410 } else {
8411 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
8412 };
8413 let q8 = moe_q8_enabled()
8414 && q8_expert_supported(m.gate_exps.qtype)
8415 && q8_expert_supported(m.up_exps.qtype)
8416 && q8_expert_supported(m.down_exps.qtype);
8417 let rows_arm = q8
8426 && t > 1
8427 && crate::spec::spec_m2()
8428 && n_ff_exp == 512
8429 && n_used <= 8
8430 && std::env::var("MEMRA_MOE_DEVQ8_GU")
8431 .map(|v| v.is_empty() || v == "v")
8432 .unwrap_or(true)
8433 && std::env::var("MEMRA_MOE_DEVQ8_DOWN")
8434 .map(|v| v.is_empty() || v == "w8h2v")
8435 .unwrap_or(true);
8436 let csr_mode = std::env::var("MEMRA_MOE_CSR")
8445 .ok()
8446 .and_then(|v| v.parse::<i32>().ok())
8447 .unwrap_or(1);
8448 let csr_nvfp4_probe = std::env::var("MEMRA_MOE_CSR_NVFP4").as_deref() == Ok("1");
8465 let csr_qt = |qt: i32| {
8466 qt == crate::QT_IQ4_XS
8467 || qt == crate::QT_IQ3_S
8468 || (csr_nvfp4_probe && qt == crate::QT_NVFP4)
8469 };
8470 let csr_t_max = if csr_nvfp4_probe { MOE_DEV_MAX_T } else { 10 };
8471 let csr_uniform = m.gate_exps.qtype == m.up_exps.qtype;
8472 let csr_arm = rows_arm
8473 && csr_mode > 0
8474 && t <= csr_t_max
8475 && csr_uniform
8476 && csr_qt(m.gate_exps.qtype)
8477 && csr_qt(m.up_exps.qtype)
8478 && csr_qt(m.down_exps.qtype);
8479 if csr_arm {
8480 if csr_mode == 2 {
8481 static ENGAGED: std::sync::Once = std::sync::Once::new();
8482 ENGAGED.call_once(|| eprintln!("[memra] moe CSR byte-compare mode ON (t={t})"));
8483 }
8484 let n_pairs = t * n_used;
8485 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
8486 let act = e.moe_gate_up_silu8_dev_q8_csr(
8487 &dev.ptr_row,
8488 &sel_d,
8489 &zq,
8490 &zd,
8491 n_pairs,
8492 n_embd,
8493 n_ff_exp,
8494 n_used,
8495 n_expert,
8496 m.gate_exps.qtype,
8497 m.up_exps.qtype,
8498 rbg_d,
8499 rbu_d,
8500 )?;
8501 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
8502 e.moe_down8_fma_dev_q8_rows(
8506 &dev.ptr_row,
8507 &sel_d,
8508 &w_d,
8509 &aq2,
8510 &ad2,
8511 &mut moe_out,
8512 t,
8513 n_ff_exp,
8514 n_embd,
8515 n_used,
8516 n_expert,
8517 m.down_exps.qtype,
8518 m.down_exps.row_bytes,
8519 )?;
8520 if csr_mode == 2 {
8521 let act_r = e.moe_gate_up_silu8_dev_q8_rows(
8523 &dev.ptr_row,
8524 &sel_d,
8525 &zq,
8526 &zd,
8527 t,
8528 n_embd,
8529 n_ff_exp,
8530 n_used,
8531 n_expert,
8532 m.gate_exps.qtype,
8533 m.up_exps.qtype,
8534 rbg_d,
8535 rbu_d,
8536 &m.dev_macros,
8537 )?;
8538 let mut out_r = e.uninit(t * n_embd)?;
8539 let (aq2r, ad2r) = e.quantize_q8_1(&act_r, n_pairs, n_ff_exp)?;
8540 e.moe_down8_fma_dev_q8_rows(
8541 &dev.ptr_row,
8542 &sel_d,
8543 &w_d,
8544 &aq2r,
8545 &ad2r,
8546 &mut out_r,
8547 t,
8548 n_ff_exp,
8549 n_embd,
8550 n_used,
8551 n_expert,
8552 m.down_exps.qtype,
8553 m.down_exps.row_bytes,
8554 )?;
8555 let (a1, a2) = (e.dtoh(&act)?, e.dtoh(&act_r)?);
8556 let (o1, o2) = (e.dtoh(&moe_out)?, e.dtoh(&out_r)?);
8557 let ba = a1
8558 .iter()
8559 .zip(&a2)
8560 .filter(|(x, y)| x.to_bits() != y.to_bits())
8561 .count();
8562 let bo = o1
8563 .iter()
8564 .zip(&o2)
8565 .filter(|(x, y)| x.to_bits() != y.to_bits())
8566 .count();
8567 if ba + bo > 0 {
8568 eprintln!(
8569 "[csr-check] il={il} t={t} ACT diffs={ba}/{} OUT diffs={bo}/{}",
8570 a1.len(),
8571 o1.len()
8572 );
8573 let sel_h = e.dtoh_i32(&sel_d)?;
8575 let mut shown = 0;
8576 for (i, (x, y)) in a1.iter().zip(&a2).enumerate() {
8577 if x.to_bits() != y.to_bits() && shown < 4 {
8578 let (p, o) = (i / n_ff_exp, i % n_ff_exp);
8579 let ex = sel_h[p];
8580 let npx = sel_h.iter().filter(|&&v| v == ex).count();
8581 eprintln!(
8582 " ACT p={p} ex={ex} np={npx} o={o} csr={x:e} rows={y:e}"
8583 );
8584 shown += 1;
8585 }
8586 }
8587 std::process::exit(3);
8588 }
8589 }
8590 } else if rows_arm {
8591 if std::env::var("MEMRA_MOE_OVERLAP").as_deref() == Ok("1") {
8594 use std::sync::atomic::{AtomicU64, Ordering};
8595 static PAIRS: AtomicU64 = AtomicU64::new(0);
8596 static UNIQ: AtomicU64 = AtomicU64::new(0);
8597 static CALLS: AtomicU64 = AtomicU64::new(0);
8598 let sel_h = e.dtoh_i32(&sel_d)?;
8599 let mut u: Vec<i32> = sel_h.clone();
8600 u.sort_unstable();
8601 u.dedup();
8602 PAIRS.fetch_add(sel_h.len() as u64, Ordering::Relaxed);
8603 UNIQ.fetch_add(u.len() as u64, Ordering::Relaxed);
8604 let c = CALLS.fetch_add(1, Ordering::Relaxed) + 1;
8605 if c % 480 == 0 {
8606 let p = PAIRS.load(Ordering::Relaxed);
8607 let q = UNIQ.load(Ordering::Relaxed);
8608 eprintln!(
8609 "[overlap] calls={c} pairs={p} unique={q} ratio={:.3} (t={t})",
8610 q as f64 / p as f64
8611 );
8612 }
8613 }
8614 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
8615 let act = e.moe_gate_up_silu8_dev_q8_rows(
8616 &dev.ptr_row,
8617 &sel_d,
8618 &zq,
8619 &zd,
8620 t,
8621 n_embd,
8622 n_ff_exp,
8623 n_used,
8624 n_expert,
8625 m.gate_exps.qtype,
8626 m.up_exps.qtype,
8627 rbg_d,
8628 rbu_d,
8629 &m.dev_macros,
8630 )?;
8631 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
8632 e.moe_down8_fma_dev_q8_rows(
8633 &dev.ptr_row,
8634 &sel_d,
8635 &w_d,
8636 &aq2,
8637 &ad2,
8638 &mut moe_out,
8639 t,
8640 n_ff_exp,
8641 n_embd,
8642 n_used,
8643 n_expert,
8644 m.down_exps.qtype,
8645 m.down_exps.row_bytes,
8646 )?;
8647 } else {
8648 for tok in 0..t {
8649 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
8650 let selt = sel_d.slice(tok * n_used..(tok + 1) * n_used);
8651 let wt = w_d.slice(tok * n_used..(tok + 1) * n_used);
8652 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
8653 if q8 {
8654 let (zq, zd) = match (t, zq8) {
8655 (1, Some((q, d))) => (q.clone(), d.clone()),
8656 _ => e.quantize_q8_1_view(&zt, 1, n_embd)?,
8657 };
8658 let act = e.moe_gate_up_silu8_dev_q8(
8659 &dev.ptr_row,
8660 &selt,
8661 &zq,
8662 &zd,
8663 n_embd,
8664 n_ff_exp,
8665 n_used,
8666 n_expert,
8667 m.gate_exps.qtype,
8668 m.up_exps.qtype,
8669 rbg_d,
8670 rbu_d,
8671 &m.dev_macros,
8672 )?;
8673 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
8674 e.moe_down8_fma_dev_q8(
8675 &dev.ptr_row,
8676 &selt,
8677 &wt,
8678 &aq2,
8679 &ad2,
8680 &mut dst,
8681 n_ff_exp,
8682 n_embd,
8683 n_used,
8684 n_expert,
8685 m.down_exps.qtype,
8686 m.down_exps.row_bytes,
8687 )?;
8688 } else {
8689 let act = e.moe_gate_up_silu8_dev(
8690 &dev.ptr_row,
8691 &selt,
8692 &zt,
8693 n_embd,
8694 n_ff_exp,
8695 n_used,
8696 n_expert,
8697 m.gate_exps.qtype,
8698 m.up_exps.qtype,
8699 rbg_d,
8700 rbu_d,
8701 &m.dev_macros,
8702 )?;
8703 e.moe_down8_fma_dev(
8704 &dev.ptr_row,
8705 &selt,
8706 &wt,
8707 &act,
8708 &mut dst,
8709 n_ff_exp,
8710 n_embd,
8711 n_used,
8712 n_expert,
8713 m.down_exps.qtype,
8714 m.down_exps.row_bytes,
8715 )?;
8716 }
8717 }
8718 }
8719 } else {
8720 let q8 = moe_q8_enabled()
8727 && q8_expert_supported(m.gate_exps.qtype)
8728 && q8_expert_supported(m.up_exps.qtype)
8729 && q8_expert_supported(m.down_exps.qtype);
8730 e.with_moe_cache(max_block, |c, eng| {
8731 let row = c
8732 .layer_dev_row(il, n_expert, eng)?
8733 .ok_or("moe_ffn_dev: layer row vanished under the lock")?;
8734 for tok in 0..t {
8735 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
8736 let selt = sel_d.slice(tok * n_used..(tok + 1) * n_used);
8737 let wt = w_d.slice(tok * n_used..(tok + 1) * n_used);
8738 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
8739 if q8 {
8740 let (zq, zd) = match (t, zq8) {
8741 (1, Some((q, d))) => (q.clone(), d.clone()),
8742 _ => eng.quantize_q8_1_view(&zt, 1, n_embd)?,
8743 };
8744 let act = eng.moe_gate_up_silu8_dev_q8(
8745 row,
8746 &selt,
8747 &zq,
8748 &zd,
8749 n_embd,
8750 n_ff_exp,
8751 n_used,
8752 n_expert,
8753 m.gate_exps.qtype,
8754 m.up_exps.qtype,
8755 m.gate_exps.row_bytes,
8756 m.up_exps.row_bytes,
8757 &m.dev_macros,
8758 )?;
8759 let (aq2, ad2) = eng.quantize_q8_1(&act, n_used, n_ff_exp)?;
8760 eng.moe_down8_fma_dev_q8(
8761 row,
8762 &selt,
8763 &wt,
8764 &aq2,
8765 &ad2,
8766 &mut dst,
8767 n_ff_exp,
8768 n_embd,
8769 n_used,
8770 n_expert,
8771 m.down_exps.qtype,
8772 m.down_exps.row_bytes,
8773 )?;
8774 } else {
8775 let act = eng.moe_gate_up_silu8_dev(
8776 row,
8777 &selt,
8778 &zt,
8779 n_embd,
8780 n_ff_exp,
8781 n_used,
8782 n_expert,
8783 m.gate_exps.qtype,
8784 m.up_exps.qtype,
8785 m.gate_exps.row_bytes,
8786 m.up_exps.row_bytes,
8787 &m.dev_macros,
8788 )?;
8789 eng.moe_down8_fma_dev(
8790 row,
8791 &selt,
8792 &wt,
8793 &act,
8794 &mut dst,
8795 n_ff_exp,
8796 n_embd,
8797 n_used,
8798 n_expert,
8799 m.down_exps.qtype,
8800 m.down_exps.row_bytes,
8801 )?;
8802 }
8803 }
8804 c.hits += (t * 3 * n_used) as u64;
8806 Ok(())
8807 })?;
8808 }
8809
8810 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
8815 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
8816 {
8817 let n_ff_sh = gate_shexp.out_features();
8818 let verify_t = t > 1 && t < PRIME_MIN_T;
8821 let (sg_gate, sg_up) = if t == 1 {
8822 shexp_gate_up_t1(e, gate_shexp, up_shexp, z, zq8)?
8823 } else if verify_t {
8824 let mut fused = None;
8828 if crate::spec::spec_fused_t()
8829 && (2..=4).contains(&t)
8830 && e.uses_q8_1_fast(gate_shexp)
8831 && e.uses_q8_1_fast(up_shexp)
8832 {
8833 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
8834 fused = e.matmul_q8_fused2_t(gate_shexp, up_shexp, &zq, &zd, t)?;
8835 }
8836 match fused {
8837 Some(pair) => pair,
8838 None => (
8839 e.matmul_decode_exact(gate_shexp, z, t)?,
8840 e.matmul_decode_exact(up_shexp, z, t)?,
8841 ),
8842 }
8843 } else {
8844 (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?)
8845 };
8846 let mut sa = e.uninit(t * n_ff_sh)?; e.silu_mul(&sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
8848 let sh = if verify_t {
8849 e.matmul_decode_exact(down_shexp, &sa, t)?
8850 } else {
8851 e.matmul(down_shexp, &sa, t)?
8852 };
8853 let g = match &m.gate_inp_shexp {
8857 Some(gate_inp_shexp) => {
8858 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
8861 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
8862 } else {
8863 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
8864 let mut g = e.uninit(t)?;
8865 e.sigmoid(&gs, &mut g, t)?;
8866 g
8867 }
8868 }
8869 None => e.htod(&vec![1.0f32; t])?,
8870 };
8871 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
8872 }
8873
8874 Ok(moe_out)
8875 }
8876
8877 #[allow(clippy::too_many_arguments)]
8887 #[allow(clippy::too_many_arguments)]
8890 fn moe_gdec_token_q8(
8891 e: &Engine,
8892 m: &MoeWeights,
8893 il: u16,
8894 max_block: usize,
8895 zq: &CudaSlice<i8>,
8896 zd: &CudaSlice<f32>,
8897 sel: &[u32],
8898 w: &[f32],
8899 moe_out: &mut CudaSlice<f32>,
8900 tok: usize,
8901 n_embd: usize,
8902 n_ff_exp: usize,
8903 n_used: usize,
8904 ) -> Result<bool, Box<dyn std::error::Error>> {
8905 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
8906 use cudarc::driver::DevicePtr;
8907 let ptrs = e.with_moe_cache(max_block, |c, eng| {
8908 let mut g = [0u64; 8];
8909 let mut u = [0u64; 8];
8910 let mut d = [0u64; 8];
8911 for (j, &ex) in sel.iter().enumerate() {
8912 let ex = ex as u16;
8913 let (Some(sg), Some(su), Some(sd)) = (
8914 c.resident(BlockId::new(il, PROJ_GATE, ex)),
8915 c.resident(BlockId::new(il, PROJ_UP, ex)),
8916 c.resident(BlockId::new(il, PROJ_DOWN, ex)),
8917 ) else {
8918 return Ok(None);
8919 };
8920 let __s = eng.stream();
8921 let (pg, _e0) = c.slot(sg).device_ptr(&__s);
8922 let (pu, _e1) = c.slot(su).device_ptr(&__s);
8923 let (pd, _e2) = c.slot(sd).device_ptr(&__s);
8924 g[j] = pg as u64;
8925 u[j] = pu as u64;
8926 d[j] = pd as u64;
8927 }
8928 if cpu_expert_profile_admit_enabled() && !c.is_frozen() {
8929 for &ex in sel {
8930 let ex = ex as u16;
8931 for proj in [PROJ_GATE, PROJ_UP, PROJ_DOWN] {
8932 c.note_profile_hit(BlockId::new(il, proj, ex));
8933 }
8934 }
8935 }
8936 c.hits += (3 * n_used) as u64;
8937 Ok(Some((g, u, d)))
8938 })?;
8939 let Some((g, u, d)) = ptrs else {
8940 return Ok(false);
8941 };
8942 let mut wv = [0f32; 8];
8943 wv[..n_used].copy_from_slice(w);
8944 let act = e.moe_gate_up_silu8_q8(
8945 crate::WPtr8(g),
8946 crate::WPtr8(u),
8947 zq,
8948 zd,
8949 n_embd,
8950 n_ff_exp,
8951 n_used,
8952 m.gate_exps.qtype,
8953 m.up_exps.qtype,
8954 m.gate_exps.row_bytes,
8955 m.up_exps.row_bytes,
8956 )?;
8957 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
8959 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
8960 e.moe_down8_fma_q8(
8961 crate::WPtr8(d),
8962 crate::F32x8(wv),
8963 &aq2,
8964 &ad2,
8965 &mut dst,
8966 n_ff_exp,
8967 n_embd,
8968 n_used,
8969 m.down_exps.qtype,
8970 m.down_exps.row_bytes,
8971 )?;
8972 Ok(true)
8973 }
8974
8975 fn moe_gdec_token(
8976 e: &Engine,
8977 m: &MoeWeights,
8978 il: u16,
8979 max_block: usize,
8980 zt: &cudarc::driver::CudaView<f32>,
8981 sel: &[u32],
8982 w: &[f32],
8983 moe_out: &mut CudaSlice<f32>,
8984 tok: usize,
8985 n_embd: usize,
8986 n_ff_exp: usize,
8987 n_used: usize,
8988 ) -> Result<bool, Box<dyn std::error::Error>> {
8989 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
8990 use cudarc::driver::DevicePtr;
8991 let ptrs = e.with_moe_cache(max_block, |c, eng| {
8993 let mut g = [0u64; 8];
8994 let mut u = [0u64; 8];
8995 let mut d = [0u64; 8];
8996 for (j, &ex) in sel.iter().enumerate() {
8997 let ex = ex as u16;
8998 let (Some(sg), Some(su), Some(sd)) = (
8999 c.resident(BlockId::new(il, PROJ_GATE, ex)),
9000 c.resident(BlockId::new(il, PROJ_UP, ex)),
9001 c.resident(BlockId::new(il, PROJ_DOWN, ex)),
9002 ) else {
9003 return Ok(None);
9004 };
9005 let __s = eng.stream();
9006 let (pg, _e0) = c.slot(sg).device_ptr(&__s);
9007 let (pu, _e1) = c.slot(su).device_ptr(&__s);
9008 let (pd, _e2) = c.slot(sd).device_ptr(&__s);
9009 g[j] = pg as u64;
9010 u[j] = pu as u64;
9011 d[j] = pd as u64;
9012 }
9013 if cpu_expert_profile_admit_enabled() && !c.is_frozen() {
9014 for &ex in sel {
9015 let ex = ex as u16;
9016 for proj in [PROJ_GATE, PROJ_UP, PROJ_DOWN] {
9017 c.note_profile_hit(BlockId::new(il, proj, ex));
9018 }
9019 }
9020 }
9021 c.hits += (3 * n_used) as u64; Ok(Some((g, u, d)))
9023 })?;
9024 let Some((g, u, d)) = ptrs else {
9025 return Ok(false);
9026 };
9027 let mut wv = [0f32; 8];
9028 wv[..n_used].copy_from_slice(w);
9029 let act = e.moe_gate_up_silu8(
9031 crate::WPtr8(g),
9032 crate::WPtr8(u),
9033 zt,
9034 n_embd,
9035 n_ff_exp,
9036 n_used,
9037 m.gate_exps.qtype,
9038 m.up_exps.qtype,
9039 m.gate_exps.row_bytes,
9040 m.up_exps.row_bytes,
9041 )?;
9042 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
9043 e.moe_down8_fma_into(
9044 crate::WPtr8(d),
9045 crate::F32x8(wv),
9046 &act,
9047 &mut dst,
9048 n_ff_exp,
9049 n_embd,
9050 n_used,
9051 m.down_exps.qtype,
9052 m.down_exps.row_bytes,
9053 )?;
9054 Ok(true)
9055 }
9056
9057 fn moe_cached_gemm_q8(
9062 e: &Engine,
9063 il: u16,
9064 proj: u8,
9065 ex: usize,
9066 m: &MoeWeights,
9067 max_block: usize,
9068 aq: &CudaSlice<i8>,
9069 ad: &CudaSlice<f32>,
9070 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9071 use crate::moe_cache::{BlockId, DispatchSlot, PROJ_GATE, PROJ_UP};
9072 let exps = match proj {
9073 PROJ_GATE => &m.gate_exps,
9074 PROJ_UP => &m.up_exps,
9075 _ => &m.down_exps,
9076 };
9077 let layout = exps.expert_layout(ex);
9078 let id = BlockId::new(il, proj, ex as u16);
9079 let source = exps.expert_source(ex);
9080 e.with_moe_cache(max_block, |c, eng| {
9081 let slot = c.dispatch_source(id, source, eng)?;
9082 let DispatchSlot::Resident(sl) = slot;
9083 let buf = c.slot(sl);
9084 eng.qmatvec_expert_q8(
9085 buf,
9086 0..layout.len,
9087 aq,
9088 ad,
9089 1,
9090 exps.in_f,
9091 exps.out_f,
9092 layout.qtype,
9093 layout.row_bytes,
9094 )
9095 })
9096 }
9097
9098 fn moe_cached_gemm(
9099 e: &Engine,
9100 il: u16,
9101 proj: u8,
9102 ex: usize,
9103 m: &MoeWeights,
9104 max_block: usize,
9105 x: &cudarc::driver::CudaView<f32>,
9106 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9107 use crate::moe_cache::{BlockId, DispatchSlot, PROJ_GATE, PROJ_UP};
9108 let exps = match proj {
9109 PROJ_GATE => &m.gate_exps,
9110 PROJ_UP => &m.up_exps,
9111 _ => &m.down_exps,
9112 };
9113 let layout = exps.expert_layout(ex);
9114 let id = BlockId::new(il, proj, ex as u16);
9115 let source = exps.expert_source(ex);
9116 e.with_moe_cache(max_block, |c, eng| {
9118 let slot = c.dispatch_source(id, source, eng)?;
9119 let DispatchSlot::Resident(sl) = slot;
9122 let buf = c.slot(sl);
9123 eng.qmatvec_view(
9124 buf,
9125 0..layout.len,
9126 x,
9127 1,
9128 exps.in_f,
9129 exps.out_f,
9130 layout.qtype,
9131 layout.row_bytes,
9132 )
9133 })
9134 }
9135
9136 fn moe_profile_admit_expert(
9140 e: &Engine,
9141 il: u16,
9142 ex: usize,
9143 m: &MoeWeights,
9144 max_block: usize,
9145 ) -> Result<(), Box<dyn std::error::Error>> {
9146 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
9147 e.with_moe_cache(max_block, |cache, eng| {
9148 for (proj, exps) in [
9149 (PROJ_GATE, &m.gate_exps),
9150 (PROJ_UP, &m.up_exps),
9151 (PROJ_DOWN, &m.down_exps),
9152 ] {
9153 let id = BlockId::new(il, proj, ex as u16);
9154 let _ = cache.dispatch_source(id, exps.expert_source(ex), eng)?;
9155 }
9156 Ok(())
9157 })
9158 }
9159
9160 #[allow(clippy::too_many_arguments)]
9163 fn moe_frozen_gemm(
9164 e: &Engine,
9165 il: u16,
9166 proj: u8,
9167 ex: usize,
9168 m: &MoeWeights,
9169 max_block: usize,
9170 x: &cudarc::driver::CudaView<f32>,
9171 scratch: &mut Option<CudaSlice<u8>>,
9172 scratch_len: usize,
9173 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9174 use crate::moe_cache::{BlockId, PROJ_GATE, PROJ_UP};
9175 let exps = match proj {
9176 PROJ_GATE => &m.gate_exps,
9177 PROJ_UP => &m.up_exps,
9178 _ => &m.down_exps,
9179 };
9180 let layout = exps.expert_layout(ex);
9181 let id = BlockId::new(il, proj, ex as u16);
9182 if let Some(output) = e.with_moe_cache(max_block, |cache, eng| {
9183 let Some(slot) = cache.resident(id) else {
9184 return Ok(None);
9185 };
9186 let buf = cache.slot(slot);
9187 Ok(Some(eng.qmatvec_view(
9188 buf,
9189 0..layout.len,
9190 x,
9191 1,
9192 exps.in_f,
9193 exps.out_f,
9194 layout.qtype,
9195 layout.row_bytes,
9196 )?))
9197 })? {
9198 return Ok(output);
9199 }
9200 if scratch.is_none() {
9201 *scratch = Some(e.alloc_u8_uninit(scratch_len)?);
9202 }
9203 let scratch = scratch.as_mut().unwrap();
9204 e.stage_expert(exps.expert_bytes(ex), scratch, 0)?;
9205 e.qmatvec_view(
9206 scratch,
9207 0..layout.len,
9208 x,
9209 1,
9210 exps.in_f,
9211 exps.out_f,
9212 layout.qtype,
9213 layout.row_bytes,
9214 )
9215 }
9216
9217 fn moe_prefetch_expert(
9218 e: &Engine,
9219 il: u16,
9220 ex: usize,
9221 m: &MoeWeights,
9222 max_block: usize,
9223 keep: &[crate::moe_cache::BlockId],
9224 ) -> Result<(), Box<dyn std::error::Error>> {
9225 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
9226 e.with_moe_cache(max_block, |c, eng| {
9227 for (proj, exps) in [
9228 (PROJ_GATE, &m.gate_exps),
9229 (PROJ_UP, &m.up_exps),
9230 (PROJ_DOWN, &m.down_exps),
9231 ] {
9232 let id = BlockId::new(il, proj, ex as u16);
9233 let _ = c.prefetch_source(id, exps.expert_source(ex), keep, eng)?;
9234 }
9235 Ok(())
9236 })
9237 }
9238
9239 fn moe_prefetch_disk_expert(
9242 e: &Engine,
9243 il: u16,
9244 ex: usize,
9245 m: &MoeWeights,
9246 max_block: usize,
9247 keep: &[crate::moe_cache::BlockId],
9248 ) -> Result<(), Box<dyn std::error::Error>> {
9249 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
9250 e.with_moe_cache(max_block, |c, eng| {
9251 for (proj, exps) in [
9252 (PROJ_GATE, &m.gate_exps),
9253 (PROJ_UP, &m.up_exps),
9254 (PROJ_DOWN, &m.down_exps),
9255 ] {
9256 let source = exps.expert_source(ex);
9257 if let crate::model::ExpertSource::Disk { .. } = &source {
9258 let id = BlockId::new(il, proj, ex as u16);
9259 let _ = c.prefetch_source(id, source, keep, eng)?;
9260 }
9261 }
9262 Ok(())
9263 })
9264 }
9265
9266 #[inline]
9267 fn moe_prefetch_host_expert(ex: usize, m: &MoeWeights) {
9268 let _ = m.gate_exps.prefetch_expert_pages(ex);
9269 let _ = m.up_exps.prefetch_expert_pages(ex);
9270 let _ = m.down_exps.prefetch_expert_pages(ex);
9271 }
9272}
9273
9274impl HybridModel {
9291 #[allow(clippy::too_many_arguments)]
9295 fn moe_ffn_grouped_resident_q8(
9296 e: &Engine,
9297 m: &MoeWeights,
9298 z: &CudaSlice<f32>,
9299 t: usize,
9300 cfg: &ModelConfig,
9301 il: u16,
9302 sel_all: &[u32],
9303 w_all: &[f32],
9304 table: &CudaSlice<u64>,
9305 gu_il: bool,
9306 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9307 let moe = cfg.moe.as_ref().unwrap();
9308 let n_embd = cfg.n_embd as usize;
9309 let n_expert = moe.expert_count as usize;
9310 let n_used = moe.expert_used_count as usize;
9311 let n_ff_exp = moe.expert_ff_length as usize;
9312 let n_pairs = t * n_used;
9313 debug_assert_eq!(sel_all.len(), n_pairs);
9314 debug_assert_eq!(w_all.len(), n_pairs);
9315 debug_assert!(
9316 m.gate_exps.macros.is_none()
9317 && m.up_exps.macros.is_none()
9318 && m.down_exps.macros.is_none(),
9319 "resident grouped q8 does not fold per-expert macro scales",
9320 );
9321
9322 if !cfg.swiglu_clamped_at(il as u32) {
9328 let sel: Vec<i32> = sel_all.iter().map(|&expert| expert as i32).collect();
9329 let sel_d = e.htod_i32(&sel)?;
9330 let w_d = e.htod(w_all)?;
9331 let (gate_row_bytes, up_row_bytes) = if gu_il {
9332 let combined = m.gate_exps.row_bytes + m.up_exps.row_bytes;
9333 (combined, combined)
9334 } else {
9335 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
9336 };
9337 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
9338 let act = e.moe_gate_up_silu8_dev_q8_rows(
9339 table,
9340 &sel_d,
9341 &zq,
9342 &zd,
9343 t,
9344 n_embd,
9345 n_ff_exp,
9346 n_used,
9347 n_expert,
9348 m.gate_exps.qtype,
9349 m.up_exps.qtype,
9350 gate_row_bytes,
9351 up_row_bytes,
9352 &m.dev_macros,
9353 )?;
9354 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
9355 let mut moe_out = e.uninit(t * n_embd)?;
9356 e.moe_down8_fma_dev_q8_rows_g(
9357 table,
9358 &sel_d,
9359 &w_d,
9360 &aq2,
9361 &ad2,
9362 &mut moe_out,
9363 t,
9364 n_ff_exp,
9365 n_embd,
9366 n_used,
9367 n_expert,
9368 m.down_exps.qtype,
9369 m.down_exps.row_bytes,
9370 )?;
9371
9372 if std::env::var("MEMRA_MOE_STATS").is_ok() {
9373 let mut counts = vec![0usize; n_expert];
9374 for &expert in sel_all {
9375 counts[expert as usize] += 1;
9376 }
9377 let mut sizes: Vec<usize> =
9378 counts.into_iter().filter(|&count| count != 0).collect();
9379 sizes.sort_unstable();
9380 let mean = sizes.iter().sum::<usize>() as f64 / sizes.len().max(1) as f64;
9381 println!(
9382 "moe-grouped il={il} t={t} dispatch=resident-q8-rows active={}/{} \
9383 m_e: min={} median={} mean={mean:.1} max={}",
9384 sizes.len(),
9385 n_expert,
9386 sizes.first().copied().unwrap_or(0),
9387 sizes.get(sizes.len() / 2).copied().unwrap_or(0),
9388 sizes.last().copied().unwrap_or(0),
9389 );
9390 }
9391 return Ok(moe_out);
9392 }
9393
9394 let pair_tok: Vec<i32> = (0..n_pairs).map(|pair| (pair / n_used) as i32).collect();
9398 let pair_ex: Vec<i32> = sel_all.iter().map(|&expert| expert as i32).collect();
9399 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
9400 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
9401
9402 let mut by_expert: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
9403 for (pair, &expert) in pair_ex.iter().enumerate() {
9404 by_expert[expert as usize].push(pair as i32);
9405 }
9406
9407 let pair_tok_d = e.htod_i32(&pair_tok)?;
9408 let pair_ex_d = e.htod_i32(&pair_ex)?;
9409 let pair_w_d = e.htod(w_all)?;
9410 let tok_off_d = e.htod_i32(&tok_off)?;
9411 let tok_ids_d = e.htod_i32(&tok_ids)?;
9412
9413 let matvec = |proj: i32,
9414 pair_rows: &CudaSlice<i32>,
9415 aq: &CudaSlice<i8>,
9416 ad: &CudaSlice<f32>,
9417 in_f: usize,
9418 out_f: usize,
9419 qtype: i32,
9420 row_bytes: usize|
9421 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9422 e.moe_pairs_matvec_q8(
9423 table, proj, pair_rows, &pair_ex_d, aq, ad, in_f, out_f, n_expert, n_pairs, qtype,
9424 row_bytes,
9425 )
9426 };
9427
9428 let (gate_row_bytes, up_row_bytes) = if gu_il {
9429 let combined = m.gate_exps.row_bytes + m.up_exps.row_bytes;
9430 (combined, combined)
9431 } else {
9432 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
9433 };
9434 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
9435 let gate = matvec(
9436 0,
9437 &pair_tok_d,
9438 &zq,
9439 &zd,
9440 n_embd,
9441 n_ff_exp,
9442 m.gate_exps.qtype,
9443 gate_row_bytes,
9444 )?;
9445 let up = matvec(
9446 1,
9447 &pair_tok_d,
9448 &zq,
9449 &zd,
9450 n_embd,
9451 n_ff_exp,
9452 m.up_exps.qtype,
9453 up_row_bytes,
9454 )?;
9455 let mut act = e.uninit(n_pairs * n_ff_exp)?;
9456 Self::ffn_act_lim(
9457 e,
9458 cfg,
9459 &gate,
9460 &up,
9461 1.0,
9462 1.0,
9463 cfg.clamp_exp_at(il as u32),
9464 &mut act,
9465 n_pairs * n_ff_exp,
9466 )?;
9467 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
9468 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
9469 let pair_self_d = e.htod_i32(&pair_self)?;
9470 let down = matvec(
9471 2,
9472 &pair_self_d,
9473 &aq2,
9474 &ad2,
9475 n_ff_exp,
9476 n_embd,
9477 m.down_exps.qtype,
9478 m.down_exps.row_bytes,
9479 )?;
9480 let mut moe_out = e.uninit(t * n_embd)?;
9481 e.moe_pairs_scatter(
9482 &down,
9483 &pair_w_d,
9484 &tok_off_d,
9485 &tok_ids_d,
9486 &mut moe_out,
9487 t,
9488 n_embd,
9489 )?;
9490
9491 if std::env::var("MEMRA_MOE_STATS").is_ok() {
9492 let mut sizes: Vec<usize> = by_expert
9493 .iter()
9494 .filter_map(|pairs| (!pairs.is_empty()).then_some(pairs.len()))
9495 .collect();
9496 sizes.sort_unstable();
9497 let mean = sizes.iter().sum::<usize>() as f64 / sizes.len().max(1) as f64;
9498 println!(
9499 "moe-grouped il={il} t={t} dispatch=resident-q8-clamped-pairs active={}/{} \
9500 m_e: min={} median={} mean={mean:.1} max={}",
9501 sizes.len(),
9502 n_expert,
9503 sizes.first().copied().unwrap_or(0),
9504 sizes.get(sizes.len() / 2).copied().unwrap_or(0),
9505 sizes.last().copied().unwrap_or(0),
9506 );
9507 }
9508 Ok(moe_out)
9509 }
9510
9511 #[allow(clippy::too_many_arguments)]
9517 fn shexp_split_matvec(
9518 e: &Engine,
9519 rank1: &Engine,
9520 wg: &CudaSlice<u8>,
9521 wu: &CudaSlice<u8>,
9522 wd: &CudaSlice<u8>,
9523 z: &CudaSlice<f32>,
9524 lim: Option<f32>,
9525 cfg: &ModelConfig,
9526 il: u16,
9527 n_embd: usize,
9528 n_ff_sh: usize,
9529 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
9530 use cudarc::driver::DevicePtr;
9531 if n_ff_sh % 2 != 0 || n_embd % 2 != 0 {
9532 return Ok(None);
9533 }
9534 let hf = n_ff_sh / 2;
9535 let nd = n_embd / 2;
9536 struct Rep {
9537 wg1: CudaSlice<u8>,
9538 wu1: CudaSlice<u8>,
9539 wd1: CudaSlice<u8>,
9540 }
9541 struct SplitWs {
9542 pin_dev: usize,
9543 gate0: CudaSlice<f32>,
9545 up0: CudaSlice<f32>,
9546 act: CudaSlice<f32>,
9547 sh_buf: CudaSlice<f32>,
9548 ev_z: cudarc::driver::CudaEvent,
9549 ev_act0: cudarc::driver::CudaEvent,
9550 z1: CudaSlice<f32>,
9552 g1: CudaSlice<f32>,
9553 u1: CudaSlice<f32>,
9554 a1h: CudaSlice<f32>,
9555 act1: CudaSlice<f32>,
9556 y1: CudaSlice<f32>,
9557 ev_act1: cudarc::driver::CudaEvent,
9558 ev_y1: cudarc::driver::CudaEvent,
9559 raw_act_e: u64,
9560 raw_sh_e: u64,
9561 raw_z1: u64,
9562 raw_a1h: u64,
9563 raw_act1: u64,
9564 raw_y1: u64,
9565 }
9566 static WS: std::sync::Mutex<Option<SplitWs>> = std::sync::Mutex::new(None);
9567 static REPS: std::sync::Mutex<Option<std::collections::HashMap<u64, Rep>>> =
9568 std::sync::Mutex::new(None);
9569 let mut guard = WS.lock().map_err(|_| "shexp split lock is poisoned")?;
9570 let mut reps_guard = REPS.lock().map_err(|_| "shexp reps lock is poisoned")?;
9571 let reps = reps_guard.get_or_insert_with(std::collections::HashMap::new);
9572 let pins = e.ctx().ordinal();
9573 if guard.as_ref().is_none_or(|w| w.pin_dev != pins) {
9574 let (gate0, up0, act, sh_buf, ev_z, ev_act0) = {
9575 let _m = e.gpu.enter_main()?;
9576 (
9577 e.htod(&vec![0.0f32; hf])?,
9578 e.htod(&vec![0.0f32; hf])?,
9579 e.htod(&vec![0.0f32; n_ff_sh])?,
9580 e.htod(&vec![0.0f32; n_embd])?,
9581 e.ctx().new_event(None)?,
9582 e.ctx().new_event(None)?,
9583 )
9584 };
9585 let (z1, g1, u1, a1h, act1, y1, ev_act1, ev_y1) = {
9586 let _r = rank1.gpu.enter_main()?;
9587 (
9588 rank1.htod(&vec![0.0f32; n_embd])?,
9589 rank1.htod(&vec![0.0f32; hf])?,
9590 rank1.htod(&vec![0.0f32; hf])?,
9591 rank1.htod(&vec![0.0f32; hf])?,
9592 rank1.htod(&vec![0.0f32; n_ff_sh])?,
9593 rank1.htod(&vec![0.0f32; nd])?,
9594 rank1.ctx().new_event(None)?,
9595 rank1.ctx().new_event(None)?,
9596 )
9597 };
9598 let (raw_act_e, raw_sh_e) = {
9599 let _m = e.gpu.enter_main()?;
9600 let stream = e.stream();
9601 let (a, _g0) = act.device_ptr(&stream);
9602 let (b, _g1) = sh_buf.device_ptr(&stream);
9603 (a as u64, b as u64)
9604 };
9605 let (raw_z1, raw_a1h, raw_act1, raw_y1) = {
9606 let _r = rank1.gpu.enter_main()?;
9607 let rs = rank1.stream();
9608 let (a, _g0) = z1.device_ptr(&rs);
9609 let (b, _g1) = a1h.device_ptr(&rs);
9610 let (c, _g2) = act1.device_ptr(&rs);
9611 let (d, _g3) = y1.device_ptr(&rs);
9612 (a as u64, b as u64, c as u64, d as u64)
9613 };
9614 *guard = Some(SplitWs {
9615 pin_dev: pins,
9616 gate0,
9617 up0,
9618 act,
9619 sh_buf,
9620 ev_z,
9621 ev_act0,
9622 z1,
9623 g1,
9624 u1,
9625 a1h,
9626 act1,
9627 y1,
9628 ev_act1,
9629 ev_y1,
9630 raw_act_e,
9631 raw_sh_e,
9632 raw_z1,
9633 raw_a1h,
9634 raw_act1,
9635 raw_y1,
9636 });
9637 }
9638 let ws = guard.as_mut().expect("armed above");
9639 let wg_pin = {
9640 let _m = e.gpu.enter_main()?;
9641 let stream = e.stream();
9642 let (p, _g) = wg.device_ptr(&stream);
9643 p as u64
9644 };
9645 if !reps.contains_key(&wg_pin) {
9646 let mut up = |src: &CudaSlice<u8>,
9648 off_bytes: usize,
9649 len: usize|
9650 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
9651 use cudarc::driver::sys;
9652 let sptr = {
9653 let _m = e.gpu.enter_main()?;
9654 let stream = e.stream();
9655 let (p, _g) = src.device_ptr(&stream);
9656 p as u64 + off_bytes as u64
9657 };
9658 let dst = {
9659 let _r = rank1.gpu.enter_main()?;
9660 rank1.alloc_u8_uninit(len)?
9661 };
9662 let dptr = {
9663 let _r = rank1.gpu.enter_main()?;
9664 let rs = rank1.stream();
9665 let (p, _g) = dst.device_ptr(&rs);
9666 p as u64
9667 };
9668 let _r = rank1.gpu.enter_main()?;
9669 let r = unsafe {
9670 sys::cuMemcpyAsync(
9671 dptr as sys::CUdeviceptr,
9672 sptr as sys::CUdeviceptr,
9673 len,
9674 rank1.stream().cu_stream() as sys::CUstream,
9675 )
9676 };
9677 if r != sys::CUresult::CUDA_SUCCESS {
9678 return Err(format!("shexp split replica upload: {r:?}").into());
9679 }
9680 rank1.stream().synchronize()?;
9681 Ok(dst)
9682 };
9683 let wg1 = up(wg, hf * n_embd * 2, hf * n_embd * 2)?;
9684 let wu1 = up(wu, hf * n_embd * 2, hf * n_embd * 2)?;
9685 let wd1 = up(wd, nd * n_ff_sh * 2, nd * n_ff_sh * 2)?;
9686 reps.insert(wg_pin, Rep { wg1, wu1, wd1 });
9687 }
9688 let _ = il;
9689 let raw_z = {
9691 let _m = e.gpu.enter_main()?;
9692 let stream = e.stream();
9693 let (p, _g) = z.device_ptr(&stream);
9694 ws.ev_z.record(&stream)?;
9695 p as u64
9696 };
9697 {
9699 let rep = reps.get(&wg_pin).expect("uploaded above");
9700 let _r = rank1.gpu.enter_main()?;
9701 rank1.stream().wait(&ws.ev_z)?;
9702 crate::tp::raw_copy_bytes(ws.raw_z1, raw_z, n_embd * 4, rank1)?;
9703 let SplitWs {
9704 z1, g1, u1, a1h, ..
9705 } = &mut *ws;
9706 rank1.matvec_bf16_dual_into(&rep.wg1, &rep.wu1, z1, g1, u1, n_embd, hf)?;
9707 Self::ffn_act_lim(rank1, cfg, g1, u1, 1.0, 1.0, lim, a1h, hf)?;
9708 crate::tp::raw_copy_bytes(ws.raw_act1 + (hf * 4) as u64, ws.raw_a1h, hf * 4, rank1)?;
9710 crate::tp::raw_copy_bytes(ws.raw_act_e + (hf * 4) as u64, ws.raw_a1h, hf * 4, rank1)?;
9711 ws.ev_act1.record(&rank1.stream())?;
9712 }
9713 {
9715 let _m = e.gpu.enter_main()?;
9716 let SplitWs {
9717 gate0, up0, act, ..
9718 } = &mut *ws;
9719 let wg_lo = wg.slice(0..hf * n_embd * 2);
9720 let wu_lo = wu.slice(0..hf * n_embd * 2);
9721 e.matvec_bf16_dual_view_into(&wg_lo, &wu_lo, z, gate0, up0, n_embd, hf)?;
9722 Self::ffn_act_lim(e, cfg, gate0, up0, 1.0, 1.0, lim, act, hf)?;
9723 ws.ev_act0.record(&e.stream())?;
9724 }
9725 {
9727 let rep = reps.get(&wg_pin).expect("uploaded above");
9728 let _r = rank1.gpu.enter_main()?;
9729 rank1.stream().wait(&ws.ev_act0)?;
9730 crate::tp::raw_copy_bytes(ws.raw_act1, ws.raw_act_e, hf * 4, rank1)?;
9731 let SplitWs { act1, y1, .. } = &mut *ws;
9732 rank1.matvec_bf16_into(&rep.wd1, act1, y1, n_ff_sh, nd)?;
9733 crate::tp::raw_copy_bytes(ws.raw_sh_e + (nd * 4) as u64, ws.raw_y1, nd * 4, rank1)?;
9734 ws.ev_y1.record(&rank1.stream())?;
9735 }
9736 {
9738 let _m = e.gpu.enter_main()?;
9739 e.stream().wait(&ws.ev_act1)?;
9740 let SplitWs { act, sh_buf, .. } = &mut *ws;
9741 let wd_lo = wd.slice(0..nd * n_ff_sh * 2);
9742 e.matvec_bf16_view_into(&wd_lo, act, sh_buf, n_ff_sh, nd)?;
9743 e.stream().wait(&ws.ev_y1)?;
9744 let mut sh = e.uninit(n_embd)?;
9745 {
9746 let mut dst = sh.slice_mut(0..n_embd);
9747 e.stream()
9748 .memcpy_dtod(&ws.sh_buf.slice(0..n_embd), &mut dst)?;
9749 }
9750 Ok(Some(sh))
9751 }
9752 }
9753
9754 fn shexp_overlap_issue(
9761 e: &Engine,
9762 m: &MoeWeights,
9763 z: &CudaSlice<f32>,
9764 cfg: &ModelConfig,
9765 il: u16,
9766 n_embd: usize,
9767 ) -> Result<bool, Box<dyn std::error::Error>> {
9768 if cfg.m3.is_some() || m.gate_inp_shexp.is_some() {
9769 return Ok(false);
9770 }
9771 let (
9772 Some(crate::model::GpuTensor::FloatBf16 { data: wg, .. }),
9773 Some(crate::model::GpuTensor::FloatBf16 { data: wu, .. }),
9774 Some(crate::model::GpuTensor::FloatBf16 { data: wd, .. }),
9775 ) = (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
9776 else {
9777 return Ok(false);
9778 };
9779 let n_ff_sh = m
9780 .gate_shexp
9781 .as_ref()
9782 .expect("matched Some above")
9783 .out_features();
9784 let lim = cfg.clamp_shexp_at(il as u32);
9785 let mut guard = SHEXP_OV_WS
9786 .lock()
9787 .map_err(|_| "shexp overlap workspace lock is poisoned")?;
9788 let pins = (e.ctx().ordinal(), n_embd, n_ff_sh);
9789 if guard
9790 .as_ref()
9791 .is_none_or(|(d, ne, nf, ..)| (*d, *ne, *nf) != pins)
9792 {
9793 *guard = Some((
9794 pins.0,
9795 pins.1,
9796 pins.2,
9797 e.uninit(n_ff_sh)?,
9798 e.uninit(n_embd)?,
9799 ));
9800 }
9801 let (_, _, _, act, sh) = guard.as_mut().expect("armed above");
9802 e.matvec_bf16_dual_silu_into(wg, wu, z, act, n_embd, n_ff_sh, lim)?;
9803 e.matvec_bf16_into(wd, act, sh, n_ff_sh, n_embd)?;
9804 drop(guard);
9805 Ok(true)
9806 }
9807
9808 #[allow(clippy::too_many_arguments)]
9814 fn shexp_dev1_issue(
9815 e: &Engine,
9816 rank1: &Engine,
9817 m: &MoeWeights,
9818 z: &CudaSlice<f32>,
9819 cfg: &ModelConfig,
9820 il: u16,
9821 n_embd: usize,
9822 ) -> Result<bool, Box<dyn std::error::Error>> {
9823 use cudarc::driver::DevicePtr;
9824 if cfg.m3.is_some() || m.gate_inp_shexp.is_some() {
9825 return Ok(false);
9826 }
9827 let (
9828 Some(crate::model::GpuTensor::FloatBf16 { data: wg, .. }),
9829 Some(crate::model::GpuTensor::FloatBf16 { data: wu, .. }),
9830 Some(crate::model::GpuTensor::FloatBf16 { data: wd, .. }),
9831 ) = (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
9832 else {
9833 return Ok(false);
9834 };
9835 let n_ff_sh = m
9836 .gate_shexp
9837 .as_ref()
9838 .expect("matched Some above")
9839 .out_features();
9840 let lim = cfg.clamp_shexp_at(il as u32);
9841 let mut ws_guard = SHEXP_D1_WS
9843 .lock()
9844 .map_err(|_| "shexp dev1 workspace lock is poisoned")?;
9845 if ws_guard
9846 .as_ref()
9847 .is_none_or(|(k, ..)| *k != (n_embd, n_ff_sh))
9848 {
9849 let (act1, z1, ev_done) = {
9850 let _r1 = rank1.gpu.enter_main()?;
9851 (
9852 rank1.htod(&vec![0.0f32; n_ff_sh])?,
9853 rank1.htod(&vec![0.0f32; n_embd])?,
9854 rank1.ctx().new_event(None)?,
9855 )
9856 };
9857 let (sh_root, ev_z) = {
9858 let _main = e.gpu.enter_main()?;
9859 (e.htod(&vec![0.0f32; n_embd])?, e.ctx().new_event(None)?)
9860 };
9861 *ws_guard = Some(((n_embd, n_ff_sh), act1, z1, sh_root, ev_z, ev_done));
9862 }
9863 let mut reps_guard = SHEXP_D1_REPS
9865 .lock()
9866 .map_err(|_| "shexp dev1 replica lock is poisoned")?;
9867 let reps = reps_guard.get_or_insert_with(Default::default);
9868 if !reps.contains_key(&il) {
9869 let (wg1, wu1, wd1) = {
9870 let _r1 = rank1.gpu.enter_main()?;
9871 (
9872 rank1.alloc_u8_uninit(n_ff_sh * n_embd * 2)?,
9873 rank1.alloc_u8_uninit(n_ff_sh * n_embd * 2)?,
9874 rank1.alloc_u8_uninit(n_embd * n_ff_sh * 2)?,
9875 )
9876 };
9877 for (src, dst) in [(wg, &wg1), (wu, &wu1), (wd, &wd1)] {
9878 let s_ptr = {
9879 let _main = e.gpu.enter_main()?;
9880 let stream = e.stream();
9881 let (p, _g) = src.device_ptr(&stream);
9882 p as u64
9883 };
9884 let d_ptr = {
9885 let _r1 = rank1.gpu.enter_main()?;
9886 let stream = rank1.stream();
9887 let (p, _g) = dst.device_ptr(&stream);
9888 p as u64
9889 };
9890 let _r1 = rank1.gpu.enter_main()?;
9891 crate::tp::raw_copy_bytes(d_ptr, s_ptr, src.len(), rank1)?;
9892 }
9893 {
9894 let _r1 = rank1.gpu.enter_main()?;
9895 rank1.stream().synchronize()?;
9896 }
9897 reps.insert(il, (wg1, wu1, wd1));
9898 }
9899 let (wg1, wu1, wd1) = reps.get(&il).expect("armed above");
9900 let (_, act1, z1, sh_root, ev_z, ev_done) = ws_guard.as_mut().expect("armed above");
9901 let (raw_z, raw_sh) = {
9904 let _main = e.gpu.enter_main()?;
9905 let stream = e.stream();
9906 let (a, _g0) = z.device_ptr(&stream);
9907 let (b, _g1) = sh_root.device_ptr(&stream);
9908 ev_z.record(&stream)?;
9909 (a as u64, b as u64)
9910 };
9911 {
9912 let _r1 = rank1.gpu.enter_main()?;
9913 rank1.stream().wait(ev_z)?;
9914 let raw_z1 = {
9915 let stream = rank1.stream();
9916 let (p, _g) = z1.device_ptr(&stream);
9917 p as u64
9918 };
9919 crate::tp::raw_copy_bytes(raw_z1, raw_z, n_embd * 4, rank1)?;
9920 rank1.matvec_bf16_dual_silu_into(wg1, wu1, z1, act1, n_embd, n_ff_sh, lim)?;
9921 rank1.matvec_bf16_raw_out(wd1, act1, raw_sh, n_ff_sh, n_embd)?;
9925 ev_done.record(&rank1.stream())?;
9926 }
9927 Ok(true)
9928 }
9929
9930 fn shexp_dev1_apply(
9932 e: &Engine,
9933 output: &mut CudaSlice<f32>,
9934 n_embd: usize,
9935 ) -> Result<(), Box<dyn std::error::Error>> {
9936 let guard = SHEXP_D1_WS
9937 .lock()
9938 .map_err(|_| "shexp dev1 workspace lock is poisoned")?;
9939 let (pin, _, _, sh_root, _, ev_done) =
9940 guard.as_ref().ok_or("shexp dev1 apply without issue")?;
9941 if pin.0 != n_embd {
9942 return Err("shexp dev1 width drifted".into());
9943 }
9944 let _main = e.gpu.enter_main()?;
9945 e.stream().wait(ev_done)?;
9946 static ONES_D1: std::sync::Mutex<Option<(usize, CudaSlice<f32>)>> =
9947 std::sync::Mutex::new(None);
9948 let mut og = ONES_D1.lock().map_err(|_| "ones lock is poisoned")?;
9949 if og.as_ref().is_none_or(|(d, _)| *d != e.ctx().ordinal()) {
9950 *og = Some((e.ctx().ordinal(), e.htod(&[1.0f32])?));
9951 }
9952 let ones = &og.as_ref().expect("armed above").1;
9953 e.add_scaled_rows(sh_root, ones, output, n_embd, 1)?;
9954 Ok(())
9955 }
9956
9957 fn shexp_overlap_tail_ptrs(
9961 e: &Engine,
9962 m: &MoeWeights,
9963 cfg: &ModelConfig,
9964 n_embd: usize,
9965 ) -> Result<Option<(u64, u64)>, Box<dyn std::error::Error>> {
9966 use cudarc::driver::DevicePtr;
9967 if cfg.m3.is_some() || m.gate_inp_shexp.is_some() {
9968 return Ok(None);
9969 }
9970 let (
9971 Some(crate::model::GpuTensor::FloatBf16 { .. }),
9972 Some(crate::model::GpuTensor::FloatBf16 { .. }),
9973 Some(crate::model::GpuTensor::FloatBf16 { .. }),
9974 ) = (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
9975 else {
9976 return Ok(None);
9977 };
9978 let n_ff_sh = m
9979 .gate_shexp
9980 .as_ref()
9981 .expect("matched Some above")
9982 .out_features();
9983 let mut guard = SHEXP_OV_WS
9984 .lock()
9985 .map_err(|_| "shexp overlap workspace lock is poisoned")?;
9986 let pins = (e.ctx().ordinal(), n_embd, n_ff_sh);
9987 if guard
9988 .as_ref()
9989 .is_none_or(|(d, ne, nf, ..)| (*d, *ne, *nf) != pins)
9990 {
9991 *guard = Some((
9992 pins.0,
9993 pins.1,
9994 pins.2,
9995 e.uninit(n_ff_sh)?,
9996 e.uninit(n_embd)?,
9997 ));
9998 }
9999 let sh_raw = {
10000 let (_, _, _, _, sh) = guard.as_ref().expect("armed above");
10001 let stream = e.stream();
10002 let (p, _g) = sh.device_ptr(&stream);
10003 p as u64
10004 };
10005 static ONES_T3: std::sync::Mutex<Option<(usize, CudaSlice<f32>)>> =
10006 std::sync::Mutex::new(None);
10007 let mut og = ONES_T3.lock().map_err(|_| "ones lock is poisoned")?;
10008 if og.as_ref().is_none_or(|(d, _)| *d != e.ctx().ordinal()) {
10009 *og = Some((e.ctx().ordinal(), e.htod(&[1.0f32])?));
10010 }
10011 let ones_raw = {
10012 let stream = e.stream();
10013 let (p, _g) = og.as_ref().expect("armed above").1.device_ptr(&stream);
10014 p as u64
10015 };
10016 Ok(Some((sh_raw, ones_raw)))
10017 }
10018
10019 fn shexp_overlap_apply(
10022 e: &Engine,
10023 output: &mut CudaSlice<f32>,
10024 n_embd: usize,
10025 ) -> Result<(), Box<dyn std::error::Error>> {
10026 let guard = SHEXP_OV_WS
10027 .lock()
10028 .map_err(|_| "shexp overlap workspace lock is poisoned")?;
10029 let (_, ne, _, _, sh) = guard.as_ref().ok_or("shexp overlap apply without issue")?;
10030 if *ne != n_embd {
10031 return Err("shexp overlap width drifted".into());
10032 }
10033 static ONES_OV: std::sync::Mutex<Option<(usize, CudaSlice<f32>)>> =
10034 std::sync::Mutex::new(None);
10035 let mut og = ONES_OV.lock().map_err(|_| "ones lock is poisoned")?;
10036 if og.as_ref().is_none_or(|(d, _)| *d != e.ctx().ordinal()) {
10037 *og = Some((e.ctx().ordinal(), e.htod(&[1.0f32])?));
10038 }
10039 let ones = &og.as_ref().expect("armed above").1;
10040 e.add_scaled_rows(sh, ones, output, n_embd, 1)?;
10041 Ok(())
10042 }
10043
10044 fn moe_ffn_grouped_add_shared(
10045 e: &Engine,
10046 m: &MoeWeights,
10047 z: &CudaSlice<f32>,
10048 t: usize,
10049 cfg: &ModelConfig,
10050 il: u16,
10051 moe_out: &mut CudaSlice<f32>,
10052 ) -> Result<(), Box<dyn std::error::Error>> {
10053 static SHEXP_NS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
10056 static SHEXP_CALLS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
10057 let shexp_timing = std::env::var("MEMRA_STEP_TP_TIMING").as_deref() == Ok("1");
10058 let shexp_started = shexp_timing.then(std::time::Instant::now);
10059 let result = Self::moe_ffn_grouped_add_shared_inner(e, m, z, t, cfg, il, moe_out);
10060 if let Some(started) = shexp_started {
10061 use std::sync::atomic::Ordering;
10062 e.stream().synchronize()?;
10063 let ns = SHEXP_NS.fetch_add(started.elapsed().as_nanos() as u64, Ordering::Relaxed)
10064 + started.elapsed().as_nanos() as u64;
10065 let calls = SHEXP_CALLS.fetch_add(1, Ordering::Relaxed) + 1;
10066 if calls % 430 == 0 {
10067 eprintln!(
10068 "[moe-shexp-timing] calls={calls} total_ms={:.1} avg_us={:.1}",
10069 ns as f64 / 1.0e6,
10070 ns as f64 / calls as f64 / 1.0e3,
10071 );
10072 }
10073 }
10074 result
10075 }
10076
10077 #[allow(clippy::too_many_arguments)]
10078 fn moe_ffn_grouped_add_shared_inner(
10079 e: &Engine,
10080 m: &MoeWeights,
10081 z: &CudaSlice<f32>,
10082 t: usize,
10083 cfg: &ModelConfig,
10084 il: u16,
10085 moe_out: &mut CudaSlice<f32>,
10086 ) -> Result<(), Box<dyn std::error::Error>> {
10087 let n_embd = cfg.n_embd as usize;
10088 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
10089 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
10090 {
10091 let n_ff_sh = gate_shexp.out_features();
10092 let lim = cfg.clamp_shexp_at(il as u32);
10093 let fused = t == 1
10100 && lim.is_none()
10101 && cfg.m3.is_none()
10102 && e.uses_q8_1_fast(gate_shexp)
10103 && e.uses_q8_1_fast(up_shexp);
10104 let bf16_dual = if t == 1 && crate::Engine::bf16_mmv_on() && n_embd % 8 == 0 {
10107 match (gate_shexp, up_shexp) {
10108 (
10109 crate::model::GpuTensor::FloatBf16 { data: wg, .. },
10110 crate::model::GpuTensor::FloatBf16 { data: wu, .. },
10111 ) => Some((wg, wu)),
10112 _ => None,
10113 }
10114 } else {
10115 None
10116 };
10117 let sh = if let Some((wg, wu)) = bf16_dual {
10118 static SHEXP_WS: std::sync::Mutex<
10122 Option<(
10123 usize,
10124 usize,
10125 usize,
10126 CudaSlice<f32>,
10127 CudaSlice<f32>,
10128 CudaSlice<f32>,
10129 CudaSlice<f32>,
10130 )>,
10131 > = std::sync::Mutex::new(None);
10132 let down_bf16 = match down_shexp {
10133 crate::model::GpuTensor::FloatBf16 { data, .. } => Some(data),
10134 _ => None,
10135 };
10136 let mut guard = SHEXP_WS
10137 .lock()
10138 .map_err(|_| "shexp workspace lock is poisoned")?;
10139 let pins = (e.ctx().ordinal(), n_embd, n_ff_sh);
10140 if guard
10141 .as_ref()
10142 .is_none_or(|(d, ne, nf, ..)| (*d, *ne, *nf) != pins)
10143 {
10144 *guard = Some((
10145 pins.0,
10146 pins.1,
10147 pins.2,
10148 e.uninit(n_ff_sh)?,
10149 e.uninit(n_ff_sh)?,
10150 e.uninit(n_ff_sh)?,
10151 e.uninit(n_embd)?,
10152 ));
10153 }
10154 {
10157 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10158 let split_on = *ON
10159 .get_or_init(|| std::env::var("MEMRA_SHEXP_SPLIT").as_deref() == Ok("1"));
10160 if split_on {
10161 if let (Some(wd), Some(rank1)) = (
10162 match down_shexp {
10163 crate::model::GpuTensor::FloatBf16 { data, .. } => Some(data),
10164 _ => None,
10165 },
10166 m.step_tp.as_ref().and_then(|st| st.runtime.rank_engine(1)),
10167 ) {
10168 if let Some(sh) = Self::shexp_split_matvec(
10169 e, rank1, wg, wu, wd, z, lim, cfg, il, n_embd, n_ff_sh,
10170 )? {
10171 drop(guard);
10172 let gate = match &m.gate_inp_shexp {
10173 Some(gate_inp_shexp) => e.sigmoid_dot_rows(
10174 z,
10175 gate_inp_shexp.float_data(),
10176 n_embd,
10177 t,
10178 )?,
10179 None => e.htod(&vec![1.0f32; t])?,
10180 };
10181 e.add_scaled_rows(&sh, &gate, moe_out, n_embd, t)?;
10182 return Ok(());
10183 }
10184 }
10185 }
10186 }
10187 let (_, _, _, gate, up, act, sh_buf) =
10188 guard.as_mut().expect("shexp workspace initialized above");
10189 if cfg.m3.is_none() {
10190 e.matvec_bf16_dual_silu_into(wg, wu, z, act, n_embd, n_ff_sh, lim)?;
10193 let _ = (&gate, &up);
10194 } else {
10195 e.matvec_bf16_dual_into(wg, wu, z, gate, up, n_embd, n_ff_sh)?;
10196 Self::ffn_act_lim(e, cfg, gate, up, 1.0, 1.0, lim, act, n_ff_sh)?;
10197 }
10198 if let Some(down) = down_bf16 {
10199 static FUSE_DA: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10206 let fuse_da = *FUSE_DA.get_or_init(|| {
10207 std::env::var("MEMRA_FUSE_DOWN_ADDSCALE").as_deref() != Ok("0")
10208 });
10209 if fuse_da && m.gate_inp_shexp.is_none() {
10210 static ONES1: std::sync::Mutex<Option<(usize, CudaSlice<f32>)>> =
10211 std::sync::Mutex::new(None);
10212 let mut og = ONES1.lock().map_err(|_| "shexp ones lock is poisoned")?;
10213 if og.as_ref().is_none_or(|(d, _)| *d != e.ctx().ordinal()) {
10214 *og = Some((e.ctx().ordinal(), e.htod(&[1.0f32])?));
10215 }
10216 let ones = &og.as_ref().expect("armed above").1;
10217 e.matvec_bf16_down_addscale_into(
10218 down, act, ones, moe_out, n_ff_sh, n_embd,
10219 )?;
10220 return Ok(());
10221 }
10222 e.matvec_bf16_into(down, act, sh_buf, n_ff_sh, n_embd)?;
10223 let sh = e.uninit(n_embd)?;
10224 let mut sh = sh;
10226 {
10227 let mut dst = sh.slice_mut(0..n_embd);
10228 e.stream().memcpy_dtod(&sh_buf.slice(0..n_embd), &mut dst)?;
10229 }
10230 sh
10231 } else {
10232 e.matmul(down_shexp, act, 1)?
10233 }
10234 } else if fused {
10235 let (zq, zd) = e.quantize_q8_1(z, 1, n_embd)?;
10236 let pair = match e.matmul_pre_dual_noscale(gate_shexp, up_shexp, &zq, &zd, 1)? {
10237 Some((gate, up)) => Some((gate, up)),
10238 None => {
10239 match (
10240 e.matmul_pre_noscale(gate_shexp, &zq, &zd, 1)?,
10241 e.matmul_pre_noscale(up_shexp, &zq, &zd, 1)?,
10242 ) {
10243 (Some(gate), Some(up)) => Some((gate, up)),
10244 _ => None,
10245 }
10246 }
10247 };
10248 match pair {
10249 Some(((gate, gs), (up, us))) => {
10250 if e.uses_q8_1_fast(down_shexp) {
10251 let (aq, ad) = e.silu_mul_scaled_q8_1(&gate, &up, gs, us, n_ff_sh)?;
10252 e.matmul_pre(down_shexp, &aq, &ad, &gate, 1)?
10253 } else {
10254 let mut act = e.uninit(n_ff_sh)?;
10255 e.silu_mul_scaled(&gate, &up, gs, us, &mut act, n_ff_sh)?;
10256 e.matmul(down_shexp, &act, 1)?
10257 }
10258 }
10259 None => {
10260 let gate = e.matmul_pre(gate_shexp, &zq, &zd, z, 1)?;
10261 let up = e.matmul_pre(up_shexp, &zq, &zd, z, 1)?;
10262 let mut act = e.uninit(n_ff_sh)?;
10263 Self::ffn_act(e, cfg, &gate, &up, &mut act, n_ff_sh)?;
10264 e.matmul(down_shexp, &act, 1)?
10265 }
10266 }
10267 } else {
10268 let sg_gate = e.matmul(gate_shexp, z, t)?;
10269 let sg_up = e.matmul(up_shexp, z, t)?;
10270 let mut sa = e.uninit(t * n_ff_sh)?;
10271 Self::ffn_act_lim(
10272 e,
10273 cfg,
10274 &sg_gate,
10275 &sg_up,
10276 1.0,
10277 1.0,
10278 lim,
10279 &mut sa,
10280 t * n_ff_sh,
10281 )?;
10282 e.matmul(down_shexp, &sa, t)?
10283 };
10284 let gate = match &m.gate_inp_shexp {
10285 Some(gate_inp_shexp) => {
10286 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
10287 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
10288 } else {
10289 let raw = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
10290 let mut gate = e.uninit(t)?;
10291 e.sigmoid(&raw, &mut gate, t)?;
10292 gate
10293 }
10294 }
10295 None if t == 1 => {
10300 static ONES: std::sync::Mutex<Option<(usize, CudaSlice<f32>)>> =
10301 std::sync::Mutex::new(None);
10302 let mut guard = ONES.lock().map_err(|_| "shexp ones lock is poisoned")?;
10303 if guard.as_ref().is_none_or(|(d, _)| *d != e.ctx().ordinal()) {
10304 *guard = Some((e.ctx().ordinal(), e.htod(&[1.0f32])?));
10305 }
10306 let ones = &guard.as_ref().expect("armed above").1;
10307 e.add_scaled_rows(&sh, ones, moe_out, n_embd, t)?;
10308 return Ok(());
10309 }
10310 None => e.htod(&vec![1.0f32; t])?,
10311 };
10312 e.add_scaled_rows(&sh, &gate, moe_out, n_embd, t)?;
10313 }
10314 Ok(())
10315 }
10316
10317 pub(crate) fn moe_ffn_grouped(
10320 e: &Engine,
10321 m: &MoeWeights,
10322 z: &CudaSlice<f32>,
10323 t: usize,
10324 cfg: &ModelConfig,
10325 il: u16,
10326 max_block: usize,
10327 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
10328 let moe = cfg.moe.as_ref().unwrap();
10329 let n_embd = cfg.n_embd as usize;
10330 let n_expert = moe.expert_count as usize;
10331 let n_used = moe.expert_used_count as usize;
10332 let n_ff_exp = moe.expert_ff_length as usize;
10333 let lim_exp = cfg.clamp_exp_at(il as u32);
10335
10336 let logits = Self::moe_router_logits(e, m, z, t, cfg)?;
10340 if let Some(sig) = cfg.sigmoid_router() {
10341 Self::trace_sigmoid_router_logits(e, il, t, n_expert, n_used, &logits, m, sig)?;
10342 }
10343 let (sel_all, w_all) = if let Some(sig) = cfg.sigmoid_router() {
10344 Self::moe_route_sigmoid_cfg(e, &logits, t, n_expert, n_used, m, sig)?
10345 } else {
10346 Self::moe_route_cfg(e, &logits, t, n_expert, n_used, m.active_experts.as_deref())?
10347 };
10348 crate::moesd::record_host_routes(il, n_expert, n_used, &sel_all)?;
10349 Self::trace_moe_routes(il, t, &sel_all, &w_all)?;
10350 Self::trace_moe_input(e, il, t, n_embd, z)?;
10351
10352 let no_exp_macros = m.gate_exps.macros.is_none()
10357 && m.up_exps.macros.is_none()
10358 && m.down_exps.macros.is_none();
10359 let resident_q8 = m.dev_exps.as_ref().filter(|dev| {
10360 m.has_uniform_expert_layout()
10361 && no_exp_macros
10362 && moe_q8_enabled()
10363 && q8_expert_supported(m.gate_exps.qtype)
10364 && q8_expert_supported(m.up_exps.qtype)
10365 && q8_expert_supported(m.down_exps.qtype)
10366 && moe_slab_enabled()
10367 && dev.dev == e.ctx().ordinal()
10368 });
10369 if let Some(dev) = resident_q8 {
10370 let mut moe_out = Self::moe_ffn_grouped_resident_q8(
10371 e,
10372 m,
10373 z,
10374 t,
10375 cfg,
10376 il,
10377 &sel_all,
10378 &w_all,
10379 &dev.ptr_row,
10380 dev.gu_il,
10381 )?;
10382 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut moe_out)?;
10383 return Ok(moe_out);
10384 }
10385
10386 struct ExpertGroup {
10390 tok_indices: Vec<i32>, slot_indices: Vec<i32>, weights: Vec<f32>, }
10394 let mut groups: Vec<ExpertGroup> = (0..n_expert)
10395 .map(|_| ExpertGroup {
10396 tok_indices: Vec::new(),
10397 slot_indices: Vec::new(),
10398 weights: Vec::new(),
10399 })
10400 .collect();
10401
10402 for tok in 0..t {
10403 for j in 0..n_used {
10404 let ex = sel_all[tok * n_used + j] as usize;
10405 let w = w_all[tok * n_used + j];
10406 groups[ex].tok_indices.push(tok as i32);
10407 groups[ex].slot_indices.push(j as i32);
10408 groups[ex].weights.push(w);
10409 }
10410 }
10411
10412 let mut slot_buf = e.zeros(t * n_used * n_embd)?;
10415 let mut wbuf = e.zeros(t * n_used)?; let g_len = m.gate_exps.max_expert_bytes();
10419 let u_len = m.up_exps.max_expert_bytes();
10420 let d_len = m.down_exps.max_expert_bytes();
10421 let moe_q8 = m.has_uniform_expert_layout()
10422 && moe_q8_enabled()
10423 && q8_expert_supported(m.gate_exps.qtype)
10424 && q8_expert_supported(m.up_exps.qtype)
10425 && q8_expert_supported(m.down_exps.qtype);
10426 let slab_local = m
10429 .dev_exps
10430 .as_ref()
10431 .filter(|dev| !dev.gu_il && moe_slab_enabled() && dev.dev == e.ctx().ordinal());
10432 let use_cache =
10433 slab_local.is_none() && Engine::moe_cache_enabled() && !e.moe_cache_frozen();
10434 let grouped_q8 = moe_q8 && (slab_local.is_some() || use_cache);
10437
10438 let (mut scratch_g, mut scratch_u, mut scratch_d) = if slab_local.is_none() && !use_cache {
10440 (
10441 Some(e.alloc_u8(g_len)?),
10442 Some(e.alloc_u8(u_len)?),
10443 Some(e.alloc_u8(d_len)?),
10444 )
10445 } else {
10446 (None, None, None)
10447 };
10448
10449 let mut order: Vec<usize> = (0..n_expert)
10460 .filter(|&ex| !groups[ex].tok_indices.is_empty())
10461 .collect();
10462 order.sort_by(|&a, &b| {
10463 groups[b]
10464 .tok_indices
10465 .len()
10466 .cmp(&groups[a].tok_indices.len())
10467 .then(a.cmp(&b))
10468 });
10469 let mut m_dist: Vec<usize> = Vec::new(); let page_window = moe_page_prefetch_window();
10471 let worker_disk_prefetch = use_cache && crate::spill_pread::worker_enabled();
10472 if worker_disk_prefetch {
10473 if let Some(first) = grouped_worker_prefetch_position(order.len(), None) {
10474 Self::moe_prefetch_disk_expert(e, il, order[first], m, max_block, &[])?;
10475 }
10476 }
10477 for (order_pos, &ex) in order.iter().enumerate() {
10478 for next in page_prefetch_positions(order_pos, order.len(), page_window) {
10479 Self::moe_prefetch_host_expert(order[next], m);
10480 }
10481 if worker_disk_prefetch {
10482 if let Some(next) = grouped_worker_prefetch_position(order.len(), Some(order_pos)) {
10483 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
10484 let keep = [
10485 BlockId::new(il, PROJ_GATE, ex as u16),
10486 BlockId::new(il, PROJ_UP, ex as u16),
10487 BlockId::new(il, PROJ_DOWN, ex as u16),
10488 ];
10489 Self::moe_prefetch_disk_expert(e, il, order[next], m, max_block, &keep)?;
10490 }
10491 }
10492 let grp = &groups[ex];
10493 let m_e = grp.tok_indices.len();
10494 m_dist.push(m_e);
10495 let gl = m.gate_exps.expert_layout(ex);
10496 let ul = m.up_exps.expert_layout(ex);
10497 let dl = m.down_exps.expert_layout(ex);
10498
10499 let tok_idx_d = e.htod_i32(&grp.tok_indices)?;
10503 let slot_idx_d = e.htod_i32(&grp.slot_indices)?;
10504 let dmac = m.down_exps.macro_scale(ex);
10505 let weight_d = if dmac == 1.0 {
10506 e.htod(&grp.weights)?
10507 } else {
10508 let scaled: Vec<f32> = grp.weights.iter().map(|&w| w * dmac).collect();
10509 e.htod(&scaled)?
10510 };
10511
10512 let mut gathered = e.zeros(m_e * n_embd)?;
10514 e.gather_rows(z, &tok_idx_d, &mut gathered, n_embd, m_e)?;
10515 let gv = gathered.slice(0..m_e * n_embd);
10516
10517 let y = if let Some(dev) = slab_local {
10520 let gate_start = ex * m.gate_exps.expert_stride;
10521 let up_start = ex * m.up_exps.expert_stride;
10522 let down_start = ex * m.down_exps.expert_stride;
10523 if grouped_q8 {
10524 let (zq, zd) = e.quantize_q8_1(&gathered, m_e, n_embd)?;
10525 let gate = e.qmatvec_expert_q8(
10526 &dev.gate,
10527 gate_start..gate_start + gl.len,
10528 &zq,
10529 &zd,
10530 m_e,
10531 m.gate_exps.in_f,
10532 m.gate_exps.out_f,
10533 gl.qtype,
10534 gl.row_bytes,
10535 )?;
10536 let up = e.qmatvec_expert_q8(
10537 &dev.up,
10538 up_start..up_start + ul.len,
10539 &zq,
10540 &zd,
10541 m_e,
10542 m.up_exps.in_f,
10543 m.up_exps.out_f,
10544 ul.qtype,
10545 ul.row_bytes,
10546 )?;
10547 let mut act = e.uninit(m_e * n_ff_exp)?;
10548 Self::ffn_act_lim(
10549 e,
10550 cfg,
10551 &gate,
10552 &up,
10553 m.gate_exps.macro_scale(ex),
10554 m.up_exps.macro_scale(ex),
10555 lim_exp,
10556 &mut act,
10557 m_e * n_ff_exp,
10558 )?;
10559 let (aq2, ad2) = e.quantize_q8_1(&act, m_e, n_ff_exp)?;
10560 e.qmatvec_expert_q8(
10561 &dev.down,
10562 down_start..down_start + dl.len,
10563 &aq2,
10564 &ad2,
10565 m_e,
10566 m.down_exps.in_f,
10567 m.down_exps.out_f,
10568 dl.qtype,
10569 dl.row_bytes,
10570 )?
10571 } else {
10572 let gate = e.qmatvec_view(
10573 &dev.gate,
10574 gate_start..gate_start + gl.len,
10575 &gv,
10576 m_e,
10577 m.gate_exps.in_f,
10578 m.gate_exps.out_f,
10579 gl.qtype,
10580 gl.row_bytes,
10581 )?;
10582 let up = e.qmatvec_view(
10583 &dev.up,
10584 up_start..up_start + ul.len,
10585 &gv,
10586 m_e,
10587 m.up_exps.in_f,
10588 m.up_exps.out_f,
10589 ul.qtype,
10590 ul.row_bytes,
10591 )?;
10592 let mut act = e.uninit(m_e * n_ff_exp)?;
10593 Self::ffn_act_lim(
10594 e,
10595 cfg,
10596 &gate,
10597 &up,
10598 m.gate_exps.macro_scale(ex),
10599 m.up_exps.macro_scale(ex),
10600 lim_exp,
10601 &mut act,
10602 m_e * n_ff_exp,
10603 )?;
10604 let actv = act.slice(0..m_e * n_ff_exp);
10605 e.qmatvec_view(
10606 &dev.down,
10607 down_start..down_start + dl.len,
10608 &actv,
10609 m_e,
10610 m.down_exps.in_f,
10611 m.down_exps.out_f,
10612 dl.qtype,
10613 dl.row_bytes,
10614 )?
10615 }
10616 } else if use_cache {
10617 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
10618 if grouped_q8 {
10619 let (zq, zd) = e.quantize_q8_1(&gathered, m_e, n_embd)?;
10620 let gate = e.with_moe_cache(max_block, |cache, eng| {
10621 let id = BlockId::new(il, PROJ_GATE, ex as u16);
10622 let slot = cache.dispatch_source(id, m.gate_exps.expert_source(ex), eng)?;
10623 eng.qmatvec_expert_q8(
10624 cache.buf(slot),
10625 0..gl.len,
10626 &zq,
10627 &zd,
10628 m_e,
10629 m.gate_exps.in_f,
10630 m.gate_exps.out_f,
10631 gl.qtype,
10632 gl.row_bytes,
10633 )
10634 })?;
10635 let up = e.with_moe_cache(max_block, |cache, eng| {
10636 let id = BlockId::new(il, PROJ_UP, ex as u16);
10637 let slot = cache.dispatch_source(id, m.up_exps.expert_source(ex), eng)?;
10638 eng.qmatvec_expert_q8(
10639 cache.buf(slot),
10640 0..ul.len,
10641 &zq,
10642 &zd,
10643 m_e,
10644 m.up_exps.in_f,
10645 m.up_exps.out_f,
10646 ul.qtype,
10647 ul.row_bytes,
10648 )
10649 })?;
10650 let mut act = e.uninit(m_e * n_ff_exp)?;
10651 Self::ffn_act_lim(
10652 e,
10653 cfg,
10654 &gate,
10655 &up,
10656 m.gate_exps.macro_scale(ex),
10657 m.up_exps.macro_scale(ex),
10658 lim_exp,
10659 &mut act,
10660 m_e * n_ff_exp,
10661 )?;
10662 let (aq2, ad2) = e.quantize_q8_1(&act, m_e, n_ff_exp)?;
10663 e.with_moe_cache(max_block, |cache, eng| {
10664 let id = BlockId::new(il, PROJ_DOWN, ex as u16);
10665 let slot = cache.dispatch_source(id, m.down_exps.expert_source(ex), eng)?;
10666 eng.qmatvec_expert_q8(
10667 cache.buf(slot),
10668 0..dl.len,
10669 &aq2,
10670 &ad2,
10671 m_e,
10672 m.down_exps.in_f,
10673 m.down_exps.out_f,
10674 dl.qtype,
10675 dl.row_bytes,
10676 )
10677 })?
10678 } else {
10679 let gate = e.with_moe_cache(max_block, |cache, eng| {
10680 let id = BlockId::new(il, PROJ_GATE, ex as u16);
10681 let slot = cache.dispatch_source(id, m.gate_exps.expert_source(ex), eng)?;
10682 eng.qmatvec_view(
10683 cache.buf(slot),
10684 0..gl.len,
10685 &gv,
10686 m_e,
10687 m.gate_exps.in_f,
10688 m.gate_exps.out_f,
10689 gl.qtype,
10690 gl.row_bytes,
10691 )
10692 })?;
10693 let up = e.with_moe_cache(max_block, |cache, eng| {
10694 let id = BlockId::new(il, PROJ_UP, ex as u16);
10695 let slot = cache.dispatch_source(id, m.up_exps.expert_source(ex), eng)?;
10696 eng.qmatvec_view(
10697 cache.buf(slot),
10698 0..ul.len,
10699 &gv,
10700 m_e,
10701 m.up_exps.in_f,
10702 m.up_exps.out_f,
10703 ul.qtype,
10704 ul.row_bytes,
10705 )
10706 })?;
10707 let mut act = e.uninit(m_e * n_ff_exp)?;
10708 Self::ffn_act_lim(
10709 e,
10710 cfg,
10711 &gate,
10712 &up,
10713 m.gate_exps.macro_scale(ex),
10714 m.up_exps.macro_scale(ex),
10715 lim_exp,
10716 &mut act,
10717 m_e * n_ff_exp,
10718 )?;
10719 let actv = act.slice(0..m_e * n_ff_exp);
10720 e.with_moe_cache(max_block, |cache, eng| {
10721 let id = BlockId::new(il, PROJ_DOWN, ex as u16);
10722 let slot = cache.dispatch_source(id, m.down_exps.expert_source(ex), eng)?;
10723 eng.qmatvec_view(
10724 cache.buf(slot),
10725 0..dl.len,
10726 &actv,
10727 m_e,
10728 m.down_exps.in_f,
10729 m.down_exps.out_f,
10730 dl.qtype,
10731 dl.row_bytes,
10732 )
10733 })?
10734 }
10735 } else {
10736 let sg = scratch_g.as_mut().unwrap();
10737 let su = scratch_u.as_mut().unwrap();
10738 let sd = scratch_d.as_mut().unwrap();
10739 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
10740 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
10741 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
10742 if grouped_q8 {
10743 let (zq, zd) = e.quantize_q8_1(&gathered, m_e, n_embd)?;
10744 let gate = e.qmatvec_expert_q8(
10745 sg,
10746 0..gl.len,
10747 &zq,
10748 &zd,
10749 m_e,
10750 m.gate_exps.in_f,
10751 m.gate_exps.out_f,
10752 gl.qtype,
10753 gl.row_bytes,
10754 )?;
10755 let up = e.qmatvec_expert_q8(
10756 su,
10757 0..ul.len,
10758 &zq,
10759 &zd,
10760 m_e,
10761 m.up_exps.in_f,
10762 m.up_exps.out_f,
10763 ul.qtype,
10764 ul.row_bytes,
10765 )?;
10766 let mut act = e.uninit(m_e * n_ff_exp)?;
10767 Self::ffn_act_lim(
10768 e,
10769 cfg,
10770 &gate,
10771 &up,
10772 m.gate_exps.macro_scale(ex),
10773 m.up_exps.macro_scale(ex),
10774 lim_exp,
10775 &mut act,
10776 m_e * n_ff_exp,
10777 )?;
10778 let (aq2, ad2) = e.quantize_q8_1(&act, m_e, n_ff_exp)?;
10779 e.qmatvec_expert_q8(
10780 sd,
10781 0..dl.len,
10782 &aq2,
10783 &ad2,
10784 m_e,
10785 m.down_exps.in_f,
10786 m.down_exps.out_f,
10787 dl.qtype,
10788 dl.row_bytes,
10789 )?
10790 } else {
10791 let gate = e.qmatvec_view(
10792 sg,
10793 0..gl.len,
10794 &gv,
10795 m_e,
10796 m.gate_exps.in_f,
10797 m.gate_exps.out_f,
10798 gl.qtype,
10799 gl.row_bytes,
10800 )?;
10801 let up = e.qmatvec_view(
10802 su,
10803 0..ul.len,
10804 &gv,
10805 m_e,
10806 m.up_exps.in_f,
10807 m.up_exps.out_f,
10808 ul.qtype,
10809 ul.row_bytes,
10810 )?;
10811 let mut act = e.uninit(m_e * n_ff_exp)?;
10812 Self::ffn_act_lim(
10813 e,
10814 cfg,
10815 &gate,
10816 &up,
10817 m.gate_exps.macro_scale(ex),
10818 m.up_exps.macro_scale(ex),
10819 lim_exp,
10820 &mut act,
10821 m_e * n_ff_exp,
10822 )?;
10823 let actv = act.slice(0..m_e * n_ff_exp);
10824 e.qmatvec_view(
10825 sd,
10826 0..dl.len,
10827 &actv,
10828 m_e,
10829 m.down_exps.in_f,
10830 m.down_exps.out_f,
10831 dl.qtype,
10832 dl.row_bytes,
10833 )?
10834 }
10835 };
10836
10837 e.scatter_slot(
10839 &y,
10840 &tok_idx_d,
10841 &slot_idx_d,
10842 &weight_d,
10843 &mut slot_buf,
10844 &mut wbuf,
10845 n_embd,
10846 n_used,
10847 m_e,
10848 )?;
10849 }
10850
10851 let mut moe_out = e.zeros(t * n_embd)?;
10853 e.reduce_slots(&slot_buf, &wbuf, &mut moe_out, n_embd, n_used, t)?;
10854
10855 if std::env::var("MEMRA_MOE_STATS").is_ok() && !m_dist.is_empty() {
10857 m_dist.sort_unstable();
10858 let active = m_dist.len();
10859 let mean = m_dist.iter().sum::<usize>() as f64 / active as f64;
10860 let median = m_dist[active / 2];
10861 let max_m = *m_dist.last().unwrap();
10862 let min_m = m_dist[0];
10863 let above16 = m_dist.iter().filter(|&&x| x >= 16).count();
10864 println!(
10865 "moe-grouped il={il} t={t} active={active}/{n_expert} \
10866 m_e: min={min_m} median={median} mean={mean:.1} max={max_m} \
10867 above_gemm_threshold(>=16)={above16}/{active}"
10868 );
10869 }
10870
10871 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut moe_out)?;
10872 Ok(moe_out)
10873 }
10874
10875 pub(crate) fn moe_ffn_lockstep(
10882 &self,
10883 e: &Engine,
10884 m: &MoeWeights,
10885 zbatch: &CudaSlice<f32>,
10886 mrows: usize,
10887 il: u16,
10888 max_block: usize,
10889 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
10890 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
10891 let cfg = &self.cfg;
10892 let moe = cfg.moe.as_ref().unwrap();
10893 let n_embd = cfg.n_embd as usize;
10894 let n_expert = moe.expert_count as usize;
10895 let n_used = moe.expert_used_count as usize;
10896 let n_ff_exp = moe.expert_ff_length as usize;
10897 let lim_exp = cfg.clamp_exp_at(il as u32);
10899 let lim_shexp = cfg.clamp_shexp_at(il as u32);
10900
10901 let logits = e.matmul(&m.gate_inp, zbatch, mrows)?;
10902 if let Some(sig) = cfg.sigmoid_router() {
10903 Self::trace_sigmoid_router_logits(e, il, mrows, n_expert, n_used, &logits, m, sig)?;
10904 }
10905 let (sel_all, w_all) = if let Some(sig) = cfg.sigmoid_router() {
10906 Self::moe_route_sigmoid_cfg(e, &logits, mrows, n_expert, n_used, m, sig)?
10907 } else {
10908 Self::moe_route_cfg(
10909 e,
10910 &logits,
10911 mrows,
10912 n_expert,
10913 n_used,
10914 m.active_experts.as_deref(),
10915 )?
10916 };
10917 Self::trace_moe_routes(il, mrows, &sel_all, &w_all)?;
10918
10919 let resident_expert: Vec<bool> = e.with_moe_cache(max_block, |c, _| {
10921 Ok((0..n_expert)
10922 .map(|ex| {
10923 [PROJ_GATE, PROJ_UP, PROJ_DOWN]
10924 .into_iter()
10925 .all(|p| c.resident(BlockId::new(il, p, ex as u16)).is_some())
10926 })
10927 .collect())
10928 })?;
10929
10930 struct Group {
10931 rows: Vec<i32>,
10932 slots: Vec<i32>,
10933 weights: Vec<f32>,
10934 }
10935 let mut groups: std::collections::HashMap<usize, Group> = Default::default();
10936 let mut cpu_rows: Vec<Vec<(usize, f32)>> = vec![Vec::new(); mrows];
10937 let mut cpu_by_expert: std::collections::HashMap<usize, Vec<(usize, f32)>> =
10938 Default::default();
10939 for row in 0..mrows {
10940 for j in 0..n_used {
10941 let ex = sel_all[row * n_used + j] as usize;
10942 let w = w_all[row * n_used + j];
10943 if resident_expert[ex] {
10944 let group = groups.entry(ex).or_insert_with(|| Group {
10945 rows: Vec::new(),
10946 slots: Vec::new(),
10947 weights: Vec::new(),
10948 });
10949 group.rows.push(row as i32);
10950 group.slots.push(j as i32);
10951 group.weights.push(w);
10952 } else {
10953 crate::cpu_experts::record_incomplete_gpu_residency(0);
10954 cpu_rows[row].push((ex, w));
10955 cpu_by_expert.entry(ex).or_default().push((row, w));
10956 }
10957 }
10958 }
10959
10960 let host_rows = e.dtoh(zbatch)?;
10966 let rows_ok = crate::cpu_experts::rows_supported();
10967 enum CpuPart {
10968 Single { row: usize },
10969 Rows { rows: Vec<usize> },
10970 }
10971 let mut tickets: Vec<(CpuPart, crate::cpu_experts::CpuExpertTicket)> = Vec::new();
10972 let mut rows_served: std::collections::HashSet<(usize, usize)> = Default::default();
10973 if rows_ok {
10974 let mut shared: Vec<(usize, Vec<(usize, f32)>)> = cpu_by_expert
10975 .into_iter()
10976 .filter(|(_, rows)| rows.len() >= 2)
10977 .collect();
10978 shared.sort_by_key(|(ex, _)| *ex);
10979 for (ex, mut row_weights) in shared {
10980 row_weights.sort_by_key(|(row, _)| *row);
10981 let inputs: Vec<(&[f32], f32)> = row_weights
10982 .iter()
10983 .map(|&(row, w)| (&host_rows[row * n_embd..(row + 1) * n_embd], w))
10984 .collect();
10985 let job = crate::cpu_experts::prepare_rows_job(m, ex, &inputs)
10986 .map_err(std::io::Error::other)?;
10987 for &(row, _) in &row_weights {
10988 rows_served.insert((row, ex));
10989 }
10990 tickets.push((
10991 CpuPart::Rows {
10992 rows: row_weights.iter().map(|&(row, _)| row).collect(),
10993 },
10994 crate::cpu_experts::submit_rows(job).map_err(std::io::Error::other)?,
10995 ));
10996 }
10997 }
10998 for (row, selected) in cpu_rows.iter().enumerate() {
10999 let leftover: Vec<(usize, f32)> = selected
11000 .iter()
11001 .copied()
11002 .filter(|&(ex, _)| !rows_served.contains(&(row, ex)))
11003 .collect();
11004 if leftover.is_empty() {
11005 continue;
11006 }
11007 let host_row = &host_rows[row * n_embd..(row + 1) * n_embd];
11008 let job = crate::cpu_experts::prepare_job(m, il, &leftover, host_row)
11009 .map_err(std::io::Error::other)?;
11010 tickets.push((
11011 CpuPart::Single { row },
11012 crate::cpu_experts::submit(job).map_err(std::io::Error::other)?,
11013 ));
11014 }
11015
11016 let mut slot_buf = e.zeros(mrows * n_used * n_embd)?;
11017 let mut wbuf = e.zeros(mrows * n_used)?;
11018 let mut order: Vec<usize> = groups.keys().copied().collect();
11019 order.sort_by(|&a, &b| {
11020 groups[&b]
11021 .rows
11022 .len()
11023 .cmp(&groups[&a].rows.len())
11024 .then(a.cmp(&b))
11025 });
11026 for &ex in &order {
11027 let group = &groups[&ex];
11028 let m_e = group.rows.len();
11029 let gl = m.gate_exps.expert_layout(ex);
11030 let ul = m.up_exps.expert_layout(ex);
11031 let dl = m.down_exps.expert_layout(ex);
11032 let row_idx_d = e.htod_i32(&group.rows)?;
11033 let slot_idx_d = e.htod_i32(&group.slots)?;
11034 let dmac = m.down_exps.macro_scale(ex);
11035 let weight_d = if dmac == 1.0 {
11036 e.htod(&group.weights)?
11037 } else {
11038 let scaled: Vec<f32> = group.weights.iter().map(|&w| w * dmac).collect();
11039 e.htod(&scaled)?
11040 };
11041 let mut gathered = e.zeros(m_e * n_embd)?;
11042 e.gather_rows(zbatch, &row_idx_d, &mut gathered, n_embd, m_e)?;
11043 let gv = gathered.slice(0..m_e * n_embd);
11044 let gate = e.with_moe_cache(max_block, |c, eng| {
11045 let slot = c
11046 .resident(BlockId::new(il, PROJ_GATE, ex as u16))
11047 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
11048 eng.qmatvec_view(
11049 c.buf(crate::moe_cache::DispatchSlot::Resident(slot)),
11050 0..gl.len,
11051 &gv,
11052 m_e,
11053 m.gate_exps.in_f,
11054 m.gate_exps.out_f,
11055 gl.qtype,
11056 gl.row_bytes,
11057 )
11058 })?;
11059 let up = e.with_moe_cache(max_block, |c, eng| {
11060 let slot = c
11061 .resident(BlockId::new(il, PROJ_UP, ex as u16))
11062 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
11063 eng.qmatvec_view(
11064 c.buf(crate::moe_cache::DispatchSlot::Resident(slot)),
11065 0..ul.len,
11066 &gv,
11067 m_e,
11068 m.up_exps.in_f,
11069 m.up_exps.out_f,
11070 ul.qtype,
11071 ul.row_bytes,
11072 )
11073 })?;
11074 let mut act = e.zeros(m_e * n_ff_exp)?;
11075 Self::ffn_act_lim(
11076 e,
11077 cfg,
11078 &gate,
11079 &up,
11080 m.gate_exps.macro_scale(ex),
11081 m.up_exps.macro_scale(ex),
11082 lim_exp,
11083 &mut act,
11084 m_e * n_ff_exp,
11085 )?;
11086 let actv = act.slice(0..m_e * n_ff_exp);
11087 let y = e.with_moe_cache(max_block, |c, eng| {
11088 let slot = c
11089 .resident(BlockId::new(il, PROJ_DOWN, ex as u16))
11090 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
11091 eng.qmatvec_view(
11092 c.buf(crate::moe_cache::DispatchSlot::Resident(slot)),
11093 0..dl.len,
11094 &actv,
11095 m_e,
11096 m.down_exps.in_f,
11097 m.down_exps.out_f,
11098 dl.qtype,
11099 dl.row_bytes,
11100 )
11101 })?;
11102 e.scatter_slot(
11103 &y,
11104 &row_idx_d,
11105 &slot_idx_d,
11106 &weight_d,
11107 &mut slot_buf,
11108 &mut wbuf,
11109 n_embd,
11110 n_used,
11111 m_e,
11112 )?;
11113 }
11114 let mut moe_out = e.zeros(mrows * n_embd)?;
11115 e.reduce_slots(&slot_buf, &wbuf, &mut moe_out, n_embd, n_used, mrows)?;
11116
11117 let mut row_sums: Vec<Option<Vec<f32>>> = vec![None; mrows];
11119 for (part, ticket) in tickets {
11120 let cpu_output = ticket.wait().map_err(std::io::Error::other)?;
11121 let mut add_row = |row: usize, chunk: &[f32]| {
11122 let sum = row_sums[row].get_or_insert_with(|| vec![0.0f32; n_embd]);
11123 for (accumulator, value) in sum.iter_mut().zip(chunk) {
11124 *accumulator += value;
11125 }
11126 };
11127 match part {
11128 CpuPart::Single { row } => add_row(row, &cpu_output),
11129 CpuPart::Rows { rows } => {
11130 for (slot, row) in rows.into_iter().enumerate() {
11131 add_row(row, &cpu_output[slot * n_embd..(slot + 1) * n_embd]);
11132 }
11133 }
11134 }
11135 }
11136 for (row, sum) in row_sums.into_iter().enumerate() {
11137 let Some(sum) = sum else { continue };
11138 let cpu_output = e.htod(&sum)?;
11139 let mut dst = moe_out.slice_mut(row * n_embd..(row + 1) * n_embd);
11140 e.axpy_into(&cpu_output, 1.0, &mut dst, n_embd)?;
11141 }
11142
11143 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
11144 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
11145 {
11146 let n_ff_sh = gate_shexp.out_features();
11147 let sg_gate = e.matmul(gate_shexp, zbatch, mrows)?;
11148 let sg_up = e.matmul(up_shexp, zbatch, mrows)?;
11149 let mut sa = e.zeros(mrows * n_ff_sh)?;
11150 Self::ffn_act_lim(
11151 e,
11152 cfg,
11153 &sg_gate,
11154 &sg_up,
11155 1.0,
11156 1.0,
11157 lim_shexp,
11158 &mut sa,
11159 mrows * n_ff_sh,
11160 )?;
11161 let sh = e.matmul(down_shexp, &sa, mrows)?;
11162 let g = match &m.gate_inp_shexp {
11165 Some(gate_inp_shexp) => {
11166 e.sigmoid_dot_rows(zbatch, gate_inp_shexp.float_data(), n_embd, mrows)?
11167 }
11168 None => e.htod(&vec![1.0f32; mrows])?,
11169 };
11170 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, mrows)?;
11171 }
11172
11173 Ok(moe_out)
11174 }
11175}
11176
11177impl HybridModel {
11183 pub(crate) fn gemma4_rope_dims(&self, il: usize) -> usize {
11197 let g = self
11198 .cfg
11199 .gemma4
11200 .as_ref()
11201 .expect("gemma4_rope_dims on a non-gemma4 config");
11202 if g.swa_pattern[il] {
11203 g.rope_dims_swa as usize
11204 } else {
11205 g.rope_dims_global as usize
11206 }
11207 }
11208
11209 pub(crate) fn gemma4_geom(&self, il: usize) -> (usize, usize, usize, f32, f32, bool) {
11210 let g = self.cfg.gemma4.as_ref().unwrap();
11211 let swa = g.swa_pattern[il];
11212 let hd = if swa {
11213 g.key_length_swa
11214 } else {
11215 g.key_length_global
11216 } as usize;
11217 (
11221 hd,
11222 g.head_count_kv[il] as usize,
11223 self.cfg.n_head as usize,
11224 if swa {
11225 g.rope_base_swa
11226 } else {
11227 g.rope_base_global
11228 },
11229 1.0,
11230 swa,
11231 )
11232 }
11233
11234 pub(crate) fn gemma4_suppress(
11238 &self,
11239 e: &Engine,
11240 ld: &mut CudaSlice<f32>,
11241 t: usize,
11242 ) -> Result<(), Box<dyn std::error::Error>> {
11243 if let Some((ids, n)) = self.gemma4_aux.as_ref().and_then(|a| a.suppress_d.as_ref()) {
11244 #[cfg(debug_assertions)]
11249 crate::debug_assert_tensor_stream_device(
11250 ids,
11251 &e.stream(),
11252 "gemma4_suppress.suppress_d",
11253 );
11254 e.mask_ids_rows(ld, ids, *n, self.output.out_features(), t)?;
11255 }
11256 Ok(())
11257 }
11258
11259 #[allow(clippy::too_many_arguments)]
11264 fn gemma_fa_one_program() -> bool {
11273 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
11274 *ON.get_or_init(|| std::env::var("MEMRA_GEMMA_FA_ONE_PROGRAM").as_deref() == Ok("1"))
11275 }
11276
11277 fn gemma4_attn_prime(
11278 &self,
11279 e: &Engine,
11280 fa: &crate::hybrid::FullAttnLayer,
11281 il: usize,
11282 h: &CudaSlice<f32>,
11283 pos_d: &CudaSlice<i32>,
11284 t: usize,
11285 cache: Option<&mut Cache>,
11286 island: Option<&CudaSlice<i32>>,
11287 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
11288 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
11289 let eps = self.cfg.rms_eps;
11290 let aux = self.gemma4_aux.as_ref().unwrap();
11291 let ones = aux.ones(e);
11292 #[cfg(debug_assertions)]
11293 crate::debug_assert_tensor_stream_device(ones, &e.stream(), "gemma4_attn_prime.ones");
11294
11295 e.mmq_act_begin();
11298 let q0 = e.matmul(&fa.wq, h, t)?; if il == 0 && std::env::var("MEMRA_G4_PRIME_TRACE").as_deref() == Ok("1") {
11300 let v = e.dtoh(&q0)?;
11301 let nan = v.iter().filter(|x| x.is_nan()).count();
11302 let amax = v.iter().fold(0f32, |a, x| a.max(x.abs()));
11303 eprintln!(
11304 "[g4-prime-trace] L0 q0: nan={nan}/{} amax={amax:.3}",
11305 v.len()
11306 );
11307 }
11308 let k0 = e.matmul(&fa.wk, h, t)?; let v0 = if swa {
11312 e.matmul(&fa.wv, h, t)?
11313 } else {
11314 e.clone_dtod(&k0)?
11315 };
11316 if il == 0 && std::env::var("MEMRA_G4_PRIME_TRACE").as_deref() == Ok("1") {
11317 for (tag, buf) in [("k0", &k0), ("v0", &v0)] {
11318 let v = e.dtoh(buf)?;
11319 let nan = v.iter().filter(|x| x.is_nan()).count();
11320 let amax = v.iter().fold(0f32, |a, x| a.max(x.abs()));
11321 eprintln!(
11322 "[g4-prime-trace] L0 {tag}: nan={nan}/{} amax={amax:.3}",
11323 v.len()
11324 );
11325 }
11326 }
11327
11328 let mut q = e.uninit(t * nh * hd)?;
11329 let mut k = e.uninit(t * nkv * hd)?;
11330 let mut v = e.uninit(t * nkv * hd)?;
11332 static EMIT: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
11336 let emit = island.is_none()
11339 && t >= 16
11340 && crate::Engine::qkvnorm_w_on_prefill(nh * t + 2 * nkv * t, hd)
11341 && *EMIT.get_or_init(|| {
11342 std::env::var("MEMRA_FA_EMIT")
11343 .map(|s| s != "0")
11344 .unwrap_or(true)
11345 });
11346 let mut qb = e.alloc_uninit::<u8>(if emit { t * nh * hd * 2 } else { 1 })?;
11347 let mut kb = e.alloc_uninit::<u8>(if emit { t * nkv * hd * 2 } else { 1 })?;
11348 let mut vb = e.alloc_uninit::<u8>(if emit { t * nkv * hd * 2 } else { 1 })?;
11349 let v_f16 = emit
11352 && crate::fa_f16pv_on()
11353 && match hd {
11354 512 => true,
11355 256 => swa && crate::faw_hp_on() && nh % 2 == 0 && (nh / nkv) % 2 == 0,
11356 _ => false,
11357 };
11358 if emit {
11359 e.rms_norm_qkv_w4b(
11360 &q0,
11361 &k0,
11362 &v0,
11363 fa.q_norm.float_data(),
11364 fa.k_norm.float_data(),
11365 ones,
11366 &mut q,
11367 &mut k,
11368 &mut v,
11369 &mut vb,
11370 hd,
11371 nh * t,
11372 nkv * t,
11373 eps,
11374 v_f16,
11375 )?;
11376 } else {
11377 e.rms_norm_qkv(
11378 &q0,
11379 &k0,
11380 &v0,
11381 fa.q_norm.float_data(),
11382 fa.k_norm.float_data(),
11383 ones,
11384 &mut q,
11385 &mut k,
11386 &mut v,
11387 hd,
11388 nh * t,
11389 nkv * t,
11390 eps,
11391 )?;
11392 }
11393
11394 let ff = if swa {
11395 None
11396 } else {
11397 Some(
11398 aux.rope_freqs(e)
11399 .expect("gemma4 global rope needs rope_freqs.weight"),
11400 )
11401 };
11402 #[cfg(debug_assertions)]
11403 if let Some(ff) = ff {
11404 crate::debug_assert_tensor_stream_device(
11405 ff,
11406 &e.stream(),
11407 "gemma4_attn_prime.rope_freqs",
11408 );
11409 }
11410 if emit {
11411 e.rope_neox2_bf16e(
11412 &mut q, &mut k, &mut qb, &mut kb, pos_d, hd, hd, nh, nkv, t, base, 1.0, ff,
11413 )?;
11414 } else {
11415 e.rope_neox2(&mut q, &mut k, pos_d, hd, hd, nh, nkv, t, base, 1.0, ff)?;
11416 }
11417
11418 if let Some(cache) = cache {
11419 let kvl = cache.kv[il].as_mut().unwrap();
11420 assert_eq!(kvl.len, 0, "gemma4 prime is fresh-prompt only (v0)");
11421 e.append_kv_quantized_rows(
11422 &k,
11423 &v,
11424 &mut kvl.k,
11425 &mut kvl.v,
11426 kvl.len,
11427 t,
11428 kvl.kv_dim_k,
11429 kvl.kv_dim_v,
11430 kvl.k_tok_bytes,
11431 kvl.v_tok_bytes,
11432 (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on()),
11433 )?;
11434 kvl.len += t;
11435 }
11436 let mut attn = e.zeros(t * nh * hd)?;
11437 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
11441 if let Some(span) = island {
11442 let w = if swa && t > win { win } else { 0 };
11447 e.sdpa_naive_island(&q, &k, &v, &mut attn, span, hd, nh, nkv, t, t, scale, w)?;
11448 } else if swa && (t > win || Self::gemma_fa_one_program()) {
11449 if hd == 256 && std::env::var("MEMRA_NOFA").is_err() {
11450 if emit {
11451 e.fa_prefill_w_pre(
11452 &qb, &kb, &vb, &mut attn, hd, nh, nkv, t, t, scale, true, win, v_f16,
11453 )?;
11454 } else {
11455 e.fa_prefill_w(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
11456 }
11457 } else {
11458 e.sdpa_naive_w(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
11459 }
11460 } else if hd == 256 && std::env::var("MEMRA_NOFA").is_err() {
11461 e.fa_prefill(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
11462 } else if hd == 512 && std::env::var("MEMRA_NOFA").is_err() {
11463 if emit {
11464 e.fa_prefill_hd512_pre(
11465 &qb, &kb, &vb, &mut attn, hd, nh, nkv, t, t, scale, true, v_f16,
11466 )?;
11467 } else {
11468 e.fa_prefill_hd512(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
11469 }
11470 } else {
11471 e.sdpa_naive(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
11472 }
11473 Ok(e.matmul(&fa.wo, &attn, t)?)
11474 }
11475
11476 fn gemma4_attn(
11478 &self,
11479 e: &Engine,
11480 fa: &crate::hybrid::FullAttnLayer,
11481 il: usize,
11482 h: &CudaSlice<f32>,
11483 pos_d: &CudaSlice<i32>,
11484 t: usize,
11485 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
11486 self.gemma4_attn_prime(e, fa, il, h, pos_d, t, None, None)
11487 }
11488
11489 fn gemma4_moe_q8(
11494 &self,
11495 e: &Engine,
11496 m: &crate::hybrid::MoeWeights,
11497 bits: &crate::hybrid::Gemma4MoeBits,
11498 mq: &(CudaSlice<i8>, CudaSlice<f32>),
11499 router_in: &CudaSlice<f32>,
11500 t: usize,
11501 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
11502 let cfg = &self.cfg;
11503 let moe = cfg.moe.as_ref().unwrap();
11504 let n_embd = cfg.n_embd as usize;
11505 let n_expert = moe.expert_count as usize;
11506 let n_used = moe.expert_used_count as usize;
11507 let n_ff_exp = moe.expert_ff_length as usize;
11508 let logits = if crate::router_kernel_on() {
11512 e.router_gemv(m.gate_inp.float_data(), router_in, n_embd, n_expert, t)?
11513 } else {
11514 e.matmul(&m.gate_inp, router_in, t)?
11515 };
11516 let dev = m.dev_exps.as_ref().unwrap();
11517 let (sel_d, w_d) =
11518 e.moe_router_topk_scaled(&logits, t, n_expert, n_used, &bits.per_expert_scale_d)?;
11519 let (zq, zd) = mq;
11520 if t == 1 {
11521 let selv = sel_d.slice(0..n_used);
11522 let wv = w_d.slice(0..n_used);
11523 let act = e.moe_gate_up_gelu8_dev_q8(
11524 &dev.ptr_row,
11525 &selv,
11526 zq,
11527 zd,
11528 n_embd,
11529 n_ff_exp,
11530 n_used,
11531 n_expert,
11532 m.gate_exps.qtype,
11533 m.up_exps.qtype,
11534 m.gate_exps.row_bytes,
11535 m.up_exps.row_bytes,
11536 )?;
11537 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
11538 let mut moe_out = e.uninit(n_embd)?;
11539 e.moe_down8_fma_dev_q8(
11540 &dev.ptr_row,
11541 &selv,
11542 &wv,
11543 &aq2,
11544 &ad2,
11545 &mut moe_out.slice_mut(0..n_embd),
11546 n_ff_exp,
11547 n_embd,
11548 n_used,
11549 n_expert,
11550 m.down_exps.qtype,
11551 m.down_exps.row_bytes,
11552 )?;
11553 return Ok(moe_out);
11554 }
11555 let csr = t <= 10 && std::env::var("MEMRA_GEMMA_CSR").as_deref() != Ok("0");
11556 let act = if csr {
11557 e.moe_gate_up_gelu8_dev_q8_csr(
11558 &dev.ptr_row,
11559 &sel_d,
11560 zq,
11561 zd,
11562 t * n_used,
11563 n_embd,
11564 n_ff_exp,
11565 n_used,
11566 n_expert,
11567 m.gate_exps.qtype,
11568 m.up_exps.qtype,
11569 m.gate_exps.row_bytes,
11570 m.up_exps.row_bytes,
11571 )?
11572 } else {
11573 e.moe_gate_up_gelu8_dev_q8_rows(
11574 &dev.ptr_row,
11575 &sel_d,
11576 zq,
11577 zd,
11578 t,
11579 n_embd,
11580 n_ff_exp,
11581 n_used,
11582 n_expert,
11583 m.gate_exps.qtype,
11584 m.up_exps.qtype,
11585 m.gate_exps.row_bytes,
11586 m.up_exps.row_bytes,
11587 )?
11588 };
11589 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
11590 let mut moe_out = e.uninit(t * n_embd)?;
11591 e.moe_down8_fma_dev_q8_rows_g(
11594 &dev.ptr_row,
11595 &sel_d,
11596 &w_d,
11597 &aq2,
11598 &ad2,
11599 &mut moe_out,
11600 t,
11601 n_ff_exp,
11602 n_embd,
11603 n_used,
11604 n_expert,
11605 m.down_exps.qtype,
11606 m.down_exps.row_bytes,
11607 )?;
11608 Ok(moe_out)
11609 }
11610
11611 fn gemma4_moe(
11615 &self,
11616 e: &Engine,
11617 m: &crate::hybrid::MoeWeights,
11618 bits: &crate::hybrid::Gemma4MoeBits,
11619 moe_in: &CudaSlice<f32>,
11620 router_in: &CudaSlice<f32>,
11621 t: usize,
11622 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
11623 let cfg = &self.cfg;
11624 let moe = cfg.moe.as_ref().unwrap();
11625 let n_embd = cfg.n_embd as usize;
11626 let n_expert = moe.expert_count as usize;
11627 let n_used = moe.expert_used_count as usize;
11628 let n_ff_exp = moe.expert_ff_length as usize;
11629
11630 let logits = if t < PRIME_MIN_T && crate::router_kernel_on() {
11634 e.router_gemv(m.gate_inp.float_data(), router_in, n_embd, n_expert, t)?
11635 } else {
11636 e.matmul(&m.gate_inp, router_in, t)?
11637 };
11638
11639 if t < PRIME_MIN_T
11644 && m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
11645 && expert_dp4a_supported(m.gate_exps.qtype)
11646 && expert_dp4a_supported(m.up_exps.qtype)
11647 && expert_dp4a_supported(m.down_exps.qtype)
11648 && std::env::var("MEMRA_GEMMA_MOE_FAST").as_deref() != Ok("0")
11649 {
11650 let dev = m.dev_exps.as_ref().unwrap();
11651 let (sel_d, w_d) =
11652 e.moe_router_topk_scaled(&logits, t, n_expert, n_used, &bits.per_expert_scale_d)?;
11653 if t == 1 {
11654 let (zq, zd) = e.quantize_q8_1(moe_in, 1, n_embd)?;
11655 let selv = sel_d.slice(0..n_used);
11656 let wv = w_d.slice(0..n_used);
11657 let act = e.moe_gate_up_gelu8_dev_q8(
11658 &dev.ptr_row,
11659 &selv,
11660 &zq,
11661 &zd,
11662 n_embd,
11663 n_ff_exp,
11664 n_used,
11665 n_expert,
11666 m.gate_exps.qtype,
11667 m.up_exps.qtype,
11668 m.gate_exps.row_bytes,
11669 m.up_exps.row_bytes,
11670 )?;
11671 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
11672 let mut moe_out = e.uninit(n_embd)?;
11673 e.moe_down8_fma_dev_q8(
11674 &dev.ptr_row,
11675 &selv,
11676 &wv,
11677 &aq2,
11678 &ad2,
11679 &mut moe_out.slice_mut(0..n_embd),
11680 n_ff_exp,
11681 n_embd,
11682 n_used,
11683 n_expert,
11684 m.down_exps.qtype,
11685 m.down_exps.row_bytes,
11686 )?;
11687 return Ok(moe_out);
11688 }
11689 let (zq, zd) = e.quantize_q8_1(moe_in, t, n_embd)?;
11694 let csr = t <= 10 && std::env::var("MEMRA_GEMMA_CSR").as_deref() != Ok("0");
11695 let act = if csr {
11696 e.moe_gate_up_gelu8_dev_q8_csr(
11697 &dev.ptr_row,
11698 &sel_d,
11699 &zq,
11700 &zd,
11701 t * n_used,
11702 n_embd,
11703 n_ff_exp,
11704 n_used,
11705 n_expert,
11706 m.gate_exps.qtype,
11707 m.up_exps.qtype,
11708 m.gate_exps.row_bytes,
11709 m.up_exps.row_bytes,
11710 )?
11711 } else {
11712 e.moe_gate_up_gelu8_dev_q8_rows(
11713 &dev.ptr_row,
11714 &sel_d,
11715 &zq,
11716 &zd,
11717 t,
11718 n_embd,
11719 n_ff_exp,
11720 n_used,
11721 n_expert,
11722 m.gate_exps.qtype,
11723 m.up_exps.qtype,
11724 m.gate_exps.row_bytes,
11725 m.up_exps.row_bytes,
11726 )?
11727 };
11728 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
11729 let mut moe_out = e.uninit(t * n_embd)?;
11730 e.moe_down8_fma_dev_q8_rows_g(
11731 &dev.ptr_row,
11732 &sel_d,
11733 &w_d,
11734 &aq2,
11735 &ad2,
11736 &mut moe_out,
11737 t,
11738 n_ff_exp,
11739 n_embd,
11740 n_used,
11741 n_expert,
11742 m.down_exps.qtype,
11743 m.down_exps.row_bytes,
11744 )?;
11745 return Ok(moe_out);
11746 }
11747
11748 let (sel_all, mut w_all) = Self::moe_route(e, &logits, t, n_expert, n_used)?;
11749 for (i, &sx) in sel_all.iter().enumerate() {
11750 w_all[i] *= bits.per_expert_scale[sx as usize];
11751 }
11752
11753 if t >= PRIME_MIN_T
11757 && m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
11758 && expert_dp4a_supported(m.gate_exps.qtype)
11759 && expert_dp4a_supported(m.up_exps.qtype)
11760 && expert_dp4a_supported(m.down_exps.qtype)
11761 && std::env::var("MEMRA_GEMMA_MOE_PAIRS").as_deref() != Ok("0")
11762 {
11763 let dev = m.dev_exps.as_ref().unwrap();
11764 let n_pairs = t * n_used;
11765 let pair_ex: Vec<i32> = sel_all.iter().map(|&x| x as i32).collect();
11766 let pair_tok: Vec<i32> = (0..n_pairs).map(|p| (p / n_used) as i32).collect();
11767 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
11768 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
11769 let pt = e.htod_i32(&pair_tok)?;
11770 let pw = e.htod(&w_all)?;
11771 let toff = e.htod_i32(&tok_off)?;
11772 let tids = e.htod_i32(&tok_ids)?;
11773 let mut by_ex: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
11774 for p in 0..n_pairs {
11775 by_ex[pair_ex[p] as usize].push(p as i32);
11776 }
11777 let mut ex_ids: Vec<i32> = Vec::new();
11778 let mut ex_off: Vec<i32> = vec![0];
11779 let mut ex_pairs: Vec<i32> = Vec::with_capacity(n_pairs);
11780 for (ex, list) in by_ex.iter().enumerate() {
11781 if list.is_empty() {
11782 continue;
11783 }
11784 ex_ids.push(ex as i32);
11785 ex_pairs.extend_from_slice(list);
11786 ex_off.push(ex_pairs.len() as i32);
11787 }
11788 let n_active = ex_ids.len();
11789 let exi = e.htod_i32(&ex_ids)?;
11790 let exo = e.htod_i32(&ex_off)?;
11791 let exp_d = e.htod_i32(&ex_pairs)?;
11792 if crate::moe_f16g_gemma_on()
11800 && f16g_proj_ok(m.gate_exps.qtype, n_embd)
11801 && f16g_proj_ok(m.up_exps.qtype, n_embd)
11802 && f16g_proj_ok(m.down_exps.qtype, n_ff_exp)
11803 {
11804 let csr_tok: Vec<i32> = ex_pairs.iter().map(|&p| p / n_used as i32).collect();
11805 let csr_tok_d = e.htod_i32(&csr_tok)?;
11806 let (z_f16, z_s) = e.moe_f16g_act(moe_in, Some(&csr_tok_d), n_embd, n_pairs)?;
11807 let g_csr = e.moe_f16_grouped(
11808 &dev.ptr_row,
11809 0,
11810 n_expert,
11811 &exi,
11812 &ex_off,
11813 &exo,
11814 &z_f16,
11815 &z_s,
11816 n_embd,
11817 n_ff_exp,
11818 n_active,
11819 n_pairs,
11820 m.gate_exps.qtype,
11821 m.gate_exps.row_bytes,
11822 )?;
11823 let u_csr = e.moe_f16_grouped(
11824 &dev.ptr_row,
11825 1,
11826 n_expert,
11827 &exi,
11828 &ex_off,
11829 &exo,
11830 &z_f16,
11831 &z_s,
11832 n_embd,
11833 n_ff_exp,
11834 n_active,
11835 n_pairs,
11836 m.up_exps.qtype,
11837 m.up_exps.row_bytes,
11838 )?;
11839 let act_csr = e.moe_pairs_gelu_mul(&g_csr, &u_csr, n_pairs * n_ff_exp)?;
11840 let (a_f16, a_s) = e.moe_f16g_act(&act_csr, None, n_ff_exp, n_pairs)?;
11841 let d_csr = e.moe_f16_grouped(
11842 &dev.ptr_row,
11843 2,
11844 n_expert,
11845 &exi,
11846 &ex_off,
11847 &exo,
11848 &a_f16,
11849 &a_s,
11850 n_ff_exp,
11851 n_embd,
11852 n_active,
11853 n_pairs,
11854 m.down_exps.qtype,
11855 m.down_exps.row_bytes,
11856 )?;
11857 let y_down = e.rows_permute(&d_csr, &exp_d, n_pairs, n_embd)?;
11858 let mut moe_out = e.uninit(t * n_embd)?;
11859 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
11860 if std::env::var("MEMRA_F16G_DEBUG").is_ok() {
11861 let scan = |v: &[f32]| v.iter().filter(|x| !x.is_finite()).count();
11862 let (yd, mo) = (e.dtoh(&y_down)?, e.dtoh(&moe_out)?);
11863 eprintln!(
11864 "[f16g-debug] post-permute bad={} post-scatter bad={}",
11865 scan(&yd),
11866 scan(&mo)
11867 );
11868 }
11869 return Ok(moe_out);
11870 }
11871 let mma =
11874 n_embd % 256 == 0 && std::env::var("MEMRA_GEMMA_MOE_MMA").as_deref() != Ok("0");
11875 let (gate, up) = if mma {
11876 let z_scr = e.mmq_iq_quantize_act(moe_in, n_embd, t)?;
11877 (
11878 e.mmq_iq_experts(
11879 &dev.ptr_row,
11880 0,
11881 n_expert,
11882 &exi,
11883 &exo,
11884 &exp_d,
11885 &pt,
11886 &z_scr,
11887 n_embd,
11888 n_ff_exp,
11889 n_active,
11890 n_pairs,
11891 t,
11892 m.gate_exps.qtype,
11893 m.gate_exps.row_bytes,
11894 )?,
11895 e.mmq_iq_experts(
11896 &dev.ptr_row,
11897 1,
11898 n_expert,
11899 &exi,
11900 &exo,
11901 &exp_d,
11902 &pt,
11903 &z_scr,
11904 n_embd,
11905 n_ff_exp,
11906 n_active,
11907 n_pairs,
11908 t,
11909 m.up_exps.qtype,
11910 m.up_exps.row_bytes,
11911 )?,
11912 )
11913 } else {
11914 let (zq, zd) = e.quantize_q8_1(moe_in, t, n_embd)?;
11915 (
11916 e.moe_pairs_matvec_q8_dec(
11917 &dev.ptr_row,
11918 0,
11919 &exi,
11920 &exo,
11921 &exp_d,
11922 &pt,
11923 &zq,
11924 &zd,
11925 n_embd,
11926 n_ff_exp,
11927 n_expert,
11928 n_active,
11929 n_pairs,
11930 m.gate_exps.qtype,
11931 m.gate_exps.row_bytes,
11932 )?,
11933 e.moe_pairs_matvec_q8_dec(
11934 &dev.ptr_row,
11935 1,
11936 &exi,
11937 &exo,
11938 &exp_d,
11939 &pt,
11940 &zq,
11941 &zd,
11942 n_embd,
11943 n_ff_exp,
11944 n_expert,
11945 n_active,
11946 n_pairs,
11947 m.up_exps.qtype,
11948 m.up_exps.row_bytes,
11949 )?,
11950 )
11951 };
11952 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
11953 let pself = e.htod_i32(&pair_self)?;
11954 let y_down = if mma {
11966 let in_pad = n_ff_exp.div_ceil(256) * 256;
11967 let a_scr = if crate::moe_fuse_actq_on() {
11968 e.mmq_iq_fused_act_quant(&gate, &up, n_ff_exp, n_pairs, 1)?
11969 } else {
11970 let act = e.moe_pairs_gelu_mul(&gate, &up, n_pairs * n_ff_exp)?;
11971 e.mmq_iq_quantize_act(&act, n_ff_exp, n_pairs)?
11972 };
11973 e.mmq_iq_experts(
11974 &dev.ptr_row,
11975 2,
11976 n_expert,
11977 &exi,
11978 &exo,
11979 &exp_d,
11980 &pself,
11981 &a_scr,
11982 in_pad,
11983 n_embd,
11984 n_active,
11985 n_pairs,
11986 n_pairs,
11987 m.down_exps.qtype,
11988 m.down_exps.row_bytes,
11989 )?
11990 } else {
11991 let act = e.moe_pairs_gelu_mul(&gate, &up, n_pairs * n_ff_exp)?;
11992 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
11993 e.moe_pairs_matvec_q8_dec(
11994 &dev.ptr_row,
11995 2,
11996 &exi,
11997 &exo,
11998 &exp_d,
11999 &pself,
12000 &aq2,
12001 &ad2,
12002 n_ff_exp,
12003 n_embd,
12004 n_expert,
12005 n_active,
12006 n_pairs,
12007 m.down_exps.qtype,
12008 m.down_exps.row_bytes,
12009 )?
12010 };
12011 let mut moe_out = e.uninit(t * n_embd)?;
12012 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
12013 return Ok(moe_out);
12014 }
12015
12016 let g_len = m.gate_exps.expert_stride;
12017 let u_len = m.up_exps.expert_stride;
12018 let d_len = m.down_exps.expert_stride;
12019 let dev = m.dev_exps.as_ref().filter(|d| !d.gu_il);
12023 let (mut sg, mut su, mut sd) = if dev.is_some() {
12024 (None, None, None)
12025 } else {
12026 (
12027 Some(e.alloc_u8_uninit(g_len)?),
12028 Some(e.alloc_u8_uninit(u_len)?),
12029 Some(e.alloc_u8_uninit(d_len)?),
12030 )
12031 };
12032 let mut moe_out = e.zeros(t * n_embd)?;
12033 for tok in 0..t {
12034 let sel = &sel_all[tok * n_used..(tok + 1) * n_used];
12035 let w = &w_all[tok * n_used..(tok + 1) * n_used];
12036 let zt = moe_in.slice(tok * n_embd..(tok + 1) * n_embd);
12037 for (j, &ex) in sel.iter().enumerate() {
12038 let ex = ex as usize;
12039 let gate = match dev {
12040 Some(d) => e.qmatvec_view(
12041 &d.gate,
12042 ex * g_len..(ex + 1) * g_len,
12043 &zt,
12044 1,
12045 m.gate_exps.in_f,
12046 m.gate_exps.out_f,
12047 m.gate_exps.qtype,
12048 m.gate_exps.row_bytes,
12049 )?,
12050 None => {
12051 let sg = sg.as_mut().unwrap();
12052 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
12053 e.qmatvec_view(
12054 sg,
12055 0..g_len,
12056 &zt,
12057 1,
12058 m.gate_exps.in_f,
12059 m.gate_exps.out_f,
12060 m.gate_exps.qtype,
12061 m.gate_exps.row_bytes,
12062 )?
12063 }
12064 };
12065 let up = match dev {
12066 Some(d) => e.qmatvec_view(
12067 &d.up,
12068 ex * u_len..(ex + 1) * u_len,
12069 &zt,
12070 1,
12071 m.up_exps.in_f,
12072 m.up_exps.out_f,
12073 m.up_exps.qtype,
12074 m.up_exps.row_bytes,
12075 )?,
12076 None => {
12077 let su = su.as_mut().unwrap();
12078 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
12079 e.qmatvec_view(
12080 su,
12081 0..u_len,
12082 &zt,
12083 1,
12084 m.up_exps.in_f,
12085 m.up_exps.out_f,
12086 m.up_exps.qtype,
12087 m.up_exps.row_bytes,
12088 )?
12089 }
12090 };
12091 let mut act = e.uninit(n_ff_exp)?;
12092 e.gelu_tanh_mul(&gate, &up, &mut act, n_ff_exp)?;
12093 let actv = act.slice(0..n_ff_exp);
12094 let y = match dev {
12095 Some(d) => e.qmatvec_view(
12096 &d.down,
12097 ex * d_len..(ex + 1) * d_len,
12098 &actv,
12099 1,
12100 m.down_exps.in_f,
12101 m.down_exps.out_f,
12102 m.down_exps.qtype,
12103 m.down_exps.row_bytes,
12104 )?,
12105 None => {
12106 let sd = sd.as_mut().unwrap();
12107 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
12108 e.qmatvec_view(
12109 sd,
12110 0..d_len,
12111 &actv,
12112 1,
12113 m.down_exps.in_f,
12114 m.down_exps.out_f,
12115 m.down_exps.qtype,
12116 m.down_exps.row_bytes,
12117 )?
12118 }
12119 };
12120 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
12121 e.axpy_into(&y, w[j], &mut dst, n_embd)?;
12122 }
12123 }
12124 Ok(moe_out)
12125 }
12126
12127 fn gemma4_layer(
12129 &self,
12130 e: &Engine,
12131 il: usize,
12132 layer: &crate::hybrid::HybridLayer,
12133 x: &CudaSlice<f32>,
12134 pos_d: &CudaSlice<i32>,
12135 t: usize,
12136 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12137 let n_embd = self.cfg.n_embd as usize;
12138 let eps = self.cfg.rms_eps;
12139
12140 let mut h = e.zeros(t * n_embd)?;
12141 e.rms_norm(x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
12142 let Mixer::Full(fa) = &layer.mixer else {
12143 panic!("gemma4 layer {il} not full-attn")
12144 };
12145 let o = self.gemma4_attn(e, fa, il, &h, pos_d, t)?;
12146 let mut cur = e.zeros(t * n_embd)?;
12148 e.rms_norm(
12149 &o,
12150 layer.post_attn_norm.float_data(),
12151 &mut cur,
12152 n_embd,
12153 t,
12154 eps,
12155 )?;
12156 self.gemma4_layer_tail_add(e, layer, &cur, x, t)
12157 }
12158
12159 fn gemma4_layer_tail_add(
12163 &self,
12164 e: &Engine,
12165 layer: &crate::hybrid::HybridLayer,
12166 cur: &CudaSlice<f32>,
12167 x: &CudaSlice<f32>,
12168 t: usize,
12169 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12170 Ok(self.gemma4_layer_tail_add_n(e, layer, cur, x, t, None)?.0)
12171 }
12172
12173 fn gemma4_layer_tail_add_n(
12176 &self,
12177 e: &Engine,
12178 layer: &crate::hybrid::HybridLayer,
12179 cur: &CudaSlice<f32>,
12180 x: &CudaSlice<f32>,
12181 t: usize,
12182 next_norm: Option<&CudaSlice<f32>>,
12183 ) -> Result<(CudaSlice<f32>, Option<CudaSlice<f32>>), Box<dyn std::error::Error>> {
12184 let n_embd = self.cfg.n_embd as usize;
12185 let bits = layer.gemma4.as_ref().unwrap();
12186 let (sn, attn_out) = self.gemma4_layer_tail_core(e, layer, cur, x, t)?;
12187 let mut xn = e.uninit(t * n_embd)?;
12188 match next_norm {
12189 Some(w) => {
12190 let mut hn = e.uninit(t * n_embd)?;
12191 e.add_scale_rms_norm(
12192 &sn,
12193 &attn_out,
12194 bits.layer_scale,
12195 w,
12196 &mut xn,
12197 &mut hn,
12198 n_embd,
12199 t,
12200 self.cfg.rms_eps,
12201 )?;
12202 Ok((xn, Some(hn)))
12203 }
12204 None => {
12205 e.add_scale(&sn, &attn_out, bits.layer_scale, &mut xn, t * n_embd)?;
12206 Ok((xn, None))
12207 }
12208 }
12209 }
12210
12211 fn gemma4_layer_tail_core(
12214 &self,
12215 e: &Engine,
12216 layer: &crate::hybrid::HybridLayer,
12217 cur: &CudaSlice<f32>,
12218 x: &CudaSlice<f32>,
12219 t: usize,
12220 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
12221 self.gemma4_layer_tail_core_pn(e, layer, cur, x, t, None, false)
12222 }
12223
12224 fn gemma4_layer_tail_core_pn(
12231 &self,
12232 e: &Engine,
12233 layer: &crate::hybrid::HybridLayer,
12234 cur: &CudaSlice<f32>,
12235 x: &CudaSlice<f32>,
12236 t: usize,
12237 pre_norm: Option<&CudaSlice<f32>>,
12238 defer_post_norm: bool,
12239 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
12240 let n_embd = self.cfg.n_embd as usize;
12241 let eps = self.cfg.rms_eps;
12242 let bits = layer.gemma4.as_ref().unwrap();
12243
12244 let Some(mbits) = bits.moe_bits.as_ref() else {
12247 let crate::hybrid::Ffn::Dense {
12248 ffn_gate,
12249 ffn_up,
12250 ffn_down,
12251 } = &layer.ffn
12252 else {
12253 panic!("gemma4 dense layer without Dense ffn")
12254 };
12255 let mut attn_out = e.uninit(t * n_embd)?;
12256 let mut zsh = e.uninit(t * n_embd)?;
12257 let mut zpair: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
12260 match pre_norm {
12261 Some(wa) if t == 1 => {
12262 zpair = Some(e.rms_pre_add_rms_norm_q8z(
12263 cur,
12264 wa,
12265 x,
12266 bits.ffn_norm.float_data(),
12267 &mut attn_out,
12268 &mut zsh,
12269 n_embd,
12270 t,
12271 eps,
12272 )?);
12273 }
12274 Some(wa) => e.rms_pre_add_rms_norm(
12275 cur,
12276 wa,
12277 x,
12278 bits.ffn_norm.float_data(),
12279 &mut attn_out,
12280 &mut zsh,
12281 n_embd,
12282 t,
12283 eps,
12284 )?,
12285 None => e.add_rms_norm(
12286 cur,
12287 x,
12288 bits.ffn_norm.float_data(),
12289 &mut attn_out,
12290 &mut zsh,
12291 n_embd,
12292 t,
12293 eps,
12294 )?,
12295 }
12296 let n_ff = ffn_gate.out_features();
12297 let (gate, up) = if t == 1 {
12303 let (zq, zd) = match zpair {
12304 Some(p) => p,
12305 None => e.quantize_q8_1(&zsh, 1, n_embd)?,
12306 };
12307 match e.matmul_q4_fused2(ffn_gate, ffn_up, &zq, &zd)? {
12308 Some(p) => p,
12309 None => match e.matmul_nvfp4_fused2(ffn_gate, ffn_up, &zq, &zd, 1)? {
12311 Some(p) => p,
12312 None => (
12313 e.matmul_pre(ffn_gate, &zq, &zd, &zsh, 1)?,
12314 e.matmul_pre(ffn_up, &zq, &zd, &zsh, 1)?,
12315 ),
12316 },
12317 }
12318 } else {
12319 static F2B: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12324 let f2b = *F2B.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
12325 let fused = if f2b {
12326 let (zq, zd) = e.quantize_q8_1(&zsh, t, n_embd)?;
12327 e.matmul_q4_fused2_batched(ffn_gate, ffn_up, &zq, &zd, t)?
12328 } else {
12329 None
12330 };
12331 match fused {
12332 Some(p) => p,
12333 None => {
12334 e.mmq_act_begin();
12336 (e.matmul(ffn_gate, &zsh, t)?, e.matmul(ffn_up, &zsh, t)?)
12337 }
12338 }
12339 };
12340 let mut act = e.uninit(t * n_ff)?;
12341 let f0 = if e.uses_q8_1_fast(ffn_down) {
12344 let upv = e.view(&up, t * n_ff);
12345 let up_all = upv.slice(0..t * n_ff);
12346 let (aq, ad) = e.gelu_tanh_mul_q8_1(&gate, &up_all, &mut act, n_ff, t)?;
12347 e.matmul_pre(ffn_down, &aq, &ad, &act, t)?
12348 } else {
12349 e.gelu_tanh_mul(&gate, &up, &mut act, t * n_ff)?;
12350 e.matmul(ffn_down, &act, t)?
12351 };
12352 if defer_post_norm {
12353 return Ok((f0, attn_out));
12354 }
12355 let mut sn = e.uninit(t * n_embd)?;
12356 e.rms_norm(
12357 &f0,
12358 bits.post_ffw_norm.float_data(),
12359 &mut sn,
12360 n_embd,
12361 t,
12362 eps,
12363 )?;
12364 return Ok((sn, attn_out));
12365 };
12366
12367 assert!(pre_norm.is_none(), "pre-norm fold is dense-entry only");
12368 let mut attn_out = e.uninit(t * n_embd)?;
12373 let mut router_in = e.uninit(t * n_embd)?;
12374 let fast_moe = match &layer.ffn {
12375 crate::hybrid::Ffn::Moe(m) => {
12376 m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
12377 && expert_dp4a_supported(m.gate_exps.qtype)
12378 && expert_dp4a_supported(m.up_exps.qtype)
12379 && expert_dp4a_supported(m.down_exps.qtype)
12380 && std::env::var("MEMRA_GEMMA_MOE_FAST").as_deref() != Ok("0")
12381 }
12382 _ => false,
12383 };
12384 let q8z = t < PRIME_MIN_T && fast_moe;
12385 let (zsh_f32, zsh_q8, moe_q8) = if q8z {
12386 let (z0, m2) = e.add_rms_norm3_q8z(
12387 cur,
12388 x,
12389 bits.ffn_norm.float_data(),
12390 &mbits.router_scale_pre,
12391 mbits.pre_ffw_norm_2.float_data(),
12392 &mut attn_out,
12393 &mut router_in,
12394 n_embd,
12395 t,
12396 eps,
12397 )?;
12398 (None, Some(z0), Some(m2))
12399 } else {
12400 let mut zsh = e.uninit(t * n_embd)?;
12401 let mut moe_in = e.uninit(t * n_embd)?;
12402 e.add_rms_norm3(
12403 cur,
12404 x,
12405 bits.ffn_norm.float_data(),
12406 &mbits.router_scale_pre,
12407 mbits.pre_ffw_norm_2.float_data(),
12408 &mut attn_out,
12409 &mut zsh,
12410 &mut router_in,
12411 &mut moe_in,
12412 n_embd,
12413 t,
12414 eps,
12415 )?;
12416 (Some((zsh, moe_in)), None, None)
12417 };
12418 let attn_out2 = attn_out;
12419 #[allow(unused_variables)]
12420 let attn_out = &attn_out2;
12421 let n_ff = mbits.shared_gate.out_features();
12422 let (gate, up) = if let Some((zq, zd)) = zsh_q8.as_ref() {
12423 if t == 1 {
12424 match e.matmul_q4_fused2(&mbits.shared_gate, &mbits.shared_up, zq, zd)? {
12425 Some(p) => p,
12426 None => match e.matmul_nvfp4_fused2(
12427 &mbits.shared_gate,
12428 &mbits.shared_up,
12429 zq,
12430 zd,
12431 1,
12432 )? {
12433 Some(p) => p,
12434 None => {
12435 let h0 = e.zeros(0)?;
12436 (
12437 e.matmul_pre(&mbits.shared_gate, zq, zd, &h0, 1)?,
12438 e.matmul_pre(&mbits.shared_up, zq, zd, &h0, 1)?,
12439 )
12440 }
12441 },
12442 }
12443 } else {
12444 let h0 = e.zeros(0)?;
12446 (
12447 e.matmul_pre(&mbits.shared_gate, zq, zd, &h0, t)?,
12448 e.matmul_pre(&mbits.shared_up, zq, zd, &h0, t)?,
12449 )
12450 }
12451 } else {
12452 let (zsh, _) = zsh_f32.as_ref().unwrap();
12453 (
12454 e.matmul(&mbits.shared_gate, zsh, t)?,
12455 e.matmul(&mbits.shared_up, zsh, t)?,
12456 )
12457 };
12458 let mut act = e.uninit(t * n_ff)?;
12459 e.gelu_tanh_mul(&gate, &up, &mut act, t * n_ff)?;
12460 let mlp0 = e.matmul(&mbits.shared_down, &act, t)?;
12461 let crate::hybrid::Ffn::Moe(m) = &layer.ffn else {
12462 panic!("gemma4 layer not MoE")
12463 };
12464 let moe0 = match (&moe_q8, &zsh_f32) {
12465 (Some(mq), _) => self.gemma4_moe_q8(e, m, mbits, mq, &router_in, t)?,
12466 (None, Some((_, moe_in))) => self.gemma4_moe(e, m, mbits, moe_in, &router_in, t)?,
12467 _ => unreachable!(),
12468 };
12469 let mut mlp = e.uninit(t * n_embd)?;
12471 let mut moe = e.uninit(t * n_embd)?;
12472 e.rms_norm2x(
12473 &mlp0,
12474 &moe0,
12475 mbits.post_ffw_norm_1.float_data(),
12476 mbits.post_ffw_norm_2.float_data(),
12477 &mut mlp,
12478 &mut moe,
12479 n_embd,
12480 t,
12481 eps,
12482 )?;
12483
12484 let mut sum = e.uninit(t * n_embd)?;
12487 let mut sn = e.uninit(t * n_embd)?;
12488 e.add_rms_norm(
12489 &mlp,
12490 &moe,
12491 bits.post_ffw_norm.float_data(),
12492 &mut sum,
12493 &mut sn,
12494 n_embd,
12495 t,
12496 eps,
12497 )?;
12498 Ok((sn, attn_out2))
12499 }
12500
12501 pub(crate) fn gemma4_layer_tail_add_nq_pn(
12511 &self,
12512 e: &Engine,
12513 layer: &crate::hybrid::HybridLayer,
12514 o: &CudaSlice<f32>,
12515 x: &CudaSlice<f32>,
12516 t: usize,
12517 next_norm: Option<&CudaSlice<f32>>,
12518 ) -> Result<(CudaSlice<f32>, Option<(CudaSlice<i8>, CudaSlice<f32>)>), Box<dyn std::error::Error>>
12519 {
12520 let n_embd = self.cfg.n_embd as usize;
12521 let eps = self.cfg.rms_eps;
12522 let bits = layer.gemma4.as_ref().unwrap();
12523 if Engine::g4_pnfold_on() && matches!(layer.ffn, crate::hybrid::Ffn::Dense { .. }) {
12524 let (f0, attn_out) = self.gemma4_layer_tail_core_pn(
12525 e,
12526 layer,
12527 o,
12528 x,
12529 t,
12530 Some(layer.post_attn_norm.float_data()),
12531 true,
12532 )?;
12533 let mut xn = e.uninit(t * n_embd)?;
12534 return match next_norm {
12535 Some(w) => {
12536 let pair = e.rms_pre_add_scale_rms_norm_q8_1(
12537 &f0,
12538 bits.post_ffw_norm.float_data(),
12539 &attn_out,
12540 bits.layer_scale,
12541 w,
12542 &mut xn,
12543 n_embd,
12544 t,
12545 eps,
12546 )?;
12547 Ok((xn, Some(pair)))
12548 }
12549 None => {
12550 let mut sn = e.uninit(t * n_embd)?;
12551 e.rms_norm(
12552 &f0,
12553 bits.post_ffw_norm.float_data(),
12554 &mut sn,
12555 n_embd,
12556 t,
12557 eps,
12558 )?;
12559 e.add_scale(&sn, &attn_out, bits.layer_scale, &mut xn, t * n_embd)?;
12560 Ok((xn, None))
12561 }
12562 };
12563 }
12564 let mut cur = e.uninit(t * n_embd)?;
12565 e.rms_norm(
12566 o,
12567 layer.post_attn_norm.float_data(),
12568 &mut cur,
12569 n_embd,
12570 t,
12571 eps,
12572 )?;
12573 self.gemma4_layer_tail_add_nq(e, layer, &cur, x, t, next_norm)
12574 }
12575
12576 pub(crate) fn gemma4_layer_tail_add_nq(
12577 &self,
12578 e: &Engine,
12579 layer: &crate::hybrid::HybridLayer,
12580 cur: &CudaSlice<f32>,
12581 x: &CudaSlice<f32>,
12582 t: usize,
12583 next_norm: Option<&CudaSlice<f32>>,
12584 ) -> Result<(CudaSlice<f32>, Option<(CudaSlice<i8>, CudaSlice<f32>)>), Box<dyn std::error::Error>>
12585 {
12586 let n_embd = self.cfg.n_embd as usize;
12587 let bits = layer.gemma4.as_ref().unwrap();
12588 let (sn, attn_out) = self.gemma4_layer_tail_core(e, layer, cur, x, t)?;
12589 let mut xn = e.uninit(t * n_embd)?;
12590 match next_norm {
12591 Some(w) => {
12592 let pair = e.add_scale_rms_norm_q8_1(
12593 &sn,
12594 &attn_out,
12595 bits.layer_scale,
12596 w,
12597 &mut xn,
12598 n_embd,
12599 t,
12600 self.cfg.rms_eps,
12601 )?;
12602 Ok((xn, Some(pair)))
12603 }
12604 None => {
12605 e.add_scale(&sn, &attn_out, bits.layer_scale, &mut xn, t * n_embd)?;
12606 Ok((xn, None))
12607 }
12608 }
12609 }
12610
12611 fn gemma4_forward(
12614 &self,
12615 e: &Engine,
12616 tokens: &[u32],
12617 last_only: bool,
12618 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
12619 if self.is_gemma4_e4b() {
12622 return self.gemma4_e4b_forward(e, tokens, last_only);
12623 }
12624 let n_embd = self.cfg.n_embd as usize;
12625 let t = tokens.len();
12626 let pos: Vec<i32> = (0..t as i32).collect();
12627 let pos_d = e.htod_i32(&pos)?;
12628
12629 let mut x = self.embed(e, tokens)?;
12630 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
12631 let probe = std::env::var("MEMRA_GEMMA_PROBE").is_ok();
12634 let stat =
12635 |e: &Engine, x: &CudaSlice<f32>, tag: &str| -> Result<(), Box<dyn std::error::Error>> {
12636 let h = e.dtoh(x)?;
12637 let bad = h.iter().filter(|v| !v.is_finite()).count();
12638 let mx = h
12639 .iter()
12640 .filter(|v| v.is_finite())
12641 .fold(0.0f32, |m, v| m.max(v.abs()));
12642 eprintln!(
12643 "[gemma-probe] {tag}: tok0_first3={:?} bad={bad} max={mx:.3e}",
12644 &h[..3]
12645 );
12646 Ok(())
12647 };
12648 if probe {
12649 stat(e, &x, "embed")?;
12650 }
12651 for (il, layer) in self.layers.iter().enumerate() {
12652 x = self.gemma4_layer(e, il, layer, &x, &pos_d, t)?;
12653 if probe {
12654 stat(e, &x, &format!("L{il}"))?;
12655 }
12656 }
12657 let mut hn = e.zeros(t * n_embd)?;
12658 e.rms_norm(
12659 &x,
12660 self.output_norm.float_data(),
12661 &mut hn,
12662 n_embd,
12663 t,
12664 self.cfg.rms_eps,
12665 )?;
12666 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
12667 let n_vocab = self.output.out_features();
12668 let logits = if last_only {
12669 let hv = e.view(&hn, t * n_embd);
12670 let last_row = hv.slice((t - 1) * n_embd..t * n_embd);
12671 let mut hlast = e.zeros(n_embd)?;
12672 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
12673 let mut ld = e.matmul(&self.output, &hlast, 1)?;
12674 e.softcap(&mut ld, cap, n_vocab)?;
12675 self.gemma4_suppress(e, &mut ld, 1)?;
12676 e.dtoh(&ld)?
12677 } else {
12678 let mut ld = e.matmul(&self.output, &hn, t)?;
12679 e.softcap(&mut ld, cap, t * n_vocab)?;
12680 self.gemma4_suppress(e, &mut ld, t)?;
12681 e.dtoh(&ld)?
12682 };
12683 Ok(logits)
12684 }
12685
12686 pub(crate) fn gemma4_prime(
12691 &self,
12692 e: &Engine,
12693 tokens: &[u32],
12694 cache: &mut Cache,
12695 overlay: Option<&crate::vision::EmbedOverlay>,
12696 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
12697 if cache.pos != 0 {
12702 return Err(
12703 "gemma4 prime v0 is fresh-prompt only (no continuation/chunked prime) \
12704 — prime the full prompt in one call or decode tokenwise"
12705 .into(),
12706 );
12707 }
12708 let n_embd = self.cfg.n_embd as usize;
12709 let eps = self.cfg.rms_eps;
12710 let t = tokens.len();
12711 let pos: Vec<i32> = (0..t as i32).collect();
12712 let pos_d = e.htod_i32(&pos)?;
12713 let mut x = self.embed(e, tokens)?;
12714 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
12715 let island: Option<CudaSlice<i32>> = match overlay {
12722 Some(ov) => {
12723 let mut span_id = vec![-1i32; t];
12724 for (i, &(pos, row_off, n_rows)) in ov.spans.iter().enumerate() {
12725 if pos + n_rows > t {
12726 return Err(format!(
12727 "gemma4 overlay span {i} [{pos}, {}) exceeds the prompt ({t})",
12728 pos + n_rows
12729 )
12730 .into());
12731 }
12732 let view = ov.rows.slice(row_off * n_embd..(row_off + n_rows) * n_embd);
12733 e.copy_view_into(&mut x, pos * n_embd, &view, n_rows * n_embd)?;
12734 for s in span_id.iter_mut().skip(pos).take(n_rows) {
12735 *s = i as i32;
12736 }
12737 }
12738 if std::env::var("MEMRA_GV_FORCE_CAUSAL").as_deref() == Ok("1") {
12742 None
12743 } else {
12744 Some(e.htod_i32(&span_id)?)
12745 }
12746 }
12747 None => None,
12748 };
12749 for (il, layer) in self.layers.iter().enumerate() {
12750 let mut h = e.zeros(t * n_embd)?;
12751 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
12752 let Mixer::Full(fa) = &layer.mixer else {
12753 panic!("gemma4 layer not full-attn")
12754 };
12755 let trace = il == 0 && std::env::var("MEMRA_G4_PRIME_TRACE").as_deref() == Ok("1");
12756 if trace {
12757 let v = e.dtoh(&h)?;
12758 let nan = v.iter().filter(|x| x.is_nan()).count();
12759 eprintln!("[g4-prime-trace] L0 post-attn_norm: nan={nan}/{}", v.len());
12760 }
12761 let o =
12762 self.gemma4_attn_prime(e, fa, il, &h, &pos_d, t, Some(cache), island.as_ref())?;
12763 if trace {
12764 let v = e.dtoh(&o)?;
12765 let nan = v.iter().filter(|x| x.is_nan()).count();
12766 eprintln!("[g4-prime-trace] L0 post-attn: nan={nan}/{}", v.len());
12767 }
12768 let mut cur = e.zeros(t * n_embd)?;
12769 e.rms_norm(
12770 &o,
12771 layer.post_attn_norm.float_data(),
12772 &mut cur,
12773 n_embd,
12774 t,
12775 eps,
12776 )?;
12777 x = self.gemma4_layer_tail_add(e, layer, &cur, &x, t)?;
12778 self.dflash_tap(e, cache, il, &x, t)?;
12779 if std::env::var("MEMRA_G4_PRIME_TRACE").as_deref() == Ok("1") {
12781 let h = e.dtoh(&x)?;
12782 let nan = h.iter().filter(|v| v.is_nan()).count();
12783 let amax = h.iter().fold(0f32, |a, v| a.max(v.abs()));
12784 eprintln!(
12785 "[g4-prime-trace] layer {il}: nan={nan}/{} amax={amax:.3}",
12786 h.len()
12787 );
12788 if nan > 0 {
12789 return Err(format!("g4-prime-trace: first NaN at layer {il}").into());
12790 }
12791 }
12792 }
12793 cache.pos += t;
12794 let hiddens = e.clone_dtod(&x)?;
12795 let xv = e.view(&x, t * n_embd);
12796 let last_row = xv.slice((t - 1) * n_embd..t * n_embd);
12797 let mut h_seed = e.zeros(n_embd)?;
12798 e.copy_view_into(&mut h_seed, 0, &last_row, n_embd)?;
12799 let mut hn = e.uninit(n_embd)?;
12800 e.rms_norm(
12801 &h_seed,
12802 self.output_norm.float_data(),
12803 &mut hn,
12804 n_embd,
12805 1,
12806 eps,
12807 )?;
12808 let mut ld = e.matmul(&self.output, &hn, 1)?;
12809 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
12810 e.softcap(&mut ld, cap, self.output.out_features())?;
12811 self.gemma4_suppress(e, &mut ld, 1)?;
12812 let logits = e.dtoh(&ld)?;
12813 Ok((logits, h_seed, hiddens))
12814 }
12815
12816 fn gemma4_decode_attn(
12821 &self,
12822 e: &Engine,
12823 fa: &crate::hybrid::FullAttnLayer,
12824 il: usize,
12825 hq: &CudaSlice<i8>,
12826 hdq: &CudaSlice<f32>,
12827 pos_d: &CudaSlice<i32>,
12828 cache: &mut Cache,
12829 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12830 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
12831 let eps = self.cfg.rms_eps;
12832 let aux = self.gemma4_aux.as_ref().unwrap();
12833 let ones = aux.ones(e);
12834 #[cfg(debug_assertions)]
12835 crate::debug_assert_tensor_stream_device(ones, &e.stream(), "gemma4_decode_attn.ones");
12836 let (hq, hdq) = (hq, hdq);
12837 let h0 = e.zeros(0)?;
12838 let h = &h0;
12839 let (q0, k0, v0) = if swa {
12840 match e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, &hq, &hdq)? {
12841 Some(t3) => t3,
12842 None => match e.matmul_nvfp4_fused2(&fa.wq, &fa.wk, &hq, &hdq, 1)? {
12845 Some((q0, k0)) => {
12846 let v0 = e.matmul_pre(&fa.wv, &hq, &hdq, h, 1)?;
12847 (q0, k0, v0)
12848 }
12849 None => (
12850 e.matmul_pre(&fa.wq, &hq, &hdq, h, 1)?,
12851 e.matmul_pre(&fa.wk, &hq, &hdq, h, 1)?,
12852 e.matmul_pre(&fa.wv, &hq, &hdq, h, 1)?,
12853 ),
12854 },
12855 }
12856 } else {
12857 let (q0, k0) = match e.matmul_q4_fused2(&fa.wq, &fa.wk, &hq, &hdq)? {
12858 Some(p) => p,
12859 None => match e.matmul_nvfp4_fused2(&fa.wq, &fa.wk, &hq, &hdq, 1)? {
12860 Some(p) => p,
12861 None => (
12862 e.matmul_pre(&fa.wq, &hq, &hdq, h, 1)?,
12863 e.matmul_pre(&fa.wk, &hq, &hdq, h, 1)?,
12864 ),
12865 },
12866 };
12867 let v0 = e.clone_dtod(&k0)?;
12868 (q0, k0, v0)
12869 };
12870 let mut q = e.uninit(nh * hd)?;
12871 let mut k = e.uninit(nkv * hd)?;
12872 let mut v = e.uninit(nkv * hd)?;
12873 let ff = if swa {
12876 None
12877 } else {
12878 Some(
12879 aux.rope_freqs(e)
12880 .expect("gemma4 global rope needs rope_freqs.weight"),
12881 )
12882 };
12883 #[cfg(debug_assertions)]
12884 if let Some(ff) = ff {
12885 crate::debug_assert_tensor_stream_device(
12886 ff,
12887 &e.stream(),
12888 "gemma4_decode_attn.rope_freqs",
12889 );
12890 }
12891 let kvl = cache.kv[il].as_mut().unwrap();
12892 let kv_fp8 = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
12893 if crate::Engine::qkv_append_on() {
12894 e.rms_norm_qkv_rope_append(
12898 &q0,
12899 &k0,
12900 &v0,
12901 fa.q_norm.float_data(),
12902 fa.k_norm.float_data(),
12903 ones,
12904 &mut q,
12905 &mut k,
12906 &mut v,
12907 hd,
12908 self.gemma4_rope_dims(il),
12909 nh,
12910 nkv,
12911 pos_d,
12912 nh,
12913 nkv,
12914 base,
12915 1.0,
12916 ff,
12917 eps,
12918 &mut kvl.k,
12919 &mut kvl.v,
12920 kvl.len,
12921 kvl.k_tok_bytes,
12922 kvl.v_tok_bytes,
12923 kv_fp8,
12924 )?;
12925 } else {
12926 e.rms_norm_qkv_rope(
12927 &q0,
12928 &k0,
12929 &v0,
12930 fa.q_norm.float_data(),
12931 fa.k_norm.float_data(),
12932 ones,
12933 &mut q,
12934 &mut k,
12935 &mut v,
12936 hd,
12937 self.gemma4_rope_dims(il),
12938 nh,
12939 nkv,
12940 pos_d,
12941 nh,
12942 nkv,
12943 base,
12944 1.0,
12945 ff,
12946 eps,
12947 )?;
12948 e.append_kv_quantized(
12949 &k,
12950 &v,
12951 &mut kvl.k,
12952 &mut kvl.v,
12953 kvl.len,
12954 kvl.kv_dim_k,
12955 kvl.kv_dim_v,
12956 kvl.k_tok_bytes,
12957 kvl.v_tok_bytes,
12958 kv_fp8,
12959 )?;
12960 }
12961 kvl.len += 1;
12962 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
12966 let mut attn = e.uninit(nh * hd)?;
12967 if !swa
12969 && hd == 512
12970 && kvl.len >= crate::fa512_min_tkv()
12971 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
12972 {
12973 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
12974 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
12975 let base = kvl.len as i32;
12977 e.i32_set_k(&mut kvl.len_d, base)?;
12978 e.fa_decode_rows(
12979 &q,
12980 &kp,
12981 &vp,
12982 &mut attn,
12983 hd,
12984 nh,
12985 nkv,
12986 kvl.len - 1,
12987 1,
12988 scale,
12989 kvl.k_tok_bytes,
12990 kvl.v_tok_bytes,
12991 Some((&kvl.len_d, -1)),
12992 false,
12993 false,
12994 None,
12995 )?;
12996 return Ok(e.matmul(&fa.wo, &attn, 1)?);
12997 }
12998 if swa
13000 && kvl.len > win
13001 && hd == 256
13002 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
13003 {
13004 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
13005 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
13006 let base = kvl.len as i32;
13007 e.i32_set_k(&mut kvl.len_d, base)?;
13008 e.fa_decode_rows_w(
13009 &q,
13010 &kp,
13011 &vp,
13012 &mut attn,
13013 hd,
13014 nh,
13015 nkv,
13016 &kvl.len_d,
13017 -1,
13018 1,
13019 scale,
13020 win,
13021 kvl.k_tok_bytes,
13022 kvl.v_tok_bytes,
13023 None,
13024 )?;
13025 return Ok(e.matmul(&fa.wo, &attn, 1)?);
13026 }
13027 let (off_tok, t_kv) = if swa && kvl.len > win {
13028 (kvl.len - win, win)
13029 } else {
13030 (0, kvl.len)
13031 };
13032 let k_view = e.view_u8_range(
13033 &kvl.k,
13034 off_tok * kvl.k_tok_bytes,
13035 (off_tok + t_kv) * kvl.k_tok_bytes,
13036 );
13037 let v_view = e.view_u8_range(
13038 &kvl.v,
13039 off_tok * kvl.v_tok_bytes,
13040 (off_tok + t_kv) * kvl.v_tok_bytes,
13041 );
13042 e.fa_decode_kvmod(
13043 &q,
13044 &k_view,
13045 &v_view,
13046 &mut attn,
13047 hd,
13048 nh,
13049 nkv,
13050 t_kv,
13051 scale,
13052 kvl.k_tok_bytes,
13053 kvl.v_tok_bytes,
13054 swa && crate::Engine::wkv_on(),
13055 )?;
13056 Ok(e.matmul(&fa.wo, &attn, 1)?)
13057 }
13058
13059 #[allow(clippy::too_many_arguments)]
13066 pub fn gemma4_decode_step_dc(
13067 &self,
13068 e: &Engine,
13069 token_d: &CudaSlice<u32>,
13070 pos_d: &mut CudaSlice<i32>,
13071 embd_gpu: &CudaSlice<u8>,
13072 embd_qt: i32,
13073 embd_rb: usize,
13074 cache: &mut Cache,
13075 n_vocab: usize,
13076 cap_bucket_max: Option<(usize, usize)>,
13077 ) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
13078 let mut tok_out = e.stream().alloc_zeros::<u32>(1)?;
13079 self.gemma4_decode_step_dc_into(
13080 e,
13081 token_d,
13082 pos_d,
13083 embd_gpu,
13084 embd_qt,
13085 embd_rb,
13086 cache,
13087 n_vocab,
13088 cap_bucket_max,
13089 &mut tok_out,
13090 )?;
13091 Ok(tok_out)
13092 }
13093
13094 #[allow(clippy::too_many_arguments)]
13097 pub fn gemma4_decode_step_dc_into(
13098 &self,
13099 e: &Engine,
13100 token_d: &CudaSlice<u32>,
13101 pos_d: &mut CudaSlice<i32>,
13102 embd_gpu: &CudaSlice<u8>,
13103 embd_qt: i32,
13104 embd_rb: usize,
13105 cache: &mut Cache,
13106 n_vocab: usize,
13107 cap_bucket_max: Option<(usize, usize)>,
13108 tok_out: &mut CudaSlice<u32>,
13109 ) -> Result<(), Box<dyn std::error::Error>> {
13110 let n_embd = self.cfg.n_embd as usize;
13111 let eps = self.cfg.rms_eps;
13112 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
13113 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
13114 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
13115 let n_layers = self.layers.len();
13116 for (il, layer) in self.layers.iter().enumerate() {
13117 let (hq, hdq) = match h_carry.take() {
13118 Some(p) => p,
13119 None => {
13120 e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, 1, eps)?
13121 }
13122 };
13123 let Mixer::Full(fa) = &layer.mixer else {
13124 panic!("gemma4 layer {il} not full-attn")
13125 };
13126 let o =
13127 self.gemma4_decode_attn_dc(e, fa, il, &hq, &hdq, pos_d, cache, cap_bucket_max)?;
13128 let next_norm = if il + 1 < n_layers {
13129 Some(self.layers[il + 1].attn_norm.float_data())
13130 } else {
13131 None
13132 };
13133 let (xn, hn) = self.gemma4_layer_tail_add_nq_pn(e, layer, &o, &x, 1, next_norm)?;
13134 x = xn;
13135 h_carry = hn;
13136 }
13137 let mut hn = e.uninit(n_embd)?;
13138 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
13139 let mut logits = e.matmul(&self.output, &hn, 1)?;
13140 self.gemma4_suppress(e, &mut logits, 1)?; e.argmax_token_device_into(&logits, tok_out, n_vocab)?;
13142 e.inc_seqlen(pos_d)?;
13143 if cap_bucket_max.is_none() {
13144 cache.pos += 1;
13145 }
13146 Ok(())
13147 }
13148
13149 pub fn g4_dc_slots(&self, e: &Engine) -> Result<G4DcSlots, Box<dyn std::error::Error>> {
13156 let n_embd = self.cfg.n_embd as usize;
13157 let n_vocab = self.output.out_features();
13158 let n_layers = self.layers.len();
13159 let (mut qmax, mut kvmax, mut ffmax) = (0usize, 0usize, 0usize);
13160 for il in 0..n_layers {
13161 let (hd, nkv, nh, _b, _s, _w) = self.gemma4_geom(il);
13162 qmax = qmax.max(nh * hd);
13163 kvmax = kvmax.max(nkv * hd);
13164 if let crate::hybrid::Ffn::Dense { ffn_gate, .. } = &self.layers[il].ffn {
13165 ffmax = ffmax.max(ffn_gate.out_features());
13166 }
13167 }
13168 Ok(G4DcSlots {
13169 x: e.uninit(n_embd)?,
13170 xn: e.uninit(n_embd)?,
13171 cur: e.uninit(n_embd)?,
13172 hq: e.alloc_i8_uninit(n_embd)?,
13173 hd_: e.uninit(n_embd / 32)?,
13174 q0: e.uninit(qmax)?,
13175 k0: e.uninit(kvmax)?,
13176 v0: e.uninit(kvmax)?,
13177 q: e.uninit(qmax)?,
13178 k: e.uninit(kvmax)?,
13179 v: e.uninit(kvmax)?,
13180 attn: e.uninit(qmax)?,
13181 o: e.uninit(n_embd)?,
13182 attn_out: e.uninit(n_embd)?,
13183 zsh: e.uninit(n_embd)?,
13184 zq: e.alloc_i8_uninit(n_embd.max(qmax))?,
13187 zd: e.uninit(n_embd.max(qmax) / 32)?,
13188 gate: e.uninit(ffmax)?,
13189 up: e.uninit(ffmax)?,
13190 act: e.uninit(ffmax)?,
13191 actq: e.alloc_i8_uninit(ffmax)?,
13192 actd: e.uninit(ffmax / 32)?,
13193 f0: e.uninit(n_embd)?,
13194 sn: e.uninit(n_embd)?,
13195 hn: e.uninit(n_embd)?,
13196 logits: e.uninit(n_vocab)?,
13197 })
13198 }
13199
13200 fn g4_matvec_m1_into(
13203 &self,
13204 e: &Engine,
13205 w: &crate::model::GpuTensor,
13206 aq: &CudaSlice<i8>,
13207 ad: &CudaSlice<f32>,
13208 y: &mut CudaSlice<f32>,
13209 ) -> Result<(), Box<dyn std::error::Error>> {
13210 use crate::model::GpuTensor;
13211 let (bytes, qtype, row_bytes, scale, rp) = match w {
13212 GpuTensor::Quant {
13213 bytes,
13214 qtype,
13215 row_bytes,
13216 scale,
13217 rp,
13218 ..
13219 } => (bytes, *qtype, *row_bytes, *scale, *rp),
13220 _ => return Err("g4_matvec_m1_into: non-quant tensor".into()),
13221 };
13222 let (mbytes, mrp) = match w {
13223 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
13224 _ => (bytes, rp),
13225 };
13226 e.qmatvec_mmvq_into(
13227 mbytes,
13228 aq,
13229 ad,
13230 1,
13231 w.in_features(),
13232 w.out_features(),
13233 qtype,
13234 row_bytes,
13235 scale,
13236 mrp,
13237 y,
13238 )
13239 }
13240
13241 #[allow(clippy::too_many_arguments)]
13245 pub fn gemma4_decode_step_dc_slotted(
13246 &self,
13247 e: &Engine,
13248 token_d: &CudaSlice<u32>,
13249 pos_d: &mut CudaSlice<i32>,
13250 embd_gpu: &CudaSlice<u8>,
13251 embd_qt: i32,
13252 embd_rb: usize,
13253 cache: &mut Cache,
13254 n_vocab: usize,
13255 cap_bucket_max: Option<(usize, usize)>,
13256 sl: &mut G4DcSlots,
13257 tok_out: &mut CudaSlice<u32>,
13258 ring: Option<(&mut CudaSlice<u32>, usize)>,
13259 ) -> Result<(), Box<dyn std::error::Error>> {
13260 let n_embd = self.cfg.n_embd as usize;
13261 let eps = self.cfg.rms_eps;
13262 e.embed_gather_device_into(embd_gpu, token_d, &mut sl.x, n_embd, embd_qt, embd_rb)?;
13263 e.scale_inplace(&mut sl.x, (n_embd as f32).sqrt(), n_embd)?;
13264 let n_layers = self.layers.len();
13265 let mut has_carry = false;
13266 for il in 0..n_layers {
13267 if !has_carry {
13268 e.rms_norm_q8_1_into(
13269 &sl.x,
13270 self.layers[il].attn_norm.float_data(),
13271 n_embd,
13272 1,
13273 eps,
13274 &mut sl.hq,
13275 &mut sl.hd_,
13276 )?;
13277 }
13278 has_carry = true;
13279 let layer = &self.layers[il];
13280 let Mixer::Full(fa) = &layer.mixer else {
13281 panic!("gemma4 layer {il} not full-attn")
13282 };
13283 self.gemma4_decode_attn_dc_slotted(e, fa, il, pos_d, cache, cap_bucket_max, sl)?;
13284 if !Engine::g4_pnfold_on() {
13287 e.rms_norm(
13288 &sl.o,
13289 layer.post_attn_norm.float_data(),
13290 &mut sl.cur,
13291 n_embd,
13292 1,
13293 eps,
13294 )?;
13295 }
13296 let next_norm = if il + 1 < n_layers {
13297 Some(self.layers[il + 1].attn_norm.float_data())
13298 } else {
13299 None
13300 };
13301 self.gemma4_layer_tail_slotted(e, layer, next_norm, sl)?;
13302 std::mem::swap(&mut sl.x, &mut sl.xn);
13303 }
13304 e.rms_norm(
13305 &sl.x,
13306 self.output_norm.float_data(),
13307 &mut sl.hn,
13308 n_embd,
13309 1,
13310 eps,
13311 )?;
13312 e.quantize_q8_1_into(&sl.hn, 1, n_embd, &mut sl.zq, &mut sl.zd)?;
13313 {
13315 let (zq, zd) = (&sl.zq, &sl.zd);
13316 let zq = unsafe { &*(zq as *const CudaSlice<i8>) };
13317 let zd = unsafe { &*(zd as *const CudaSlice<f32>) };
13318 self.g4_matvec_m1_into(e, &self.output, zq, zd, &mut sl.logits)?;
13319 }
13320 self.gemma4_suppress(e, &mut sl.logits, 1)?;
13321 e.argmax_token_device_into(&sl.logits, tok_out, n_vocab)?;
13322 if let Some((ring, base)) = ring {
13323 e.plain_tok_ring(tok_out, pos_d, base, ring)?;
13327 }
13328 e.inc_seqlen(pos_d)?;
13329 if cap_bucket_max.is_none() {
13330 cache.pos += 1;
13331 }
13332 Ok(())
13333 }
13334
13335 #[allow(clippy::too_many_arguments)]
13337 fn gemma4_decode_attn_dc_slotted(
13338 &self,
13339 e: &Engine,
13340 fa: &crate::hybrid::FullAttnLayer,
13341 il: usize,
13342 pos_d: &CudaSlice<i32>,
13343 cache: &mut Cache,
13344 cap_bucket_max: Option<(usize, usize)>,
13345 sl: &mut G4DcSlots,
13346 ) -> Result<(), Box<dyn std::error::Error>> {
13347 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
13348 let eps = self.cfg.rms_eps;
13349 let aux = self.gemma4_aux.as_ref().unwrap();
13350 let ones = aux.ones(e);
13351 #[cfg(debug_assertions)]
13352 crate::debug_assert_tensor_stream_device(
13353 ones,
13354 &e.stream(),
13355 "gemma4_decode_attn_dc_slotted.ones",
13356 );
13357 {
13358 let hq = unsafe { &*(&sl.hq as *const CudaSlice<i8>) };
13359 let hdq = unsafe { &*(&sl.hd_ as *const CudaSlice<f32>) };
13360 if swa {
13361 if !e.matmul_q4_fused3_into(
13362 &fa.wq, &fa.wk, &fa.wv, hq, hdq, &mut sl.q0, &mut sl.k0, &mut sl.v0,
13363 )? {
13364 if e.matmul_nvfp4_fused2_into(&fa.wq, &fa.wk, hq, hdq, &mut sl.q0, &mut sl.k0)?
13368 {
13369 self.g4_matvec_m1_into(e, &fa.wv, hq, hdq, &mut sl.v0)?;
13370 } else {
13371 return Err("slotted step: fused3 unavailable (non-uniform trunk)".into());
13372 }
13373 }
13374 } else {
13375 if !e.matmul_q4_fused2_into(&fa.wq, &fa.wk, hq, hdq, &mut sl.q0, &mut sl.k0)?
13376 && !e
13377 .matmul_nvfp4_fused2_into(&fa.wq, &fa.wk, hq, hdq, &mut sl.q0, &mut sl.k0)?
13378 {
13379 return Err("slotted step: fused2 unavailable".into());
13380 }
13381 let k0r = unsafe { &*(&sl.k0 as *const CudaSlice<f32>) };
13382 e.copy_into(&mut sl.v0, 0, k0r, nkv * hd)?;
13383 }
13384 }
13385 let ff = if swa {
13388 None
13389 } else {
13390 Some(
13391 aux.rope_freqs(e)
13392 .expect("gemma4 global rope needs rope_freqs.weight"),
13393 )
13394 };
13395 #[cfg(debug_assertions)]
13396 if let Some(ff) = ff {
13397 crate::debug_assert_tensor_stream_device(
13398 ff,
13399 &e.stream(),
13400 "gemma4_decode_attn_dc_slotted.rope_freqs",
13401 );
13402 }
13403 let kvl = cache.kv[il].as_mut().unwrap();
13404 let kv_fp8 = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
13405 if crate::Engine::qkv_append_on() {
13406 e.rms_norm_qkv_rope_append_dc(
13408 &sl.q0,
13409 &sl.k0,
13410 &sl.v0,
13411 fa.q_norm.float_data(),
13412 fa.k_norm.float_data(),
13413 ones,
13414 &mut sl.q,
13415 &mut sl.k,
13416 &mut sl.v,
13417 hd,
13418 self.gemma4_rope_dims(il),
13419 nh,
13420 nkv,
13421 pos_d,
13422 nh,
13423 nkv,
13424 base,
13425 1.0,
13426 ff,
13427 eps,
13428 &mut kvl.k,
13429 &mut kvl.v,
13430 &kvl.len_d,
13431 kvl.k_tok_bytes,
13432 kvl.v_tok_bytes,
13433 kv_fp8,
13434 )?;
13435 } else {
13436 e.rms_norm_qkv_rope(
13437 &sl.q0,
13438 &sl.k0,
13439 &sl.v0,
13440 fa.q_norm.float_data(),
13441 fa.k_norm.float_data(),
13442 ones,
13443 &mut sl.q,
13444 &mut sl.k,
13445 &mut sl.v,
13446 hd,
13447 self.gemma4_rope_dims(il),
13448 nh,
13449 nkv,
13450 pos_d,
13451 nh,
13452 nkv,
13453 base,
13454 1.0,
13455 ff,
13456 eps,
13457 )?;
13458 e.append_kv_quantized_dc(
13459 &sl.k,
13460 &sl.v,
13461 &mut kvl.k,
13462 &mut kvl.v,
13463 &kvl.len_d,
13464 kvl.kv_dim_k,
13465 kvl.kv_dim_v,
13466 kvl.k_tok_bytes,
13467 kvl.v_tok_bytes,
13468 kv_fp8,
13469 )?;
13470 }
13471 e.inc_seqlen(&mut kvl.len_d)?;
13472 let (b_swa, b_glob) = cap_bucket_max.expect("slotted step is capture-only");
13473 let k_view = e.view_u8(&kvl.k, kvl.k.len());
13474 let v_view = e.view_u8(&kvl.v, kvl.v.len());
13475 let rows_on = std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0");
13476 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
13477 let mut fa_q8 = false;
13481 if !swa && hd == 512 && b_glob >= crate::fa512_min_tkv() && rows_on {
13482 e.fa_decode_rows(
13483 &sl.q,
13484 &k_view,
13485 &v_view,
13486 &mut sl.attn,
13487 hd,
13488 nh,
13489 nkv,
13490 b_glob - 1,
13491 1,
13492 scale,
13493 kvl.k_tok_bytes,
13494 kvl.v_tok_bytes,
13495 Some((&kvl.len_d, -1)),
13496 false,
13497 false,
13498 Some((&mut sl.zq, &mut sl.zd)),
13499 )?;
13500 fa_q8 = true;
13501 } else if swa && b_swa > win && hd == 256 && rows_on {
13502 e.fa_decode_rows_w(
13503 &sl.q,
13504 &k_view,
13505 &v_view,
13506 &mut sl.attn,
13507 hd,
13508 nh,
13509 nkv,
13510 &kvl.len_d,
13511 -1,
13512 1,
13513 scale,
13514 win,
13515 kvl.k_tok_bytes,
13516 kvl.v_tok_bytes,
13517 Some((&mut sl.zq, &mut sl.zd)),
13518 )?;
13519 fa_q8 = true;
13520 } else {
13521 let b = if swa { b_swa } else { b_glob };
13522 e.fa_decode_dc(
13523 &sl.q,
13524 &k_view,
13525 &v_view,
13526 &mut sl.attn,
13527 hd,
13528 nh,
13529 nkv,
13530 &kvl.len_d,
13531 b,
13532 scale,
13533 kvl.k_tok_bytes,
13534 kvl.v_tok_bytes,
13535 swa && crate::Engine::wkv_on(),
13536 )?;
13537 }
13538 if !fa_q8 {
13539 let aq = unsafe { &*(&sl.attn as *const CudaSlice<f32>) };
13540 e.quantize_q8_1_into(aq, 1, nh * hd, &mut sl.zq, &mut sl.zd)?;
13541 }
13542 {
13543 let zq = unsafe { &*(&sl.zq as *const CudaSlice<i8>) };
13544 let zd = unsafe { &*(&sl.zd as *const CudaSlice<f32>) };
13545 self.g4_matvec_m1_into(e, &fa.wo, zq, zd, &mut sl.o)?;
13546 }
13547 Ok(())
13548 }
13549
13550 fn gemma4_layer_tail_slotted(
13553 &self,
13554 e: &Engine,
13555 layer: &crate::hybrid::HybridLayer,
13556 next_norm: Option<&CudaSlice<f32>>,
13557 sl: &mut G4DcSlots,
13558 ) -> Result<(), Box<dyn std::error::Error>> {
13559 let n_embd = self.cfg.n_embd as usize;
13560 let eps = self.cfg.rms_eps;
13561 let bits = layer.gemma4.as_ref().unwrap();
13562 let crate::hybrid::Ffn::Dense {
13563 ffn_gate,
13564 ffn_up,
13565 ffn_down,
13566 } = &layer.ffn
13567 else {
13568 return Err("slotted tail: dense ffn only".into());
13569 };
13570 let pnfold = Engine::g4_pnfold_on();
13571 if pnfold {
13572 let or = unsafe { &*(&sl.o as *const CudaSlice<f32>) };
13575 let xr = unsafe { &*(&sl.x as *const CudaSlice<f32>) };
13576 e.rms_pre_add_rms_norm_q8z_into(
13577 or,
13578 layer.post_attn_norm.float_data(),
13579 xr,
13580 bits.ffn_norm.float_data(),
13581 &mut sl.attn_out,
13582 &mut sl.zsh,
13583 n_embd,
13584 1,
13585 eps,
13586 &mut sl.zq,
13587 &mut sl.zd,
13588 )?;
13589 } else {
13590 e.add_rms_norm(
13591 &sl.cur,
13592 &sl.x,
13593 bits.ffn_norm.float_data(),
13594 &mut sl.attn_out,
13595 &mut sl.zsh,
13596 n_embd,
13597 1,
13598 eps,
13599 )?;
13600 }
13601 let n_ff = ffn_gate.out_features();
13602 if !pnfold {
13603 let zshr = unsafe { &*(&sl.zsh as *const CudaSlice<f32>) };
13604 e.quantize_q8_1_into(zshr, 1, n_embd, &mut sl.zq, &mut sl.zd)?;
13605 }
13606 {
13607 let zq = unsafe { &*(&sl.zq as *const CudaSlice<i8>) };
13608 let zd = unsafe { &*(&sl.zd as *const CudaSlice<f32>) };
13609 if !e.matmul_q4_fused2_into(ffn_gate, ffn_up, zq, zd, &mut sl.gate, &mut sl.up)?
13610 && !e.matmul_nvfp4_fused2_into(
13611 ffn_gate,
13612 ffn_up,
13613 zq,
13614 zd,
13615 &mut sl.gate,
13616 &mut sl.up,
13617 )?
13618 {
13619 return Err("slotted tail: ffn fused2 unavailable".into());
13620 }
13621 }
13622 debug_assert!(e.uses_q8_1_fast(ffn_down));
13623 {
13624 let upr = unsafe { &*(&sl.up as *const CudaSlice<f32>) };
13625 let upv = e.view(upr, n_ff);
13626 let up_all = upv.slice(0..n_ff);
13627 let gr = unsafe { &*(&sl.gate as *const CudaSlice<f32>) };
13628 e.gelu_tanh_mul_q8_1_into(
13629 gr,
13630 &up_all,
13631 &mut sl.act,
13632 n_ff,
13633 1,
13634 &mut sl.actq,
13635 &mut sl.actd,
13636 )?;
13637 }
13638 {
13639 let aq = unsafe { &*(&sl.actq as *const CudaSlice<i8>) };
13640 let ad = unsafe { &*(&sl.actd as *const CudaSlice<f32>) };
13641 self.g4_matvec_m1_into(e, ffn_down, aq, ad, &mut sl.f0)?;
13642 }
13643 if pnfold {
13644 if let Some(w) = next_norm {
13647 let f0r = unsafe { &*(&sl.f0 as *const CudaSlice<f32>) };
13648 let aor = unsafe { &*(&sl.attn_out as *const CudaSlice<f32>) };
13649 e.rms_pre_add_scale_rms_norm_q8_1_into(
13650 f0r,
13651 bits.post_ffw_norm.float_data(),
13652 aor,
13653 bits.layer_scale,
13654 w,
13655 &mut sl.xn,
13656 n_embd,
13657 1,
13658 eps,
13659 &mut sl.hq,
13660 &mut sl.hd_,
13661 )?;
13662 return Ok(());
13663 }
13664 }
13665 e.rms_norm(
13666 &sl.f0,
13667 bits.post_ffw_norm.float_data(),
13668 &mut sl.sn,
13669 n_embd,
13670 1,
13671 eps,
13672 )?;
13673 match next_norm {
13674 Some(w) => {
13675 e.add_scale_rms_norm_q8_1_into(
13676 &sl.sn,
13677 &sl.attn_out,
13678 bits.layer_scale,
13679 w,
13680 &mut sl.xn,
13681 n_embd,
13682 1,
13683 eps,
13684 &mut sl.hq,
13685 &mut sl.hd_,
13686 )?;
13687 }
13688 None => {
13689 e.add_scale(&sl.sn, &sl.attn_out, bits.layer_scale, &mut sl.xn, n_embd)?;
13690 }
13691 }
13692 Ok(())
13693 }
13694
13695 #[allow(clippy::too_many_arguments)]
13697 fn gemma4_decode_attn_dc(
13698 &self,
13699 e: &Engine,
13700 fa: &crate::hybrid::FullAttnLayer,
13701 il: usize,
13702 hq: &CudaSlice<i8>,
13703 hdq: &CudaSlice<f32>,
13704 pos_d: &CudaSlice<i32>,
13705 cache: &mut Cache,
13706 cap_bucket_max: Option<(usize, usize)>,
13707 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
13708 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
13709 let eps = self.cfg.rms_eps;
13710 let aux = self.gemma4_aux.as_ref().unwrap();
13711 let ones = aux.ones(e);
13712 #[cfg(debug_assertions)]
13713 crate::debug_assert_tensor_stream_device(ones, &e.stream(), "gemma4_decode_attn_dc.ones");
13714 let (q0, k0, v0) = if swa {
13715 match e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)? {
13716 Some(t3) => t3,
13717 None => match e.matmul_nvfp4_fused2(&fa.wq, &fa.wk, hq, hdq, 1)? {
13719 Some((q0, k0)) => {
13720 let h0 = e.zeros(0)?;
13721 let v0 = e.matmul_pre(&fa.wv, hq, hdq, &h0, 1)?;
13722 (q0, k0, v0)
13723 }
13724 None => {
13725 let h0 = e.zeros(0)?;
13726 (
13727 e.matmul_pre(&fa.wq, hq, hdq, &h0, 1)?,
13728 e.matmul_pre(&fa.wk, hq, hdq, &h0, 1)?,
13729 e.matmul_pre(&fa.wv, hq, hdq, &h0, 1)?,
13730 )
13731 }
13732 },
13733 }
13734 } else {
13735 let (q0, k0) = match e.matmul_q4_fused2(&fa.wq, &fa.wk, hq, hdq)? {
13736 Some(p) => p,
13737 None => match e.matmul_nvfp4_fused2(&fa.wq, &fa.wk, hq, hdq, 1)? {
13738 Some(p) => p,
13739 None => {
13740 let h0 = e.zeros(0)?;
13741 (
13742 e.matmul_pre(&fa.wq, hq, hdq, &h0, 1)?,
13743 e.matmul_pre(&fa.wk, hq, hdq, &h0, 1)?,
13744 )
13745 }
13746 },
13747 };
13748 let v0 = e.clone_dtod(&k0)?;
13749 (q0, k0, v0)
13750 };
13751 let mut q = e.uninit(nh * hd)?;
13752 let mut k = e.uninit(nkv * hd)?;
13753 let mut v = e.uninit(nkv * hd)?;
13754 let ff = if swa {
13756 None
13757 } else {
13758 Some(
13759 aux.rope_freqs(e)
13760 .expect("gemma4 global rope needs rope_freqs.weight"),
13761 )
13762 };
13763 #[cfg(debug_assertions)]
13764 if let Some(ff) = ff {
13765 crate::debug_assert_tensor_stream_device(
13766 ff,
13767 &e.stream(),
13768 "gemma4_decode_attn_dc.rope_freqs",
13769 );
13770 }
13771 let kvl = cache.kv[il].as_mut().unwrap();
13772 let kv_fp8 = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
13773 if crate::Engine::qkv_append_on() {
13774 e.rms_norm_qkv_rope_append_dc(
13776 &q0,
13777 &k0,
13778 &v0,
13779 fa.q_norm.float_data(),
13780 fa.k_norm.float_data(),
13781 ones,
13782 &mut q,
13783 &mut k,
13784 &mut v,
13785 hd,
13786 self.gemma4_rope_dims(il),
13787 nh,
13788 nkv,
13789 pos_d,
13790 nh,
13791 nkv,
13792 base,
13793 1.0,
13794 ff,
13795 eps,
13796 &mut kvl.k,
13797 &mut kvl.v,
13798 &kvl.len_d,
13799 kvl.k_tok_bytes,
13800 kvl.v_tok_bytes,
13801 kv_fp8,
13802 )?;
13803 } else {
13804 e.rms_norm_qkv_rope(
13805 &q0,
13806 &k0,
13807 &v0,
13808 fa.q_norm.float_data(),
13809 fa.k_norm.float_data(),
13810 ones,
13811 &mut q,
13812 &mut k,
13813 &mut v,
13814 hd,
13815 self.gemma4_rope_dims(il),
13816 nh,
13817 nkv,
13818 pos_d,
13819 nh,
13820 nkv,
13821 base,
13822 1.0,
13823 ff,
13824 eps,
13825 )?;
13826 e.append_kv_quantized_dc(
13827 &k,
13828 &v,
13829 &mut kvl.k,
13830 &mut kvl.v,
13831 &kvl.len_d,
13832 kvl.kv_dim_k,
13833 kvl.kv_dim_v,
13834 kvl.k_tok_bytes,
13835 kvl.v_tok_bytes,
13836 kv_fp8,
13837 )?;
13838 }
13839 e.inc_seqlen(&mut kvl.len_d)?;
13840 let mut attn = e.uninit(nh * hd)?;
13841 let mut fa_q8: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
13844 match cap_bucket_max {
13849 None => {
13850 kvl.len += 1;
13854 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
13855 if !swa
13856 && hd == 512
13857 && kvl.len >= crate::fa512_min_tkv()
13858 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
13859 {
13860 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
13863 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
13864 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
13865 e.fa_decode_rows(
13866 &q,
13867 &kp,
13868 &vp,
13869 &mut attn,
13870 hd,
13871 nh,
13872 nkv,
13873 kvl.len - 1,
13874 1,
13875 scale,
13876 kvl.k_tok_bytes,
13877 kvl.v_tok_bytes,
13878 Some((&kvl.len_d, -1)),
13879 false,
13880 false,
13881 Some((&mut aq8, &mut ad8)),
13882 )?;
13883 fa_q8 = Some((aq8, ad8));
13884 } else if swa
13885 && kvl.len > win
13886 && hd == 256
13887 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
13888 {
13889 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
13891 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
13892 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
13893 e.fa_decode_rows_w(
13894 &q,
13895 &kp,
13896 &vp,
13897 &mut attn,
13898 hd,
13899 nh,
13900 nkv,
13901 &kvl.len_d,
13902 -1,
13903 1,
13904 scale,
13905 win,
13906 kvl.k_tok_bytes,
13907 kvl.v_tok_bytes,
13908 Some((&mut aq8, &mut ad8)),
13909 )?;
13910 fa_q8 = Some((aq8, ad8));
13911 } else {
13912 let (off_tok, t_kv) = if swa && kvl.len > win {
13913 (kvl.len - win, win)
13914 } else {
13915 (0, kvl.len)
13916 };
13917 let k_view = e.view_u8_range(
13918 &kvl.k,
13919 off_tok * kvl.k_tok_bytes,
13920 (off_tok + t_kv) * kvl.k_tok_bytes,
13921 );
13922 let v_view = e.view_u8_range(
13923 &kvl.v,
13924 off_tok * kvl.v_tok_bytes,
13925 (off_tok + t_kv) * kvl.v_tok_bytes,
13926 );
13927 e.fa_decode_kvmod(
13928 &q,
13929 &k_view,
13930 &v_view,
13931 &mut attn,
13932 hd,
13933 nh,
13934 nkv,
13935 t_kv,
13936 scale,
13937 kvl.k_tok_bytes,
13938 kvl.v_tok_bytes,
13939 swa && crate::Engine::wkv_on(),
13940 )?;
13941 }
13942 }
13943 Some((b_swa, b_glob)) => {
13944 let k_view = e.view_u8(&kvl.k, kvl.k.len());
13950 let v_view = e.view_u8(&kvl.v, kvl.v.len());
13951 let rows_on = std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0");
13952 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
13953 if !swa && hd == 512 && b_glob >= crate::fa512_min_tkv() && rows_on {
13954 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
13955 e.fa_decode_rows(
13956 &q,
13957 &k_view,
13958 &v_view,
13959 &mut attn,
13960 hd,
13961 nh,
13962 nkv,
13963 b_glob - 1,
13964 1,
13965 scale,
13966 kvl.k_tok_bytes,
13967 kvl.v_tok_bytes,
13968 Some((&kvl.len_d, -1)),
13969 false,
13970 false,
13971 Some((&mut aq8, &mut ad8)),
13972 )?;
13973 fa_q8 = Some((aq8, ad8));
13974 } else if swa && b_swa > win && hd == 256 && rows_on {
13975 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
13976 e.fa_decode_rows_w(
13977 &q,
13978 &k_view,
13979 &v_view,
13980 &mut attn,
13981 hd,
13982 nh,
13983 nkv,
13984 &kvl.len_d,
13985 -1,
13986 1,
13987 scale,
13988 win,
13989 kvl.k_tok_bytes,
13990 kvl.v_tok_bytes,
13991 Some((&mut aq8, &mut ad8)),
13992 )?;
13993 fa_q8 = Some((aq8, ad8));
13994 } else {
13995 let b = if swa { b_swa } else { b_glob };
13996 e.fa_decode_dc(
13997 &q,
13998 &k_view,
13999 &v_view,
14000 &mut attn,
14001 hd,
14002 nh,
14003 nkv,
14004 &kvl.len_d,
14005 b,
14006 scale,
14007 kvl.k_tok_bytes,
14008 kvl.v_tok_bytes,
14009 swa && crate::Engine::wkv_on(),
14010 )?;
14011 }
14012 }
14013 }
14014 if let Some((aq8, ad8)) = fa_q8 {
14017 let mut y = e.uninit(fa.wo.out_features())?;
14018 self.g4_matvec_m1_into(e, &fa.wo, &aq8, &ad8, &mut y)?;
14019 return Ok(y);
14020 }
14021 Ok(e.matmul(&fa.wo, &attn, 1)?)
14022 }
14023
14024 pub fn gemma4_generate_graph(
14029 &self,
14030 e: &Engine,
14031 prompt_pos: usize,
14032 first_token: u32,
14033 cache: &mut Cache,
14034 max_new: usize,
14035 eos: &[u32],
14036 mut on_token: impl FnMut(u32) -> bool,
14037 ) -> Result<(Vec<u32>, crate::decode::StopReason), Box<dyn std::error::Error>> {
14038 if self.is_gemma4_e4b() {
14039 return Err(
14040 "E4B graph serving is unwired (HANDOVER-E4B.md) — dc-eager is the serving arm"
14041 .into(),
14042 );
14043 }
14044 use crate::decode::StopReason;
14045 let n_vocab = self.output.out_features();
14046 let n_embd = self.cfg.n_embd as usize;
14047 let embd_gpu = self
14048 .embd_gpu
14049 .get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload"));
14050 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
14051 for kvl in cache.kv.iter_mut().flatten() {
14052 e.set_i32_one(&mut kvl.len_d, kvl.len as i32)?;
14053 }
14054 let mut token_d = e.stream().clone_htod(&[first_token])?;
14055 let mut pos_d = e.htod_i32(&[prompt_pos as i32])?;
14056 let g4 = self.cfg.gemma4.as_ref().unwrap();
14057 let (hd_s, hd_g) = (g4.key_length_swa as usize, g4.key_length_global as usize);
14058 let nkv_s = g4
14060 .head_count_kv
14061 .iter()
14062 .zip(g4.swa_pattern.iter())
14063 .find(|p| *p.1)
14064 .map(|p| *p.0 as usize)
14065 .unwrap_or(8);
14066 let nkv_g = g4
14067 .head_count_kv
14068 .iter()
14069 .zip(g4.swa_pattern.iter())
14070 .find(|p| !*p.1)
14071 .map(|p| *p.0 as usize)
14072 .unwrap_or(2);
14073 let mut graphs: std::collections::HashMap<
14074 ((bool, usize), (bool, usize), bool, bool),
14075 (
14076 cudarc::driver::CudaGraph,
14077 Vec<Box<dyn std::any::Any + Send>>,
14078 ),
14079 > = Default::default();
14080 let mut slots = self.g4_dc_slots(e)?;
14083 const RING: usize = 64;
14086 const DRAIN: usize = 1;
14092 let mut ring = e.stream().alloc_zeros::<u32>(RING)?;
14093 let ring_base = prompt_pos;
14094 let mut out = Vec::with_capacity(max_new);
14095 let mut reason = StopReason::MaxNew;
14096 let mut next = first_token;
14097 let mut captures = 0usize;
14098 for _ in 0..max_new {
14099 out.push(next);
14100 if eos.contains(&next) {
14101 reason = StopReason::Eos;
14102 break;
14103 }
14104 if !on_token(next) {
14105 reason = StopReason::Callback;
14106 break;
14107 }
14108 let t_kv = cache.pos + 1;
14109 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
14117 let f512 = crate::fa512_min_tkv();
14118 let key_s = if t_kv > win {
14119 (true, usize::MAX)
14120 } else {
14121 e.fa_bucket_key(t_kv, hd_s, nkv_s, crate::Engine::wkv_on())
14122 };
14123 let (key_g, rung_end) = if t_kv >= f512 {
14124 let end = (t_kv + 1).next_power_of_two().max(f512 * 2);
14127 ((true, end), end)
14128 } else {
14129 (e.fa_bucket_key(t_kv, hd_g, nkv_g, false), t_kv)
14130 };
14131 let key = (key_s, key_g, t_kv >= f512, t_kv > win);
14132 if !graphs.contains_key(&key) {
14133 let bucket_max = (t_kv, rung_end);
14134 let snap = cache.snapshot(e)?;
14136 let pos_save = e.dtoh_i32_one(&pos_d)?;
14137 let len_save: Vec<Option<i32>> = cache
14138 .kv
14139 .iter()
14140 .map(|k| k.as_ref().map(|kvl| e.dtoh_i32_one(&kvl.len_d).unwrap()))
14141 .collect();
14142 let tok_save = e.dtoh_u32_one(&token_d)?;
14143 let graph = {
14148 let tok_ref = &mut token_d;
14149 let pos_ref = &mut pos_d;
14150 let cache_ref = &mut *cache;
14151 let slots_ref = &mut slots;
14152 let ring_ref = &mut ring;
14153 e.capture_graph_retained_flags(
14154 cudarc::driver::sys::CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_USE_NODE_PRIORITY,
14155 |e| {
14156 let tok_in = unsafe { &*(tok_ref as *const CudaSlice<u32>) };
14158 let sl = unsafe { &mut *(slots_ref as *mut G4DcSlots) };
14159 let rg = unsafe { &mut *(ring_ref as *mut CudaSlice<u32>) };
14160 self.gemma4_decode_step_dc_slotted(e, tok_in, pos_ref, embd_gpu, qt, rb,
14161 cache_ref, n_vocab, Some(bucket_max),
14162 sl, tok_ref, Some((rg, ring_base)))
14163 })?
14164 };
14165 cache.rollback(e, &snap, 0)?;
14166 e.set_i32_one(&mut pos_d, pos_save)?;
14167 for (il, ls) in len_save.iter().enumerate() {
14168 if let (Some(kvl), Some(v)) = (cache.kv[il].as_mut(), ls) {
14169 e.set_i32_one(&mut kvl.len_d, *v)?;
14170 }
14171 }
14172 e.set_u32_one(&mut token_d, tok_save)?;
14173 if std::env::var("MEMRA_GRAPH_CENSUS").as_deref() == Ok("1") {
14174 if let Ok(c) = crate::graph_update::node_census(&graph.0) {
14175 eprintln!("[graph-census] {c:?}");
14176 }
14177 }
14178 graphs.insert(key, graph);
14179 captures += 1;
14180 }
14181 let mut chunk = 1usize;
14186 let drain_cap: usize = std::env::var("MEMRA_GRAPH_DRAIN")
14187 .ok()
14188 .and_then(|v| v.parse().ok())
14189 .unwrap_or(DRAIN);
14190 while chunk < drain_cap && out.len() + chunk < max_new {
14191 let t_next = cache.pos + 1 + chunk;
14192 let key_s2 = if t_next > win {
14193 (true, usize::MAX)
14194 } else {
14195 e.fa_bucket_key(t_next, hd_s, nkv_s, crate::Engine::wkv_on())
14196 };
14197 let key_g2 = if t_next >= f512 {
14198 (true, (t_next + 1).next_power_of_two().max(f512 * 2))
14199 } else {
14200 e.fa_bucket_key(t_next, hd_g, nkv_g, false)
14201 };
14202 if (key_s2, key_g2, t_next >= f512, t_next > win) != key {
14203 break;
14204 }
14205 chunk += 1;
14206 }
14207 let g = &graphs.get(&key).unwrap().0;
14208 for _ in 0..chunk {
14209 g.launch()?;
14210 }
14211 e.stream().synchronize()?;
14212 let ringh = e.dtoh_u32(&ring)?;
14213 for j in 0..chunk {
14214 let pos_j = cache.pos + j;
14215 let tok_j = ringh[(pos_j - ring_base) % RING];
14216 cache.pos += 0; if j + 1 == chunk {
14218 next = tok_j;
14219 } else {
14220 out.push(tok_j);
14221 if eos.contains(&tok_j) || !on_token(tok_j) {
14222 reason = if eos.contains(&tok_j) {
14223 StopReason::Eos
14224 } else {
14225 StopReason::Callback
14226 };
14227 let keep = cache.pos + j + 1;
14229 e.set_i32_one(&mut pos_d, keep as i32)?;
14230 for kvl in cache.kv.iter_mut().filter_map(|k| k.as_mut()) {
14231 e.set_i32_one(&mut kvl.len_d, keep as i32)?;
14232 kvl.len = keep;
14233 }
14234 cache.pos = keep;
14235 if std::env::var("MEMRA_GRAPH_STATS").is_ok() {
14236 eprintln!("[gemma-graph] captures={captures} buckets={}", graphs.len());
14237 }
14238 return Ok((out, reason));
14239 }
14240 }
14241 }
14242 cache.pos += chunk;
14243 for kvl in cache.kv.iter_mut().filter_map(|k| k.as_mut()) {
14244 kvl.len += chunk;
14245 }
14246 }
14247 if std::env::var("MEMRA_GRAPH_STATS").is_ok() {
14248 eprintln!("[gemma-graph] captures={captures} buckets={}", graphs.len());
14249 }
14250 Ok((out, reason))
14251 }
14252
14253 pub(crate) fn gemma4_decode_step_t(
14259 &self,
14260 e: &Engine,
14261 tokens: &[u32],
14262 pos0: usize,
14263 cache: &mut Cache,
14264 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
14265 Ok(self.gemma4_decode_step_t_h(e, tokens, pos0, cache)?.0)
14266 }
14267
14268 pub(crate) fn gemma4_decode_step_t_am(
14272 &self,
14273 e: &Engine,
14274 tokens: &[u32],
14275 pos0: usize,
14276 cache: &mut Cache,
14277 ) -> Result<(Vec<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
14278 let (ld, hn) = self.gemma4_verify_trunk(e, tokens, pos0, cache, None)?;
14279 let t = tokens.len();
14280 let n_vocab = self.output.out_features();
14281 let mut toks = e.stream().alloc_zeros::<u32>(t)?;
14282 for i in 0..t {
14283 e.argmax_token_device_col(&ld, i, n_vocab, &mut toks, i)?;
14284 }
14285 Ok((e.dtoh_u32(&toks)?, hn))
14286 }
14287
14288 pub(crate) fn gemma4_decode_step_t_am_dev(
14291 &self,
14292 e: &Engine,
14293 tok_d: &CudaSlice<u32>,
14294 t: usize,
14295 pos0: usize,
14296 cache: &mut Cache,
14297 ) -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
14298 let (ld, hn) = self.gemma4_verify_trunk(e, &vec![0u32; t], pos0, cache, Some(tok_d))?;
14299 let n_vocab = self.output.out_features();
14300 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
14301 for i in 0..t {
14302 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
14303 }
14304 Ok((vam, hn))
14305 }
14306
14307 pub(crate) fn gemma4_decode_step_t_h(
14310 &self,
14311 e: &Engine,
14312 tokens: &[u32],
14313 pos0: usize,
14314 cache: &mut Cache,
14315 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
14316 let (mut ld, hn) = self.gemma4_verify_trunk(e, tokens, pos0, cache, None)?;
14317 let t = tokens.len();
14318 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
14319 e.softcap(&mut ld, cap, t * self.output.out_features())?;
14320 Ok((e.dtoh(&ld)?, hn))
14321 }
14322
14323 pub(crate) fn verify_stream_scratch(
14326 &self,
14327 e: &Engine,
14328 cap: usize,
14329 ) -> Result<VerifyStreamScratch, Box<dyn std::error::Error>> {
14330 Ok(VerifyStreamScratch {
14331 pos_d: e.htod_i32(&vec![0i32; cap])?,
14332 row_ctrs: (0..cap)
14333 .map(|_| e.htod_i32(&[0]))
14334 .collect::<Result<_, _>>()?,
14335 })
14336 }
14337
14338 pub(crate) fn gemma4_verify_t_am_stream(
14346 &self,
14347 e: &Engine,
14348 tok_d: &CudaSlice<u32>,
14349 t: usize,
14350 ctr: &CudaSlice<i32>,
14351 hint: usize,
14352 cache: &mut Cache,
14353 scr: &mut VerifyStreamScratch,
14354 ) -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
14355 let n_embd = self.cfg.n_embd as usize;
14356 let eps = self.cfg.rms_eps;
14357 assert!(t <= scr.row_ctrs.len() && t <= 64);
14358 e.i32_iota_from(ctr, &mut scr.pos_d, t)?;
14359 for i in 0..t {
14360 e.i32_copy_add(ctr, &mut scr.row_ctrs[i], (i + 1) as i32)?;
14361 }
14362 let (pos_d, row_ctrs) = (&scr.pos_d, &scr.row_ctrs);
14363 let embd_gpu = self
14364 .embd_gpu
14365 .get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload"));
14366 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
14367 let mut x = e.embed_gather_device_td(embd_gpu, tok_d, t, n_embd, qt, rb)?;
14368 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
14369 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
14370 let n_layers = self.layers.len();
14371 for (il, layer) in self.layers.iter().enumerate() {
14372 let (hq, hdq) = match h_carry.take() {
14373 Some(p) => p,
14374 None => {
14375 e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, t, eps)?
14376 }
14377 };
14378 let Mixer::Full(fa) = &layer.mixer else {
14379 panic!("gemma4 layer {il} not full-attn")
14380 };
14381 let o = self
14382 .gemma4_verify_attn_stream(e, fa, il, &hq, &hdq, pos_d, t, cache, hint, row_ctrs)?;
14383 let next_norm = if il + 1 < n_layers {
14384 Some(self.layers[il + 1].attn_norm.float_data())
14385 } else {
14386 None
14387 };
14388 let (xn, hn) = self.gemma4_layer_tail_add_nq_pn(e, layer, &o, &x, t, next_norm)?;
14389 x = xn;
14390 h_carry = hn;
14391 self.dflash_tap(e, cache, il, &x, t)?;
14392 }
14393 let mut hn = e.uninit(t * n_embd)?;
14394 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
14395 let ld = e.matmul(&self.output, &hn, t)?;
14396 let n_vocab = self.output.out_features();
14397 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
14398 for i in 0..t {
14399 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
14400 }
14401 Ok((vam, hn))
14402 }
14403
14404 pub(crate) fn dflash_tap(
14411 &self,
14412 e: &Engine,
14413 cache: &mut Cache,
14414 il: usize,
14415 x: &CudaSlice<f32>,
14416 t: usize,
14417 ) -> Result<(), Box<dyn std::error::Error>> {
14418 let Some(taps) = cache.dflash_taps.as_mut() else {
14419 return Ok(());
14420 };
14421 let Some(slot) = taps.layer_ids.iter().position(|&l| l == il) else {
14422 return Ok(());
14423 };
14424 let h = taps.hidden;
14425 let n_taps = taps.layer_ids.len();
14426 let base = taps.base;
14427 debug_assert!(
14428 base + t <= taps.t,
14429 "tap window {base}+{t} exceeds sink {}",
14430 taps.t
14431 );
14432 let xv = e.view(x, t * h);
14433 for r in 0..t {
14434 let row = xv.slice(r * h..(r + 1) * h);
14435 e.copy_view_into(&mut taps.buf, (base + r) * n_taps * h + slot * h, &row, h)?;
14436 }
14437 Ok(())
14438 }
14439
14440 fn gemma4_verify_trunk(
14441 &self,
14442 e: &Engine,
14443 tokens: &[u32],
14444 pos0: usize,
14445 cache: &mut Cache,
14446 tok_dev: Option<&CudaSlice<u32>>,
14447 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
14448 let n_embd = self.cfg.n_embd as usize;
14449 let eps = self.cfg.rms_eps;
14450 let t = tokens.len();
14451 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
14452 let pos_d = e.htod_i32(&pos)?;
14453 let mut x = match tok_dev {
14454 Some(td) => {
14455 let embd_gpu = self
14456 .embd_gpu
14457 .get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload"));
14458 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
14459 e.embed_gather_device_td(embd_gpu, td, t, n_embd, qt, rb)?
14460 }
14461 None => e.htod(&self.embd.gather(n_embd, tokens))?,
14462 };
14463 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
14464 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
14465 let n_layers = self.layers.len();
14466 for (il, layer) in self.layers.iter().enumerate() {
14467 let (hq, hdq) = match h_carry.take() {
14468 Some(p) => p,
14469 None => {
14470 e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, t, eps)?
14471 }
14472 };
14473 let Mixer::Full(fa) = &layer.mixer else {
14474 panic!("gemma4 layer {il} not full-attn")
14475 };
14476 let o = self.gemma4_verify_attn(e, fa, il, &hq, &hdq, &pos_d, t, cache)?;
14477 let next_norm = if il + 1 < n_layers {
14478 Some(self.layers[il + 1].attn_norm.float_data())
14479 } else {
14480 None
14481 };
14482 let (xn, hn) = self.gemma4_layer_tail_add_nq_pn(e, layer, &o, &x, t, next_norm)?;
14483 x = xn;
14484 h_carry = hn;
14485 self.dflash_tap(e, cache, il, &x, t)?;
14486 }
14487 let mut hn = e.uninit(t * n_embd)?;
14488 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
14489 let mut ld = e.matmul(&self.output, &hn, t)?;
14490 self.gemma4_suppress(e, &mut ld, t)?; cache.pos += t;
14492 Ok((ld, hn))
14493 }
14494
14495 #[allow(clippy::too_many_arguments)]
14503 fn gemma4_verify_attn_stream(
14504 &self,
14505 e: &Engine,
14506 fa: &crate::hybrid::FullAttnLayer,
14507 il: usize,
14508 hq: &CudaSlice<i8>,
14509 hdq: &CudaSlice<f32>,
14510 pos_d: &CudaSlice<i32>,
14511 t: usize,
14512 cache: &mut Cache,
14513 hint: usize,
14514 row_ctrs: &[CudaSlice<i32>],
14515 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
14516 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
14517 let eps = self.cfg.rms_eps;
14518 let aux = self.gemma4_aux.as_ref().unwrap();
14519 let ones = aux.ones(e);
14520 #[cfg(debug_assertions)]
14521 crate::debug_assert_tensor_stream_device(
14522 ones,
14523 &e.stream(),
14524 "gemma4_verify_attn_stream.ones",
14525 );
14526 let h0 = e.zeros(0)?;
14527 let h = &h0;
14528 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
14531 let f2b = *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
14532 let fused_qkv = if f2b {
14533 if swa {
14534 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
14535 .map(|(a, b, c)| (a, b, Some(c)))
14536 } else {
14537 e.matmul_q4_fused2_batched(&fa.wq, &fa.wk, hq, hdq, t)?
14538 .map(|(a, b)| (a, b, None))
14539 }
14540 } else {
14541 None
14542 };
14543 let (q0, k0, v0) = match fused_qkv {
14544 Some((a, b, cv)) => {
14545 let v = match cv {
14546 Some(c) => c,
14547 None => e.clone_dtod(&b)?,
14548 };
14549 (a, b, v)
14550 }
14551 None => {
14552 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
14553 let k0 = e.matmul_pre(&fa.wk, hq, hdq, h, t)?;
14554 let v0 = if swa {
14555 e.matmul_pre(&fa.wv, hq, hdq, h, t)?
14556 } else {
14557 e.clone_dtod(&k0)?
14558 };
14559 (q0, k0, v0)
14560 }
14561 };
14562 let mut q = e.uninit(t * nh * hd)?;
14563 let mut k = e.uninit(t * nkv * hd)?;
14564 let mut v = e.uninit(t * nkv * hd)?;
14565 let ff = if swa {
14568 None
14569 } else {
14570 Some(
14571 aux.rope_freqs(e)
14572 .expect("gemma4 global rope needs rope_freqs.weight"),
14573 )
14574 };
14575 #[cfg(debug_assertions)]
14576 if let Some(ff) = ff {
14577 crate::debug_assert_tensor_stream_device(
14578 ff,
14579 &e.stream(),
14580 "gemma4_verify_attn_stream.rope_freqs",
14581 );
14582 }
14583 e.rms_norm_qkv_rope(
14584 &q0,
14585 &k0,
14586 &v0,
14587 fa.q_norm.float_data(),
14588 fa.k_norm.float_data(),
14589 ones,
14590 &mut q,
14591 &mut k,
14592 &mut v,
14593 hd,
14594 self.gemma4_rope_dims(il),
14595 nh * t,
14596 nkv * t,
14597 pos_d,
14598 nh,
14599 nkv,
14600 base,
14601 1.0,
14602 ff,
14603 eps,
14604 )?;
14605 let kvl = cache.kv[il].as_mut().unwrap();
14606 e.append_kv_quantized_rows_dc(
14608 &k,
14609 &v,
14610 &mut kvl.k,
14611 &mut kvl.v,
14612 &kvl.len_d,
14613 t,
14614 kvl.kv_dim_k,
14615 kvl.kv_dim_v,
14616 kvl.k_tok_bytes,
14617 kvl.v_tok_bytes,
14618 (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on()),
14619 )?;
14620 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
14623 let mut attn = e.uninit(t * nh * hd)?;
14624 let k_view = e.view_u8(&kvl.k, kvl.k.len());
14625 let v_view = e.view_u8(&kvl.v, kvl.v.len());
14626 if swa && hint + 1 >= win {
14629 e.fa_decode_rows_w(
14632 &q,
14633 &k_view,
14634 &v_view,
14635 &mut attn,
14636 hd,
14637 nh,
14638 nkv,
14639 &kvl.len_d,
14640 0,
14641 t,
14642 scale,
14643 win,
14644 kvl.k_tok_bytes,
14645 kvl.v_tok_bytes,
14646 None,
14647 )?;
14648 } else if hd == 512 && hint + t < crate::fa512_min_tkv() {
14649 let bucket = (hint + t + 2)
14662 .next_power_of_two()
14663 .min(crate::fa512_min_tkv().saturating_sub(1));
14664 let qv = e.view(&q, t * nh * hd);
14665 for i in 0..t {
14666 let q_row = qv.slice(i * nh * hd..(i + 1) * nh * hd);
14667 let mut q_one = e.uninit(nh * hd)?;
14668 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
14669 let mut a_one = e.uninit(nh * hd)?;
14670 e.fa_decode_dc(
14671 &q_one,
14672 &k_view,
14673 &v_view,
14674 &mut a_one,
14675 hd,
14676 nh,
14677 nkv,
14678 &row_ctrs[i],
14679 bucket,
14680 scale,
14681 kvl.k_tok_bytes,
14682 kvl.v_tok_bytes,
14683 false,
14684 )?;
14685 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
14686 }
14687 } else if hd == 512 {
14688 e.fa_decode_rows(
14691 &q,
14692 &k_view,
14693 &v_view,
14694 &mut attn,
14695 hd,
14696 nh,
14697 nkv,
14698 hint,
14699 t,
14700 scale,
14701 kvl.k_tok_bytes,
14702 kvl.v_tok_bytes,
14703 Some((&kvl.len_d, 0)),
14704 false,
14705 false,
14706 None,
14707 )?;
14708 } else {
14709 e.fa_decode_rows_dc(
14711 &q,
14712 &k_view,
14713 &v_view,
14714 &mut attn,
14715 hd,
14716 nh,
14717 nkv,
14718 &kvl.len_d,
14719 hint + t,
14720 t,
14721 scale,
14722 kvl.k_tok_bytes,
14723 kvl.v_tok_bytes,
14724 0,
14725 swa && crate::Engine::wkv_on(),
14726 )?;
14727 }
14728 Ok(e.matmul(&fa.wo, &attn, t)?)
14729 }
14730
14731 fn gemma4_verify_attn(
14732 &self,
14733 e: &Engine,
14734 fa: &crate::hybrid::FullAttnLayer,
14735 il: usize,
14736 hq: &CudaSlice<i8>,
14737 hdq: &CudaSlice<f32>,
14738 pos_d: &CudaSlice<i32>,
14739 t: usize,
14740 cache: &mut Cache,
14741 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
14742 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
14743 let eps = self.cfg.rms_eps;
14744 let aux = self.gemma4_aux.as_ref().unwrap();
14745 let ones = aux.ones(e);
14746 #[cfg(debug_assertions)]
14747 crate::debug_assert_tensor_stream_device(ones, &e.stream(), "gemma4_verify_attn.ones");
14748 let n_embd = self.cfg.n_embd as usize;
14749 let _ = n_embd;
14750
14751 let h0 = e.zeros(0)?;
14752 let h = &h0;
14753 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
14756 let f2b = *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
14757 let fused_qkv = if f2b {
14758 if swa {
14759 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
14760 .map(|(a, b, c)| (a, b, Some(c)))
14761 } else {
14762 e.matmul_q4_fused2_batched(&fa.wq, &fa.wk, hq, hdq, t)?
14763 .map(|(a, b)| (a, b, None))
14764 }
14765 } else {
14766 None
14767 };
14768 let (q0, k0, v0) = match fused_qkv {
14769 Some((a, b, cv)) => {
14770 let v = match cv {
14771 Some(c) => c,
14772 None => e.clone_dtod(&b)?,
14773 };
14774 (a, b, v)
14775 }
14776 None => {
14777 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
14778 let k0 = e.matmul_pre(&fa.wk, hq, hdq, h, t)?;
14779 let v0 = if swa {
14780 e.matmul_pre(&fa.wv, hq, hdq, h, t)?
14781 } else {
14782 e.clone_dtod(&k0)?
14783 };
14784 (q0, k0, v0)
14785 }
14786 };
14787 let mut q = e.uninit(t * nh * hd)?;
14788 let mut k = e.uninit(t * nkv * hd)?;
14789 let mut v = e.uninit(t * nkv * hd)?;
14790 let ff = if swa {
14793 None
14794 } else {
14795 Some(
14796 aux.rope_freqs(e)
14797 .expect("gemma4 global rope needs rope_freqs.weight"),
14798 )
14799 };
14800 #[cfg(debug_assertions)]
14801 if let Some(ff) = ff {
14802 crate::debug_assert_tensor_stream_device(
14803 ff,
14804 &e.stream(),
14805 "gemma4_verify_attn.rope_freqs",
14806 );
14807 }
14808 e.rms_norm_qkv_rope(
14809 &q0,
14810 &k0,
14811 &v0,
14812 fa.q_norm.float_data(),
14813 fa.k_norm.float_data(),
14814 ones,
14815 &mut q,
14816 &mut k,
14817 &mut v,
14818 hd,
14819 self.gemma4_rope_dims(il),
14820 nh * t,
14821 nkv * t,
14822 pos_d,
14823 nh,
14824 nkv,
14825 base,
14826 1.0,
14827 ff,
14828 eps,
14829 )?;
14830 let kvl = cache.kv[il].as_mut().unwrap();
14831 let base_len = kvl.len;
14832 e.append_kv_quantized_rows(
14833 &k,
14834 &v,
14835 &mut kvl.k,
14836 &mut kvl.v,
14837 base_len,
14838 t,
14839 kvl.kv_dim_k,
14840 kvl.kv_dim_v,
14841 kvl.k_tok_bytes,
14842 kvl.v_tok_bytes,
14843 (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on()),
14844 )?;
14845 kvl.len += t;
14846 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
14847 let mut attn = e.uninit(t * nh * hd)?;
14848 let rows_ok = (hd == 256 && base_len + 1 >= crate::fa_vec_min_tkv())
14851 || (hd == 512 && !swa && base_len + 1 >= crate::fa512_min_tkv());
14854 if rows_ok && (!swa || base_len + t <= win) {
14855 let k_view = e.view_u8(&kvl.k, (base_len + t) * kvl.k_tok_bytes);
14856 let v_view = e.view_u8(&kvl.v, (base_len + t) * kvl.v_tok_bytes);
14857 if hd == 512 {
14858 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
14860 e.fa_decode_rows(
14861 &q,
14862 &k_view,
14863 &v_view,
14864 &mut attn,
14865 hd,
14866 nh,
14867 nkv,
14868 base_len,
14869 t,
14870 scale,
14871 kvl.k_tok_bytes,
14872 kvl.v_tok_bytes,
14873 Some((&kvl.len_d, 0)),
14874 false,
14875 swa && crate::Engine::wkv_on(),
14876 None,
14877 )?;
14878 } else {
14879 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
14883 e.fa_decode_rows_dc(
14884 &q,
14885 &k_view,
14886 &v_view,
14887 &mut attn,
14888 hd,
14889 nh,
14890 nkv,
14891 &kvl.len_d,
14892 base_len + t,
14893 t,
14894 scale,
14895 kvl.k_tok_bytes,
14896 kvl.v_tok_bytes,
14897 0,
14898 swa && crate::Engine::wkv_on(),
14899 )?;
14900 }
14901 return Ok(e.matmul(&fa.wo, &attn, t)?);
14902 }
14903 if hd == 256
14911 && swa
14912 && base_len + 1 >= win
14913 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
14914 {
14915 let k_view = e.view_u8(&kvl.k, (base_len + t) * kvl.k_tok_bytes);
14916 let v_view = e.view_u8(&kvl.v, (base_len + t) * kvl.v_tok_bytes);
14917 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
14918 e.fa_decode_rows_w(
14919 &q,
14920 &k_view,
14921 &v_view,
14922 &mut attn,
14923 hd,
14924 nh,
14925 nkv,
14926 &kvl.len_d,
14927 0,
14928 t,
14929 scale,
14930 win,
14931 kvl.k_tok_bytes,
14932 kvl.v_tok_bytes,
14933 None,
14934 )?;
14935 return Ok(e.matmul(&fa.wo, &attn, t)?);
14936 }
14937 for i in 0..t {
14938 let avail = base_len + i + 1;
14939 let (off_tok, t_kv) = if swa && avail > win {
14940 (avail - win, win)
14941 } else {
14942 (0, avail)
14943 };
14944 let k_view = e.view_u8_range(
14945 &kvl.k,
14946 off_tok * kvl.k_tok_bytes,
14947 (off_tok + t_kv) * kvl.k_tok_bytes,
14948 );
14949 let v_view = e.view_u8_range(
14950 &kvl.v,
14951 off_tok * kvl.v_tok_bytes,
14952 (off_tok + t_kv) * kvl.v_tok_bytes,
14953 );
14954 let qi = e.view(&q, t * nh * hd);
14955 let q_row = qi.slice(i * nh * hd..(i + 1) * nh * hd);
14956 let mut q_one = e.uninit(nh * hd)?;
14957 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
14958 let mut a_one = e.uninit(nh * hd)?;
14959 if swa
14963 && avail > win
14964 && hd == 256
14965 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
14966 {
14967 let kp = e.view_u8(&kvl.k, avail * kvl.k_tok_bytes);
14968 let vp = e.view_u8(&kvl.v, avail * kvl.v_tok_bytes);
14969 e.i32_set_k(&mut kvl.len_d, (avail - 1) as i32)?;
14970 e.fa_decode_rows_w(
14971 &q_one,
14972 &kp,
14973 &vp,
14974 &mut a_one,
14975 hd,
14976 nh,
14977 nkv,
14978 &kvl.len_d,
14979 0,
14980 1,
14981 scale,
14982 win,
14983 kvl.k_tok_bytes,
14984 kvl.v_tok_bytes,
14985 None,
14986 )?;
14987 } else if !swa
14988 && hd == 512
14989 && avail >= crate::fa512_min_tkv()
14990 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
14991 {
14992 let kp = e.view_u8(&kvl.k, avail * kvl.k_tok_bytes);
14993 let vp = e.view_u8(&kvl.v, avail * kvl.v_tok_bytes);
14994 e.i32_set_k(&mut kvl.len_d, (avail - 1) as i32)?;
14995 e.fa_decode_rows(
14996 &q_one,
14997 &kp,
14998 &vp,
14999 &mut a_one,
15000 hd,
15001 nh,
15002 nkv,
15003 avail - 1,
15004 1,
15005 scale,
15006 kvl.k_tok_bytes,
15007 kvl.v_tok_bytes,
15008 Some((&kvl.len_d, 0)),
15009 false,
15010 false,
15011 None,
15012 )?;
15013 } else {
15014 e.fa_decode_kvmod(
15015 &q_one,
15016 &k_view,
15017 &v_view,
15018 &mut a_one,
15019 hd,
15020 nh,
15021 nkv,
15022 t_kv,
15023 scale,
15024 kvl.k_tok_bytes,
15025 kvl.v_tok_bytes,
15026 swa && crate::Engine::wkv_on(),
15027 )?;
15028 }
15029 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
15030 }
15031 Ok(e.matmul(&fa.wo, &attn, t)?)
15032 }
15033
15034 pub(crate) fn gemma4_decode_step_h(
15037 &self,
15038 e: &Engine,
15039 token: u32,
15040 cache: &mut Cache,
15041 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
15042 if let Some(split) = crate::pp::pp2_split(self.layers.len()) {
15047 return self.gemma4_decode_step_h_pp2(e, token, cache, split);
15048 }
15049 if crate::pp::pp_cuts(self.layers.len()).is_some() {
15050 crate::pp::warn_unwired_once("gemma4 eager decode (N>2)");
15051 }
15052 let n_embd = self.cfg.n_embd as usize;
15053 let eps = self.cfg.rms_eps;
15054 let pos_d = e.htod_i32(&[cache.pos as i32])?;
15055 let mut x = e.htod(&self.embd.gather(n_embd, &[token]))?;
15056 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
15057 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
15060 let n_layers = self.layers.len();
15061 for (il, layer) in self.layers.iter().enumerate() {
15062 let (hq, hdq) = match h_carry.take() {
15063 Some(p) => p,
15064 None => {
15065 e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, 1, eps)?
15066 }
15067 };
15068 let Mixer::Full(fa) = &layer.mixer else {
15069 panic!("gemma4 layer {il} not full-attn")
15070 };
15071 let o = self.gemma4_decode_attn(e, fa, il, &hq, &hdq, &pos_d, cache)?;
15072 let next_norm = if il + 1 < n_layers {
15073 Some(self.layers[il + 1].attn_norm.float_data())
15074 } else {
15075 None
15076 };
15077 let (xn, hn) = self.gemma4_layer_tail_add_nq_pn(e, layer, &o, &x, 1, next_norm)?;
15078 x = xn;
15079 h_carry = hn;
15080 }
15081 let mut hn = e.uninit(n_embd)?;
15082 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
15083 let h_seed = e.clone_dtod(&x)?;
15084 let mut ld = e.matmul(&self.output, &hn, 1)?;
15085 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
15086 e.softcap(&mut ld, cap, self.output.out_features())?; self.gemma4_suppress(e, &mut ld, 1)?;
15088 let logits = e.dtoh(&ld)?;
15089 cache.pos += 1;
15090 Ok((logits, h_seed))
15091 }
15092
15093 fn gemma4_decode_layers(
15101 &self,
15102 e: &Engine,
15103 mut x: CudaSlice<f32>,
15104 lo: usize,
15105 hi: usize,
15106 pos_d: &CudaSlice<i32>,
15107 cache: &mut Cache,
15108 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
15109 let n_embd = self.cfg.n_embd as usize;
15110 let eps = self.cfg.rms_eps;
15111 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
15112 for il in lo..hi {
15113 let layer = &self.layers[il];
15114 let (hq, hdq) = match h_carry.take() {
15115 Some(p) => p,
15116 None => {
15118 e.rms_norm_q8_1(&x, self.layers[il].attn_norm.float_data(), n_embd, 1, eps)?
15119 }
15120 };
15121 let Mixer::Full(fa) = &layer.mixer else {
15122 panic!("gemma4 layer {il} not full-attn")
15123 };
15124 let o = self.gemma4_decode_attn(e, fa, il, &hq, &hdq, pos_d, cache)?;
15125 let next_norm = if il + 1 < hi {
15126 Some(self.layers[il + 1].attn_norm.float_data())
15127 } else {
15128 None
15129 };
15130 let (xn, hn) = self.gemma4_layer_tail_add_nq_pn(e, layer, &o, &x, 1, next_norm)?;
15131 x = xn;
15132 h_carry = hn;
15133 }
15134 Ok(x)
15135 }
15136
15137 fn gemma4_decode_step_h_pp2(
15145 &self,
15146 e: &Engine,
15147 token: u32,
15148 cache: &mut Cache,
15149 split: usize,
15150 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
15151 if crate::pp::pp2_streams_off() {
15152 return self.gemma4_decode_step_h_pp2_samestream(e, token, cache, split);
15153 }
15154 let rt = crate::pp::Pp2Rt::get(e)?;
15155 let e0 = rt.engine(0, e);
15156 let e1 = rt.engine(1, e);
15157 let n_embd = self.cfg.n_embd as usize;
15158 let eps = self.cfg.rms_eps;
15159 let pos = cache.pos as i32;
15160
15161 let slot = {
15163 let _st0 = rt.enter(0);
15164 let pos_d = e0.htod_i32(&[pos])?;
15165 #[cfg(debug_assertions)]
15166 crate::debug_assert_tensor_stream_device(
15167 &pos_d,
15168 &e0.stream(),
15169 "gemma4_decode_step_h_pp2.stage0.pos_d",
15170 );
15171 let mut x = e0.htod(&self.embd.gather(n_embd, &[token]))?;
15172 e0.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
15173 let x = self.gemma4_decode_layers(e0, x, 0, split, &pos_d, cache)?;
15174 rt.tx(0, &x, n_embd)?
15175 };
15176
15177 let _st1 = rt.enter(1);
15179 let pos_d = e1.htod_i32(&[pos])?;
15180 #[cfg(debug_assertions)]
15181 crate::debug_assert_tensor_stream_device(
15182 &pos_d,
15183 &e1.stream(),
15184 "gemma4_decode_step_h_pp2.stage1.pos_d",
15185 );
15186 let x = rt.rx(0, slot, n_embd)?;
15187 let x = self.gemma4_decode_layers(e1, x, split, self.layers.len(), &pos_d, cache)?;
15188
15189 let mut hn = e1.uninit(n_embd)?;
15190 e1.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
15191 let h_seed = e1.clone_dtod(&x)?;
15192 let mut ld = e1.matmul(&self.output, &hn, 1)?;
15193 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
15194 e1.softcap(&mut ld, cap, self.output.out_features())?;
15195 self.gemma4_suppress(e1, &mut ld, 1)?;
15196 let logits = e1.dtoh(&ld)?;
15197 cache.pos += 1;
15198 Ok((logits, h_seed))
15199 }
15200
15201 fn gemma4_decode_step_h_pp2_samestream(
15204 &self,
15205 e: &Engine,
15206 token: u32,
15207 cache: &mut Cache,
15208 split: usize,
15209 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
15210 let n_embd = self.cfg.n_embd as usize;
15211 let eps = self.cfg.rms_eps;
15212 let pos_d = e.htod_i32(&[cache.pos as i32])?;
15213
15214 let mut x = e.htod(&self.embd.gather(n_embd, &[token]))?;
15216 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
15217 let x = self.gemma4_decode_layers(e, x, 0, split, &pos_d, cache)?;
15218
15219 let boundary_tx = e.clone_dtod(&x)?;
15221 let boundary_rx = e.clone_dtod(&boundary_tx)?;
15222
15223 let x =
15225 self.gemma4_decode_layers(e, boundary_rx, split, self.layers.len(), &pos_d, cache)?;
15226
15227 let mut hn = e.uninit(n_embd)?;
15228 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
15229 let h_seed = e.clone_dtod(&x)?;
15230 let mut ld = e.matmul(&self.output, &hn, 1)?;
15231 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
15232 e.softcap(&mut ld, cap, self.output.out_features())?;
15233 self.gemma4_suppress(e, &mut ld, 1)?;
15234 let logits = e.dtoh(&ld)?;
15235 cache.pos += 1;
15236 Ok((logits, h_seed))
15237 }
15238}
15239
15240impl HybridModel {
15259 pub(crate) fn step35_geom(&self, il: usize) -> memra_gguf::config::LayerGeometry {
15262 let geometry = self
15263 .cfg
15264 .layer_geometry(il as u32)
15265 .unwrap_or_else(|| panic!("step35 layer {il} has no geometry-table row"));
15266 debug_assert_eq!(
15267 geometry.attention_gate,
15268 memra_gguf::config::AttentionGateKind::SeparateHead
15269 );
15270 geometry
15271 }
15272
15273 #[allow(clippy::too_many_arguments)]
15333 fn step35_attn_pre_wo(
15334 &self,
15335 e: &Engine,
15336 fa: &FullAttnLayer,
15337 mut g3: Vec<CudaSlice<f32>>,
15338 hg: Option<&CudaSlice<f32>>,
15339 gt_pre: Option<&CudaSlice<f32>>,
15340 pos_d: &CudaSlice<i32>,
15341 t: usize,
15342 cache: Option<&mut Cache>,
15343 il: usize,
15344 seq_end: usize,
15345 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
15346 let geometry = self.step35_geom(il);
15347 let hd = geometry.head_dim_k as usize;
15348 let nkv = geometry.n_head_kv as usize;
15349 let nh = geometry.n_head as usize;
15350 let rbase = geometry.rope_base;
15351 let scale = geometry.attention_scale();
15352 let swa = geometry.window.is_some();
15353 let eps = self.cfg.rms_eps;
15354 let win = geometry.window.unwrap_or(0) as usize;
15355 let n_rot = geometry.n_rot as usize;
15356
15357 let v = g3.pop().unwrap();
15358 let k0 = g3.pop().unwrap();
15359 let q0 = g3.pop().unwrap();
15360
15361 let mut q = e.uninit(t * nh * hd)?;
15365 e.rms_norm(&q0, fa.q_norm.float_data(), &mut q, hd, nh * t, eps)?;
15366 let mut k = e.uninit(t * nkv * hd)?;
15367 e.rms_norm(&k0, fa.k_norm.float_data(), &mut k, hd, nkv * t, eps)?;
15368 let ff = if geometry.rope_factors {
15369 self.step35_aux.as_ref().and_then(|a| a.rope_freqs(e))
15370 } else {
15371 None
15372 };
15373 #[cfg(debug_assertions)]
15374 if let Some(ff) = ff {
15375 crate::debug_assert_tensor_stream_device(
15376 ff,
15377 &e.stream(),
15378 "step35_attn_pre_wo.rope_freqs",
15379 );
15380 }
15381 e.rope_neox2(&mut q, &mut k, pos_d, hd, n_rot, nh, nkv, t, rbase, 1.0, ff)?;
15382
15383 let mut attn = e.uninit(t * nh * hd)?;
15384 match cache {
15385 Some(cache) => {
15386 let base_len = cache.kv[il].as_ref().unwrap().len;
15387 let legacy_tkv = std::env::var("MEMRA_STEP35_SWA_TKV").as_deref() == Ok("1");
15389 let legacy_calllocal = std::env::var("MEMRA_PRIME_CALLLOCAL").as_deref() == Ok("1");
15390 let off = if swa {
15391 let raw = base_len.saturating_sub(win - 1);
15392 if legacy_tkv || legacy_calllocal {
15393 raw
15394 } else {
15395 raw & !31usize
15396 }
15397 } else {
15398 0
15399 };
15400 {
15401 let kvl = cache.kv[il].as_mut().unwrap();
15402 assert!(kvl.len + t <= cache.max_ctx, "step35 prime: KV overflow");
15403 let write_row = e.prepare_kv_append(kvl, off, t)?;
15404 e.append_kv_quantized_rows(
15405 &k,
15406 &v,
15407 &mut kvl.k,
15408 &mut kvl.v,
15409 write_row,
15410 t,
15411 kvl.kv_dim_k,
15412 kvl.kv_dim_v,
15413 kvl.k_tok_bytes,
15414 kvl.v_tok_bytes,
15415 crate::Engine::kv_fp8_on(),
15416 )?;
15417 kvl.len += t;
15418 let new_len = kvl.len as i32;
15419 e.set_i32_one(&mut kvl.len_d, new_len)?;
15420 }
15421 let kvl = cache.kv[il].as_ref().unwrap();
15422 let t_kv = base_len + t - off;
15445 let physical = kvl.physical_rows(off, off + t_kv)?;
15446 let k_view = e.view_u8_range(
15447 &kvl.k,
15448 physical.start * kvl.k_tok_bytes,
15449 physical.end * kvl.k_tok_bytes,
15450 );
15451 let v_view = e.view_u8_range(
15452 &kvl.v,
15453 physical.start * kvl.v_tok_bytes,
15454 physical.end * kvl.v_tok_bytes,
15455 );
15456 let swa_naive = if legacy_tkv {
15468 t_kv > win
15469 } else {
15470 seq_end > win
15471 };
15472 if swa && swa_naive {
15473 if std::env::var("MEMRA_STEP35_SWA_FA").as_deref() == Ok("0") {
15486 e.sdpa_naive_w_quantized_view(
15487 &q,
15488 &k_view,
15489 &v_view,
15490 &mut attn,
15491 hd,
15492 nh,
15493 nkv,
15494 t,
15495 t_kv,
15496 scale,
15497 true,
15498 win,
15499 kvl.k_tok_bytes,
15500 kvl.v_tok_bytes,
15501 )?;
15502 } else {
15503 e.fa_prefill_view_ws_w_hd128(
15504 &q,
15505 &k_view,
15506 &v_view,
15507 &mut attn,
15508 hd,
15509 nh,
15510 nkv,
15511 t,
15512 t_kv,
15513 scale,
15514 true,
15515 win,
15516 kvl.k_tok_bytes,
15517 kvl.v_tok_bytes,
15518 )?;
15519 }
15520 } else if std::env::var("MEMRA_NOFA").is_ok() {
15521 e.sdpa_naive_quantized_view(
15522 &q,
15523 &k_view,
15524 &v_view,
15525 &mut attn,
15526 hd,
15527 nh,
15528 nkv,
15529 t,
15530 t_kv,
15531 scale,
15532 true,
15533 kvl.k_tok_bytes,
15534 kvl.v_tok_bytes,
15535 )?;
15536 } else {
15537 e.fa_prefill_view_ws(
15542 &q,
15543 &k_view,
15544 &v_view,
15545 &mut attn,
15546 hd,
15547 nh,
15548 nkv,
15549 t,
15550 t_kv,
15551 scale,
15552 true,
15553 kvl.k_tok_bytes,
15554 kvl.v_tok_bytes,
15555 crate::Engine::kv_fp8_on(),
15556 )?;
15557 }
15558 }
15559 None => {
15560 debug_assert_eq!(
15565 seq_end, t,
15566 "step35 cacheless prefill is monolithic (seq_end == t)"
15567 );
15568 if swa && seq_end > win {
15569 e.sdpa_naive_w(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
15570 } else if std::env::var("MEMRA_NOFA").is_ok() {
15571 e.sdpa_naive(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
15572 } else {
15573 e.fa_prefill(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
15574 }
15575 }
15576 }
15577
15578 let gw = fa
15581 .attn_gate
15582 .as_ref()
15583 .ok_or("step35 layer is missing attn_gate.weight (head-wise attention gate)")?;
15584 let gt_owned = if gt_pre.is_none() {
15585 Some(e.matmul(
15586 gw,
15587 hg.ok_or("step35 attention needs hg when gt_pre is absent")?,
15588 t,
15589 )?)
15590 } else {
15591 None
15592 };
15593 let gt = gt_pre.or(gt_owned.as_ref()).unwrap();
15594 let mut ag = e.uninit(t * nh * hd)?;
15595 e.attn_head_gate(&attn, gt, &mut ag, None, hd, nh, t)?;
15596 Ok(ag)
15597 }
15598
15599 pub(crate) fn step35_attn(
15602 &self,
15603 e: &Engine,
15604 fa: &FullAttnLayer,
15605 h: &CudaSlice<f32>,
15606 pos_d: &CudaSlice<i32>,
15607 t: usize,
15608 il: usize,
15609 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
15610 let g3 = match self.step35_tp_qkv(e, fa, h, t)? {
15611 Some(g3) => g3,
15612 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
15613 };
15614 let ag = self.step35_attn_pre_wo(e, fa, g3, Some(h), None, pos_d, t, None, il, t)?;
15616 self.step35_o(e, fa, &ag, t)
15617 }
15618
15619 #[allow(clippy::too_many_arguments)]
15626 pub(crate) fn step35_attn_prime(
15627 &self,
15628 e: &Engine,
15629 fa: &FullAttnLayer,
15630 h: &CudaSlice<f32>,
15631 hx: Option<&CudaSlice<u8>>,
15632 pos_d: &CudaSlice<i32>,
15633 t: usize,
15634 cache: &mut Cache,
15635 il: usize,
15636 seq_end: usize,
15637 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
15638 if step_tp_prefill_enabled()? && fa.step_tp_qkv.is_some() {
15639 if hx.is_some() {
15640 return Err(
15641 "rank-local Step prefill preserves BF16 activations and refuses the q8_1 \
15642 pre-quantized prime path"
15643 .into(),
15644 );
15645 }
15646 return self.step35_tp_prefill_attn_resident(e, fa, il, h, pos_d, t, cache, seq_end);
15647 }
15648 let g3 = if fa.step_tp_qkv.is_some() {
15649 if hx.is_some() {
15650 return Err(
15651 "Step Q/K/V TP preserves BF16 activations and refuses the q8_1 \
15652 pre-quantized prime path"
15653 .into(),
15654 );
15655 }
15656 self.step35_tp_qkv(e, fa, h, t)?
15657 .expect("Step Q/K/V TP disappeared after the presence check")
15658 } else {
15659 match hx {
15660 Some(xh) => e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], h, xh, t)?,
15661 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
15662 }
15663 };
15664 let ag =
15665 self.step35_attn_pre_wo(e, fa, g3, Some(h), None, pos_d, t, Some(cache), il, seq_end)?;
15666 self.step35_o(e, fa, &ag, t)
15667 }
15668
15669 fn ensure_step_tp_kv_cache(
15670 &self,
15671 e: &Engine,
15672 fa: &FullAttnLayer,
15673 il: usize,
15674 cache: &mut Cache,
15675 ) -> Result<bool, Box<dyn std::error::Error>> {
15676 let tp = fa
15677 .step_tp_qkv
15678 .as_ref()
15679 .ok_or("Step TP cache hydration lost its resident projections")?;
15680 let geometry = self.step35_geom(il);
15681 let window = geometry.window.map(|window| window as usize);
15682 let ranks = tp.runtime.devices().len();
15683 let head_dim = geometry.head_dim_k as usize;
15684 let kv_heads = geometry.n_head_kv as usize;
15685 let max_ctx = cache.max_ctx;
15686
15687 if cache.tp_kv[il].is_some() {
15688 return Ok(false);
15689 }
15690 let local = cache.kv[il]
15691 .as_ref()
15692 .ok_or_else(|| format!("Step TP layer {il} has no owning-stage KV cache"))?;
15693 if local.kv_dim_k != kv_heads * head_dim || local.kv_dim_v != kv_heads * head_dim {
15694 return Err(format!(
15695 "Step TP layer {il} local KV geometry k={} v={} != {}",
15696 local.kv_dim_k,
15697 local.kv_dim_v,
15698 kv_heads * head_dim
15699 )
15700 .into());
15701 }
15702 let resident_start = window
15703 .map(|window| local.len.saturating_sub(window.saturating_sub(1)) & !31usize)
15704 .unwrap_or(0);
15705 let resident_rows = local.len - resident_start;
15706 let physical = local.physical_rows(resident_start, local.len)?;
15707 let k_rows = if resident_rows == 0 {
15708 Vec::new()
15709 } else {
15710 e.dtoh_u8_view(&e.view_u8_range(
15711 &local.k,
15712 physical.start * local.k_tok_bytes,
15713 physical.end * local.k_tok_bytes,
15714 ))?
15715 };
15716 let v_rows = if resident_rows == 0 {
15717 Vec::new()
15718 } else {
15719 e.dtoh_u8_view(&e.view_u8_range(
15720 &local.v,
15721 physical.start * local.v_tok_bytes,
15722 physical.end * local.v_tok_bytes,
15723 ))?
15724 };
15725 let mut distributed = match window {
15726 Some(window) => tp.runtime.allocate_tp_swa_kv_cache(
15727 kv_heads * head_dim,
15728 kv_heads * head_dim,
15729 max_ctx,
15730 window,
15731 )?,
15732 None => tp.runtime.allocate_tp_kv_cache(
15733 kv_heads * head_dim,
15734 kv_heads * head_dim,
15735 max_ctx,
15736 )?,
15737 };
15738 if distributed.k_tok_bytes() * ranks != local.k_tok_bytes
15739 || distributed.v_tok_bytes() * ranks != local.v_tok_bytes
15740 {
15741 return Err(format!(
15742 "Step TP layer {il} distributed/local KV token bytes disagree: \
15743 k={}x{ranks}/{} v={}x{ranks}/{}",
15744 distributed.k_tok_bytes(),
15745 local.k_tok_bytes,
15746 distributed.v_tok_bytes(),
15747 local.v_tok_bytes,
15748 )
15749 .into());
15750 }
15751 tp.runtime.hydrate_tp_kv_cache_from(
15752 &mut distributed,
15753 local.len,
15754 resident_start,
15755 &k_rows,
15756 &v_rows,
15757 )?;
15758 cache.tp_kv[il] = Some(distributed);
15759 Ok(true)
15760 }
15761
15762 #[allow(clippy::too_many_arguments)]
15763 fn step35_tp_prefill_attn_resident(
15764 &self,
15765 e: &Engine,
15766 fa: &FullAttnLayer,
15767 il: usize,
15768 h: &CudaSlice<f32>,
15769 pos_d: &CudaSlice<i32>,
15770 tokens: usize,
15771 cache: &mut Cache,
15772 seq_end: usize,
15773 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
15774 let tp = fa
15775 .step_tp_qkv
15776 .as_ref()
15777 .ok_or("Step TP prefill lost its resident projections")?;
15778 let attention = tp
15779 .attention
15780 .as_ref()
15781 .ok_or("Step TP prefill lost its resident attention auxiliaries")?;
15782 let ranks = tp.runtime.devices().len();
15783 if !step_tp_prefill_shape(
15784 true,
15785 tokens,
15786 ranks,
15787 tp.runtime.native_p2p(),
15788 true,
15789 crate::Engine::kv_fp8_on(),
15790 ) {
15791 return Err(format!(
15792 "rank-local Step prefill requires tokens>={PRIME_MIN_T}, TP2/TP4 native P2P, \
15793 rank-local attention, and q8_0/q5_1 KV; got tokens={tokens} ranks={ranks} \
15794 native_p2p={} fp8_kv={}",
15795 tp.runtime.native_p2p(),
15796 crate::Engine::kv_fp8_on(),
15797 )
15798 .into());
15799 }
15800 for seam in [
15801 "MEMRA_STEP35_SWA_TKV",
15802 "MEMRA_PRIME_CALLLOCAL",
15803 "MEMRA_PRIME_F32CHUNK0",
15804 ] {
15805 if std::env::var(seam).as_deref() == Ok("1") {
15806 return Err(format!(
15807 "rank-local Step prefill has not qualified the legacy seam {seam}=1"
15808 )
15809 .into());
15810 }
15811 }
15812
15813 let geometry = self.step35_geom(il);
15814 let window = geometry.window.map(|window| window as usize);
15815 let head_dim = geometry.head_dim_k as usize;
15816 let heads = geometry.n_head as usize;
15817 let kv_heads = geometry.n_head_kv as usize;
15818 if heads % ranks != 0 || kv_heads % ranks != 0 {
15819 return Err(format!(
15820 "Step attention heads q={heads} kv={kv_heads} are not divisible by TP={ranks}"
15821 )
15822 .into());
15823 }
15824 let local_heads = heads / ranks;
15825 let local_kv_heads = kv_heads / ranks;
15826 let local_kv_dim = local_kv_heads * head_dim;
15827 let hidden = self.cfg.n_embd as usize;
15828 let expected_input = tokens
15829 .checked_mul(hidden)
15830 .ok_or("Step TP prefill input size overflow")?;
15831 if h.len() < expected_input {
15832 return Err(format!(
15833 "Step TP prefill input {} is shorter than {tokens}x{hidden}",
15834 h.len()
15835 )
15836 .into());
15837 }
15838 let positions = e.dtoh_i32(pos_d)?;
15839 if positions.len() != tokens {
15840 return Err(format!(
15841 "rank-local Step prefill positions {} != tokens {tokens}",
15842 positions.len()
15843 )
15844 .into());
15845 }
15846
15847 let mut active_input = e.uninit(expected_input)?;
15848 e.copy_view_into(
15849 &mut active_input,
15850 0,
15851 &h.slice(0..expected_input),
15852 expected_input,
15853 )?;
15854 let mut input = tp.runtime.allocate_replicated_device_rows(tokens, hidden)?;
15855 e.stream().synchronize()?;
15860 tp.runtime
15861 .refresh_replicated_device_rows_from_root(&mut input, &active_input)?;
15862 let q_raw = tp
15863 .runtime
15864 .bf16_column_parallel_resident_replicated_device_shards(&tp.q, &input)?;
15865 let k_raw = tp
15866 .runtime
15867 .bf16_column_parallel_resident_replicated_device_shards(&tp.k, &input)?;
15868 let v_raw = tp
15869 .runtime
15870 .bf16_column_parallel_resident_replicated_device_shards(&tp.v, &input)?;
15871 let mut q = Vec::with_capacity(ranks);
15872 let mut k = Vec::with_capacity(ranks);
15873 for rank in 0..ranks {
15874 let engine = tp
15875 .runtime
15876 .rank_engine(rank)
15877 .ok_or_else(|| format!("Step TP layer {il} has no engine for rank {rank}"))?;
15878 let _main = engine.gpu.enter_main()?;
15879 let mut q_rank = engine.uninit(tokens * local_heads * head_dim)?;
15880 engine.rms_norm(
15881 &q_raw[rank],
15882 &attention.q_norm[rank],
15883 &mut q_rank,
15884 head_dim,
15885 tokens * local_heads,
15886 self.cfg.rms_eps,
15887 )?;
15888 let mut k_rank = engine.uninit(tokens * local_kv_dim)?;
15889 engine.rms_norm(
15890 &k_raw[rank],
15891 &attention.k_norm[rank],
15892 &mut k_rank,
15893 head_dim,
15894 tokens * local_kv_heads,
15895 self.cfg.rms_eps,
15896 )?;
15897 let position = engine.htod_i32(&positions)?;
15898 let rope_freqs = if geometry.rope_factors {
15899 self.step35_aux
15900 .as_ref()
15901 .and_then(|aux| aux.rope_freqs(engine))
15902 } else {
15903 None
15904 };
15905 engine.rope_neox2(
15906 &mut q_rank,
15907 &mut k_rank,
15908 &position,
15909 head_dim,
15910 geometry.n_rot as usize,
15911 local_heads,
15912 local_kv_heads,
15913 tokens,
15914 geometry.rope_base,
15915 1.0,
15916 rope_freqs,
15917 )?;
15918 q.push(q_rank);
15919 k.push(k_rank);
15920 }
15921
15922 let gate_weight = fa
15923 .attn_gate
15924 .as_ref()
15925 .ok_or("step35 layer is missing attn_gate.weight")?;
15926 let gate = e.dtoh(&e.matmul(gate_weight, h, tokens)?)?;
15927 if gate.len() != tokens * heads {
15928 return Err(format!(
15929 "Step TP layer {il} gate output {} != {tokens}x{heads}",
15930 gate.len()
15931 )
15932 .into());
15933 }
15934
15935 let hydrated = self.ensure_step_tp_kv_cache(e, fa, il, cache)?;
15936 let base_len = cache.kv[il]
15937 .as_ref()
15938 .ok_or_else(|| format!("Step TP layer {il} lost its local KV cache"))?
15939 .len;
15940 let distributed = cache.tp_kv[il]
15941 .as_ref()
15942 .ok_or_else(|| format!("Step TP layer {il} lost its distributed KV cache"))?;
15943 if distributed.committed_len() != base_len || distributed.staged_len() != base_len {
15944 return Err(format!(
15945 "Step TP layer {il} cache lengths diverged before prefill: \
15946 local={base_len} distributed={}/{}",
15947 distributed.committed_len(),
15948 distributed.staged_len()
15949 )
15950 .into());
15951 }
15952 let target_len = base_len
15953 .checked_add(tokens)
15954 .ok_or("Step TP prefill cache length overflow")?;
15955 if target_len > cache.max_ctx {
15956 return Err(format!(
15957 "Step TP layer {il} prefill exceeds cache: {base_len}+{tokens}>{}",
15958 cache.max_ctx
15959 )
15960 .into());
15961 }
15962 if seq_end < target_len {
15963 return Err(format!(
15964 "Step TP layer {il} request end {seq_end} precedes chunk end {target_len}"
15965 )
15966 .into());
15967 }
15968
15969 let transaction = cache.tp_kv[il]
15970 .as_mut()
15971 .expect("distributed cache checked above")
15972 .begin_transaction()?;
15973 if let Err(error) = tp.runtime.append_tp_kv_transaction(
15974 cache.tp_kv[il]
15975 .as_mut()
15976 .expect("distributed cache checked above"),
15977 transaction,
15978 &k,
15979 &v_raw,
15980 tokens,
15981 ) {
15982 let _ = tp.runtime.rollback_tp_kv_transaction(
15983 cache.tp_kv[il]
15984 .as_mut()
15985 .expect("distributed cache checked above"),
15986 transaction,
15987 );
15988 return Err(error);
15989 }
15990
15991 let staged = (|| -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
15992 let distributed = cache.tp_kv[il]
15993 .as_ref()
15994 .expect("distributed cache checked above");
15995 let staged_len = distributed.staged_len();
15996 let view_start = window
15997 .map(|window| base_len.saturating_sub(window.saturating_sub(1)) & !31usize)
15998 .unwrap_or(0);
15999 let physical = distributed.physical_range(view_start, staged_len)?;
16000 let t_kv = staged_len - view_start;
16001 let swa_naive = window.is_some_and(|window| seq_end > window);
16002 let mut gated = Vec::with_capacity(ranks);
16003 for rank in 0..ranks {
16004 let engine = tp
16005 .runtime
16006 .rank_engine(rank)
16007 .ok_or_else(|| format!("Step TP layer {il} has no engine for rank {rank}"))?;
16008 let _main = engine.gpu.enter_main()?;
16009 let rank_cache = distributed
16010 .rank(rank)
16011 .ok_or_else(|| format!("Step TP layer {il} has no KV cache rank {rank}"))?;
16012 let k_view = engine.view_u8_range(
16013 rank_cache.k(),
16014 physical.start * distributed.k_tok_bytes(),
16015 physical.end * distributed.k_tok_bytes(),
16016 );
16017 let v_view = engine.view_u8_range(
16018 rank_cache.v(),
16019 physical.start * distributed.v_tok_bytes(),
16020 physical.end * distributed.v_tok_bytes(),
16021 );
16022 let mut attention_out = engine.uninit(tokens * local_heads * head_dim)?;
16023 if swa_naive {
16024 let window = window.expect("SWA predicate requires a window");
16025 if std::env::var("MEMRA_STEP35_SWA_FA").as_deref() == Ok("0") {
16026 engine.sdpa_naive_w_quantized_view(
16027 &q[rank],
16028 &k_view,
16029 &v_view,
16030 &mut attention_out,
16031 head_dim,
16032 local_heads,
16033 local_kv_heads,
16034 tokens,
16035 t_kv,
16036 geometry.attention_scale(),
16037 true,
16038 window,
16039 distributed.k_tok_bytes(),
16040 distributed.v_tok_bytes(),
16041 )?;
16042 } else {
16043 engine.fa_prefill_view_ws_w_hd128(
16044 &q[rank],
16045 &k_view,
16046 &v_view,
16047 &mut attention_out,
16048 head_dim,
16049 local_heads,
16050 local_kv_heads,
16051 tokens,
16052 t_kv,
16053 geometry.attention_scale(),
16054 true,
16055 window,
16056 distributed.k_tok_bytes(),
16057 distributed.v_tok_bytes(),
16058 )?;
16059 }
16060 } else if std::env::var("MEMRA_NOFA").is_ok() {
16061 engine.sdpa_naive_quantized_view(
16062 &q[rank],
16063 &k_view,
16064 &v_view,
16065 &mut attention_out,
16066 head_dim,
16067 local_heads,
16068 local_kv_heads,
16069 tokens,
16070 t_kv,
16071 geometry.attention_scale(),
16072 true,
16073 distributed.k_tok_bytes(),
16074 distributed.v_tok_bytes(),
16075 )?;
16076 } else {
16077 engine.fa_prefill_view_ws(
16078 &q[rank],
16079 &k_view,
16080 &v_view,
16081 &mut attention_out,
16082 head_dim,
16083 local_heads,
16084 local_kv_heads,
16085 tokens,
16086 t_kv,
16087 geometry.attention_scale(),
16088 true,
16089 distributed.k_tok_bytes(),
16090 distributed.v_tok_bytes(),
16091 false,
16092 )?;
16093 }
16094
16095 let gate_start = rank * local_heads;
16096 let mut gate_rank = Vec::with_capacity(tokens * local_heads);
16097 for token in 0..tokens {
16098 let start = token * heads + gate_start;
16099 gate_rank.extend_from_slice(&gate[start..start + local_heads]);
16100 }
16101 let gate_rank = engine.htod(&gate_rank)?;
16102 let mut gated_rank = engine.uninit(tokens * local_heads * head_dim)?;
16103 engine.attn_head_gate(
16104 &attention_out,
16105 &gate_rank,
16106 &mut gated_rank,
16107 None,
16108 head_dim,
16109 local_heads,
16110 tokens,
16111 )?;
16112 gated.push(gated_rank);
16113 }
16114 for rank in 1..ranks {
16115 let engine = tp
16116 .runtime
16117 .rank_engine(rank)
16118 .ok_or_else(|| format!("Step TP layer {il} has no engine for rank {rank}"))?;
16119 let _main = engine.gpu.enter_main()?;
16120 engine.stream().synchronize()?;
16121 }
16122
16123 let (output, k_shadow, v_shadow) = if tp.runtime.bulk_p2p() {
16124 let output = tp
16125 .runtime
16126 .step_bf16_row_parallel_resident_root_device(&tp.o, &gated, tokens)?;
16127 let k_shadow =
16128 tp.runtime
16129 .gather_native_column_shards_device(&k, tokens, local_kv_dim)?;
16130 let v_shadow =
16131 tp.runtime
16132 .gather_native_column_shards_device(&v_raw, tokens, local_kv_dim)?;
16133 let root = tp
16134 .runtime
16135 .rank_engine(0)
16136 .ok_or("Step TP prefill lost its root engine")?;
16137 let _main = root.gpu.enter_main()?;
16138 root.stream().synchronize()?;
16139 (output, k_shadow, v_shadow)
16140 } else {
16141 let attention = tp.runtime.gather_native_column_shards(
16142 &gated,
16143 tokens,
16144 local_heads * head_dim,
16145 )?;
16146 let output = tp
16147 .runtime
16148 .step_bf16_row_parallel_resident_native(&tp.o, &attention, tokens)?;
16149 let k_shadow = tp
16150 .runtime
16151 .gather_native_column_shards(&k, tokens, local_kv_dim)?;
16152 let v_shadow =
16153 tp.runtime
16154 .gather_native_column_shards(&v_raw, tokens, local_kv_dim)?;
16155 (e.htod(&output)?, e.htod(&k_shadow)?, e.htod(&v_shadow)?)
16156 };
16157 let local = cache.kv[il]
16158 .as_mut()
16159 .ok_or_else(|| format!("Step TP layer {il} lost its local KV cache"))?;
16160 if local.len != base_len {
16161 return Err(format!(
16162 "Step TP layer {il} local cache changed during prefill: \
16163 len={} base={base_len}",
16164 local.len
16165 )
16166 .into());
16167 }
16168 let retain_from = window
16169 .map(|window| {
16170 let staged_retain = staged_len.saturating_sub(window) & !31usize;
16171 let rollback_retain =
16172 base_len.saturating_sub(window.saturating_sub(1)) & !31usize;
16173 staged_retain.min(rollback_retain)
16174 })
16175 .unwrap_or(0);
16176 let write_row = e.prepare_kv_append(local, retain_from, tokens)?;
16177 e.append_kv_quantized_rows(
16178 &k_shadow,
16179 &v_shadow,
16180 &mut local.k,
16181 &mut local.v,
16182 write_row,
16183 tokens,
16184 local.kv_dim_k,
16185 local.kv_dim_v,
16186 local.k_tok_bytes,
16187 local.v_tok_bytes,
16188 false,
16189 )?;
16190 local.len = staged_len;
16191 e.set_i32_one(&mut local.len_d, staged_len as i32)?;
16192 Ok(output)
16193 })();
16194
16195 let output = match staged {
16196 Ok(output) => output,
16197 Err(error) => {
16198 let _ = tp.runtime.rollback_tp_kv_transaction(
16199 cache.tp_kv[il]
16200 .as_mut()
16201 .expect("distributed cache checked above"),
16202 transaction,
16203 );
16204 if let Some(local) = cache.kv[il].as_mut() {
16205 local.len = base_len;
16206 let _ = e.set_i32_one(&mut local.len_d, base_len as i32);
16207 }
16208 return Err(error);
16209 }
16210 };
16211 if let Err(error) = tp.runtime.commit_tp_kv_transaction(
16212 cache.tp_kv[il]
16213 .as_mut()
16214 .expect("distributed cache checked above"),
16215 transaction,
16216 tokens,
16217 ) {
16218 let _ = tp.runtime.rollback_tp_kv_transaction(
16219 cache.tp_kv[il]
16220 .as_mut()
16221 .expect("distributed cache checked above"),
16222 transaction,
16223 );
16224 let local = cache.kv[il].as_mut().expect("local cache checked above");
16225 local.len = base_len;
16226 e.set_i32_one(&mut local.len_d, base_len as i32)?;
16227 return Err(error);
16228 }
16229
16230 let committed = cache.tp_kv[il]
16231 .as_ref()
16232 .expect("distributed cache checked above")
16233 .committed_len();
16234 let local_len = cache.kv[il]
16235 .as_ref()
16236 .expect("local cache checked above")
16237 .len;
16238 if committed != local_len {
16239 return Err(format!(
16240 "Step TP layer {il} committed cache length {committed} != local shadow {local_len}"
16241 )
16242 .into());
16243 }
16244 eprintln!(
16245 "[step-tp-prefill-attn] execute layer={} devices={:?} tokens={tokens} \
16246 qkv_tensor_parallel=true qk_norm_rank_local=true rope_rank_local=true \
16247 kv_cache_distributed=true kv_cache_hydrated={} attention_tensor_parallel=true \
16248 attention_scope={} input_path=root-device-replicated gate_tensor_parallel=false \
16249 gate_shards=host-canonical o_tensor_parallel=true local_cache_shadow=true \
16250 cache_commit=chunk transport={} native_p2p=true bulk_p2p={} \
16251 output={} performance_claim=false",
16252 tp.layer,
16253 tp.devices,
16254 hydrated,
16255 if window.is_some() {
16256 "rank-local-swa-ring"
16257 } else {
16258 "rank-local-global"
16259 },
16260 tp.runtime.transport_label(),
16261 tp.runtime.bulk_p2p(),
16262 if tp.runtime.bulk_p2p() {
16263 "root-device"
16264 } else {
16265 "root-readback"
16266 },
16267 );
16268 Ok(output)
16269 }
16270
16271 fn step35_tp_decode_attn_resident(
16272 &self,
16273 e: &Engine,
16274 fa: &FullAttnLayer,
16275 il: usize,
16276 h: &CudaSlice<f32>,
16277 pos_d: &CudaSlice<i32>,
16278 cache: &mut Cache,
16279 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
16280 static ATTN_NS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
16284 static ATTN_CALLS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
16285 let timing = std::env::var("MEMRA_STEP_TP_TIMING").as_deref() == Ok("1");
16286 let started = timing.then(std::time::Instant::now);
16287 let result = if crate::tp::step_tp_decode_v2_enabled()? {
16288 self.step35_tp_decode_attn_resident_v2(e, fa, il, h, pos_d, cache)
16289 } else {
16290 self.step35_tp_decode_attn_resident_inner(e, fa, il, h, pos_d, cache)
16291 };
16292 if let Some(started) = started {
16293 use std::sync::atomic::Ordering;
16294 let ns = ATTN_NS.fetch_add(started.elapsed().as_nanos() as u64, Ordering::Relaxed)
16295 + started.elapsed().as_nanos() as u64;
16296 let calls = ATTN_CALLS.fetch_add(1, Ordering::Relaxed) + 1;
16297 if calls % 430 == 0 {
16298 eprintln!(
16299 "[step-tp-attn-timing] calls={calls} total_ms={:.1} avg_us={:.1}",
16300 ns as f64 / 1.0e6,
16301 ns as f64 / calls as f64 / 1.0e3,
16302 );
16303 }
16304 }
16305 result
16306 }
16307
16308 #[allow(clippy::too_many_arguments)]
16309 fn step35_tp_decode_attn_resident_inner(
16310 &self,
16311 e: &Engine,
16312 fa: &FullAttnLayer,
16313 il: usize,
16314 h: &CudaSlice<f32>,
16315 pos_d: &CudaSlice<i32>,
16316 cache: &mut Cache,
16317 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
16318 static T_POS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
16323 static T_QKV: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
16324 static T_NORMROPE: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
16325 static T_GATE: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
16326 static T_APPEND: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
16327 static T_ATTN: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
16328 static T_OPROJ: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
16329 static T_SHADOW: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
16330 static T_PHASE_CALLS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
16331 let timing = std::env::var("MEMRA_STEP_TP_TIMING").as_deref() == Ok("1");
16332 fn lap(
16333 runtime: &crate::tp::TpE4m3HostBounce,
16334 e: &Engine,
16335 timer: &std::sync::atomic::AtomicU64,
16336 started: &mut Option<std::time::Instant>,
16337 ) -> Result<(), Box<dyn std::error::Error>> {
16338 let Some(start) = started.as_mut() else {
16339 return Ok(());
16340 };
16341 for rank in 0..runtime.devices().len() {
16342 if let Some(engine) = runtime.rank_engine(rank) {
16343 let _main = engine.gpu.enter_main()?;
16344 engine.stream().synchronize()?;
16345 }
16346 }
16347 e.stream().synchronize()?;
16348 timer.fetch_add(
16349 start.elapsed().as_nanos() as u64,
16350 std::sync::atomic::Ordering::Relaxed,
16351 );
16352 *start = std::time::Instant::now();
16353 Ok(())
16354 }
16355 let tp = fa
16356 .step_tp_qkv
16357 .as_ref()
16358 .ok_or("Step TP decode lost its resident projections")?;
16359 let attention = tp
16360 .attention
16361 .as_ref()
16362 .ok_or("Step TP decode lost its resident attention auxiliaries")?;
16363 if !tp.runtime.native_p2p() {
16364 return Err("rank-local Step attention requires native P2P".into());
16365 }
16366 if crate::Engine::kv_fp8_on() {
16367 return Err("rank-local Step attention has not qualified the FP8 KV cache".into());
16368 }
16369
16370 let geometry = self.step35_geom(il);
16371 let window = geometry.window.map(|window| window as usize);
16372 let ranks = tp.runtime.devices().len();
16373 let head_dim = geometry.head_dim_k as usize;
16374 let heads = geometry.n_head as usize;
16375 let kv_heads = geometry.n_head_kv as usize;
16376 if heads % ranks != 0 || kv_heads % ranks != 0 {
16377 return Err(format!(
16378 "Step attention heads q={heads} kv={kv_heads} are not divisible by TP={ranks}"
16379 )
16380 .into());
16381 }
16382 let local_heads = heads / ranks;
16383 let local_kv_heads = kv_heads / ranks;
16384 let local_kv_dim = local_kv_heads * head_dim;
16385 let max_ctx = cache.max_ctx;
16386
16387 let hydrated = self.ensure_step_tp_kv_cache(e, fa, il, cache)?;
16388
16389 let base_len = cache.kv[il]
16390 .as_ref()
16391 .ok_or_else(|| format!("Step TP layer {il} lost its local KV cache"))?
16392 .len;
16393 let distributed = cache.tp_kv[il]
16394 .as_ref()
16395 .ok_or_else(|| format!("Step TP layer {il} lost its distributed KV cache"))?;
16396 if distributed.committed_len() != base_len || distributed.staged_len() != base_len {
16397 return Err(format!(
16398 "Step TP layer {il} cache lengths diverged before decode: \
16399 local={base_len} distributed={}/{}",
16400 distributed.committed_len(),
16401 distributed.staged_len()
16402 )
16403 .into());
16404 }
16405
16406 let mut lap_start = timing.then(std::time::Instant::now);
16407 let positions = e.dtoh_i32(pos_d)?;
16408 if positions.len() != 1 {
16409 return Err(format!(
16410 "rank-local Step decode requires one position, got {}",
16411 positions.len()
16412 )
16413 .into());
16414 }
16415 lap(&tp.runtime, e, &T_POS, &mut lap_start)?;
16416 let (q_raw, k_raw, v_raw, input_path) = if let Some(decode_input) =
16417 attention.decode_input.as_ref()
16418 {
16419 let mut decode_input = decode_input
16420 .lock()
16421 .map_err(|_| "Step TP replicated decode input lock is poisoned")?;
16422 e.stream().synchronize()?;
16426 tp.runtime
16427 .refresh_replicated_device_rows_from_root(&mut decode_input, h)?;
16428 let q_raw = tp
16429 .runtime
16430 .bf16_column_parallel_resident_replicated_device_shards(&tp.q, &decode_input)?;
16431 let k_raw = tp
16432 .runtime
16433 .bf16_column_parallel_resident_replicated_device_shards(&tp.k, &decode_input)?;
16434 let v_raw = tp
16435 .runtime
16436 .bf16_column_parallel_resident_replicated_device_shards(&tp.v, &decode_input)?;
16437 (q_raw, k_raw, v_raw, "root-device-replicated")
16438 } else {
16439 let activation = e.dtoh(h)?;
16440 let q_raw =
16441 tp.runtime
16442 .bf16_column_parallel_resident_device_shards(&tp.q, &activation, 1)?;
16443 let k_raw =
16444 tp.runtime
16445 .bf16_column_parallel_resident_device_shards(&tp.k, &activation, 1)?;
16446 let v_raw =
16447 tp.runtime
16448 .bf16_column_parallel_resident_device_shards(&tp.v, &activation, 1)?;
16449 (q_raw, k_raw, v_raw, "host-replicated")
16450 };
16451 lap(&tp.runtime, e, &T_QKV, &mut lap_start)?;
16452 let mut q = Vec::with_capacity(ranks);
16453 let mut k = Vec::with_capacity(ranks);
16454 for rank in 0..ranks {
16455 let engine = tp
16456 .runtime
16457 .rank_engine(rank)
16458 .ok_or_else(|| format!("Step TP layer {il} has no engine for rank {rank}"))?;
16459 let _main = engine.gpu.enter_main()?;
16460 let mut q_rank = engine.uninit(local_heads * head_dim)?;
16461 engine.rms_norm(
16462 &q_raw[rank],
16463 &attention.q_norm[rank],
16464 &mut q_rank,
16465 head_dim,
16466 local_heads,
16467 self.cfg.rms_eps,
16468 )?;
16469 let mut k_rank = engine.uninit(local_kv_dim)?;
16470 engine.rms_norm(
16471 &k_raw[rank],
16472 &attention.k_norm[rank],
16473 &mut k_rank,
16474 head_dim,
16475 local_kv_heads,
16476 self.cfg.rms_eps,
16477 )?;
16478 let position = engine.htod_i32(&positions)?;
16479 let rope_freqs = if geometry.rope_factors {
16480 self.step35_aux
16481 .as_ref()
16482 .and_then(|aux| aux.rope_freqs(engine))
16483 } else {
16484 None
16485 };
16486 engine.rope_neox2(
16487 &mut q_rank,
16488 &mut k_rank,
16489 &position,
16490 head_dim,
16491 geometry.n_rot as usize,
16492 local_heads,
16493 local_kv_heads,
16494 1,
16495 geometry.rope_base,
16496 1.0,
16497 rope_freqs,
16498 )?;
16499 q.push(q_rank);
16500 k.push(k_rank);
16501 }
16502 lap(&tp.runtime, e, &T_NORMROPE, &mut lap_start)?;
16503
16504 let gate_weight = fa
16505 .attn_gate
16506 .as_ref()
16507 .ok_or("step35 layer is missing attn_gate.weight")?;
16508 let gate = e.matmul(gate_weight, h, 1)?;
16509 let gate = e.dtoh(&gate)?;
16510 if gate.len() != heads {
16511 return Err(format!("Step TP layer {il} gate output {} != {heads}", gate.len()).into());
16512 }
16513 lap(&tp.runtime, e, &T_GATE, &mut lap_start)?;
16514
16515 let transaction = cache.tp_kv[il]
16516 .as_mut()
16517 .expect("distributed cache checked above")
16518 .begin_transaction()?;
16519 if let Err(error) = tp.runtime.append_tp_kv_transaction(
16520 cache.tp_kv[il]
16521 .as_mut()
16522 .expect("distributed cache checked above"),
16523 transaction,
16524 &k,
16525 &v_raw,
16526 1,
16527 ) {
16528 let _ = tp.runtime.rollback_tp_kv_transaction(
16529 cache.tp_kv[il]
16530 .as_mut()
16531 .expect("distributed cache checked above"),
16532 transaction,
16533 );
16534 return Err(error);
16535 }
16536 lap(&tp.runtime, e, &T_APPEND, &mut lap_start)?;
16537
16538 let staged = (|| -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
16539 let distributed = cache.tp_kv[il]
16540 .as_ref()
16541 .expect("distributed cache checked above");
16542 let staged_len = distributed.staged_len();
16543 let view_start = window
16544 .map(|window| staged_len.saturating_sub(window))
16545 .unwrap_or(0);
16546 let physical = distributed.physical_range(view_start, staged_len)?;
16547 let t_kv = staged_len - view_start;
16548 let mut gated = Vec::with_capacity(ranks);
16549 for rank in 0..ranks {
16550 let engine = tp
16551 .runtime
16552 .rank_engine(rank)
16553 .ok_or_else(|| format!("Step TP layer {il} has no engine for rank {rank}"))?;
16554 let _main = engine.gpu.enter_main()?;
16555 let rank_cache = distributed
16556 .rank(rank)
16557 .ok_or_else(|| format!("Step TP layer {il} has no KV cache rank {rank}"))?;
16558 let k_view = engine.view_u8_range(
16559 rank_cache.k(),
16560 physical.start * distributed.k_tok_bytes(),
16561 physical.end * distributed.k_tok_bytes(),
16562 );
16563 let v_view = engine.view_u8_range(
16564 rank_cache.v(),
16565 physical.start * distributed.v_tok_bytes(),
16566 physical.end * distributed.v_tok_bytes(),
16567 );
16568 let mut attention_out = engine.uninit(local_heads * head_dim)?;
16569 engine.fa_decode_kvmod(
16570 &q[rank],
16571 &k_view,
16572 &v_view,
16573 &mut attention_out,
16574 head_dim,
16575 local_heads,
16576 local_kv_heads,
16577 t_kv,
16578 geometry.attention_scale(),
16579 distributed.k_tok_bytes(),
16580 distributed.v_tok_bytes(),
16581 false,
16582 )?;
16583 let gate_start = rank * local_heads;
16584 let gate_rank = engine.htod(&gate[gate_start..gate_start + local_heads])?;
16585 let mut gated_rank = engine.uninit(local_heads * head_dim)?;
16586 engine.attn_head_gate(
16587 &attention_out,
16588 &gate_rank,
16589 &mut gated_rank,
16590 None,
16591 head_dim,
16592 local_heads,
16593 1,
16594 )?;
16595 gated.push(gated_rank);
16596 }
16597 lap(&tp.runtime, e, &T_ATTN, &mut lap_start)?;
16598
16599 let gathered =
16600 tp.runtime
16601 .gather_native_column_shards(&gated, 1, local_heads * head_dim)?;
16602 let output = tp
16603 .runtime
16604 .step_bf16_row_parallel_resident_native(&tp.o, &gathered, 1)?;
16605 let output = e.htod(&output)?;
16606 lap(&tp.runtime, e, &T_OPROJ, &mut lap_start)?;
16607
16608 let k_shadow = tp
16609 .runtime
16610 .gather_native_column_shards(&k, 1, local_kv_dim)?;
16611 let v_shadow = tp
16612 .runtime
16613 .gather_native_column_shards(&v_raw, 1, local_kv_dim)?;
16614 let k_shadow = e.htod(&k_shadow)?;
16615 let v_shadow = e.htod(&v_shadow)?;
16616 let local = cache.kv[il]
16617 .as_mut()
16618 .ok_or_else(|| format!("Step TP layer {il} lost its local KV cache"))?;
16619 if local.len != base_len || base_len + 1 > max_ctx {
16620 return Err(format!(
16621 "Step TP layer {il} local cache changed during decode: \
16622 len={} base={base_len} max={max_ctx}",
16623 local.len
16624 )
16625 .into());
16626 }
16627 let retain_from = window
16628 .map(|window| {
16629 let staged_retain = (base_len + 1).saturating_sub(window) & !31usize;
16630 let rollback_retain =
16631 base_len.saturating_sub(window.saturating_sub(1)) & !31usize;
16632 staged_retain.min(rollback_retain)
16633 })
16634 .unwrap_or(0);
16635 let write_row = e.prepare_kv_append(local, retain_from, 1)?;
16636 e.append_kv_quantized(
16637 &k_shadow,
16638 &v_shadow,
16639 &mut local.k,
16640 &mut local.v,
16641 write_row,
16642 local.kv_dim_k,
16643 local.kv_dim_v,
16644 local.k_tok_bytes,
16645 local.v_tok_bytes,
16646 false,
16647 )?;
16648 local.len = base_len + 1;
16649 e.set_i32_one(&mut local.len_d, local.len as i32)?;
16650 Ok(output)
16651 })();
16652
16653 let output = match staged {
16654 Ok(output) => output,
16655 Err(error) => {
16656 let _ = tp.runtime.rollback_tp_kv_transaction(
16657 cache.tp_kv[il]
16658 .as_mut()
16659 .expect("distributed cache checked above"),
16660 transaction,
16661 );
16662 if let Some(local) = cache.kv[il].as_mut() {
16663 local.len = base_len;
16664 let _ = e.set_i32_one(&mut local.len_d, base_len as i32);
16665 }
16666 return Err(error);
16667 }
16668 };
16669 if let Err(error) = tp.runtime.commit_tp_kv_transaction(
16670 cache.tp_kv[il]
16671 .as_mut()
16672 .expect("distributed cache checked above"),
16673 transaction,
16674 1,
16675 ) {
16676 let _ = tp.runtime.rollback_tp_kv_transaction(
16677 cache.tp_kv[il]
16678 .as_mut()
16679 .expect("distributed cache checked above"),
16680 transaction,
16681 );
16682 let local = cache.kv[il].as_mut().expect("local cache checked above");
16683 local.len = base_len;
16684 e.set_i32_one(&mut local.len_d, base_len as i32)?;
16685 return Err(error);
16686 }
16687
16688 let committed = cache.tp_kv[il]
16689 .as_ref()
16690 .expect("distributed cache checked above")
16691 .committed_len();
16692 let local_len = cache.kv[il]
16693 .as_ref()
16694 .expect("local cache checked above")
16695 .len;
16696 if committed != local_len {
16697 return Err(format!(
16698 "Step TP layer {il} committed cache length {committed} != local shadow {local_len}"
16699 )
16700 .into());
16701 }
16702 lap(&tp.runtime, e, &T_SHADOW, &mut lap_start)?;
16703 if timing {
16704 use std::sync::atomic::Ordering;
16705 let calls = T_PHASE_CALLS.fetch_add(1, Ordering::Relaxed) + 1;
16706 if calls % 430 == 0 {
16707 let avg = |t: &std::sync::atomic::AtomicU64| {
16708 t.load(Ordering::Relaxed) as f64 / calls as f64 / 1.0e3
16709 };
16710 eprintln!(
16711 "[step-tp-attn-phase] calls={calls} avg_us pos={:.1} qkv={:.1} \
16712 normrope={:.1} gate={:.1} append={:.1} attn={:.1} oproj={:.1} shadow={:.1}",
16713 avg(&T_POS),
16714 avg(&T_QKV),
16715 avg(&T_NORMROPE),
16716 avg(&T_GATE),
16717 avg(&T_APPEND),
16718 avg(&T_ATTN),
16719 avg(&T_OPROJ),
16720 avg(&T_SHADOW),
16721 );
16722 }
16723 }
16724 eprintln!(
16725 "[step-tp-attn] execute layer={} devices={:?} tokens=1 \
16726 qkv_tensor_parallel=true qk_norm_rank_local=true rope_rank_local=true \
16727 kv_cache_distributed=true kv_cache_hydrated={} attention_tensor_parallel=true \
16728 attention_scope={} input_path={} kv_physical_rows={} \
16729 gate_tensor_parallel=false gate_shards=host-canonical o_tensor_parallel=true \
16730 local_cache_shadow=true cache_commit=immediate transport={} native_p2p=true \
16731 bulk_p2p={} output=root-readback performance_claim=false",
16732 tp.layer,
16733 tp.devices,
16734 hydrated,
16735 if window.is_some() {
16736 "rank-local-swa-ring"
16737 } else {
16738 "rank-local-global"
16739 },
16740 input_path,
16741 cache.tp_kv[il]
16742 .as_ref()
16743 .expect("distributed cache checked above")
16744 .physical_capacity(),
16745 tp.runtime.transport_label(),
16746 tp.runtime.bulk_p2p(),
16747 );
16748 Ok(output)
16749 }
16750
16751 #[allow(clippy::too_many_arguments)]
16758 pub(crate) fn step35_verify_qkv_precompute(
16763 &self,
16764 e: &Engine,
16765 il: usize,
16766 h_t: &CudaSlice<f32>,
16767 t: usize,
16768 ) -> Result<bool, Box<dyn std::error::Error>> {
16769 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
16770 return Ok(false);
16771 };
16772 let Some(tp) = fa.step_tp_qkv.as_ref() else {
16773 return Ok(false);
16774 };
16775 let Some(attention) = tp.attention.as_ref() else {
16776 return Ok(false);
16777 };
16778 if !tp.runtime.native_p2p() || !crate::tp::step_tp_qkv_fused_enabled()? {
16779 return Ok(false);
16780 }
16781 let geometry = self.step35_geom(il);
16782 let heads = geometry.n_head as usize;
16783 let ws_index = tp
16784 .runtime
16785 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
16786 let gate_shards = attention
16787 .gate_shards_bf16
16788 .as_deref()
16789 .map(crate::tp::StepTpGateShards::Bf16);
16790 tp.runtime.decode_v2_input_qkv_tcol(
16791 ws_index,
16792 e,
16793 h_t,
16794 t,
16795 &tp.q,
16796 &tp.k,
16797 &tp.v,
16798 gate_shards,
16799 )?;
16800 Ok(true)
16801 }
16802
16803 pub(crate) fn step35_verify_oproj_tcol(
16808 &self,
16809 e: &Engine,
16810 il: usize,
16811 t: usize,
16812 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
16813 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
16814 return Err("tcol o_proj join expects full attention".into());
16815 };
16816 let tp = fa
16817 .step_tp_qkv
16818 .as_ref()
16819 .ok_or("tcol o_proj join lost its resident projections")?;
16820 let heads = self.step35_geom(il).n_head as usize;
16821 let ws_index = tp
16822 .runtime
16823 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
16824 tp.runtime.decode_v2_oproj_tcol(ws_index, e, &tp.o, t)
16825 }
16826
16827 pub(crate) fn step35_spec_fa2_precheck(
16834 &self,
16835 cache: &Cache,
16836 il: usize,
16837 pos0: usize,
16838 ) -> Result<bool, Box<dyn std::error::Error>> {
16839 fn nope(clause: &str, il: usize, pos0: usize) -> bool {
16842 static DBG: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16843 static SEEN: std::sync::Mutex<Vec<&'static str>> = std::sync::Mutex::new(Vec::new());
16844 if *DBG.get_or_init(|| std::env::var("MEMRA_SPEC_FA2_DEBUG").as_deref() == Ok("1")) {
16845 let mut seen = SEEN.lock().unwrap();
16846 if !seen.iter().any(|c| *c == clause) {
16847 seen.push(Box::leak(clause.to_string().into_boxed_str()));
16849 eprintln!("[spec-fa2] precheck FAIL clause={clause} il={il} pos0={pos0}");
16850 }
16851 }
16852 false
16853 }
16854 static ONLY: std::sync::OnceLock<Option<usize>> = std::sync::OnceLock::new();
16856 if let Some(only) =
16857 ONLY.get_or_init(|| std::env::var("MEMRA_SPEC_FA2_LAYER").ok()?.parse().ok())
16858 {
16859 if *only != il {
16860 return Ok(false);
16861 }
16862 }
16863 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
16864 return Ok(nope("mixer", il, pos0));
16865 };
16866 let Some(tp) = fa.step_tp_qkv.as_ref() else {
16867 return Ok(nope("step_tp", il, pos0));
16868 };
16869 let Some(attention) = tp.attention.as_ref() else {
16870 return Ok(nope("attention", il, pos0));
16871 };
16872 if !tp.runtime.native_p2p()
16873 || crate::Engine::kv_fp8_on()
16874 || !crate::tp::step_tp_dcw_enabled()?
16875 || !crate::tp::step_tp_qkv_fused_enabled()?
16876 || (attention.gate_shards.is_none() && attention.gate_shards_bf16.is_none())
16877 {
16878 return Ok(nope("runtime-doors", il, pos0));
16879 }
16880 let geometry = self.step35_geom(il);
16881 let head_dim = geometry.head_dim_k as usize;
16882 if head_dim > 256 || head_dim % 32 != 0 || !crate::fa_v3_on() {
16883 return Ok(nope("fa-class", il, pos0));
16884 }
16885 let Some(distributed) = cache.tp_kv[il].as_ref() else {
16886 return Ok(nope("tp-kv", il, pos0));
16887 };
16888 if distributed.staged_len() != pos0 {
16889 return Ok(nope("staged-len", il, pos0));
16890 }
16891 let (_, would_rebase) = distributed.peek_append_ring(2)?;
16894 if would_rebase {
16895 return Ok(nope("rebase", il, pos0));
16896 }
16897 let window = geometry.window.map(|w| w as usize);
16898 if let Some(w) = window {
16906 if pos0 + 2 > w {
16907 return Ok(nope("swa-capped", il, pos0));
16908 }
16909 }
16910 let (t0, t1) = (pos0 + 1, pos0 + 2);
16914 if t0 < 96 {
16915 return Ok(nope("dcw-floor", il, pos0));
16916 }
16917 if std::env::var("MEMRA_NO_FA_VEC").is_ok() || t0 < crate::fa_vec_min_tkv() {
16918 return Ok(nope("vec-floor", il, pos0));
16919 }
16920 let ranks = tp.runtime.devices().len();
16926 let local_kv_heads = (geometry.n_head_kv as usize / ranks).max(1);
16927 let sp0 = crate::fa_split_keys_pub(t0, local_kv_heads);
16928 let sp1 = crate::fa_split_keys_pub(t1, local_kv_heads);
16929 if sp0 != sp1 {
16930 return Ok(nope("partition-sp", il, pos0));
16931 }
16932 let (ns0, ns1) = (t0.div_ceil(sp0), t1.div_ceil(sp1));
16933 if ns0 != ns1 {
16934 return Ok(nope("partition-ns", il, pos0));
16935 }
16936 if t0.div_ceil(ns0) != t1.div_ceil(ns1) {
16937 return Ok(nope("partition-per", il, pos0));
16938 }
16939 Ok(true)
16940 }
16941
16942 pub(crate) fn step35_fa_rows_precheck(
16947 &self,
16948 cache: &Cache,
16949 il: usize,
16950 pos0: usize,
16951 t: usize,
16952 ) -> Result<bool, Box<dyn std::error::Error>> {
16953 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
16954 return Ok(false);
16955 };
16956 let Some(tp) = fa.step_tp_qkv.as_ref() else {
16957 return Ok(false);
16958 };
16959 let Some(attention) = tp.attention.as_ref() else {
16960 return Ok(false);
16961 };
16962 if !tp.runtime.native_p2p()
16963 || crate::Engine::kv_fp8_on()
16964 || !crate::tp::step_tp_dcw_enabled()?
16965 || !crate::tp::step_tp_qkv_fused_enabled()?
16966 || (attention.gate_shards.is_none() && attention.gate_shards_bf16.is_none())
16967 {
16968 return Ok(false);
16969 }
16970 let geometry = self.step35_geom(il);
16971 let head_dim = geometry.head_dim_k as usize;
16972 if head_dim > 256 || head_dim % 32 != 0 || !crate::fa_v3_on() {
16973 return Ok(false);
16974 }
16975 if crate::fa_sm_count() < 128
16976 || std::env::var("MEMRA_FA_SPLIT").is_ok()
16977 || std::env::var("MEMRA_FA_SP_SHORT").is_ok()
16978 || std::env::var("MEMRA_FA_SP16").is_ok()
16979 || std::env::var("MEMRA_NO_FA_VEC").is_ok()
16980 {
16981 return Ok(false);
16982 }
16983 let Some(distributed) = cache.tp_kv[il].as_ref() else {
16984 return Ok(false);
16985 };
16986 if distributed.staged_len() != pos0 {
16987 return Ok(false);
16988 }
16989 let (_, would_rebase) = distributed.peek_append_ring(t)?;
16990 if would_rebase {
16991 return Ok(false);
16992 }
16993 let window = geometry.window.map(|w| w as usize);
16996 let t0 = window.map(|w| (pos0 + 1).min(w)).unwrap_or(pos0 + 1);
16997 if t0 < 96 || t0 < crate::fa_vec_min_tkv() {
16998 return Ok(false);
16999 }
17000 Ok(true)
17001 }
17002
17003 pub(crate) fn step35_verify_fa_rows_join(
17007 &self,
17008 e: &Engine,
17009 il: usize,
17010 cache: &Cache,
17011 pos0: usize,
17012 t: usize,
17013 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
17014 use cudarc::driver::DevicePtr;
17015 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
17016 return Err("fa rows join expects full attention".into());
17017 };
17018 let tp = fa
17019 .step_tp_qkv
17020 .as_ref()
17021 .ok_or("fa rows join lost its resident projections")?;
17022 let geometry = self.step35_geom(il);
17023 let heads = geometry.n_head as usize;
17024 let head_dim = geometry.head_dim_k as usize;
17025 let window = geometry.window.map(|w| w as usize);
17026 let distributed = cache.tp_kv[il]
17027 .as_ref()
17028 .ok_or("fa rows join lost its distributed KV cache")?;
17029 let (k_tok_bytes, v_tok_bytes) = (distributed.k_tok_bytes(), distributed.v_tok_bytes());
17030 let ladder = |t_kv: usize| -> usize {
17032 if t_kv <= 2048 {
17033 16
17034 } else if t_kv <= 16384 {
17035 64
17036 } else {
17037 128
17038 }
17039 };
17040 let mut max_ns = 1usize;
17041 for r in 0..t {
17042 let t_kv = window
17043 .map(|w| (pos0 + r + 1).min(w))
17044 .unwrap_or(pos0 + r + 1);
17045 max_ns = max_ns.max(t_kv.div_ceil(ladder(t_kv)));
17046 }
17047 let ranks = tp.runtime.devices().len();
17052 let mut tables = Vec::with_capacity(ranks);
17053 for rank in 0..ranks {
17054 let engine = tp
17055 .runtime
17056 .rank_engine(rank)
17057 .ok_or("fa rows join lost a rank engine")?;
17058 let rank_cache = distributed
17059 .rank(rank)
17060 .ok_or("fa rows join lost a KV cache rank")?;
17061 let _main = engine.gpu.enter_main()?;
17062 let s = engine.stream();
17063 let (kp, _g0) = rank_cache.k().device_ptr(&s);
17064 let (vp, _g1) = rank_cache.v().device_ptr(&s);
17065 let (lp, _g2) = rank_cache.len_d().device_ptr(&s);
17066 let bp = match rank_cache.base_d() {
17067 Some(b) => {
17068 let (p, _g) = b.device_ptr(&s);
17069 p as u64
17070 }
17071 None => 0u64,
17072 };
17073 let mut host = Vec::with_capacity(t * 6);
17074 for r in 0..t {
17075 host.extend_from_slice(&[
17076 kp as u64,
17077 vp as u64,
17078 lp as u64,
17079 bp,
17080 0u64,
17081 (t - 1 - r) as u64,
17082 ]);
17083 }
17084 tables.push(engine.stream().clone_htod(&host)?);
17085 }
17086 let tabs: Vec<&CudaSlice<u64>> = tables.iter().collect();
17087 let ws_index = tp
17088 .runtime
17089 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
17090 tp.runtime.decode_v2_fa_rows_join(
17091 ws_index,
17092 e,
17093 &tp.o,
17094 &tabs,
17095 t,
17096 head_dim,
17097 window.unwrap_or(0),
17098 max_ns,
17099 geometry.attention_scale(),
17100 k_tok_bytes,
17101 v_tok_bytes,
17102 )
17103 }
17104
17105 pub(crate) fn step35_batch_fa_rows_precheck(
17109 &self,
17110 caches: &[&mut Cache],
17111 row_to_cache: impl Fn(usize) -> usize,
17112 positions: &[i32],
17113 il: usize,
17114 ) -> Result<bool, Box<dyn std::error::Error>> {
17115 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
17116 return Ok(false);
17117 };
17118 let Some(tp) = fa.step_tp_qkv.as_ref() else {
17119 return Ok(false);
17120 };
17121 let Some(attention) = tp.attention.as_ref() else {
17122 return Ok(false);
17123 };
17124 if !tp.runtime.native_p2p()
17125 || crate::Engine::kv_fp8_on()
17126 || !crate::tp::step_tp_dcw_enabled()?
17127 || !crate::tp::step_tp_qkv_fused_enabled()?
17128 || (attention.gate_shards.is_none() && attention.gate_shards_bf16.is_none())
17129 {
17130 return Ok(false);
17131 }
17132 let geometry = self.step35_geom(il);
17133 let head_dim = geometry.head_dim_k as usize;
17134 if head_dim > 256 || head_dim % 32 != 0 || !crate::fa_v3_on() {
17135 return Ok(false);
17136 }
17137 if crate::fa_sm_count() < 128
17138 || std::env::var("MEMRA_FA_SPLIT").is_ok()
17139 || std::env::var("MEMRA_FA_SP_SHORT").is_ok()
17140 || std::env::var("MEMRA_FA_SP16").is_ok()
17141 || std::env::var("MEMRA_NO_FA_VEC").is_ok()
17142 {
17143 return Ok(false);
17144 }
17145 let window = geometry.window.map(|w| w as usize);
17146 for (r, &pos) in positions.iter().enumerate() {
17147 let cache = &caches[row_to_cache(r)];
17148 let Some(distributed) = cache.tp_kv[il].as_ref() else {
17149 return Ok(false);
17150 };
17151 if distributed.staged_len() != pos as usize {
17152 return Ok(false);
17153 }
17154 if distributed.peek_append_ring(1)?.1 {
17155 return Ok(false);
17156 }
17157 let t0 = window
17158 .map(|w| (pos as usize + 1).min(w))
17159 .unwrap_or(pos as usize + 1);
17160 if t0 < 96 || t0 < crate::fa_vec_min_tkv() {
17161 return Ok(false);
17162 }
17163 }
17164 Ok(true)
17165 }
17166
17167 pub(crate) fn step35_verify_rope_fa_pass(
17173 &self,
17174 e: &Engine,
17175 il: usize,
17176 cache: &Cache,
17177 pos0: usize,
17178 t: usize,
17179 stage_pos: bool,
17180 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
17181 use cudarc::driver::DevicePtr;
17182 if !crate::tp::fuse_rope_append_on() {
17183 return Ok(None);
17184 }
17185 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
17186 return Ok(None);
17187 };
17188 let Some(tp) = fa.step_tp_qkv.as_ref() else {
17189 return Ok(None);
17190 };
17191 let Some(attention) = tp.attention.as_ref() else {
17192 return Ok(None);
17193 };
17194 let geometry = self.step35_geom(il);
17195 let head_dim = geometry.head_dim_k as usize;
17196 if head_dim != 128 {
17197 return Ok(None);
17198 }
17199 let heads = geometry.n_head as usize;
17200 let window = geometry.window.map(|w| w as usize);
17201 let ranks = tp.runtime.devices().len();
17202 let Some(distributed) = cache.tp_kv[il].as_ref() else {
17203 return Ok(None);
17204 };
17205 if distributed.kv_dim_k() != distributed.kv_dim_v() {
17206 return Ok(None);
17207 }
17208 {
17209 let rank0 = distributed.rank(0).ok_or("verify rope pass lost rank 0")?;
17210 if rank0.base_d().is_none()
17211 && distributed.staged_len() + t > distributed.physical_capacity()
17212 {
17213 return Ok(None);
17214 }
17215 }
17216 let mut rope_freqs = Vec::with_capacity(ranks);
17217 for rank in 0..ranks {
17218 let engine = tp
17219 .runtime
17220 .rank_engine(rank)
17221 .ok_or("verify rope pass lost a rank engine")?;
17222 rope_freqs.push(if geometry.rope_factors {
17223 match self
17224 .step35_aux
17225 .as_ref()
17226 .and_then(|aux| aux.rope_freqs(engine))
17227 {
17228 Some(f) => Some(f),
17229 None => return Ok(None),
17230 }
17231 } else {
17232 None
17233 });
17234 }
17235 let ladder = |t_kv: usize| -> usize {
17236 if t_kv <= 2048 {
17237 16
17238 } else if t_kv <= 16384 {
17239 64
17240 } else {
17241 128
17242 }
17243 };
17244 let (k_tok_bytes, v_tok_bytes) = (distributed.k_tok_bytes(), distributed.v_tok_bytes());
17245 let mut max_ns = 1usize;
17246 let mut positions = Vec::with_capacity(t);
17247 for r in 0..t {
17248 positions.push((pos0 + r) as i32);
17249 let t_kv = window
17250 .map(|w| (pos0 + r + 1).min(w))
17251 .unwrap_or(pos0 + r + 1);
17252 max_ns = max_ns.max(t_kv.div_ceil(ladder(t_kv)));
17253 }
17254 let mut session_parts: Vec<Vec<[u64; 4]>> = vec![Vec::with_capacity(t); ranks];
17255 let mut tab_keys = vec![0u64; ranks];
17256 for rank in 0..ranks {
17257 let engine = tp
17258 .runtime
17259 .rank_engine(rank)
17260 .ok_or("verify rope pass lost a rank engine")?;
17261 let rank_cache = distributed
17262 .rank(rank)
17263 .ok_or("verify rope pass lost a KV cache rank")?;
17264 let _main = engine.gpu.enter_main()?;
17265 let s = engine.stream();
17266 let (kp, _g0) = rank_cache.k().device_ptr(&s);
17267 let (vp, _g1) = rank_cache.v().device_ptr(&s);
17268 let (lp, _g2) = rank_cache.len_d().device_ptr(&s);
17269 let bp = match rank_cache.base_d() {
17270 Some(b) => {
17271 let (p, _g) = b.device_ptr(&s);
17272 p as u64
17273 }
17274 None => 0u64,
17275 };
17276 tab_keys[rank] = (kp as u64)
17277 .rotate_left(17)
17278 .wrapping_add(bp)
17279 .wrapping_add((il as u64) << 32)
17280 .wrapping_add(t as u64)
17281 .wrapping_add(1 << 63);
17282 for _r in 0..t {
17283 session_parts[rank].push([kp as u64, vp as u64, lp as u64, bp]);
17284 }
17285 }
17286 let ws_index = tp
17287 .runtime
17288 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
17289 tp.runtime
17290 .decode_v2_rope_fa_rows(
17291 ws_index,
17292 e,
17293 &tp.o,
17294 &session_parts,
17295 &tab_keys,
17296 &positions,
17297 stage_pos,
17298 true,
17299 &attention.q_norm,
17300 &attention.k_norm,
17301 &rope_freqs,
17302 t,
17303 head_dim,
17304 geometry.n_rot as usize,
17305 window.unwrap_or(0),
17306 max_ns,
17307 geometry.attention_scale(),
17308 k_tok_bytes,
17309 v_tok_bytes,
17310 self.cfg.rms_eps,
17311 geometry.rope_base,
17312 )
17313 .map(Some)
17314 }
17315
17316 #[allow(clippy::too_many_arguments)]
17321 pub(crate) fn step35_batch_rope_fa_pass(
17322 &self,
17323 e: &Engine,
17324 il: usize,
17325 caches: &[&mut Cache],
17326 row_to_cache: impl Fn(usize) -> usize,
17327 positions: &[i32],
17328 t: usize,
17329 stage_pos: bool,
17330 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
17331 use cudarc::driver::DevicePtr;
17332 if !crate::tp::fuse_rope_append_on() {
17333 return Ok(None);
17334 }
17335 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
17336 return Ok(None);
17337 };
17338 let Some(tp) = fa.step_tp_qkv.as_ref() else {
17339 return Ok(None);
17340 };
17341 let Some(attention) = tp.attention.as_ref() else {
17342 return Ok(None);
17343 };
17344 let geometry = self.step35_geom(il);
17345 let head_dim = geometry.head_dim_k as usize;
17346 if head_dim != 128 {
17347 return Ok(None);
17348 }
17349 let heads = geometry.n_head as usize;
17350 let window = geometry.window.map(|w| w as usize);
17351 let ranks = tp.runtime.devices().len();
17352 for r in 0..t {
17355 let cache = &caches[row_to_cache(r)];
17356 let Some(distributed) = cache.tp_kv[il].as_ref() else {
17357 return Ok(None);
17358 };
17359 if distributed.kv_dim_k() != distributed.kv_dim_v() {
17360 return Ok(None);
17361 }
17362 let rank0 = distributed.rank(0).ok_or("rope fa pass lost rank 0")?;
17363 if rank0.base_d().is_none()
17364 && distributed.staged_len() + t > distributed.physical_capacity()
17365 {
17366 return Ok(None);
17367 }
17368 }
17369 let mut rope_freqs = Vec::with_capacity(ranks);
17370 for rank in 0..ranks {
17371 let engine = tp
17372 .runtime
17373 .rank_engine(rank)
17374 .ok_or("rope fa pass lost a rank engine")?;
17375 rope_freqs.push(if geometry.rope_factors {
17376 match self
17377 .step35_aux
17378 .as_ref()
17379 .and_then(|aux| aux.rope_freqs(engine))
17380 {
17381 Some(f) => Some(f),
17382 None => return Ok(None),
17383 }
17384 } else {
17385 None
17386 });
17387 }
17388 let ladder = |t_kv: usize| -> usize {
17389 if t_kv <= 2048 {
17390 16
17391 } else if t_kv <= 16384 {
17392 64
17393 } else {
17394 128
17395 }
17396 };
17397 let (mut max_ns, mut k_tok_bytes, mut v_tok_bytes) = (1usize, 0usize, 0usize);
17398 let mut session_parts: Vec<Vec<[u64; 4]>> = vec![Vec::with_capacity(t); ranks];
17399 let mut tab_keys = vec![0u64; ranks];
17400 for (r, &pos) in positions.iter().enumerate().take(t) {
17401 let cache = &caches[row_to_cache(r)];
17402 let distributed = cache.tp_kv[il]
17403 .as_ref()
17404 .ok_or("rope fa pass lost a distributed KV cache")?;
17405 (k_tok_bytes, v_tok_bytes) = (distributed.k_tok_bytes(), distributed.v_tok_bytes());
17406 let t_kv = window
17407 .map(|w| (pos as usize + 1).min(w))
17408 .unwrap_or(pos as usize + 1);
17409 max_ns = max_ns.max(t_kv.div_ceil(ladder(t_kv)));
17410 for rank in 0..ranks {
17411 let engine = tp
17412 .runtime
17413 .rank_engine(rank)
17414 .ok_or("rope fa pass lost a rank engine")?;
17415 let rank_cache = distributed
17416 .rank(rank)
17417 .ok_or("rope fa pass lost a KV cache rank")?;
17418 let _main = engine.gpu.enter_main()?;
17419 let s = engine.stream();
17420 let (kp, _g0) = rank_cache.k().device_ptr(&s);
17421 let (vp, _g1) = rank_cache.v().device_ptr(&s);
17422 let (lp, _g2) = rank_cache.len_d().device_ptr(&s);
17423 let bp = match rank_cache.base_d() {
17424 Some(b) => {
17425 let (p, _g) = b.device_ptr(&s);
17426 p as u64
17427 }
17428 None => 0u64,
17429 };
17430 tab_keys[rank] = tab_keys[rank]
17431 .rotate_left(9)
17432 .wrapping_add(kp as u64)
17433 .wrapping_add(bp)
17434 .wrapping_add(il as u64);
17435 session_parts[rank].push([kp as u64, vp as u64, lp as u64, bp]);
17436 }
17437 }
17438 let ws_index = tp
17439 .runtime
17440 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
17441 tp.runtime
17442 .decode_v2_rope_fa_rows(
17443 ws_index,
17444 e,
17445 &tp.o,
17446 &session_parts,
17447 &tab_keys,
17448 positions,
17449 stage_pos,
17450 false,
17451 &attention.q_norm,
17452 &attention.k_norm,
17453 &rope_freqs,
17454 t,
17455 head_dim,
17456 geometry.n_rot as usize,
17457 window.unwrap_or(0),
17458 max_ns,
17459 geometry.attention_scale(),
17460 k_tok_bytes,
17461 v_tok_bytes,
17462 self.cfg.rms_eps,
17463 geometry.rope_base,
17464 )
17465 .map(Some)
17466 }
17467
17468 #[allow(clippy::too_many_arguments)]
17472 pub(crate) fn step35_batch_fa_rows_join(
17473 &self,
17474 e: &Engine,
17475 il: usize,
17476 caches: &[&mut Cache],
17477 row_to_cache: impl Fn(usize) -> usize,
17478 positions: &[i32],
17479 t: usize,
17480 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
17481 use cudarc::driver::DevicePtr;
17482 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
17483 return Err("batch fa rows join expects full attention".into());
17484 };
17485 let tp = fa
17486 .step_tp_qkv
17487 .as_ref()
17488 .ok_or("batch fa rows join lost its resident projections")?;
17489 let geometry = self.step35_geom(il);
17490 let heads = geometry.n_head as usize;
17491 let head_dim = geometry.head_dim_k as usize;
17492 let window = geometry.window.map(|w| w as usize);
17493 let ladder = |t_kv: usize| -> usize {
17494 if t_kv <= 2048 {
17495 16
17496 } else if t_kv <= 16384 {
17497 64
17498 } else {
17499 128
17500 }
17501 };
17502 let (mut max_ns, mut k_tok_bytes, mut v_tok_bytes) = (1usize, 0usize, 0usize);
17503 for (r, &pos) in positions.iter().enumerate() {
17504 let cache = &caches[row_to_cache(r)];
17505 let distributed = cache.tp_kv[il]
17506 .as_ref()
17507 .ok_or("batch fa rows join lost a distributed KV cache")?;
17508 (k_tok_bytes, v_tok_bytes) = (distributed.k_tok_bytes(), distributed.v_tok_bytes());
17509 let t_kv = window
17510 .map(|w| (pos as usize + 1).min(w))
17511 .unwrap_or(pos as usize + 1);
17512 max_ns = max_ns.max(t_kv.div_ceil(ladder(t_kv)));
17513 }
17514 let ranks = tp.runtime.devices().len();
17518 let mut tables = Vec::with_capacity(ranks);
17519 for rank in 0..ranks {
17520 let engine = tp
17521 .runtime
17522 .rank_engine(rank)
17523 .ok_or("batch fa rows join lost a rank engine")?;
17524 let _main = engine.gpu.enter_main()?;
17525 let s = engine.stream();
17526 let mut host = Vec::with_capacity(t * 6);
17527 for r in 0..t {
17528 let cache = &caches[row_to_cache(r)];
17529 let distributed = cache.tp_kv[il]
17530 .as_ref()
17531 .ok_or("batch fa rows join lost a distributed KV cache")?;
17532 let rank_cache = distributed
17533 .rank(rank)
17534 .ok_or("batch fa rows join lost a KV cache rank")?;
17535 let (kp, _g0) = rank_cache.k().device_ptr(&s);
17536 let (vp, _g1) = rank_cache.v().device_ptr(&s);
17537 let (lp, _g2) = rank_cache.len_d().device_ptr(&s);
17538 let bp = match rank_cache.base_d() {
17539 Some(b) => {
17540 let (p, _g) = b.device_ptr(&s);
17541 p as u64
17542 }
17543 None => 0u64,
17544 };
17545 host.extend_from_slice(&[kp as u64, vp as u64, lp as u64, bp, 0u64, 0u64]);
17546 }
17547 tables.push(engine.stream().clone_htod(&host)?);
17548 }
17549 let tabs: Vec<&CudaSlice<u64>> = tables.iter().collect();
17550 let ws_index = tp
17551 .runtime
17552 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
17553 tp.runtime.decode_v2_fa_rows_join(
17554 ws_index,
17555 e,
17556 &tp.o,
17557 &tabs,
17558 t,
17559 head_dim,
17560 window.unwrap_or(0),
17561 max_ns,
17562 geometry.attention_scale(),
17563 k_tok_bytes,
17564 v_tok_bytes,
17565 )
17566 }
17567
17568 pub(crate) fn step35_verify_spec_fa2_join(
17572 &self,
17573 e: &Engine,
17574 il: usize,
17575 cache: &Cache,
17576 pos0: usize,
17577 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
17578 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
17579 return Err("spec fa2 join expects full attention".into());
17580 };
17581 let tp = fa
17582 .step_tp_qkv
17583 .as_ref()
17584 .ok_or("spec fa2 join lost its resident projections")?;
17585 let geometry = self.step35_geom(il);
17586 let heads = geometry.n_head as usize;
17587 let head_dim = geometry.head_dim_k as usize;
17588 let window = geometry.window.map(|w| w as usize);
17589 let bucket = window.map(|w| (pos0 + 2).min(w)).unwrap_or(pos0 + 2);
17592 let distributed = cache.tp_kv[il]
17593 .as_ref()
17594 .ok_or("spec fa2 join lost its distributed KV cache")?;
17595 let ws_index = tp
17596 .runtime
17597 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
17598 tp.runtime.decode_v2_spec_fa2_join(
17599 ws_index,
17600 e,
17601 &tp.o,
17602 distributed,
17603 head_dim,
17604 window.unwrap_or(0),
17605 bucket,
17606 geometry.attention_scale(),
17607 )
17608 }
17609
17610 fn step35_tp_decode_attn_resident_v2(
17611 &self,
17612 e: &Engine,
17613 fa: &FullAttnLayer,
17614 il: usize,
17615 h: &CudaSlice<f32>,
17616 pos_d: &CudaSlice<i32>,
17617 cache: &mut Cache,
17618 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
17619 let tp = fa
17620 .step_tp_qkv
17621 .as_ref()
17622 .ok_or("Step TP decode lost its resident projections")?;
17623 let attention = tp
17624 .attention
17625 .as_ref()
17626 .ok_or("Step TP decode lost its resident attention auxiliaries")?;
17627 if !tp.runtime.native_p2p() {
17628 return Err("rank-local Step attention requires native P2P".into());
17629 }
17630 if crate::Engine::kv_fp8_on() {
17631 return Err("rank-local Step attention has not qualified the FP8 KV cache".into());
17632 }
17633
17634 let geometry = self.step35_geom(il);
17635 let window = geometry.window.map(|window| window as usize);
17636 let ranks = tp.runtime.devices().len();
17637 let head_dim = geometry.head_dim_k as usize;
17638 let heads = geometry.n_head as usize;
17639 let kv_heads = geometry.n_head_kv as usize;
17640 if heads % ranks != 0 || kv_heads % ranks != 0 {
17641 return Err(format!(
17642 "Step attention heads q={heads} kv={kv_heads} are not divisible by TP={ranks}"
17643 )
17644 .into());
17645 }
17646 let local_heads = heads / ranks;
17647 let local_kv_heads = kv_heads / ranks;
17648 let max_ctx = cache.max_ctx;
17649
17650 let hydrated = self.ensure_step_tp_kv_cache(e, fa, il, cache)?;
17651
17652 let base_len = cache.kv[il]
17653 .as_ref()
17654 .ok_or_else(|| format!("Step TP layer {il} lost its local KV cache"))?
17655 .len;
17656 {
17657 let distributed = cache.tp_kv[il]
17658 .as_ref()
17659 .ok_or_else(|| format!("Step TP layer {il} lost its distributed KV cache"))?;
17660 if distributed.committed_len() != base_len || distributed.staged_len() != base_len {
17661 return Err(format!(
17662 "Step TP layer {il} cache lengths diverged before decode: \
17663 local={base_len} distributed={}/{}",
17664 distributed.committed_len(),
17665 distributed.staged_len()
17666 )
17667 .into());
17668 }
17669 }
17670 if pos_d.len() != 1 {
17671 return Err(format!(
17672 "rank-local Step decode requires one position, got {}",
17673 pos_d.len()
17674 )
17675 .into());
17676 }
17677
17678 let decode_input = attention
17679 .decode_input
17680 .as_ref()
17681 .ok_or("Step TP decode v2 requires the replicated decode input")?;
17682 let mut decode_input = decode_input
17683 .lock()
17684 .map_err(|_| "Step TP replicated decode input lock is poisoned")?;
17685
17686 let use_gate_shards = (attention.gate_shards.is_some()
17691 || attention.gate_shards_bf16.is_some())
17692 && crate::tp::step_tp_qkv_fused_enabled()?;
17693 let gate_raw = if use_gate_shards {
17694 None
17695 } else {
17696 let gate_weight = fa
17697 .attn_gate
17698 .as_ref()
17699 .ok_or("step35 layer is missing attn_gate.weight")?;
17700 let gate_raw = e.matmul(gate_weight, h, 1)?;
17701 if gate_raw.len() != heads {
17702 return Err(format!(
17703 "Step TP layer {il} gate output {} != {heads}",
17704 gate_raw.len()
17705 )
17706 .into());
17707 }
17708 Some(gate_raw)
17709 };
17710
17711 let ws_index = tp
17712 .runtime
17713 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
17714 let mut ws_guard = tp
17715 .runtime
17716 .decode_v2_workspace()
17717 .lock()
17718 .map_err(|_| "Step TP decode v2 workspace lock is poisoned")?;
17719 let ws = ws_guard
17720 .get_mut(ws_index)
17721 .ok_or("Step TP decode v2 workspace missing after ensure")?;
17722
17723 let mut rope_freqs = Vec::with_capacity(ranks);
17724 for rank in 0..ranks {
17725 let engine = tp
17726 .runtime
17727 .rank_engine(rank)
17728 .ok_or_else(|| format!("Step TP layer {il} has no engine for rank {rank}"))?;
17729 rope_freqs.push(if geometry.rope_factors {
17730 self.step35_aux
17731 .as_ref()
17732 .and_then(|aux| aux.rope_freqs(engine))
17733 } else {
17734 None
17735 });
17736 }
17737 let staged_next = base_len + 1;
17744 let t_kv_eff = window
17745 .map(|window| staged_next.min(window))
17746 .unwrap_or(staged_next);
17747 let dcw = crate::tp::step_tp_dcw_enabled()? && use_gate_shards && t_kv_eff >= 96 && {
17748 let (write_row, would_rebase) = cache.tp_kv[il]
17749 .as_ref()
17750 .expect("distributed cache checked above")
17751 .peek_append_ring(1)?;
17752 if !would_rebase {
17753 let base = (base_len - write_row) as i32;
17755 let distributed = cache.tp_kv[il]
17756 .as_mut()
17757 .expect("distributed cache checked above");
17758 for rank in 0..ranks {
17759 let engine = tp.runtime.rank_engine(rank).ok_or_else(|| {
17760 format!("Step TP layer {il} has no engine for rank {rank}")
17761 })?;
17762 let _main = engine.gpu.enter_main()?;
17763 let rank_cache = distributed
17764 .rank_mut(rank)
17765 .ok_or_else(|| format!("Step TP layer {il} has no KV cache rank {rank}"))?;
17766 if rank_cache.base_d().is_none() {
17767 rank_cache.arm_base_d(engine.htod_i32(&[base])?);
17768 }
17769 }
17770 }
17771 !would_rebase
17772 };
17773 let fuse_rope = dcw
17774 && crate::tp::fuse_rope_append_on()
17775 && head_dim == 128
17776 && cache.tp_kv[il]
17777 .as_ref()
17778 .map(|d| d.kv_dim_k() == d.kv_dim_v() && d.kv_dim_k() == local_kv_heads * head_dim)
17779 .unwrap_or(false);
17780
17781 let tcol_col = crate::tp::take_verify_tcol();
17782 let fa2_col = crate::tp::take_spec_fa2_defer();
17789 tp.runtime.decode_v2_input_qkv(
17790 ws,
17791 e,
17792 h,
17793 pos_d,
17794 gate_raw.as_ref(),
17795 if !use_gate_shards {
17796 None
17797 } else if let Some(shards) = attention.gate_shards.as_deref() {
17798 Some(crate::tp::StepTpGateShards::F32(shards))
17799 } else {
17800 attention
17801 .gate_shards_bf16
17802 .as_deref()
17803 .map(crate::tp::StepTpGateShards::Bf16)
17804 },
17805 &mut decode_input,
17806 &tp.q,
17807 &tp.k,
17808 &tp.v,
17809 &attention.q_norm,
17810 &attention.k_norm,
17811 head_dim,
17812 geometry.n_rot as usize,
17813 geometry.rope_base,
17814 &rope_freqs,
17815 self.cfg.rms_eps,
17816 fuse_rope,
17817 tcol_col,
17818 )?;
17819
17820 let transaction = cache.tp_kv[il]
17821 .as_mut()
17822 .expect("distributed cache checked above")
17823 .begin_transaction()?;
17824 let append_result = tp.runtime.append_tp_kv_transaction_inner(
17825 cache.tp_kv[il]
17826 .as_mut()
17827 .expect("distributed cache checked above"),
17828 transaction,
17829 &ws.k,
17830 &ws.v_raw,
17831 1,
17832 dcw,
17833 );
17834 if let Err(error) = append_result {
17835 let _ = tp.runtime.rollback_tp_kv_transaction(
17836 cache.tp_kv[il]
17837 .as_mut()
17838 .expect("distributed cache checked above"),
17839 transaction,
17840 );
17841 return Err(error);
17842 }
17843
17844 let staged = (|| -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
17845 let (staged_len, physical, k_tok_bytes_c, v_tok_bytes_c, capacity) = {
17848 let distributed = cache.tp_kv[il]
17849 .as_ref()
17850 .expect("distributed cache checked above");
17851 let staged_len = distributed.staged_len();
17852 let view_start = window
17853 .map(|window| staged_len.saturating_sub(window))
17854 .unwrap_or(0);
17855 (
17856 staged_len,
17857 distributed.physical_range(view_start, staged_len)?,
17858 distributed.k_tok_bytes(),
17859 distributed.v_tok_bytes(),
17860 distributed.physical_capacity(),
17861 )
17862 };
17863 let view_start = window
17864 .map(|window| staged_len.saturating_sub(window))
17865 .unwrap_or(0);
17866 let t_kv = staged_len - view_start;
17867 for rank in 0..ranks {
17868 let engine = tp
17869 .runtime
17870 .rank_engine(rank)
17871 .ok_or_else(|| format!("Step TP layer {il} has no engine for rank {rank}"))?;
17872 let _main = engine.gpu.enter_main()?;
17873 if dcw {
17874 {
17878 let distributed_mut = cache.tp_kv[il]
17879 .as_mut()
17880 .expect("distributed cache checked above");
17881 let (kv_dim_k, kv_dim_v) =
17882 (distributed_mut.kv_dim_k(), distributed_mut.kv_dim_v());
17883 let (k_tok_bytes, v_tok_bytes) =
17884 (distributed_mut.k_tok_bytes(), distributed_mut.v_tok_bytes());
17885 let rank_cache = distributed_mut.rank_mut(rank).ok_or_else(|| {
17886 format!("Step TP layer {il} has no KV cache rank {rank}")
17887 })?;
17888 let (k_plane, v_plane, len_d, base_d) =
17889 rank_cache.planes_and_counters_mut();
17890 if fuse_rope {
17891 let same_dev = engine.ctx().ordinal() == e.ctx().ordinal();
17894 let crate::tp::StepTpDecodeV2Ws {
17895 q_raw,
17896 k_raw,
17897 v_raw,
17898 q,
17899 k,
17900 pos,
17901 pos_stage,
17902 fuse_ctr,
17903 ..
17904 } = &mut *ws;
17905 let pos_ref: &CudaSlice<i32> = if same_dev {
17909 pos_stage
17910 .as_ref()
17911 .ok_or("step TP decode v2 pos stage not armed")?
17912 } else {
17913 &pos[rank]
17914 };
17915 engine.qk_norm_rope_append_inc_dcw(
17916 &q_raw[rank],
17917 &k_raw[rank],
17918 &v_raw[rank],
17919 &attention.q_norm[rank],
17920 &attention.k_norm[rank],
17921 &mut q[rank],
17922 &mut k[rank],
17923 pos_ref,
17924 k_plane,
17925 v_plane,
17926 len_d,
17927 base_d,
17928 &mut fuse_ctr[rank],
17929 kv_dim_k,
17930 kv_dim_v,
17931 k_tok_bytes,
17932 v_tok_bytes,
17933 head_dim,
17934 geometry.n_rot as usize,
17935 local_heads,
17936 local_kv_heads,
17937 self.cfg.rms_eps,
17938 geometry.rope_base,
17939 1.0,
17940 rope_freqs[rank],
17941 )?;
17942 } else {
17943 engine.append_kv_quantized_dcw(
17944 &ws.k[rank],
17945 &ws.v_raw[rank],
17946 k_plane,
17947 v_plane,
17948 len_d,
17949 base_d,
17950 kv_dim_k,
17951 kv_dim_v,
17952 k_tok_bytes,
17953 v_tok_bytes,
17954 )?;
17955 }
17956 if !fuse_rope {
17957 let rank_cache = distributed_mut.rank_mut(rank).ok_or_else(|| {
17958 format!("Step TP layer {il} has no KV cache rank {rank}")
17959 })?;
17960 engine.inc_i32(rank_cache.len_d_mut())?;
17961 }
17962 }
17963 let distributed = cache.tp_kv[il]
17964 .as_ref()
17965 .expect("distributed cache checked above");
17966 let rank_cache = distributed
17967 .rank(rank)
17968 .ok_or_else(|| format!("Step TP layer {il} has no KV cache rank {rank}"))?;
17969 let k_ring = engine.view_u8_range(rank_cache.k(), 0, capacity * k_tok_bytes_c);
17970 let v_ring = engine.view_u8_range(rank_cache.v(), 0, capacity * v_tok_bytes_c);
17971 if fa2_col.is_some() {
17972 continue;
17975 }
17976 {
17977 let crate::tp::StepTpDecodeV2Ws { q, gate, gated, .. } = &mut *ws;
17980 engine.fa_decode_dcw(
17981 &q[rank],
17982 &k_ring,
17983 &v_ring,
17984 &mut gated[rank],
17985 head_dim,
17986 local_heads,
17987 local_kv_heads,
17988 rank_cache.len_d(),
17989 rank_cache.base_d(),
17990 window.unwrap_or(0),
17991 t_kv,
17992 geometry.attention_scale(),
17993 k_tok_bytes_c,
17994 v_tok_bytes_c,
17995 Some(&gate[rank]),
17996 )?;
17997 }
17998 continue;
17999 }
18000 let distributed = cache.tp_kv[il]
18001 .as_ref()
18002 .expect("distributed cache checked above");
18003 let rank_cache = distributed
18004 .rank(rank)
18005 .ok_or_else(|| format!("Step TP layer {il} has no KV cache rank {rank}"))?;
18006 let k_view = engine.view_u8_range(
18007 rank_cache.k(),
18008 physical.start * k_tok_bytes_c,
18009 physical.end * k_tok_bytes_c,
18010 );
18011 let v_view = engine.view_u8_range(
18012 rank_cache.v(),
18013 physical.start * v_tok_bytes_c,
18014 physical.end * v_tok_bytes_c,
18015 );
18016 engine.fa_decode_kvmod(
18017 &ws.q[rank],
18018 &k_view,
18019 &v_view,
18020 &mut ws.attn_out[rank],
18021 head_dim,
18022 local_heads,
18023 local_kv_heads,
18024 t_kv,
18025 geometry.attention_scale(),
18026 k_tok_bytes_c,
18027 v_tok_bytes_c,
18028 false,
18029 )?;
18030 engine.attn_head_gate(
18031 &ws.attn_out[rank],
18032 &ws.gate[rank],
18033 &mut ws.gated[rank],
18034 None,
18035 head_dim,
18036 local_heads,
18037 1,
18038 )?;
18039 }
18040
18041 let output = if let Some(col) = fa2_col.filter(|_| dcw) {
18048 tp.runtime.decode_v2_stash_fa2(ws, e, col)?;
18052 crate::tp::set_spec_fa2_stashed();
18053 e.uninit(ws.o_out)?
18054 } else if let Some(col) = crate::tp::take_tcol_oproj_defer() {
18055 if tp.runtime.decode_v2_oproj_tcol_eligible(ws, &tp.o) {
18056 tp.runtime.decode_v2_stash_gated(ws, e, col)?;
18057 crate::tp::set_tcol_oproj_stashed();
18058 e.uninit(ws.o_out)?
18059 } else {
18060 tp.runtime.decode_v2_finish(ws, e, &tp.o)?
18061 }
18062 } else {
18063 tp.runtime.decode_v2_finish(ws, e, &tp.o)?
18064 };
18065
18066 let local = cache.kv[il]
18070 .as_mut()
18071 .ok_or_else(|| format!("Step TP layer {il} lost its local KV cache"))?;
18072 if local.len != base_len || base_len + 1 > max_ctx {
18073 return Err(format!(
18074 "Step TP layer {il} local cache changed during decode: \
18075 len={} base={base_len} max={max_ctx}",
18076 local.len
18077 )
18078 .into());
18079 }
18080 if crate::tp::no_local_shadow_on() {
18081 local.len = base_len + 1;
18084 if !crate::tp::len_mirror_lazy_on() {
18088 e.set_i32_one(&mut local.len_d, local.len as i32)?;
18089 }
18090 } else {
18091 let retain_from = window
18092 .map(|window| {
18093 let staged_retain = (base_len + 1).saturating_sub(window) & !31usize;
18094 let rollback_retain =
18095 base_len.saturating_sub(window.saturating_sub(1)) & !31usize;
18096 staged_retain.min(rollback_retain)
18097 })
18098 .unwrap_or(0);
18099 let write_row = e.prepare_kv_append(local, retain_from, 1)?;
18100 e.append_kv_quantized(
18101 &ws.k_shadow,
18102 &ws.v_shadow,
18103 &mut local.k,
18104 &mut local.v,
18105 write_row,
18106 local.kv_dim_k,
18107 local.kv_dim_v,
18108 local.k_tok_bytes,
18109 local.v_tok_bytes,
18110 false,
18111 )?;
18112 local.len = base_len + 1;
18113 e.set_i32_one(&mut local.len_d, local.len as i32)?;
18114 }
18115 Ok(output)
18116 })();
18117
18118 let output = match staged {
18119 Ok(output) => output,
18120 Err(error) => {
18121 let _ = tp.runtime.rollback_tp_kv_transaction(
18122 cache.tp_kv[il]
18123 .as_mut()
18124 .expect("distributed cache checked above"),
18125 transaction,
18126 );
18127 if let Some(local) = cache.kv[il].as_mut() {
18128 local.len = base_len;
18129 let _ = e.set_i32_one(&mut local.len_d, base_len as i32);
18130 }
18131 return Err(error);
18132 }
18133 };
18134 let lazy_commit = fuse_rope && crate::tp::len_mirror_lazy_on();
18139 if lazy_commit {
18140 if let Err(error) = tp.runtime.commit_tp_kv_transaction_external(
18141 cache.tp_kv[il]
18142 .as_mut()
18143 .expect("distributed cache checked above"),
18144 transaction,
18145 1,
18146 ) {
18147 let _ = tp.runtime.rollback_tp_kv_transaction(
18148 cache.tp_kv[il]
18149 .as_mut()
18150 .expect("distributed cache checked above"),
18151 transaction,
18152 );
18153 let local = cache.kv[il].as_mut().expect("local cache checked above");
18154 local.len = base_len;
18155 e.set_i32_one(&mut local.len_d, base_len as i32)?;
18156 return Err(error);
18157 }
18158 } else if let Err(error) = tp.runtime.commit_tp_kv_transaction(
18159 cache.tp_kv[il]
18160 .as_mut()
18161 .expect("distributed cache checked above"),
18162 transaction,
18163 1,
18164 ) {
18165 let _ = tp.runtime.rollback_tp_kv_transaction(
18166 cache.tp_kv[il]
18167 .as_mut()
18168 .expect("distributed cache checked above"),
18169 transaction,
18170 );
18171 let local = cache.kv[il].as_mut().expect("local cache checked above");
18172 local.len = base_len;
18173 e.set_i32_one(&mut local.len_d, base_len as i32)?;
18174 return Err(error);
18175 }
18176
18177 let committed = cache.tp_kv[il]
18178 .as_ref()
18179 .expect("distributed cache checked above")
18180 .committed_len();
18181 let local_len = cache.kv[il]
18182 .as_ref()
18183 .expect("local cache checked above")
18184 .len;
18185 if committed != local_len {
18186 return Err(format!(
18187 "Step TP layer {il} committed cache length {committed} != local shadow {local_len}"
18188 )
18189 .into());
18190 }
18191 static V2_LOGGED: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false);
18192 if !V2_LOGGED.swap(true, std::sync::atomic::Ordering::Relaxed) {
18193 eprintln!(
18194 "[step-tp-attn-v2] execute layer={} devices={:?} tokens=1 driver=v2 \
18195 qkv_tensor_parallel=true qk_norm_rank_local=true rope_rank_local=true \
18196 kv_cache_distributed=true kv_cache_hydrated={hydrated} \
18197 attention_tensor_parallel=true attention_scope={} \
18198 input_path=root-device-replicated gate_tensor_parallel=false \
18199 gate_shards=device-staged o_tensor_parallel=true o_reduce=root-device \
18200 local_cache_shadow=true cache_commit=immediate transport={} native_p2p=true \
18201 bulk_p2p={} workspace=persistent ordering=evented output=e-device \
18202 performance_claim=false (logged once; every decode layer runs this driver)",
18203 tp.layer,
18204 tp.devices,
18205 if window.is_some() {
18206 "rank-local-swa-ring"
18207 } else {
18208 "rank-local-global"
18209 },
18210 tp.runtime.transport_label(),
18211 tp.runtime.bulk_p2p(),
18212 );
18213 }
18214 Ok(output)
18215 }
18216
18217 #[allow(clippy::too_many_arguments)]
18227 pub(crate) fn step35_decode_attn(
18228 &self,
18229 e: &Engine,
18230 fa: &FullAttnLayer,
18231 il: usize,
18232 h: &CudaSlice<f32>,
18233 pre_q: Option<(&CudaSlice<i8>, &CudaSlice<f32>)>,
18234 pos_d: &CudaSlice<i32>,
18235 cache: &mut Cache,
18236 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
18237 if fa
18238 .step_tp_qkv
18239 .as_ref()
18240 .is_some_and(|tp| tp.attention.is_some())
18241 {
18242 if pre_q.is_some() {
18243 return Err(
18244 "rank-local Step attention preserves BF16 activations and refuses the q8_1 \
18245 pre-quantized decode path"
18246 .into(),
18247 );
18248 }
18249 return self.step35_tp_decode_attn_resident(e, fa, il, h, pos_d, cache);
18250 }
18251
18252 let geometry = self.step35_geom(il);
18253 let hd = geometry.head_dim_k as usize;
18254 let nkv = geometry.n_head_kv as usize;
18255 let nh = geometry.n_head as usize;
18256 let rbase = geometry.rope_base;
18257 let scale = geometry.attention_scale();
18258 let swa = geometry.window.is_some();
18259 let eps = self.cfg.rms_eps;
18260 let win = geometry.window.unwrap_or(0) as usize;
18261 let n_rot = geometry.n_rot as usize;
18262 let n_embd = self.cfg.n_embd as usize;
18263 let gw = fa
18264 .attn_gate
18265 .as_ref()
18266 .ok_or("step35 layer is missing attn_gate.weight (head-wise attention gate)")?;
18267
18268 let tp_qkv = if fa.step_tp_qkv.is_some() {
18269 if pre_q.is_some() {
18270 return Err(
18271 "Step Q/K/V TP preserves BF16 activations and refuses the q8_1 \
18272 pre-quantized decode path"
18273 .into(),
18274 );
18275 }
18276 self.step35_tp_qkv(e, fa, h, 1)?
18277 } else {
18278 None
18279 };
18280
18281 let (q0, k0, v0, gt) = match tp_qkv {
18282 Some(mut g3) => {
18283 let v = g3.pop().unwrap();
18284 let k = g3.pop().unwrap();
18285 let q = g3.pop().unwrap();
18286 let gt = e.matmul(gw, h, 1)?;
18287 (q, k, v, gt)
18288 }
18289 None => match pre_q {
18290 Some((hq, hdq)) => {
18291 debug_assert!(
18292 e.uses_q8_1_fast(gw),
18293 "step35 pre-quantized decode requires attn_gate on the q8_1 fast path \
18294 (h is a zero-length placeholder here) — see mixer_in_q8_1_fast"
18295 );
18296 let (a, b, c) = match e.matmul_q8_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)? {
18297 Some(t3) => t3,
18298 None => (
18299 e.matmul_pre(&fa.wq, hq, hdq, h, 1)?,
18300 e.matmul_pre(&fa.wk, hq, hdq, h, 1)?,
18301 e.matmul_pre(&fa.wv, hq, hdq, h, 1)?,
18302 ),
18303 };
18304 let gt = e.matmul_pre(gw, hq, hdq, h, 1)?;
18305 (a, b, c, gt)
18306 }
18307 None => {
18308 if e.uses_q8_1_fast(&fa.wq)
18309 && e.uses_q8_1_fast(&fa.wk)
18310 && e.uses_q8_1_fast(&fa.wv)
18311 && e.uses_q8_1_fast(gw)
18312 {
18313 let (hq, hdq) = e.quantize_q8_1(h, 1, n_embd)?;
18314 let (a, b, c) =
18315 match e.matmul_q8_fused3(&fa.wq, &fa.wk, &fa.wv, &hq, &hdq)? {
18316 Some(t3) => t3,
18317 None => (
18318 e.matmul_pre(&fa.wq, &hq, &hdq, h, 1)?,
18319 e.matmul_pre(&fa.wk, &hq, &hdq, h, 1)?,
18320 e.matmul_pre(&fa.wv, &hq, &hdq, h, 1)?,
18321 ),
18322 };
18323 let gt = e.matmul_pre(gw, &hq, &hdq, h, 1)?;
18324 (a, b, c, gt)
18325 } else {
18326 (
18327 e.matmul(&fa.wq, h, 1)?,
18328 e.matmul(&fa.wk, h, 1)?,
18329 e.matmul(&fa.wv, h, 1)?,
18330 e.matmul(gw, h, 1)?,
18331 )
18332 }
18333 }
18334 },
18335 };
18336
18337 let mut q = e.uninit(nh * hd)?;
18338 e.rms_norm(&q0, fa.q_norm.float_data(), &mut q, hd, nh, eps)?;
18339 let mut k = e.uninit(nkv * hd)?;
18340 e.rms_norm(&k0, fa.k_norm.float_data(), &mut k, hd, nkv, eps)?;
18341 let ff = if swa {
18342 None
18343 } else {
18344 self.step35_aux.as_ref().and_then(|a| a.rope_freqs(e))
18345 };
18346 #[cfg(debug_assertions)]
18347 if let Some(ff) = ff {
18348 crate::debug_assert_tensor_stream_device(
18349 ff,
18350 &e.stream(),
18351 "step35_decode_attn.rope_freqs",
18352 );
18353 }
18354 e.rope_neox2(&mut q, &mut k, pos_d, hd, n_rot, nh, nkv, 1, rbase, 1.0, ff)?;
18355
18356 if std::env::var("MEMRA_NOFA").is_ok() {
18357 return Err(
18358 "MEMRA_NOFA (naive f32 SDPA) is incompatible with the quantized KV \
18359 cache; unset MEMRA_NOFA to use fa_decode"
18360 .into(),
18361 );
18362 }
18363 let kvl = cache.kv[il].as_mut().unwrap();
18364 let next_len = kvl.len + 1;
18365 let (off, t_kv) = if swa && next_len > win {
18366 (next_len - win, win)
18367 } else {
18368 (0, next_len)
18369 };
18370 let write_row = e.prepare_kv_append(kvl, off & !31usize, 1)?;
18371 e.append_kv_quantized(
18372 &k,
18373 &v0,
18374 &mut kvl.k,
18375 &mut kvl.v,
18376 write_row,
18377 kvl.kv_dim_k,
18378 kvl.kv_dim_v,
18379 kvl.k_tok_bytes,
18380 kvl.v_tok_bytes,
18381 crate::Engine::kv_fp8_on(),
18382 )?;
18383 kvl.len = next_len;
18384 let physical = kvl.physical_rows(off, off + t_kv)?;
18385 let k_view = e.view_u8_range(
18386 &kvl.k,
18387 physical.start * kvl.k_tok_bytes,
18388 physical.end * kvl.k_tok_bytes,
18389 );
18390 let v_view = e.view_u8_range(
18391 &kvl.v,
18392 physical.start * kvl.v_tok_bytes,
18393 physical.end * kvl.v_tok_bytes,
18394 );
18395 let mut attn = e.uninit(nh * hd)?;
18396 e.fa_decode_kvmod(
18397 &q,
18398 &k_view,
18399 &v_view,
18400 &mut attn,
18401 hd,
18402 nh,
18403 nkv,
18404 t_kv,
18405 scale,
18406 kvl.k_tok_bytes,
18407 kvl.v_tok_bytes,
18408 crate::Engine::kv_fp8_on(),
18409 )?;
18410
18411 let mut ag = e.uninit(nh * hd)?;
18412 e.attn_head_gate(&attn, >, &mut ag, None, hd, nh, 1)?;
18413 self.step35_o(e, fa, &ag, 1)
18414 }
18415}
18416
18417impl HybridModel {
18426 pub fn is_gemma4_e4b(&self) -> bool {
18427 self.gemma4_aux.as_ref().is_some_and(|a| a.e4b.is_some())
18428 }
18429
18430 fn gemma4_e4b_geom(&self, il: usize) -> (usize, usize, usize, f32, f32, bool) {
18434 let g = self.cfg.gemma4.as_ref().unwrap();
18435 let swa = g.swa_pattern[il];
18436 let hd = if swa {
18437 g.key_length_swa
18438 } else {
18439 g.key_length_global
18440 } as usize;
18441 let Mixer::Full(fa) = &self.layers[il].mixer else {
18442 panic!("e4b layer {il} not full-attn")
18443 };
18444 let nh = fa.wq.out_features() / hd;
18445 let nkv = fa.wk.out_features() / hd;
18446 (
18447 hd,
18448 nkv,
18449 nh,
18450 if swa {
18451 g.rope_base_swa
18452 } else {
18453 g.rope_base_global
18454 },
18455 1.0,
18456 swa,
18457 )
18458 }
18459
18460 fn gemma4_e4b_kv_target(&self, il: usize) -> Option<usize> {
18462 self.layers[il]
18463 .gemma4
18464 .as_ref()
18465 .and_then(|b| b.e4b.as_ref())
18466 .and_then(|e4| e4.kv_share.map(|t| t as usize))
18467 }
18468
18469 fn gemma4_e4b_inp_pl(
18474 &self,
18475 e: &Engine,
18476 tokens: &[u32],
18477 x_scaled: &CudaSlice<f32>,
18478 t: usize,
18479 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
18480 let tok_d = e.stream().clone_htod(&tokens.to_vec())?;
18481 self.gemma4_e4b_inp_pl_dev(e, &tok_d, x_scaled, t)
18482 }
18483
18484 fn gemma4_e4b_inp_pl_dev(
18486 &self,
18487 e: &Engine,
18488 tok_d: &CudaSlice<u32>,
18489 x_scaled: &CudaSlice<f32>,
18490 t: usize,
18491 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
18492 let aux = self.gemma4_aux.as_ref().unwrap();
18493 let m = aux.e4b.as_ref().unwrap();
18494 let n_embd = self.cfg.n_embd as usize;
18495 let n_layer = self.layers.len();
18496 let width = m.n_epl * n_layer;
18497 let tbl = m.tok_tbl_gpu.get_or_init(|| {
18498 e.upload_u8(&m.tok_embd_bytes)
18499 .expect("e4b per-layer token table upload")
18500 });
18501 let mut a =
18502 e.embed_gather_device_td(tbl, tok_d, t, width, m.tok_embd_qt, m.tok_embd_row_bytes)?;
18503 e.scale_inplace(&mut a, (m.n_epl as f32).sqrt(), t * width)?;
18504 let mut p = e.matmul(&m.model_proj, x_scaled, t)?;
18505 e.scale_inplace(&mut p, 1.0 / (n_embd as f32).sqrt(), t * width)?;
18506 let mut pn = e.uninit(t * width)?;
18507 e.rms_norm(
18508 &p,
18509 m.proj_norm.float_data(),
18510 &mut pn,
18511 m.n_epl,
18512 t * n_layer,
18513 self.cfg.rms_eps,
18514 )?;
18515 let mut out = e.uninit(t * width)?;
18516 e.add_scale(&a, &pn, 1.0 / 2f32.sqrt(), &mut out, t * width)?;
18517 Ok(out)
18518 }
18519
18520 #[allow(clippy::too_many_arguments)]
18525 fn gemma4_e4b_attn(
18526 &self,
18527 e: &Engine,
18528 il: usize,
18529 hq: &CudaSlice<i8>,
18530 hdq: &CudaSlice<f32>,
18531 pos_d: &CudaSlice<i32>,
18532 t: usize,
18533 cache: &mut Cache,
18534 dc_bucket: Option<usize>,
18535 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
18536 let (hd, nkv, nh, base, scale, swa) = self.gemma4_e4b_geom(il);
18537 let eps = self.cfg.rms_eps;
18538 let aux = self.gemma4_aux.as_ref().unwrap();
18539 let ones = aux.ones(e);
18540 #[cfg(debug_assertions)]
18541 crate::debug_assert_tensor_stream_device(ones, &e.stream(), "gemma4_e4b_attn.ones");
18542 let Mixer::Full(fa) = &self.layers[il].mixer else {
18543 unreachable!()
18544 };
18545 let h0 = e.zeros(0)?;
18549 let h = &h0;
18550
18551 let ff = if swa {
18552 None
18553 } else {
18554 Some(
18555 aux.rope_freqs(e)
18556 .expect("e4b global rope needs rope_freqs.weight"),
18557 )
18558 };
18559 #[cfg(debug_assertions)]
18560 if let Some(ff) = ff {
18561 crate::debug_assert_tensor_stream_device(ff, &e.stream(), "gemma4_e4b_attn.rope_freqs");
18562 }
18563 let share = self.gemma4_e4b_kv_target(il);
18564 let mut kv_f32: Option<(CudaSlice<f32>, CudaSlice<f32>)> = None;
18566 let mut q;
18567 if let Some(_tgt) = share {
18568 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
18569 q = e.uninit(t * nh * hd)?;
18570 let mut kdummy = e.uninit(1)?;
18573 let mut vdummy = e.uninit(1)?;
18574 e.rms_norm_qkv_rope(
18575 &q0,
18576 &q0,
18577 &q0,
18578 fa.q_norm.float_data(),
18579 fa.q_norm.float_data(),
18580 ones,
18581 &mut q,
18582 &mut kdummy,
18583 &mut vdummy,
18584 hd,
18585 self.gemma4_rope_dims(il),
18586 nh * t,
18587 0,
18588 pos_d,
18589 nh,
18590 1,
18591 base,
18592 1.0,
18593 ff,
18594 eps,
18595 )?;
18596 } else {
18597 let e4bits = self.layers[il].gemma4.as_ref().and_then(|g| g.e4b.as_ref());
18601 let cat = e4bits.and_then(|e4| e4.qkv_cat.as_ref());
18602 q = e.uninit(t * nh * hd)?;
18603 let mut k = e.uninit(t * nkv * hd)?;
18604 let mut v = e.uninit(t * nkv * hd)?;
18605 if t == 1 && cat.is_some() {
18606 let qkv0 = e.matmul_pre(cat.unwrap(), hq, hdq, h, 1)?;
18607 e.rms_norm_qkv_rope_cat(
18608 &qkv0,
18609 fa.q_norm.float_data(),
18610 fa.k_norm.float_data(),
18611 ones,
18612 &mut q,
18613 &mut k,
18614 &mut v,
18615 hd,
18616 self.gemma4_rope_dims(il),
18617 nh,
18618 nkv,
18619 pos_d,
18620 nh,
18621 nkv,
18622 base,
18623 1.0,
18624 ff,
18625 eps,
18626 )?;
18627 } else {
18628 let (q0, k0, v0) = match if t == 1 {
18629 e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)?
18630 } else {
18631 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
18634 if *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0")) {
18635 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
18636 } else {
18637 None
18638 }
18639 } {
18640 Some(triple) => triple,
18641 None => (
18642 e.matmul_pre(&fa.wq, hq, hdq, h, t)?,
18643 e.matmul_pre(&fa.wk, hq, hdq, h, t)?,
18644 e.matmul_pre(&fa.wv, hq, hdq, h, t)?,
18645 ), };
18647 e.rms_norm_qkv_rope(
18650 &q0,
18651 &k0,
18652 &v0,
18653 fa.q_norm.float_data(),
18654 fa.k_norm.float_data(),
18655 ones,
18656 &mut q,
18657 &mut k,
18658 &mut v,
18659 hd,
18660 self.gemma4_rope_dims(il),
18661 nh * t,
18662 nkv * t,
18663 pos_d,
18664 nh,
18665 nkv,
18666 base,
18667 1.0,
18668 ff,
18669 eps,
18670 )?;
18671 }
18672 let kvl = cache.kv[il].as_mut().unwrap();
18673 let cls = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
18677 if dc_bucket.is_some() {
18678 debug_assert!(t == 1);
18683 e.append_kv_quantized_row_dc_inc(
18685 &k,
18686 &v,
18687 &mut kvl.k,
18688 &mut kvl.v,
18689 &mut kvl.len_d,
18690 kvl.kv_dim_k,
18691 kvl.kv_dim_v,
18692 kvl.k_tok_bytes,
18693 kvl.v_tok_bytes,
18694 cls,
18695 )?;
18696 } else {
18697 e.append_kv_quantized_rows(
18698 &k,
18699 &v,
18700 &mut kvl.k,
18701 &mut kvl.v,
18702 kvl.len,
18703 t,
18704 kvl.kv_dim_k,
18705 kvl.kv_dim_v,
18706 kvl.k_tok_bytes,
18707 kvl.v_tok_bytes,
18708 cls,
18709 )?;
18710 kvl.len += t;
18711 }
18712 kv_f32 = Some((k, v));
18713 }
18714 let kvl_idx = share.unwrap_or(il);
18717 let kvl = cache.kv[kvl_idx].as_ref().unwrap();
18718 let base_len = kvl.len - t; let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
18720 let mut attn = e.uninit(t * nh * hd)?;
18721 if t > 1 && base_len == 0 && std::env::var("MEMRA_NOFA").is_err() {
18733 if let Some((kf, vf)) = &kv_f32 {
18734 if hd == 256 && t <= win {
18735 e.fa_prefill(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale, true)?;
18736 return Ok(e.matmul(&fa.wo, &attn, t)?);
18737 }
18738 if hd == 256 && swa && t > win {
18739 e.fa_prefill_w(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
18740 return Ok(e.matmul(&fa.wo, &attn, t)?);
18741 }
18742 if hd == 512 && !swa {
18743 e.fa_prefill_hd512(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale, true)?;
18744 return Ok(e.matmul(&fa.wo, &attn, t)?);
18745 }
18746 } else if share.is_some() {
18747 let g = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
18748 let k_view = e.view_u8(&kvl.k, kvl.k.len());
18749 let v_view = e.view_u8(&kvl.v, kvl.v.len());
18750 if hd == 256 && (!swa || t <= win) {
18751 e.fa_prefill_view(
18753 &q,
18754 &k_view,
18755 &v_view,
18756 &mut attn,
18757 hd,
18758 nh,
18759 nkv,
18760 t,
18761 t,
18762 scale,
18763 true,
18764 kvl.k_tok_bytes,
18765 kvl.v_tok_bytes,
18766 g,
18767 )?;
18768 return Ok(e.matmul(&fa.wo, &attn, t)?);
18769 }
18770 let kv_dim = nkv * hd;
18773 let mut kf = e.uninit(t * kv_dim)?;
18774 let mut vf = e.uninit(t * kv_dim)?;
18775 e.fa_dequant_kv_view_f32(
18776 &k_view,
18777 &v_view,
18778 &mut kf,
18779 &mut vf,
18780 kv_dim,
18781 kv_dim,
18782 t,
18783 kvl.k_tok_bytes,
18784 kvl.v_tok_bytes,
18785 g,
18786 )?;
18787 if hd == 512 {
18788 e.fa_prefill_hd512(&q, &kf, &vf, &mut attn, hd, nh, nkv, t, t, scale, true)?;
18789 } else {
18790 e.fa_prefill_w(&q, &kf, &vf, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
18791 }
18792 return Ok(e.matmul(&fa.wo, &attn, t)?);
18793 }
18794 }
18795 if let Some(bucket) = dc_bucket {
18796 assert!(t == 1);
18801 let bucket = if hd == 512 && win <= crate::fa512_min_tkv() {
18807 bucket.min(crate::fa512_min_tkv().saturating_sub(1))
18808 } else {
18809 bucket
18810 };
18811 let k_view = e.view_u8(&kvl.k, kvl.k.len());
18812 let v_view = e.view_u8(&kvl.v, kvl.v.len());
18813 let g = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
18814 if crate::Engine::wpf_level() >= 1 {
18822 e.prefetch_weight_l2(&fa.wo)?;
18823 }
18824 if e.uses_q8_1_fast(&fa.wo) {
18827 let mut oq = e.alloc_i8_uninit(nh * hd)?;
18828 let mut od = e.zeros(nh * hd / 32)?;
18829 e.fa_decode_dc_q8(
18830 &q,
18831 &k_view,
18832 &v_view,
18833 &mut attn,
18834 hd,
18835 nh,
18836 nkv,
18837 &kvl.len_d,
18838 bucket,
18839 scale,
18840 kvl.k_tok_bytes,
18841 kvl.v_tok_bytes,
18842 g,
18843 Some((&mut oq, &mut od)),
18844 )?;
18845 return Ok(e.matmul_pre(&fa.wo, &oq, &od, &attn, t)?);
18846 }
18847 e.fa_decode_dc(
18848 &q,
18849 &k_view,
18850 &v_view,
18851 &mut attn,
18852 hd,
18853 nh,
18854 nkv,
18855 &kvl.len_d,
18856 bucket,
18857 scale,
18858 kvl.k_tok_bytes,
18859 kvl.v_tok_bytes,
18860 g,
18861 )?;
18862 return Ok(e.matmul(&fa.wo, &attn, t)?);
18863 }
18864 for i in 0..t {
18865 let avail = base_len + i + 1;
18866 let (off_tok, t_kv) = if swa && avail > win {
18867 (avail - win, win)
18868 } else {
18869 (0, avail)
18870 };
18871 let k_view = e.view_u8_range(
18872 &kvl.k,
18873 off_tok * kvl.k_tok_bytes,
18874 (off_tok + t_kv) * kvl.k_tok_bytes,
18875 );
18876 let v_view = e.view_u8_range(
18877 &kvl.v,
18878 off_tok * kvl.v_tok_bytes,
18879 (off_tok + t_kv) * kvl.v_tok_bytes,
18880 );
18881 let qv = e.view(&q, t * nh * hd);
18882 let q_row = qv.slice(i * nh * hd..(i + 1) * nh * hd);
18883 let mut q_one = e.uninit(nh * hd)?;
18884 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
18885 let mut a_one = e.uninit(nh * hd)?;
18886 e.fa_decode_kvmod(
18890 &q_one,
18891 &k_view,
18892 &v_view,
18893 &mut a_one,
18894 hd,
18895 nh,
18896 nkv,
18897 t_kv,
18898 scale,
18899 kvl.k_tok_bytes,
18900 kvl.v_tok_bytes,
18901 (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on()),
18902 )?;
18903 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
18904 }
18905 Ok(e.matmul(&fa.wo, &attn, t)?)
18906 }
18907
18908 fn gemma4_e4b_trunk(
18913 &self,
18914 e: &Engine,
18915 tokens: &[u32],
18916 pos0: usize,
18917 cache: &mut Cache,
18918 head_last: bool,
18919 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18920 let n_embd = self.cfg.n_embd as usize;
18921 let t = tokens.len();
18922 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
18923 let pos_d = e.htod_i32(&pos)?;
18924 let mut x = e.htod(&self.embd.gather(n_embd, tokens))?;
18925 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
18926 let inp_pl = self.gemma4_e4b_inp_pl(e, tokens, &x, t)?;
18927 self.gemma4_e4b_trunk_core(e, x, inp_pl, &pos_d, t, cache, None, true, head_last)
18928 }
18929
18930 fn gemma4_e4b_trunk_core(
18934 &self,
18935 e: &Engine,
18936 x_in: CudaSlice<f32>,
18937 inp_pl: CudaSlice<f32>,
18938 pos_d: &CudaSlice<i32>,
18939 t: usize,
18940 cache: &mut Cache,
18941 dc_bucket: Option<usize>,
18942 cap_logits: bool,
18943 head_last: bool,
18944 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18945 let n_embd = self.cfg.n_embd as usize;
18946 let eps = self.cfg.rms_eps;
18947 let n_layer = self.layers.len();
18948 let mut x = x_in;
18949 let aux_e4b = self.gemma4_aux.as_ref().unwrap().e4b.as_ref().unwrap();
18950 let n_epl = aux_e4b.n_epl;
18951
18952 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
18958 for il in 0..n_layer {
18959 let layer = &self.layers[il];
18960 let (hq, hdq) = match h_carry.take() {
18961 Some(p) => p,
18962 None => e.rms_norm_q8_1(&x, layer.attn_norm.float_data(), n_embd, t, eps)?,
18963 };
18964 let o = self.gemma4_e4b_attn(e, il, &hq, &hdq, pos_d, t, cache, dc_bucket)?;
18965 let bits = layer.gemma4.as_ref().unwrap();
18968 let e4b = bits.e4b.as_ref().expect("e4b layer bits");
18969 let fuse_exit = e.uses_q8_1_fast(&e4b.inp_gate);
18980 let (sn, attn_out) = self.gemma4_layer_tail_core_pn(
18981 e,
18982 layer,
18983 &o,
18984 &x,
18985 t,
18986 Some(layer.post_attn_norm.float_data()),
18987 fuse_exit,
18988 )?;
18989 let mut resid = e.uninit(t * n_embd)?;
18990 let g = if fuse_exit {
18996 let (rq, rd) = e.rms_pre_add_q8_1(
18998 &sn,
18999 bits.post_ffw_norm.float_data(),
19000 &attn_out,
19001 &mut resid,
19002 n_embd,
19003 t,
19004 self.cfg.rms_eps,
19005 )?;
19006 e.matmul_pre(&e4b.inp_gate, &rq, &rd, &resid, t)?
19007 } else {
19008 e.add(&sn, &attn_out, &mut resid, t * n_embd)?;
19009 e.matmul(&e4b.inp_gate, &resid, t)?
19010 };
19011 let mut act = e.uninit(t * n_epl)?;
19012 let y = if t == 1 && e.uses_q8_1_fast(&e4b.proj) {
19013 let ipv = e.view(&inp_pl, n_epl * n_layer);
19014 let row = ipv.slice(il * n_epl..(il + 1) * n_epl);
19015 let (aq, ad) = e.gelu_tanh_mul_q8_1(&g, &row, &mut act, n_epl, 1)?;
19016 e.matmul_pre(&e4b.proj, &aq, &ad, &act, t)?
19017 } else {
19018 let mut inp_this = e.uninit(t * n_epl)?;
19019 e.copy_rows_strided(
19020 &inp_pl,
19021 &mut inp_this,
19022 n_epl,
19023 t,
19024 n_epl * n_layer,
19025 il * n_epl,
19026 )?;
19027 e.gelu_tanh_mul(&g, &inp_this, &mut act, t * n_epl)?;
19028 e.matmul(&e4b.proj, &act, t)?
19029 };
19030 let next_norm = if il + 1 < n_layer {
19033 self.layers[il + 1].attn_norm.float_data()
19034 } else {
19035 self.output_norm.float_data()
19036 };
19037 let mut xn = e.uninit(t * n_embd)?;
19038 let pair = e.rms_pre_add_scale_rms_norm_q8_1(
19039 &y,
19040 e4b.post_norm.float_data(),
19041 &resid,
19042 bits.layer_scale,
19043 next_norm,
19044 &mut xn,
19045 n_embd,
19046 t,
19047 eps,
19048 )?;
19049 h_carry = Some(pair);
19050 x = xn;
19051 }
19052 let (oq, odq) = h_carry.take().unwrap();
19056 let h0 = e.zeros(0)?;
19057 let hm = if head_last { 1 } else { t };
19058 let (hq, hd) = if head_last && t > 1 {
19059 let mut q1 = e.uninit_i8(n_embd)?;
19060 e.dtod_copy_view_i8(&oq.slice((t - 1) * n_embd..t * n_embd), &mut q1)?;
19061 let nb = n_embd / 32;
19062 let mut d1 = e.uninit(nb)?;
19063 e.dtod_copy_view(&odq.slice((t - 1) * nb..t * nb), &mut d1)?;
19064 (q1, d1)
19065 } else {
19066 (oq, odq)
19067 };
19068 let mut ld = e.matmul_pre(&self.output, &hq, &hd, &h0, hm)?;
19069 if cap_logits {
19073 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
19074 e.softcap(&mut ld, cap, hm * self.output.out_features())?;
19075 }
19076 self.gemma4_suppress(e, &mut ld, hm)?; Ok((ld, x))
19078 }
19079
19080 pub fn gemma4_e4b_decode_step_t_am_dev(
19087 &self,
19088 e: &Engine,
19089 tok_d: &CudaSlice<u32>,
19090 t: usize,
19091 pos0: usize,
19092 cache: &mut Cache,
19093 ) -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
19094 let n_embd = self.cfg.n_embd as usize;
19095 let eps = self.cfg.rms_eps;
19096 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
19097 let pos_d = e.htod_i32(&pos)?;
19098 let embd_gpu = self
19099 .embd_gpu
19100 .get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload"));
19101 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
19102 let mut x = e.embed_gather_device_td(embd_gpu, tok_d, t, n_embd, qt, rb)?;
19103 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
19104 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, tok_d, &x, t)?;
19105 let (ld, xp) =
19106 self.gemma4_e4b_trunk_core(e, x, inp_pl, &pos_d, t, cache, None, true, false)?;
19107 let n_vocab = self.output.out_features();
19110 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
19111 for i in 0..t {
19112 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
19113 }
19114 let mut hn = e.uninit(t * n_embd)?;
19115 e.rms_norm(&xp, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
19116 cache.pos += t;
19117 Ok((vam, hn))
19118 }
19119
19120 pub(crate) fn gemma4_e4b_decode_step_t_h(
19123 &self,
19124 e: &Engine,
19125 tokens: &[u32],
19126 pos0: usize,
19127 cache: &mut Cache,
19128 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
19129 let n_embd = self.cfg.n_embd as usize;
19130 let eps = self.cfg.rms_eps;
19131 let t = tokens.len();
19132 let (ld, xp) = self.gemma4_e4b_trunk(e, tokens, pos0, cache, false)?;
19133 let mut hn = e.uninit(t * n_embd)?;
19134 e.rms_norm(&xp, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
19135 cache.pos += t;
19136 Ok((e.dtoh(&ld)?, hn))
19137 }
19138
19139 pub fn gemma4_e4b_decode_step_dcg(
19145 &self,
19146 e: &Engine,
19147 token_d: &mut CudaSlice<u32>,
19148 pos_d: &mut CudaSlice<i32>,
19149 embd_gpu: &CudaSlice<u8>,
19150 embd_qt: i32,
19151 embd_rb: usize,
19152 cache: &mut Cache,
19153 n_vocab: usize,
19154 bucket: usize,
19155 ) -> Result<(), Box<dyn std::error::Error>> {
19156 let n_embd = self.cfg.n_embd as usize;
19157 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
19158 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
19159 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, token_d, &x, 1)?;
19160 let (ld, _x) =
19161 self.gemma4_e4b_trunk_core(e, x, inp_pl, pos_d, 1, cache, Some(bucket), false, false)?;
19162 e.argmax_token_device_into(&ld, token_d, n_vocab)?;
19163 e.inc_seqlen(pos_d)?;
19164 Ok(())
19165 }
19166
19167 #[allow(clippy::too_many_arguments)]
19175 pub fn gemma4_e4b_decode_step_dc(
19176 &self,
19177 e: &Engine,
19178 token_d: &CudaSlice<u32>,
19179 pos_d: &mut CudaSlice<i32>,
19180 embd_gpu: &CudaSlice<u8>,
19181 embd_qt: i32,
19182 embd_rb: usize,
19183 cache: &mut Cache,
19184 n_vocab: usize,
19185 ) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
19186 let n_embd = self.cfg.n_embd as usize;
19187 let eps = self.cfg.rms_eps;
19188 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
19189 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
19190 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, token_d, &x, 1)?;
19191 let (ld, _x) =
19192 self.gemma4_e4b_trunk_core(e, x, inp_pl, pos_d, 1, cache, None, false, false)?;
19193 let mut tok_out = e.stream().alloc_zeros::<u32>(1)?;
19194 e.argmax_token_device_into(&ld, &mut tok_out, n_vocab)?;
19195 e.inc_seqlen(pos_d)?;
19196 cache.pos += 1;
19197 let _ = eps;
19198 Ok(tok_out)
19199 }
19200
19201 pub(crate) fn gemma4_e4b_decode_step_h(
19204 &self,
19205 e: &Engine,
19206 token: u32,
19207 cache: &mut Cache,
19208 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
19209 let (ld, x) = self.gemma4_e4b_trunk(e, &[token], cache.pos, cache, false)?;
19210 let logits = e.dtoh(&ld)?;
19211 cache.pos += 1;
19212 Ok((logits, x))
19213 }
19214
19215 pub(crate) fn gemma4_e4b_prime(
19219 &self,
19220 e: &Engine,
19221 tokens: &[u32],
19222 cache: &mut Cache,
19223 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
19224 if cache.pos != 0 {
19227 return Err(
19228 "e4b prime is fresh-prompt only (v0) — prime the full prompt in one \
19229 call or decode tokenwise"
19230 .into(),
19231 );
19232 }
19233 let n_embd = self.cfg.n_embd as usize;
19234 let t = tokens.len();
19235 let (ld, x) = self.gemma4_e4b_trunk(e, tokens, 0, cache, true)?;
19236 cache.pos += t;
19237 let last = e.dtoh(&ld)?; let xv = e.view(&x, t * n_embd);
19239 let row = xv.slice((t - 1) * n_embd..t * n_embd);
19240 let mut h_seed = e.uninit(n_embd)?;
19241 e.copy_view_into(&mut h_seed, 0, &row, n_embd)?;
19242 Ok((last, h_seed, x))
19243 }
19244
19245 pub(crate) fn gemma4_e4b_forward(
19247 &self,
19248 e: &Engine,
19249 tokens: &[u32],
19250 last_only: bool,
19251 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
19252 let mut cache = Cache::new(e, &self.cfg, tokens.len() + 8)?;
19253 let (ld, _x) = self.gemma4_e4b_trunk(e, tokens, 0, &mut cache, last_only)?;
19254 Ok(e.dtoh(&ld)?) }
19256}
19257
19258#[cfg(test)]
19259mod prime_chunk_schedule_tests {
19260 use super::{
19261 PRIME_MIN_T, PRIME_PIPE_MIN_CHUNK, active_matrix_values, align_prime_ranges_to_gdn,
19262 dynamic_prime_chunk_ranges, fixed_prime_chunk_ranges, fixed_prime_chunk_ranges_for_ring,
19263 parse_step_ep_grouped_prefill, parse_step_tp_prefill, step_grouped_decode_shape,
19264 step_grouped_prefill_shape, step_tp_prefill_shape, validate_step_prime_batch_modes,
19265 };
19266
19267 fn sizes(ranges: &[(usize, usize)]) -> Vec<usize> {
19268 ranges.iter().map(|(start, end)| end - start).collect()
19269 }
19270
19271 fn auto_chunk(t: usize) -> usize {
19272 t.div_ceil(8).max(PRIME_PIPE_MIN_CHUNK).min(4096)
19273 }
19274
19275 #[test]
19280 fn auto_prime_ranges_align_to_the_gdn_grid() {
19281 let c = 32usize; let assert_covers = |ranges: &[(usize, usize)], t: usize| {
19283 assert_eq!(ranges.first().map(|&(s, _)| s), Some(0));
19284 assert_eq!(ranges.last().map(|&(_, e)| e), Some(t));
19285 for w in ranges.windows(2) {
19286 assert_eq!(w[0].1, w[1].0, "ranges must stay contiguous");
19287 }
19288 assert!(ranges.iter().all(|&(s, e)| e > s), "no empty range");
19289 };
19290
19291 let t = 9510usize;
19294 let fill = auto_chunk(t);
19295 let fixed = fixed_prime_chunk_ranges(t, fill);
19296 assert!(
19297 fixed[..fixed.len() - 1].iter().any(|&(_, e)| e % c != 0),
19298 "broken arm vanished: fixed auto boundaries all landed on-grid"
19299 );
19300 let dynamic = dynamic_prime_chunk_ranges(t, fill, &fixed);
19301 assert!(
19302 dynamic[..dynamic.len() - 1]
19303 .iter()
19304 .any(|&(_, e)| e % c != 0),
19305 "broken arm vanished: dynamic auto boundaries all landed on-grid"
19306 );
19307
19308 for ranges in [&fixed, &dynamic] {
19309 let aligned = align_prime_ranges_to_gdn(ranges, t, c);
19310 assert_covers(&aligned, t);
19311 for &(_, e) in &aligned[..aligned.len() - 1] {
19312 assert_eq!(e % c, 0, "internal boundary {e} off the {c}-grid");
19313 }
19314 for (&(_, a), &(_, b)) in aligned.iter().zip(ranges.iter()) {
19316 assert!(a <= b && b - a < c);
19317 }
19318 }
19319
19320 let tight = vec![(0usize, 33usize), (33, 40), (40, 200)];
19323 let aligned = align_prime_ranges_to_gdn(&tight, 200, c);
19324 assert_covers(&aligned, 200);
19325 assert_eq!(aligned, vec![(0, 32), (32, 200)]);
19326
19327 assert_eq!(align_prime_ranges_to_gdn(&[(0, 200)], 200, c), [(0, 200)]);
19329 assert_eq!(align_prime_ranges_to_gdn(&tight, 200, 0), tight.as_slice());
19330 let on_grid = vec![(0usize, 128usize), (128, 256), (256, 300)];
19331 assert_eq!(
19332 align_prime_ranges_to_gdn(&on_grid, 300, c),
19333 on_grid.as_slice()
19334 );
19335 }
19336
19337 #[test]
19338 fn active_matrix_prefix_scopes_reused_prime_slabs() {
19339 assert_eq!(
19340 active_matrix_values(40 * 4096, 29, 4096, "activation").unwrap(),
19341 29 * 4096
19342 );
19343 assert_eq!(
19344 active_matrix_values(29 * 4096, 29, 4096, "activation").unwrap(),
19345 29 * 4096
19346 );
19347 assert_eq!(
19348 active_matrix_values(29 * 4096, 24, 4096, "activation").unwrap(),
19349 24 * 4096
19350 );
19351 assert!(active_matrix_values(28 * 4096, 29, 4096, "activation").is_err());
19352 assert!(active_matrix_values(usize::MAX, usize::MAX, 2, "activation").is_err());
19353 }
19354
19355 #[test]
19356 fn step_tp_prefill_batch_refuses_before_scheduler_fallback() {
19357 assert!(validate_step_prime_batch_modes(false, false).is_ok());
19358
19359 let grouped_without_tp = validate_step_prime_batch_modes(false, true).unwrap_err();
19360 assert!(grouped_without_tp.contains("requires MEMRA_STEP_TP_PREFILL=1"));
19361
19362 for grouped in [false, true] {
19363 let err = validate_step_prime_batch_modes(true, grouped).unwrap_err();
19364 assert!(err.contains("did not clear the live-server performance gate"));
19365 assert!(err.contains("per-session grouped prefill"));
19366 }
19367 }
19368
19369 #[test]
19370 fn step_grouped_path_is_eager_single_token_only() {
19371 assert!(step_grouped_decode_shape(false, 1));
19372 assert!(!step_grouped_decode_shape(true, 1));
19373 assert!(!step_grouped_decode_shape(false, 2));
19374 assert!(!step_grouped_decode_shape(true, 2));
19375 }
19376
19377 #[test]
19378 fn step_grouped_prefill_door_is_strict_and_capacity_bounded() {
19379 assert!(!parse_step_ep_grouped_prefill(None).unwrap());
19380 assert!(!parse_step_ep_grouped_prefill(Some("")).unwrap());
19381 assert!(!parse_step_ep_grouped_prefill(Some("0")).unwrap());
19382 assert!(parse_step_ep_grouped_prefill(Some("1")).unwrap());
19383 assert!(parse_step_ep_grouped_prefill(Some("true")).is_err());
19384 assert!(parse_step_ep_grouped_prefill(Some("2")).is_err());
19385
19386 assert!(step_grouped_prefill_shape(true, true, PRIME_MIN_T));
19387 assert!(step_grouped_prefill_shape(
19388 true,
19389 true,
19390 crate::cache::PRIME_CHUNK_MAX_TOKENS,
19391 ));
19392 assert!(!step_grouped_prefill_shape(true, true, PRIME_MIN_T - 1,));
19393 assert!(!step_grouped_prefill_shape(
19394 true,
19395 true,
19396 crate::cache::PRIME_CHUNK_MAX_TOKENS + 1,
19397 ));
19398 assert!(!step_grouped_prefill_shape(false, true, PRIME_MIN_T));
19399 assert!(!step_grouped_prefill_shape(true, false, PRIME_MIN_T));
19400 }
19401
19402 #[test]
19403 fn step_tp_prefill_door_is_strict_and_default_off() {
19404 assert!(!parse_step_tp_prefill(None).unwrap());
19405 assert!(!parse_step_tp_prefill(Some("")).unwrap());
19406 assert!(!parse_step_tp_prefill(Some("0")).unwrap());
19407 assert!(parse_step_tp_prefill(Some("1")).unwrap());
19408 assert!(parse_step_tp_prefill(Some("true")).is_err());
19409 assert!(parse_step_tp_prefill(Some("2")).is_err());
19410 }
19411
19412 #[test]
19413 fn step_tp_prefill_requires_a_qualified_even_rank_shape() {
19414 assert!(step_tp_prefill_shape(
19415 true,
19416 PRIME_MIN_T,
19417 4,
19418 true,
19419 true,
19420 false,
19421 ));
19422 assert!(!step_tp_prefill_shape(
19423 false,
19424 PRIME_MIN_T,
19425 4,
19426 true,
19427 true,
19428 false,
19429 ));
19430 assert!(!step_tp_prefill_shape(
19431 true,
19432 PRIME_MIN_T - 1,
19433 4,
19434 true,
19435 true,
19436 false,
19437 ));
19438 assert!(step_tp_prefill_shape(
19440 true,
19441 PRIME_MIN_T,
19442 2,
19443 true,
19444 true,
19445 false
19446 ));
19447 assert!(!step_tp_prefill_shape(
19448 true,
19449 PRIME_MIN_T,
19450 1,
19451 true,
19452 true,
19453 false
19454 ));
19455 assert!(!step_tp_prefill_shape(
19456 true,
19457 PRIME_MIN_T,
19458 3,
19459 true,
19460 true,
19461 false
19462 ));
19463 assert!(!step_tp_prefill_shape(
19464 true,
19465 PRIME_MIN_T,
19466 4,
19467 false,
19468 true,
19469 false,
19470 ));
19471 assert!(!step_tp_prefill_shape(
19472 true,
19473 PRIME_MIN_T,
19474 4,
19475 true,
19476 false,
19477 false,
19478 ));
19479 assert!(!step_tp_prefill_shape(
19480 true,
19481 PRIME_MIN_T,
19482 4,
19483 true,
19484 true,
19485 true,
19486 ));
19487 }
19488
19489 #[test]
19490 fn fixed_schedule_retains_measured_geometry() {
19491 assert_eq!(
19492 sizes(&fixed_prime_chunk_ranges(461, 128)),
19493 vec![128, 128, 128, 77]
19494 );
19495 assert_eq!(
19496 sizes(&fixed_prime_chunk_ranges(1833, 230)),
19497 vec![230, 230, 230, 230, 230, 230, 230, 223]
19498 );
19499 assert_eq!(sizes(&fixed_prime_chunk_ranges(4096, 512)), vec![512; 8]);
19500 let capped = sizes(&fixed_prime_chunk_ranges_for_ring(8200, 4096, true));
19501 assert_eq!(capped, vec![4096, 4088, 16]);
19502 assert!(capped.iter().all(|&rows| rows <= 4096));
19503 assert_eq!(
19504 sizes(&fixed_prime_chunk_ranges_for_ring(4100, 4096, false)),
19505 vec![4100],
19506 "flag-off schedule remains byte-for-byte the legacy monolithic tail",
19507 );
19508 }
19509
19510 #[test]
19511 fn dynamic_schedule_matches_registered_shapes() {
19512 let cases = [
19513 (461, vec![64, 141, 132, 124]),
19514 (1833, vec![115, 269, 260, 252, 244, 237, 231, 225]),
19515 (4096, vec![256, 602, 580, 563, 545, 531, 516, 503]),
19516 ];
19517 for (t, expected) in cases {
19518 let chunk = auto_chunk(t);
19519 let fixed = fixed_prime_chunk_ranges(t, chunk);
19520 assert_eq!(
19521 sizes(&dynamic_prime_chunk_ranges(t, chunk, &fixed)),
19522 expected
19523 );
19524 }
19525 }
19526
19527 #[test]
19528 fn dynamic_schedule_covers_exactly_and_shrinks_after_fill() {
19529 for t in 256..=8192 {
19530 let chunk = auto_chunk(t);
19531 let fixed = fixed_prime_chunk_ranges(t, chunk);
19532 let dynamic = dynamic_prime_chunk_ranges(t, chunk, &fixed);
19533 assert_eq!(dynamic.len(), fixed.len(), "T={t}");
19534 assert_eq!(dynamic.first().unwrap().0, 0, "T={t}");
19535 assert_eq!(dynamic.last().unwrap().1, t, "T={t}");
19536 for pair in dynamic.windows(2) {
19537 assert_eq!(pair[0].1, pair[1].0, "T={t}");
19538 }
19539 assert!(
19540 dynamic
19541 .iter()
19542 .all(|(start, end)| end - start >= PRIME_MIN_T),
19543 "T={t} sizes={:?}",
19544 sizes(&dynamic)
19545 );
19546 if dynamic.len() >= 3 {
19547 let chunk_sizes = sizes(&dynamic);
19548 assert!(
19549 chunk_sizes[0] < chunk_sizes[1],
19550 "T={t} sizes={chunk_sizes:?}"
19551 );
19552 assert!(
19553 chunk_sizes[1..].windows(2).all(|pair| pair[0] >= pair[1]),
19554 "T={t} sizes={chunk_sizes:?}"
19555 );
19556 }
19557 }
19558 }
19559}
19560
19561#[cfg(test)]
19562mod page_prefetch_tests {
19563 use super::{
19564 grouped_worker_prefetch_position, page_prefetch_positions,
19565 page_prefetch_window_from_values, worker_prefetch_positions,
19566 };
19567
19568 #[test]
19569 fn page_prefetch_window_keeps_existing_opt_in_default() {
19570 assert_eq!(page_prefetch_window_from_values(false, None), 0);
19571 assert_eq!(page_prefetch_window_from_values(false, Some("8")), 0);
19572 assert_eq!(page_prefetch_window_from_values(true, None), 1);
19573 assert_eq!(page_prefetch_window_from_values(true, Some("bad")), 1);
19574 assert_eq!(page_prefetch_window_from_values(true, Some("0")), 0);
19575 assert_eq!(page_prefetch_window_from_values(true, Some("8")), 8);
19576 }
19577
19578 #[test]
19579 fn rolling_page_prefetch_advises_each_future_expert_once() {
19580 let advised: Vec<_> = (0..7)
19581 .flat_map(|position| page_prefetch_positions(position, 7, 3))
19582 .collect();
19583 assert_eq!(advised, vec![1, 2, 3, 4, 5, 6]);
19584
19585 let one_ahead: Vec<_> = (0..4)
19586 .flat_map(|position| page_prefetch_positions(position, 4, 1))
19587 .collect();
19588 assert_eq!(one_ahead, vec![1, 2, 3]);
19589 assert!(page_prefetch_positions(0, 4, 0).is_empty());
19590 }
19591
19592 #[test]
19593 fn grouped_worker_prefetch_primes_first_then_each_known_next_once() {
19594 assert_eq!(grouped_worker_prefetch_position(0, None), None);
19595 let positions: Vec<_> = std::iter::once(grouped_worker_prefetch_position(4, None).unwrap())
19596 .chain(
19597 (0..4).filter_map(|position| grouped_worker_prefetch_position(4, Some(position))),
19598 )
19599 .collect();
19600 assert_eq!(positions, vec![0, 1, 2, 3]);
19601 assert_eq!(grouped_worker_prefetch_position(1, Some(0)), None);
19602 }
19603
19604 #[test]
19605 fn rolling_worker_prefetch_primes_current_and_each_future_expert_once() {
19606 let queued: Vec<_> = (0..8)
19607 .flat_map(|position| worker_prefetch_positions(position, 8, 5))
19608 .collect();
19609 assert_eq!(queued, (0..8).collect::<Vec<_>>());
19610
19611 let one_at_a_time: Vec<_> = (0..4)
19612 .flat_map(|position| worker_prefetch_positions(position, 4, 1))
19613 .collect();
19614 assert_eq!(one_at_a_time, vec![0, 1, 2, 3]);
19615 assert!(worker_prefetch_positions(0, 4, 0).is_empty());
19616 }
19617}
19618
19619pub struct G4DcSlots {
19620 x: CudaSlice<f32>,
19621 xn: CudaSlice<f32>,
19622 cur: CudaSlice<f32>,
19623 hq: CudaSlice<i8>,
19624 hd_: CudaSlice<f32>,
19625 q0: CudaSlice<f32>,
19626 k0: CudaSlice<f32>,
19627 v0: CudaSlice<f32>,
19628 q: CudaSlice<f32>,
19629 k: CudaSlice<f32>,
19630 v: CudaSlice<f32>,
19631 attn: CudaSlice<f32>,
19632 o: CudaSlice<f32>,
19633 attn_out: CudaSlice<f32>,
19634 zsh: CudaSlice<f32>,
19635 zq: CudaSlice<i8>,
19636 zd: CudaSlice<f32>,
19637 gate: CudaSlice<f32>,
19638 up: CudaSlice<f32>,
19639 act: CudaSlice<f32>,
19640 actq: CudaSlice<i8>,
19641 actd: CudaSlice<f32>,
19642 f0: CudaSlice<f32>,
19643 sn: CudaSlice<f32>,
19644 hn: CudaSlice<f32>,
19645 logits: CudaSlice<f32>,
19646}
19647
19648pub struct Step35TokenGraphState {
19653 pub graphs: Vec<(usize, crate::tp::TokenGraph)>,
19655 pub token_d: cudarc::driver::CudaSlice<u32>,
19656 pub pos_d: cudarc::driver::CudaSlice<i32>,
19657 pub logits_stage: cudarc::driver::CudaSlice<f32>,
19658 pub x: cudarc::driver::CudaSlice<f32>,
19663 pub x1: cudarc::driver::CudaSlice<f32>,
19664 pub mixed_stage: cudarc::driver::CudaSlice<f32>,
19665 pub sh_stage: cudarc::driver::CudaSlice<f32>,
19666 pub k_shadow_stage: cudarc::driver::CudaSlice<f32>,
19667 pub v_shadow_stage: cudarc::driver::CudaSlice<f32>,
19668 pub router_logits: cudarc::driver::CudaSlice<f32>,
19671 pub shexp_gate: cudarc::driver::CudaSlice<f32>,
19672 pub shexp_up: cudarc::driver::CudaSlice<f32>,
19673 pub shexp_act: cudarc::driver::CudaSlice<f32>,
19674 pub gate_sig: cudarc::driver::CudaSlice<f32>,
19675 pub dense_z: cudarc::driver::CudaSlice<f32>,
19676 pub dense_gate: cudarc::driver::CudaSlice<f32>,
19677 pub dense_up: cudarc::driver::CudaSlice<f32>,
19678 pub dense_act: cudarc::driver::CudaSlice<f32>,
19679 pub hn: cudarc::driver::CudaSlice<f32>,
19680 pub probe_mixed: cudarc::driver::CudaSlice<f32>,
19683 pub probe_x: cudarc::driver::CudaSlice<f32>,
19684 pub token_hist: cudarc::driver::CudaSlice<u32>,
19687 pub hist_idx: cudarc::driver::CudaSlice<i32>,
19688}
19689
19690impl HybridModel {
19691 pub(crate) fn step35_token_graph_step(
19702 &self,
19703 e: &Engine,
19704 token: u32,
19705 cache: &mut Cache,
19706 ) -> Result<Option<(Vec<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
19707 if !self.uses_sliding_gated_moe_program()
19708 || !crate::tp::step_tp_graph_enabled()?
19709 || !crate::tp::step_tp_dcw_enabled()?
19710 || !crate::tp::step_tp_qkv_fused_enabled()?
19711 || !crate::tp::step_tp_dev_router_enabled()?
19712 || !crate::tp::step_nvfp4_dev_routes_enabled()?
19713 {
19714 return Ok(None);
19715 }
19716 if !crate::spec::graph_launch_headroom_ok(e) {
19722 static NOTED: std::sync::Once = std::sync::Once::new();
19723 NOTED.call_once(|| crate::spec::graph_replay_suspended_note("step-tp-token"));
19724 return Ok(None);
19725 }
19726 let n_embd = self.cfg.n_embd as usize;
19727 let n_vocab = self.cfg.n_vocab as usize;
19728 let eps = self.cfg.rms_eps;
19729 let n_layers = self.layers.len();
19730 let pos = cache.pos;
19731 let staged_next = pos + 1;
19732 if staged_next < 96 {
19733 return Ok(None); }
19735
19736 for il in 0..n_layers {
19739 let Some(tp_kv) = cache.tp_kv[il].as_ref() else {
19740 return Ok(None); };
19742 if tp_kv.peek_append_ring(1)?.1 {
19743 return Ok(None);
19744 }
19745 }
19746
19747 let (fa_vec, n_splits) = e.fa_geom_eager(staged_next, 128, 8, false);
19750 if !fa_vec {
19751 return Ok(None);
19752 }
19753 let sp = crate::fa_split_keys(staged_next, 8);
19754 let bucket_max = (n_splits * sp).max(staged_next);
19755
19756 let mut state_guard = self
19757 .step35_token_graph
19758 .lock()
19759 .map_err(|_| "step35 token graph lock is poisoned")?;
19760 if state_guard.is_none() {
19761 let _main = e.gpu.enter_main()?;
19762 let n_expert = self
19763 .cfg
19764 .moe
19765 .as_ref()
19766 .map(|m| m.expert_count as usize)
19767 .unwrap_or(0);
19768 let n_ff_sh = self
19769 .layers
19770 .iter()
19771 .find_map(|l| match &l.ffn {
19772 crate::hybrid::Ffn::Moe(m) => m.gate_shexp.as_ref().map(|g| g.out_features()),
19773 _ => None,
19774 })
19775 .unwrap_or(0);
19776 let n_ff_dense = self
19777 .layers
19778 .iter()
19779 .find_map(|l| match &l.ffn {
19780 crate::hybrid::Ffn::Dense { ffn_gate, .. } => Some(ffn_gate.out_features()),
19781 _ => None,
19782 })
19783 .unwrap_or(0);
19784 *state_guard = Some(Step35TokenGraphState {
19785 graphs: Vec::new(),
19786 token_d: e.stream().clone_htod(&[0u32])?,
19787 pos_d: e.htod_i32(&[pos as i32])?,
19788 logits_stage: e.htod(&vec![0.0f32; n_vocab])?,
19789 x: e.htod(&vec![0.0f32; n_embd])?,
19790 x1: e.htod(&vec![0.0f32; n_embd])?,
19791 mixed_stage: e.htod(&vec![0.0f32; n_embd])?,
19792 sh_stage: e.htod(&vec![0.0f32; n_embd])?,
19793 k_shadow_stage: e.htod(&vec![0.0f32; 2048])?,
19794 v_shadow_stage: e.htod(&vec![0.0f32; 2048])?,
19795 router_logits: e.htod(&vec![0.0f32; n_expert.max(1)])?,
19796 shexp_gate: e.htod(&vec![0.0f32; n_ff_sh.max(1)])?,
19797 shexp_up: e.htod(&vec![0.0f32; n_ff_sh.max(1)])?,
19798 shexp_act: e.htod(&vec![0.0f32; n_ff_sh.max(1)])?,
19799 gate_sig: e.htod(&vec![1.0f32; 1])?,
19800 dense_z: e.htod(&vec![0.0f32; n_embd])?,
19801 dense_gate: e.htod(&vec![0.0f32; n_ff_dense.max(1)])?,
19802 dense_up: e.htod(&vec![0.0f32; n_ff_dense.max(1)])?,
19803 dense_act: e.htod(&vec![0.0f32; n_ff_dense.max(1)])?,
19804 hn: e.htod(&vec![0.0f32; n_embd])?,
19805 probe_mixed: e.htod(&vec![0.0f32; n_embd])?,
19806 probe_x: e.htod(&vec![0.0f32; n_embd])?,
19807 token_hist: e.stream().clone_htod(&[0u32; 16])?,
19808 hist_idx: e.htod_i32(&[0])?,
19809 });
19810 }
19811 let state = state_guard.as_mut().expect("state armed above");
19812 {
19816 let _main = e.gpu.enter_main()?;
19817 let Step35TokenGraphState {
19818 logits_stage,
19819 token_d,
19820 ..
19821 } = &mut *state;
19822 e.argmax_token_device_into(logits_stage, token_d, n_vocab)?;
19823 }
19824
19825 if state.graphs.is_empty() {
19830 self.step35_token_graph_build(e, cache, state, bucket_max)?;
19833 }
19834 {
19835 let (b, g) = state.graphs.first_mut().expect("graph built above");
19836 if *b != bucket_max {
19837 g.retarget_bucket(bucket_max)?;
19838 *b = bucket_max;
19839 }
19840 }
19841 let graph = state
19842 .graphs
19843 .first()
19844 .map(|(_, g)| g)
19845 .expect("graph built above");
19846
19847 let tg_timing = std::env::var("MEMRA_STEP_TP_TIMING").as_deref() == Ok("1");
19848 let t_fence = tg_timing.then(std::time::Instant::now);
19849 {
19854 let fa0 = match &self.layers[0].mixer {
19855 Mixer::Full(fa) => fa,
19856 _ => return Err("step35 token graph expects full-attention layers".into()),
19857 };
19858 let tp0 = fa0
19859 .step_tp_qkv
19860 .as_ref()
19861 .ok_or("step35 token graph lost its TP state")?;
19862 for rank in 0..tp0.runtime.devices().len() {
19863 let engine = tp0
19864 .runtime
19865 .rank_engine(rank)
19866 .ok_or("step35 token graph lost a rank engine")?;
19867 let _main = engine.gpu.enter_main()?;
19868 engine.stream().synchronize()?;
19869 }
19870 }
19871
19872 {
19874 let _main = e.gpu.enter_main()?;
19875 e.set_u32_one(&mut state.token_d, token)?;
19876 e.set_i32_one(&mut state.pos_d, pos as i32)?;
19877 }
19878 let t_launch = tg_timing.then(std::time::Instant::now);
19879 graph.launch(e)?;
19880 let t_book = tg_timing.then(std::time::Instant::now);
19881 for il in 0..n_layers {
19886 let tp_kv = cache.tp_kv[il].as_mut().expect("eligibility checked above");
19887 let transaction = tp_kv.begin_transaction()?;
19888 let fa = match &self.layers[il].mixer {
19889 Mixer::Full(fa) => fa,
19890 _ => return Err("step35 token graph expects full-attention layers".into()),
19891 };
19892 let tp = fa
19893 .step_tp_qkv
19894 .as_ref()
19895 .ok_or("step35 token graph lost its TP state")?;
19896 let empty: [CudaSlice<f32>; 0] = [];
19899 tp.runtime.append_tp_kv_transaction_inner(
19900 tp_kv,
19901 transaction,
19902 &empty,
19903 &empty,
19904 1,
19905 true,
19906 )?;
19907 tp.runtime
19908 .commit_tp_kv_transaction_external(tp_kv, transaction, 1)?;
19909 if let Some(local) = cache.kv[il].as_mut() {
19911 local.len = pos + 1;
19912 let _main = e.gpu.enter_main()?;
19913 e.set_i32_one(&mut local.len_d, (pos + 1) as i32)?;
19914 }
19915 }
19916 cache.pos = pos + 1;
19917 let t_sync = tg_timing.then(std::time::Instant::now);
19918 let (logits, h_seed) = {
19919 let _main = e.gpu.enter_main()?;
19920 e.stream().synchronize()?;
19921 (e.dtoh(&state.logits_stage)?, e.clone_dtod(&state.x)?)
19922 };
19923 if let (Some(f), Some(l), Some(b), Some(sy)) = (t_fence, t_launch, t_book, t_sync) {
19924 use std::sync::atomic::{AtomicU64, Ordering};
19925 static NS: [AtomicU64; 5] = [
19926 AtomicU64::new(0),
19927 AtomicU64::new(0),
19928 AtomicU64::new(0),
19929 AtomicU64::new(0),
19930 AtomicU64::new(0),
19931 ];
19932 static CALLS: AtomicU64 = AtomicU64::new(0);
19933 let now = std::time::Instant::now();
19934 NS[0].fetch_add((l - f).as_nanos() as u64, Ordering::Relaxed); NS[1].fetch_add((b - l).as_nanos() as u64, Ordering::Relaxed); NS[2].fetch_add((sy - b).as_nanos() as u64, Ordering::Relaxed); NS[3].fetch_add((now - sy).as_nanos() as u64, Ordering::Relaxed); NS[4].fetch_add((now - f).as_nanos() as u64, Ordering::Relaxed); let calls = CALLS.fetch_add(1, Ordering::Relaxed) + 1;
19940 if calls % 100 == 0 {
19941 let avg = |i: usize| NS[i].load(Ordering::Relaxed) as f64 / calls as f64 / 1e3;
19942 eprintln!(
19943 "[tg-timing] calls={calls} fence_us={:.0} launch_us={:.0} book_us={:.0} \
19944 syncdtoh_us={:.0} total_us={:.0}",
19945 avg(0),
19946 avg(1),
19947 avg(2),
19948 avg(3),
19949 avg(4)
19950 );
19951 }
19952 }
19953 if std::env::var("MEMRA_TG_PROBE_LAYER").is_ok() {
19955 use std::io::Write;
19956 let (pm, px) = {
19957 let _main = e.gpu.enter_main()?;
19958 (e.dtoh(&state.probe_mixed)?, e.dtoh(&state.probe_x)?)
19959 };
19960 for (path, data) in [
19961 ("/root/tg-probe-mixed.bin", &pm),
19962 ("/root/tg-probe-x.bin", &px),
19963 ] {
19964 let mut fo = std::fs::OpenOptions::new()
19965 .create(true)
19966 .append(true)
19967 .open(path)?;
19968 for v in data {
19969 fo.write_all(&v.to_le_bytes())?;
19970 }
19971 }
19972 }
19973 if let Ok(path) = std::env::var("MEMRA_DUMP_HN") {
19976 let hh = {
19977 let _main = e.gpu.enter_main()?;
19978 e.dtoh(&state.hn)?
19979 };
19980 use std::io::Write;
19981 let mut fo = std::fs::OpenOptions::new()
19982 .create(true)
19983 .append(true)
19984 .open(path)?;
19985 for v in &hh {
19986 fo.write_all(&v.to_le_bytes())?;
19987 }
19988 }
19989 if std::env::var("MEMRA_STEP_TP_GRAPH_DEBUG").as_deref() == Ok("1") {
19992 for il in [0usize, 1, 44] {
19993 let tp_kv = cache.tp_kv[il].as_ref().expect("eligibility checked above");
19994 let host_len = tp_kv.staged_len();
19995 let fa = match &self.layers[il].mixer {
19996 Mixer::Full(fa) => fa,
19997 _ => continue,
19998 };
19999 let tp = fa
20000 .step_tp_qkv
20001 .as_ref()
20002 .ok_or("step35 token graph lost its TP state")?;
20003 for rank in 0..tp.runtime.devices().len() {
20004 let engine = tp
20005 .runtime
20006 .rank_engine(rank)
20007 .ok_or("step35 token graph lost a rank engine")?;
20008 let rank_cache = tp_kv.rank(rank).ok_or("debug rank cache missing")?;
20009 let _main = engine.gpu.enter_main()?;
20010 engine.stream().synchronize()?;
20011 let len_d = engine.dtoh_i32_one(rank_cache.len_d())?;
20012 let base_d = match rank_cache.base_d() {
20013 Some(b) => engine.dtoh_i32_one(b)?,
20014 None => -1,
20015 };
20016 eprintln!(
20017 "[graph-debug] pos={pos} il={il} rank={rank} host_len={host_len} \
20018 len_d={len_d} base_d={base_d}"
20019 );
20020 }
20021 }
20022 }
20023 Ok(Some((logits, h_seed)))
20024 }
20025
20026 pub(crate) fn head_split_matvec(
20032 &self,
20033 e: &Engine,
20034 hn: &CudaSlice<f32>,
20035 ) -> Result<Option<Vec<f32>>, Box<dyn std::error::Error>> {
20036 if self.head_split_fill_device(e, hn)?.is_none() {
20037 return Ok(None);
20038 }
20039 let guard = HEAD_SPLIT_WS
20040 .lock()
20041 .map_err(|_| "head split lock is poisoned")?;
20042 let ws = guard.as_ref().expect("filled above");
20043 let _main = e.gpu.enter_main()?;
20044 Ok(Some(e.dtoh(&ws.logits_e)?))
20045 }
20046
20047 fn head_split_fill_device(
20051 &self,
20052 e: &Engine,
20053 hn: &CudaSlice<f32>,
20054 ) -> Result<Option<()>, Box<dyn std::error::Error>> {
20055 use cudarc::driver::DevicePtr;
20056 let crate::model::GpuTensor::FloatBf16 { data: head, .. } = &self.output else {
20057 return Ok(None);
20058 };
20059 let Some(rank1) = self.layers.first().and_then(|l| match &l.mixer {
20060 Mixer::Full(fa) => fa
20061 .step_tp_qkv
20062 .as_ref()
20063 .and_then(|tp| tp.runtime.rank_engine(1)),
20064 _ => None,
20065 }) else {
20066 return Ok(None);
20067 };
20068 let n_embd = self.cfg.n_embd as usize;
20069 let n_vocab = self.cfg.n_vocab as usize;
20070 let half = n_vocab / 2;
20071 let mut guard = HEAD_SPLIT_WS
20072 .lock()
20073 .map_err(|_| "head split lock is poisoned")?;
20074 let pin = {
20075 let _main = e.gpu.enter_main()?;
20076 let stream = e.stream();
20077 let (ptr, _g) = head.device_ptr(&stream);
20078 ptr as u64
20079 };
20080 if guard.as_ref().is_none_or(|ws| ws.pin != pin) {
20081 let hi_rows = n_vocab - half;
20083 let (w1, hn1, y1, ev_done) = {
20084 let _r1 = rank1.gpu.enter_main()?;
20085 (
20086 rank1.alloc_u8_uninit(hi_rows * n_embd * 2)?,
20087 rank1.htod(&vec![0.0f32; n_embd])?,
20088 rank1.htod(&vec![0.0f32; hi_rows])?,
20089 rank1.ctx().new_event(None)?,
20090 )
20091 };
20092 {
20093 use cudarc::driver::sys;
20094 let src = pin + (half * n_embd * 2) as u64;
20095 let dst = {
20096 let _r1 = rank1.gpu.enter_main()?;
20097 let rstream = rank1.stream();
20098 let (d, _g) = w1.device_ptr(&rstream);
20099 d as u64
20100 };
20101 let _r1 = rank1.gpu.enter_main()?;
20102 let r = unsafe {
20103 sys::cuMemcpyAsync(
20104 dst as sys::CUdeviceptr,
20105 src as sys::CUdeviceptr,
20106 hi_rows * n_embd * 2,
20107 rank1.stream().cu_stream() as sys::CUstream,
20108 )
20109 };
20110 if r != sys::CUresult::CUDA_SUCCESS {
20111 return Err(format!("head split replica upload: {r:?}").into());
20112 }
20113 rank1.stream().synchronize()?;
20114 }
20115 let (logits_e, ev_hn) = {
20116 let _main = e.gpu.enter_main()?;
20117 (e.htod(&vec![0.0f32; n_vocab])?, e.ctx().new_event(None)?)
20118 };
20119 let (raw_hn1, raw_y1) = {
20120 let _r1 = rank1.gpu.enter_main()?;
20121 let rstream = rank1.stream();
20122 let (a, _g0) = hn1.device_ptr(&rstream);
20123 let (b, _g1) = y1.device_ptr(&rstream);
20124 (a as u64, b as u64)
20125 };
20126 let raw_logits_hi = {
20127 let _main = e.gpu.enter_main()?;
20128 let stream = e.stream();
20129 let (l, _g) = logits_e.device_ptr(&stream);
20130 l as u64 + (half * 4) as u64
20131 };
20132 *guard = Some(HeadSplit {
20133 pin,
20134 w1,
20135 hn1,
20136 y1,
20137 logits_e,
20138 ev_hn,
20139 ev_done,
20140 raw_hn1,
20141 raw_y1,
20142 raw_logits_hi,
20143 samp: None,
20144 });
20145 }
20146 let ws = guard.as_mut().expect("armed above");
20147 let hi_rows = n_vocab - half;
20148 let raw_hn = {
20150 let _main = e.gpu.enter_main()?;
20151 let stream = e.stream();
20152 let (h, _g) = hn.device_ptr(&stream);
20153 ws.ev_hn.record(&stream)?;
20154 h as u64
20155 };
20156 {
20157 let _r1 = rank1.gpu.enter_main()?;
20158 rank1.stream().wait(&ws.ev_hn)?;
20159 crate::tp::raw_copy_bytes(ws.raw_hn1, raw_hn, n_embd * 4, rank1)?;
20160 let HeadSplit { w1, hn1, y1, .. } = &mut *ws;
20161 rank1.matvec_bf16_into(w1, hn1, y1, n_embd, hi_rows)?;
20162 crate::tp::raw_copy_bytes(ws.raw_logits_hi, ws.raw_y1, hi_rows * 4, rank1)?;
20163 ws.ev_done.record(&rank1.stream())?;
20164 }
20165 {
20166 let _main = e.gpu.enter_main()?;
20167 let head_lo = head.slice(0..half * n_embd * 2);
20168 let HeadSplit { logits_e, .. } = &mut *ws;
20169 e.matvec_bf16_view_into(&head_lo, hn, logits_e, n_embd, half)?;
20171 e.stream().wait(&ws.ev_done)?;
20172 Ok(Some(()))
20173 }
20174 }
20175
20176 pub(crate) fn head_split_argmax_device(
20181 &self,
20182 e: &Engine,
20183 hn: &CudaSlice<f32>,
20184 token_d: &mut CudaSlice<u32>,
20185 ) -> Result<bool, Box<dyn std::error::Error>> {
20186 if self.head_split_fill_device(e, hn)?.is_none() {
20187 return Ok(false);
20188 }
20189 let n_vocab = self.cfg.n_vocab as usize;
20190 let guard = HEAD_SPLIT_WS
20191 .lock()
20192 .map_err(|_| "head split lock is poisoned")?;
20193 let ws = guard.as_ref().expect("filled above");
20194 let _main = e.gpu.enter_main()?;
20195 e.argmax_token_device_into(&ws.logits_e, token_d, n_vocab)?;
20196 Ok(true)
20197 }
20198
20199 pub(crate) fn head_split_sample_device(
20205 &self,
20206 e: &Engine,
20207 hn: &CudaSlice<f32>,
20208 token_d: &mut CudaSlice<u32>,
20209 samp: &crate::decode_batch::DevSamp,
20210 ctr: u32,
20211 ) -> Result<bool, Box<dyn std::error::Error>> {
20212 if self.head_split_fill_device(e, hn)?.is_none() {
20213 return Ok(false);
20214 }
20215 let n_vocab = self.cfg.n_vocab as usize;
20216 let guard = HEAD_SPLIT_WS
20217 .lock()
20218 .map_err(|_| "head split lock is poisoned")?;
20219 let mut guard = guard;
20220 let ws = guard.as_mut().expect("filled above");
20221 let _main = e.gpu.enter_main()?;
20222 if ws.samp.is_none() {
20223 ws.samp = Some(SampScratch {
20224 pb: e.zeros(n_vocab)?,
20225 th: e.zeros(1)?,
20226 z: e.zeros(1)?,
20227 mx: e.zeros(1)?,
20228 rows: e.htod_i32(&[0i32])?,
20229 });
20230 }
20231 let filtered = samp.top_k > 0 || samp.top_p < 1.0 || samp.min_p > 0.0;
20232 let HeadSplit {
20233 logits_e,
20234 samp: scratch,
20235 ..
20236 } = &mut *ws;
20237 let sc = scratch.as_mut().expect("armed above");
20238 if filtered {
20239 e.filter_stats(
20240 logits_e, n_vocab, &sc.rows, &mut sc.th, &mut sc.z, &mut sc.mx, n_vocab, 1,
20241 samp.temp, samp.top_k, samp.top_p, samp.min_p,
20242 )?;
20243 let SampScratch { pb, th, mx, .. } = sc;
20244 e.gumbel_perturb_filtered_col(
20245 logits_e, 0, pb, n_vocab, samp.seed, ctr, samp.temp, mx, th, 0,
20246 )?;
20247 } else {
20248 e.gumbel_perturb_col(logits_e, 0, &mut sc.pb, n_vocab, samp.seed, ctr, samp.temp)?;
20249 }
20250 e.argmax_token_device_col(&sc.pb, 0, n_vocab, token_d, 0)?;
20251 Ok(true)
20252 }
20253
20254 pub(crate) fn head_split_logits_dtoh(
20257 &self,
20258 e: &Engine,
20259 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
20260 let guard = HEAD_SPLIT_WS
20261 .lock()
20262 .map_err(|_| "head split lock is poisoned")?;
20263 let ws = guard.as_ref().ok_or("head split logits not armed")?;
20264 let _main = e.gpu.enter_main()?;
20265 Ok(e.dtoh(&ws.logits_e)?)
20266 }
20267
20268 pub fn step35_token_graph_chunk(
20277 &self,
20278 e: &Engine,
20279 token: u32,
20280 k_target: usize,
20281 cache: &mut Cache,
20282 ) -> Result<Option<(Vec<u32>, Vec<f32>)>, Box<dyn std::error::Error>> {
20283 if !self.uses_sliding_gated_moe_program()
20284 || !crate::tp::step_tp_graph_enabled()?
20285 || !crate::tp::step_tp_dcw_enabled()?
20286 || !crate::tp::step_tp_qkv_fused_enabled()?
20287 || !crate::tp::step_tp_dev_router_enabled()?
20288 || !crate::tp::step_nvfp4_dev_routes_enabled()?
20289 {
20290 return Ok(None);
20291 }
20292 if !crate::spec::graph_launch_headroom_ok(e) {
20295 static NOTED: std::sync::Once = std::sync::Once::new();
20296 NOTED.call_once(|| crate::spec::graph_replay_suspended_note("step-tp-token"));
20297 return Ok(None);
20298 }
20299 let n_layers = self.layers.len();
20300 let pos = cache.pos;
20301 let staged_next = pos + 1;
20302 if staged_next < 96 {
20303 return Ok(None);
20304 }
20305 let (fa_vec, n_splits) = e.fa_geom_eager(staged_next, 128, 8, false);
20308 if !fa_vec {
20309 return Ok(None);
20310 }
20311 let sp = crate::fa_split_keys(staged_next, 8);
20312 let bucket_max = (n_splits * sp).max(staged_next);
20313 let to_boundary = bucket_max.saturating_sub(staged_next) + 1;
20314 let mut k = k_target.min(to_boundary).min(16);
20315 if k < 2 {
20316 return Ok(None);
20317 }
20318 for il in 0..n_layers {
20320 let Some(tp_kv) = cache.tp_kv[il].as_ref() else {
20321 return Ok(None);
20322 };
20323 while k >= 2 && tp_kv.peek_append_ring(k)?.1 {
20324 k -= 1;
20325 }
20326 if k < 2 {
20327 return Ok(None);
20328 }
20329 }
20330
20331 let mut state_guard = self
20332 .step35_token_graph
20333 .lock()
20334 .map_err(|_| "step35 token graph lock is poisoned")?;
20335 let Some(state) = state_guard.as_mut() else {
20336 return Ok(None); };
20338 if state.graphs.is_empty() {
20339 return Ok(None);
20340 }
20341 {
20342 let (b, g) = state.graphs.first_mut().expect("checked above");
20343 if *b != bucket_max {
20344 g.retarget_bucket(bucket_max)?;
20345 *b = bucket_max;
20346 }
20347 }
20348 let graph = state.graphs.first().map(|(_, g)| g).expect("checked above");
20349
20350 {
20352 let fa0 = match &self.layers[0].mixer {
20353 Mixer::Full(fa) => fa,
20354 _ => return Err("step35 token graph expects full-attention layers".into()),
20355 };
20356 let tp0 = fa0
20357 .step_tp_qkv
20358 .as_ref()
20359 .ok_or("step35 token graph lost its TP state")?;
20360 for rank in 0..tp0.runtime.devices().len() {
20361 let engine = tp0
20362 .runtime
20363 .rank_engine(rank)
20364 .ok_or("step35 token graph lost a rank engine")?;
20365 let _main = engine.gpu.enter_main()?;
20366 engine.stream().synchronize()?;
20367 }
20368 }
20369
20370 {
20373 let _main = e.gpu.enter_main()?;
20374 e.set_u32_one(&mut state.token_d, token)?;
20375 e.set_i32_one(&mut state.pos_d, pos as i32)?;
20376 e.set_i32_one(&mut state.hist_idx, 0)?;
20377 }
20378 for _ in 0..k {
20379 graph.launch(e)?;
20380 }
20381 for il in 0..n_layers {
20383 let tp_kv = cache.tp_kv[il].as_mut().expect("eligibility checked above");
20384 let transaction = tp_kv.begin_transaction()?;
20385 let fa = match &self.layers[il].mixer {
20386 Mixer::Full(fa) => fa,
20387 _ => return Err("step35 token graph expects full-attention layers".into()),
20388 };
20389 let tp = fa
20390 .step_tp_qkv
20391 .as_ref()
20392 .ok_or("step35 token graph lost its TP state")?;
20393 let empty: [CudaSlice<f32>; 0] = [];
20394 tp.runtime.append_tp_kv_transaction_inner(
20395 tp_kv,
20396 transaction,
20397 &empty,
20398 &empty,
20399 k,
20400 true,
20401 )?;
20402 tp.runtime
20403 .commit_tp_kv_transaction_external(tp_kv, transaction, k)?;
20404 if let Some(local) = cache.kv[il].as_mut() {
20405 local.len = pos + k;
20406 let _main = e.gpu.enter_main()?;
20407 e.set_i32_one(&mut local.len_d, (pos + k) as i32)?;
20408 }
20409 }
20410 cache.pos = pos + k;
20411 let (hist, logits) = {
20412 let _main = e.gpu.enter_main()?;
20413 e.stream().synchronize()?;
20414 (e.dtoh_u32(&state.token_hist)?, e.dtoh(&state.logits_stage)?)
20415 };
20416 Ok(Some((hist[..k].to_vec(), logits)))
20417 }
20418}
20419
20420impl HybridModel {
20421 #[allow(clippy::too_many_arguments)]
20426 fn step35_token_graph_build(
20427 &self,
20428 e: &Engine,
20429 cache: &mut Cache,
20430 state: &mut Step35TokenGraphState,
20431 bucket_max: usize,
20432 ) -> Result<(), Box<dyn std::error::Error>> {
20433 use cudarc::driver::DevicePtr;
20434 let n_embd = self.cfg.n_embd as usize;
20435 let eps = self.cfg.rms_eps;
20436 let n_layers = self.layers.len();
20437 let started = std::time::Instant::now();
20438 if !crate::router_kernel_on() {
20439 return Err(
20440 "step35 token graph requires the router kernel (MEMRA_ROUTER_KERNEL=0)".into(),
20441 );
20442 }
20443 if !Engine::bf16_mmv_on() || n_embd % 8 != 0 {
20444 return Err("step35 token graph requires MEMRA_BF16_MMV bf16-resident matvecs".into());
20445 }
20446
20447 let embd_gpu = self
20449 .embd_gpu_try(e)
20450 .ok_or("step35 token graph could not upload the device embed table")?;
20451 let embd_qtype = match self.embd.ggml_type {
20452 memra_gguf::GgmlType::BF16 => crate::QT_BF16,
20453 memra_gguf::GgmlType::Q8_0 => crate::QT_Q8_0,
20454 other => return Err(format!("token graph embed dtype {other:?} unhandled").into()),
20455 };
20456 let embd_row_bytes = self.embd.raw.len() / self.cfg.n_vocab as usize;
20457
20458 let (p_mixed, p_kshadow, p_vshadow) = {
20460 let _main = e.gpu.enter_main()?;
20461 let stream = e.stream();
20462 let (a, _g) = state.mixed_stage.device_ptr(&stream);
20463 let (b, _g) = state.k_shadow_stage.device_ptr(&stream);
20464 let (c, _g) = state.v_shadow_stage.device_ptr(&stream);
20465 (a as u64, b as u64, c as u64)
20466 };
20467
20468 crate::tp::token_graph_build_begin()?;
20469 let mut group_id: u32 = 0;
20470 for il in 0..n_layers {
20471 let layer = &self.layers[il];
20472 let fa = match &layer.mixer {
20473 Mixer::Full(fa) => fa,
20474 _ => return Err("step35 token graph expects full-attention layers".into()),
20475 };
20476 let tp = fa
20477 .step_tp_qkv
20478 .as_ref()
20479 .ok_or("step35 token graph lost its TP state")?;
20480 let attention = tp
20481 .attention
20482 .as_ref()
20483 .ok_or("step35 token graph lost its attention aux")?;
20484 let geometry = self.step35_geom(il);
20485 let window = geometry.window.map(|w| w as usize);
20486 let head_dim = geometry.head_dim_k as usize;
20487 let heads = geometry.n_head as usize;
20488 let kv_heads = geometry.n_head_kv as usize;
20489 let ranks = tp.runtime.devices().len();
20490 let local_heads = heads / ranks;
20491 let local_kv_heads = kv_heads / ranks;
20492 let layer_bucket = window.map(|w| bucket_max.min(w)).unwrap_or(bucket_max);
20493 let use_gate_shards =
20494 attention.gate_shards.is_some() || attention.gate_shards_bf16.is_some();
20495 if !use_gate_shards {
20496 return Err("step35 token graph requires the fused gate shards".into());
20497 }
20498
20499 let ws_index = tp
20500 .runtime
20501 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
20502 let ws_mutex = tp.runtime.decode_v2_workspace();
20503 let mut ws_guard = ws_mutex
20504 .lock()
20505 .map_err(|_| "step TP decode v2 workspace lock is poisoned")?;
20506 let ws = ws_guard
20507 .get_mut(ws_index)
20508 .ok_or("step TP decode v2 workspace missing after ensure")?;
20509 tp.runtime
20510 .decode_v2_arm_token_mirrors(ws, p_mixed, (p_kshadow, p_vshadow))?;
20511 let mut rope_freqs = Vec::with_capacity(ranks);
20512 for rank in 0..ranks {
20513 let engine = tp
20514 .runtime
20515 .rank_engine(rank)
20516 .ok_or("step35 token graph lost a rank engine")?;
20517 rope_freqs.push(if geometry.rope_factors {
20518 self.step35_aux
20519 .as_ref()
20520 .and_then(|aux| aux.rope_freqs(engine))
20521 } else {
20522 None
20523 });
20524 }
20525 let gate_shards_arg = if let Some(shards) = attention.gate_shards.as_deref() {
20526 Some(crate::tp::StepTpGateShards::F32(shards))
20527 } else {
20528 attention
20529 .gate_shards_bf16
20530 .as_deref()
20531 .map(crate::tp::StepTpGateShards::Bf16)
20532 };
20533
20534 let decode_input = attention
20536 .decode_input
20537 .as_ref()
20538 .ok_or("step35 token graph requires the replicated decode input")?;
20539 let mut decode_input = decode_input
20540 .lock()
20541 .map_err(|_| "replicated decode input lock is poisoned")?;
20542 if ws.h_stage.is_none() {
20544 return Err(
20545 "step35 token graph requires the stage flow armed (run eager dcw first)".into(),
20546 );
20547 }
20548 {
20549 let state_x = &mut state.x;
20550 let token_d = &state.token_d;
20551 let pos_d = &state.pos_d;
20552 crate::tp::graph_section(e, None, || {
20553 let _main = e.gpu.enter_main()?;
20554 if il == 0 {
20555 e.embed_gather_device_into(
20556 embd_gpu,
20557 token_d,
20558 state_x,
20559 n_embd,
20560 embd_qtype,
20561 embd_row_bytes,
20562 )?;
20563 }
20564 {
20565 let h_stage = ws.h_stage.as_mut().expect("stage armed checked above");
20566 e.rms_norm(
20567 state_x,
20568 layer.attn_norm.float_data(),
20569 h_stage,
20570 n_embd,
20571 1,
20572 eps,
20573 )?;
20574 }
20575 {
20576 let pos_stage = ws.pos_stage.as_mut().expect("stage armed above");
20577 let mut dst = pos_stage.slice_mut(0..1);
20578 e.stream().memcpy_dtod(&pos_d.slice(0..1), &mut dst)?;
20579 }
20580 Ok(())
20581 })?;
20582 }
20583
20584 group_id += 1;
20586 for rank in 0..ranks {
20587 let engine = tp
20588 .runtime
20589 .rank_engine(rank)
20590 .ok_or("step35 token graph lost a rank engine")?;
20591 {
20592 let ceiling = window
20597 .map(|w| cache.max_ctx.min(w))
20598 .unwrap_or(cache.max_ctx);
20599 let _main = engine.gpu.enter_main()?;
20600 engine.fa_dcw_pool_ensure(
20601 head_dim,
20602 local_heads,
20603 local_kv_heads,
20604 ceiling.min(2048),
20605 )?;
20606 engine.fa_dcw_pool_ensure(head_dim, local_heads, local_kv_heads, ceiling)?;
20607 engine.fa_dcw_pool_ensure(
20608 head_dim,
20609 local_heads,
20610 local_kv_heads,
20611 layer_bucket,
20612 )?;
20613 }
20614 let runtime = &tp.runtime;
20615 let q_norm = &attention.q_norm;
20616 let k_norm = &attention.k_norm;
20617 let gate_ref = gate_shards_arg.as_ref();
20618 crate::tp::graph_section(engine, Some(group_id), || {
20619 runtime.decode_v2_input_qkv_rank(
20620 ws,
20621 &state.pos_d,
20622 &mut decode_input,
20623 &tp.q,
20624 &tp.k,
20625 &tp.v,
20626 q_norm,
20627 k_norm,
20628 head_dim,
20629 geometry.n_rot as usize,
20630 geometry.rope_base,
20631 &rope_freqs,
20632 eps,
20633 gate_ref,
20634 true,
20635 false,
20636 rank,
20637 None,
20638 )?;
20639 let distributed = cache.tp_kv[il]
20642 .as_mut()
20643 .ok_or("step35 token graph lost a TP cache")?;
20644 let (kv_dim_k, kv_dim_v) = (distributed.kv_dim_k(), distributed.kv_dim_v());
20645 let (ktb, vtb) = (distributed.k_tok_bytes(), distributed.v_tok_bytes());
20646 let capacity = distributed.physical_capacity();
20647 {
20648 let rank_cache = distributed
20649 .rank_mut(rank)
20650 .ok_or("step35 token graph lost a rank cache")?;
20651 let (k_plane, v_plane, len_d, base_d) =
20652 rank_cache.planes_and_counters_mut();
20653 engine.append_kv_quantized_dcw(
20654 &ws.k[rank],
20655 &ws.v_raw[rank],
20656 k_plane,
20657 v_plane,
20658 len_d,
20659 base_d,
20660 kv_dim_k,
20661 kv_dim_v,
20662 ktb,
20663 vtb,
20664 )?;
20665 }
20666 {
20667 let rank_cache = distributed
20668 .rank_mut(rank)
20669 .ok_or("step35 token graph lost a rank cache")?;
20670 engine.inc_i32(rank_cache.len_d_mut())?;
20671 }
20672 let rank_cache = distributed
20673 .rank(rank)
20674 .ok_or("step35 token graph lost a rank cache")?;
20675 let k_ring = engine.view_u8_range(rank_cache.k(), 0, capacity * ktb);
20676 let v_ring = engine.view_u8_range(rank_cache.v(), 0, capacity * vtb);
20677 engine.fa_decode_dcw(
20682 &ws.q[rank],
20683 &k_ring,
20684 &v_ring,
20685 &mut ws.attn_out[rank],
20686 head_dim,
20687 local_heads,
20688 local_kv_heads,
20689 rank_cache.len_d(),
20690 rank_cache.base_d(),
20691 window.unwrap_or(0),
20692 layer_bucket,
20693 geometry.attention_scale(),
20694 ktb,
20695 vtb,
20696 None,
20697 )?;
20698 engine.attn_head_gate(
20699 &ws.attn_out[rank],
20700 &ws.gate[rank],
20701 &mut ws.gated[rank],
20702 None,
20703 head_dim,
20704 local_heads,
20705 1,
20706 )?;
20707 runtime.decode_v2_finish_rank_partial(ws, &tp.o, true, rank)?;
20708 Ok(())
20709 })?;
20710 }
20711
20712 {
20714 let root = tp
20715 .runtime
20716 .rank_engine(0)
20717 .ok_or("step35 token graph lost the root engine")?;
20718 let runtime = &tp.runtime;
20719 crate::tp::graph_section(root, None, || runtime.decode_v2_finish_root_fused(ws))?;
20720 }
20721 drop(ws_guard);
20722 drop(decode_input);
20723
20724 let probe_layer: Option<usize> = std::env::var("MEMRA_TG_PROBE_LAYER")
20725 .ok()
20726 .and_then(|v| v.parse().ok());
20727 if probe_layer == Some(il) {
20728 let Step35TokenGraphState {
20729 mixed_stage,
20730 probe_mixed,
20731 ..
20732 } = &mut *state;
20733 crate::tp::graph_section(e, None, || {
20734 let _main = e.gpu.enter_main()?;
20735 let mut dst = probe_mixed.slice_mut(0..n_embd);
20736 e.stream()
20737 .memcpy_dtod(&mixed_stage.slice(0..n_embd), &mut dst)?;
20738 Ok(())
20739 })?;
20740 }
20741
20742 match &layer.ffn {
20744 crate::hybrid::Ffn::Dense {
20745 ffn_gate,
20746 ffn_up,
20747 ffn_down,
20748 } => {
20749 let n_ff = ffn_gate.out_features();
20750 let lim = self.cfg.clamp_shexp_at(il as u32);
20751 if lim.is_some() {
20755 return Err("step35 token graph dense FFN with clamp unsupported".into());
20756 }
20757 let (wg_d, wu_d, wd_d) = match (ffn_gate, ffn_up, ffn_down) {
20758 (
20759 crate::model::GpuTensor::FloatBf16 { data: wg, .. },
20760 crate::model::GpuTensor::FloatBf16 { data: wu, .. },
20761 crate::model::GpuTensor::FloatBf16 { data: wd, .. },
20762 ) => (wg, wu, wd),
20763 _ => {
20764 return Err(
20765 "step35 token graph dense FFN requires bf16-resident weights"
20766 .into(),
20767 );
20768 }
20769 };
20770 crate::tp::graph_section(e, None, || {
20771 let _main = e.gpu.enter_main()?;
20772 let Step35TokenGraphState {
20773 x,
20774 x1,
20775 mixed_stage,
20776 dense_z,
20777 dense_gate,
20778 dense_up,
20779 dense_act,
20780 sh_stage,
20781 ..
20782 } = &mut *state;
20783 e.add_rms_norm(
20784 x,
20785 mixed_stage,
20786 layer.post_attn_norm.float_data(),
20787 x1,
20788 dense_z,
20789 n_embd,
20790 1,
20791 eps,
20792 )?;
20793 e.matvec_bf16_into(wg_d, dense_z, dense_gate, n_embd, n_ff)?;
20797 e.matvec_bf16_into(wu_d, dense_z, dense_up, n_embd, n_ff)?;
20798 Self::ffn_act_lim(
20799 e, &self.cfg, dense_gate, dense_up, 1.0, 1.0, lim, dense_act, n_ff,
20800 )?;
20801 e.matvec_bf16_into(wd_d, dense_act, sh_stage, n_ff, n_embd)?;
20802 e.add(x1, sh_stage, x, n_embd)?;
20803 Ok(())
20804 })?;
20805 }
20806 crate::hybrid::Ffn::Moe(m) => {
20807 let moe = self
20808 .cfg
20809 .moe
20810 .as_ref()
20811 .ok_or("step35 token graph needs moe cfg")?;
20812 let n_expert = moe.expert_count as usize;
20813 let n_used = moe.expert_used_count as usize;
20814 let sigmoid = self
20815 .cfg
20816 .sigmoid_router()
20817 .ok_or("step35 token graph needs the sigmoid router")?;
20818 let step_tp = m
20819 .step_tp
20820 .as_ref()
20821 .ok_or("step35 token graph needs TP experts")?;
20822 let bank = match &step_tp.experts {
20823 crate::hybrid::StepTpExpertBank::Nvfp4(bank) => bank,
20824 _ => return Err("step35 token graph needs the NVFP4 bank".into()),
20825 };
20826 let routes_ws_mutex = bank.device_workspace_handle();
20827 let mut routes_guard = routes_ws_mutex
20828 .lock()
20829 .map_err(|_| "routes workspace lock is poisoned")?;
20830 let routes_ws = routes_guard
20831 .as_mut()
20832 .ok_or("step35 token graph requires the routes workspace warmed")?;
20833 routes_ws.arm_stages(e, bank.input_width, n_used)?;
20834 step_tp.runtime.routes_arm_raw(bank, routes_ws)?;
20835 let p_z = {
20836 let root = step_tp
20837 .runtime
20838 .rank_engine(0)
20839 .ok_or("routes root engine missing")?;
20840 let _main = root.gpu.enter_main()?;
20841 let stream = root.stream();
20842 let in_stage = routes_ws
20843 .in_stage_handle()
20844 .ok_or("routes in stage not armed")?;
20845 let (a, _g) = in_stage.device_ptr(&stream);
20846 a as u64
20847 };
20848 let local_out = bank.expert_width / ranks;
20849
20850 crate::tp::graph_section(e, None, || {
20852 let _main = e.gpu.enter_main()?;
20853 {
20854 let in_stage = routes_ws
20855 .in_stage_mut()
20856 .ok_or("routes in stage not armed")?;
20857 let Step35TokenGraphState {
20858 x, x1, mixed_stage, ..
20859 } = &mut *state;
20860 e.add_rms_norm(
20861 x,
20862 mixed_stage,
20863 layer.post_attn_norm.float_data(),
20864 x1,
20865 in_stage,
20866 n_embd,
20867 1,
20868 eps,
20869 )?;
20870 }
20871 {
20872 let z_ref = routes_ws
20873 .in_stage_handle()
20874 .ok_or("routes in stage not armed")?;
20875 e.router_gemv_into(
20876 m.gate_inp.float_data(),
20877 z_ref,
20878 &mut state.router_logits,
20879 n_embd,
20880 n_expert,
20881 1,
20882 )?;
20883 }
20884 let (sel_e, w_e) = routes_ws
20885 .dev_route_e_mut()
20886 .ok_or("routes staging not armed")?;
20887 e.moe_router_sigmoid_topk_into(
20888 &state.router_logits,
20889 1,
20890 n_expert,
20891 n_used,
20892 m.active_count(),
20893 &m.exp_probs_b_dev,
20894 &m.active_experts_dev,
20895 sigmoid.0,
20896 sigmoid.1,
20897 sel_e,
20898 w_e,
20899 )?;
20900 Ok(())
20901 })?;
20902
20903 group_id += 1;
20905 for rank in 0..ranks {
20906 let engine = step_tp
20907 .runtime
20908 .rank_engine(rank)
20909 .ok_or("routes rank engine missing")?;
20910 let runtime = &step_tp.runtime;
20911 crate::tp::graph_section(engine, Some(group_id), || {
20912 runtime.routes_rank_section(
20913 bank,
20914 routes_ws,
20915 p_z,
20916 local_out,
20917 n_used,
20918 step_tp.activation_limit,
20919 rank,
20920 )
20921 })?;
20922 }
20923
20924 {
20926 let root = step_tp
20927 .runtime
20928 .rank_engine(0)
20929 .ok_or("routes root engine missing")?;
20930 let runtime = &step_tp.runtime;
20931 crate::tp::graph_section(root, None, || {
20932 runtime.routes_root_section(bank, routes_ws)
20933 })?;
20934 }
20935
20936 let lim_sh = self.cfg.clamp_shexp_at(il as u32);
20940 let (wg_sh, wu_sh, wd_sh) = match (&m.gate_shexp, &m.up_shexp, &m.down_shexp) {
20941 (
20942 Some(crate::model::GpuTensor::FloatBf16 { data: wg, .. }),
20943 Some(crate::model::GpuTensor::FloatBf16 { data: wu, .. }),
20944 Some(crate::model::GpuTensor::FloatBf16 { data: wd, .. }),
20945 ) => (wg, wu, wd),
20946 _ => {
20947 return Err(
20948 "step35 token graph shexp requires bf16-resident weights".into()
20949 );
20950 }
20951 };
20952 let n_ff_sh = m
20953 .gate_shexp
20954 .as_ref()
20955 .expect("matched Some above")
20956 .out_features();
20957 let gate_inp_shexp = m.gate_inp_shexp.as_ref();
20960 crate::tp::graph_section(e, None, || {
20961 let _main = e.gpu.enter_main()?;
20962 let (z_ref, out_stage) = routes_ws
20963 .in_and_out_stages_mut()
20964 .ok_or("routes stages not armed")?;
20965 let Step35TokenGraphState {
20966 x,
20967 x1,
20968 sh_stage,
20969 shexp_gate,
20970 shexp_up,
20971 shexp_act,
20972 gate_sig,
20973 ..
20974 } = &mut *state;
20975 e.matvec_bf16_dual_into(
20976 wg_sh, wu_sh, z_ref, shexp_gate, shexp_up, n_embd, n_ff_sh,
20977 )?;
20978 Self::ffn_act_lim(
20979 e, &self.cfg, shexp_gate, shexp_up, 1.0, 1.0, lim_sh, shexp_act,
20980 n_ff_sh,
20981 )?;
20982 e.matvec_bf16_into(wd_sh, shexp_act, sh_stage, n_ff_sh, n_embd)?;
20983 if let Some(gate_w) = gate_inp_shexp {
20984 e.sigmoid_dot_rows_into(
20985 z_ref,
20986 gate_w.float_data(),
20987 gate_sig,
20988 n_embd,
20989 1,
20990 )?;
20991 }
20992 e.add_scaled_rows(sh_stage, gate_sig, out_stage, n_embd, 1)?;
20993 e.add(x1, out_stage, x, n_embd)?;
20994 Ok(())
20995 })?;
20996 }
20997 }
20998 if probe_layer == Some(il) {
20999 let Step35TokenGraphState { x, probe_x, .. } = &mut *state;
21000 crate::tp::graph_section(e, None, || {
21001 let _main = e.gpu.enter_main()?;
21002 let mut dst = probe_x.slice_mut(0..n_embd);
21003 e.stream().memcpy_dtod(&x.slice(0..n_embd), &mut dst)?;
21004 Ok(())
21005 })?;
21006 }
21007 }
21008
21009 let head = match &self.output {
21011 crate::model::GpuTensor::FloatBf16 { data, .. } => data,
21012 _ => return Err("step35 token graph head requires the bf16-resident output".into()),
21013 };
21014 crate::tp::graph_section(e, None, || {
21015 let _main = e.gpu.enter_main()?;
21016 let Step35TokenGraphState {
21017 x,
21018 hn,
21019 logits_stage,
21020 token_d,
21021 pos_d,
21022 token_hist,
21023 hist_idx,
21024 ..
21025 } = &mut *state;
21026 e.rms_norm(x, self.output_norm.float_data(), hn, n_embd, 1, eps)?;
21027 e.matvec_bf16_into(head, hn, logits_stage, n_embd, self.cfg.n_vocab as usize)?;
21028 e.argmax_token_device_into(logits_stage, token_d, self.cfg.n_vocab as usize)?;
21034 e.u32_hist_append(token_d, token_hist, hist_idx)?;
21035 e.inc_i32(pos_d)?;
21036 Ok(())
21037 })?;
21038
21039 let graph = crate::tp::token_graph_build_finish()?;
21040 state.graphs.push((bucket_max, graph));
21041 eprintln!(
21042 "[step35-token-graph] built bucket={bucket_max} layers={n_layers} \
21043 build_ms={:.0} performance_claim=false",
21044 started.elapsed().as_secs_f64() * 1e3
21045 );
21046 Ok(())
21047 }
21048}