1#![allow(clippy::needless_range_loop)]
8
9use crate::Engine;
10use crate::cache::Cache;
11use cudarc::driver::CudaSlice;
12use memra_gguf::config::{ModelConfig, SwigluClamp};
13
14pub struct PrimeSlabs {
18 pub t_cap: usize,
19 pub h: CudaSlice<f32>,
20 pub x1: CudaSlice<f32>,
21 pub z: CudaSlice<f32>,
22 pub act: CudaSlice<f32>,
23 pub xa: CudaSlice<f32>,
24 pub xb: CudaSlice<f32>,
25 pub h16: CudaSlice<u8>,
26 pub z16: CudaSlice<u8>,
27 pub gate: CudaSlice<f32>, pub up: CudaSlice<f32>, pub ffn_out: CudaSlice<f32>, pub seg_glue: Vec<Option<cudarc::driver::CudaGraph>>,
37 pub mixed: CudaSlice<f32>,
41 pub seg_mid: Vec<Option<cudarc::driver::CudaGraph>>,
42 pub seg_t: usize,
43}
44
45unsafe impl Send for PrimeSlabs {}
48
49fn shexp_gate_up_t1(
58 e: &Engine,
59 gate_shexp: &crate::model::GpuTensor,
60 up_shexp: &crate::model::GpuTensor,
61 z: &CudaSlice<f32>,
62 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
63) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
64 let is_nvfp4 = |w: &crate::model::GpuTensor| matches!(w, crate::model::GpuTensor::Quant { qtype, .. } if *qtype == crate::QT_NVFP4);
65 if is_nvfp4(gate_shexp) && is_nvfp4(up_shexp) {
66 let pair = match zq8 {
70 Some((zq, zd)) => e.matmul_nvfp4_fused2(gate_shexp, up_shexp, zq, zd, 1)?,
71 None => {
72 let (zq, zd) = e.quantize_q8_1(z, 1, gate_shexp.in_features())?;
73 e.matmul_nvfp4_fused2(gate_shexp, up_shexp, &zq, &zd, 1)?
74 }
75 };
76 if let Some(pair) = pair {
77 return Ok(pair);
78 }
79 }
80 match e.matmul_q8_fused2_x(gate_shexp, up_shexp, z)? {
81 Some(pair) => Ok(pair),
82 None => Ok((e.matmul(gate_shexp, z, 1)?, e.matmul(up_shexp, z, 1)?)),
83 }
84}
85
86fn active_matrix_values(
87 available: usize,
88 rows: usize,
89 columns: usize,
90 label: &str,
91) -> Result<usize, String> {
92 let required = rows
93 .checked_mul(columns)
94 .ok_or_else(|| format!("{label} shape overflows: {rows}x{columns}"))?;
95 if available < required {
96 return Err(format!(
97 "{label} has {available} values, fewer than the active {rows}x{columns} ({required})"
98 ));
99 }
100 Ok(required)
101}
102
103fn step_grouped_decode_shape(prefill: bool, tokens: usize) -> bool {
104 !prefill && tokens == 1
105}
106
107fn parse_step_ep_grouped_prefill(value: Option<&str>) -> Result<bool, String> {
108 match value {
109 None | Some("") | Some("0") => Ok(false),
110 Some("1") => Ok(true),
111 Some(value) => Err(format!(
112 "MEMRA_STEP_EP_GROUPED_PREFILL={value:?} is invalid; expected 0 or 1"
113 )),
114 }
115}
116
117fn step_ep_grouped_prefill_enabled() -> Result<bool, String> {
118 parse_step_ep_grouped_prefill(
119 std::env::var("MEMRA_STEP_EP_GROUPED_PREFILL")
120 .ok()
121 .as_deref(),
122 )
123}
124
125fn step_grouped_prefill_shape(enabled: bool, prefill: bool, tokens: usize) -> bool {
126 enabled && prefill && (PRIME_MIN_T..=crate::cache::PRIME_CHUNK_MAX_TOKENS).contains(&tokens)
127}
128
129fn parse_step_tp_prefill(value: Option<&str>) -> Result<bool, String> {
130 match value {
131 None | Some("") | Some("0") => Ok(false),
132 Some("1") => Ok(true),
133 Some(value) => Err(format!(
134 "MEMRA_STEP_TP_PREFILL={value:?} is invalid; expected 0 or 1"
135 )),
136 }
137}
138
139fn step_tp_prefill_enabled() -> Result<bool, String> {
140 parse_step_tp_prefill(std::env::var("MEMRA_STEP_TP_PREFILL").ok().as_deref())
141}
142
143fn validate_step_prime_batch_modes(tp_prefill: bool, grouped_prefill: bool) -> Result<(), String> {
144 if grouped_prefill && !tp_prefill {
145 return Err("MEMRA_STEP_EP_GROUPED_PREFILL=1 requires MEMRA_STEP_TP_PREFILL=1".into());
146 }
147 if tp_prefill {
148 return Err(
149 "Step TP4 cross-request prime batching did not clear the live-server performance \
150 gate; use per-session grouped prefill"
151 .into(),
152 );
153 }
154 Ok(())
155}
156
157fn step_tp_prefill_shape(
158 enabled: bool,
159 tokens: usize,
160 ranks: usize,
161 native_p2p: bool,
162 has_rank_local_attention: bool,
163 fp8_kv: bool,
164) -> bool {
165 enabled
173 && tokens >= PRIME_MIN_T
174 && matches!(ranks, 2 | 4)
175 && native_p2p
176 && has_rank_local_attention
177 && !fp8_kv
178}
179
180fn empty_cache_layers<T>(n: usize) -> Vec<Option<T>> {
181 std::iter::repeat_with(|| None).take(n).collect()
182}
183
184fn prime_cache_stage_for_layer(fence: &[usize], layer: usize) -> usize {
185 debug_assert!(fence.len() >= 3);
186 match fence[1..fence.len() - 1].binary_search(&layer) {
187 Ok(index) => index + 1,
188 Err(index) => index,
189 }
190}
191
192fn move_prime_cache_layers<T>(
193 parent: &mut [Option<T>],
194 stages: &mut [Vec<Option<T>>],
195 fence: &[usize],
196) {
197 assert_eq!(stages.len() + 1, fence.len());
198 assert!(stages.iter().all(|stage| stage.len() == parent.len()));
199 for (layer, value) in parent.iter_mut().enumerate() {
200 let stage = prime_cache_stage_for_layer(fence, layer);
201 debug_assert!(stages[stage][layer].is_none());
202 stages[stage][layer] = value.take();
203 }
204}
205
206#[cfg(test)]
207fn restore_prime_cache_layers<T>(
208 parent: &mut [Option<T>],
209 stages: &mut [Vec<Option<T>>],
210 fence: &[usize],
211) {
212 assert_eq!(stages.len() + 1, fence.len());
213 assert!(stages.iter().all(|stage| stage.len() == parent.len()));
214 for (layer, value) in parent.iter_mut().enumerate() {
215 let stage = prime_cache_stage_for_layer(fence, layer);
216 debug_assert!(value.is_none());
217 *value = stages[stage][layer].take();
218 }
219}
220
221struct PrimeCacheStages<'a> {
226 parent: &'a mut Cache,
227 fence: Vec<usize>,
228 stages: Vec<std::sync::Mutex<Cache>>,
229 committed: bool,
230}
231
232impl<'a> PrimeCacheStages<'a> {
233 fn new(parent: &'a mut Cache, fence: &[usize]) -> Self {
234 let n = parent.kv.len();
235 assert_eq!(parent.recur.len(), n, "cache layer vectors disagree");
236 assert_eq!(parent.tp_kv.len(), n, "cache layer vectors disagree");
237 assert_eq!(parent.latent.len(), n, "cache layer vectors disagree");
238 let n_stages = fence.len().checked_sub(1).expect("PP cache fence is empty");
239 assert!((2..=4).contains(&n_stages), "PP cache needs 2..=4 stages");
240 assert_eq!(fence[0], 0, "PP cache fence must start at layer zero");
241 assert!(
242 fence.windows(2).all(|pair| pair[0] < pair[1]),
243 "PP cache fence must be strictly increasing"
244 );
245 assert!(fence[n_stages] <= n, "PP cache fence exceeds {n} layers");
246
247 let mut latent: Vec<_> = (0..n_stages).map(|_| empty_cache_layers(n)).collect();
248 let mut g5_recur: Vec<_> = (0..n_stages).map(|_| empty_cache_layers(n)).collect();
249 let mut g5_latent: Vec<_> = (0..n_stages).map(|_| empty_cache_layers(n)).collect();
250 let mut kv: Vec<_> = (0..n_stages).map(|_| empty_cache_layers(n)).collect();
251 let mut tp_kv: Vec<_> = (0..n_stages).map(|_| empty_cache_layers(n)).collect();
252 let mut recur: Vec<_> = (0..n_stages).map(|_| empty_cache_layers(n)).collect();
253 move_prime_cache_layers(&mut parent.kv, &mut kv, fence);
254 move_prime_cache_layers(&mut parent.tp_kv, &mut tp_kv, fence);
255 move_prime_cache_layers(&mut parent.recur, &mut recur, fence);
256
257 move_prime_cache_layers(&mut parent.latent, &mut latent, fence);
258 move_prime_cache_layers(&mut parent.glm5_tp_recur, &mut g5_recur, fence);
259 move_prime_cache_layers(&mut parent.glm5_tp_latent_peer, &mut g5_latent, fence);
260 let pos = parent.pos;
261 let max_ctx = parent.max_ctx;
262 let stages = (0..n_stages)
266 .map(|stage| {
267 std::sync::Mutex::new(Cache {
268 kv: std::mem::take(&mut kv[stage]),
269 tp_kv: std::mem::take(&mut tp_kv[stage]),
270 recur: std::mem::take(&mut recur[stage]),
271 latent: std::mem::take(&mut latent[stage]),
272 glm5_tp_recur: std::mem::take(&mut g5_recur[stage]),
273 glm5_tp_latent_peer: std::mem::take(&mut g5_latent[stage]),
274 pos,
275 max_ctx,
276 tainted: false,
277 last_logits_dev: None,
278 dflash_taps: None,
279 hc_taps: None,
280 })
281 })
282 .collect();
283 Self {
284 parent,
285 fence: fence.to_vec(),
286 stages,
287 committed: false,
288 }
289 }
290
291 fn pp2_parts(&mut self) -> (&mut Cache, &mut Cache) {
292 assert_eq!(self.stages.len(), 2);
293 let (stage0, stage1) = self.stages.split_at_mut(1);
294 (
295 stage0[0]
296 .get_mut()
297 .unwrap_or_else(|poisoned| poisoned.into_inner()),
298 stage1[0]
299 .get_mut()
300 .unwrap_or_else(|poisoned| poisoned.into_inner()),
301 )
302 }
303
304 fn stages(&self) -> &[std::sync::Mutex<Cache>] {
305 &self.stages
306 }
307
308 fn commit(&mut self) {
309 self.committed = true;
310 }
311}
312
313impl Drop for PrimeCacheStages<'_> {
314 fn drop(&mut self) {
315 let n = self.parent.kv.len();
316 for i in 0..n {
317 let stage = prime_cache_stage_for_layer(&self.fence, i);
318 let source = self.stages[stage]
319 .get_mut()
320 .unwrap_or_else(|poisoned| poisoned.into_inner());
321 debug_assert!(self.parent.kv[i].is_none());
322 debug_assert!(self.parent.tp_kv[i].is_none());
323 debug_assert!(self.parent.recur[i].is_none());
324 debug_assert!(self.parent.latent[i].is_none());
325 self.parent.kv[i] = source.kv[i].take();
326 self.parent.tp_kv[i] = source.tp_kv[i].take();
327 self.parent.recur[i] = source.recur[i].take();
328 self.parent.latent[i] = source.latent[i].take();
329 self.parent.glm5_tp_recur[i] = source.glm5_tp_recur[i].take();
330 self.parent.glm5_tp_latent_peer[i] = source.glm5_tp_latent_peer[i].take();
331 }
332 self.parent.pos = self
333 .stages
334 .iter_mut()
335 .map(|stage| {
336 stage
337 .get_mut()
338 .unwrap_or_else(|poisoned| poisoned.into_inner())
339 .pos
340 })
341 .min()
342 .unwrap_or(self.parent.pos);
343 if !self.committed {
344 self.parent.mark_tainted();
345 }
346 }
347}
348
349struct CacheTaintGuard {
353 caches: Vec<*mut Cache>,
354 committed: bool,
355}
356
357impl CacheTaintGuard {
358 fn arm(caches: &mut [&mut Cache]) -> Self {
359 Self {
360 caches: caches
361 .iter_mut()
362 .map(|cache| *cache as *mut Cache)
363 .collect(),
364 committed: false,
365 }
366 }
367
368 fn commit(&mut self) {
369 self.committed = true;
370 }
371}
372
373impl Drop for CacheTaintGuard {
374 fn drop(&mut self) {
375 if self.committed {
376 return;
377 }
378 for cache in &self.caches {
379 unsafe { (&mut **cache).mark_tainted() };
382 }
383 }
384}
385
386#[derive(Debug, Clone, Copy, PartialEq, Eq)]
387struct PrimePpWaveSlot {
388 wave: usize,
389 slot: usize,
390}
391
392#[derive(Debug)]
393enum PrimePpSignal {
394 Slot(PrimePpWaveSlot),
395 Error(String),
396}
397
398#[derive(Default)]
399struct PrimePpWaveCredits {
400 next_wave: usize,
401 pending: std::collections::VecDeque<PrimePpWaveSlot>,
402}
403
404impl PrimePpWaveCredits {
405 fn release_required(&self) -> Option<PrimePpWaveSlot> {
406 (self.pending.len() == 2).then(|| self.pending[0])
407 }
408
409 fn record_release(&mut self, released: PrimePpWaveSlot) -> Result<(), String> {
410 let expected =
411 self.pending.front().copied().ok_or_else(|| {
412 "prime PP received a slot release with no pending wave".to_string()
413 })?;
414 if released != expected {
415 return Err(format!(
416 "prime PP slot release {:?} does not match oldest pending {:?}",
417 released, expected
418 ));
419 }
420 self.pending.pop_front();
421 Ok(())
422 }
423
424 fn record_send(&mut self, sent: PrimePpWaveSlot) -> Result<(), String> {
425 if sent.wave != self.next_wave {
426 return Err(format!(
427 "prime PP sent wave {} while wave {} was next",
428 sent.wave, self.next_wave
429 ));
430 }
431 if sent.slot >= 2 {
432 return Err(format!(
433 "prime PP boundary returned invalid slot {}",
434 sent.slot
435 ));
436 }
437 if self.pending.iter().any(|pending| pending.slot == sent.slot) {
438 return Err(format!(
439 "prime PP reused slot {} before its exact-wave release",
440 sent.slot
441 ));
442 }
443 self.pending.push_back(sent);
444 self.next_wave += 1;
445 Ok(())
446 }
447}
448
449fn recv_prime_pp_signal(
450 receiver: &std::sync::mpsc::Receiver<PrimePpSignal>,
451 expected: PrimePpWaveSlot,
452 exact_slot: bool,
453 label: &str,
454) -> Result<PrimePpWaveSlot, String> {
455 match receiver.recv() {
456 Ok(PrimePpSignal::Error(error)) => Err(error),
457 Ok(PrimePpSignal::Slot(received))
458 if received.wave == expected.wave
459 && (!exact_slot || received.slot == expected.slot) =>
460 {
461 if received.slot >= 2 {
462 Err(format!(
463 "{label}: wave {} carried invalid slot {}",
464 received.wave, received.slot
465 ))
466 } else {
467 Ok(received)
468 }
469 }
470 Ok(PrimePpSignal::Slot(received)) => Err(format!(
471 "{label}: expected wave/slot {:?}, received {:?}",
472 expected, received
473 )),
474 Err(_) => Err(format!(
475 "{label}: channel closed while waiting for wave {}",
476 expected.wave
477 )),
478 }
479}
480
481fn send_prime_pp_signal(
482 sender: &std::sync::mpsc::Sender<PrimePpSignal>,
483 signal: PrimePpSignal,
484 label: &str,
485) -> Result<(), String> {
486 sender
487 .send(signal)
488 .map_err(|_| format!("{label}: channel closed"))
489}
490
491struct PrimePpWave<'a> {
492 start: usize,
493 end: usize,
494 tokens: &'a [u32],
495}
496
497struct PrimePpStageChannels {
498 incoming: Option<std::sync::mpsc::Receiver<PrimePpSignal>>,
499 release_upstream: Option<std::sync::mpsc::Sender<PrimePpSignal>>,
500 outgoing: std::sync::mpsc::Sender<PrimePpSignal>,
501 released_downstream: std::sync::mpsc::Receiver<PrimePpSignal>,
502}
503
504impl PrimePpStageChannels {
505 fn notify_failure(&self, error: &str) {
506 if let Some(upstream) = &self.release_upstream {
507 let _ = upstream.send(PrimePpSignal::Error(error.to_string()));
508 }
509 let _ = self.outgoing.send(PrimePpSignal::Error(error.to_string()));
510 }
511}
512
513pub struct IndexerPlanes<'a> {
527 pub state: &'a mut CudaSlice<f32>,
528 pub pool_keys: &'a mut Option<CudaSlice<f32>>,
529 pub ready: &'a mut usize,
530 pub state_ring_rows: usize,
535 pub capacity_tokens: usize,
539}
540
541pub(crate) struct AttnPre {
543 pub q: cudarc::driver::CudaSlice<f32>,
544 pub k: cudarc::driver::CudaSlice<f32>,
545 pub v: cudarc::driver::CudaSlice<f32>,
546 pub gate: Option<cudarc::driver::CudaSlice<f32>>,
547}
548
549pub(crate) struct GdnPrep {
551 pub hk: usize,
552 pub q_l2: cudarc::driver::CudaSlice<f32>,
553 pub k_l2: cudarc::driver::CudaSlice<f32>,
554 pub v_g: cudarc::driver::CudaSlice<f32>,
555 pub beta: cudarc::driver::CudaSlice<f32>,
556 pub g_log: cudarc::driver::CudaSlice<f32>,
557 pub kb16: Option<cudarc::driver::CudaSlice<u8>>,
558 pub qb16: Option<cudarc::driver::CudaSlice<u8>>,
559}
560
561pub(crate) struct VerifyStreamScratch {
563 pub pos_d: CudaSlice<i32>,
564 pub row_ctrs: Vec<CudaSlice<i32>>,
565}
566use crate::hybrid::{FullAttnLayer, HybridModel, LinearAttnLayer, Mixer, MoeWeights};
567
568struct MoeInputTraceWriter {
569 dir: std::path::PathBuf,
570 index: std::fs::File,
571 payloads: std::collections::HashMap<u16, (std::fs::File, u64)>,
572}
573
574static MOE_INPUT_TRACE_WRITER: std::sync::OnceLock<std::sync::Mutex<Option<MoeInputTraceWriter>>> =
575 std::sync::OnceLock::new();
576
577fn gdec_enabled() -> bool {
580 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
581 *E.get_or_init(|| {
582 std::env::var("MEMRA_MOE_GDEC")
583 .map(|v| v != "0")
584 .unwrap_or(true)
585 })
586}
587
588fn moe_slab_enabled() -> bool {
599 std::env::var("MEMRA_MOE_SLAB").as_deref() != Ok("0")
600}
601
602fn moe_fused_epi_enabled() -> bool {
615 std::env::var("MEMRA_MOE_FUSED_EPI")
616 .map(|v| v != "0")
617 .unwrap_or(false)
618}
619
620fn hyper_decode_ws_on() -> bool {
634 std::env::var("MEMRA_HC_DECODE_WS").as_deref() == Ok("1")
635}
636
637pub static HC_DECODE_WS_DISPATCHES: std::sync::atomic::AtomicU64 =
640 std::sync::atomic::AtomicU64::new(0);
641
642fn mla_tc_prefill_enabled() -> bool {
665 std::env::var("MEMRA_MLA_TC_PREFILL")
666 .map(|v| v != "0")
667 .unwrap_or(true)
668}
669
670fn moe_grouped_enabled(_cfg: &ModelConfig, _prefill: bool) -> bool {
674 std::env::var("MEMRA_MOE_GROUPED")
675 .map(|value| value != "0")
676 .unwrap_or(false)
677}
678
679fn moe_grouped_prefill_enabled() -> bool {
700 std::env::var("MEMRA_MOE_GROUPED_PREFILL")
701 .map(|v| v != "0")
702 .unwrap_or(true)
703}
704
705fn moe_prefetch_enabled() -> bool {
708 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
709 *E.get_or_init(|| {
710 std::env::var("MEMRA_MOE_PREFETCH").as_deref() == Ok("1")
711 || crate::spill_pread::worker_enabled()
712 })
713}
714
715fn moe_page_prefetch_window() -> usize {
720 static W: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
721 *W.get_or_init(|| {
722 page_prefetch_window_from_values(
723 std::env::var("MEMRA_MOE_PAGE_PREFETCH").as_deref() == Ok("1"),
724 std::env::var("MEMRA_MOE_PAGE_PREFETCH_WINDOW")
725 .ok()
726 .as_deref(),
727 )
728 })
729}
730
731fn page_prefetch_window_from_values(enabled: bool, raw_window: Option<&str>) -> usize {
732 if !enabled {
733 return 0;
734 }
735 raw_window.and_then(|value| value.parse().ok()).unwrap_or(1)
736}
737
738fn page_prefetch_positions(position: usize, len: usize, window: usize) -> std::ops::Range<usize> {
742 if window == 0 || position >= len {
743 return len..len;
744 }
745 let (start, count) = if position == 0 {
746 (1, window)
747 } else {
748 (position.saturating_add(window), 1)
749 };
750 let start = start.min(len);
751 start..start.saturating_add(count).min(len)
752}
753
754fn grouped_worker_prefetch_position(order_len: usize, current: Option<usize>) -> Option<usize> {
757 let position = current.map_or(0, |position| position.saturating_add(1));
758 (position < order_len).then_some(position)
759}
760
761fn worker_prefetch_window() -> usize {
766 static WINDOW: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
767 *WINDOW.get_or_init(|| {
768 let automatic = crate::spill_pread::configured_depth().saturating_sub(1) / 3;
769 std::env::var("MEMRA_SPILL_WORKER_EXPERT_WINDOW")
770 .ok()
771 .and_then(|value| value.parse::<usize>().ok())
772 .unwrap_or(automatic.max(1))
773 })
774}
775
776fn worker_prefetch_positions(position: usize, len: usize, window: usize) -> std::ops::Range<usize> {
780 if window == 0 || position >= len {
781 return len..len;
782 }
783 let (start, count) = if position == 0 {
784 (0, window)
785 } else {
786 (position.saturating_add(window).saturating_sub(1), 1)
787 };
788 let start = start.min(len);
789 start..start.saturating_add(count).min(len)
790}
791
792fn moe_dev_enabled() -> bool {
797 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
798 *E.get_or_init(|| {
799 std::env::var("MEMRA_MOE_DEV")
800 .map(|v| v != "0")
801 .unwrap_or(true)
802 && !matches!(std::env::var("MEMRA_FUSED_ROUTER").as_deref(), Ok("0"))
803 })
804}
805
806enum VrowsSel<'a> {
815 Host(&'a [u32], &'a [f32]),
816 Dev(&'a CudaSlice<i32>, &'a CudaSlice<f32>),
817}
818
819fn sigmoid_router_enabled() -> bool {
820 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
821 *E.get_or_init(|| {
822 std::env::var("MEMRA_SIG_ROUTER")
823 .map(|v| v != "0")
824 .unwrap_or(true)
825 })
826}
827
828fn moe_q8_enabled() -> bool {
833 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
834 *E.get_or_init(|| {
835 std::env::var("MEMRA_MOE_Q8")
836 .map(|v| v != "0")
837 .unwrap_or(true)
838 })
839}
840
841fn expert_dp4a_supported(qt: i32) -> bool {
844 qt == crate::QT_Q4_0
845 || qt == crate::QT_IQ3_S
846 || qt == crate::QT_IQ4_XS
847 || qt == crate::QT_Q3_K
848 || qt == crate::QT_Q4_K
849 || qt == crate::QT_Q6_K
850}
851
852fn q8_expert_supported(qt: i32) -> bool {
853 static KQ: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
859 let kq = *KQ.get_or_init(|| {
860 std::env::var("MEMRA_MOE_Q8_KQ")
861 .map(|v| v != "0")
862 .unwrap_or(true)
863 });
864 let nvfp4_q8 = std::env::var("MEMRA_MOE_Q8_NVFP4")
871 .map(|v| v != "0")
872 .unwrap_or(true);
873 qt == crate::QT_IQ3_S
874 || qt == crate::QT_IQ4_XS
875 || (nvfp4_q8 && qt == crate::QT_NVFP4)
876 || (kq && (qt == crate::QT_Q3_K || qt == crate::QT_Q4_K || qt == crate::QT_Q6_K))
877}
878
879fn q8_expert_supported_for_model(cfg: &ModelConfig, qt: i32) -> bool {
883 let weight_only_nvfp4 = cfg.hy3.as_ref().is_some_and(|hy3| hy3.weight_only_nvfp4);
884 q8_expert_supported(qt) && !(weight_only_nvfp4 && qt == crate::QT_NVFP4)
885}
886
887fn moe_q8_enabled_for_model(cfg: &ModelConfig, m: &MoeWeights) -> bool {
888 m.has_uniform_expert_layout()
889 && moe_q8_enabled()
890 && q8_expert_supported_for_model(cfg, m.gate_exps.qtype)
891 && q8_expert_supported_for_model(cfg, m.up_exps.qtype)
892 && q8_expert_supported_for_model(cfg, m.down_exps.qtype)
893}
894
895#[cfg(test)]
896mod w4a16_dispatch_tests {
897 use super::q8_expert_supported_for_model;
898 use memra_gguf::config::{HfConfig, ModelConfig};
899
900 #[test]
901 fn hy3_w4a16_never_admits_q8_activations() {
902 let hf = HfConfig::parse(
903 r#"{"model_type":"hy_v3","num_hidden_layers":2,"hidden_size":8,
904 "num_attention_heads":2,"intermediate_size":16,"vocab_size":32,
905 "max_position_embeddings":32,
906 "quantization_config":{"quant_method":"modelopt","quant_algo":"W4A16_NVFP4"}}"#,
907 );
908 let cfg = ModelConfig::from_hf(&hf);
909 assert!(!q8_expert_supported_for_model(&cfg, crate::QT_NVFP4));
910 assert!(q8_expert_supported_for_model(&cfg, crate::QT_IQ4_XS));
911 }
912}
913
914fn q8_expert_dec_supported(qt: i32) -> bool {
917 qt == crate::QT_IQ3_S || qt == crate::QT_IQ4_XS || qt == crate::QT_Q4_0
918}
919
920fn f16g_proj_ok(qt: i32, in_f: usize) -> bool {
926 match qt {
927 crate::QT_Q4_0 => in_f.is_multiple_of(32),
928 crate::QT_IQ4_XS | crate::QT_IQ3_S | crate::QT_Q3_K | crate::QT_Q4_K | crate::QT_Q6_K => {
929 in_f.is_multiple_of(256)
930 }
931 crate::QT_NVFP4 => in_f.is_multiple_of(64),
936 _ => false,
937 }
938}
939
940fn moe_prewarm_enabled() -> bool {
943 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
944 *E.get_or_init(|| {
945 std::env::var("MEMRA_MOE_PREWARM")
946 .map(|v| v != "0")
947 .unwrap_or(true)
948 })
949}
950
951fn cpu_expert_profile_admit_enabled() -> bool {
955 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
956 *E.get_or_init(|| std::env::var("MEMRA_CPU_EXPERT_FREEZE_PROFILE_ADMIT").as_deref() == Ok("1"))
957}
958
959pub const PRIME_MIN_T: usize = 16;
963
964pub const CUDA_GRID_YZ_MAX: usize = 65_535;
977
978pub const PRIME_CHUNK_LAUNCH_CAP: usize = CUDA_GRID_YZ_MAX - (PRIME_MIN_T - 1);
986
987fn explicit_prime_chunk(parsed: usize, ring_on: bool) -> usize {
994 if ring_on {
995 if parsed == 0 {
996 crate::cache::PRIME_CHUNK_MAX_TOKENS
997 } else {
998 parsed.min(crate::cache::PRIME_CHUNK_MAX_TOKENS)
999 }
1000 } else if parsed == 0 {
1001 PRIME_CHUNK_LAUNCH_CAP
1002 } else {
1003 parsed.min(PRIME_CHUNK_LAUNCH_CAP)
1004 }
1005}
1006
1007fn step_gemm_prime_suffix_on() -> bool {
1037 std::env::var("MEMRA_STEP_GEMM_PRIME_SUFFIX").as_deref() != Ok("0")
1038}
1039
1040const MOE_DEV_MAX_T: usize = 16;
1049const PRIME_PIPE_MICROBATCHES: usize = 8;
1050const PRIME_PIPE_MIN_CHUNK: usize = 128;
1051const PRIME_PIPE_EDGE_MIN_CHUNK: usize = 64;
1052const PRIME_PIPE_LINEAR_WORK: usize = 8;
1053
1054fn prime_pp2_auto_geometry(n_layers: usize) -> bool {
1055 crate::pp::prime_pp_on()
1056 && !crate::pp::pp2_streams_off()
1057 && crate::pp::pp_cuts(n_layers).is_some_and(|cuts| cuts.len() == 3)
1058}
1059
1060fn prime_ppn_wave_auto_geometry(n_layers: usize) -> bool {
1061 crate::pp::prime_pp_on()
1062 && !crate::pp::pp2_streams_off()
1063 && crate::pp::pp_wave_on() == Ok(true)
1064 && crate::pp::pp_cuts(n_layers).is_some_and(|cuts| matches!(cuts.len(), 4 | 5))
1065}
1066
1067fn prime_pipeline_auto_geometry(n_layers: usize) -> bool {
1068 prime_pp2_auto_geometry(n_layers) || prime_ppn_wave_auto_geometry(n_layers)
1069}
1070
1071pub fn prime_chunk_tokens(t: usize, n_layers: usize) -> usize {
1078 if let Ok(value) = std::env::var("MEMRA_PRIME_CHUNK") {
1079 let parsed = value
1080 .parse::<usize>()
1081 .unwrap_or(crate::cache::PRIME_CHUNK_MAX_TOKENS);
1082 return explicit_prime_chunk(parsed, crate::cache::swa_ring_on());
1083 }
1084 let chunk = crate::cache::PRIME_CHUNK_MAX_TOKENS;
1085 if prime_pipeline_auto_geometry(n_layers) && t >= 2 * PRIME_PIPE_MIN_CHUNK {
1086 chunk.min(
1087 t.div_ceil(PRIME_PIPE_MICROBATCHES)
1088 .max(PRIME_PIPE_MIN_CHUNK),
1089 )
1090 } else {
1091 chunk
1092 }
1093}
1094
1095fn fixed_prime_chunk_ranges(t: usize, chunk: usize) -> Vec<(usize, usize)> {
1096 fixed_prime_chunk_ranges_for_ring(t, chunk, crate::cache::swa_ring_on())
1097}
1098
1099fn fixed_prime_chunk_ranges_for_ring(t: usize, chunk: usize, ring_on: bool) -> Vec<(usize, usize)> {
1100 if chunk == 0 || t <= chunk {
1101 return vec![(0, t)];
1102 }
1103 let mut ranges = Vec::with_capacity(t.div_ceil(chunk));
1104 let mut start = 0usize;
1105 while start < t {
1106 let mut end = (start + chunk).min(t);
1107 if t - end > 0 && t - end < PRIME_MIN_T {
1108 if ring_on {
1109 let shifted = t - PRIME_MIN_T;
1110 end = if shifted > start { shifted } else { t };
1111 } else {
1112 end = t;
1113 }
1114 }
1115 ranges.push((start, end));
1116 start = end;
1117 }
1118 ranges
1119}
1120
1121fn prime_chunk_work(prefix: usize, total: usize) -> u128 {
1122 let prefix = prefix as u128;
1123 prefix * (prefix + (PRIME_PIPE_LINEAR_WORK as u128) * (total as u128))
1124}
1125
1126fn dynamic_prime_chunk_ranges(
1127 t: usize,
1128 fixed_chunk: usize,
1129 fixed: &[(usize, usize)],
1130) -> Vec<(usize, usize)> {
1131 let n = fixed.len();
1132 if n < 3 {
1133 return fixed.to_vec();
1134 }
1135
1136 let max_first = t - (n - 1) * PRIME_MIN_T;
1137 let first = fixed_chunk
1138 .div_ceil(2)
1139 .max(PRIME_PIPE_EDGE_MIN_CHUNK)
1140 .min(max_first);
1141 let mut ranges = Vec::with_capacity(n);
1142 ranges.push((0, first));
1143
1144 let first_work = prime_chunk_work(first, t);
1145 let work_span = prime_chunk_work(t, t) - first_work;
1146 let denominator = (n - 1) as u128;
1147 let mut previous = first;
1148 for boundary in 1..n - 1 {
1149 let target = first_work * denominator + work_span * (boundary as u128);
1150 let remaining = n - 1 - boundary;
1151 let mut low = previous + PRIME_MIN_T;
1152 let mut high = t - remaining * PRIME_MIN_T;
1153 while low < high {
1154 let mid = low + (high - low) / 2;
1155 if prime_chunk_work(mid, t) * denominator >= target {
1156 high = mid;
1157 } else {
1158 low = mid + 1;
1159 }
1160 }
1161 ranges.push((previous, low));
1162 previous = low;
1163 }
1164 ranges.push((previous, t));
1165 ranges
1166}
1167
1168pub fn prime_chunk_ranges(t: usize, n_layers: usize, gdn_grid: bool) -> Vec<(usize, usize)> {
1178 let explicit_chunk = std::env::var_os("MEMRA_PRIME_CHUNK").is_some();
1179 let chunk = prime_chunk_tokens(t, n_layers);
1180 let fixed = fixed_prime_chunk_ranges(t, chunk);
1181 let dynamic = match std::env::var("MEMRA_PRIME_CHUNK_SCHED") {
1182 Ok(value) => value == "dynamic",
1183 Err(_) => true,
1184 };
1185 if explicit_chunk {
1186 return fixed;
1187 }
1188 let ranges = if !dynamic || !prime_pipeline_auto_geometry(n_layers) {
1189 fixed
1190 } else {
1191 dynamic_prime_chunk_ranges(t, chunk, &fixed)
1192 };
1193 if gdn_grid && std::env::var("MEMRA_PRIME_GRID_ALIGN").as_deref() != Ok("0") {
1197 align_prime_ranges_to_gdn(&ranges, t, Engine::gdn_chunk_size())
1198 } else {
1199 ranges
1200 }
1201}
1202
1203pub fn hyper_prime_ranges(t: usize, n_layers: usize, gdn_grid: bool) -> Vec<(usize, usize)> {
1259 prime_chunk_ranges(t, n_layers, gdn_grid)
1260}
1261
1262pub fn hyper_prime_call_rows(t: usize, n_layers: usize, gdn_grid: bool) -> usize {
1267 hyper_prime_ranges(t, n_layers, gdn_grid)
1268 .iter()
1269 .map(|&(start, end)| end - start)
1270 .max()
1271 .unwrap_or(0)
1272}
1273
1274#[derive(Debug, Clone, Copy)]
1286pub struct HyperPrimeWorkspaceShape {
1287 pub chunk_token_bytes: usize,
1297 pub prompt_bytes_per_token: usize,
1301 pub kpool_score_pool: usize,
1306 pub n_layers: usize,
1309 pub gdn_grid: bool,
1311}
1312
1313impl HyperPrimeWorkspaceShape {
1314 pub fn admission_bytes(&self, prompt_rows: usize) -> usize {
1325 let rows = hyper_prime_call_rows(prompt_rows, self.n_layers, self.gdn_grid);
1326 let chunk = self.chunk_token_bytes.saturating_mul(rows);
1327 let score = prompt_rows
1328 .checked_div(self.kpool_score_pool)
1329 .map(|pools| rows.saturating_mul(pools).saturating_mul(size_of::<f32>()))
1330 .unwrap_or(0);
1331 chunk
1332 .saturating_add(score)
1333 .saturating_add(self.prompt_bytes_per_token.saturating_mul(prompt_rows))
1334 }
1335}
1336
1337pub fn align_prime_ranges_to_gdn(
1357 ranges: &[(usize, usize)],
1358 t: usize,
1359 c: usize,
1360) -> Vec<(usize, usize)> {
1361 if c == 0 || ranges.len() < 2 {
1362 return ranges.to_vec();
1363 }
1364 let mut out: Vec<(usize, usize)> = Vec::with_capacity(ranges.len());
1365 let mut start = 0usize;
1366 for (i, &(_, end)) in ranges.iter().enumerate() {
1367 let e = if i + 1 == ranges.len() {
1368 t
1369 } else {
1370 end / c * c
1371 };
1372 if e > start {
1373 out.push((start, e));
1374 start = e;
1375 } }
1377 debug_assert_eq!(out.last().map(|&(_, e)| e), Some(t));
1378 out
1379}
1380
1381struct HeadSplit {
1382 pin: u64,
1383 w1: CudaSlice<u8>,
1384 hn1: CudaSlice<f32>,
1385 y1: CudaSlice<f32>,
1386 logits_e: CudaSlice<f32>,
1387 ev_hn: cudarc::driver::CudaEvent,
1388 ev_done: cudarc::driver::CudaEvent,
1389 raw_hn1: u64,
1390 raw_y1: u64,
1391 raw_logits_hi: u64,
1392 samp: Option<SampScratch>,
1397}
1398
1399struct SampScratch {
1400 pb: CudaSlice<f32>,
1401 th: CudaSlice<f32>,
1402 z: CudaSlice<f32>,
1403 mx: CudaSlice<f32>,
1404 rows: CudaSlice<i32>,
1405}
1406static HEAD_SPLIT_WS: std::sync::Mutex<Option<HeadSplit>> = std::sync::Mutex::new(None);
1408
1409#[allow(clippy::type_complexity)]
1413static DEV1_ROUTER_REPS: std::sync::Mutex<
1414 Option<(
1415 std::collections::HashMap<u16, (CudaSlice<f32>, CudaSlice<f32>, CudaSlice<u8>)>,
1416 Option<CudaSlice<f32>>,
1417 )>,
1418> = std::sync::Mutex::new(None);
1419
1420#[allow(clippy::type_complexity)]
1424static SHEXP_D1_REPS: std::sync::Mutex<
1425 Option<std::collections::HashMap<u16, (CudaSlice<u8>, CudaSlice<u8>, CudaSlice<u8>)>>,
1426> = std::sync::Mutex::new(None);
1427#[allow(clippy::type_complexity)]
1428static SHEXP_D1_WS: std::sync::Mutex<
1429 Option<(
1430 (usize, usize),
1431 CudaSlice<f32>,
1432 CudaSlice<f32>,
1433 CudaSlice<f32>,
1434 cudarc::driver::CudaEvent,
1435 cudarc::driver::CudaEvent,
1436 )>,
1437> = std::sync::Mutex::new(None);
1438
1439#[allow(clippy::type_complexity)] static SHEXP_OV_WS: std::sync::Mutex<
1442 Option<(usize, usize, usize, CudaSlice<f32>, CudaSlice<f32>)>,
1443> = std::sync::Mutex::new(None);
1444
1445impl HybridModel {
1446 pub fn gdn_prime_grid_on(&self) -> bool {
1452 Engine::gdn_chunked_enabled()
1453 && self
1454 .layers
1455 .iter()
1456 .any(|l| matches!(l.mixer, crate::hybrid::Mixer::Linear(_)))
1457 }
1458
1459 fn full_attn_tp_device_resident(e: &Engine, tp: &crate::hybrid::StepTpQkv) -> bool {
1464 tp.runtime.native_p2p() && tp.runtime.root_shares_ctx(e)
1465 }
1466
1467 pub(crate) fn full_attn_tp_qkv(
1468 &self,
1469 e: &Engine,
1470 fa: &FullAttnLayer,
1471 h: &CudaSlice<f32>,
1472 t: usize,
1473 ) -> Result<Option<Vec<CudaSlice<f32>>>, Box<dyn std::error::Error>> {
1474 let Some(tp) = fa.step_tp_qkv.as_ref() else {
1475 return Ok(None);
1476 };
1477 let values = active_matrix_values(
1478 h.len(),
1479 t,
1480 self.cfg.n_embd as usize,
1481 "Step TP QKV activation",
1482 )?;
1483 if Self::full_attn_tp_device_resident(e, tp) {
1492 e.stream().synchronize()?;
1495 let q = tp
1496 .runtime
1497 .bf16_column_parallel_resident_native_device(&tp.q, h, t)?;
1498 let k = tp
1499 .runtime
1500 .bf16_column_parallel_resident_native_device(&tp.k, h, t)?;
1501 let v = tp
1502 .runtime
1503 .bf16_column_parallel_resident_native_device(&tp.v, h, t)?;
1504 Self::full_attn_tp_log_once(tp, "qkv", "device-resident");
1505 return Ok(Some(vec![q, k, v]));
1506 }
1507 let host = e.dtoh_view(&h.slice(0..values))?;
1508 let q = if tp.runtime.native_p2p() {
1509 tp.runtime
1510 .bf16_column_parallel_resident_native(&tp.q, &host, t)?
1511 } else {
1512 tp.runtime
1513 .bf16_column_parallel_resident(&tp.q, &host, t)?
1514 .gathered
1515 };
1516 let k = if tp.runtime.native_p2p() {
1517 tp.runtime
1518 .bf16_column_parallel_resident_native(&tp.k, &host, t)?
1519 } else {
1520 tp.runtime
1521 .bf16_column_parallel_resident(&tp.k, &host, t)?
1522 .gathered
1523 };
1524 let v = if tp.runtime.native_p2p() {
1525 tp.runtime
1526 .bf16_column_parallel_resident_native(&tp.v, &host, t)?
1527 } else {
1528 tp.runtime
1529 .bf16_column_parallel_resident(&tp.v, &host, t)?
1530 .gathered
1531 };
1532 Self::full_attn_tp_log_once(tp, "qkv", "host-canonical");
1533 Ok(Some(vec![e.htod(&q)?, e.htod(&k)?, e.htod(&v)?]))
1534 }
1535
1536 fn full_attn_tp_log_once(tp: &crate::hybrid::StepTpQkv, proj: &str, activation: &'static str) {
1540 use std::sync::atomic::{AtomicBool, Ordering};
1541 static LOGGED: [AtomicBool; 4] = [
1542 AtomicBool::new(false),
1543 AtomicBool::new(false),
1544 AtomicBool::new(false),
1545 AtomicBool::new(false),
1546 ];
1547 let idx = 2 * usize::from(proj == "o") + usize::from(activation == "device-resident");
1548 if LOGGED[idx].swap(true, Ordering::Relaxed) {
1549 return;
1550 }
1551 eprintln!(
1552 "[step-tp-{proj}] execute layer={} devices={:?} projections={proj} \
1553 tensor_parallel=true attention_local=true kv_local=true transport={} \
1554 native_p2p={} bulk_p2p={} activation={activation} \
1555 output={} performance_claim=false (logged once per transport)",
1556 tp.layer,
1557 tp.devices,
1558 tp.runtime.transport_label(),
1559 tp.runtime.native_p2p(),
1560 tp.runtime.bulk_p2p(),
1561 if activation == "device-resident" {
1562 "root-resident"
1563 } else {
1564 "root-readback"
1565 },
1566 );
1567 }
1568
1569 pub(crate) fn full_attn_tp_o(
1570 &self,
1571 e: &Engine,
1572 fa: &FullAttnLayer,
1573 activation: &CudaSlice<f32>,
1574 tokens: usize,
1575 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
1576 let Some(tp) = fa.step_tp_qkv.as_ref() else {
1577 return Ok(None);
1578 };
1579 if Self::full_attn_tp_device_resident(e, tp) {
1583 e.stream().synchronize()?; let output = tp
1585 .runtime
1586 .step_bf16_row_parallel_resident_native_device(&tp.o, activation, tokens)?;
1587 Self::full_attn_tp_log_once(tp, "o", "device-resident");
1588 return Ok(Some(output));
1589 }
1590 let host = e.dtoh(activation)?;
1591 let output = if tp.runtime.native_p2p() {
1592 tp.runtime
1593 .step_bf16_row_parallel_resident_native(&tp.o, &host, tokens)?
1594 } else {
1595 tp.runtime
1596 .step_bf16_row_parallel_resident(&tp.o, &host, tokens)?
1597 };
1598 Self::full_attn_tp_log_once(tp, "o", "host-canonical");
1599 Ok(Some(e.htod(&output)?))
1600 }
1601
1602 fn full_attn_o(
1603 &self,
1604 e: &Engine,
1605 fa: &FullAttnLayer,
1606 activation: &CudaSlice<f32>,
1607 tokens: usize,
1608 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1609 match self.full_attn_tp_o(e, fa, activation, tokens)? {
1610 Some(output) => Ok(output),
1611 None => e.matmul(&fa.wo, activation, tokens),
1612 }
1613 }
1614
1615 fn prime_trace_path() -> Option<&'static str> {
1620 static P: std::sync::OnceLock<Option<String>> = std::sync::OnceLock::new();
1621 P.get_or_init(|| std::env::var("MEMRA_PRIME_TRACE").ok())
1622 .as_deref()
1623 }
1624
1625 fn prime_anatomy_on() -> bool {
1631 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1632 *E.get_or_init(|| std::env::var("MEMRA_PRIME_ANATOMY").as_deref() == Ok("1"))
1633 }
1634
1635 fn prime_anatomy_slots() -> &'static [std::sync::atomic::AtomicU64; 5] {
1636 static S: [std::sync::atomic::AtomicU64; 5] = [
1637 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), ];
1643 &S
1644 }
1645
1646 pub(crate) fn refuse_hyper(&self, path: &str) -> Result<(), Box<dyn std::error::Error>> {
1653 if let Some(topology) = self.hyper.as_ref() {
1654 return Err(format!(
1655 "{path} runs a serial residual, but this model's ModelPlan declares \
1656 ResidualTopology::HyperConnections{{ streams: {}, collapse: {:?} }}. Refusing: \
1657 that path would compute a different model. Converted paths: forward, \
1658 forward_last, prime_cache, decode_step, and the batched serving chain \
1659 decode_step_batch / _sampled / _lean / _masked.",
1660 topology.streams, topology.collapse
1661 )
1662 .into());
1663 }
1664 Ok(())
1665 }
1666
1667 fn hyper_ffn_branch(
1673 &self,
1674 e: &Engine,
1675 layer: &crate::hybrid::HybridLayer,
1676 z: &CudaSlice<f32>,
1677 t: usize,
1678 il: usize,
1679 prefill: bool,
1680 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1681 match &layer.ffn {
1682 crate::hybrid::Ffn::Dense {
1683 ffn_gate,
1684 ffn_up,
1685 ffn_down,
1686 } => {
1687 let n_ff = ffn_gate.out_features();
1688 let mut g2 = e.matmul_group(&[ffn_gate, ffn_up], z, t)?;
1689 let up = g2.pop().unwrap();
1690 let gate = g2.pop().unwrap();
1691 let mut act = e.uninit(t * n_ff)?;
1692 Self::ffn_act_lim(
1694 e,
1695 &self.cfg,
1696 &gate,
1697 &up,
1698 1.0,
1699 1.0,
1700 self.cfg.clamp_shexp_at(il as u32),
1701 &mut act,
1702 t * n_ff,
1703 )?;
1704 e.matmul(ffn_down, &act, t)
1705 }
1706 crate::hybrid::Ffn::Moe(m) => {
1707 if prefill {
1708 self.moe_ffn_il_prefill(e, m, z, t, il as u16)
1709 } else {
1710 self.moe_ffn_il_zq8(e, m, z, None, t, il as u16)
1711 }
1712 }
1713 }
1714 }
1715
1716 fn forward_hyper(
1721 &self,
1722 e: &Engine,
1723 tokens: &[u32],
1724 last_only: bool,
1725 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
1726 let topology = *self
1727 .hyper
1728 .as_ref()
1729 .ok_or("forward_hyper on a model with no HyperConnections topology")?;
1730 if let Some(fence) = crate::pp::pp_cuts(self.layers.len()) {
1736 if !self.rewrite_allowed(memra_gguf::execution_manifest::RewriteSurface::Pipeline) {
1737 return Err("pipeline rewrite is not qualified for this ModelPlan".into());
1738 }
1739 return self.forward_hyper_ppn(e, tokens, last_only, &topology, &fence);
1740 }
1741 let n_embd = self.cfg.n_embd as usize;
1742 let t = tokens.len();
1743 let eps = self.cfg.rms_eps;
1744 let pos: Vec<i32> = (0..t as i32).collect();
1745 let pos_d = e.htod_i32(&pos)?;
1746
1747 let embedded = self.embed(e, tokens)?;
1748 let mut x = crate::hyper::expand(e, &topology, &embedded, t, n_embd)?;
1749 let trace = memra_reference::hidden_trace::enabled();
1750 if trace {
1751 memra_reference::hidden_trace::emit_tokens(tokens);
1752 let streams = x.len() / (t * n_embd);
1753 memra_reference::hidden_trace::emit_last_row(
1754 "expand",
1755 -1,
1756 t,
1757 streams * n_embd,
1758 &e.dtoh(&x)?,
1759 );
1760 }
1761
1762 x = self.hyper_range_forward(e, &topology, x, 0, self.layers.len(), &pos_d, t, trace)?;
1763
1764 self.hyper_head_logits(e, &topology, &x, t, n_embd, eps, last_only)
1768 }
1769
1770 #[allow(clippy::type_complexity)] fn prime_cache_hyper(
1782 &self,
1783 e: &Engine,
1784 tokens: &[u32],
1785 cache: &mut Cache,
1786 queued_after: usize,
1787 overlay: Option<&crate::vision::EmbedOverlay>,
1788 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
1789 let topology = *self
1790 .hyper
1791 .as_ref()
1792 .ok_or("prime_cache_hyper on a model with no HyperConnections topology")?;
1793 if let Some(fence) = crate::pp::pp_cuts(self.layers.len()) {
1816 if !self.rewrite_allowed(memra_gguf::execution_manifest::RewriteSurface::Pipeline) {
1817 return Err("pipeline rewrite is not qualified for this ModelPlan".into());
1818 }
1819 let n_embd = self.cfg.n_embd as usize;
1828 let t = tokens.len();
1829 if cache.pos + t > cache.max_ctx {
1830 return Err("prime_cache: prompt exceeds cache max_ctx".into());
1831 }
1832 let ranges = hyper_prime_ranges(t, self.layers.len(), self.gdn_prime_grid_on());
1833 if ranges.len() == 1 {
1834 return self.prime_cache_hyper_ppn(
1835 e,
1836 tokens,
1837 cache,
1838 queued_after,
1839 &topology,
1840 &fence,
1841 overlay,
1842 );
1843 }
1844 let mut hiddens = e.uninit(t * n_embd)?;
1845 let mut last: Option<(Vec<f32>, CudaSlice<f32>)> = None;
1846 for &(start, end) in &ranges {
1847 let ov = overlay.and_then(|o| o.window(start, end - start));
1848 let (l, hs, x) = self.prime_cache_hyper_ppn(
1849 e,
1850 &tokens[start..end],
1851 cache,
1852 queued_after + (t - end),
1853 &topology,
1854 &fence,
1855 ov.as_ref(),
1856 )?;
1857 e.copy_into(&mut hiddens, start * n_embd, &x, (end - start) * n_embd)?;
1858 last = Some((l, hs));
1859 }
1860 let (logits, h_seed) =
1861 last.expect("hyper_prime_ranges never returns an empty schedule");
1862 return Ok((logits, h_seed, hiddens));
1863 }
1864 let n_embd = self.cfg.n_embd as usize;
1865 let t = tokens.len();
1866 if cache.pos + t > cache.max_ctx {
1867 return Err("prime_cache: prompt exceeds cache max_ctx".into());
1868 }
1869 let seq_end = cache.pos + t + queued_after;
1873 let ranges = hyper_prime_ranges(t, self.layers.len(), self.gdn_prime_grid_on());
1874 if ranges.len() == 1 {
1875 return self.prime_chunk_hyper(e, tokens, cache, seq_end, 0, overlay);
1876 }
1877 let mut hiddens = e.uninit(t * n_embd)?;
1878 let mut last: Option<(Vec<f32>, CudaSlice<f32>)> = None;
1879 for &(start, end) in &ranges {
1880 let (l, hs, x) =
1881 self.prime_chunk_hyper(e, &tokens[start..end], cache, seq_end, start, overlay)?;
1882 e.copy_into(&mut hiddens, start * n_embd, &x, (end - start) * n_embd)?;
1883 last = Some((l, hs));
1884 }
1885 let (logits, h_seed) = last.expect("hyper_prime_ranges never returns an empty schedule");
1886 Ok((logits, h_seed, hiddens))
1887 }
1888
1889 pub fn hyper_prime_workspace_shape(&self) -> Option<HyperPrimeWorkspaceShape> {
1914 let topology = self.hyper.as_ref()?;
1915 let h = self.cfg.n_embd as usize;
1916 let s = topology.streams;
1917 let f32b = std::mem::size_of::<f32>();
1918 let mut chunk_token_bytes = 4 * s * h * f32b + 4 * h * f32b + 2 * h * f32b;
1920 if let Some(moe) = self.cfg.moe.as_ref() {
1921 let u = moe.expert_used_count as usize;
1922 let f = moe.expert_ff_length as usize;
1923 chunk_token_bytes += u * (10 * h + 14 * f) + 4 * h;
1924 }
1925 let mut kpool_score_pool = 0;
1926 if let Some(glm5) = self.cfg.glm5.as_ref() {
1927 let heads = self.cfg.n_head as usize;
1928 chunk_token_bytes +=
1929 heads * (glm5.qk_head_dim as usize + glm5.v_head_dim as usize) * f32b;
1930 if glm5.index_kpool > 0 {
1931 chunk_token_bytes += (glm5.index_topk as usize / glm5.index_kpool as usize + 1)
1932 * std::mem::size_of::<i32>();
1933 kpool_score_pool = glm5.index_kpool as usize;
1934 }
1935 }
1936 Some(HyperPrimeWorkspaceShape {
1937 chunk_token_bytes,
1938 prompt_bytes_per_token: h * f32b,
1939 kpool_score_pool,
1940 n_layers: self.layers.len(),
1941 gdn_grid: self.gdn_prime_grid_on(),
1942 })
1943 }
1944
1945 pub(crate) fn decode_step_hyper(
1950 &self,
1951 e: &Engine,
1952 token: u32,
1953 cache: &mut Cache,
1954 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
1955 let topology = *self
1956 .hyper
1957 .as_ref()
1958 .ok_or("decode_step_hyper on a model with no HyperConnections topology")?;
1959 if let Some(fence) = crate::pp::pp_cuts(self.layers.len()) {
1963 if !self.rewrite_allowed(memra_gguf::execution_manifest::RewriteSurface::Pipeline) {
1964 return Err("pipeline rewrite is not qualified for this ModelPlan".into());
1965 }
1966 return self.decode_step_hyper_ppn(e, token, cache, &topology, &fence);
1967 }
1968 let n_embd = self.cfg.n_embd as usize;
1969 let eps = self.cfg.rms_eps;
1970 let pos = cache.pos;
1971 let pos_d = e.htod_i32(&[pos as i32])?;
1972
1973 let embedded = e.htod(&self.embd.gather(n_embd, &[token]))?;
1974 let mut x = crate::hyper::expand(e, &topology, &embedded, 1, n_embd)?;
1975
1976 x = self.hyper_range_decode(e, &topology, x, 0, self.layers.len(), &pos_d, pos, cache)?;
1977
1978 self.hyper_decode_tail(e, &topology, &x, n_embd, eps, cache)
1980 }
1981
1982 #[allow(clippy::too_many_arguments)]
1989 fn hyper_range_forward(
1990 &self,
1991 e: &Engine,
1992 topology: &crate::hyper::HyperTopology,
1993 mut x: CudaSlice<f32>,
1994 lo: usize,
1995 hi: usize,
1996 pos_d: &CudaSlice<i32>,
1997 t: usize,
1998 trace: bool,
1999 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2000 let n_embd = self.cfg.n_embd as usize;
2001 let eps = self.cfg.rms_eps;
2002 for il in lo..hi {
2003 let layer = &self.layers[il];
2004 let hyper = layer.hyper.as_ref().ok_or_else(|| {
2005 format!("layer {il} carries no hyper-connection weights under an hc plan")
2006 })?;
2007
2008 let (y, mix) = crate::hyper::pre(e, topology, &hyper.attn, &x, t, n_embd)?;
2009 let mut h = e.uninit(t * n_embd)?;
2010 e.rms_norm(&y, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
2011 let mixed = match &layer.mixer {
2012 Mixer::Full(fa) => self.full_attn(e, fa, &h, pos_d, t, il)?,
2013 Mixer::Linear(la) => self.linear_attn(e, la, &h, t)?,
2014 Mixer::Mla(mla) => self.mla_attn(e, mla, &h, pos_d, t, il)?,
2015 Mixer::Kda(la) => crate::kda::kda_attn(e, la, &h, t, eps)?,
2016 };
2017 x = crate::hyper::post(e, topology, &mixed, &x, &mix, t, n_embd)?;
2018 if trace {
2019 let index = il as i64;
2020 memra_reference::hidden_trace::emit_last_row(
2021 "mixer",
2022 index,
2023 t,
2024 n_embd,
2025 &e.dtoh(&mixed)?,
2026 );
2027 let streams = x.len() / (t * n_embd);
2028 memra_reference::hidden_trace::emit_last_row(
2029 "attn",
2030 index,
2031 t,
2032 streams * n_embd,
2033 &e.dtoh(&x)?,
2034 );
2035 }
2036
2037 let (y, mix) = crate::hyper::pre(e, topology, &hyper.mlp, &x, t, n_embd)?;
2038 let mut z = e.uninit(t * n_embd)?;
2039 e.rms_norm(
2040 &y,
2041 layer.post_attn_norm.float_data(),
2042 &mut z,
2043 n_embd,
2044 t,
2045 eps,
2046 )?;
2047 let ffn_out = self.hyper_ffn_branch(e, layer, &z, t, il, true)?;
2048 x = crate::hyper::post(e, topology, &ffn_out, &x, &mix, t, n_embd)?;
2049 if trace {
2050 let index = il as i64;
2051 memra_reference::hidden_trace::emit_last_row(
2052 "ffn",
2053 index,
2054 t,
2055 n_embd,
2056 &e.dtoh(&ffn_out)?,
2057 );
2058 let streams = x.len() / (t * n_embd);
2059 memra_reference::hidden_trace::emit_last_row(
2060 "layer",
2061 index,
2062 t,
2063 streams * n_embd,
2064 &e.dtoh(&x)?,
2065 );
2066 }
2067 }
2068 Ok(x)
2069 }
2070
2071 #[allow(clippy::too_many_arguments)]
2076 fn hyper_range_prime(
2077 &self,
2078 e: &Engine,
2079 topology: &crate::hyper::HyperTopology,
2080 mut x: CudaSlice<f32>,
2081 lo: usize,
2082 hi: usize,
2083 pos_d: &CudaSlice<i32>,
2084 t: usize,
2085 cache: &mut Cache,
2086 seq_end: usize,
2087 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2088 let n_embd = self.cfg.n_embd as usize;
2089 let eps = self.cfg.rms_eps;
2090 for il in lo..hi {
2091 let layer = &self.layers[il];
2092 let hyper = layer.hyper.as_ref().ok_or_else(|| {
2093 format!("layer {il} carries no hyper-connection weights under an hc plan")
2094 })?;
2095
2096 let (y, mix) = crate::hyper::pre(e, topology, &hyper.attn, &x, t, n_embd)?;
2097 let mut h = e.uninit(t * n_embd)?;
2098 e.rms_norm(&y, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
2099 let mixed = match &layer.mixer {
2100 Mixer::Full(fa) => {
2101 self.full_attn_prime(e, fa, &h, None, pos_d, t, cache, il, seq_end)?
2102 }
2103 Mixer::Linear(la) => self.linear_attn_prime(e, la, &h, None, t, cache, il)?,
2104 Mixer::Mla(mla) if mla.tp.is_some() => {
2105 self.mla_tp_attn_cached(e, mla, &h, pos_d, t, il, cache, false)?
2106 }
2107 Mixer::Mla(mla) => self.mla_attn_cached(e, mla, &h, pos_d, t, il, cache)?,
2108 Mixer::Kda(la) if la.tp.is_some() => crate::glm5_tp::kda_tp_cached(
2109 e,
2110 la,
2111 &h,
2112 t,
2113 eps,
2114 cache,
2115 il,
2116 crate::kda::ConvArm::Prefill,
2117 )?,
2118 Mixer::Kda(la) => crate::kda::kda_prime_cached(e, la, &h, t, eps, cache, il)?,
2119 };
2120 x = crate::hyper::post(e, topology, &mixed, &x, &mix, t, n_embd)?;
2121
2122 let (y, mix) = crate::hyper::pre(e, topology, &hyper.mlp, &x, t, n_embd)?;
2123 let mut z = e.uninit(t * n_embd)?;
2124 e.rms_norm(
2125 &y,
2126 layer.post_attn_norm.float_data(),
2127 &mut z,
2128 n_embd,
2129 t,
2130 eps,
2131 )?;
2132 let ffn_out = self.hyper_ffn_branch(e, layer, &z, t, il, true)?;
2133 x = crate::hyper::post(e, topology, &ffn_out, &x, &mix, t, n_embd)?;
2134 self.glm5_hc_tap(e, cache, topology, il, &x, t)?;
2137 }
2138 Ok(x)
2139 }
2140
2141 #[allow(clippy::too_many_arguments)]
2143 fn hyper_range_decode(
2144 &self,
2145 e: &Engine,
2146 topology: &crate::hyper::HyperTopology,
2147 mut x: CudaSlice<f32>,
2148 lo: usize,
2149 hi: usize,
2150 pos_d: &CudaSlice<i32>,
2151 pos: usize,
2152 cache: &mut Cache,
2153 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2154 if hyper_decode_ws_on() {
2160 return self.hyper_range_decode_ws(e, topology, x, lo, hi, pos_d, pos, cache);
2161 }
2162 let n_embd = self.cfg.n_embd as usize;
2163 let eps = self.cfg.rms_eps;
2164 for il in lo..hi {
2165 let layer = &self.layers[il];
2166 let hyper = layer.hyper.as_ref().ok_or_else(|| {
2167 format!("layer {il} carries no hyper-connection weights under an hc plan")
2168 })?;
2169
2170 let (y, mix) = crate::hyper::pre(e, topology, &hyper.attn, &x, 1, n_embd)?;
2171 let mut h = e.uninit(n_embd)?;
2172 e.rms_norm(&y, layer.attn_norm.float_data(), &mut h, n_embd, 1, eps)?;
2173 let mixed = match &layer.mixer {
2174 Mixer::Full(fa) => self.full_attn_decode(e, fa, &h, pos_d, pos, cache, il)?,
2175 Mixer::Linear(la) => self.linear_attn_decode(e, la, &h, cache, il)?,
2176 Mixer::Mla(mla) if mla.tp.is_some() => {
2177 self.mla_tp_attn_cached(e, mla, &h, pos_d, 1, il, cache, false)?
2178 }
2179 Mixer::Mla(mla) => self.mla_attn_cached(e, mla, &h, pos_d, 1, il, cache)?,
2180 Mixer::Kda(la) if la.tp.is_some() => crate::glm5_tp::kda_tp_cached(
2181 e,
2182 la,
2183 &h,
2184 1,
2185 eps,
2186 cache,
2187 il,
2188 crate::kda::ConvArm::Decode,
2189 )?,
2190 Mixer::Kda(la) => crate::kda::kda_decode_cached(e, la, &h, eps, cache, il)?,
2191 };
2192 x = crate::hyper::post(e, topology, &mixed, &x, &mix, 1, n_embd)?;
2193
2194 let (y, mix) = crate::hyper::pre(e, topology, &hyper.mlp, &x, 1, n_embd)?;
2195 let mut z = e.uninit(n_embd)?;
2196 e.rms_norm(
2197 &y,
2198 layer.post_attn_norm.float_data(),
2199 &mut z,
2200 n_embd,
2201 1,
2202 eps,
2203 )?;
2204 let ffn_out = self.hyper_ffn_branch(e, layer, &z, 1, il, false)?;
2205 x = crate::hyper::post(e, topology, &ffn_out, &x, &mix, 1, n_embd)?;
2206 }
2207 Ok(x)
2208 }
2209
2210 #[allow(clippy::too_many_arguments)]
2217 fn hyper_range_decode_ws(
2218 &self,
2219 e: &Engine,
2220 topology: &crate::hyper::HyperTopology,
2221 x: CudaSlice<f32>,
2222 lo: usize,
2223 hi: usize,
2224 pos_d: &CudaSlice<i32>,
2225 pos: usize,
2226 cache: &mut Cache,
2227 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2228 let n_embd = self.cfg.n_embd as usize;
2229 let mut ws = match e.hyper_ws_take() {
2230 Some(ws) if ws.matches(topology, n_embd) => ws,
2231 _ => crate::hyper::HyperDecodeWs::new(e, topology, n_embd)?,
2232 };
2233 if HC_DECODE_WS_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed) == 0 {
2234 eprintln!(
2235 "[hc-decode-ws] engaged streams={} hidden={n_embd} (persistent hc-glue \
2236 workspace, per-engine pool; MEMRA_HC_DECODE_WS=1)",
2237 topology.streams
2238 );
2239 }
2240 let out =
2241 self.hyper_range_decode_ws_body(e, topology, x, lo, hi, pos_d, pos, cache, &mut ws);
2242 e.hyper_ws_put(ws);
2243 out
2244 }
2245
2246 #[allow(clippy::too_many_arguments)]
2251 fn hyper_range_decode_ws_body(
2252 &self,
2253 e: &Engine,
2254 topology: &crate::hyper::HyperTopology,
2255 mut x: CudaSlice<f32>,
2256 lo: usize,
2257 hi: usize,
2258 pos_d: &CudaSlice<i32>,
2259 pos: usize,
2260 cache: &mut Cache,
2261 ws: &mut crate::hyper::HyperDecodeWs,
2262 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2263 let n_embd = self.cfg.n_embd as usize;
2264 let eps = self.cfg.rms_eps;
2265 for il in lo..hi {
2266 let layer = &self.layers[il];
2267 let hyper = layer.hyper.as_ref().ok_or_else(|| {
2268 format!("layer {il} carries no hyper-connection weights under an hc plan")
2269 })?;
2270
2271 crate::hyper::pre_t1_ws(e, topology, &hyper.attn, &x, ws, n_embd)?;
2272 e.rms_norm(
2273 &ws.y,
2274 layer.attn_norm.float_data(),
2275 &mut ws.h,
2276 n_embd,
2277 1,
2278 eps,
2279 )?;
2280 let mixed = match &layer.mixer {
2281 Mixer::Full(fa) => self.full_attn_decode(e, fa, &ws.h, pos_d, pos, cache, il)?,
2282 Mixer::Linear(la) => self.linear_attn_decode(e, la, &ws.h, cache, il)?,
2283 Mixer::Mla(mla) => self.mla_attn_cached(e, mla, &ws.h, pos_d, 1, il, cache)?,
2284 Mixer::Kda(la) => crate::kda::kda_decode_cached(e, la, &ws.h, eps, cache, il)?,
2285 };
2286 crate::hyper::post_t1_ws(e, topology, &mixed, &x, ws, n_embd)?;
2287 std::mem::swap(&mut x, &mut ws.xb);
2288
2289 crate::hyper::pre_t1_ws(e, topology, &hyper.mlp, &x, ws, n_embd)?;
2290 e.rms_norm(
2291 &ws.y,
2292 layer.post_attn_norm.float_data(),
2293 &mut ws.z,
2294 n_embd,
2295 1,
2296 eps,
2297 )?;
2298 let ffn_out = self.hyper_ffn_branch(e, layer, &ws.z, 1, il, false)?;
2299 crate::hyper::post_t1_ws(e, topology, &ffn_out, &x, ws, n_embd)?;
2300 std::mem::swap(&mut x, &mut ws.xb);
2301 }
2302 Ok(x)
2303 }
2304
2305 #[allow(clippy::too_many_arguments)]
2340 pub(crate) fn hyper_batch_range_decode(
2341 &self,
2342 e: &Engine,
2343 topology: &crate::hyper::HyperTopology,
2344 mut x: CudaSlice<f32>,
2345 lo: usize,
2346 hi: usize,
2347 pos_rows: &[CudaSlice<i32>],
2348 caches: &mut [&mut Cache],
2349 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2350 let b_n = caches.len();
2351 assert_eq!(
2352 pos_rows.len(),
2353 b_n,
2354 "hyper_batch_range_decode: pos_rows built for a different batch width"
2355 );
2356 let n_embd = self.cfg.n_embd as usize;
2357 let eps = self.cfg.rms_eps;
2358 for il in lo..hi {
2359 let layer = &self.layers[il];
2360 let hyper = layer.hyper.as_ref().ok_or_else(|| {
2361 format!("layer {il} carries no hyper-connection weights under an hc plan")
2362 })?;
2363
2364 let (y, mix) = crate::hyper::pre_exact(e, topology, &hyper.attn, &x, b_n, n_embd)?;
2365 let mut h = e.uninit(b_n * n_embd)?;
2366 e.rms_norm(&y, layer.attn_norm.float_data(), &mut h, n_embd, b_n, eps)?;
2367 let mut mixed = e.uninit(b_n * n_embd)?;
2371 for bi in 0..b_n {
2372 let mut h_row = e.uninit(n_embd)?;
2373 e.dtod_copy_view(&h.slice(bi * n_embd..(bi + 1) * n_embd), &mut h_row)?;
2374 let cache: &mut Cache = &mut *caches[bi];
2375 let pos = cache.pos;
2376 let out_row = match &layer.mixer {
2377 Mixer::Full(fa) => {
2378 self.full_attn_decode(e, fa, &h_row, &pos_rows[bi], pos, cache, il)?
2379 }
2380 Mixer::Linear(la) => self.linear_attn_decode(e, la, &h_row, cache, il)?,
2381 Mixer::Mla(mla) => {
2382 self.mla_attn_cached(e, mla, &h_row, &pos_rows[bi], 1, il, cache)?
2383 }
2384 Mixer::Kda(la) => crate::kda::kda_decode_cached(e, la, &h_row, eps, cache, il)?,
2385 };
2386 e.copy_into(&mut mixed, bi * n_embd, &out_row, n_embd)?;
2387 }
2388 x = crate::hyper::post(e, topology, &mixed, &x, &mix, b_n, n_embd)?;
2389
2390 let (y, mix) = crate::hyper::pre_exact(e, topology, &hyper.mlp, &x, b_n, n_embd)?;
2391 let mut z = e.uninit(b_n * n_embd)?;
2392 e.rms_norm(
2393 &y,
2394 layer.post_attn_norm.float_data(),
2395 &mut z,
2396 n_embd,
2397 b_n,
2398 eps,
2399 )?;
2400 let ffn_out = self.hyper_ffn_branch_batch(e, layer, &z, b_n, il, false)?;
2401 x = crate::hyper::post(e, topology, &ffn_out, &x, &mix, b_n, n_embd)?;
2402 }
2403 Ok(x)
2404 }
2405
2406 pub(crate) fn hyper_ffn_branch_batch(
2426 &self,
2427 e: &Engine,
2428 layer: &crate::hybrid::HybridLayer,
2429 z: &CudaSlice<f32>,
2430 b_n: usize,
2431 il: usize,
2432 vrows: bool,
2433 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2434 let n_embd = self.cfg.n_embd as usize;
2435 match &layer.ffn {
2436 crate::hybrid::Ffn::Dense { .. } => {
2437 let mut out = e.uninit(b_n * n_embd)?;
2438 for bi in 0..b_n {
2439 let mut z_row = e.uninit(n_embd)?;
2440 e.dtod_copy_view(&z.slice(bi * n_embd..(bi + 1) * n_embd), &mut z_row)?;
2441 let row = self.hyper_ffn_branch(e, layer, &z_row, 1, il, false)?;
2442 e.copy_into(&mut out, bi * n_embd, &row, n_embd)?;
2443 }
2444 Ok(out)
2445 }
2446 crate::hybrid::Ffn::Moe(m) => {
2447 if vrows {
2448 self.moe_ffn_il_zq8_vrows(e, m, z, b_n, il as u16)
2449 } else {
2450 self.moe_ffn_il_zq8(e, m, z, None, b_n, il as u16)
2451 }
2452 }
2453 }
2454 }
2455
2456 fn forward_hyper_ppn(
2494 &self,
2495 e: &Engine,
2496 tokens: &[u32],
2497 last_only: bool,
2498 topology: &crate::hyper::HyperTopology,
2499 fence: &[usize],
2500 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
2501 let n_embd = self.cfg.n_embd as usize;
2502 let eps = self.cfg.rms_eps;
2503 let t = tokens.len();
2504 let width = topology.streams * n_embd;
2505 let trace = memra_reference::hidden_trace::enabled();
2506 if trace {
2507 memra_reference::hidden_trace::emit_tokens(tokens);
2508 }
2509 let pos: Vec<i32> = (0..t as i32).collect();
2510
2511 if crate::pp::pp2_streams_off() {
2512 let pos_d = e.htod_i32(&pos)?;
2516 let embedded = self.embed(e, tokens)?;
2517 let mut x = crate::hyper::expand(e, topology, &embedded, t, n_embd)?;
2518 x = self.hyper_range_forward(e, topology, x, fence[0], fence[1], &pos_d, t, trace)?;
2519 for s in 1..fence.len() - 1 {
2520 let boundary_tx = e.clone_dtod(&x)?;
2521 let boundary_rx = e.clone_dtod(&boundary_tx)?;
2522 x = self.hyper_range_forward(
2523 e,
2524 topology,
2525 boundary_rx,
2526 fence[s],
2527 fence[s + 1],
2528 &pos_d,
2529 t,
2530 trace,
2531 )?;
2532 }
2533 return self.hyper_head_logits(e, topology, &x, t, n_embd, eps, last_only);
2534 }
2535
2536 let rt = crate::pp::PpNRt::get(e)?;
2537 let n_st = fence.len() - 1;
2538 assert_eq!(
2539 rt.n_stages(),
2540 n_st,
2541 "PpNRt stage count {} != fence stages {n_st}",
2542 rt.n_stages()
2543 );
2544 rt.fence_stages_behind(&e.stream())?;
2547
2548 let mut slot = {
2549 let _st0 = rt.enter(0);
2550 let e0 = rt.engine(0, e);
2551 let pos_d = e0.htod_i32(&pos)?;
2554 let embedded = self.embed(e0, tokens)?;
2555 let x = crate::hyper::expand(e0, topology, &embedded, t, n_embd)?;
2556 let x =
2557 self.hyper_range_forward(e0, topology, x, fence[0], fence[1], &pos_d, t, trace)?;
2558 rt.tx(0, &x, t * width)?
2559 };
2560 for s in 1..n_st - 1 {
2561 let _st = rt.enter(s);
2562 let es = rt.engine(s, e);
2563 let pos_d = es.htod_i32(&pos)?;
2564 let x = rt.rx(s - 1, slot, t * width)?;
2565 let x = self.hyper_range_forward(
2566 es,
2567 topology,
2568 x,
2569 fence[s],
2570 fence[s + 1],
2571 &pos_d,
2572 t,
2573 trace,
2574 )?;
2575 slot = rt.tx(s, &x, t * width)?;
2576 }
2577 let _stl = rt.enter(n_st - 1);
2578 let el = rt.engine(n_st - 1, e);
2579 let pos_d = el.htod_i32(&pos)?;
2580 let x = rt.rx(n_st - 2, slot, t * width)?;
2581 let x = self.hyper_range_forward(
2582 el,
2583 topology,
2584 x,
2585 fence[n_st - 1],
2586 fence[n_st],
2587 &pos_d,
2588 t,
2589 trace,
2590 )?;
2591 self.hyper_head_logits(el, topology, &x, t, n_embd, eps, last_only)
2592 }
2593
2594 #[allow(clippy::too_many_arguments)]
2599 fn hyper_head_logits(
2600 &self,
2601 e: &Engine,
2602 topology: &crate::hyper::HyperTopology,
2603 x: &CudaSlice<f32>,
2604 t: usize,
2605 n_embd: usize,
2606 eps: f32,
2607 last_only: bool,
2608 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
2609 let collapsed =
2610 crate::hyper::collapse(e, topology, self.hyper_head.as_ref(), x, t, n_embd)?;
2611 if memra_reference::hidden_trace::enabled() {
2612 memra_reference::hidden_trace::emit_last_row(
2613 "collapse",
2614 -1,
2615 t,
2616 n_embd,
2617 &e.dtoh(&collapsed)?,
2618 );
2619 }
2620 let mut hn = e.uninit(t * n_embd)?;
2621 e.rms_norm(
2622 &collapsed,
2623 self.output_norm.float_data(),
2624 &mut hn,
2625 n_embd,
2626 t,
2627 eps,
2628 )?;
2629 let logits = if last_only {
2630 let last = e.view(&hn, t * n_embd);
2631 let last_row = last.slice((t - 1) * n_embd..t * n_embd);
2632 let mut hlast = e.uninit(n_embd)?;
2633 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
2634 e.matmul(&self.output, &hlast, 1)?
2635 } else {
2636 e.matmul(&self.output, &hn, t)?
2637 };
2638 e.dtoh(&logits)
2639 }
2640
2641 #[allow(clippy::type_complexity)] #[allow(clippy::too_many_arguments)]
2646 fn prime_cache_hyper_ppn(
2647 &self,
2648 e: &Engine,
2649 tokens: &[u32],
2650 cache: &mut Cache,
2651 queued_after: usize,
2652 topology: &crate::hyper::HyperTopology,
2653 fence: &[usize],
2654 overlay: Option<&crate::vision::EmbedOverlay>,
2655 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
2656 let n_embd = self.cfg.n_embd as usize;
2657 let eps = self.cfg.rms_eps;
2658 let t = tokens.len();
2659 let width = topology.streams * n_embd;
2660 if cache.pos + t > cache.max_ctx {
2661 return Err("prime_cache: prompt exceeds cache max_ctx".into());
2662 }
2663 let seq_end = cache.pos + t + queued_after;
2664 let pos: Vec<i32> = (cache.pos as i32..(cache.pos + t) as i32).collect();
2665 if let Some(sink) = cache.hc_taps.as_mut() {
2668 sink.base = cache.pos;
2669 }
2670
2671 if crate::pp::pp2_streams_off() {
2672 let pos_d = e.htod_i32(&pos)?;
2673 let mut embedded = self.embed(e, tokens)?;
2674 if let Some(ov) = overlay {
2675 ov.splice_into(e, &mut embedded, 0, t, n_embd)?;
2679 }
2680 let mut x = crate::hyper::expand(e, topology, &embedded, t, n_embd)?;
2681 x = self.hyper_range_prime(
2682 e, topology, x, fence[0], fence[1], &pos_d, t, cache, seq_end,
2683 )?;
2684 for s in 1..fence.len() - 1 {
2685 let boundary_tx = e.clone_dtod(&x)?;
2686 let boundary_rx = e.clone_dtod(&boundary_tx)?;
2687 x = self.hyper_range_prime(
2688 e,
2689 topology,
2690 boundary_rx,
2691 fence[s],
2692 fence[s + 1],
2693 &pos_d,
2694 t,
2695 cache,
2696 seq_end,
2697 )?;
2698 }
2699 return self.hyper_prime_tail(e, topology, &x, t, n_embd, eps, cache);
2700 }
2701
2702 {
2703 let rt = crate::pp::PpNRt::get(e)?;
2704 let n_st = fence.len() - 1;
2705 assert_eq!(
2706 rt.n_stages(),
2707 n_st,
2708 "PpNRt stage count {} != fence stages {n_st}",
2709 rt.n_stages()
2710 );
2711 let caller_stream = e.stream();
2712 rt.fence_stages_behind(&caller_stream)?;
2713 if let Some(ov) = overlay
2737 && !ov.resident_in(rt.engine(0, e))
2738 {
2739 return Err(format!(
2740 "vision embedding overlay rows are resident on dev{} but pp stage 0's \
2741 embedding intake runs on dev{}: the overlay must be published into the \
2742 intake engine's context (build it with EmbedOverlay::new_published; \
2743 MEMRA_VISION_OVERLAY_PUBLISH=0 pins the pre-publication program, whose \
2744 only vision-capable shape is MEMRA_PP_STREAMS=0)",
2745 ov.ctx().ordinal(),
2746 rt.engine(0, e).ctx().ordinal(),
2747 )
2748 .into());
2749 }
2750 let mut slot = {
2751 let _st0 = rt.enter(0);
2752 let e0 = rt.engine(0, e);
2753 let pos_d = e0.htod_i32(&pos)?;
2754 let mut embedded = self.embed(e0, tokens)?;
2755 if let Some(ov) = overlay {
2756 ov.splice_into(e0, &mut embedded, 0, t, n_embd)?;
2762 }
2763 let x = crate::hyper::expand(e0, topology, &embedded, t, n_embd)?;
2764 let x = self.hyper_range_prime(
2765 e0, topology, x, fence[0], fence[1], &pos_d, t, cache, seq_end,
2766 )?;
2767 rt.tx(0, &x, t * width)?
2768 };
2769 for s in 1..n_st - 1 {
2770 let _st = rt.enter(s);
2771 let es = rt.engine(s, e);
2772 let pos_d = es.htod_i32(&pos)?;
2773 let x = rt.rx(s - 1, slot, t * width)?;
2774 let x = self.hyper_range_prime(
2775 es,
2776 topology,
2777 x,
2778 fence[s],
2779 fence[s + 1],
2780 &pos_d,
2781 t,
2782 cache,
2783 seq_end,
2784 )?;
2785 slot = rt.tx(s, &x, t * width)?;
2786 }
2787 let out = {
2788 let _stl = rt.enter(n_st - 1);
2789 let el = rt.engine(n_st - 1, e);
2790 let pos_d = el.htod_i32(&pos)?;
2791 let x = rt.rx(n_st - 2, slot, t * width)?;
2792 let x = self.hyper_range_prime(
2793 el,
2794 topology,
2795 x,
2796 fence[n_st - 1],
2797 fence[n_st],
2798 &pos_d,
2799 t,
2800 cache,
2801 seq_end,
2802 )?;
2803 self.hyper_prime_tail(el, topology, &x, t, n_embd, eps, cache)?
2804 };
2805 rt.publish_all_to(&caller_stream)?;
2813 Ok(out)
2814 }
2815 }
2816
2817 #[allow(clippy::type_complexity)] fn prime_chunk_hyper(
2823 &self,
2824 e: &Engine,
2825 tokens: &[u32],
2826 cache: &mut Cache,
2827 seq_end: usize,
2828 chunk_off: usize,
2829 overlay: Option<&crate::vision::EmbedOverlay>,
2830 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
2831 let topology = *self
2832 .hyper
2833 .as_ref()
2834 .ok_or("prime_chunk_hyper on a model with no HyperConnections topology")?;
2835 let n_embd = self.cfg.n_embd as usize;
2836 let t = tokens.len();
2837 let eps = self.cfg.rms_eps;
2838 let pos: Vec<i32> = (cache.pos as i32..(cache.pos + t) as i32).collect();
2839 let pos_d = e.htod_i32(&pos)?;
2840 if let Some(sink) = cache.hc_taps.as_mut() {
2843 sink.base = cache.pos;
2844 }
2845
2846 let mut embedded = self.embed(e, tokens)?;
2847 if let Some(ov) = overlay {
2848 ov.splice_into(e, &mut embedded, chunk_off, t, n_embd)?;
2854 }
2855 let mut x = crate::hyper::expand(e, &topology, &embedded, t, n_embd)?;
2856
2857 for (il, layer) in self.layers.iter().enumerate() {
2858 let hyper = layer.hyper.as_ref().ok_or_else(|| {
2859 format!("layer {il} carries no hyper-connection weights under an hc plan")
2860 })?;
2861
2862 let (y, mix) = crate::hyper::pre(e, &topology, &hyper.attn, &x, t, n_embd)?;
2863 let mut h = e.uninit(t * n_embd)?;
2864 e.rms_norm(&y, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
2865 let mixed = match &layer.mixer {
2866 Mixer::Full(fa) => {
2867 self.full_attn_prime(e, fa, &h, None, &pos_d, t, cache, il, seq_end)?
2868 }
2869 Mixer::Linear(la) => self.linear_attn_prime(e, la, &h, None, t, cache, il)?,
2870 Mixer::Mla(mla) if mla.tp.is_some() => {
2871 self.mla_tp_attn_cached(e, mla, &h, &pos_d, t, il, cache, false)?
2872 }
2873 Mixer::Mla(mla) => self.mla_attn_cached(e, mla, &h, &pos_d, t, il, cache)?,
2874 Mixer::Kda(la) if la.tp.is_some() => crate::glm5_tp::kda_tp_cached(
2875 e,
2876 la,
2877 &h,
2878 t,
2879 eps,
2880 cache,
2881 il,
2882 crate::kda::ConvArm::Prefill,
2883 )?,
2884 Mixer::Kda(la) => crate::kda::kda_prime_cached(e, la, &h, t, eps, cache, il)?,
2885 };
2886 x = crate::hyper::post(e, &topology, &mixed, &x, &mix, t, n_embd)?;
2887
2888 let (y, mix) = crate::hyper::pre(e, &topology, &hyper.mlp, &x, t, n_embd)?;
2889 let mut z = e.uninit(t * n_embd)?;
2890 e.rms_norm(
2891 &y,
2892 layer.post_attn_norm.float_data(),
2893 &mut z,
2894 n_embd,
2895 t,
2896 eps,
2897 )?;
2898 let ffn_out = self.hyper_ffn_branch(e, layer, &z, t, il, true)?;
2899 x = crate::hyper::post(e, &topology, &ffn_out, &x, &mix, t, n_embd)?;
2900 self.glm5_hc_tap(e, cache, &topology, il, &x, t)?;
2903 }
2904
2905 let hiddens =
2906 crate::hyper::collapse(e, &topology, self.hyper_head.as_ref(), &x, t, n_embd)?;
2907 let mut hn = e.uninit(t * n_embd)?;
2908 e.rms_norm(
2909 &hiddens,
2910 self.output_norm.float_data(),
2911 &mut hn,
2912 n_embd,
2913 t,
2914 eps,
2915 )?;
2916 let last = e.view(&hn, t * n_embd);
2917 let last_row = last.slice((t - 1) * n_embd..t * n_embd);
2918 let mut hlast = e.uninit(n_embd)?;
2919 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
2920 let logits = e.matmul(&self.output, &hlast, 1)?;
2921 let host = e.dtoh(&logits)?;
2922
2923 let stack = e.view(&hiddens, t * n_embd);
2926 let seed_row = stack.slice((t - 1) * n_embd..t * n_embd);
2927 let mut h_seed = e.uninit(n_embd)?;
2928 e.copy_view_into(&mut h_seed, 0, &seed_row, n_embd)?;
2929 cache.pos += t;
2930 Ok((host, h_seed, hiddens))
2931 }
2932
2933 #[allow(clippy::too_many_arguments)]
2937 #[allow(clippy::type_complexity)] fn hyper_prime_tail(
2939 &self,
2940 e: &Engine,
2941 topology: &crate::hyper::HyperTopology,
2942 x: &CudaSlice<f32>,
2943 t: usize,
2944 n_embd: usize,
2945 eps: f32,
2946 cache: &mut Cache,
2947 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
2948 let hiddens = crate::hyper::collapse(e, topology, self.hyper_head.as_ref(), x, t, n_embd)?;
2949 let mut hn = e.uninit(t * n_embd)?;
2950 e.rms_norm(
2951 &hiddens,
2952 self.output_norm.float_data(),
2953 &mut hn,
2954 n_embd,
2955 t,
2956 eps,
2957 )?;
2958 let last = e.view(&hn, t * n_embd);
2959 let last_row = last.slice((t - 1) * n_embd..t * n_embd);
2960 let mut hlast = e.uninit(n_embd)?;
2961 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
2962 let logits = e.matmul(&self.output, &hlast, 1)?;
2963 let host = e.dtoh(&logits)?;
2964 let stack = e.view(&hiddens, t * n_embd);
2965 let seed_row = stack.slice((t - 1) * n_embd..t * n_embd);
2966 let mut h_seed = e.uninit(n_embd)?;
2967 e.copy_view_into(&mut h_seed, 0, &seed_row, n_embd)?;
2968 cache.pos += t;
2969 Ok((host, h_seed, hiddens))
2970 }
2971
2972 fn decode_step_hyper_ppn(
2977 &self,
2978 e: &Engine,
2979 token: u32,
2980 cache: &mut Cache,
2981 topology: &crate::hyper::HyperTopology,
2982 fence: &[usize],
2983 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
2984 let n_embd = self.cfg.n_embd as usize;
2985 let eps = self.cfg.rms_eps;
2986 let pos = cache.pos;
2987 let width = topology.streams * n_embd;
2988
2989 if crate::pp::pp2_streams_off() {
2990 let pos_d = e.htod_i32(&[pos as i32])?;
2991 let embedded = e.htod(&self.embd.gather(n_embd, &[token]))?;
2992 let mut x = crate::hyper::expand(e, topology, &embedded, 1, n_embd)?;
2993 x = self.hyper_range_decode(e, topology, x, fence[0], fence[1], &pos_d, pos, cache)?;
2994 for s in 1..fence.len() - 1 {
2995 let boundary_tx = e.clone_dtod(&x)?;
2996 let boundary_rx = e.clone_dtod(&boundary_tx)?;
2997 x = self.hyper_range_decode(
2998 e,
2999 topology,
3000 boundary_rx,
3001 fence[s],
3002 fence[s + 1],
3003 &pos_d,
3004 pos,
3005 cache,
3006 )?;
3007 }
3008 return self.hyper_decode_tail(e, topology, &x, n_embd, eps, cache);
3009 }
3010
3011 let rt = crate::pp::PpNRt::get(e)?;
3012 let n_st = fence.len() - 1;
3013 assert_eq!(
3014 rt.n_stages(),
3015 n_st,
3016 "PpNRt stage count {} != fence stages {n_st}",
3017 rt.n_stages()
3018 );
3019 rt.fence_stages_behind(&e.stream())?;
3020
3021 let mut slot = {
3022 let _st0 = rt.enter(0);
3023 let e0 = rt.engine(0, e);
3024 let pos_d = e0.htod_i32(&[pos as i32])?;
3025 let embedded = e0.htod(&self.embd.gather(n_embd, &[token]))?;
3026 let x = crate::hyper::expand(e0, topology, &embedded, 1, n_embd)?;
3027 let x =
3028 self.hyper_range_decode(e0, topology, x, fence[0], fence[1], &pos_d, pos, cache)?;
3029 rt.tx(0, &x, width)?
3030 };
3031 for s in 1..n_st - 1 {
3032 let _st = rt.enter(s);
3033 let es = rt.engine(s, e);
3034 let pos_d = es.htod_i32(&[pos as i32])?;
3035 let x = rt.rx(s - 1, slot, width)?;
3036 let x = self.hyper_range_decode(
3037 es,
3038 topology,
3039 x,
3040 fence[s],
3041 fence[s + 1],
3042 &pos_d,
3043 pos,
3044 cache,
3045 )?;
3046 slot = rt.tx(s, &x, width)?;
3047 }
3048 let _stl = rt.enter(n_st - 1);
3049 let el = rt.engine(n_st - 1, e);
3050 let pos_d = el.htod_i32(&[pos as i32])?;
3051 let x = rt.rx(n_st - 2, slot, width)?;
3052 let x = self.hyper_range_decode(
3053 el,
3054 topology,
3055 x,
3056 fence[n_st - 1],
3057 fence[n_st],
3058 &pos_d,
3059 pos,
3060 cache,
3061 )?;
3062 self.hyper_decode_tail(el, topology, &x, n_embd, eps, cache)
3063 }
3064
3065 fn hyper_decode_tail(
3069 &self,
3070 e: &Engine,
3071 topology: &crate::hyper::HyperTopology,
3072 x: &CudaSlice<f32>,
3073 n_embd: usize,
3074 eps: f32,
3075 cache: &mut Cache,
3076 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
3077 let h_seed = crate::hyper::collapse(e, topology, self.hyper_head.as_ref(), x, 1, n_embd)?;
3078 let mut hn = e.uninit(n_embd)?;
3079 e.rms_norm(
3080 &h_seed,
3081 self.output_norm.float_data(),
3082 &mut hn,
3083 n_embd,
3084 1,
3085 eps,
3086 )?;
3087 let logits = e.matmul(&self.output, &hn, 1)?;
3088 let host = e.dtoh(&logits)?;
3089 cache.pos += 1;
3090 Ok((host, h_seed))
3091 }
3092
3093 pub fn forward(
3095 &self,
3096 e: &Engine,
3097 tokens: &[u32],
3098 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
3099 if self.hyper.is_some() {
3100 return self.forward_hyper(e, tokens, false);
3101 }
3102 if self.is_gemma4_e4b() {
3103 return self.gemma4_e4b_forward(e, tokens, false);
3104 }
3105 if self.uses_gemma_program() {
3106 return self.gemma4_forward(e, tokens, false);
3107 }
3108 let cfg = &self.cfg;
3109 let n_embd = cfg.n_embd as usize;
3110 let t = tokens.len();
3111 let eps = cfg.rms_eps;
3112 let pos: Vec<i32> = (0..t as i32).collect();
3113 let pos_d = e.htod_i32(&pos)?;
3114
3115 let mut x = self.embed(e, tokens)?; for (il, layer) in self.layers.iter().enumerate() {
3118 let mut h = e.uninit(t * n_embd)?;
3120 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
3121
3122 let mixed = match &layer.mixer {
3123 Mixer::Full(fa) => self.full_attn(e, fa, &h, &pos_d, t, il)?,
3124 Mixer::Linear(la) => self.linear_attn(e, la, &h, t)?,
3125 Mixer::Mla(mla) => self.mla_attn(e, mla, &h, &pos_d, t, il)?,
3126 Mixer::Kda(la) => crate::kda::kda_attn(e, la, &h, t, eps)?,
3127 };
3128
3129 let mut x1 = e.uninit(t * n_embd)?;
3131 e.add(&x, &mixed, &mut x1, t * n_embd)?;
3132
3133 let mut z = e.uninit(t * n_embd)?;
3135 e.rms_norm(
3136 &x1,
3137 layer.post_attn_norm.float_data(),
3138 &mut z,
3139 n_embd,
3140 t,
3141 eps,
3142 )?;
3143 let ffn_out = match &layer.ffn {
3144 crate::hybrid::Ffn::Dense {
3145 ffn_gate,
3146 ffn_up,
3147 ffn_down,
3148 } => {
3149 let n_ff = ffn_gate.out_features();
3150 let mut g2 = e.matmul_group(&[ffn_gate, ffn_up], &z, t)?;
3151 let up = g2.pop().unwrap();
3152 let gate = g2.pop().unwrap();
3153 let mut act = e.uninit(t * n_ff)?;
3154 Self::ffn_act_lim(
3159 e,
3160 &self.cfg,
3161 &gate,
3162 &up,
3163 1.0,
3164 1.0,
3165 self.cfg.clamp_shexp_at(il as u32),
3166 &mut act,
3167 t * n_ff,
3168 )?;
3169 e.matmul(ffn_down, &act, t)?
3170 }
3171 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il_prefill(e, m, &z, t, il as u16)?,
3172 };
3173 let mut x2 = e.uninit(t * n_embd)?;
3174 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
3175 x = x2;
3176 }
3177
3178 let mut hn = e.uninit(t * n_embd)?;
3179 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
3180 let logits = e.matmul(&self.output, &hn, t)?;
3181 e.dtoh(&logits)
3182 }
3183
3184 pub fn forward_last(
3190 &self,
3191 e: &Engine,
3192 tokens: &[u32],
3193 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
3194 if self.hyper.is_some() {
3195 return self.forward_hyper(e, tokens, true);
3196 }
3197 if self.uses_gemma_program() {
3198 return self.gemma4_forward(e, tokens, true);
3199 }
3200 let cfg = &self.cfg;
3201 let n_embd = cfg.n_embd as usize;
3202 let t = tokens.len();
3203 let eps = cfg.rms_eps;
3204 let pos: Vec<i32> = (0..t as i32).collect();
3205 let pos_d = e.htod_i32(&pos)?;
3206
3207 let mut x = self.embed(e, tokens)?; let probe = std::env::var("MEMRA_LAYER_PROBE").is_ok();
3211 let anat = Self::prime_anatomy_on();
3212 let mut anat_last = if anat {
3213 e.stream().synchronize()?;
3214 Some(std::time::Instant::now())
3215 } else {
3216 None
3217 };
3218 macro_rules! anat_mark {
3219 ($slot:expr) => {
3220 if let Some(ts) = anat_last.as_mut() {
3221 e.stream().synchronize()?;
3222 Self::prime_anatomy_slots()[$slot].fetch_add(
3223 ts.elapsed().as_nanos() as u64,
3224 std::sync::atomic::Ordering::Relaxed,
3225 );
3226 *ts = std::time::Instant::now();
3227 }
3228 };
3229 }
3230 for (il, layer) in self.layers.iter().enumerate() {
3231 let mut h = e.uninit(t * n_embd)?;
3232 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
3233 if probe {
3234 e.stream().synchronize()?;
3235 eprintln!("[probe] L{il} norm ok");
3236 }
3237 anat_mark!(4);
3238 let mixed = match &layer.mixer {
3239 Mixer::Full(fa) => {
3240 let y = self.full_attn(e, fa, &h, &pos_d, t, il)?;
3241 anat_mark!(0);
3242 y
3243 }
3244 Mixer::Linear(la) => {
3245 let y = self.linear_attn(e, la, &h, t)?;
3246 anat_mark!(1);
3247 y
3248 }
3249 Mixer::Mla(mla) => self.mla_attn(e, mla, &h, &pos_d, t, il)?,
3250 Mixer::Kda(la) => {
3251 let y = crate::kda::kda_attn(e, la, &h, t, eps)?;
3252 anat_mark!(1);
3254 y
3255 }
3256 };
3257 if probe {
3258 e.stream().synchronize()?;
3259 eprintln!("[probe] L{il} mixer ok");
3260 }
3261 let mut x1 = e.uninit(t * n_embd)?;
3262 e.add(&x, &mixed, &mut x1, t * n_embd)?;
3263 let mut z = e.uninit(t * n_embd)?;
3264 e.rms_norm(
3265 &x1,
3266 layer.post_attn_norm.float_data(),
3267 &mut z,
3268 n_embd,
3269 t,
3270 eps,
3271 )?;
3272 anat_mark!(4);
3273 let ffn_out = match &layer.ffn {
3274 crate::hybrid::Ffn::Dense {
3275 ffn_gate,
3276 ffn_up,
3277 ffn_down,
3278 } => {
3279 let n_ff = ffn_gate.out_features();
3280 let mut g2 = e.matmul_group(&[ffn_gate, ffn_up], &z, t)?;
3281 let up = g2.pop().unwrap();
3282 let gate = g2.pop().unwrap();
3283 let mut act = e.uninit(t * n_ff)?;
3284 Self::ffn_act_lim(
3286 e,
3287 &self.cfg,
3288 &gate,
3289 &up,
3290 1.0,
3291 1.0,
3292 self.cfg.clamp_shexp_at(il as u32),
3293 &mut act,
3294 t * n_ff,
3295 )?;
3296 let y = e.matmul(ffn_down, &act, t)?;
3297 anat_mark!(3);
3298 y
3299 }
3300 crate::hybrid::Ffn::Moe(m) => {
3301 let y = self.moe_ffn_il_prefill(e, m, &z, t, il as u16)?;
3302 anat_mark!(2);
3303 y
3304 }
3305 };
3306 if probe {
3307 e.stream().synchronize()?;
3308 eprintln!("[probe] L{il} ffn ok");
3309 }
3310 let mut x2 = e.uninit(t * n_embd)?;
3311 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
3312 x = x2;
3313 }
3314 if anat {
3315 let s = Self::prime_anatomy_slots();
3316 let ms = |i: usize| s[i].load(std::sync::atomic::Ordering::Relaxed) as f64 / 1.0e6;
3317 eprintln!(
3318 "[prime-anatomy] cumulative ms: attn_full={:.1} gdn_linear={:.1} moe={:.1} \
3319 dense={:.1} norms_adds={:.1} (t={t}, forward_last)",
3320 ms(0),
3321 ms(1),
3322 ms(2),
3323 ms(3),
3324 ms(4)
3325 );
3326 }
3327 let mut hn = e.uninit(t * n_embd)?;
3329 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
3330 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)?;
3333 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
3334 let logits = e.matmul(&self.output, &hlast, 1)?; e.dtoh(&logits)
3336 }
3337
3338 #[allow(clippy::type_complexity)] pub fn prime_cache(
3371 &self,
3372 e: &Engine,
3373 tokens: &[u32],
3374 cache: &mut Cache,
3375 queued_after: usize,
3376 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
3377 self.prime_cache_overlaid(e, tokens, cache, queued_after, None)
3378 }
3379
3380 pub fn vision_intake_engine<'a>(
3395 &self,
3396 e: &'a Engine,
3397 ) -> Result<&'a Engine, Box<dyn std::error::Error>> {
3398 if self.hyper.is_some()
3399 && !crate::pp::pp2_streams_off()
3400 && crate::pp::pp_cuts(self.layers.len()).is_some()
3401 {
3402 let rt = crate::pp::PpNRt::get(e)?;
3403 return Ok(rt.engine(0, e));
3404 }
3405 Ok(e)
3406 }
3407
3408 #[allow(clippy::type_complexity)] pub fn prime_cache_overlaid(
3417 &self,
3418 e: &Engine,
3419 tokens: &[u32],
3420 cache: &mut Cache,
3421 queued_after: usize,
3422 overlay: Option<&crate::vision::EmbedOverlay>,
3423 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
3424 cache.ensure_usable("prime_cache")?;
3425 if self.hyper.is_some() {
3426 return self.prime_cache_hyper(e, tokens, cache, queued_after, overlay);
3430 }
3431 let _pp_walk =
3432 if crate::pp::pp_cuts(self.layers.len()).is_some() && !crate::pp::pp2_streams_off() {
3433 let rt = crate::pp::PpNRt::get(e)?;
3434 Some(rt.acquire_walk("prime_cache")?)
3435 } else {
3436 None
3437 };
3438 let n_embd = self.cfg.n_embd as usize;
3439 let t = tokens.len();
3440 let legacy_calllocal = std::env::var("MEMRA_PRIME_CALLLOCAL").as_deref() == Ok("1");
3469 let seq_end = if legacy_calllocal {
3470 cache.pos + t
3471 } else {
3472 cache.pos + t + queued_after
3473 };
3474 if overlay.is_none()
3479 && (cache.pos == 0 || step_gemm_prime_suffix_on())
3480 && t >= PRIME_MIN_T
3481 && crate::step_gemm_prime_on()
3482 && self.uses_sliding_gated_moe_program()
3483 {
3484 let n_embd = self.cfg.n_embd as usize;
3485 let base = cache.pos;
3486 let width = crate::cache::PRIME_CHUNK_MAX_TOKENS;
3487 let mut hiddens = e.uninit(t * n_embd)?;
3488 let mut last: Option<(Vec<f32>, CudaSlice<f32>)> = None;
3489 let mut start = 0usize;
3490 while start < t {
3491 let mut end = (start + width).min(t);
3494 if t - end > 0 && t - end < PRIME_MIN_T {
3495 end = t;
3496 }
3497 let mut out = self.step35_prime_cache_batch(
3498 e,
3499 &[&tokens[start..end]],
3500 &mut [cache],
3501 &[seq_end],
3502 )?;
3503 if out.len() != 1 {
3504 return Err("B=1 batched prime returned a non-singleton".into());
3505 }
3506 let (logits, h_seed, hidden) = out.remove(0);
3507 e.copy_into(
3508 &mut hiddens,
3509 start * n_embd,
3510 &hidden,
3511 (end - start) * n_embd,
3512 )?;
3513 last = Some((logits, h_seed));
3514 start = end;
3515 }
3516 let (logits, h_seed) = last.expect("prime produced no chunk");
3517 eprintln!(
3522 "[gemm-prime] ENGAGED t={t} base={base} seq_end={seq_end} chunks<={width} (GEMM trunk + grouped MoE)"
3523 );
3524 return Ok((logits, h_seed, hiddens));
3525 }
3526 if self.uses_sliding_gated_moe_program() {
3527 eprintln!(
3528 "[gemm-prime] WALK t={t} base={} seq_end={seq_end} (batched prime declined)",
3529 cache.pos
3530 );
3531 }
3532 if overlay.is_none()
3533 && let Some(out) = self.step35_prime_trows(e, tokens, cache)?
3534 {
3535 return Ok(out);
3536 }
3537 assert!(
3541 t >= PRIME_MIN_T,
3542 "prime_cache needs T >= {PRIME_MIN_T} (caller gates)"
3543 );
3544 assert!(
3545 cache.pos + t <= cache.max_ctx,
3546 "prime_cache: prompt exceeds cache max_ctx"
3547 );
3548
3549 if self.is_gemma4_e4b() || self.uses_gemma_program() {
3561 if self.is_gemma4_e4b() {
3562 if overlay.is_some() {
3563 return Err(
3564 "vision embedding overlay is unsupported on gemma4 E4B (PLE prime)".into(),
3565 );
3566 }
3567 return self.gemma4_e4b_prime(e, tokens, cache);
3568 }
3569 return self.gemma4_prime(e, tokens, cache, overlay);
3574 }
3575 if crate::pp::prime_pipe_on()
3576 && crate::pp::prime_pp_on()
3577 && !crate::pp::pp2_streams_off()
3578 && crate::pp::pp_cuts(self.layers.len())
3579 .is_some_and(|fence| matches!(fence.len(), 4 | 5))
3580 {
3581 crate::pp::pp_wave_on()
3582 .map_err(|reason| -> Box<dyn std::error::Error> { reason.into() })?;
3583 }
3584 let ranges = prime_chunk_ranges(t, self.layers.len(), self.gdn_prime_grid_on());
3585 if ranges.len() == 1 {
3623 return self.prime_chunk(e, tokens, cache, seq_end, 0, overlay);
3624 }
3625 if crate::pp::prime_pipe_on() && crate::pp::prime_pp_on() && !crate::pp::pp2_streams_off() {
3629 if let Some(fence) = crate::pp::pp_cuts(self.layers.len()).filter(|f| f.len() == 3) {
3630 if overlay.is_some() {
3631 return Err(
3632 "vision embedding overlay + pipelined PP prime unsupported (v1); \
3633 run the serial prime (single device or MEMRA_PRIME_PIPE=0)"
3634 .into(),
3635 );
3636 }
3637 if crate::pp::pp_multi_stream_same_device() {
3638 return Err(
3639 "prime chunk pipeline refused with 2 stage streams on one device — \
3640 that concurrent-stream placement remains quarantined by the deferred \
3641 pp flake record. Use one device per stage or MEMRA_PRIME_PIPE=0 for \
3642 the serial split."
3643 .into(),
3644 );
3645 }
3646 return self.prime_cache_pp2_pipelined(e, tokens, cache, seq_end, &ranges, &fence);
3647 }
3648 if let Some(fence) =
3649 crate::pp::pp_cuts(self.layers.len()).filter(|f| matches!(f.len(), 4 | 5))
3650 {
3651 let wave_on = crate::pp::pp_wave_on()
3652 .map_err(|reason| -> Box<dyn std::error::Error> { reason.into() })?;
3653 let stages = fence.len() - 1;
3654 if crate::pp::pp_wave_route_enabled(
3655 wave_on,
3656 crate::pp::pp2_overlap(),
3657 stages,
3658 ranges.len(),
3659 ) {
3660 if overlay.is_some() {
3661 return Err(
3662 "vision embedding overlay + pipelined PP prime unsupported; \
3663 run the serial prime (MEMRA_PP_WAVE=0 or MEMRA_PRIME_PIPE=0)"
3664 .into(),
3665 );
3666 }
3667 let rt = crate::pp::PpNRt::get(e)?;
3668 let double_slot = crate::pp::pp2_overlap();
3669 crate::pp::pp_wave_eligibility(
3670 stages,
3671 double_slot,
3672 rt.host_bounce_active(),
3673 rt.repeated_stage_device(),
3674 )
3675 .map_err(|reason| -> Box<dyn std::error::Error> { reason.into() })?;
3676 return self
3677 .prime_cache_ppn_pipelined(e, tokens, cache, seq_end, &ranges, &fence);
3678 }
3679 }
3680 }
3681 let mut hiddens = e.uninit(t * n_embd)?;
3682 let mut last: Option<(Vec<f32>, CudaSlice<f32>)> = None;
3683 for &(start, end) in &ranges {
3684 if let Some(taps) = cache.dflash_taps.as_mut() {
3686 taps.base = start;
3687 }
3688 let (l, hs, x) =
3689 self.prime_chunk(e, &tokens[start..end], cache, seq_end, start, overlay)?;
3690 e.copy_into(&mut hiddens, start * n_embd, &x, (end - start) * n_embd)?;
3691 last = Some((l, hs));
3692 }
3693 let (logits, h_seed) = last.unwrap();
3694 Ok((logits, h_seed, hiddens))
3695 }
3696
3697 #[allow(clippy::type_complexity)] fn prime_cache_pp2_pipelined(
3703 &self,
3704 e: &Engine,
3705 tokens: &[u32],
3706 cache: &mut Cache,
3707 seq_end: usize,
3708 ranges: &[(usize, usize)],
3709 fence: &[usize],
3710 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
3711 debug_assert_eq!(fence.len(), 3);
3712 debug_assert!(ranges.len() >= 2);
3713 let rt = crate::pp::PpNRt::get(e)?;
3714 assert_eq!(
3715 rt.n_stages(),
3716 2,
3717 "prime pipeline requires exactly two PP stages"
3718 );
3719 let n_embd = self.cfg.n_embd as usize;
3720 let t = tokens.len();
3721 let initial_base = cache.pos;
3722 let caller_stream = e.stream();
3723
3724 rt.fence_stages_behind(&caller_stream)?;
3729 let max_payload = ranges.iter().map(|(s, e)| (e - s) * n_embd).max().unwrap();
3730 rt.prepare_overlap_slots(0, max_payload)?;
3731
3732 let mut hiddens = e.uninit(t * n_embd)?;
3733 let mut last: Option<(Vec<f32>, CudaSlice<f32>)> = None;
3734 let mut stage_caches = PrimeCacheStages::new(cache, fence);
3735 let (cache0, cache1) = stage_caches.pp2_parts();
3736 let (first_start, first_end) = ranges[0];
3737 let mut slot = self.prime_pp2_stage0_enqueue(
3738 e,
3739 rt,
3740 &tokens[first_start..first_end],
3741 cache0,
3742 seq_end,
3743 fence,
3744 initial_base + first_start,
3745 true,
3746 )?;
3747 cache0.pos = initial_base + first_end;
3748
3749 for (i, &(start, end)) in ranges.iter().enumerate() {
3750 let base = initial_base + start;
3751 debug_assert_eq!(
3752 cache1.pos, base,
3753 "stage 1 must drain chunks in original position order"
3754 );
3755 let (out, next_slot) = if let Some(&(next_start, next_end)) = ranges.get(i + 1) {
3756 let next_base = initial_base + next_start;
3757 debug_assert_eq!(
3758 cache0.pos, next_base,
3759 "stage 0 must issue chunks in original position order"
3760 );
3761 let cache0_stage = &mut *cache0;
3762 std::thread::scope(|scope| -> Result<_, Box<dyn std::error::Error>> {
3767 let stage0 = scope.spawn(move || -> Result<usize, String> {
3768 let next = self
3769 .prime_pp2_stage0_enqueue(
3770 e,
3771 rt,
3772 &tokens[next_start..next_end],
3773 cache0_stage,
3774 seq_end,
3775 fence,
3776 next_base,
3777 true,
3778 )
3779 .map_err(|err| err.to_string())?;
3780 cache0_stage.pos = initial_base + next_end;
3781 Ok(next)
3782 });
3783 let x = self.prime_pp2_stage1_enqueue(
3784 e,
3785 rt,
3786 slot,
3787 end - start,
3788 cache1,
3789 seq_end,
3790 fence,
3791 base,
3792 true,
3793 )?;
3794 let out = {
3795 rt.bind_stage(1)?;
3796 let _st1 = rt.enter(1);
3797 let e1 = rt.engine(1, e);
3798 self.prime_chunk_epilogue(e1, x, end - start, cache1)?
3799 };
3800 let next = match stage0.join() {
3801 Ok(result) => {
3802 result.map_err(|err| -> Box<dyn std::error::Error> { err.into() })?
3803 }
3804 Err(payload) => std::panic::resume_unwind(payload),
3805 };
3806 Ok((out, Some(next)))
3807 })?
3808 } else {
3809 let x = self.prime_pp2_stage1_enqueue(
3810 e,
3811 rt,
3812 slot,
3813 end - start,
3814 cache1,
3815 seq_end,
3816 fence,
3817 base,
3818 true,
3819 )?;
3820 let out = {
3821 rt.bind_stage(1)?;
3822 let _st1 = rt.enter(1);
3823 let e1 = rt.engine(1, e);
3824 self.prime_chunk_epilogue(e1, x, end - start, cache1)?
3825 };
3826 (out, None)
3827 };
3828
3829 rt.publish_to(1, &caller_stream)?;
3830 e.copy_into(&mut hiddens, start * n_embd, &out.2, (end - start) * n_embd)?;
3831 last = Some((out.0, out.1));
3832 crate::pp::PRIME_SPLIT_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
3833
3834 if let Some(next) = next_slot {
3835 rt.fence_stages_behind(&caller_stream)?;
3840 slot = next;
3841 }
3842 }
3843
3844 debug_assert_eq!(cache0.pos, initial_base + t);
3845 debug_assert_eq!(cache1.pos, initial_base + t);
3846 let (logits, h_seed) = last.unwrap();
3847 stage_caches.commit();
3848 Ok((logits, h_seed, hiddens))
3849 }
3850
3851 #[allow(clippy::type_complexity)] fn prime_cache_ppn_pipelined(
3861 &self,
3862 e: &Engine,
3863 tokens: &[u32],
3864 cache: &mut Cache,
3865 seq_end: usize,
3866 ranges: &[(usize, usize)],
3867 fence: &[usize],
3868 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
3869 let stages = fence.len().saturating_sub(1);
3870 debug_assert!((3..=4).contains(&stages));
3871 debug_assert!(ranges.len() >= 2);
3872 let rt = crate::pp::PpNRt::get(e)?;
3873 assert_eq!(
3874 rt.n_stages(),
3875 stages,
3876 "prime wavefront PpNRt/fence stage mismatch"
3877 );
3878 let n_embd = self.cfg.n_embd as usize;
3879 let initial_base = cache.pos;
3880 let caller_stream = e.stream();
3881 let primary_context = crate::pp::PrimaryContextRestore::new(e);
3882
3883 rt.fence_stages_behind(&caller_stream)?;
3884 let max_payload = ranges
3885 .iter()
3886 .map(|(start, end)| (end - start) * n_embd)
3887 .max()
3888 .unwrap_or(0);
3889 for boundary in 0..stages - 1 {
3890 rt.prepare_overlap_slots(boundary, max_payload)?;
3891 }
3892
3893 let mut stage_caches = PrimeCacheStages::new(cache, fence);
3894 let waves: Vec<_> = ranges
3895 .iter()
3896 .map(|&(start, end)| PrimePpWave {
3897 start,
3898 end,
3899 tokens: &tokens[start..end],
3900 })
3901 .collect();
3902 let mut forward_senders = Vec::with_capacity(stages - 1);
3903 let mut forward_receivers = Vec::with_capacity(stages - 1);
3904 let mut release_senders = Vec::with_capacity(stages - 1);
3905 let mut release_receivers = Vec::with_capacity(stages - 1);
3906 for _ in 0..stages - 1 {
3907 let (forward_sender, forward_receiver) = std::sync::mpsc::channel();
3908 let (release_sender, release_receiver) = std::sync::mpsc::channel();
3909 forward_senders.push(Some(forward_sender));
3910 forward_receivers.push(Some(forward_receiver));
3911 release_senders.push(Some(release_sender));
3912 release_receivers.push(Some(release_receiver));
3913 }
3914 let mut stage_channels = Vec::with_capacity(stages - 1);
3915 for stage in 0..stages - 1 {
3916 stage_channels.push(Some(PrimePpStageChannels {
3917 incoming: (stage > 0).then(|| forward_receivers[stage - 1].take().unwrap()),
3918 release_upstream: (stage > 0).then(|| release_senders[stage - 1].take().unwrap()),
3919 outgoing: forward_senders[stage].take().unwrap(),
3920 released_downstream: release_receivers[stage].take().unwrap(),
3921 }));
3922 }
3923 let head_incoming = forward_receivers[stages - 2].take().unwrap();
3924 let head_release = release_senders[stages - 2].take().unwrap();
3925 #[allow(clippy::type_complexity)]
3928 let mut results: Vec<Option<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>)>> =
3929 std::iter::repeat_with(|| None).take(waves.len()).collect();
3930 let walk_result = std::thread::scope(|scope| -> Result<(), Box<dyn std::error::Error>> {
3931 let waves_ref = &waves;
3932 let mut handles = Vec::with_capacity(stages - 1);
3933 #[allow(clippy::needless_range_loop)]
3937 for stage in 0..stages - 1 {
3938 let channels = stage_channels[stage].take().unwrap();
3939 let cache_state = &stage_caches.stages()[stage];
3940 handles.push(scope.spawn(move || -> Result<(), String> {
3941 let result = self.prime_ppn_wave_worker(
3942 e,
3943 rt,
3944 waves_ref,
3945 cache_state,
3946 channels.incoming.as_ref(),
3947 channels.release_upstream.as_ref(),
3948 &channels.outgoing,
3949 &channels.released_downstream,
3950 stage,
3951 seq_end,
3952 fence,
3953 initial_base,
3954 );
3955 if let Err(error) = &result {
3956 channels.notify_failure(&error.to_string());
3957 }
3958 result.map_err(|error| error.to_string())
3959 }));
3960 }
3961
3962 let mut head_panic = None;
3963 let head_result = match std::panic::catch_unwind(std::panic::AssertUnwindSafe(
3964 || -> Result<(), Box<dyn std::error::Error>> {
3965 let mut head_cache = stage_caches.stages()[stages - 1]
3966 .lock()
3967 .map_err(|_| "prime PP head cache lock poisoned")?;
3968 for (wave_index, wave) in waves_ref.iter().enumerate() {
3969 let incoming = recv_prime_pp_signal(
3970 &head_incoming,
3971 PrimePpWaveSlot {
3972 wave: wave_index,
3973 slot: 0,
3974 },
3975 false,
3976 "prime PP head input",
3977 )?;
3978 results[wave_index] = Some(self.prime_ppn_wave_final(
3979 e,
3980 rt,
3981 wave,
3982 &mut head_cache,
3983 incoming,
3984 &head_release,
3985 seq_end,
3986 fence,
3987 initial_base,
3988 )?);
3989 }
3990 Ok(())
3991 },
3992 )) {
3993 Ok(result) => result,
3994 Err(payload) => {
3995 head_panic = Some(payload);
3996 Err("prime PP head-stage host walker panicked".into())
3997 }
3998 };
3999 if let Err(error) = &head_result {
4000 let _ = head_release.send(PrimePpSignal::Error(error.to_string()));
4001 }
4002 let mut first_error = head_result.err().map(|error| error.to_string());
4003 let mut worker_panic = None;
4004 for handle in handles {
4005 match handle.join() {
4006 Ok(Ok(())) => {}
4007 Ok(Err(error)) => {
4008 first_error.get_or_insert(error);
4009 }
4010 Err(payload) => {
4011 if worker_panic.is_none() {
4012 worker_panic = Some(payload);
4013 }
4014 }
4015 }
4016 }
4017 if let Some(payload) = head_panic {
4018 std::panic::resume_unwind(payload);
4019 }
4020 if let Some(payload) = worker_panic {
4021 std::panic::resume_unwind(payload);
4022 }
4023 if let Some(error) = first_error {
4024 return Err(error.into());
4025 }
4026 Ok(())
4027 });
4028 let publish_result = if walk_result.is_ok() {
4029 Some(rt.publish_to(stages - 1, &caller_stream))
4030 } else {
4031 None
4032 };
4033 let restore_result = primary_context.restore();
4034 walk_result?;
4035 if let Some(result) = publish_result {
4036 result?;
4037 }
4038 restore_result?;
4039 let mut hiddens = e.uninit(tokens.len() * n_embd)?;
4040 let mut last: Option<(Vec<f32>, CudaSlice<f32>)> = None;
4041 for (wave_index, (wave, result)) in waves.iter().zip(results).enumerate() {
4042 debug_assert_eq!((wave.start, wave.end), ranges[wave_index]);
4043 let out = result.ok_or("prime PP wavefront completed without a head-stage result")?;
4044 e.copy_into(
4045 &mut hiddens,
4046 wave.start * n_embd,
4047 &out.2,
4048 (wave.end - wave.start) * n_embd,
4049 )?;
4050 last = Some((out.0, out.1));
4051 crate::pp::PRIME_SPLIT_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
4052 }
4053 stage_caches.commit();
4054 drop(stage_caches);
4055
4056 static LOGGED: std::sync::Once = std::sync::Once::new();
4057 LOGGED.call_once(|| {
4058 eprintln!(
4059 "[pp-wave] PP{stages} prime wavefront engaged: microchunks={} \
4060 (experimental, MEMRA_PP_WAVE=1)",
4061 ranges.len(),
4062 );
4063 });
4064 let (logits, h_seed) = last.expect("prime PP wavefront produced no microchunk");
4065 crate::pp::record_pp_wave_tick();
4066 Ok((logits, h_seed, hiddens))
4067 }
4068
4069 #[allow(clippy::too_many_arguments)]
4070 fn prime_ppn_wave_worker(
4071 &self,
4072 e: &Engine,
4073 rt: &crate::pp::PpNRt,
4074 waves: &[PrimePpWave<'_>],
4075 cache: &std::sync::Mutex<Cache>,
4076 incoming: Option<&std::sync::mpsc::Receiver<PrimePpSignal>>,
4077 release_upstream: Option<&std::sync::mpsc::Sender<PrimePpSignal>>,
4078 outgoing: &std::sync::mpsc::Sender<PrimePpSignal>,
4079 released_downstream: &std::sync::mpsc::Receiver<PrimePpSignal>,
4080 stage: usize,
4081 seq_end: usize,
4082 fence: &[usize],
4083 initial_base: usize,
4084 ) -> Result<(), Box<dyn std::error::Error>> {
4085 debug_assert_eq!(incoming.is_some(), stage > 0);
4086 debug_assert_eq!(release_upstream.is_some(), stage > 0);
4087 let mut cache = cache
4088 .lock()
4089 .map_err(|_| "prime PP cache stage lock poisoned")?;
4090 let mut credits = PrimePpWaveCredits::default();
4091 for (wave_index, wave) in waves.iter().enumerate() {
4092 let incoming = match incoming {
4093 Some(receiver) => Some(recv_prime_pp_signal(
4094 receiver,
4095 PrimePpWaveSlot {
4096 wave: wave_index,
4097 slot: 0,
4098 },
4099 false,
4100 "prime PP stage input",
4101 )?),
4102 None => None,
4103 };
4104 let sent = self.prime_ppn_wave_stage(
4105 e,
4106 rt,
4107 wave,
4108 &mut cache,
4109 stage,
4110 incoming,
4111 release_upstream,
4112 &mut credits,
4113 released_downstream,
4114 seq_end,
4115 fence,
4116 initial_base,
4117 )?;
4118 send_prime_pp_signal(outgoing, PrimePpSignal::Slot(sent), "prime PP stage output")?;
4119 }
4120 while let Some(expected) = credits.pending.front().copied() {
4121 let released = recv_prime_pp_signal(
4122 released_downstream,
4123 expected,
4124 true,
4125 "prime PP final slot release",
4126 )?;
4127 credits.record_release(released)?;
4128 }
4129 Ok(())
4130 }
4131
4132 #[allow(clippy::too_many_arguments)]
4133 fn prime_ppn_wave_stage(
4134 &self,
4135 e: &Engine,
4136 rt: &crate::pp::PpNRt,
4137 wave: &PrimePpWave<'_>,
4138 cache: &mut Cache,
4139 stage: usize,
4140 incoming: Option<PrimePpWaveSlot>,
4141 release_upstream: Option<&std::sync::mpsc::Sender<PrimePpSignal>>,
4142 credits: &mut PrimePpWaveCredits,
4143 released_downstream: &std::sync::mpsc::Receiver<PrimePpSignal>,
4144 seq_end: usize,
4145 fence: &[usize],
4146 initial_base: usize,
4147 ) -> Result<PrimePpWaveSlot, Box<dyn std::error::Error>> {
4148 debug_assert!(stage + 1 < fence.len() - 1);
4149 let t = wave.end - wave.start;
4150 let base = initial_base + wave.start;
4151 debug_assert_eq!(cache.pos, base, "prime PP stage advanced out of order");
4152 let n_embd = self.cfg.n_embd as usize;
4153 let payload = t * n_embd;
4154 let positions: Vec<i32> = (base as i32..(base + t) as i32).collect();
4155 rt.bind_stage(stage)?;
4156 let _stage = rt.enter(stage);
4157 let engine = rt.engine(stage, e);
4158 let positions_d = engine.htod_i32(&positions)?;
4159 let x = if stage == 0 {
4160 debug_assert!(incoming.is_none());
4161 self.embed(engine, wave.tokens)?
4162 } else {
4163 let incoming = incoming.ok_or("prime PP stage has no incoming boundary slot")?;
4164 let x = rt.rx(stage - 1, incoming.slot, payload)?;
4165 send_prime_pp_signal(
4166 release_upstream.ok_or("prime PP stage has no upstream release channel")?,
4167 PrimePpSignal::Slot(incoming),
4168 "prime PP upstream slot release",
4169 )?;
4170 x
4171 };
4172 let x = {
4173 let _wave_cell = crate::pp::enter_pp_wave_cell();
4174 let _overlap = crate::pp::enter_prime_pipe_stage();
4175 self.prime_layers(
4176 engine,
4177 x,
4178 fence[stage],
4179 fence[stage + 1],
4180 &positions_d,
4181 t,
4182 base,
4183 cache,
4184 seq_end,
4185 )?
4186 };
4187 if let Some(expected) = credits.release_required() {
4188 let released =
4189 recv_prime_pp_signal(released_downstream, expected, true, "prime PP slot credit")?;
4190 credits.record_release(released)?;
4191 }
4192 let sent = PrimePpWaveSlot {
4193 wave: credits.next_wave,
4194 slot: rt.tx_pipelined(stage, &x, payload)?,
4195 };
4196 credits.record_send(sent)?;
4197 cache.pos = base + t;
4198 Ok(sent)
4199 }
4200
4201 #[allow(clippy::too_many_arguments)]
4202 #[allow(clippy::type_complexity)] fn prime_ppn_wave_final(
4204 &self,
4205 e: &Engine,
4206 rt: &crate::pp::PpNRt,
4207 wave: &PrimePpWave<'_>,
4208 cache: &mut Cache,
4209 incoming: PrimePpWaveSlot,
4210 release_upstream: &std::sync::mpsc::Sender<PrimePpSignal>,
4211 seq_end: usize,
4212 fence: &[usize],
4213 initial_base: usize,
4214 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4215 let stage = fence.len() - 2;
4216 let t = wave.end - wave.start;
4217 let base = initial_base + wave.start;
4218 debug_assert_eq!(cache.pos, base, "prime PP head stage advanced out of order");
4219 let n_embd = self.cfg.n_embd as usize;
4220 let payload = t * n_embd;
4221 let positions: Vec<i32> = (base as i32..(base + t) as i32).collect();
4222 rt.bind_stage(stage)?;
4223 let _stage = rt.enter(stage);
4224 let engine = rt.engine(stage, e);
4225 let positions_d = engine.htod_i32(&positions)?;
4226 let x = rt.rx(stage - 1, incoming.slot, payload)?;
4227 send_prime_pp_signal(
4228 release_upstream,
4229 PrimePpSignal::Slot(incoming),
4230 "prime PP head slot release",
4231 )?;
4232 let _wave_cell = crate::pp::enter_pp_wave_cell();
4233 let _overlap = crate::pp::enter_prime_pipe_stage();
4234 let x = self.prime_layers(
4235 engine,
4236 x,
4237 fence[stage],
4238 fence[stage + 1],
4239 &positions_d,
4240 t,
4241 base,
4242 cache,
4243 seq_end,
4244 )?;
4245 self.prime_chunk_epilogue(engine, x, t, cache)
4246 }
4247
4248 fn gdn_hk(e: &Engine, t: usize, num_v: usize, num_k: usize) -> usize {
4255 if Engine::gdn_db_on()
4256 && Engine::gdn_chunked_enabled()
4257 && t >= 16
4258 && e.gdn_mma_enabled(Engine::gdn_chunk_size())
4259 && num_k * 2 == num_v
4260 {
4261 num_k
4262 } else {
4263 num_v
4264 }
4265 }
4266
4267 fn f16out_on(e: &Engine, t: usize) -> bool {
4272 crate::f16_ffi::pp_f16_enabled()
4273 && t >= 16
4274 && !e.verify_exact_on()
4275 && std::env::var("MEMRA_F16OUT").as_deref() != Ok("0")
4276 }
4277
4278 pub fn prime_slabs_get(
4286 &self,
4287 e: &Engine,
4288 t: usize,
4289 n_embd: usize,
4290 n_ff_max: usize,
4291 ) -> Result<std::sync::Arc<std::sync::Mutex<PrimeSlabs>>, Box<dyn std::error::Error>> {
4292 let mut slabs = self.prime_slabs.lock().unwrap();
4293 let dev = e.ctx().ordinal();
4294 let need_new = match slabs.get(&dev) {
4295 None => true,
4296 Some(sl) => sl.lock().unwrap().t_cap < t,
4297 };
4298 if need_new {
4299 slabs.insert(
4300 dev,
4301 std::sync::Arc::new(std::sync::Mutex::new(PrimeSlabs {
4302 t_cap: t,
4303 h: e.uninit(t * n_embd)?,
4304 x1: e.uninit(t * n_embd)?,
4305 z: e.uninit(t * n_embd)?,
4306 act: e.uninit(t * n_ff_max)?,
4307 xa: e.uninit(t * n_embd)?,
4308 xb: e.uninit(t * n_embd)?,
4309 h16: e.alloc_u8_uninit(t * n_embd * 2)?,
4310 z16: e.alloc_u8_uninit(t * n_embd * 2)?,
4311 gate: e.uninit(t * n_ff_max)?,
4312 up: e.uninit(t * n_ff_max)?,
4313 ffn_out: e.uninit(t * n_embd)?,
4314 seg_glue: Vec::new(),
4315 mixed: e.uninit(t * n_embd)?,
4316 seg_mid: Vec::new(),
4317 seg_t: 0,
4318 })),
4319 );
4320 }
4321 Ok(slabs.get(&dev).expect("prime slab inserted").clone())
4322 }
4323
4324 #[allow(clippy::type_complexity)] fn prime_chunk(
4329 &self,
4330 e: &Engine,
4331 tokens: &[u32],
4332 cache: &mut Cache,
4333 seq_end: usize,
4334 chunk_off: usize,
4335 overlay: Option<&crate::vision::EmbedOverlay>,
4336 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4337 if crate::pp::pp_host_bounce_active()
4338 && (self.uses_gemma_program() || !crate::pp::prime_pp_on())
4339 {
4340 return Err(
4341 "prime_chunk: refused with MEMRA_PP_HOST_BOUNCE=1 because this configuration \
4342 has no active prime stage split and would peer-read remote weights; keep \
4343 MEMRA_PRIME_PP enabled and use a PP-prime-supported model"
4344 .into(),
4345 );
4346 }
4347 if !self.uses_gemma_program()
4356 && !crate::pp::pp2_streams_off()
4357 && crate::pp::prime_pp_on()
4358 && let Some(fence) = crate::pp::pp_cuts(self.layers.len())
4359 {
4360 if overlay.is_some() {
4361 return Err("vision embedding overlay + PP prime unsupported (v1); \
4362 run single-device or MEMRA_PRIME_PP=0"
4363 .into());
4364 }
4365 return self.prime_chunk_ppn(e, tokens, cache, seq_end, &fence);
4366 }
4367 if crate::pp::pp_host_bounce_active() {
4368 return Err(
4369 "prime_chunk: MEMRA_PP_HOST_BOUNCE=1 found no valid prime stage split; \
4370 refusing an unsplit remote-weight walk"
4371 .into(),
4372 );
4373 }
4374 let t = tokens.len();
4375 let base = cache.pos;
4376 debug_assert!(
4377 seq_end >= base + t,
4378 "prime_chunk: seq_end must cover this chunk"
4379 );
4380 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
4381 let pos_d = e.htod_i32(&pos)?;
4382
4383 let mut x_embed = self.embed(e, tokens)?; if let Some(ov) = overlay {
4385 ov.splice_into(e, &mut x_embed, chunk_off, t, self.cfg.n_embd as usize)?;
4394 }
4395 let x = self.prime_layers(
4396 e,
4397 x_embed,
4398 0,
4399 self.layers.len(),
4400 &pos_d,
4401 t,
4402 base,
4403 cache,
4404 seq_end,
4405 )?;
4406 self.prime_chunk_epilogue(e, x, t, cache)
4407 }
4408
4409 #[allow(clippy::too_many_arguments)]
4425 fn prime_layers(
4426 &self,
4427 e: &Engine,
4428 x_in: CudaSlice<f32>,
4429 lo: usize,
4430 hi: usize,
4431 pos_d: &CudaSlice<i32>,
4432 t: usize,
4433 base: usize,
4434 cache: &mut Cache,
4435 seq_end: usize,
4436 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4437 let cfg = &self.cfg;
4438 let n_embd = cfg.n_embd as usize;
4439 let eps = cfg.rms_eps;
4440 let f16fuse = crate::f16_ffi::pp_f16_enabled() && t >= 16;
4444 let n_ff_max = self
4451 .layers
4452 .iter()
4453 .map(|l| match &l.ffn {
4454 crate::hybrid::Ffn::Dense { ffn_gate, .. } => ffn_gate.out_features(),
4455 _ => n_embd,
4456 })
4457 .max()
4458 .unwrap_or(n_embd)
4459 .max(n_embd);
4460 let use_slabs = std::env::var("MEMRA_PRIME_SLABS").as_deref() != Ok("0");
4461 let slab = if use_slabs {
4462 Some(self.prime_slabs_get(e, t, n_embd, n_ff_max)?)
4463 } else {
4464 None
4465 };
4466 let mut slab_guard = slab.as_ref().map(|sl| sl.lock().unwrap());
4467 let mut x_own; type SlabRefs<'a> = (
4469 &'a mut CudaSlice<f32>,
4470 &'a mut CudaSlice<f32>,
4471 &'a mut CudaSlice<f32>,
4472 &'a mut CudaSlice<f32>,
4473 &'a mut CudaSlice<u8>,
4474 &'a mut CudaSlice<u8>,
4475 &'a mut CudaSlice<f32>,
4476 &'a mut CudaSlice<f32>,
4477 &'a mut CudaSlice<f32>,
4478 );
4479 let (mut x_cur, mut x_nxt, sl): (
4480 &mut CudaSlice<f32>,
4481 &mut CudaSlice<f32>,
4482 Option<SlabRefs>,
4483 );
4484 #[allow(clippy::type_complexity)]
4485 let mut seg: Option<(
4487 &mut Vec<Option<cudarc::driver::CudaGraph>>,
4488 &mut Vec<Option<cudarc::driver::CudaGraph>>,
4489 &mut CudaSlice<f32>,
4490 &mut usize,
4491 )> = None;
4492 let mut x_own2;
4493 match slab_guard.as_mut() {
4494 Some(g) => {
4495 let slabs = &mut **g;
4496 e.copy_into(&mut slabs.xa, 0, &x_in, t * n_embd)?;
4497 let PrimeSlabs {
4498 xa,
4499 xb,
4500 h,
4501 x1,
4502 z,
4503 act,
4504 h16,
4505 z16,
4506 gate,
4507 up,
4508 ffn_out,
4509 seg_glue,
4510 mixed,
4511 seg_mid,
4512 seg_t,
4513 ..
4514 } = slabs;
4515 x_cur = xa;
4516 x_nxt = xb;
4517 seg = Some((seg_glue, seg_mid, mixed, seg_t));
4518 sl = Some((h, x1, z, act, h16, z16, gate, up, ffn_out));
4519 }
4520 None => {
4521 x_own = x_in;
4522 x_own2 = e.uninit(t * n_embd)?;
4523 x_cur = &mut x_own;
4524 x_nxt = &mut x_own2;
4525 sl = None;
4526 }
4527 }
4528 let mut alloc_h;
4529 let mut alloc_x1;
4530 let mut alloc_z;
4531 let mut alloc_act;
4532 let mut alloc_h16;
4533 let mut alloc_z16;
4534 let mut alloc_gate;
4535 let mut alloc_up;
4536 let mut alloc_fo;
4537 let (h, x1, z, act): (
4538 &mut CudaSlice<f32>,
4539 &mut CudaSlice<f32>,
4540 &mut CudaSlice<f32>,
4541 &mut CudaSlice<f32>,
4542 );
4543 let (h16, z16): (&mut CudaSlice<u8>, &mut CudaSlice<u8>);
4544 let (sl_gate, sl_up, sl_fo): (
4545 &mut CudaSlice<f32>,
4546 &mut CudaSlice<f32>,
4547 &mut CudaSlice<f32>,
4548 );
4549 match sl {
4550 Some((a, b, c, d, e16, f16b, g, u, fo)) => {
4551 h = a;
4552 x1 = b;
4553 z = c;
4554 act = d;
4555 h16 = e16;
4556 z16 = f16b;
4557 sl_gate = g;
4558 sl_up = u;
4559 sl_fo = fo;
4560 }
4561 None => {
4562 alloc_h = e.uninit(t * n_embd)?;
4563 alloc_x1 = e.uninit(t * n_embd)?;
4564 alloc_z = e.uninit(t * n_embd)?;
4565 alloc_act = e.uninit(t * n_ff_max)?;
4566 alloc_h16 = e.alloc_u8_uninit(t * n_embd * 2)?;
4567 alloc_z16 = e.alloc_u8_uninit(t * n_embd * 2)?;
4568 alloc_gate = e.uninit(t * n_ff_max)?;
4569 alloc_up = e.uninit(t * n_ff_max)?;
4570 alloc_fo = e.uninit(t * n_embd)?;
4571 h = &mut alloc_h;
4572 x1 = &mut alloc_x1;
4573 z = &mut alloc_z;
4574 act = &mut alloc_act;
4575 h16 = &mut alloc_h16;
4576 z16 = &mut alloc_z16;
4577 sl_gate = &mut alloc_gate;
4578 sl_up = &mut alloc_up;
4579 sl_fo = &mut alloc_fo;
4580 }
4581 }
4582 let n_layers = self.layers.len();
4587 let use_seg = f16fuse
4597 && seg.is_some()
4598 && !self.uses_sliding_gated_moe_program()
4599 && lo == 0
4600 && hi == n_layers
4601 && std::env::var("MEMRA_PRIME_SEG").as_deref() == Ok("1")
4602 && {
4608 let ok = crate::spec::graph_launch_headroom_ok(e);
4609 if !ok {
4610 static NOTED: std::sync::Once = std::sync::Once::new();
4611 NOTED.call_once(|| crate::spec::graph_replay_suspended_note("prime-seg"));
4612 }
4613 ok
4614 };
4615 if let Some((sg, sm, _, st)) = seg.as_mut()
4616 && **st != t
4617 {
4618 sg.clear();
4619 sg.extend((0..n_layers).map(|_| None));
4620 sm.clear();
4621 sm.extend((0..n_layers).map(|_| None));
4622 **st = t;
4623 }
4624 {
4625 let layer_lo = &self.layers[lo];
4626 if f16fuse {
4627 e.rms_norm_f16out(
4628 x_cur,
4629 layer_lo.attn_norm.float_data(),
4630 h,
4631 h16,
4632 n_embd,
4633 t,
4634 eps,
4635 )?;
4636 } else {
4637 e.rms_norm(x_cur, layer_lo.attn_norm.float_data(), h, n_embd, t, eps)?;
4638 }
4639 }
4640 let anat = Self::prime_anatomy_on();
4641 let mut anat_last = if anat {
4642 e.stream().synchronize()?;
4643 Some(std::time::Instant::now())
4644 } else {
4645 None
4646 };
4647 macro_rules! anat_mark {
4649 ($slot:expr) => {
4650 if let Some(ts) = anat_last.as_mut() {
4651 e.stream().synchronize()?;
4652 Self::prime_anatomy_slots()[$slot].fetch_add(
4653 ts.elapsed().as_nanos() as u64,
4654 std::sync::atomic::Ordering::Relaxed,
4655 );
4656 *ts = std::time::Instant::now();
4657 }
4658 };
4659 }
4660 for il in lo..hi {
4661 let layer = &self.layers[il];
4662 let hx16 = if f16fuse { Some(&*h16) } else { None };
4663 if use_seg {
4664 let (pre, pre16, w_out) = match &layer.mixer {
4667 Mixer::Full(fa) => {
4668 let g3 = match hx16 {
4669 Some(xh) => e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], h, xh, t)?,
4670 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
4671 };
4672 let (pre, pre16) =
4673 self.full_attn_prime_core_inner(e, fa, g3, pos_d, t, cache, il)?;
4674 (pre, pre16, &fa.wo)
4675 }
4676 Mixer::Mla(_) => crate::hybrid::mla_path_unimplemented("core-split prime"),
4677 Mixer::Kda(_) => {
4678 crate::hybrid::kda_path_unimplemented("core-split captured prime")
4679 }
4680 Mixer::Linear(la) => {
4681 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
4682 let g4 = match hx16 {
4683 Some(xh) => e.matmul_group_xh(&ws, h, xh, t)?,
4684 None => e.matmul_group(&ws, h, t)?,
4685 };
4686 let (pre, pre16) =
4687 self.linear_attn_prime_core_pad_inner(e, la, g4, t, cache, il, None)?;
4688 (pre, pre16, &la.ssm_out)
4689 }
4690 };
4691 {
4692 let (_, sm, mslab, _) = seg.as_mut().unwrap();
4693 let pre_n = pre.len() / t;
4694 let xh_pre = match pre16 {
4695 Some(x) => x,
4696 None => e.f16_act(&pre, t * pre_n, pre_n)?,
4697 };
4698 if !e.try_f16_gemm_pre_into(w_out, &xh_pre, t, mslab)? {
4699 let y = e.matmul(w_out, &pre, t)?;
4700 e.copy_into(mslab, 0, &y, t * n_embd)?;
4701 }
4702 if sm[il].is_none() {
4703 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
4704 let w_post = layer.post_attn_norm.float_data();
4705 e.stream().synchronize()?;
4706 e.stream()
4707 .begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
4708 let r = (|| -> Result<(), Box<dyn std::error::Error>> {
4709 e.add(x_cur, mslab, x1, t * n_embd)?;
4710 e.rms_norm_f16out(x1, w_post, z, z16, n_embd, t, eps)?;
4711 Ok(())
4712 })();
4713 let g = e.stream().end_capture(
4714 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH);
4715 r?;
4716 sm[il] = Some(g?.ok_or("S-mid capture produced no graph")?);
4717 }
4718 sm[il].as_ref().unwrap().launch()?;
4719 }
4720 } else {
4721 let mixed = match &layer.mixer {
4722 Mixer::Full(fa) => {
4723 let y =
4724 self.full_attn_prime(e, fa, h, hx16, pos_d, t, cache, il, seq_end)?;
4725 anat_mark!(0);
4726 y
4727 }
4728 Mixer::Linear(la) => {
4729 let y = self.linear_attn_prime(e, la, h, hx16, t, cache, il)?;
4730 anat_mark!(1);
4731 y
4732 }
4733 Mixer::Mla(mla) => self.mla_attn_cached(e, mla, h, pos_d, t, il, cache)?,
4734 Mixer::Kda(la) => {
4735 let y = crate::kda::kda_prime_cached(e, la, h, t, eps, cache, il)?;
4736 anat_mark!(1);
4737 y
4738 }
4739 };
4740 if f16fuse {
4741 e.add_rms_norm_f16out(
4744 x_cur,
4745 &mixed,
4746 layer.post_attn_norm.float_data(),
4747 x1,
4748 z,
4749 z16,
4750 n_embd,
4751 t,
4752 eps,
4753 )?;
4754 } else {
4755 e.add(x_cur, &mixed, x1, t * n_embd)?;
4756 e.rms_norm(x1, layer.post_attn_norm.float_data(), z, n_embd, t, eps)?;
4757 }
4758 anat_mark!(4);
4759 }
4760 let zx16 = if f16fuse { Some(&*z16) } else { None };
4761 match &layer.ffn {
4762 crate::hybrid::Ffn::Dense {
4763 ffn_gate,
4764 ffn_up,
4765 ffn_down,
4766 } => {
4767 let n_ff = ffn_gate.out_features();
4768 let mut into_ok = false;
4771 if let Some(xh) = zx16 {
4772 into_ok = e.try_f16_gemm_pre_into(ffn_gate, xh, t, sl_gate)?
4773 && e.try_f16_gemm_pre_into(ffn_up, xh, t, sl_up)?;
4774 }
4775 if !into_ok {
4776 let mut g2 = match zx16 {
4777 Some(xh) => e.matmul_group_xh(&[ffn_gate, ffn_up], z, xh, t)?,
4778 None => e.matmul_group(&[ffn_gate, ffn_up], z, t)?,
4779 };
4780 let up_y = g2.pop().unwrap();
4781 let gate_y = g2.pop().unwrap();
4782 e.copy_into(sl_gate, 0, &gate_y, t * n_ff)?;
4783 e.copy_into(sl_up, 0, &up_y, t * n_ff)?;
4784 }
4785 let d_lim = self.cfg.clamp_shexp_at(il as u32);
4790 let act16 = if Self::f16out_on(e, t) && self.cfg.m3.is_none() && d_lim.is_none()
4791 {
4792 let mut a16 = e.alloc_u8_uninit(t * n_ff * 2)?;
4793 e.silu_mul_f16out(sl_gate, sl_up, act, &mut a16, t * n_ff)?;
4794 Some(a16)
4795 } else {
4796 Self::ffn_act_lim(
4797 e,
4798 &self.cfg,
4799 sl_gate,
4800 sl_up,
4801 1.0,
4802 1.0,
4803 d_lim,
4804 act,
4805 t * n_ff,
4806 )?;
4807 None
4808 };
4809 let xh_act = match act16 {
4811 Some(x) => x,
4812 None => e.f16_act(act, t * n_ff, n_ff)?,
4813 };
4814 if !e.try_f16_gemm_pre_into(ffn_down, &xh_act, t, sl_fo)? {
4815 let y = e.matmul(ffn_down, &*act, t)?;
4816 e.copy_into(sl_fo, 0, &y, t * n_embd)?;
4817 }
4818 }
4819 crate::hybrid::Ffn::Moe(m) => {
4820 let y = self.moe_ffn_il_prefill(e, m, z, t, il as u16)?;
4821 e.copy_into(sl_fo, 0, &y, t * n_embd)?;
4822 anat_mark!(2);
4823 }
4824 }
4825 if let (crate::hybrid::Ffn::Dense { .. }, true) = (&layer.ffn, anat) {
4826 anat_mark!(3);
4827 }
4828 if use_seg && il + 1 < hi {
4829 let w_next = self.layers[il + 1].attn_norm.float_data();
4831 let (sg, _, _, _) = seg.as_mut().unwrap();
4832 if sg[il].is_none() {
4833 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
4834 e.stream().synchronize()?;
4835 e.stream()
4836 .begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
4837 let r = (|| -> Result<(), Box<dyn std::error::Error>> {
4838 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
4839 e.rms_norm_f16out(x_nxt, w_next, h, h16, n_embd, t, eps)?;
4840 Ok(())
4841 })();
4842 let g = e.stream().end_capture(
4843 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
4844 );
4845 r?;
4846 sg[il] = Some(g?.ok_or("S-glue capture produced no graph")?);
4847 }
4848 sg[il].as_ref().unwrap().launch()?;
4849 } else {
4850 if il + 1 < hi {
4851 let w_next = self.layers[il + 1].attn_norm.float_data();
4852 if f16fuse {
4853 e.add_rms_norm_f16out(x1, sl_fo, w_next, x_nxt, h, h16, n_embd, t, eps)?;
4854 } else {
4855 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
4856 e.rms_norm(x_nxt, w_next, h, n_embd, t, eps)?;
4857 }
4858 } else {
4859 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
4860 }
4861 }
4862 anat_mark!(4);
4863 if let Some(path) = Self::prime_trace_path() {
4869 let row = base + t - 1;
4870 let host = e.dtoh(x_nxt)?;
4871 let last = &host[(t - 1) * n_embd..t * n_embd];
4872 use std::io::Write as _;
4873 let mut f = std::fs::OpenOptions::new()
4874 .create(true)
4875 .append(true)
4876 .open(path)?;
4877 let mut h64: u64 = 0xcbf29ce484222325;
4878 for v in last {
4879 h64 ^= v.to_bits() as u64;
4880 h64 = h64.wrapping_mul(0x100000001b3);
4881 }
4882 writeln!(
4883 f,
4884 "{{\"pos\":{row},\"layer\":{il},\"t\":{t},\"base\":{base},\
4885 \"hash\":\"{h64:016x}\",\"v0\":{:.9e},\"v1\":{:.9e},\"v2\":{:.9e}}}",
4886 last[0], last[1], last[2]
4887 )?;
4888 }
4889 self.dflash_tap(e, cache, il, x_nxt, t)?;
4892 std::mem::swap(&mut x_cur, &mut x_nxt);
4893 }
4894 if anat {
4895 let s = Self::prime_anatomy_slots();
4896 let ms = |i: usize| s[i].load(std::sync::atomic::Ordering::Relaxed) as f64 / 1.0e6;
4897 eprintln!(
4898 "[prime-anatomy] cumulative ms: attn_full={:.1} gdn_linear={:.1} moe={:.1} \
4899 dense={:.1} norms_adds={:.1} (t={t}, layers {lo}..{hi})",
4900 ms(0),
4901 ms(1),
4902 ms(2),
4903 ms(3),
4904 ms(4)
4905 );
4906 }
4907 let mut x = e.uninit(t * n_embd)?;
4909 e.copy_into(&mut x, 0, x_cur, t * n_embd)?;
4910 drop(slab_guard);
4911 Ok(x)
4912 }
4913
4914 #[allow(clippy::type_complexity)] fn prime_chunk_epilogue(
4920 &self,
4921 e: &Engine,
4922 x: CudaSlice<f32>,
4923 t: usize,
4924 cache: &mut Cache,
4925 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4926 let n_embd = self.cfg.n_embd as usize;
4927 let eps = self.cfg.rms_eps;
4928 let mut h_seed = e.uninit(n_embd)?;
4932 if !crate::spec::spec_hpost() {
4933 e.copy_view_into(
4934 &mut h_seed,
4935 0,
4936 &x.slice((t - 1) * n_embd..t * n_embd),
4937 n_embd,
4938 )?;
4939 }
4940 let mut hn = e.uninit(t * n_embd)?;
4942 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
4943 if crate::spec::spec_hpost() {
4944 e.copy_view_into(
4945 &mut h_seed,
4946 0,
4947 &hn.slice((t - 1) * n_embd..t * n_embd),
4948 n_embd,
4949 )?;
4950 }
4951 let last = e.view(&hn, t * n_embd);
4952 let last_row = last.slice((t - 1) * n_embd..t * n_embd);
4953 let mut hlast = e.uninit(n_embd)?;
4954 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
4955 let logits = e.matmul(&self.output, &hlast, 1)?;
4956 cache.pos += t;
4957 Ok((
4960 e.dtoh(&logits)?,
4961 h_seed,
4962 if crate::spec::spec_hpost() { hn } else { x },
4963 ))
4964 }
4965
4966 pub fn hidden_postnorm_row(
4972 &self,
4973 e: &Engine,
4974 hiddens: &CudaSlice<f32>,
4975 row: usize,
4976 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
4977 let n_embd = self.cfg.n_embd as usize;
4978 let mut x1 = e.uninit(n_embd)?;
4979 e.copy_view_into(
4980 &mut x1,
4981 0,
4982 &hiddens.slice(row * n_embd..(row + 1) * n_embd),
4983 n_embd,
4984 )?;
4985 if crate::spec::spec_hpost() {
4986 return e.dtoh(&x1);
4987 }
4988 let mut hn = e.uninit(n_embd)?;
4989 e.rms_norm(
4990 &x1,
4991 self.output_norm.float_data(),
4992 &mut hn,
4993 n_embd,
4994 1,
4995 self.cfg.rms_eps,
4996 )?;
4997 e.dtoh(&hn)
4998 }
4999
5000 #[allow(clippy::type_complexity)] fn prime_chunk_ppn(
5025 &self,
5026 e: &Engine,
5027 tokens: &[u32],
5028 cache: &mut Cache,
5029 seq_end: usize,
5030 fence: &[usize],
5031 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5032 let rt = crate::pp::PpNRt::get(e)?;
5033 let n_st = fence.len() - 1;
5034 assert_eq!(
5035 rt.n_stages(),
5036 n_st,
5037 "PpNRt stage count {} != fence stages {n_st}",
5038 rt.n_stages()
5039 );
5040 let n_embd = self.cfg.n_embd as usize;
5041 let t = tokens.len();
5042 let base = cache.pos;
5043 debug_assert!(
5044 seq_end >= base + t,
5045 "prime_chunk_ppn: seq_end must cover this chunk"
5046 );
5047 let payload = t * n_embd;
5048 let caller_stream = e.stream();
5052 rt.fence_stages_behind(&caller_stream)?;
5053
5054 if n_st == 2 {
5055 let slot =
5056 self.prime_pp2_stage0_enqueue(e, rt, tokens, cache, seq_end, fence, base, false)?;
5057 let x =
5058 self.prime_pp2_stage1_enqueue(e, rt, slot, t, cache, seq_end, fence, base, false)?;
5059 let out = {
5060 rt.bind_stage(1)?;
5061 let _st1 = rt.enter(1);
5062 let e1 = rt.engine(1, e);
5063 self.prime_chunk_epilogue(e1, x, t, cache)?
5064 };
5065 rt.publish_to(1, &caller_stream)?;
5066 crate::pp::PRIME_SPLIT_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
5067 return Ok(out);
5068 }
5069
5070 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
5071
5072 let mut slot = {
5074 let _st0 = rt.enter(0);
5075 let e0 = rt.engine(0, e);
5076 let pos_d = e0.htod_i32(&pos)?;
5077 let x = self.embed(e0, tokens)?;
5078 let x =
5079 self.prime_layers(e0, x, fence[0], fence[1], &pos_d, t, base, cache, seq_end)?;
5080 rt.tx(0, &x, payload)?
5081 };
5083
5084 for s in 1..n_st - 1 {
5086 let _st = rt.enter(s);
5087 let es = rt.engine(s, e);
5088 let pos_d = es.htod_i32(&pos)?;
5089 let x = rt.rx(s - 1, slot, payload)?;
5090 let x = self.prime_layers(
5091 es,
5092 x,
5093 fence[s],
5094 fence[s + 1],
5095 &pos_d,
5096 t,
5097 base,
5098 cache,
5099 seq_end,
5100 )?;
5101 slot = rt.tx(s, &x, payload)?;
5102 }
5103
5104 let _stl = rt.enter(n_st - 1);
5106 let el = rt.engine(n_st - 1, e);
5107 let pos_d = el.htod_i32(&pos)?;
5108 let x = rt.rx(n_st - 2, slot, payload)?;
5109 let x = self.prime_layers(
5110 el,
5111 x,
5112 fence[n_st - 1],
5113 fence[n_st],
5114 &pos_d,
5115 t,
5116 base,
5117 cache,
5118 seq_end,
5119 )?;
5120 let out = self.prime_chunk_epilogue(el, x, t, cache)?;
5121 rt.publish_to(n_st - 1, &caller_stream)?;
5127 crate::pp::PRIME_SPLIT_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
5128 Ok(out)
5129 }
5130
5131 #[allow(clippy::too_many_arguments)] fn prime_pp2_stage0_enqueue(
5133 &self,
5134 e: &Engine,
5135 rt: &crate::pp::PpNRt,
5136 tokens: &[u32],
5137 cache: &mut Cache,
5138 seq_end: usize,
5139 fence: &[usize],
5140 base: usize,
5141 pipelined: bool,
5142 ) -> Result<usize, Box<dyn std::error::Error>> {
5143 let t = tokens.len();
5144 let n_embd = self.cfg.n_embd as usize;
5145 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
5146 rt.bind_stage(0)?;
5147 let _st0 = rt.enter(0);
5148 let e0 = rt.engine(0, e);
5149 let pos_d = e0.htod_i32(&pos)?;
5150 let x = self.embed(e0, tokens)?;
5151 let _overlap = pipelined.then(crate::pp::enter_prime_pipe_stage);
5152 let x = self.prime_layers(e0, x, fence[0], fence[1], &pos_d, t, base, cache, seq_end)?;
5153 if pipelined {
5154 rt.tx_pipelined(0, &x, t * n_embd)
5155 } else {
5156 rt.tx(0, &x, t * n_embd)
5157 }
5158 }
5159
5160 #[allow(clippy::too_many_arguments)] fn prime_pp2_stage1_enqueue(
5162 &self,
5163 e: &Engine,
5164 rt: &crate::pp::PpNRt,
5165 slot: usize,
5166 t: usize,
5167 cache: &mut Cache,
5168 seq_end: usize,
5169 fence: &[usize],
5170 base: usize,
5171 pipelined: bool,
5172 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5173 let n_embd = self.cfg.n_embd as usize;
5174 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
5175 rt.bind_stage(1)?;
5176 let _st1 = rt.enter(1);
5177 let e1 = rt.engine(1, e);
5178 let pos_d = e1.htod_i32(&pos)?;
5179 let x = rt.rx(0, slot, t * n_embd)?;
5180 let _overlap = pipelined.then(crate::pp::enter_prime_pipe_stage);
5181 self.prime_layers(e1, x, fence[1], fence[2], &pos_d, t, base, cache, seq_end)
5182 }
5183
5184 #[allow(clippy::too_many_arguments)] pub fn prime_chunk_captured(
5201 &self,
5202 e: &Engine,
5203 x_in: &CudaSlice<f32>,
5204 pos_d: &CudaSlice<i32>,
5205 t: usize,
5206 cache: &mut Cache,
5207 len_d: &CudaSlice<i32>,
5208 logits_out: &mut CudaSlice<f32>,
5209 h_seed_out: &mut CudaSlice<f32>,
5210 ) -> Result<(), Box<dyn std::error::Error>> {
5211 self.refuse_hyper("prime_chunk_captured")?;
5212 cache.ensure_usable("prime_chunk_captured")?;
5213 let cfg = &self.cfg;
5214 let n_embd = cfg.n_embd as usize;
5215 let eps = cfg.rms_eps;
5216 let f16fuse = crate::f16_ffi::pp_f16_enabled() && t >= 16;
5217 let mut x = e.uninit(t * n_embd)?;
5218 e.copy_into(&mut x, 0, x_in, t * n_embd)?;
5219 for (il, layer) in self.layers.iter().enumerate() {
5220 let mut h = e.uninit(t * n_embd)?;
5221 let mut hx16: Option<CudaSlice<u8>> = None;
5222 if f16fuse {
5223 let mut b16 = e.alloc_u8_uninit(t * n_embd * 2)?;
5224 e.rms_norm_f16out(
5225 &x,
5226 layer.attn_norm.float_data(),
5227 &mut h,
5228 &mut b16,
5229 n_embd,
5230 t,
5231 eps,
5232 )?;
5233 hx16 = Some(b16);
5234 } else {
5235 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
5236 }
5237 let mixed = match &layer.mixer {
5238 Mixer::Full(fa) => {
5242 self.full_attn_prime(e, fa, &h, hx16.as_ref(), pos_d, t, cache, il, t)?
5243 }
5244 Mixer::Mla(_) => crate::hybrid::mla_path_unimplemented("captured-graph prime"),
5245 Mixer::Kda(_) => crate::hybrid::kda_path_unimplemented("captured prime chunk"),
5246 Mixer::Linear(la) => {
5247 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
5248 let g4 = match hx16.as_ref() {
5249 Some(xh) => e.matmul_group_xh(&ws, &h, xh, t)?,
5250 None => e.matmul_group(&ws, &h, t)?,
5251 };
5252 self.linear_attn_prime_core_pad(e, la, g4, t, cache, il, Some(len_d))?
5253 }
5254 };
5255 let mut x1 = e.uninit(t * n_embd)?;
5256 e.add(&x, &mixed, &mut x1, t * n_embd)?;
5257 let mut z = e.uninit(t * n_embd)?;
5258 let mut zx16: Option<CudaSlice<u8>> = None;
5259 if f16fuse {
5260 let mut b16 = e.alloc_u8_uninit(t * n_embd * 2)?;
5261 e.rms_norm_f16out(
5262 &x1,
5263 layer.post_attn_norm.float_data(),
5264 &mut z,
5265 &mut b16,
5266 n_embd,
5267 t,
5268 eps,
5269 )?;
5270 zx16 = Some(b16);
5271 } else {
5272 e.rms_norm(
5273 &x1,
5274 layer.post_attn_norm.float_data(),
5275 &mut z,
5276 n_embd,
5277 t,
5278 eps,
5279 )?;
5280 }
5281 let ffn_out = match &layer.ffn {
5282 crate::hybrid::Ffn::Dense {
5283 ffn_gate,
5284 ffn_up,
5285 ffn_down,
5286 } => {
5287 let n_ff = ffn_gate.out_features();
5288 let mut g2 = match &zx16 {
5289 Some(xh) => e.matmul_group_xh(&[ffn_gate, ffn_up], &z, xh, t)?,
5290 None => e.matmul_group(&[ffn_gate, ffn_up], &z, t)?,
5291 };
5292 let up = g2.pop().unwrap();
5293 let gate = g2.pop().unwrap();
5294 let mut act = e.uninit(t * n_ff)?;
5295 Self::ffn_act_lim(
5297 e,
5298 &self.cfg,
5299 &gate,
5300 &up,
5301 1.0,
5302 1.0,
5303 self.cfg.clamp_shexp_at(il as u32),
5304 &mut act,
5305 t * n_ff,
5306 )?;
5307 e.matmul(ffn_down, &act, t)?
5308 }
5309 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il_prefill(e, m, &z, t, il as u16)?,
5310 };
5311 let mut x2 = e.uninit(t * n_embd)?;
5312 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
5313 x = x2;
5314 }
5315 if !crate::spec::spec_hpost() {
5317 e.row_gather_dev(&x, h_seed_out, len_d, n_embd)?;
5318 }
5319 let mut hn = e.uninit(t * n_embd)?;
5320 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
5321 if crate::spec::spec_hpost() {
5322 e.row_gather_dev(&hn, h_seed_out, len_d, n_embd)?;
5323 }
5324 let mut hlast = e.uninit(n_embd)?;
5325 e.row_gather_dev(&hn, &mut hlast, len_d, n_embd)?;
5326 let logits = e.matmul(&self.output, &hlast, 1)?;
5327 let nv = logits.len();
5328 e.copy_into(logits_out, 0, &logits, nv)?;
5329 Ok(())
5330 }
5331
5332 fn step35_prime_batch_on() -> bool {
5333 std::env::var("MEMRA_STEP35_PRIME_BATCH").as_deref() != Ok("0")
5334 }
5335
5336 #[allow(clippy::too_many_arguments)]
5339 #[allow(clippy::too_many_arguments)]
5344 fn step35_prime_batch_layers(
5345 &self,
5346 e: &Engine,
5347 mut x: CudaSlice<f32>,
5348 lo: usize,
5349 hi: usize,
5350 ts: &[usize],
5351 offs: &[usize],
5352 seq_ends: &[usize],
5353 pos_ds: &[CudaSlice<i32>],
5354 caches: &mut [&mut Cache],
5355 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5356 let cfg = &self.cfg;
5357 let n_embd = cfg.n_embd as usize;
5358 let eps = cfg.rms_eps;
5359 let b = ts.len();
5360 let total: usize = ts.iter().sum();
5361 let f16fuse = crate::f16_ffi::pp_f16_enabled() && total >= 16;
5362
5363 let split = |e: &Engine,
5364 y: &CudaSlice<f32>,
5365 dim: usize|
5366 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
5367 let mut out = Vec::with_capacity(b);
5368 for s in 0..b {
5369 let mut ys = e.uninit(ts[s] * dim)?;
5370 e.copy_view_into(
5371 &mut ys,
5372 0,
5373 &y.slice(offs[s] * dim..(offs[s] + ts[s]) * dim),
5374 ts[s] * dim,
5375 )?;
5376 out.push(ys);
5377 }
5378 Ok(out)
5379 };
5380
5381 let prof = std::env::var("MEMRA_PRIME_PROF").as_deref() == Ok("1");
5386 let mut ph = [0f64; 4]; let mark = |e: &Engine, acc: usize, t0: &mut std::time::Instant, ph: &mut [f64; 4]| {
5388 if prof {
5389 let _ = e.stream().synchronize();
5390 ph[acc] += t0.elapsed().as_secs_f64() * 1e3;
5391 *t0 = std::time::Instant::now();
5392 }
5393 };
5394 let mut pt = std::time::Instant::now();
5395 for il in lo..hi {
5396 let layer = &self.layers[il];
5397 let Mixer::Full(fa) = &layer.mixer else {
5398 return Err(format!("step35 layer {il} is not full-attn — corrupt config").into());
5399 };
5400
5401 let mut h = e.uninit(total * n_embd)?;
5402 let mut hx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
5403 if f16fuse {
5404 e.rms_norm_f16out(
5405 &x,
5406 layer.attn_norm.float_data(),
5407 &mut h,
5408 &mut hx16,
5409 n_embd,
5410 total,
5411 eps,
5412 )?;
5413 } else {
5414 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, total, eps)?;
5415 }
5416
5417 let gate_w = fa
5421 .attn_gate
5422 .as_ref()
5423 .ok_or("step35 layer is missing attn_gate.weight (head-wise attention gate)")?;
5424 let mut g4 = if f16fuse {
5425 e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv, gate_w], &h, &hx16, total)?
5426 } else {
5427 e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv, gate_w], &h, total)?
5428 };
5429 let gate = g4.pop().unwrap();
5430 let mut parts: Vec<Vec<CudaSlice<f32>>> =
5431 (0..b).map(|_| Vec::with_capacity(3)).collect();
5432 for (w, y) in [&fa.wq, &fa.wk, &fa.wv].iter().zip(g4) {
5433 for (s, ys) in split(e, &y, w.out_features())?.into_iter().enumerate() {
5434 parts[s].push(ys);
5435 }
5436 }
5437 let gates = split(e, &gate, gate_w.out_features())?;
5438 let geometry = self.step35_geom(il);
5439 let hd = geometry.head_dim_k as usize;
5440 let nh = geometry.n_head as usize;
5441 let mut ag_cat = e.uninit(total * nh * hd)?;
5442 for (s, (g3s, gate)) in parts.into_iter().zip(gates).enumerate() {
5443 mark(e, 0, &mut pt, &mut ph);
5444 let ag = self.step35_attn_pre_wo(
5445 e,
5446 fa,
5447 g3s,
5448 None,
5449 Some(&gate),
5450 &pos_ds[s],
5451 ts[s],
5452 Some(&mut *caches[s]),
5453 il,
5454 seq_ends[s],
5455 )?;
5456 e.copy_into(&mut ag_cat, offs[s] * nh * hd, &ag, ts[s] * nh * hd)?;
5457 }
5458 mark(e, 1, &mut pt, &mut ph);
5459 let mixed = e.matmul(&fa.wo, &ag_cat, total)?;
5460
5461 let mut x1 = e.uninit(total * n_embd)?;
5462 let mut z = e.uninit(total * n_embd)?;
5463 let mut zx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
5464 if f16fuse {
5465 e.add_rms_norm_f16out(
5466 &x,
5467 &mixed,
5468 layer.post_attn_norm.float_data(),
5469 &mut x1,
5470 &mut z,
5471 &mut zx16,
5472 n_embd,
5473 total,
5474 eps,
5475 )?;
5476 } else {
5477 e.add(&x, &mixed, &mut x1, total * n_embd)?;
5478 e.rms_norm(
5479 &x1,
5480 layer.post_attn_norm.float_data(),
5481 &mut z,
5482 n_embd,
5483 total,
5484 eps,
5485 )?;
5486 }
5487
5488 mark(e, 2, &mut pt, &mut ph);
5489 let ffn_out = match &layer.ffn {
5490 crate::hybrid::Ffn::Dense {
5491 ffn_gate,
5492 ffn_up,
5493 ffn_down,
5494 } => {
5495 let n_ff = ffn_gate.out_features();
5496 let mut g2 = if f16fuse {
5497 e.matmul_group_xh(&[ffn_gate, ffn_up], &z, &zx16, total)?
5498 } else {
5499 e.matmul_group(&[ffn_gate, ffn_up], &z, total)?
5500 };
5501 let up = g2.pop().unwrap();
5502 let gate = g2.pop().unwrap();
5503 let mut act = e.uninit(total * n_ff)?;
5504 let d_lim = cfg.clamp_shexp_at(il as u32);
5505 if Self::f16out_on(e, total) && cfg.m3.is_none() && d_lim.is_none() {
5506 let mut a16 = e.alloc_u8_uninit(total * n_ff * 2)?;
5507 e.silu_mul_f16out(&gate, &up, &mut act, &mut a16, total * n_ff)?;
5508 match e.try_f16_gemm_pre(ffn_down, &a16, total)? {
5509 Some(y) => y,
5510 None => e.matmul(ffn_down, &act, total)?,
5511 }
5512 } else {
5513 Self::ffn_act_lim(
5514 e,
5515 cfg,
5516 &gate,
5517 &up,
5518 1.0,
5519 1.0,
5520 d_lim,
5521 &mut act,
5522 total * n_ff,
5523 )?;
5524 e.matmul(ffn_down, &act, total)?
5525 }
5526 }
5527 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il(e, m, &z, total, il as u16)?,
5528 };
5529 let mut x2 = e.uninit(total * n_embd)?;
5530 e.add(&x1, &ffn_out, &mut x2, total * n_embd)?;
5531 x = x2;
5532 mark(e, 3, &mut pt, &mut ph);
5533 }
5534 if prof {
5535 eprintln!(
5536 "[prime-prof] t={total} layers={} norm+qkv={:.0}ms attn={:.0}ms o_proj={:.0}ms moe={:.0}ms",
5537 hi - lo,
5538 ph[0],
5539 ph[1],
5540 ph[2],
5541 ph[3]
5542 );
5543 }
5544 Ok(x)
5545 }
5546
5547 #[allow(clippy::type_complexity)] fn step35_prime_batch_epilogue(
5549 &self,
5550 e: &Engine,
5551 x: CudaSlice<f32>,
5552 ts: &[usize],
5553 offs: &[usize],
5554 caches: &mut [&mut Cache],
5555 ) -> Result<Vec<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
5556 let n_embd = self.cfg.n_embd as usize;
5557 let total: usize = ts.iter().sum();
5558 let mut hn = e.uninit(total * n_embd)?;
5559 e.rms_norm(
5560 &x,
5561 self.output_norm.float_data(),
5562 &mut hn,
5563 n_embd,
5564 total,
5565 self.cfg.rms_eps,
5566 )?;
5567
5568 let hidden_src = if crate::spec::spec_hpost() { &hn } else { &x };
5569 let mut out = Vec::with_capacity(ts.len());
5570 for s in 0..ts.len() {
5571 let mut hidden = e.uninit(ts[s] * n_embd)?;
5572 e.copy_view_into(
5573 &mut hidden,
5574 0,
5575 &hidden_src.slice(offs[s] * n_embd..(offs[s] + ts[s]) * n_embd),
5576 ts[s] * n_embd,
5577 )?;
5578 let last0 = (offs[s] + ts[s] - 1) * n_embd;
5579 let mut h_seed = e.uninit(n_embd)?;
5580 e.copy_view_into(
5581 &mut h_seed,
5582 0,
5583 &hidden_src.slice(last0..last0 + n_embd),
5584 n_embd,
5585 )?;
5586 let mut hlast = e.uninit(n_embd)?;
5588 e.copy_view_into(&mut hlast, 0, &hn.slice(last0..last0 + n_embd), n_embd)?;
5589 let logits = e.dtoh(&e.matmul(&self.output, &hlast, 1)?)?;
5590 caches[s].pos += ts[s];
5591 out.push((logits, h_seed, hidden));
5592 }
5593 Ok(out)
5594 }
5595
5596 #[allow(clippy::type_complexity)] fn step35_prime_cache_batch(
5602 &self,
5603 e: &Engine,
5604 prompts: &[&[u32]],
5605 caches: &mut [&mut Cache],
5606 seq_ends: &[usize],
5607 ) -> Result<Vec<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
5608 assert_eq!(
5609 seq_ends.len(),
5610 prompts.len(),
5611 "step35 batched prime: one seq_end per sequence"
5612 );
5613 validate_step_prime_batch_modes(
5614 step_tp_prefill_enabled()?,
5615 step_ep_grouped_prefill_enabled()?,
5616 )?;
5617 if crate::pp::pp_host_bounce_active()
5618 && (!crate::pp::prime_pp_on() || crate::pp::pp_cuts(self.layers.len()).is_none())
5619 {
5620 return Err(
5621 "step35_prime_cache_batch: MEMRA_PP_HOST_BOUNCE=1 requires a valid prime \
5622 stage split; refusing an unsplit remote-weight walk"
5623 .into(),
5624 );
5625 }
5626 if !Self::step35_prime_batch_on() {
5627 return Err("step35 batched prime is disabled (MEMRA_STEP35_PRIME_BATCH=0)".into());
5628 }
5629 if prompts.len() > 1 && caches.iter().any(|c| c.pos != 0) {
5634 return Err(
5635 "step35 batched prime supports continuation only at B=1; a cross-request batch \
5636 at mixed positions requires per-request queued_after"
5637 .into(),
5638 );
5639 }
5640
5641 let ts: Vec<usize> = prompts.iter().map(|p| p.len()).collect();
5642 for &t in &ts {
5643 assert!(
5644 t >= PRIME_MIN_T,
5645 "step35 batched prime needs T >= {PRIME_MIN_T}"
5646 );
5647 }
5648 for (s, c) in caches.iter().enumerate() {
5649 assert!(
5652 c.pos + ts[s] <= c.max_ctx,
5653 "step35 batched prime exceeds cache max_ctx"
5654 );
5655 assert!(
5656 seq_ends[s] >= c.pos + ts[s],
5657 "step35 batched prime: seq_end must cover this chunk"
5658 );
5659 }
5660 let mut transaction = CacheTaintGuard::arm(caches);
5661 let legacy_tsend = std::env::var("MEMRA_STEP35_PRIME_BATCH_TSEND").as_deref() == Ok("1");
5668 let seq_ends_eff: Vec<usize> = if legacy_tsend {
5669 ts.clone()
5670 } else {
5671 seq_ends.to_vec()
5672 };
5673 let offs: Vec<usize> = ts
5674 .iter()
5675 .scan(0usize, |a, &t| {
5676 let o = *a;
5677 *a += t;
5678 Some(o)
5679 })
5680 .collect();
5681 let total: usize = ts.iter().sum();
5682 let payload = total * self.cfg.n_embd as usize;
5683 let cat_tokens: Vec<u32> = prompts.iter().flat_map(|p| p.iter().copied()).collect();
5684 let positions: Vec<Vec<i32>> = ts
5690 .iter()
5691 .zip(caches.iter())
5692 .map(|(&t, c)| {
5693 let base = c.pos as i32;
5694 (0..t as i32).map(|i| base + i).collect()
5695 })
5696 .collect();
5697 let upload_positions =
5698 |e: &Engine| -> Result<Vec<CudaSlice<i32>>, Box<dyn std::error::Error>> {
5699 positions
5700 .iter()
5701 .map(|p| e.htod_i32(p))
5702 .collect::<Result<_, _>>()
5703 };
5704
5705 static ONCE: std::sync::Once = std::sync::Once::new();
5706 ONCE.call_once(|| {
5707 eprintln!(
5708 "[step35-prime-batch] first concat prime: B={} tokens={total}",
5709 prompts.len()
5710 );
5711 });
5712
5713 let out = if !crate::pp::pp2_streams_off() && crate::pp::prime_pp_on() {
5714 if let Some(fence) = crate::pp::pp_cuts(self.layers.len()) {
5715 let rt = crate::pp::PpNRt::get(e)?;
5716 let n_st = fence.len() - 1;
5717 assert_eq!(
5718 rt.n_stages(),
5719 n_st,
5720 "step35 prime batch stage count mismatch"
5721 );
5722 let caller_stream = e.stream();
5723 rt.fence_stages_behind(&caller_stream)?;
5724
5725 let mut slot = {
5726 let _st0 = rt.enter(0);
5727 let e0 = rt.engine(0, e);
5728 let pos_ds = upload_positions(e0)?;
5729 let x = self.embed(e0, &cat_tokens)?;
5730 let x = self.step35_prime_batch_layers(
5731 e0,
5732 x,
5733 fence[0],
5734 fence[1],
5735 &ts,
5736 &offs,
5737 &seq_ends_eff,
5738 &pos_ds,
5739 caches,
5740 )?;
5741 rt.tx(0, &x, payload)?
5742 };
5743 for s in 1..n_st - 1 {
5744 let _st = rt.enter(s);
5745 let es = rt.engine(s, e);
5746 let pos_ds = upload_positions(es)?;
5747 let x = rt.rx(s - 1, slot, payload)?;
5748 let x = self.step35_prime_batch_layers(
5749 es,
5750 x,
5751 fence[s],
5752 fence[s + 1],
5753 &ts,
5754 &offs,
5755 &seq_ends_eff,
5756 &pos_ds,
5757 caches,
5758 )?;
5759 slot = rt.tx(s, &x, payload)?;
5760 }
5761
5762 let _stl = rt.enter(n_st - 1);
5763 let el = rt.engine(n_st - 1, e);
5764 let pos_ds = upload_positions(el)?;
5765 let x = rt.rx(n_st - 2, slot, payload)?;
5766 let x = self.step35_prime_batch_layers(
5767 el,
5768 x,
5769 fence[n_st - 1],
5770 fence[n_st],
5771 &ts,
5772 &offs,
5773 &seq_ends_eff,
5774 &pos_ds,
5775 caches,
5776 )?;
5777 let out = self.step35_prime_batch_epilogue(el, x, &ts, &offs, caches)?;
5778 rt.publish_to(n_st - 1, &caller_stream)?;
5779 crate::pp::STEP35_PRIME_BATCH_SPLITS
5780 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
5781 out
5782 } else {
5783 let pos_ds = upload_positions(e)?;
5784 let x = self.embed(e, &cat_tokens)?;
5785 let x = self.step35_prime_batch_layers(
5786 e,
5787 x,
5788 0,
5789 self.layers.len(),
5790 &ts,
5791 &offs,
5792 &seq_ends_eff,
5793 &pos_ds,
5794 caches,
5795 )?;
5796 self.step35_prime_batch_epilogue(e, x, &ts, &offs, caches)?
5797 }
5798 } else {
5799 let pos_ds = upload_positions(e)?;
5800 let x = self.embed(e, &cat_tokens)?;
5801 let x = self.step35_prime_batch_layers(
5802 e,
5803 x,
5804 0,
5805 self.layers.len(),
5806 &ts,
5807 &offs,
5808 &seq_ends_eff,
5809 &pos_ds,
5810 caches,
5811 )?;
5812 self.step35_prime_batch_epilogue(e, x, &ts, &offs, caches)?
5813 };
5814 crate::pp::STEP35_PRIME_BATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
5815 transaction.commit();
5816 Ok(out)
5817 }
5818
5819 #[allow(clippy::type_complexity)] pub fn prime_cache_batch(
5837 &self,
5838 e: &Engine,
5839 prompts: &[&[u32]],
5840 caches: &mut [&mut Cache],
5841 ) -> Result<Vec<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
5842 self.refuse_hyper("prime_cache_batch")?;
5843 for cache in caches.iter() {
5844 cache.ensure_usable("prime_cache_batch")?;
5845 }
5846 if crate::pp::pp_cuts(self.layers.len()).is_some()
5847 && !self.rewrite_allowed(memra_gguf::execution_manifest::RewriteSurface::Pipeline)
5848 {
5849 return Err("pipeline rewrite is not qualified for batched prime".into());
5850 }
5851 if !self.rewrite_allowed(memra_gguf::execution_manifest::RewriteSurface::CarriedPrime) {
5852 if !self.rewrite_allowed(memra_gguf::execution_manifest::RewriteSurface::DecodeEager) {
5853 return Err("neither batched-prime nor eager rewrite is qualified".into());
5854 }
5855 if prompts.len() != caches.len() {
5856 return Err("prime fallback prompt/cache shape mismatch".into());
5857 }
5858 static ONCE: std::sync::Once = std::sync::Once::new();
5859 ONCE.call_once(|| {
5860 eprintln!(
5861 "[rewrite] carried-prime.v1 unqualified; using individual native eager primes"
5862 );
5863 });
5864 let mut transaction = CacheTaintGuard::arm(caches);
5865 let result: Result<Vec<_>, Box<dyn std::error::Error>> = prompts
5866 .iter()
5867 .copied()
5868 .zip(caches.iter_mut())
5869 .map(|(prompt, cache)| self.prime_cache(e, prompt, cache, 0))
5870 .collect();
5871 if result.is_ok() {
5872 transaction.commit();
5873 }
5874 return result;
5875 }
5876 let _pp_walk =
5877 if crate::pp::pp_cuts(self.layers.len()).is_some() && !crate::pp::pp2_streams_off() {
5878 let rt = crate::pp::PpNRt::get(e)?;
5879 Some(rt.acquire_walk("prime_cache_batch")?)
5880 } else {
5881 None
5882 };
5883 let cfg = &self.cfg;
5884 let n_embd = cfg.n_embd as usize;
5885 let eps = cfg.rms_eps;
5886 let b = prompts.len();
5887 assert!(b >= 1 && b == caches.len());
5888 let pos0s: Vec<usize> = caches.iter().map(|c| c.pos).collect();
5889 let carried = pos0s.iter().any(|&p| p > 0);
5890 if self.uses_gemma_program() {
5896 return Err(
5897 "prime_cache_batch: gemma4 has no batched prime core (per-layer \
5898 swa/global geometry, softcapped head) — use gemma4_prime per sequence"
5899 .into(),
5900 );
5901 }
5902 if self.uses_sliding_gated_moe_program() {
5905 let seq_ends: Vec<usize> = caches
5910 .iter()
5911 .zip(prompts.iter())
5912 .map(|(c, p)| c.pos + p.len())
5913 .collect();
5914 return self.step35_prime_cache_batch(e, prompts, caches, &seq_ends);
5915 }
5916 if crate::pp::pp_cuts(self.layers.len()).is_some() && !crate::pp::pp2_streams_off() {
5917 let rt = crate::pp::PpNRt::get(e)?;
5918 if rt.cross_device() {
5919 return Err(
5920 "prime_cache_batch: generic dense concat prime has no cross-device PP split; use individual prime_cache calls"
5921 .into(),
5922 );
5923 }
5924 }
5925 let ts: Vec<usize> = prompts.iter().map(|p| p.len()).collect();
5926 for &t in &ts {
5927 assert!(
5928 t >= PRIME_MIN_T,
5929 "prime_cache_batch needs T >= {PRIME_MIN_T}"
5930 );
5931 }
5932 for (s, c) in caches.iter().enumerate() {
5933 assert!(
5934 c.pos + ts[s] <= c.max_ctx,
5935 "prime_cache_batch: prompt exceeds cache max_ctx"
5936 );
5937 }
5938 let mut transaction = CacheTaintGuard::arm(caches);
5939 let total: usize = ts.iter().sum();
5940 let offs: Vec<usize> = ts
5941 .iter()
5942 .scan(0usize, |a, &t| {
5943 let o = *a;
5944 *a += t;
5945 Some(o)
5946 })
5947 .collect();
5948 let pos_ds: Vec<CudaSlice<i32>> = ts
5950 .iter()
5951 .zip(&pos0s)
5952 .map(|(&t, &p0)| e.htod_i32(&(p0 as i32..(p0 + t) as i32).collect::<Vec<_>>()))
5953 .collect::<Result<_, _>>()?;
5954 let split = |e: &Engine,
5956 y: &CudaSlice<f32>,
5957 dim: usize|
5958 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
5959 let mut out = Vec::with_capacity(b);
5960 for s in 0..b {
5961 let mut ys = e.uninit(ts[s] * dim)?;
5962 e.copy_view_into(
5963 &mut ys,
5964 0,
5965 &y.slice(offs[s] * dim..(offs[s] + ts[s]) * dim),
5966 ts[s] * dim,
5967 )?;
5968 out.push(ys);
5969 }
5970 Ok(out)
5971 };
5972
5973 let cat_tokens: Vec<u32> = prompts.iter().flat_map(|p| p.iter().copied()).collect();
5974 let mut x = self.embed(e, &cat_tokens)?; for (il, layer) in self.layers.iter().enumerate() {
5976 let mut h = e.uninit(total * n_embd)?;
5977 let mut hx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
5978 e.rms_norm_f16out(
5979 &x,
5980 layer.attn_norm.float_data(),
5981 &mut h,
5982 &mut hx16,
5983 n_embd,
5984 total,
5985 eps,
5986 )?;
5987 let mut mixed = e.uninit(total * n_embd)?;
5989 match &layer.mixer {
5990 Mixer::Full(fa) => {
5991 let g3 = e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], &h, &hx16, total)?;
5992 let geometry = self.cfg.full_attention_geometry_at(il as u32);
5998 let (n_head, n_head_kv, head_dim) = (
5999 geometry.n_head as usize,
6000 geometry.n_head_kv as usize,
6001 geometry.head_dim_k as usize,
6002 );
6003 let fa_scale = geometry.attention_scale();
6004 let use_favl = !carried
6005 && (2..=8).contains(&b)
6006 && (head_dim == 256 || head_dim == 128)
6007 && geometry.attention_gate == memra_gguf::config::AttentionGateKind::FusedQ
6008 && std::env::var("MEMRA_NOFA").is_err()
6009 && std::env::var("MEMRA_FA_FLOOR").is_err()
6010 && std::env::var("MEMRA_FA_PP_W2").as_deref() != Ok("1")
6011 && std::env::var("MEMRA_FA_BF16KV").as_deref() != Ok("0")
6012 && std::env::var("MEMRA_FA_VL").as_deref() != Ok("0");
6013 if use_favl {
6014 let (qf_w, kf_w, vf_w) = (
6015 fa.wq.out_features(),
6016 fa.wk.out_features(),
6017 fa.wv.out_features(),
6018 );
6019 memra_gguf::config::check_fused_q_gate_extent(qf_w, head_dim, n_head, 1)?;
6024 struct APre {
6025 q: CudaSlice<f32>,
6026 gate: Option<CudaSlice<f32>>,
6027 qn: CudaSlice<f32>,
6028 kn: CudaSlice<f32>,
6029 }
6030 let mut aps = Vec::with_capacity(b);
6031 for &t in ts.iter().take(b) {
6032 aps.push(APre {
6033 q: e.uninit(t * n_head * head_dim)?,
6034 gate: Some(e.uninit(t * n_head * head_dim)?),
6035 qn: e.uninit(t * n_head * head_dim)?,
6036 kn: e.uninit(t * n_head_kv * head_dim)?,
6037 });
6038 }
6039 let (kv_dim_k, kv_dim_v, ktb, vtb) = {
6040 let kvl = caches[0].kv[il].as_ref().unwrap();
6041 (kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes)
6042 };
6043 let pargs: Vec<crate::AttnPreVl> = (0..b)
6044 .map(|s| {
6045 let (o, t) = (offs[s], ts[s]);
6046 let kvl = caches[s].kv[il].as_ref().unwrap();
6047 assert!(
6048 kvl.len == 0 && kvl.len + t <= caches[s].max_ctx,
6049 "prime_cache_batch attn vl: fresh + capacity"
6050 );
6051 crate::AttnPreVl {
6052 qf: e.addr_f32v(&g3[0].slice(o * qf_w..(o + t) * qf_w)),
6053 kf: e.addr_f32v(&g3[1].slice(o * kf_w..(o + t) * kf_w)),
6054 vf: e.addr_f32v(&g3[2].slice(o * vf_w..(o + t) * vf_w)),
6055 q: e.addr_f32(&aps[s].q),
6056 gate: e.addr_f32(aps[s].gate.as_ref().unwrap()),
6057 qn: e.addr_f32(&aps[s].qn),
6058 kn: e.addr_f32(&aps[s].kn),
6059 kc: e.addr_u8(&kvl.k),
6060 vc: e.addr_u8(&kvl.v),
6061 t: t as i32,
6062 pad: 0,
6063 }
6064 })
6065 .collect();
6066 e.attn_pre_vl8(
6067 &pargs,
6068 fa.q_norm.float_data(),
6069 fa.k_norm.float_data(),
6070 head_dim,
6071 geometry.n_rot as usize,
6072 n_head,
6073 n_head_kv,
6074 self.cfg.rms_eps,
6075 geometry.rope_base,
6076 1.0,
6077 kv_dim_k,
6078 kv_dim_v,
6079 ktb,
6080 vtb,
6081 )?;
6082 for s in 0..b {
6083 let kvl = caches[s].kv[il].as_mut().unwrap();
6084 kvl.len += ts[s];
6085 let new_len = kvl.len as i32;
6086 e.set_i32_one(&mut kvl.len_d, new_len)?;
6087 }
6088 let mut attns = Vec::with_capacity(b);
6089 let mut mirrors = Vec::with_capacity(b);
6090 for &t in ts.iter().take(b) {
6091 attns.push(e.uninit(t * n_head * head_dim)?);
6092 let n = t * n_head_kv * head_dim;
6093 mirrors.push((e.alloc_u8_uninit(n * 2)?, e.alloc_u8_uninit(n * 2)?));
6094 }
6095 let fa3_on = match std::env::var("MEMRA_FA3").as_deref() {
6098 Ok("0") => false,
6099 Ok("1") => {
6103 crate::refuse_portable_force(
6104 "MEMRA_FA3=1",
6105 "the sm_90a fa3/bf16 kernels",
6106 );
6107 true
6108 }
6109 _ => cfg!(memra_hopper_mma),
6110 };
6111 if fa3_on {
6112 let mut q16s = Vec::with_capacity(b);
6113 let mut v16s = Vec::with_capacity(b);
6114 for s in 0..b {
6115 let t = ts[s];
6116 let mut q16 = e.alloc_u8_uninit(t * n_head * head_dim * 2)?;
6117 e.f32_to_bf16_into(&aps[s].qn, &mut q16, t * n_head * head_dim)?;
6118 let mut k16 = e.alloc_u8_uninit(t * n_head_kv * head_dim * 2)?;
6119 e.f32_to_bf16_into(&aps[s].kn, &mut k16, t * n_head_kv * head_dim)?;
6120 let mut v16 = e.alloc_u8_uninit(t * n_head_kv * head_dim * 2)?;
6121 e.f32_to_bf16_v(
6122 &g3[2].slice(offs[s] * vf_w..(offs[s] + t) * vf_w),
6123 &mut v16,
6124 t * n_head_kv * head_dim,
6125 )?;
6126 q16s.push(q16);
6127 v16s.push((k16, v16));
6128 }
6129 let mut qp = [core::ptr::null::<core::ffi::c_void>(); 8];
6130 let mut kp = qp;
6131 let mut vp = qp;
6132 let mut op = [core::ptr::null_mut::<f32>(); 8];
6133 let mut tsv = [0i32; 8];
6134 for s in 0..b {
6135 qp[s] = e.addr_u8(&q16s[s]) as *const core::ffi::c_void;
6136 kp[s] = e.addr_u8(&v16s[s].0) as *const core::ffi::c_void;
6137 vp[s] = e.addr_u8(&v16s[s].1) as *const core::ffi::c_void;
6138 op[s] = e.addr_f32(&attns[s]) as *mut f32;
6139 tsv[s] = ts[s] as i32;
6140 }
6141 let rc = unsafe {
6142 crate::fa3_vl_raw(
6143 qp.as_ptr(),
6144 kp.as_ptr(),
6145 vp.as_ptr(),
6146 op.as_ptr(),
6147 tsv.as_ptr(),
6148 b as i32,
6149 n_head as i32,
6150 n_head_kv as i32,
6151 head_dim as i32,
6152 fa_scale,
6153 e.stream().cu_stream() as *mut core::ffi::c_void,
6154 )
6155 };
6156 if rc != 0 {
6157 return Err(format!("memra_fa3_vl rc={rc}").into());
6158 }
6159 } else {
6160 let fargs: Vec<crate::FaSeqVl> = (0..b)
6161 .map(|s| crate::FaSeqVl {
6162 q: e.addr_f32(&aps[s].qn),
6163 k16: e.addr_u8(&mirrors[s].0),
6164 v16: e.addr_u8(&mirrors[s].1),
6165 o: e.addr_f32(&attns[s]),
6166 kf: e.addr_f32(&aps[s].kn),
6167 vf: e.addr_f32v(
6168 &g3[2].slice(offs[s] * vf_w..(offs[s] + ts[s]) * vf_w),
6169 ),
6170 t: ts[s] as i32,
6171 pad: 0,
6172 })
6173 .collect();
6174 e.fa_prefill_vl8(&fargs, head_dim, n_head, n_head_kv, fa_scale)?;
6175 }
6176 for (s, attn) in attns.into_iter().enumerate() {
6177 let (attn_g, ag16) = self.full_attn_prime_post_fa(
6178 e,
6179 attn,
6180 &aps[s].gate,
6181 ts[s],
6182 n_head,
6183 head_dim,
6184 )?;
6185 let mut done = false;
6186 if let Some(xh) = &ag16 {
6187 done = e.try_f16_gemm_pre_into_off(
6188 &fa.wo,
6189 xh,
6190 ts[s],
6191 &mut mixed,
6192 offs[s] * n_embd,
6193 )?;
6194 }
6195 if !done {
6196 let m = e.matmul(&fa.wo, &attn_g, ts[s])?;
6197 e.copy_into(&mut mixed, offs[s] * n_embd, &m, ts[s] * n_embd)?;
6198 }
6199 }
6200 } else {
6201 let mut parts: Vec<Vec<CudaSlice<f32>>> =
6202 (0..b).map(|_| Vec::new()).collect();
6203 for (w, y) in [&fa.wq, &fa.wk, &fa.wv].iter().zip(g3) {
6204 for (s, ys) in split(e, &y, w.out_features())?.into_iter().enumerate() {
6205 parts[s].push(ys);
6206 }
6207 }
6208 for (s, g3s) in parts.into_iter().enumerate() {
6209 let (attn_g, ag16) = self.full_attn_prime_core_inner(
6211 e, fa, g3s, &pos_ds[s], ts[s], caches[s], il,
6212 )?;
6213 let mut done = false;
6214 if let Some(xh) = &ag16 {
6215 done = e.try_f16_gemm_pre_into_off(
6216 &fa.wo,
6217 xh,
6218 ts[s],
6219 &mut mixed,
6220 offs[s] * n_embd,
6221 )?;
6222 }
6223 if !done {
6224 let m = e.matmul(&fa.wo, &attn_g, ts[s])?;
6225 e.copy_into(&mut mixed, offs[s] * n_embd, &m, ts[s] * n_embd)?;
6226 }
6227 }
6228 }
6229 }
6230 Mixer::Mla(_) => crate::hybrid::mla_path_unimplemented("batched cache prime"),
6231 Mixer::Kda(_) => crate::hybrid::kda_path_unimplemented("batched prime"),
6232 Mixer::Linear(la) => {
6233 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
6238 let g4 = e.matmul_group_xh(&ws, &h, &hx16, total)?;
6239 let outs =
6240 self.linear_attn_prime_core_batch(e, la, &g4, &offs, &ts, caches, il)?;
6241 for (s, (gn, gn16)) in outs.into_iter().enumerate() {
6242 let (o, t) = (offs[s], ts[s]);
6243 let mut done = false;
6244 if let Some(xh) = &gn16 {
6245 done = e.try_f16_gemm_pre_into_off(
6246 &la.ssm_out,
6247 xh,
6248 t,
6249 &mut mixed,
6250 o * n_embd,
6251 )?;
6252 }
6253 if !done {
6254 let m = e.matmul(&la.ssm_out, &gn, t)?;
6255 e.copy_into(&mut mixed, o * n_embd, &m, t * n_embd)?;
6256 }
6257 }
6258 }
6259 }
6260 let mut x1 = e.uninit(total * n_embd)?;
6261 let mut z = e.uninit(total * n_embd)?;
6262 let mut zx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
6263 e.add_rms_norm_f16out(
6264 &x,
6265 &mixed,
6266 layer.post_attn_norm.float_data(),
6267 &mut x1,
6268 &mut z,
6269 &mut zx16,
6270 n_embd,
6271 total,
6272 eps,
6273 )?;
6274 let ffn_out = match &layer.ffn {
6275 crate::hybrid::Ffn::Dense {
6276 ffn_gate,
6277 ffn_up,
6278 ffn_down,
6279 } => {
6280 let n_ff = ffn_gate.out_features();
6281 let mut g2 = e.matmul_group_xh(&[ffn_gate, ffn_up], &z, &zx16, total)?;
6282 let up = g2.pop().unwrap();
6283 let gate = g2.pop().unwrap();
6284 let mut act = e.uninit(total * n_ff)?;
6285 let d_lim = self.cfg.clamp_shexp_at(il as u32);
6289 if Self::f16out_on(e, total) && self.cfg.m3.is_none() && d_lim.is_none() {
6290 let mut a16 = e.alloc_u8_uninit(total * n_ff * 2)?;
6291 e.silu_mul_f16out(&gate, &up, &mut act, &mut a16, total * n_ff)?;
6292 match e.try_f16_gemm_pre(ffn_down, &a16, total)? {
6293 Some(y) => y,
6294 None => e.matmul(ffn_down, &act, total)?,
6295 }
6296 } else {
6297 Self::ffn_act_lim(
6298 e,
6299 &self.cfg,
6300 &gate,
6301 &up,
6302 1.0,
6303 1.0,
6304 d_lim,
6305 &mut act,
6306 total * n_ff,
6307 )?;
6308 e.matmul(ffn_down, &act, total)?
6309 }
6310 }
6311 crate::hybrid::Ffn::Moe(m) => {
6312 self.moe_ffn_il_prefill(e, m, &z, total, il as u16)?
6313 }
6314 };
6315 let mut x2 = e.uninit(total * n_embd)?;
6316 e.add(&x1, &ffn_out, &mut x2, total * n_embd)?;
6317 x = x2;
6318 }
6319 let mut hn = e.uninit(total * n_embd)?;
6321 e.rms_norm(
6322 &x,
6323 self.output_norm.float_data(),
6324 &mut hn,
6325 n_embd,
6326 total,
6327 eps,
6328 )?;
6329 let mut hcat = e.uninit(b * n_embd)?;
6335 for s in 0..b {
6336 let last0 = (offs[s] + ts[s] - 1) * n_embd;
6337 e.copy_view_into(
6338 &mut hcat,
6339 s * n_embd,
6340 &hn.slice(last0..last0 + n_embd),
6341 n_embd,
6342 )?;
6343 }
6344 let logits_cat = if b >= 2 {
6345 e.try_f16_gemm(&self.output, &hcat, b)?
6346 } else {
6347 None
6348 };
6349 let logits_host: Option<Vec<f32>> = match &logits_cat {
6350 Some(lc) => Some(e.dtoh(lc)?),
6351 None => None,
6352 };
6353 let n_vocab = self.output.out_features();
6354 let mut hidden_all = if crate::spec::spec_hpost() {
6355 split(e, &hn, n_embd)?
6356 } else {
6357 split(e, &x, n_embd)?
6358 };
6359 let mut out = Vec::with_capacity(b);
6360 for s in 0..b {
6361 let last0 = (offs[s] + ts[s] - 1) * n_embd;
6362 let mut h_seed = e.uninit(n_embd)?;
6363 if !crate::spec::spec_hpost() {
6364 e.copy_view_into(&mut h_seed, 0, &x.slice(last0..last0 + n_embd), n_embd)?;
6365 } else {
6366 e.copy_view_into(&mut h_seed, 0, &hn.slice(last0..last0 + n_embd), n_embd)?;
6367 }
6368 let logits = match &logits_host {
6369 Some(lh) => lh[s * n_vocab..(s + 1) * n_vocab].to_vec(),
6370 None => {
6371 let mut hlast = e.uninit(n_embd)?;
6372 e.copy_view_into(&mut hlast, 0, &hn.slice(last0..last0 + n_embd), n_embd)?;
6373 e.dtoh(&e.matmul(&self.output, &hlast, 1)?)?
6374 }
6375 };
6376 caches[s].pos += ts[s];
6377 out.push((logits, h_seed, hidden_all.remove(0)));
6378 }
6379 transaction.commit();
6380 Ok(out)
6381 }
6382
6383 #[allow(clippy::too_many_arguments)]
6394 fn full_attn_prime(
6395 &self,
6396 e: &Engine,
6397 fa: &FullAttnLayer,
6398 h: &CudaSlice<f32>,
6399 hx: Option<&CudaSlice<u8>>,
6400 pos_d: &CudaSlice<i32>,
6401 t: usize,
6402 cache: &mut Cache,
6403 il: usize,
6404 seq_end: usize,
6405 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6406 if self.uses_sliding_gated_moe_program() {
6407 return self.step35_attn_prime(e, fa, h, hx, pos_d, t, cache, il, seq_end);
6408 }
6409 let g3 = match hx {
6414 Some(xh) => e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], h, xh, t)?,
6415 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
6416 };
6417 self.full_attn_prime_core(e, fa, g3, pos_d, t, cache, il)
6418 }
6419
6420 #[allow(clippy::too_many_arguments)] fn full_attn_prime_core(
6425 &self,
6426 e: &Engine,
6427 fa: &FullAttnLayer,
6428 g3: Vec<CudaSlice<f32>>,
6429 pos_d: &CudaSlice<i32>,
6430 t: usize,
6431 cache: &mut Cache,
6432 il: usize,
6433 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6434 let (attn_g, ag16) = self.full_attn_prime_core_inner(e, fa, g3, pos_d, t, cache, il)?;
6435 if let Some(xh) = &ag16
6436 && let Some(y) = e.try_f16_gemm_pre(&fa.wo, xh, t)?
6437 {
6438 return Ok(y);
6439 }
6440 e.matmul(&fa.wo, &attn_g, t)
6441 }
6442
6443 #[allow(clippy::type_complexity)] #[allow(clippy::too_many_arguments)] fn full_attn_prime_core_inner(
6446 &self,
6447 e: &Engine,
6448 fa: &FullAttnLayer,
6449 g3: Vec<CudaSlice<f32>>,
6450 pos_d: &CudaSlice<i32>,
6451 t: usize,
6452 cache: &mut Cache,
6453 il: usize,
6454 ) -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
6455 let cfg = &self.cfg;
6456 let geometry = cfg.full_attention_geometry_at(il as u32);
6457 let n_head = geometry.n_head as usize;
6458 let n_head_kv = geometry.n_head_kv as usize;
6459 let head_dim = geometry.head_dim_k as usize;
6460 let scale = geometry.attention_scale();
6461 let (pre, base_len) = self.full_attn_prime_pre_fa(e, fa, g3, pos_d, t, cache, il)?;
6462 let AttnPre { q, k, v, gate } = pre;
6463 let mut attn = e.uninit(t * n_head * head_dim)?;
6464 self.full_attn_prime_fa_dispatch(
6465 e, &q, &k, &v, &mut attn, base_len, t, cache, il, head_dim, n_head, n_head_kv, scale,
6466 )?;
6467 self.full_attn_prime_post_fa(e, attn, &gate, t, n_head, head_dim)
6468 }
6469
6470 #[allow(clippy::type_complexity)]
6474 #[allow(clippy::too_many_arguments)] fn full_attn_prime_pre_fa(
6476 &self,
6477 e: &Engine,
6478 fa: &FullAttnLayer,
6479 mut g3: Vec<CudaSlice<f32>>,
6480 pos_d: &CudaSlice<i32>,
6481 t: usize,
6482 cache: &mut Cache,
6483 il: usize,
6484 ) -> Result<(AttnPre, usize), Box<dyn std::error::Error>> {
6485 let cfg = &self.cfg;
6486 let geometry = cfg.full_attention_geometry_at(il as u32);
6487 let n_head = geometry.n_head as usize;
6488 let n_head_kv = geometry.n_head_kv as usize;
6489 let head_dim = geometry.head_dim_k as usize;
6490 let eps = cfg.rms_eps;
6491
6492 let gated = geometry.attention_gate == memra_gguf::config::AttentionGateKind::FusedQ;
6496 let v = g3.pop().unwrap();
6497 let mut k = g3.pop().unwrap();
6498 let qf = g3.pop().unwrap();
6499 let (mut q, gate) = if gated {
6500 let mut q = e.uninit(t * n_head * head_dim)?;
6501 let mut gate = e.uninit(t * n_head * head_dim)?;
6502 e.q_gate_split(&qf, &mut q, &mut gate, head_dim, n_head, t)?;
6503 (q, Some(gate))
6504 } else {
6505 (qf, None)
6506 };
6507
6508 let mut qn = e.uninit(t * n_head * head_dim)?;
6509 e.rms_norm(
6510 &q,
6511 fa.q_norm.float_data(),
6512 &mut qn,
6513 head_dim,
6514 n_head * t,
6515 eps,
6516 )?;
6517 q = qn;
6518 let mut kn = e.uninit(t * n_head_kv * head_dim)?;
6519 e.rms_norm(
6520 &k,
6521 fa.k_norm.float_data(),
6522 &mut kn,
6523 head_dim,
6524 n_head_kv * t,
6525 eps,
6526 )?;
6527 k = kn;
6528 let rope_dims = geometry.n_rot as usize;
6529 e.rope_neox(
6530 &mut q,
6531 pos_d,
6532 head_dim,
6533 rope_dims,
6534 n_head,
6535 t,
6536 geometry.rope_base,
6537 1.0,
6538 )?;
6539 e.rope_neox(
6540 &mut k,
6541 pos_d,
6542 head_dim,
6543 rope_dims,
6544 n_head_kv,
6545 t,
6546 geometry.rope_base,
6547 1.0,
6548 )?;
6549
6550 {
6553 let kvl = cache.kv[il].as_mut().unwrap();
6554 assert!(kvl.len + t <= cache.max_ctx, "prime_cache: KV overflow");
6555 e.append_kv_quantized_rows(
6556 &k,
6557 &v,
6558 &mut kvl.k,
6559 &mut kvl.v,
6560 kvl.len,
6561 t,
6562 kvl.kv_dim_k,
6563 kvl.kv_dim_v,
6564 kvl.k_tok_bytes,
6565 kvl.v_tok_bytes,
6566 crate::Engine::kv_fp8_on(),
6567 )?;
6568 kvl.len += t;
6569 let new_len = kvl.len as i32;
6570 e.set_i32_one(&mut kvl.len_d, new_len)?;
6571 }
6572
6573 let base_len = {
6574 let kvl = cache.kv[il].as_ref().unwrap();
6575 kvl.len - t };
6577 Ok((AttnPre { q, k, v, gate }, base_len))
6578 }
6579
6580 #[allow(clippy::too_many_arguments)]
6587 fn full_attn_prime_fa_dispatch(
6588 &self,
6589 e: &Engine,
6590 q: &CudaSlice<f32>,
6591 k: &CudaSlice<f32>,
6592 v: &CudaSlice<f32>,
6593 attn: &mut CudaSlice<f32>,
6594 base_len: usize,
6595 t: usize,
6596 cache: &mut Cache,
6597 il: usize,
6598 head_dim: usize,
6599 n_head: usize,
6600 n_head_kv: usize,
6601 scale: f32,
6602 ) -> Result<(), Box<dyn std::error::Error>> {
6603 if base_len == 0 && std::env::var("MEMRA_PRIME_F32CHUNK0").as_deref() == Ok("1") {
6616 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
6617 e.sdpa_naive(
6618 q, k, v, attn, head_dim, n_head, n_head_kv, t, t, scale, true,
6619 )?;
6620 } else {
6621 e.fa_prefill(
6622 q, k, v, attn, head_dim, n_head, n_head_kv, t, t, scale, true,
6623 )?;
6624 }
6625 return Ok(());
6626 }
6627 let kvl = cache.kv[il].as_ref().unwrap();
6628 let t_kv = base_len + t;
6629 let k_view = e.view_u8(&kvl.k, t_kv * kvl.k_tok_bytes);
6630 let v_view = e.view_u8(&kvl.v, t_kv * kvl.v_tok_bytes);
6631 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
6635 e.sdpa_naive_quantized_view(
6636 q,
6637 &k_view,
6638 &v_view,
6639 attn,
6640 head_dim,
6641 n_head,
6642 n_head_kv,
6643 t,
6644 t_kv,
6645 scale,
6646 true,
6647 kvl.k_tok_bytes,
6648 kvl.v_tok_bytes,
6649 )?;
6650 return Ok(());
6651 }
6652 let deqw = std::env::var("MEMRA_PRIME_DEQW")
6660 .map(|v| v != "0")
6661 .unwrap_or(true);
6662 if deqw {
6663 e.fa_prefill_view_ws(
6664 q,
6665 &k_view,
6666 &v_view,
6667 attn,
6668 head_dim,
6669 n_head,
6670 n_head_kv,
6671 t,
6672 t_kv,
6673 scale,
6674 true,
6675 kvl.k_tok_bytes,
6676 kvl.v_tok_bytes,
6677 crate::Engine::kv_fp8_on(),
6678 )?;
6679 } else {
6680 e.fa_prefill_view(
6681 q,
6682 &k_view,
6683 &v_view,
6684 attn,
6685 head_dim,
6686 n_head,
6687 n_head_kv,
6688 t,
6689 t_kv,
6690 scale,
6691 true,
6692 kvl.k_tok_bytes,
6693 kvl.v_tok_bytes,
6694 crate::Engine::kv_fp8_on(),
6695 )?;
6696 }
6697 Ok(())
6698 }
6699
6700 #[allow(clippy::type_complexity)] fn full_attn_prime_post_fa(
6704 &self,
6705 e: &Engine,
6706 attn: CudaSlice<f32>,
6707 gate: &Option<CudaSlice<f32>>,
6708 t: usize,
6709 n_head: usize,
6710 head_dim: usize,
6711 ) -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
6712 let (attn_g, ag16) = match gate {
6713 Some(gate) => {
6714 let n = t * n_head * head_dim;
6715 let mut ag = e.uninit(n)?;
6716 if Self::f16out_on(e, t) {
6717 let mut a16 = e.alloc_u8_uninit(n * 2)?;
6718 e.sig_mul_f16out(&attn, gate, &mut ag, &mut a16, n)?;
6719 (ag, Some(a16))
6720 } else {
6721 let mut gsig = e.uninit(n)?;
6722 e.sigmoid(gate, &mut gsig, n)?;
6723 e.mul(&attn, &gsig, &mut ag, n)?;
6724 (ag, None)
6725 }
6726 }
6727 None => (attn, None),
6728 };
6729 Ok((attn_g, ag16))
6730 }
6731
6732 #[allow(clippy::too_many_arguments)] fn linear_attn_prime(
6740 &self,
6741 e: &Engine,
6742 la: &LinearAttnLayer,
6743 h: &CudaSlice<f32>,
6744 hx: Option<&CudaSlice<u8>>,
6745 t: usize,
6746 cache: &mut Cache,
6747 il: usize,
6748 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6749 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
6751 let g4 = match hx {
6752 Some(xh) => e.matmul_group_xh(&ws, h, xh, t)?,
6753 None => e.matmul_group(&ws, h, t)?,
6754 };
6755 self.linear_attn_prime_core(e, la, g4, t, cache, il)
6756 }
6757
6758 fn linear_attn_prime_core(
6760 &self,
6761 e: &Engine,
6762 la: &LinearAttnLayer,
6763 mut g4: Vec<CudaSlice<f32>>,
6764 t: usize,
6765 cache: &mut Cache,
6766 il: usize,
6767 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6768 self.linear_attn_prime_core_pad(e, la, std::mem::take(&mut g4), t, cache, il, None)
6769 }
6770
6771 #[allow(clippy::too_many_arguments)]
6775 #[allow(clippy::type_complexity)] fn linear_attn_prime_core_pad_inner(
6777 &self,
6778 e: &Engine,
6779 la: &LinearAttnLayer,
6780 mut g4: Vec<CudaSlice<f32>>,
6781 t: usize,
6782 cache: &mut Cache,
6783 il: usize,
6784 pad_len: Option<&CudaSlice<i32>>,
6785 ) -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
6786 let geometry = la.geometry;
6788 let d_state = geometry.key_head_dim as usize;
6789 let num_k = geometry.key_heads as usize;
6790 let num_v = geometry.value_heads as usize;
6791 let key_dim = d_state * num_k;
6792 let value_dim = geometry.value_head_dim as usize * num_v;
6793 let conv_dim = key_dim * 2 + value_dim;
6794 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(
6799 e,
6800 la,
6801 &qkv_mixed.slice(0..t * conv_dim),
6802 &z.slice(0..t * value_dim),
6803 &beta_raw.slice(0..t * num_v),
6804 &alpha.slice(0..t * num_v),
6805 t,
6806 cache,
6807 il,
6808 pad_len,
6809 )
6810 }
6811
6812 #[allow(clippy::too_many_arguments)]
6815 fn linear_attn_gdn_prep(
6816 &self,
6817 e: &Engine,
6818 la: &LinearAttnLayer,
6819 qkv_mixed: &cudarc::driver::CudaView<f32>,
6820 beta_raw: &cudarc::driver::CudaView<f32>,
6821 alpha: &cudarc::driver::CudaView<f32>,
6822 t: usize,
6823 cache: &mut Cache,
6824 il: usize,
6825 pad_len: Option<&CudaSlice<i32>>,
6826 ) -> Result<GdnPrep, Box<dyn std::error::Error>> {
6827 let cfg = &self.cfg;
6828 let geometry = la.geometry;
6829 let d_state = geometry.key_head_dim as usize;
6830 let num_k = geometry.key_heads as usize;
6831 let num_v = geometry.value_heads as usize;
6832 let d_conv = geometry.conv_kernel as usize;
6833 let key_dim = d_state * num_k; let value_dim = geometry.value_head_dim as usize * num_v;
6835 let conv_dim = key_dim * 2 + value_dim; let eps = cfg.rms_eps;
6837 debug_assert!(
6838 t >= d_conv - 1,
6839 "stateful conv needs T >= pad (PRIME_MIN_T gates)"
6840 );
6841
6842 let rl = cache.recur[il].as_mut().unwrap();
6847 let hk = Self::gdn_hk(e, t, num_v, num_k);
6848 let conv_fuse = std::env::var("MEMRA_CONV_FUSE").as_deref() != Ok("0");
6849 let hk = if conv_fuse { hk } else { num_v }; let mut q_g = e.uninit(d_state * hk * t)?;
6851 let mut k_g = e.uninit(d_state * hk * t)?;
6852 let mut v_g = e.uninit(d_state * num_v * t)?;
6853 if conv_fuse {
6854 e.ssm_conv1d_gdn_state_pad(
6855 qkv_mixed,
6856 &mut rl.conv_state,
6857 la.ssm_conv1d.float_data(),
6858 &mut q_g,
6859 &mut k_g,
6860 &mut v_g,
6861 conv_dim,
6862 t,
6863 d_conv,
6864 d_state,
6865 num_v,
6866 num_k,
6867 key_dim,
6868 hk,
6869 pad_len,
6870 )?;
6871 } else {
6872 let mut conv_out = e.uninit(conv_dim * t)?; e.ssm_conv1d_tm_state_pad_v(
6874 qkv_mixed,
6875 &mut rl.conv_state,
6876 la.ssm_conv1d.float_data(),
6877 &mut conv_out,
6878 conv_dim,
6879 t,
6880 d_conv,
6881 pad_len,
6882 )?;
6883 e.qkv_to_gdn_repack(
6884 &conv_out, &mut q_g, &mut k_g, &mut v_g, d_state, num_v, num_k, key_dim, t,
6885 )?;
6886 }
6887 let mut q_l2 = e.uninit(d_state * hk * t)?;
6888 let qb16 = if Engine::l2_v2_on(d_state) && e.gdn_wgmma_on(32) {
6892 let mut qb = e.alloc_u8_uninit(d_state * hk * t * 2)?;
6893 e.l2_norm_pp(&q_g, &mut q_l2, Some(&mut qb), d_state, hk * t, eps)?;
6894 Some(qb)
6895 } else {
6896 e.l2_norm_pp(&q_g, &mut q_l2, None, d_state, hk * t, eps)?;
6897 None
6898 };
6899 let mut k_l2 = e.uninit(d_state * hk * t)?;
6900 let kb16 = if Engine::l2_v2_on(d_state) {
6902 let mut kb = e.alloc_u8_uninit(d_state * hk * t * 2)?;
6903 e.l2_norm_pp(&k_g, &mut k_l2, Some(&mut kb), d_state, hk * t, eps)?;
6904 Some(kb)
6905 } else {
6906 e.l2_norm_pp(&k_g, &mut k_l2, None, d_state, hk * t, eps)?;
6907 None
6908 };
6909 let mut beta = e.uninit(t * num_v)?;
6910 e.sigmoid_v(beta_raw, &mut beta, t * num_v)?;
6911 let mut g_log = e.uninit(t * num_v)?;
6912 e.gdn_glog_v(
6913 alpha,
6914 la.ssm_dt.float_data(),
6915 la.ssm_a.float_data(),
6916 &mut g_log,
6917 num_v,
6918 t,
6919 )?;
6920 if let Some(len_d) = pad_len {
6921 e.gdn_pad_mask(&mut beta, &mut g_log, len_d, num_v, t)?;
6922 }
6923 Ok(GdnPrep {
6924 hk,
6925 q_l2,
6926 k_l2,
6927 v_g,
6928 beta,
6929 g_log,
6930 kb16,
6931 qb16,
6932 })
6933 }
6934
6935 #[allow(clippy::too_many_arguments)]
6940 #[allow(clippy::type_complexity)] fn linear_attn_prime_core_batch(
6942 &self,
6943 e: &Engine,
6944 la: &LinearAttnLayer,
6945 g4: &[CudaSlice<f32>],
6946 offs: &[usize],
6947 ts: &[usize],
6948 caches: &mut [&mut Cache],
6949 il: usize,
6950 ) -> Result<Vec<(CudaSlice<f32>, Option<CudaSlice<u8>>)>, Box<dyn std::error::Error>> {
6951 let geometry = la.geometry;
6952 let d_state = geometry.key_head_dim as usize;
6953 let num_k = geometry.key_heads as usize;
6954 let num_v = geometry.value_heads as usize;
6955 let d_conv = geometry.conv_kernel as usize;
6956 let key_dim = d_state * num_k;
6957 let value_dim = geometry.value_head_dim as usize * num_v;
6958 let conv_dim = key_dim * 2 + value_dim;
6959 let eps = self.cfg.rms_eps;
6960 let scale = 1.0 / (d_state as f32).sqrt();
6961 let b = ts.len();
6962 let c = Engine::gdn_chunk_size();
6963 let carried = caches.iter().any(|c| c.pos > 0);
6966 let use_vl = !carried
6967 && (2..=8).contains(&b)
6968 && Engine::gdn_chunked_enabled()
6969 && ts.iter().all(|&t| t >= 16)
6970 && e.gdn_mma_enabled(c)
6971 && std::env::var("MEMRA_GDN_VL").as_deref() != Ok("0");
6972 if !use_vl {
6973 return (0..b)
6974 .map(|s| {
6975 let (o, t) = (offs[s], ts[s]);
6976 self.linear_attn_prime_core_pad_view(
6977 e,
6978 la,
6979 &g4[0].slice(o * conv_dim..(o + t) * conv_dim),
6980 &g4[1].slice(o * value_dim..(o + t) * value_dim),
6981 &g4[2].slice(o * num_v..(o + t) * num_v),
6982 &g4[3].slice(o * num_v..(o + t) * num_v),
6983 t,
6984 caches[s],
6985 il,
6986 None,
6987 )
6988 })
6989 .collect();
6990 }
6991 struct SeqBufs {
6995 conv_out: CudaSlice<f32>,
6996 q_g: CudaSlice<f32>,
6997 k_g: CudaSlice<f32>,
6998 v_g: CudaSlice<f32>,
6999 q_l2: CudaSlice<f32>,
7000 k_l2: CudaSlice<f32>,
7001 beta: CudaSlice<f32>,
7002 g_log: CudaSlice<f32>,
7003 gn: CudaSlice<f32>,
7004 gn16: CudaSlice<u8>,
7005 }
7006 let f16o = Self::f16out_on(e, 16);
7007 let hk = Self::gdn_hk(e, 16, num_v, num_k); let mut sb = Vec::with_capacity(b);
7009 let mut pres = Vec::with_capacity(b);
7010 for &t in ts.iter().take(b) {
7011 sb.push(SeqBufs {
7012 conv_out: e.uninit(conv_dim * t)?,
7013 q_g: e.uninit(d_state * hk * t)?,
7014 k_g: e.uninit(d_state * hk * t)?,
7015 v_g: e.uninit(d_state * num_v * t)?,
7016 q_l2: e.uninit(d_state * hk * t)?,
7017 k_l2: e.uninit(d_state * hk * t)?,
7018 beta: e.uninit(t * num_v)?,
7019 g_log: e.uninit(t * num_v)?,
7020 gn: e.uninit(d_state * num_v * t)?,
7021 gn16: e.alloc_u8_uninit(d_state * num_v * t * 2)?,
7022 });
7023 pres.push(e.gdn_chunk_alloc(num_v, t, c, hk)?);
7024 }
7025 let prep_args: Vec<crate::GdnPrepVl> = (0..b)
7026 .map(|s| {
7027 let (o, t) = (offs[s], ts[s]);
7028 let rl = caches[s].recur[il].as_ref().unwrap();
7029 crate::GdnPrepVl {
7030 qkv: e.addr_f32v(&g4[0].slice(o * conv_dim..(o + t) * conv_dim)),
7031 conv_state: e.addr_f32(&rl.conv_state),
7032 conv_out: e.addr_f32(&sb[s].conv_out),
7033 q_g: e.addr_f32(&sb[s].q_g),
7034 k_g: e.addr_f32(&sb[s].k_g),
7035 v_g: e.addr_f32(&sb[s].v_g),
7036 q_l2: e.addr_f32(&sb[s].q_l2),
7037 k_l2: e.addr_f32(&sb[s].k_l2),
7038 beta_raw: e.addr_f32v(&g4[2].slice(o * num_v..(o + t) * num_v)),
7039 alpha: e.addr_f32v(&g4[3].slice(o * num_v..(o + t) * num_v)),
7040 beta: e.addr_f32(&sb[s].beta),
7041 g_log: e.addr_f32(&sb[s].g_log),
7042 o: e.addr_f32(&pres[s].o),
7043 z: e.addr_f32v(&g4[1].slice(o * value_dim..(o + t) * value_dim)),
7044 gn: e.addr_f32(&sb[s].gn),
7045 gn16: e.addr_u8(&sb[s].gn16),
7046 kb16: if Engine::l2_v2_on(d_state) {
7047 e.addr_u8(&pres[s].kb16)
7048 } else {
7049 0
7050 },
7051 qb16: if Engine::l2_v2_on(d_state) && e.gdn_wgmma_on(c) {
7052 e.addr_u8(&pres[s].qb16)
7053 } else {
7054 0
7055 },
7056 t: t as i32,
7057 pad: 0,
7058 }
7059 })
7060 .collect();
7061 let args: Vec<crate::GdnSeqVl> = (0..b)
7062 .map(|s| {
7063 let rl = caches[s].recur[il].as_ref().unwrap();
7064 crate::GdnSeqVl {
7065 kb16: e.addr_u8(&pres[s].kb16),
7066 gcum: e.addr_f32(&pres[s].gcum),
7067 beta: e.addr_f32(&sb[s].beta),
7068 u: e.addr_f32(&pres[s].u),
7069 wb16: e.addr_u8(&pres[s].wb16),
7070 y: e.addr_u8(&pres[s].y16),
7071 ssnap: e.addr_u8(&pres[s].ssnap16),
7072 state_in: e.addr_f32(&rl.ssm_state),
7073 state_out: e.addr_f32(&rl.ssm_state_alt),
7074 q: e.addr_f32(&sb[s].q_l2),
7075 p: e.addr_f32(&pres[s].p),
7076 o: e.addr_f32(&pres[s].o),
7077 k: e.addr_f32(&sb[s].k_l2),
7078 v: e.addr_f32(&sb[s].v_g),
7079 g: e.addr_f32(&sb[s].g_log),
7080 a: e.addr_f32(&pres[s].a),
7081 w: e.addr_f32(&pres[s].w),
7082 t: ts[s] as i32,
7083 nc: pres[s].nc as i32,
7084 }
7085 })
7086 .collect();
7087 e.gdn_prep_vl8(
7088 &prep_args,
7089 la.ssm_conv1d.float_data(),
7090 la.ssm_dt.float_data(),
7091 la.ssm_a.float_data(),
7092 conv_dim,
7093 d_conv,
7094 d_state,
7095 num_v,
7096 num_k,
7097 key_dim,
7098 hk,
7099 eps,
7100 )?;
7101 if !Engine::l2_v2_on(d_state) {
7104 e.gdn_mirror_vl8(&args, num_v, 0, hk)?;
7105 }
7106 let wq8: Option<crate::GdnWVl8> = if e.gdn_wgmma_on(c) {
7108 if !Engine::l2_v2_on(d_state) {
7110 for s in 0..b {
7111 e.f32_to_bf16_into(&sb[s].q_l2, &mut pres[s].qb16, d_state * hk * ts[s])?;
7112 }
7113 }
7114 let mut wa = [crate::GdnWVl::default(); 8];
7115 for s in 0..b {
7116 wa[s] = crate::GdnWVl {
7117 qb16: e.addr_u8(&pres[s].qb16),
7118 pb16: e.addr_u8(&pres[s].pb16),
7119 };
7120 }
7121 Some(crate::GdnWVl8(wa))
7122 } else {
7123 None
7124 };
7125 e.gdn_chunk_k123_vl8(&args, num_v, hk, wq8.as_ref())?;
7126 e.gdn_chunk_vl8(&args, num_v, scale, hk, wq8.as_ref())?;
7127 if f16o {
7128 e.gdn_tail_vl8(&prep_args, la.ssm_norm.float_data(), d_state, num_v, eps)?;
7129 }
7130 let mut out = Vec::with_capacity(b);
7132 for (s, bufs) in sb.into_iter().enumerate() {
7133 let rl = caches[s].recur[il].as_mut().unwrap();
7134 std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
7135 let (o, t) = (offs[s], ts[s]);
7136 let SeqBufs { mut gn, gn16, .. } = bufs;
7137 if f16o {
7138 out.push((gn, Some(gn16)));
7139 } else {
7140 let z_v = g4[1].slice(o * value_dim..(o + t) * value_dim);
7141 e.gated_rmsnorm_zv(
7142 &pres[s].o,
7143 la.ssm_norm.float_data(),
7144 &z_v,
7145 &mut gn,
7146 d_state,
7147 num_v * t,
7148 eps,
7149 )?;
7150 out.push((gn, None));
7151 }
7152 }
7153 Ok(out)
7154 }
7155
7156 #[allow(clippy::too_many_arguments)]
7160 #[allow(clippy::type_complexity)] fn linear_attn_prime_core_pad_view(
7162 &self,
7163 e: &Engine,
7164 la: &LinearAttnLayer,
7165 qkv_mixed: &cudarc::driver::CudaView<f32>,
7166 z: &cudarc::driver::CudaView<f32>,
7167 beta_raw: &cudarc::driver::CudaView<f32>,
7168 alpha: &cudarc::driver::CudaView<f32>,
7169 t: usize,
7170 cache: &mut Cache,
7171 il: usize,
7172 pad_len: Option<&CudaSlice<i32>>,
7173 ) -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
7174 let cfg = &self.cfg;
7175 let geometry = la.geometry;
7176 let d_state = geometry.key_head_dim as usize;
7177 let num_v = geometry.value_heads as usize;
7178 let eps = cfg.rms_eps;
7179 let scale = 1.0 / (d_state as f32).sqrt();
7180
7181 let prep =
7182 self.linear_attn_gdn_prep(e, la, qkv_mixed, beta_raw, alpha, t, cache, il, pad_len)?;
7183
7184 let mut o = e.uninit(d_state * num_v * t)?;
7190 let rl = cache.recur[il].as_mut().unwrap();
7191 {
7192 let crate::cache::RecurLayer {
7193 ssm_state,
7194 ssm_state_alt,
7195 ..
7196 } = rl;
7197 e.gdn_scan_prefill(
7198 &prep.q_l2,
7199 &prep.k_l2,
7200 &prep.v_g,
7201 &prep.g_log,
7202 &prep.beta,
7203 prep.kb16.as_ref(),
7204 prep.qb16.as_ref(),
7205 ssm_state,
7206 ssm_state_alt,
7207 &mut o,
7208 num_v,
7209 t,
7210 scale,
7211 prep.hk,
7212 )?;
7213 }
7214 std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
7215
7216 let mut gn = e.uninit(d_state * num_v * t)?;
7219 let gn16 = if Self::f16out_on(e, t) {
7220 let mut g16 = e.alloc_u8_uninit(d_state * num_v * t * 2)?;
7221 e.gated_rmsnorm_f16out_zv(
7222 &o,
7223 la.ssm_norm.float_data(),
7224 z,
7225 &mut gn,
7226 &mut g16,
7227 d_state,
7228 num_v * t,
7229 eps,
7230 )?;
7231 Some(g16)
7232 } else {
7233 e.gated_rmsnorm_zv(
7234 &o,
7235 la.ssm_norm.float_data(),
7236 z,
7237 &mut gn,
7238 d_state,
7239 num_v * t,
7240 eps,
7241 )?;
7242 None
7243 };
7244 Ok((gn, gn16))
7245 }
7246
7247 #[allow(clippy::too_many_arguments)]
7249 fn linear_attn_prime_core_pad(
7250 &self,
7251 e: &Engine,
7252 la: &LinearAttnLayer,
7253 g4: Vec<CudaSlice<f32>>,
7254 t: usize,
7255 cache: &mut Cache,
7256 il: usize,
7257 pad_len: Option<&CudaSlice<i32>>,
7258 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7259 let (gn, gn16) = self.linear_attn_prime_core_pad_inner(e, la, g4, t, cache, il, pad_len)?;
7260 if let Some(xh) = &gn16
7261 && let Some(y) = e.try_f16_gemm_pre(&la.ssm_out, xh, t)?
7262 {
7263 return Ok(y);
7264 }
7265 e.matmul(&la.ssm_out, &gn, t)
7266 }
7267
7268 pub fn full_attn(
7273 &self,
7274 e: &Engine,
7275 fa: &FullAttnLayer,
7276 h: &CudaSlice<f32>,
7277 pos_d: &CudaSlice<i32>,
7278 t: usize,
7279 il: usize,
7280 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7281 if self.uses_sliding_gated_moe_program() {
7282 return self.step35_attn(e, fa, h, pos_d, t, il);
7283 }
7284 let cfg = &self.cfg;
7285 let _n_embd = cfg.n_embd as usize;
7286 let geometry = cfg.full_attention_geometry_at(il as u32);
7287 let n_head = geometry.n_head as usize;
7288 let n_head_kv = geometry.n_head_kv as usize;
7289 let head_dim = geometry.head_dim_k as usize;
7290 let eps = cfg.rms_eps;
7291 let scale = geometry.attention_scale();
7292
7293 let gated = geometry.attention_gate == memra_gguf::config::AttentionGateKind::FusedQ;
7296 let mut g3 = match self.full_attn_tp_qkv(e, fa, h, t)? {
7300 Some(g3) => g3,
7301 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
7302 };
7303 let v = g3.pop().unwrap();
7304 let mut k = g3.pop().unwrap();
7305 let qf = g3.pop().unwrap();
7306 let (mut q, gate) = if gated {
7307 let mut q = e.uninit(t * n_head * head_dim)?;
7308 let mut gate = e.uninit(t * n_head * head_dim)?;
7309 e.q_gate_split(&qf, &mut q, &mut gate, head_dim, n_head, t)?;
7310 (q, Some(gate))
7311 } else {
7312 (qf, None)
7313 };
7314
7315 let mut qn = e.uninit(t * n_head * head_dim)?;
7317 e.rms_norm(
7318 &q,
7319 fa.q_norm.float_data(),
7320 &mut qn,
7321 head_dim,
7322 n_head * t,
7323 eps,
7324 )?;
7325 q = qn;
7326 let mut kn = e.uninit(t * n_head_kv * head_dim)?;
7327 e.rms_norm(
7328 &k,
7329 fa.k_norm.float_data(),
7330 &mut kn,
7331 head_dim,
7332 n_head_kv * t,
7333 eps,
7334 )?;
7335 k = kn;
7336 let rope_dims = geometry.n_rot as usize;
7337 e.rope_neox(
7338 &mut q,
7339 pos_d,
7340 head_dim,
7341 rope_dims,
7342 n_head,
7343 t,
7344 geometry.rope_base,
7345 1.0,
7346 )?;
7347 e.rope_neox(
7348 &mut k,
7349 pos_d,
7350 head_dim,
7351 rope_dims,
7352 n_head_kv,
7353 t,
7354 geometry.rope_base,
7355 1.0,
7356 )?;
7357
7358 let mut attn = e.uninit(t * n_head * head_dim)?;
7360 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
7363 e.sdpa_naive(
7365 &q, &k, &v, &mut attn, head_dim, n_head, n_head_kv, t, t, scale, true,
7366 )?;
7367 } else {
7368 e.fa_prefill(
7369 &q, &k, &v, &mut attn, head_dim, n_head, n_head_kv, t, t, scale, true,
7370 )?;
7371 }
7372
7373 let attn_g = match &gate {
7375 Some(gate) => {
7376 let mut gsig = e.uninit(t * n_head * head_dim)?;
7377 e.sigmoid(gate, &mut gsig, t * n_head * head_dim)?;
7378 let mut ag = e.uninit(t * n_head * head_dim)?;
7379 e.mul(&attn, &gsig, &mut ag, t * n_head * head_dim)?;
7380 ag
7381 }
7382 None => attn,
7383 };
7384
7385 self.full_attn_o(e, fa, &attn_g, t)
7387 }
7388
7389 fn mla_split_operand<'w>(
7397 w: &'w crate::model::GpuTensor,
7398 name: &str,
7399 il: usize,
7400 ) -> &'w CudaSlice<f32> {
7401 match w {
7402 crate::model::GpuTensor::Float { data, .. } => data,
7403 _ => panic!(
7404 "layer {il}: MLA conversion-split operand {name} is not f32-resident. The 3D \
7405 (d_nope|kv_rank, kv_rank|d_v, n_head) splits have no quantized resident layout: \
7406 a quantized 3D tensor mis-derives row_bytes in the generic 2D Quant arm, so the \
7407 source must dequantize the fused kv_b_proj (TensorTransform::SplitMlaKv). \
7408 Reaching this means both the loader rank guard and MlaAttnLayer::load's \
7409 residency audit were bypassed"
7410 ),
7411 }
7412 }
7413
7414 #[allow(clippy::too_many_arguments)]
7422 fn mla_attn_core(
7426 &self,
7427 e: &Engine,
7428 mla: &crate::hybrid::MlaAttnLayer,
7429 h: &CudaSlice<f32>,
7430 pos_d: &CudaSlice<i32>,
7431 t: usize,
7432 il: usize,
7433 latent: &mut CudaSlice<f32>,
7434 index_plane: Option<IndexerPlanes<'_>>,
7435 slot: usize,
7436 rows_exact: bool,
7437 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7438 let attn = self.mla_attn_core_pre_wo(
7439 e,
7440 mla,
7441 h,
7442 pos_d,
7443 t,
7444 il,
7445 latent,
7446 index_plane,
7447 slot,
7448 rows_exact,
7449 )?;
7450 if rows_exact {
7454 e.matmul_rows_exact(&mla.wo, &attn, t)
7455 } else {
7456 e.matmul(&mla.wo, &attn, t)
7457 }
7458 }
7459
7460 #[allow(clippy::too_many_arguments)]
7466 fn mla_attn_core_pre_wo(
7467 &self,
7468 e: &Engine,
7469 mla: &crate::hybrid::MlaAttnLayer,
7470 h: &CudaSlice<f32>,
7471 pos_d: &CudaSlice<i32>,
7472 t: usize,
7473 il: usize,
7474 latent: &mut CudaSlice<f32>,
7475 index_plane: Option<IndexerPlanes<'_>>,
7476 slot: usize,
7477 rows_exact: bool,
7478 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7479 let g = mla.geom;
7480 let cfg = &self.cfg;
7481 let eps = cfg.rms_eps;
7482 let base = cfg.rope_freq_base;
7483 let (nh, dn, dr, dv, r) = (g.n_head, g.d_nope, g.d_rope, g.d_v, g.kv_rank);
7484 assert_eq!(
7485 g.latent_dim,
7486 r + dr,
7487 "layer {il}: MlaGeom latent_dim disagrees with kv_rank + d_rope"
7488 );
7489 let t_kv = slot + t;
7490 let mm = |w: &crate::model::GpuTensor,
7495 x: &CudaSlice<f32>|
7496 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7497 if rows_exact {
7498 e.matmul_rows_exact(w, x, t)
7499 } else {
7500 e.matmul(w, x, t)
7501 }
7502 };
7503
7504 let q_a = mm(&mla.wq_a, h)?;
7506 let q_lora = mla.wq_b.in_features();
7507 let mut q_an = e.uninit(t * q_lora)?;
7508 e.rms_norm(&q_a, mla.q_a_norm.float_data(), &mut q_an, q_lora, t, eps)?;
7509 let q = mm(&mla.wq_b, &q_an)?;
7510 let mut q_nope = e.uninit(t * nh * dn)?;
7513 let mut q_pe = e.uninit((t * nh * dr).max(1))?;
7514 e.mla_split_latent(&q, &mut q_nope, &mut q_pe, t * nh, dn, dr)?;
7515 e.mla_rope_interleaved(&mut q_pe, pos_d, t, nh, dr, base)?;
7518
7519 let kv = mm(&mla.wkv_a, h)?;
7521 let mut c_kv = e.uninit(t * r)?;
7522 let mut k_pe = e.uninit((t * dr).max(1))?;
7523 e.mla_split_latent(&kv, &mut c_kv, &mut k_pe, t, r, dr)?;
7524 let mut c_kv_n = e.uninit(t * r)?;
7525 e.rms_norm(&c_kv, mla.kv_a_norm.float_data(), &mut c_kv_n, r, t, eps)?;
7526 e.mla_rope_interleaved(&mut k_pe, pos_d, t, 1, dr, base)?;
7527 e.mla_append_latent(latent, &c_kv_n, &k_pe, slot, t, r, dr)?;
7528
7529 let gathered = match (&mla.index, index_plane) {
7533 (Some(indexer), Some(plane)) => {
7534 Some(self.mla_kpool_select(e, indexer, h, &q_an, plane, t, slot, il, rows_exact)?)
7535 }
7536 (Some(_), None) => {
7537 return Err(format!(
7538 "layer {il} declares a DSA k-pool indexer but no indexer state plane was \
7539 supplied — the ModelPlan must declare StatePlan::LatentKvCache with a \
7540 non-zero index_width for it"
7541 )
7542 .into());
7543 }
7544 (None, _) => None,
7545 };
7546
7547 let wk_b = Self::mla_split_operand(&mla.wk_b, "attn_k_b", il);
7549 let wv_b = Self::mla_split_operand(&mla.wv_b, "attn_v_b", il);
7550
7551 if mla.tp_shard
7574 && gathered.is_some()
7575 && dr == 0
7576 && r == 512
7577 && t >= 16
7578 && !crate::portable_mma_gated()
7579 && mla_tc_prefill_enabled()
7580 {
7581 static TP_TC_DECLINE: std::sync::Once = std::sync::Once::new();
7582 TP_TC_DECLINE.call_once(|| {
7583 eprintln!(
7584 "[mla-tc-prefill] DECLINED on glm5-TP head shards: the door's gate ran \
7585 on full-head geometry; shards ride the f32 prefill kernels until the \
7586 TP composition gate lands (pin MEMRA_MLA_TC_PREFILL=0 to silence)"
7587 );
7588 });
7589 }
7590 if let Some((idx, slots)) = &gathered
7591 && dr == 0
7592 && r == 512
7593 && t >= 16
7594 && !rows_exact && !crate::portable_mma_gated()
7596 && !mla.tp_shard
7597 && mla_tc_prefill_enabled()
7598 && let Some(attn) = self.mla_tc_prefill_chain(
7599 e, wk_b, wv_b, &q_nope, latent, idx, *slots, t, t_kv, nh, dn, dv, r, g.scale,
7600 )?
7601 {
7602 return Ok(attn);
7603 }
7604
7605 let mut q_lat = e.uninit(t * nh * r)?;
7606 e.mla_absorb_q(&q_nope, wk_b, &mut q_lat, t, nh, dn, r)?;
7607 let mut o_lat = e.uninit(t * nh * r)?;
7608 match &gathered {
7609 Some((idx, slots)) => e.mla_attn_gathered(
7610 &q_lat, &q_pe, latent, idx, &mut o_lat, nh, r, dr, t, *slots, g.scale,
7611 )?,
7612 None => e.mla_attn_absorbed(
7613 &q_lat, &q_pe, latent, &mut o_lat, nh, r, dr, t, t_kv, g.scale,
7614 )?,
7615 }
7616 let mut attn = e.uninit(t * nh * dv)?;
7617 e.mla_decompress_v(&o_lat, wv_b, &mut attn, t, nh, dv, r)?;
7618
7619 Ok(attn)
7620 }
7621
7622 #[allow(clippy::too_many_arguments)]
7646 fn mla_tc_prefill_chain(
7647 &self,
7648 e: &Engine,
7649 wk_b: &CudaSlice<f32>,
7650 wv_b: &CudaSlice<f32>,
7651 q_nope: &CudaSlice<f32>,
7652 latent: &CudaSlice<f32>,
7653 idx: &CudaSlice<i32>,
7654 width: usize,
7655 t: usize,
7656 t_kv: usize,
7657 nh: usize,
7658 dn: usize,
7659 dv: usize,
7660 r: usize,
7661 scale: f32,
7662 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
7663 fn declined(stage: &str, m: usize, n: usize, k: usize, batch: usize) {
7666 type ShapeSet = std::collections::HashSet<(usize, usize, usize, usize)>;
7667 static SAID: std::sync::Mutex<Option<ShapeSet>> = std::sync::Mutex::new(None);
7668 let mut g = SAID.lock().unwrap();
7669 if g.get_or_insert_with(std::collections::HashSet::new)
7670 .insert((m, n, k, batch))
7671 {
7672 eprintln!(
7673 "[mla-tc-prefill] DECLINED at {stage} m={m} n={n} k={k} batch={batch} \
7674 (no cuBLASLt heuristic) — this call falls back to the f32 MLA kernels"
7675 );
7676 }
7677 }
7678 for (name, n) in [
7682 ("wk_b", nh * r * dn),
7683 ("wv_b", nh * dv * r),
7684 ("q_nope", t * nh * dn),
7685 ("latent", t_kv * r),
7686 ] {
7687 debug_assert!(
7688 n.is_multiple_of(4),
7689 "mla-tc-prefill: {name} elems {n} % 4 != 0"
7690 );
7691 let _ = (name, n);
7692 }
7693 let wk_bf = e.f32_to_bf16(wk_b, nh * r * dn)?;
7694 let wv_bf = e.f32_to_bf16(wv_b, nh * dv * r)?;
7695 let qn_bf = e.f32_to_bf16(q_nope, t * nh * dn)?;
7696 let mut q_lat_bf = e.alloc_u8_uninit(t * nh * r * 2)?;
7700 if !e.mla_bf16_gemm_sb_bf16out(
7701 &wk_bf,
7702 &qn_bf,
7703 &mut q_lat_bf,
7704 t,
7705 r,
7706 dn,
7707 nh * dn,
7708 dn,
7709 nh * r,
7710 r,
7711 nh,
7712 )? {
7713 declined("absorb", t, r, dn, nh);
7714 return Ok(None);
7715 }
7716 let cache_bf = e.f32_to_bf16(latent, t_kv * r)?;
7718 let mut o_lat = e.uninit(t * nh * r)?;
7719 e.mla_attn_gathered_tc(
7720 &q_lat_bf, &cache_bf, idx, &mut o_lat, nh, r, t, width, scale,
7721 )?;
7722 let o_bf = e.f32_to_bf16(&o_lat, t * nh * r)?;
7725 let mut attn = e.uninit(t * nh * dv)?;
7726 if !e.mla_bf16_gemm_sb_f32out(
7727 &wv_bf,
7728 &o_bf,
7729 &mut attn,
7730 t,
7731 dv,
7732 r,
7733 nh * r,
7734 r,
7735 nh * dv,
7736 dv,
7737 nh,
7738 )? {
7739 declined("decompress", t, dv, r, nh);
7740 return Ok(None);
7741 }
7742 crate::MLA_TC_PREFILL_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
7743 {
7744 static ANNOUNCED: std::sync::Once = std::sync::Once::new();
7745 ANNOUNCED.call_once(|| {
7746 eprintln!(
7747 "[mla-tc-prefill] engaged: absorb/decompress = strided-batched bf16 TC \
7748 GEMMs, attention = fa_mla_gathered_bf16 (t={t}, t_kv={t_kv}, nh={nh}, \
7749 width={width}); dispatches counted in MLA_TC_PREFILL_DISPATCHES"
7750 );
7751 });
7752 }
7753 Ok(Some(attn))
7754 }
7755
7756 #[allow(clippy::too_many_arguments)]
7758 fn mla_kpool_select(
7759 &self,
7760 e: &Engine,
7761 indexer: &crate::hybrid::MlaIndexer,
7762 h: &CudaSlice<f32>,
7763 q_resid: &CudaSlice<f32>,
7764 plane: IndexerPlanes<'_>,
7765 t: usize,
7766 slot: usize,
7767 il: usize,
7768 rows_exact: bool,
7769 ) -> Result<(CudaSlice<i32>, usize), Box<dyn std::error::Error>> {
7770 Self::mla_kpool_indices_ex(e, indexer, h, q_resid, plane, t, slot, rows_exact).map_err(
7771 |source| -> Box<dyn std::error::Error> {
7772 format!("layer {il}: DSA k-pool selection failed: {source}").into()
7773 },
7774 )
7775 }
7776
7777 #[allow(clippy::too_many_arguments)]
7794 pub fn mla_kpool_indices(
7795 e: &Engine,
7796 indexer: &crate::hybrid::MlaIndexer,
7797 h: &CudaSlice<f32>,
7798 q_resid: &CudaSlice<f32>,
7799 plane: IndexerPlanes<'_>,
7800 t: usize,
7801 slot: usize,
7802 ) -> Result<(CudaSlice<i32>, usize), Box<dyn std::error::Error>> {
7803 Self::mla_kpool_indices_ex(e, indexer, h, q_resid, plane, t, slot, false)
7804 }
7805
7806 #[allow(clippy::too_many_arguments)]
7811 pub fn mla_kpool_indices_ex(
7812 e: &Engine,
7813 indexer: &crate::hybrid::MlaIndexer,
7814 h: &CudaSlice<f32>,
7815 q_resid: &CudaSlice<f32>,
7816 plane: IndexerPlanes<'_>,
7817 t: usize,
7818 slot: usize,
7819 rows_exact: bool,
7820 ) -> Result<(CudaSlice<i32>, usize), Box<dyn std::error::Error>> {
7821 let mm = |w: &crate::model::GpuTensor,
7822 x: &CudaSlice<f32>|
7823 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7824 if rows_exact {
7825 e.matmul_rows_exact(w, x, t)
7826 } else {
7827 e.matmul(w, x, t)
7828 }
7829 };
7830 const INDEX_NORM_EPS: f32 = 1e-5;
7834
7835 let ig = indexer.geom;
7836 let d = ig.head_dim;
7837 let t_kv = slot + t;
7838 let IndexerPlanes {
7839 state: plane,
7840 pool_keys: pool_key_plane,
7841 ready: pools_ready,
7842 state_ring_rows,
7843 capacity_tokens,
7844 } = plane;
7845
7846 let ring = if state_ring_rows == 0 {
7854 0
7855 } else {
7856 state_ring_rows / ig.pool * ig.pool
7857 };
7858 if state_ring_rows > 0 && ring == 0 {
7859 return Err(format!(
7860 "indexer tail ring of {state_ring_rows} rows cannot hold one pool of {}; \
7861 raise MEMRA_DSA_INDEX_RING or set it to 0 for the flat plane",
7862 ig.pool
7863 )
7864 .into());
7865 }
7866 if *pools_ready > slot / ig.pool {
7871 return Err(format!(
7872 "resident k-pool key plane claims {} finished pools but the cache holds only {} \
7873 complete pools before this call ({slot} rows / pool {}) — a rewind reduced the \
7874 latent length without clamping index_pools_ready",
7875 *pools_ready,
7876 slot / ig.pool,
7877 ig.pool
7878 )
7879 .into());
7880 }
7881
7882 let k_raw = mm(&indexer.wk, h)?;
7885 let mut k_norm = e.uninit(t * d)?;
7886 e.layer_norm_bias(
7887 &k_raw,
7888 indexer.k_norm_w.float_data(),
7889 indexer.k_norm_b.float_data(),
7890 &mut k_norm,
7891 d,
7892 t,
7893 INDEX_NORM_EPS,
7894 )?;
7895 let gate = mm(&indexer.kpool_gate, h)?;
7896
7897 let n_pools = t_kv / ig.pool;
7903 let select_k = ig.select_k(n_pools);
7904 let width = ig.index_width(n_pools);
7905 let capacity_pools = capacity_tokens / ig.pool;
7913 let need = (capacity_pools * d).max(n_pools * d).max(1);
7914 if pool_key_plane.as_ref().is_none_or(|k| k.len() < need) {
7915 *pool_key_plane = Some(e.uninit(need)?);
7916 *pools_ready = 0;
7917 }
7918 let pool_keys = pool_key_plane
7919 .as_mut()
7920 .expect("resident pool-key plane just allocated");
7921
7922 let ape = indexer.kpool_ape.float_data();
7935 let mut cur = slot;
7936 let mut appended = 0usize;
7937 while appended < t {
7938 let take =
7939 crate::cache::index_ring_take(ring, ig.pool, *pools_ready, cur, t - appended)
7940 .ok_or_else(|| -> Box<dyn std::error::Error> {
7941 format!(
7942 "indexer tail ring lapped: {ring} rows cannot hold the {} rows still \
7943 owed to unbuilt pools at row {cur} (pools_ready {}, pool {}, slot \
7944 {slot}, t {t}). The pool-key plane was reset or the cache rewound \
7945 without clamping index_pools_ready, so rows this call must read were \
7946 already overwritten. Raise MEMRA_DSA_INDEX_RING, or set \
7947 MEMRA_DSA_INDEX_RING=0 for the flat plane",
7948 cur.saturating_sub((*pools_ready).saturating_mul(ig.pool)),
7949 *pools_ready,
7950 ig.pool
7951 )
7952 .into()
7953 })?;
7954 debug_assert!(take > 0 && appended + take <= t);
7955 e.mla_index_append(plane, &k_norm, &gate, appended, cur, take, d, d, ring)?;
7956 cur += take;
7957 appended += take;
7958 let ready_now = cur / ig.pool;
7959 e.mla_kpool_pool_keys(
7960 plane,
7961 ape,
7962 pool_keys,
7963 (*pools_ready).min(ready_now),
7964 ready_now,
7965 ig.pool,
7966 d,
7967 ring,
7968 )?;
7969 *pools_ready = ready_now;
7970 }
7971 debug_assert!(t == 0 || *pools_ready == n_pools);
7972 let pool_keys = &*pool_keys;
7973
7974 let q_index = mm(&indexer.wq_b, q_resid)?;
7976 let head_weights = mm(&indexer.weights_proj, h)?;
7977 let mut score = e.uninit((t * n_pools).max(1))?;
7978 e.mla_kpool_score(
7979 &q_index,
7980 pool_keys,
7981 &head_weights,
7982 &mut score,
7983 t,
7984 ig.heads,
7985 d,
7986 n_pools,
7987 ig.pool,
7988 slot,
7989 (d as f32).powf(-0.5),
7990 (ig.heads as f32).powf(-0.5),
7991 )?;
7992 let mut idx = e.uninit_i32(t * width)?;
7993 e.mla_kpool_select(
7994 &score,
7995 &mut idx,
7996 t,
7997 n_pools,
7998 ig.pool,
7999 select_k,
8000 width,
8001 slot,
8002 ig.always_select_tail,
8003 )?;
8004 Ok((idx, width))
8005 }
8006
8007 pub fn mla_attn(
8010 &self,
8011 e: &Engine,
8012 mla: &crate::hybrid::MlaAttnLayer,
8013 h: &CudaSlice<f32>,
8014 pos_d: &CudaSlice<i32>,
8015 t: usize,
8016 il: usize,
8017 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8018 if mla.tp.is_some() {
8019 return Err(format!(
8020 "layer {il}: MLA layer is glm5-TP-sharded (MEMRA_GLM5_TP): the stateless \
8021 mixer path is unwired for a head shard"
8022 )
8023 .into());
8024 }
8025 let mut latent = e.uninit(t * mla.geom.latent_dim)?;
8026 let mut index_plane = match mla.index.as_ref() {
8027 Some(indexer) => Some(e.uninit(t * indexer.geom.state_width())?),
8028 None => None,
8029 };
8030 let mut pool_keys = None;
8033 let mut pools_ready = 0usize;
8034 let planes = index_plane.as_mut().map(|state| IndexerPlanes {
8035 state,
8036 pool_keys: &mut pool_keys,
8037 ready: &mut pools_ready,
8038 state_ring_rows: 0,
8040 capacity_tokens: t,
8041 });
8042 self.mla_attn_core(e, mla, h, pos_d, t, il, &mut latent, planes, 0, false)
8043 }
8044
8045 #[allow(clippy::too_many_arguments)] pub fn mla_attn_cached(
8050 &self,
8051 e: &Engine,
8052 mla: &crate::hybrid::MlaAttnLayer,
8053 h: &CudaSlice<f32>,
8054 pos_d: &CudaSlice<i32>,
8055 t: usize,
8056 il: usize,
8057 cache: &mut Cache,
8058 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8059 self.mla_attn_cached_inner(e, mla, h, pos_d, t, il, cache, false)
8060 }
8061
8062 #[allow(clippy::too_many_arguments)] pub fn mla_attn_cached_rows_exact(
8073 &self,
8074 e: &Engine,
8075 mla: &crate::hybrid::MlaAttnLayer,
8076 h: &CudaSlice<f32>,
8077 pos_d: &CudaSlice<i32>,
8078 t: usize,
8079 il: usize,
8080 cache: &mut Cache,
8081 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8082 self.mla_attn_cached_inner(e, mla, h, pos_d, t, il, cache, true)
8083 }
8084
8085 #[allow(clippy::too_many_arguments)] fn mla_attn_cached_inner(
8087 &self,
8088 e: &Engine,
8089 mla: &crate::hybrid::MlaAttnLayer,
8090 h: &CudaSlice<f32>,
8091 pos_d: &CudaSlice<i32>,
8092 t: usize,
8093 il: usize,
8094 cache: &mut Cache,
8095 rows_exact: bool,
8096 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8097 if mla.tp.is_some() {
8102 return Err(format!(
8103 "layer {il}: MLA layer is glm5-TP-sharded (MEMRA_GLM5_TP): the plain mixer \
8104 path is unwired for a head shard — only the TP decode/prime walk may \
8105 execute it (rows_exact={rows_exact})"
8106 )
8107 .into());
8108 }
8109 let max_ctx = cache.max_ctx;
8112 let layer = cache.latent[il].as_mut().ok_or_else(|| {
8113 format!(
8114 "layer {il} is Mixer::Mla but the cache has no latent plane — the ModelPlan \
8115 must declare StatePlan::LatentKvCache for it"
8116 )
8117 })?;
8118 let attn =
8119 self.mla_attn_cached_pre_wo(e, mla, h, pos_d, t, il, layer, max_ctx, rows_exact)?;
8120 if rows_exact {
8123 e.matmul_rows_exact(&mla.wo, &attn, t)
8124 } else {
8125 e.matmul(&mla.wo, &attn, t)
8126 }
8127 }
8128
8129 #[allow(clippy::too_many_arguments)]
8134 pub(crate) fn mla_attn_cached_pre_wo(
8135 &self,
8136 e: &Engine,
8137 mla: &crate::hybrid::MlaAttnLayer,
8138 h: &CudaSlice<f32>,
8139 pos_d: &CudaSlice<i32>,
8140 t: usize,
8141 il: usize,
8142 layer: &mut memra_kv::LatentKvLayer,
8143 max_ctx: usize,
8144 rows_exact: bool,
8145 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8146 let slot = layer.len;
8147 let width = layer.width;
8148 assert_eq!(
8149 width, mla.geom.latent_dim,
8150 "layer {il}: cache latent width {width} != MlaGeom latent_dim {}",
8151 mla.geom.latent_dim
8152 );
8153 let capacity = layer.rows.len() / width;
8154 if slot + t > capacity {
8155 return Err(format!(
8156 "layer {il}: latent cache overflow — {slot} + {t} rows exceeds capacity {capacity}"
8157 )
8158 .into());
8159 }
8160 if mla.index.is_some() && layer.index_rows.is_none() {
8161 return Err(format!(
8162 "layer {il} loaded a DSA k-pool indexer but its latent cache carries no indexer \
8163 state plane — StatePlan::LatentKvCache declared index_width 0 for a layer whose \
8164 SparseIndexPlan is Own {{ kpool: Some(..) }}"
8165 )
8166 .into());
8167 }
8168 let mut rows = std::mem::replace(&mut layer.rows, e.uninit(0)?);
8173 let mut index_rows = layer.index_rows.take();
8174 let mut pool_keys = layer.index_pool_keys.take();
8175 let mut pools_ready = layer.index_pools_ready;
8176 let index_ring_rows = layer.index_ring_rows.unwrap_or(0);
8177 let planes = index_rows.as_mut().map(|state| IndexerPlanes {
8178 state,
8179 pool_keys: &mut pool_keys,
8180 ready: &mut pools_ready,
8181 state_ring_rows: index_ring_rows,
8182 capacity_tokens: max_ctx,
8183 });
8184 let out =
8185 self.mla_attn_core_pre_wo(e, mla, h, pos_d, t, il, &mut rows, planes, slot, rows_exact);
8186 layer.rows = rows;
8187 layer.index_rows = index_rows;
8188 layer.index_pool_keys = pool_keys;
8189 layer.index_pools_ready = if out.is_ok() {
8194 pools_ready
8195 } else if let Some(indexer) = mla.index.as_ref() {
8196 pools_ready.min(layer.len / indexer.geom.pool)
8197 } else {
8198 pools_ready
8199 };
8200 let out = out?;
8201 if let Some(indexer) = mla.index.as_ref() {
8206 let pool = indexer.geom.pool;
8207 if layer.index_pool != 0 && layer.index_pool != pool {
8208 return Err(format!(
8209 "layer {il}: resident indexer pool {} != loaded geometry pool {pool}",
8210 layer.index_pool,
8211 )
8212 .into());
8213 }
8214 layer.index_pool = pool;
8215 }
8216 layer.len = slot + t;
8217 let len_i32 = i32::try_from(layer.len).map_err(|_| "latent length exceeds i32 mirror")?;
8218 e.i32_mirror_store(&mut layer.len_d, len_i32)?;
8221 Ok(out)
8222 }
8223
8224 #[allow(clippy::too_many_arguments)]
8232 pub(crate) fn mla_tp_attn_cached(
8233 &self,
8234 e: &Engine,
8235 mla: &crate::hybrid::MlaAttnLayer,
8236 h: &CudaSlice<f32>,
8237 pos_d: &CudaSlice<i32>,
8238 t: usize,
8239 il: usize,
8240 cache: &mut Cache,
8241 rows_exact: bool,
8242 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8243 let tp = mla
8244 .tp
8245 .as_ref()
8246 .ok_or("mla_tp_attn_cached called on an unsharded layer")?;
8247 let rt = &tp.rt;
8248 let ranks = tp.ranks();
8249 let g = mla.geom; let hl = g.n_head;
8251 let dv = g.d_v;
8252 let full_heads = tp.full_heads;
8253 let n_embd = tp.n_embd;
8254 let hh = n_embd / ranks;
8255 let max_ctx = cache.max_ctx;
8256
8257 let hop = rt.hop(e);
8261 let h_peers = crate::tp_transport::fanout_f32(&hop, h, h.len())?;
8262 let pos_peers = crate::tp_transport::fanout_i32(&hop, pos_d, pos_d.len())?;
8263
8264 {
8266 let canonical = cache.latent[il].as_ref().ok_or_else(|| {
8267 format!("layer {il}: glm5 TP MLA walk found no canonical latent plane")
8268 })?;
8269 crate::glm5_tp::ensure_mla_peer_latent(
8270 rt,
8271 canonical,
8272 &mut cache.glm5_tp_latent_peer[il],
8273 )?;
8274 }
8275
8276 let mut attn: Vec<Option<CudaSlice<f32>>> = (0..ranks).map(|_| None).collect();
8283 for r in 1..ranks {
8284 let layer = &mut cache.glm5_tp_latent_peer[il].as_mut().unwrap()[r - 1];
8285 attn[r] = Some(self.mla_attn_cached_pre_wo(
8286 &rt.peers[r - 1],
8287 &tp.peers[r - 1],
8288 &h_peers[r - 1],
8289 &pos_peers[r - 1],
8290 t,
8291 il,
8292 layer,
8293 max_ctx,
8294 rows_exact,
8295 )?);
8296 }
8297 attn[0] = {
8298 let layer = cache.latent[il].as_mut().unwrap();
8299 Some(self.mla_attn_cached_pre_wo(e, mla, h, pos_d, t, il, layer, max_ctx, rows_exact)?)
8300 };
8301
8302 let part = hl * dv;
8305 debug_assert_eq!(full_heads * dv, ranks * part);
8306 let attn_refs: Vec<&CudaSlice<f32>> = attn
8307 .iter()
8308 .map(|a| a.as_ref().expect("filled above"))
8309 .collect();
8310 let fulls = crate::tp_transport::gather_parts(&hop, &attn_refs, t, part)?;
8311
8312 let mut ys = Vec::with_capacity(ranks);
8315 if rows_exact {
8316 ys.push(e.matmul_rows_exact(&mla.wo, &fulls[0], t)?);
8317 for r in 1..ranks {
8318 ys.push(rt.peers[r - 1].matmul_rows_exact(&tp.peers[r - 1].wo, &fulls[r], t)?);
8319 }
8320 } else {
8321 ys.push(e.matmul(&mla.wo, &fulls[0], t)?);
8322 for r in 1..ranks {
8323 ys.push(rt.peers[r - 1].matmul(&tp.peers[r - 1].wo, &fulls[r], t)?);
8324 }
8325 }
8326 debug_assert_eq!(n_embd, ranks * hh);
8328 let y_refs: Vec<&CudaSlice<f32>> = ys.iter().collect();
8329 crate::tp_transport::concat_parts_on_root(&hop, &y_refs, t, hh)
8330 }
8331
8332 pub fn linear_attn(
8334 &self,
8335 e: &Engine,
8336 la: &LinearAttnLayer,
8337 h: &CudaSlice<f32>,
8338 t: usize,
8339 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8340 let cfg = &self.cfg;
8341 let _n_embd = cfg.n_embd as usize;
8342 let geometry = la.geometry;
8343 let d_state = geometry.key_head_dim as usize;
8344 let num_k = geometry.key_heads as usize;
8345 let num_v = geometry.value_heads as usize;
8346 let d_conv = geometry.conv_kernel as usize;
8347 let head_k = d_state;
8348 let head_v = geometry.value_head_dim as usize;
8349 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;
8353 let scale = 1.0 / (d_state as f32).sqrt();
8354
8355 let mut g4 = e.matmul_group(
8358 &[&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha],
8359 h,
8360 t,
8361 )?;
8362 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);
8374 let mut q_g = e.uninit(d_state * num_v * t)?;
8375 let mut k_g = e.uninit(d_state * num_v * t)?;
8376 let mut v_g = e.uninit(d_state * num_v * t)?;
8377 e.ssm_conv1d_gdn(
8378 &qkv_mixed,
8379 la.ssm_conv1d.float_data(),
8380 &mut q_g,
8381 &mut k_g,
8382 &mut v_g,
8383 conv_dim,
8384 t,
8385 d_conv,
8386 d_state,
8387 num_v,
8388 num_k,
8389 key_dim,
8390 )?;
8391 let mut q_l2 = e.uninit(d_state * num_v * t)?;
8393 e.l2_norm(&q_g, &mut q_l2, d_state, num_v * t, eps)?;
8394 let mut k_l2 = e.uninit(d_state * num_v * t)?;
8395 e.l2_norm(&k_g, &mut k_l2, d_state, num_v * t, eps)?;
8396 let v_gd = v_g;
8397
8398 let mut beta = e.uninit(t * num_v)?;
8401 e.sigmoid(&beta_raw, &mut beta, t * num_v)?;
8402 let mut g_log = e.uninit(t * num_v)?;
8404 e.gdn_glog(
8405 &alpha,
8406 la.ssm_dt.float_data(),
8407 la.ssm_a.float_data(),
8408 &mut g_log,
8409 num_v,
8410 t,
8411 )?;
8412
8413 let state_in = e.zeros(d_state * d_state * num_v)?; let mut state_out = e.zeros(d_state * d_state * num_v)?;
8416 let mut o = e.uninit(d_state * num_v * t)?;
8417 e.gdn_scan_prefill(
8418 &q_l2,
8419 &k_l2,
8420 &v_gd,
8421 &g_log,
8422 &beta,
8423 None,
8424 None,
8425 &state_in,
8426 &mut state_out,
8427 &mut o,
8428 num_v,
8429 t,
8430 scale,
8431 num_v,
8432 )?;
8433
8434 let mut gn = e.uninit(d_state * num_v * t)?;
8439 e.gated_rmsnorm(
8440 &o,
8441 la.ssm_norm.float_data(),
8442 &z,
8443 &mut gn,
8444 d_state,
8445 num_v * t,
8446 eps,
8447 )?;
8448
8449 let out = e.matmul(&la.ssm_out, &gn, t)?;
8453 Ok(out)
8454 }
8455}
8456
8457impl HybridModel {
8458 pub fn moe_ffn_il(
8469 &self,
8470 e: &Engine,
8471 m: &MoeWeights,
8472 z: &CudaSlice<f32>,
8473 t: usize,
8474 il: u16,
8475 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8476 Self::moe_ffn_inner(
8477 e,
8478 m,
8479 z,
8480 None,
8481 t,
8482 &self.cfg,
8483 il,
8484 self.max_moe_block(),
8485 false,
8486 None,
8487 self.uses_sliding_gated_moe_program(),
8488 false,
8489 )
8490 }
8491
8492 pub fn moe_ffn_il_prefill(
8495 &self,
8496 e: &Engine,
8497 m: &MoeWeights,
8498 z: &CudaSlice<f32>,
8499 t: usize,
8500 il: u16,
8501 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8502 Self::moe_ffn_inner(
8503 e,
8504 m,
8505 z,
8506 None,
8507 t,
8508 &self.cfg,
8509 il,
8510 self.max_moe_block(),
8511 true,
8512 Some(&self.step_grouped_prefill),
8513 self.uses_sliding_gated_moe_program(),
8514 false,
8515 )
8516 }
8517
8518 pub fn moe_ffn_il_zq8(
8522 &self,
8523 e: &Engine,
8524 m: &MoeWeights,
8525 z: &CudaSlice<f32>,
8526 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
8527 t: usize,
8528 il: u16,
8529 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8530 Self::moe_ffn_inner(
8531 e,
8532 m,
8533 z,
8534 zq8,
8535 t,
8536 &self.cfg,
8537 il,
8538 self.max_moe_block(),
8539 false,
8540 None,
8541 self.uses_sliding_gated_moe_program(),
8542 false,
8543 )
8544 }
8545
8546 pub(crate) fn moe_ffn_il_zq8_vrows(
8551 &self,
8552 e: &Engine,
8553 m: &MoeWeights,
8554 z: &CudaSlice<f32>,
8555 t: usize,
8556 il: u16,
8557 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8558 Self::moe_ffn_inner(
8559 e,
8560 m,
8561 z,
8562 None,
8563 t,
8564 &self.cfg,
8565 il,
8566 self.max_moe_block(),
8567 false,
8568 None,
8569 self.uses_sliding_gated_moe_program(),
8570 true,
8571 )
8572 }
8573
8574 pub(crate) fn moe_ffn(
8582 e: &Engine,
8583 m: &MoeWeights,
8584 z: &CudaSlice<f32>,
8585 t: usize,
8586 cfg: &ModelConfig,
8587 il: u16,
8588 max_block: usize,
8589 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8590 Self::moe_ffn_inner(
8591 e, m, z, None, t, cfg, il, max_block, false, None, false, false,
8592 )
8593 }
8594
8595 #[allow(clippy::too_many_arguments)]
8596 #[allow(clippy::map_entry)] pub(crate) fn moe_ffn_inner(
8598 e: &Engine,
8599 m: &MoeWeights,
8600 z: &CudaSlice<f32>,
8601 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
8602 t: usize,
8603 cfg: &ModelConfig,
8604 il: u16,
8605 max_block: usize,
8606 prefill: bool,
8607 grouped_prefill: Option<&std::sync::Mutex<crate::hybrid::StepEpGroupedPrefill>>,
8608 sliding_gated_moe: bool,
8609 vrows: bool,
8610 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8611 let worker_io = crate::spill_pread::worker_enabled();
8612 let epoch_lfu = std::env::var_os("MEMRA_MOE_LFU_DECAY").is_some();
8613 if Engine::moe_cache_enabled() && (worker_io || epoch_lfu) {
8614 e.with_moe_cache(max_block, |cache, _| {
8615 cache.begin_forward_epoch(il, t);
8616 if worker_io {
8617 cache.begin_worker_scope();
8618 }
8619 Ok(())
8620 })?;
8621 }
8622 if let Some(ep) = &m.glm5_ep {
8623 return Self::moe_ffn_glm5_ep(e, m, ep, z, zq8, t, cfg, il, prefill);
8628 }
8629 if m.step_ep.is_some() || m.step_tp.is_some() {
8630 let moe = cfg
8631 .moe
8632 .as_ref()
8633 .ok_or("Step distributed execution requires MoE model metadata")?;
8634 let n_embd = cfg.n_embd as usize;
8635 let n_expert = moe.expert_count as usize;
8636 let n_used = moe.expert_used_count as usize;
8637 let sigmoid = cfg
8638 .sigmoid_router()
8639 .ok_or("Step distributed execution requires the Step sigmoid router")?;
8640 let logits = Self::moe_router_logits(e, m, z, t, cfg)?;
8641 Self::trace_sigmoid_router_logits(e, il, t, n_expert, n_used, &logits, m, sigmoid)?;
8642 let grouped_prefill_requested = prefill && step_ep_grouped_prefill_enabled()?;
8643 if grouped_prefill_requested && !step_tp_prefill_enabled()? {
8644 return Err(
8645 "MEMRA_STEP_EP_GROUPED_PREFILL=1 requires MEMRA_STEP_TP_PREFILL=1".into(),
8646 );
8647 }
8648 if grouped_prefill_requested && !step_grouped_prefill_shape(true, prefill, t) {
8649 return Err(format!(
8650 "Step grouped prefill tokens {t} are outside the qualified {}..={} range",
8651 PRIME_MIN_T,
8652 crate::cache::PRIME_CHUNK_MAX_TOKENS,
8653 )
8654 .into());
8655 }
8656 let grouped_decode_shape = step_grouped_decode_shape(prefill, t);
8657 let grouped_prefill_shape =
8658 step_grouped_prefill_shape(grouped_prefill_requested, prefill, t);
8659 if let Some(ep) = m.step_ep.as_ref().filter(|ep| {
8660 ep.grouped_decode.is_some() && (grouped_decode_shape || grouped_prefill_shape)
8661 }) {
8662 let (selected, route_weights) =
8663 Self::moe_route_sigmoid_cfg(e, &logits, t, n_expert, n_used, m, sigmoid)?;
8664 crate::moesd::record_host_routes(il, n_expert, n_used, &selected)?;
8665 Self::trace_moe_routes(il, t, &selected, &route_weights)?;
8666 Self::trace_moe_input(e, il, t, n_embd, z)?;
8667 let selected = selected
8668 .iter()
8669 .map(|&expert| expert as usize)
8670 .collect::<Vec<_>>();
8671
8672 e.stream().synchronize()?;
8675 let execute = |state: &mut crate::hybrid::StepEpGroupedDecode| {
8676 state.projection.set_activation_limit(ep.activation_limit)?;
8677 ep.runtime
8678 .refresh_step_grouped_expert_parallel_gate_from_root_device(
8679 ep.experts.e4m3()?,
8680 &mut state.projection,
8681 z,
8682 t,
8683 &selected,
8684 )?;
8685 ep.runtime.refresh_step_grouped_expert_parallel_combine(
8686 &state.projection,
8687 &mut state.combine,
8688 &route_weights,
8689 )?;
8690 ep.runtime.execute_step_grouped_expert_parallel_gate(
8691 ep.experts.e4m3()?,
8692 &mut state.projection,
8693 )?;
8694 ep.runtime.execute_step_grouped_expert_parallel_combine(
8695 &state.projection,
8696 &mut state.combine,
8697 )?;
8698 let mut output = ep.runtime.copy_step_grouped_expert_parallel_combine_root(
8699 &state.projection,
8700 &state.combine,
8701 e,
8702 )?;
8703 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut output)?;
8704 if prefill {
8705 e.stream().synchronize()?;
8708 }
8709 eprintln!(
8710 "[step-tp-ep-grouped] execute layer={il} tokens={t} devices={:?} \
8711 attention_layout=tensor-parallel expert_layout=expert-parallel \
8712 expert_transport={} native_p2p=true route_control=host-narrow \
8713 input=root-device projection_workspaces=persistent \
8714 combine=root-device output=owning-stage-device \
8715 prefill={prefill} batched_decode=false capacity={} \
8716 performance_claim=false",
8717 ep.devices,
8718 ep.runtime.transport_label(),
8719 state.projection.max_tokens(),
8720 );
8721 Ok::<_, Box<dyn std::error::Error>>(output)
8722 };
8723
8724 if grouped_prefill_shape {
8725 let grouped_prefill = grouped_prefill
8726 .ok_or("Step grouped prefill has no model-scoped executor")?;
8727 let mut shared = grouped_prefill
8728 .lock()
8729 .map_err(|_| "Step grouped prefill state lock is poisoned")?;
8730 let needs_prepare = shared.state.as_ref().is_none_or(|state| {
8731 state.devices != ep.devices
8732 || state.grouped.projection.max_tokens() < t
8733 || state.grouped.projection.input_width() != n_embd
8734 || state.grouped.projection.expert_width()
8735 != moe.expert_ff_length as usize
8736 });
8737 if needs_prepare {
8738 let seed_input = vec![0.0f32; n_embd];
8739 let seed_selected = &selected[..n_used];
8740 let seed_weights = &route_weights[..n_used];
8741 let projection = ep
8742 .runtime
8743 .prepare_step_grouped_expert_parallel_gate_with_capacity(
8744 ep.experts.e4m3()?,
8745 &seed_input,
8746 1,
8747 seed_selected,
8748 ep.activation_limit,
8749 t,
8750 )?;
8751 let combine = ep.runtime.prepare_step_grouped_expert_parallel_combine(
8752 &projection,
8753 seed_weights,
8754 )?;
8755 shared.state = Some(crate::hybrid::StepEpGroupedPrefillState {
8756 devices: ep.devices.clone(),
8757 grouped: crate::hybrid::StepEpGroupedDecode {
8758 projection,
8759 combine,
8760 },
8761 });
8762 eprintln!(
8763 "[step-tp-ep-grouped-prefill] prepare capacity={t} devices={:?} \
8764 shared_across_layers=true performance_claim=false",
8765 ep.devices,
8766 );
8767 }
8768 return execute(
8769 &mut shared
8770 .state
8771 .as_mut()
8772 .expect("Step grouped prefill state prepared above")
8773 .grouped,
8774 );
8775 }
8776
8777 let mut grouped = ep
8778 .grouped_decode
8779 .as_ref()
8780 .expect("grouped decode presence checked above")
8781 .lock()
8782 .map_err(|_| "Step grouped decode state lock is poisoned")?;
8783 return execute(&mut grouped);
8784 }
8785 if grouped_prefill_shape {
8786 return Err(
8787 "Step grouped prefill requires native-P2P expert-owner device arithmetic"
8788 .into(),
8789 );
8790 }
8791 if t >= 16
8806 && crate::step_gemm_prime_on()
8807 && let Some(tp) = &m.step_tp
8808 && let crate::hybrid::StepTpExpertBank::Nvfp4(bank) = &tp.experts
8809 {
8810 let mprof = std::env::var("MEMRA_PRIME_PROF").as_deref() == Ok("1") && t >= 16;
8818 let mut mt = std::time::Instant::now();
8819 let (selected, route_weights) =
8820 Self::moe_route_sigmoid_cfg(e, &logits, t, n_expert, n_used, m, sigmoid)?;
8821 let sel_i32: Vec<i32> = selected.iter().map(|&x| x as i32).collect();
8822 let d_router = if mprof {
8823 let _ = e.stream().synchronize();
8824 let v = mt.elapsed().as_secs_f64() * 1e3;
8825 mt = std::time::Instant::now();
8826 v
8827 } else {
8828 0.0
8829 };
8830 let mdet =
8844 std::env::var("MEMRA_MOE_DETERM").as_deref() == Ok("1") && t >= 16 && il < 4;
8845 if mdet {
8846 let a = tp.runtime.run_tensor_parallel_routes_nvfp4_prime_grouped(
8847 bank,
8848 e,
8849 z,
8850 t,
8851 &sel_i32,
8852 &route_weights,
8853 n_used,
8854 tp.activation_limit,
8855 )?;
8856 let b = tp.runtime.run_tensor_parallel_routes_nvfp4_prime_grouped(
8857 bank,
8858 e,
8859 z,
8860 t,
8861 &sel_i32,
8862 &route_weights,
8863 n_used,
8864 tp.activation_limit,
8865 )?;
8866 let (ha, hb) = (e.dtoh(&a)?, e.dtoh(&b)?);
8867 let mut md = 0.0f32;
8868 let mut ndiff = 0usize;
8869 for (x, y) in ha.iter().zip(hb.iter()) {
8870 let d = (x - y).abs();
8871 if d > 0.0 {
8872 ndiff += 1;
8873 }
8874 if d > md {
8875 md = d;
8876 }
8877 }
8878 eprintln!(
8879 "[moe-determ] il={il} t={t} maxdiff={md:.3e} \
8880 differing={ndiff}/{} -> {}",
8881 ha.len(),
8882 if ndiff == 0 {
8883 "IDENTICAL"
8884 } else {
8885 "NONDETERMINISTIC"
8886 }
8887 );
8888 }
8889 let mut output = tp.runtime.run_tensor_parallel_routes_nvfp4_prime_grouped(
8890 bank,
8891 e,
8892 z,
8893 t,
8894 &sel_i32,
8895 &route_weights,
8896 n_used,
8897 tp.activation_limit,
8898 )?;
8899 let d_gemm = if mprof {
8900 let _ = e.stream().synchronize();
8901 let v = mt.elapsed().as_secs_f64() * 1e3;
8902 mt = std::time::Instant::now();
8903 v
8904 } else {
8905 0.0
8906 };
8907 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut output)?;
8908 if mprof {
8909 let _ = e.stream().synchronize();
8910 let d_shared = mt.elapsed().as_secs_f64() * 1e3;
8911 eprintln!(
8915 "[moe-prof] il={il} t={t} router={d_router:.1}ms \
8916 gemm={d_gemm:.1}ms shared={d_shared:.1}ms"
8917 );
8918 }
8919 return Ok(output);
8920 }
8921 if t == 1
8922 && crate::tp::step_nvfp4_dev_routes_enabled()?
8923 && crate::tp::step_tp_dev_router_enabled()?
8924 && let Some(tp) = &m.step_tp
8925 && let crate::hybrid::StepTpExpertBank::Nvfp4(bank) = &tp.experts
8926 {
8927 let (sf, route_norm) = sigmoid;
8928 static D1_ROUTER: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8935 let d1_router = *D1_ROUTER
8936 .get_or_init(|| std::env::var("MEMRA_DEV1_ROUTER").as_deref() == Ok("1"));
8937 if d1_router {
8938 let (sf_h, rn_h) = sigmoid;
8939 let n_ex = m.gate_exps.n_expert;
8940 let act_ct = m.active_count();
8941 let _ = tp.runtime.nvfp4_routes_prestage_with(
8942 bank,
8943 e,
8944 z,
8945 |rank1, in1, sel1, w1| {
8946 let mut guard = DEV1_ROUTER_REPS
8947 .lock()
8948 .map_err(|_| "dev1 router replica lock")?;
8949 let (reps, scratch) =
8950 guard.get_or_insert_with(|| (Default::default(), None));
8951 if !reps.contains_key(&il) {
8952 use cudarc::driver::DevicePtr;
8953 let (g1, p1, a1) = (
8954 rank1.htod(&vec![0.0f32; n_ex * n_embd])?,
8955 rank1.htod(&vec![0.0f32; n_ex])?,
8956 rank1.alloc_u8_uninit(n_ex)?,
8957 );
8958 for (src, dst_len, dst) in [
8959 (
8960 {
8961 let s = e.stream();
8962 let (p, _g) = m.gate_inp.float_data().device_ptr(&s);
8963 p
8964 },
8965 n_ex * n_embd * 4,
8966 {
8967 let s = rank1.stream();
8968 let (p, _g) = g1.device_ptr(&s);
8969 p
8970 },
8971 ),
8972 (
8973 {
8974 let s = e.stream();
8975 let (p, _g) = m.exp_probs_b_dev.device_ptr(&s);
8976 p
8977 },
8978 n_ex * 4,
8979 {
8980 let s = rank1.stream();
8981 let (p, _g) = p1.device_ptr(&s);
8982 p
8983 },
8984 ),
8985 (
8986 {
8987 let s = e.stream();
8988 let (p, _g) = m.active_experts_dev.device_ptr(&s);
8989 p
8990 },
8991 n_ex,
8992 {
8993 let s = rank1.stream();
8994 let (p, _g) = a1.device_ptr(&s);
8995 p
8996 },
8997 ),
8998 ] {
8999 crate::tp::raw_copy_bytes(dst, src, dst_len, rank1)?;
9000 }
9001 rank1.stream().synchronize()?;
9002 reps.insert(il, (g1, p1, a1));
9003 }
9004 if scratch.is_none() {
9005 *scratch = Some(rank1.htod(&vec![0.0f32; n_ex])?);
9006 }
9007 let (g1, p1, a1) = reps.get(&il).expect("armed above");
9008 let logits1 = scratch.as_mut().expect("armed above");
9009 rank1.router_gemv_into(g1, in1, logits1, n_embd, n_ex, 1)?;
9010 rank1.moe_router_sigmoid_topk_into(
9011 logits1, 1, n_ex, n_used, act_ct, p1, a1, sf_h, rn_h, sel1, w1,
9012 )?;
9013 Ok(true)
9014 },
9015 )?;
9016 } else {
9017 let _ = tp.runtime.nvfp4_routes_prestage(bank, e, z)?;
9018 }
9019 #[allow(clippy::type_complexity)] static SELW: std::sync::Mutex<Option<(usize, CudaSlice<i32>, CudaSlice<f32>)>> =
9024 std::sync::Mutex::new(None);
9025 let mut selw = SELW.lock().map_err(|_| "selw lock poisoned")?;
9026 if selw.as_ref().is_none_or(|(d, ..)| *d != e.ctx().ordinal()) {
9027 *selw = Some((
9028 e.ctx().ordinal(),
9029 e.htod_i32(&vec![0i32; n_used])?,
9030 e.htod(&vec![0.0f32; n_used])?,
9031 ));
9032 }
9033 let (_, sel_d, w_d) = selw.as_mut().expect("armed above");
9034 e.moe_router_sigmoid_topk_into(
9035 &logits,
9036 t,
9037 n_expert,
9038 n_used,
9039 m.active_count(),
9040 &m.exp_probs_b_dev,
9041 &m.active_experts_dev,
9042 sf,
9043 route_norm,
9044 sel_d,
9045 w_d,
9046 )?;
9047 crate::moesd::record_device_routes(e, il, n_expert, n_used, sel_d)?;
9048 if std::env::var("MEMRA_MOE_TRACE").is_ok()
9057 || std::env::var("MEMRA_MOE_WEIGHT_TRACE").is_ok()
9058 {
9059 return Err("MEMRA_MOE_TRACE/MEMRA_MOE_WEIGHT_TRACE cannot trace the \
9060 device-routed step TP walk (selection never returns to host; \
9061 tracing would add a new sync). Route through the host-router \
9062 arm — refused rather than silently dropping rows"
9063 .into());
9064 }
9065 static SHEXP_OV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9069 let shexp_ov = *SHEXP_OV
9070 .get_or_init(|| std::env::var("MEMRA_SHEXP_OVERLAP").as_deref() == Ok("1"));
9071 static SHEXP_D1: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9075 let shexp_d1 = *SHEXP_D1
9076 .get_or_init(|| std::env::var("MEMRA_SHEXP_DEV1").as_deref() == Ok("1"))
9077 && tp.runtime.rank_engine(1).is_some();
9078 static TAIL3: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9082 let tail3 =
9083 *TAIL3.get_or_init(|| std::env::var("MEMRA_TAIL_ADD3").as_deref() != Ok("0"));
9084 let mut ov_issued = false;
9085 let mut d1_issued = false;
9086 let mut tail_folded = false;
9087 let mut output = if shexp_d1 {
9088 let rank1 = tp.runtime.rank_engine(1).expect("checked above");
9089 tp.runtime
9090 .run_tensor_parallel_routes_nvfp4_device_routed_prejoin(
9091 bank,
9092 e,
9093 z,
9094 sel_d,
9095 w_d,
9096 n_used,
9097 tp.activation_limit,
9098 || {
9099 d1_issued =
9100 Self::shexp_dev1_issue(e, rank1, m, z, cfg, il, n_embd)?;
9101 Ok(())
9102 },
9103 )?
9104 } else if shexp_ov {
9105 let post_add = if tail3 {
9110 Self::shexp_overlap_tail_ptrs(e, m, cfg, n_embd)?
9111 } else {
9112 None
9113 };
9114 let used_post = post_add.is_some();
9115 let out = tp
9116 .runtime
9117 .run_tensor_parallel_routes_nvfp4_device_routed_prejoin_add3(
9118 bank,
9119 e,
9120 z,
9121 sel_d,
9122 w_d,
9123 n_used,
9124 tp.activation_limit,
9125 || {
9126 ov_issued = Self::shexp_overlap_issue(e, m, z, cfg, il, n_embd)?;
9127 Ok(())
9128 },
9129 post_add,
9130 )?;
9131 if used_post && ov_issued {
9136 tail_folded = true; }
9138 out
9139 } else {
9140 tp.runtime.run_tensor_parallel_routes_nvfp4_device_routed(
9141 bank,
9142 e,
9143 z,
9144 sel_d,
9145 w_d,
9146 n_used,
9147 tp.activation_limit,
9148 )?
9149 };
9150 if output.len() != t * n_embd {
9151 return Err(format!(
9152 "Step tp routed output has {} values, expected {t}x{n_embd}",
9153 output.len()
9154 )
9155 .into());
9156 }
9157 if tail_folded {
9158 } else if d1_issued {
9160 Self::shexp_dev1_apply(e, &mut output, n_embd)?;
9161 } else if ov_issued {
9162 Self::shexp_overlap_apply(e, &mut output, n_embd)?;
9163 } else {
9164 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut output)?;
9165 }
9166 static DR_LOGGED: std::sync::atomic::AtomicU64 =
9167 std::sync::atomic::AtomicU64::new(0);
9168 let layer_bit = 1u64 << (il as u64 % 64);
9169 if DR_LOGGED.fetch_or(layer_bit, std::sync::atomic::Ordering::Relaxed) & layer_bit
9170 == 0
9171 {
9172 eprintln!(
9173 "[step-tp] execute layer={il} tokens={t} devices={:?} \
9174 expert_transport={} native_p2p={} router=device \
9175 activation=host-canonical accumulation=host-canonical \
9176 output=e-device io=device performance_claim=false \
9177 (logged once per layer)",
9178 tp.devices,
9179 tp.runtime.transport_label(),
9180 tp.runtime.native_p2p(),
9181 );
9182 }
9183 return Ok(output);
9184 }
9185 let automatic_ep_device_router = crate::tp::parallel_ep_device_router_enabled()?;
9186 let automatic_ep_q8_act = crate::tp::parallel_ep_q8_act_enabled()?;
9187 let automatic_ep_q8_scope = crate::tp::parallel_ep_q8_scope()?;
9188 crate::tp::parallel_ep_q8_gu_paired_enabled(
9189 automatic_ep_q8_act,
9190 automatic_ep_q8_scope,
9191 )?;
9192 let automatic_ep_q8_active =
9193 automatic_ep_q8_act && t <= crate::tp::NVFP4_EP_Q8_BATCH_CAP;
9194 if automatic_ep_q8_scope.is_some() && !automatic_ep_q8_act {
9195 return Err(
9196 "MEMRA_PARALLEL_EP_Q8_SCOPE requires MEMRA_PARALLEL_EP_Q8_ACT=1".into(),
9197 );
9198 }
9199 if automatic_ep_q8_act && !automatic_ep_device_router {
9200 return Err(
9201 "MEMRA_PARALLEL_EP_Q8_ACT=1 requires MEMRA_PARALLEL_EP_DEVICE_ROUTER=1".into(),
9202 );
9203 }
9204 if automatic_ep_q8_act && m.step_ep.as_ref().is_none_or(|ep| !ep.nvfp4_device_routes) {
9205 return Err(
9206 "MEMRA_PARALLEL_EP_Q8_ACT=1 requires automatic W4A16 whole-expert EP".into(),
9207 );
9208 }
9209 if t <= crate::tp::NVFP4_EP_DEVICE_ROUTER_BATCH_CAP
9210 && automatic_ep_device_router
9211 && let Some(ep) = &m.step_ep
9212 && ep.nvfp4_device_routes
9213 {
9214 let bank = match &ep.experts {
9215 crate::hybrid::StepEpExpertBank::Nvfp4(bank) => bank,
9216 crate::hybrid::StepEpExpertBank::E4m3(_) => {
9217 return Err("W4A16 device-routed EP reached an E4M3 expert bank".into());
9218 }
9219 };
9220 let pairs = t
9221 .checked_mul(n_used)
9222 .ok_or("W4A16 device-routed EP pair count overflow")?;
9223 let capacity = crate::tp::NVFP4_EP_DEVICE_BATCH_CAP * n_used;
9224 type EpSelwByDevice =
9228 std::collections::HashMap<usize, (usize, CudaSlice<i32>, CudaSlice<f32>)>;
9229 static EP_SELW: std::sync::Mutex<Option<EpSelwByDevice>> =
9230 std::sync::Mutex::new(None);
9231 let mut selw = EP_SELW
9232 .lock()
9233 .map_err(|_| "automatic EP device-router workspace lock poisoned")?;
9234 let device = e.ctx().ordinal();
9235 let workspaces = selw.get_or_insert_with(Default::default);
9236 if workspaces
9237 .get(&device)
9238 .is_none_or(|(cap, ..)| *cap < capacity)
9239 {
9240 workspaces.insert(
9241 device,
9242 (
9243 capacity,
9244 e.htod_i32(&vec![0i32; capacity])?,
9245 e.htod(&vec![0.0f32; capacity])?,
9246 ),
9247 );
9248 }
9249 let (_, sel_d, w_d) = workspaces.get_mut(&device).expect("armed above");
9250 let (sf, route_norm) = sigmoid;
9251 e.moe_router_sigmoid_topk_into(
9252 &logits,
9253 t,
9254 n_expert,
9255 n_used,
9256 m.active_count(),
9257 &m.exp_probs_b_dev,
9258 &m.active_experts_dev,
9259 sf,
9260 route_norm,
9261 sel_d,
9262 w_d,
9263 )?;
9264 crate::moesd::record_device_routes(e, il, n_expert, n_used, sel_d)?;
9265 static SHEXP_OV_AUTO: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9266 let shexp_ov = t == 1
9267 && *SHEXP_OV_AUTO
9268 .get_or_init(|| std::env::var("MEMRA_SHEXP_OVERLAP").as_deref() == Ok("1"));
9269 let mut ov_issued = false;
9270 let mut output = if shexp_ov {
9271 ep.runtime
9272 .run_routed_experts_nvfp4_w4a16_device_routed_prejoin(
9273 bank,
9274 e,
9275 z,
9276 sel_d,
9277 w_d,
9278 t,
9279 n_used,
9280 ep.activation_limit,
9281 || {
9282 ov_issued = Self::shexp_overlap_issue(e, m, z, cfg, il, n_embd)?;
9283 Ok(())
9284 },
9285 )?
9286 } else {
9287 ep.runtime.run_routed_experts_nvfp4_w4a16_device_routed(
9288 bank,
9289 e,
9290 z,
9291 sel_d,
9292 w_d,
9293 t,
9294 n_used,
9295 ep.activation_limit,
9296 )?
9297 };
9298 if output.len() != t * n_embd {
9299 return Err(format!(
9300 "W4A16 device-routed EP output has {} values, expected \
9301 {t}x{n_embd}={}",
9302 output.len(),
9303 t * n_embd,
9304 )
9305 .into());
9306 }
9307 if ov_issued {
9308 Self::shexp_overlap_apply(e, &mut output, n_embd)?;
9309 } else {
9310 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut output)?;
9311 }
9312 static DEVICE_ROUTER_LOGGED: std::sync::atomic::AtomicU64 =
9313 std::sync::atomic::AtomicU64::new(0);
9314 let layer_bit = 1u64 << (il as u64 % 64);
9315 if DEVICE_ROUTER_LOGGED.fetch_or(layer_bit, std::sync::atomic::Ordering::Relaxed)
9316 & layer_bit
9317 == 0
9318 {
9319 eprintln!(
9320 "[parallel-ep] execute layer={il} tokens={t} devices={:?} \
9321 router=device expert_transport={} native_p2p={} \
9322 activation=bf16-rounded accumulation={} output=e-device \
9323 performance_claim=false (logged once per layer)",
9324 ep.devices,
9325 ep.runtime.transport_label(),
9326 ep.runtime.native_p2p(),
9327 if automatic_ep_q8_active {
9328 "token-slot-order-q8"
9329 } else {
9330 "token-slot-order"
9331 },
9332 );
9333 }
9334 debug_assert!(pairs <= capacity);
9335 return Ok(output);
9336 }
9337 static ROUTE_NS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
9341 static ROUTE_CALLS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
9342 let route_timing = std::env::var("MEMRA_STEP_TP_TIMING").as_deref() == Ok("1");
9343 let route_started = route_timing.then(std::time::Instant::now);
9344 let (selected, route_weights, input) = Self::moe_route_sigmoid_with_input(
9345 e,
9346 &logits,
9347 z,
9348 t,
9349 n_embd,
9350 n_expert,
9351 n_used,
9352 m.exp_probs_b.as_deref(),
9353 sigmoid,
9354 m.active_experts.as_deref(),
9355 )?;
9356 if let Some(started) = route_started {
9357 use std::sync::atomic::Ordering;
9358 let ns = ROUTE_NS.fetch_add(started.elapsed().as_nanos() as u64, Ordering::Relaxed)
9359 + started.elapsed().as_nanos() as u64;
9360 let calls = ROUTE_CALLS.fetch_add(1, Ordering::Relaxed) + 1;
9361 if calls.is_multiple_of(430) {
9362 eprintln!(
9363 "[moe-route-timing] calls={calls} total_ms={:.1} avg_us={:.1}",
9364 ns as f64 / 1.0e6,
9365 ns as f64 / calls as f64 / 1.0e3,
9366 );
9367 }
9368 }
9369 crate::moesd::record_host_routes(il, n_expert, n_used, &selected)?;
9370 Self::trace_moe_routes(il, t, &selected, &route_weights)?;
9371 Self::trace_moe_input(e, il, t, n_embd, z)?;
9372 let selected = selected
9373 .iter()
9374 .map(|&expert| expert as usize)
9375 .collect::<Vec<_>>();
9376 if t <= crate::tp::NVFP4_EP_DEVICE_BATCH_CAP
9377 && let Some(ep) = &m.step_ep
9378 && ep.nvfp4_device_routes
9379 {
9380 let bank = match &ep.experts {
9381 crate::hybrid::StepEpExpertBank::Nvfp4(bank) => bank,
9382 crate::hybrid::StepEpExpertBank::E4m3(_) => {
9383 return Err("W4A16 NVFP4 device EP reached an E4M3 expert bank".into());
9384 }
9385 };
9386 let mut output = ep.runtime.run_routed_experts_nvfp4_w4a16_device_io(
9387 bank,
9388 e,
9389 z,
9390 t,
9391 &selected,
9392 &route_weights,
9393 n_used,
9394 ep.activation_limit,
9395 )?;
9396 if output.len() != t * n_embd {
9397 return Err(format!(
9398 "W4A16 NVFP4 EP routed output has {} values, expected \
9399 {t}x{n_embd}={}",
9400 output.len(),
9401 t * n_embd,
9402 )
9403 .into());
9404 }
9405 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut output)?;
9406 static W4A16_EP_LOGGED: std::sync::atomic::AtomicU64 =
9407 std::sync::atomic::AtomicU64::new(0);
9408 let layer_bit = 1u64 << (il as u64 % 64);
9409 if W4A16_EP_LOGGED.fetch_or(layer_bit, std::sync::atomic::Ordering::Relaxed)
9410 & layer_bit
9411 == 0
9412 {
9413 eprintln!(
9414 "[step-ep] execute layer={il} tokens={t} devices={:?} \
9415 expert_transport={} native_p2p={} activation=bf16-rounded \
9416 accumulation={} output=e-device \
9417 performance_claim=false (logged once per layer)",
9418 ep.devices,
9419 ep.runtime.transport_label(),
9420 ep.runtime.native_p2p(),
9421 if t == 1 {
9422 "owner-grouped-rank-order"
9423 } else {
9424 "token-slot-order"
9425 },
9426 );
9427 }
9428 return Ok(output);
9429 }
9430 if t == 1
9435 && crate::tp::step_nvfp4_dev_routes_enabled()?
9436 && let Some(tp) = &m.step_tp
9437 && let crate::hybrid::StepTpExpertBank::Nvfp4(bank) = &tp.experts
9438 {
9439 let mut output = tp.runtime.run_tensor_parallel_routes_nvfp4_device_io(
9440 bank,
9441 e,
9442 z,
9443 &selected,
9444 &route_weights,
9445 n_used,
9446 tp.activation_limit,
9447 )?;
9448 if output.len() != t * n_embd {
9449 return Err(format!(
9450 "Step tp routed output has {} values, expected {t}x{n_embd}",
9451 output.len()
9452 )
9453 .into());
9454 }
9455 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut output)?;
9456 static IO_LOGGED: std::sync::atomic::AtomicU64 =
9457 std::sync::atomic::AtomicU64::new(0);
9458 let layer_bit = 1u64 << (il as u64 % 64);
9459 if IO_LOGGED.fetch_or(layer_bit, std::sync::atomic::Ordering::Relaxed) & layer_bit
9460 == 0
9461 {
9462 eprintln!(
9463 "[step-tp] execute layer={il} tokens={t} devices={:?} \
9464 expert_transport={} native_p2p={} activation=host-canonical \
9465 accumulation=host-canonical output=e-device io=device \
9466 performance_claim=false (logged once per layer)",
9467 tp.devices,
9468 tp.runtime.transport_label(),
9469 tp.runtime.native_p2p(),
9470 );
9471 }
9472 return Ok(output);
9473 }
9474 let (routed, mode, devices, transport, native_p2p) = if let Some(tp) = &m.step_tp {
9475 (
9476 match &tp.experts {
9477 crate::hybrid::StepTpExpertBank::E4m3(bank) => {
9478 tp.runtime.run_tensor_parallel_routes(
9479 bank,
9480 &input,
9481 t,
9482 &selected,
9483 &route_weights,
9484 n_used,
9485 )?
9486 }
9487 crate::hybrid::StepTpExpertBank::Nvfp4(bank) => {
9488 if t == 1 && crate::tp::step_nvfp4_dev_routes_enabled()? {
9489 tp.runtime.run_tensor_parallel_routes_nvfp4_device(
9490 bank,
9491 &input,
9492 &selected,
9493 &route_weights,
9494 n_used,
9495 tp.activation_limit,
9496 )?
9497 } else {
9498 tp.runtime.run_tensor_parallel_routes_nvfp4(
9499 bank,
9500 &input,
9501 t,
9502 &selected,
9503 &route_weights,
9504 n_used,
9505 tp.activation_limit,
9506 )?
9507 }
9508 }
9509 },
9510 "tp",
9511 &tp.devices,
9512 tp.runtime.transport_label(),
9513 tp.runtime.native_p2p(),
9514 )
9515 } else {
9516 let ep = m
9517 .step_ep
9518 .as_ref()
9519 .ok_or("Step distributed runtime has no EP or TP state")?;
9520 (
9521 match &ep.experts {
9522 crate::hybrid::StepEpExpertBank::E4m3(bank) => {
9523 ep.runtime.run_routed_experts(
9524 bank,
9525 &input,
9526 t,
9527 &selected,
9528 &route_weights,
9529 n_used,
9530 ep.activation_limit,
9531 )?
9532 }
9533 crate::hybrid::StepEpExpertBank::Nvfp4(bank) => {
9534 ep.runtime.run_routed_experts_nvfp4(
9535 bank,
9536 &input,
9537 t,
9538 &selected,
9539 &route_weights,
9540 n_used,
9541 ep.activation_limit,
9542 )?
9543 }
9544 },
9545 if ep.configured_by_tp { "tp-ep" } else { "ep" },
9546 &ep.devices,
9547 ep.runtime.transport_label(),
9548 ep.runtime.native_p2p(),
9549 )
9550 };
9551 if routed.len() != t * n_embd {
9552 return Err(format!(
9553 "Step {mode} routed output has {} values, expected {t}x{n_embd}",
9554 routed.len()
9555 )
9556 .into());
9557 }
9558 let mut output = e.htod(&routed)?;
9559 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut output)?;
9560 static STEP_LOGGED: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
9563 let layer_bit = 1u64 << (il as u64 % 64);
9564 if STEP_LOGGED.fetch_or(layer_bit, std::sync::atomic::Ordering::Relaxed) & layer_bit
9565 == 0
9566 {
9567 eprintln!(
9568 "[step-{mode}] execute layer={il} tokens={t} devices={devices:?} \
9569 expert_transport={transport} native_p2p={native_p2p} \
9570 activation={} accumulation={} output={} \
9571 performance_claim=false (logged once per layer)",
9572 if let Some(ep) = &m.step_ep {
9573 ep.runtime.expert_activation_label()
9574 } else {
9575 "host-canonical"
9576 },
9577 if let Some(ep) = &m.step_ep {
9578 ep.runtime.expert_accumulation_label()
9579 } else {
9580 "host-canonical"
9581 },
9582 if let Some(ep) = &m.step_ep {
9583 ep.runtime.expert_output_label()
9584 } else {
9585 "host-accumulated"
9586 },
9587 );
9588 if let Some(ep) = &m.step_ep
9589 && let Some(limit) = ep.activation_limit
9590 {
9591 eprintln!(
9592 "[step-ep-clamp] execute layer={il} tokens={t} routed_clamp={limit} \
9593 formula=min-silu-times-clamped-up performance_claim=false"
9594 );
9595 }
9596 }
9597 return Ok(output);
9598 }
9599 if Self::sigmoid_resident_dev_eligible(e, m, cfg, sliding_gated_moe) {
9600 let moe = cfg.moe.as_ref().unwrap();
9601 let n_expert = moe.expert_count as usize;
9602 let n_used = moe.expert_used_count as usize;
9603 let sigmoid = cfg.sigmoid_router().unwrap();
9604 let logits = Self::moe_router_logits(e, m, z, t, cfg)?;
9605 Self::trace_sigmoid_router_logits(e, il, t, n_expert, n_used, &logits, m, sigmoid)?;
9606 return Self::moe_ffn_sigmoid_dev(e, m, z, zq8, &logits, t, cfg, il, sigmoid);
9607 }
9608 if prefill && t > MOE_DEV_MAX_T && cfg.sigmoid_router().is_some() && cfg.glm5.is_some() {
9630 static GPF_ANNOUNCED: std::sync::atomic::AtomicU8 = std::sync::atomic::AtomicU8::new(0);
9633 let enabled = moe_grouped_prefill_enabled();
9634 let bit = 1u8 << u8::from(enabled);
9635 if GPF_ANNOUNCED.fetch_or(bit, std::sync::atomic::Ordering::Relaxed) & bit == 0 {
9636 eprintln!(
9637 "[moe-grouped-prefill] flag={} t={t} il={il} (announce printed in both \
9638 arms; engagement is the per-layer execute line + the dispatch counter)",
9639 if enabled { "on" } else { "off" },
9640 );
9641 }
9642 if enabled
9643 && let Some(out) = Self::moe_ffn_grouped_prefill_sigmoid(e, m, z, t, cfg, il)?
9644 {
9645 return Ok(out);
9646 }
9647 }
9648 if t > 1 && moe_grouped_enabled(cfg, prefill) {
9651 let grouped_out = Self::moe_ffn_grouped(e, m, z, t, cfg, il, max_block)?;
9652 if std::env::var("MEMRA_MOE_GATE").is_ok() {
9657 let seq_out = Self::moe_ffn_sequential(e, m, z, t, cfg, il, max_block)?;
9658 let g_host = e.dtoh(&grouped_out)?;
9659 let s_host = e.dtoh(&seq_out)?;
9660 let g_bytes: &[u8] = unsafe {
9661 std::slice::from_raw_parts(g_host.as_ptr() as *const u8, g_host.len() * 4)
9662 };
9663 let s_bytes: &[u8] = unsafe {
9664 std::slice::from_raw_parts(s_host.as_ptr() as *const u8, s_host.len() * 4)
9665 };
9666 if g_bytes == s_bytes {
9667 println!("moe-gate il={il} t={t} BYTE-IDENTICAL");
9668 } else {
9669 let diffs = g_host
9670 .iter()
9671 .zip(s_host.iter())
9672 .enumerate()
9673 .filter(|(_, (a, b))| a != b)
9674 .count();
9675 let maxdiff = g_host
9676 .iter()
9677 .zip(s_host.iter())
9678 .map(|(a, b)| (a - b).abs())
9679 .fold(0.0f32, f32::max);
9680 panic!(
9681 "moe-gate il={il} t={t} MISMATCH: {diffs}/{} elems differ, maxdiff={maxdiff:.6e}",
9682 g_host.len()
9683 );
9684 }
9685 }
9686 return Ok(grouped_out);
9687 }
9688 Self::moe_ffn_sequential_zq8(e, m, z, zq8, t, cfg, il, max_block, vrows)
9689 }
9690
9691 fn sigmoid_resident_dev_eligible(
9692 e: &Engine,
9693 m: &MoeWeights,
9694 cfg: &ModelConfig,
9695 sliding_gated_moe: bool,
9696 ) -> bool {
9697 let Some(moe) = cfg.moe.as_ref() else {
9698 return false;
9699 };
9700 static OBSERVATION_MODE: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9703 let observation_mode = *OBSERVATION_MODE.get_or_init(|| {
9704 std::env::var("MEMRA_MOE_STATS").is_ok()
9705 || std::env::var("MEMRA_MOE_TRACE").is_ok()
9706 || std::env::var("MEMRA_MOE_WEIGHT_TRACE").is_ok()
9707 || std::env::var("MEMRA_MOE_INPUT_TRACE_DIR").is_ok()
9708 || std::env::var("MEMRA_MOE_GATE").is_ok()
9709 });
9710 let resident_layout_supported = m.dev_exps.as_ref().is_some_and(|dev| {
9711 if dev.dev != e.ctx().ordinal() {
9712 return false;
9713 }
9714 let q8 = moe_q8_enabled_for_model(cfg, m);
9715 let fp8 = dev.fp8_blk.is_some()
9716 && m.gate_exps.qtype == crate::QT_F8_E4M3_BLK
9717 && m.up_exps.qtype == crate::QT_F8_E4M3_BLK
9718 && m.down_exps.qtype == crate::QT_F8_E4M3_BLK;
9719 q8 || fp8
9720 });
9721 sliding_gated_moe
9722 && sigmoid_router_enabled()
9723 && moe_dev_enabled()
9724 && moe_slab_enabled()
9725 && !observation_mode
9726 && moe.expert_used_count <= 8
9727 && m.has_uniform_expert_layout()
9728 && m.gate_exps.macros.is_none()
9729 && m.up_exps.macros.is_none()
9730 && m.down_exps.macros.is_none()
9731 && !m.has_macros
9732 && resident_layout_supported
9733 }
9734
9735 pub(crate) fn moe_ffn_sequential(
9737 e: &Engine,
9738 m: &MoeWeights,
9739 z: &CudaSlice<f32>,
9740 t: usize,
9741 cfg: &ModelConfig,
9742 il: u16,
9743 max_block: usize,
9744 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9745 Self::moe_ffn_sequential_zq8(e, m, z, None, t, cfg, il, max_block, false)
9746 }
9747
9748 fn moe_router_logits(
9752 e: &Engine,
9753 m: &MoeWeights,
9754 z: &CudaSlice<f32>,
9755 t: usize,
9756 cfg: &ModelConfig,
9757 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9758 if t < PRIME_MIN_T {
9759 if crate::router_kernel_on() {
9761 e.router_gemv(
9762 m.gate_inp.float_data(),
9763 z,
9764 cfg.n_embd as usize,
9765 m.gate_exps.n_expert,
9766 t,
9767 )
9768 } else {
9769 e.matmul_decode_exact(&m.gate_inp, z, t)
9770 }
9771 } else if crate::router_prefill_exact_on() && crate::router_kernel_on() {
9772 e.router_gemv(
9773 m.gate_inp.float_data(),
9774 z,
9775 cfg.n_embd as usize,
9776 m.gate_exps.n_expert,
9777 t,
9778 )
9779 } else {
9780 e.matmul(&m.gate_inp, z, t)
9781 }
9782 }
9783
9784 fn trace_moe_routes(
9792 il: u16,
9793 t: usize,
9794 sel_all: &[u32],
9795 weights: &[f32],
9796 ) -> Result<(), Box<dyn std::error::Error>> {
9797 use std::io::Write as _;
9798 if let Ok(path) = std::env::var("MEMRA_MOE_TRACE") {
9799 let mut f = std::fs::OpenOptions::new()
9800 .create(true)
9801 .append(true)
9802 .open(path)?;
9803 let ids: Vec<String> = sel_all.iter().map(|s| s.to_string()).collect();
9804 writeln!(f, "{} {} {}", il, t, ids.join(","))?;
9805 }
9806 if let Ok(path) = std::env::var("MEMRA_MOE_WEIGHT_TRACE") {
9807 let mut f = std::fs::OpenOptions::new()
9808 .create(true)
9809 .append(true)
9810 .open(path)?;
9811 let pairs: Vec<String> = sel_all
9812 .iter()
9813 .zip(weights)
9814 .map(|(expert, weight)| format!("{expert}:{weight:.9}"))
9815 .collect();
9816 writeln!(f, "{} {} {}", il, t, pairs.join(","))?;
9817 }
9818 Ok(())
9819 }
9820
9821 #[allow(clippy::too_many_arguments)]
9822 fn trace_sigmoid_router_logits(
9823 e: &Engine,
9824 il: u16,
9825 t: usize,
9826 n_expert: usize,
9827 n_used: usize,
9828 logits: &CudaSlice<f32>,
9829 m: &MoeWeights,
9830 (scaling_factor, route_norm): (f32, bool),
9831 ) -> Result<(), Box<dyn std::error::Error>> {
9832 if !crate::sigrouter_contract::served_logit_trace_enabled() || t != 1 {
9833 return Ok(());
9834 }
9835 let logits = e.dtoh(logits)?;
9836 let active: Vec<u8> = m
9837 .active_experts
9838 .as_ref()
9839 .map(|mask| mask.iter().map(|&enabled| u8::from(enabled)).collect())
9840 .unwrap_or_else(|| vec![1; n_expert]);
9841 let bias = m.exp_probs_b.clone().unwrap_or_else(|| vec![0.0; n_expert]);
9842 crate::sigrouter_contract::capture_served_logits(
9843 il as u32,
9844 t,
9845 n_expert,
9846 n_used,
9847 scaling_factor,
9848 route_norm,
9849 &active,
9850 &bias,
9851 &logits,
9852 )?;
9853 Ok(())
9854 }
9855
9856 fn trace_moe_input(
9861 e: &Engine,
9862 il: u16,
9863 t: usize,
9864 n_embd: usize,
9865 z: &CudaSlice<f32>,
9866 ) -> Result<(), Box<dyn std::error::Error>> {
9867 use std::io::Write as _;
9868 let Ok(dir) = std::env::var("MEMRA_MOE_INPUT_TRACE_DIR") else {
9869 return Ok(());
9870 };
9871 let values = active_matrix_values(z.len(), t, n_embd, "MoE input trace activation")?;
9872 let host = e.dtoh_view(&z.slice(0..values))?;
9873 let bytes = unsafe {
9874 std::slice::from_raw_parts(
9875 host.as_ptr().cast::<u8>(),
9876 host.len() * std::mem::size_of::<f32>(),
9877 )
9878 };
9879 let state = MOE_INPUT_TRACE_WRITER.get_or_init(|| std::sync::Mutex::new(None));
9880 let mut state = state
9881 .lock()
9882 .map_err(|_| "MoE input trace writer lock is poisoned")?;
9883 if state.is_none() {
9884 let dir = std::path::PathBuf::from(&dir);
9885 std::fs::create_dir_all(&dir)?;
9886 let index = std::fs::OpenOptions::new()
9887 .create(true)
9888 .append(true)
9889 .open(dir.join("index.jsonl"))?;
9890 *state = Some(MoeInputTraceWriter {
9891 dir,
9892 index,
9893 payloads: std::collections::HashMap::new(),
9894 });
9895 }
9896 let writer = state.as_mut().unwrap();
9897 if writer.dir != std::path::Path::new(&dir) {
9898 return Err("MEMRA_MOE_INPUT_TRACE_DIR changed after capture started".into());
9899 }
9900 let file_name = format!("layer-{il:03}.f32");
9901 if !writer.payloads.contains_key(&il) {
9902 let payload = std::fs::OpenOptions::new()
9903 .create(true)
9904 .append(true)
9905 .open(writer.dir.join(&file_name))?;
9906 let offset = payload.metadata()?.len();
9907 writer.payloads.insert(il, (payload, offset));
9908 }
9909 let (payload, offset) = writer.payloads.get_mut(&il).unwrap();
9910 let row_offset = *offset;
9911 payload.write_all(bytes)?;
9912 *offset += bytes.len() as u64;
9913 writeln!(
9914 writer.index,
9915 "{{\"format\":\"memra-moe-input-trace-v1\",\"layer\":{il},\"tokens\":{t},\
9916 \"hidden_size\":{n_embd},\"file\":\"{file_name}\",\"offset\":{row_offset},\
9917 \"payload_bytes\":{}}}",
9918 bytes.len()
9919 )?;
9920 Ok(())
9921 }
9922
9923 #[allow(clippy::too_many_arguments)]
9924 #[allow(clippy::too_many_arguments)]
9925 pub(crate) fn moe_ffn_sequential_zq8(
9927 e: &Engine,
9928 m: &MoeWeights,
9929 z: &CudaSlice<f32>,
9930 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
9931 t: usize,
9932 cfg: &ModelConfig,
9933 il: u16,
9934 max_block: usize,
9935 vrows: bool,
9936 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9937 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
9938 let moe = cfg.moe.as_ref().unwrap();
9939 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);
9946 debug_assert_eq!(m.gate_exps.out_f, n_ff_exp);
9947 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);
9950
9951 let lim_exp = cfg.clamp_exp_at(il as u32);
9954 let lim_shexp = cfg.clamp_shexp_at(il as u32);
9955 let use_cache = Engine::moe_cache_enabled();
9956 let uniform_experts = m.has_uniform_expert_layout();
9957 let moe_q8 = uniform_experts && moe_q8_enabled_for_model(cfg, m);
9958 let cpu_expert_requested = crate::cpu_experts::configured();
9965 if cpu_expert_requested && (cfg.hy3.is_none() || cfg.m3.is_some()) {
9966 return Err(std::io::Error::other(
9967 "MEMRA_CPU_EXPERT_LIB is experimental and currently gated to Hy3",
9968 )
9969 .into());
9970 }
9971 let cpu_hybrid = cpu_expert_requested && t < PRIME_MIN_T && m.dev_exps.is_none();
9972 let freeze_cpu_residency = cpu_expert_requested
9978 && std::env::var("MEMRA_CPU_EXPERT_FREEZE_CACHE").as_deref() == Ok("1");
9979 let caller_warms_before_freeze = std::env::var("MEMRA_CPU_EXPERT_FREEZE_WARMUP_TOKENS")
9980 .ok()
9981 .and_then(|value| value.parse::<usize>().ok())
9982 .is_some_and(|tokens| tokens > 0);
9983 if cpu_hybrid && freeze_cpu_residency && !caller_warms_before_freeze {
9984 e.freeze_moe_cache();
9985 }
9986 let cache_frozen = use_cache && e.moe_cache_frozen();
9987 let cache_dispatch = use_cache && (!cache_frozen || cpu_hybrid);
9988
9989 let logits = Self::moe_router_logits(e, m, z, t, cfg)?;
9992 if let Some(sig) = cfg.sigmoid_router() {
9993 Self::trace_sigmoid_router_logits(e, il, t, n_expert, n_used, &logits, m, sig)?;
9994 }
9995
9996 let no_exp_macros = m.gate_exps.macros.is_none()
10035 && m.up_exps.macros.is_none()
10036 && m.down_exps.macros.is_none();
10037 if cfg.sigmoid_router().is_none()
10041 && cfg.m3.is_none()
10042 && cfg.hy3.is_none()
10043 && !cfg.swiglu_clamped_at(il as u32)
10044 && no_exp_macros
10045 && t > MOE_DEV_MAX_T
10049 && m.dev_exps.is_some()
10050 && moe_q8_enabled_for_model(cfg, m)
10051 && std::env::var("MEMRA_MOE_PAIRS")
10052 .map(|v| v != "0")
10053 .unwrap_or(true)
10054 && std::env::var("MEMRA_MOE_STATS").is_err()
10055 {
10056 return Self::moe_ffn_pairs(e, m, z, &logits, t, cfg);
10057 }
10058
10059 let dev_ok = uniform_experts
10077 && cfg.sigmoid_router().is_none()
10078 && cfg.m3.is_none()
10079 && cfg.hy3.is_none()
10080 && !cfg.swiglu_clamped_at(il as u32);
10081 let observe_routes = std::env::var("MEMRA_MOE_STATS").is_ok()
10085 || std::env::var("MEMRA_MOE_TRACE").is_ok()
10086 || std::env::var("MEMRA_MOE_WEIGHT_TRACE").is_ok()
10087 || std::env::var("MEMRA_MOE_INPUT_TRACE_DIR").is_ok();
10088 if dev_ok
10089 && t <= MOE_DEV_MAX_T
10090 && m.dev_exps.is_some()
10091 && n_used <= 8
10092 && moe_dev_enabled()
10093 && !observe_routes
10094 {
10095 return Self::moe_ffn_dev(e, m, z, zq8, &logits, t, cfg, il, max_block);
10096 }
10097 if dev_ok && use_cache && n_used <= 8 && moe_dev_enabled() && !observe_routes {
10098 let row_ok = e.with_moe_cache(max_block, |c, eng| {
10099 if moe_prewarm_enabled() {
10100 c.prewarm_layer(il, m, eng)?;
10101 }
10102 Ok(c.layer_dev_row(il, n_expert, eng)?.is_some())
10103 })?;
10104 if row_ok {
10105 return Self::moe_ffn_dev(e, m, z, zq8, &logits, t, cfg, il, max_block);
10106 }
10107 }
10108
10109 let slab_local = m
10115 .dev_exps
10116 .as_ref()
10117 .filter(|d| !d.gu_il && moe_slab_enabled() && d.dev == e.ctx().ordinal());
10118 let slab_bases = slab_local.map(|d| {
10119 use cudarc::driver::DevicePtr;
10120 let s = e.stream();
10121 let (pg, _g0) = d.gate.device_ptr(&s);
10122 let (pu, _g1) = d.up.device_ptr(&s);
10123 let (pd, _g2) = d.down.device_ptr(&s);
10124 (pg, pu, pd)
10125 });
10126 let vrows_dev = vrows
10142 && t >= 2
10143 && crate::moe_vrows_dev_tables_on()
10144 && slab_bases.is_some()
10145 && moe_q8
10146 && uniform_experts
10147 && n_used <= 8
10148 && cfg.sigmoid_router().is_some()
10149 && matches!(lim_exp, Some(SwigluClamp::Pre(l)) if l > 1e-6)
10150 && !cpu_hybrid
10151 && sigmoid_router_enabled()
10152 && !observe_routes
10153 && !memra_reference::hidden_trace::enabled()
10154 && !crate::moesd::capture_active();
10155 let mut sel_dev: Option<(CudaSlice<i32>, CudaSlice<f32>)> = None;
10158 let (sel_all, w_all, routed_cpu_input) = if vrows_dev {
10159 let (sf, route_norm) = cfg
10160 .sigmoid_router()
10161 .expect("vrows_dev carries cfg.sigmoid_router().is_some()");
10162 sel_dev = Some(e.moe_router_sigmoid_topk(
10163 &logits,
10164 t,
10165 n_expert,
10166 n_used,
10167 m.active_count(),
10168 &m.exp_probs_b_dev,
10169 &m.active_experts_dev,
10170 sf,
10171 route_norm,
10172 )?);
10173 crate::MOE_VROWS_ROUTER_SYNCS_AVOIDED
10174 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
10175 (Vec::new(), Vec::new(), None)
10176 } else if let Some(sig) = cfg.sigmoid_router() {
10177 if cpu_hybrid {
10178 let (sel, w, input) = Self::moe_route_sigmoid_with_input(
10179 e,
10180 &logits,
10181 z,
10182 t,
10183 n_embd,
10184 n_expert,
10185 n_used,
10186 m.exp_probs_b.as_deref(),
10187 sig,
10188 m.active_experts.as_deref(),
10189 )?;
10190 (sel, w, Some(input))
10191 } else {
10192 let (sel, w) =
10193 Self::moe_route_sigmoid_cfg(e, &logits, t, n_expert, n_used, m, sig)?;
10194 (sel, w, None)
10195 }
10196 } else {
10197 let (sel, w) =
10198 Self::moe_route_cfg(e, &logits, t, n_expert, n_used, m.active_experts.as_deref())?;
10199 (sel, w, None)
10200 };
10201 crate::moesd::record_host_routes(il, n_expert, n_used, &sel_all)?;
10202 if memra_reference::hidden_trace::enabled() {
10203 memra_reference::hidden_trace::emit_last_row(
10204 "router",
10205 il as i64,
10206 t,
10207 n_expert,
10208 &e.dtoh(&logits)?,
10209 );
10210 let last = (t - 1) * n_used;
10211 let mut route = Vec::with_capacity(n_used * 2);
10212 for slot in 0..n_used {
10213 route.push(sel_all[last + slot] as f32);
10214 route.push(w_all[last + slot]);
10215 }
10216 memra_reference::hidden_trace::emit_last_row("route", il as i64, 1, n_used * 2, &route);
10217 }
10218
10219 Self::trace_moe_routes(il, t, &sel_all, &w_all)?;
10223 Self::trace_moe_input(e, il, t, n_embd, z)?;
10224
10225 let worker_disk_prefetch =
10237 cache_dispatch && crate::spill_pread::worker_enabled() && !cpu_hybrid;
10238 let promote_worker_h2d =
10239 t == 1 && worker_disk_prefetch && crate::spill_pread::copy_h2d_enabled();
10240 if promote_worker_h2d {
10241 let mut selected_blocks = Vec::with_capacity(n_used * 3);
10242 for &ex in sel_all.iter().take(n_used) {
10243 let ex = ex as u16;
10244 selected_blocks.extend([
10245 BlockId::new(il, PROJ_GATE, ex),
10246 BlockId::new(il, PROJ_UP, ex),
10247 BlockId::new(il, PROJ_DOWN, ex),
10248 ]);
10249 }
10250 for &ex in sel_all.iter().take(n_used) {
10251 Self::moe_prefetch_disk_expert(e, il, ex as usize, m, max_block, &selected_blocks)?;
10252 }
10253 e.with_moe_cache(max_block, |cache, eng| {
10254 cache.promote_worker_reads_at_safe_boundary(
10255 &selected_blocks,
10256 &selected_blocks,
10257 eng,
10258 )?;
10259 Ok(())
10260 })?;
10261 }
10262
10263 if t > 1 && std::env::var("MEMRA_MOE_STATS").is_ok() {
10266 let mut cnt = vec![0u32; n_expert];
10267 for &s in sel_all.iter() {
10268 cnt[s as usize] += 1;
10269 }
10270 let total = sel_all.len() as f64;
10271 let mut h = 0.0f64;
10272 let mut active = 0usize;
10273 for &c in &cnt {
10274 if c > 0 {
10275 active += 1;
10276 let p = c as f64 / total;
10277 h -= p * p.log2();
10278 }
10279 }
10280 let maxc = cnt.iter().copied().max().unwrap_or(0);
10281 println!(
10282 "moe-stats il={} t={} assignments={} active={}/{} entropy={:.3}b (max {:.3}b) mean_tok_per_active={:.2} max_tok_per_expert={}",
10283 il,
10284 t,
10285 sel_all.len(),
10286 active,
10287 n_expert,
10288 h,
10289 (n_expert as f64).log2(),
10290 total / active.max(1) as f64,
10291 maxc
10292 );
10293 }
10294
10295 let gdec_may_fire = uniform_experts
10308 && use_cache
10309 && n_used <= 8
10310 && gdec_enabled()
10311 && !cfg.swiglu_clamped_at(il as u32);
10312 let slab_fused_may_fire = slab_bases.is_some()
10340 && n_used <= 8
10341 && gdec_enabled()
10342 && !cfg.swiglu_clamped_at(il as u32)
10343 && cfg.m3.is_none()
10344 && no_exp_macros
10345 && moe_q8;
10346 let fused_epi_common = n_used <= 8
10381 && moe_q8
10382 && cfg.m3.is_none()
10383 && cfg.sigmoid_router().is_some()
10384 && matches!(lim_exp, Some(SwigluClamp::Pre(l)) if l > 1e-6)
10385 && moe_fused_epi_enabled();
10386 let fused_epi_may_fire = fused_epi_common
10387 && uniform_experts
10388 && use_cache
10389 && cache_dispatch
10390 && slab_local.is_none();
10391 let fused_epi_slab_may_fire = fused_epi_common && slab_bases.is_some();
10392 let vrows_fires = vrows
10408 && t >= 2
10409 && slab_bases.is_some()
10410 && moe_q8
10411 && uniform_experts
10412 && n_used <= 8
10413 && cfg.sigmoid_router().is_some()
10414 && matches!(lim_exp, Some(SwigluClamp::Pre(l)) if l > 1e-6)
10415 && !cpu_hybrid;
10416 let mut moe_out = if gdec_may_fire
10421 || slab_fused_may_fire
10422 || fused_epi_may_fire
10423 || fused_epi_slab_may_fire
10424 || vrows_fires
10425 {
10426 e.uninit(t * n_embd)?
10427 } else {
10428 e.zeros(t * n_embd)?
10429 };
10430 let cpu_input = if cpu_hybrid {
10433 Some(routed_cpu_input.ok_or("CPU expert routing did not return the MoE input")?)
10434 } else {
10435 None
10436 };
10437
10438 if vrows_dev && !vrows_fires {
10443 return Err(
10444 "MEMRA_MOE_VROWS_DEV_TABLES routed device-only but the verify-rows arm did not \
10445 fire: the door-D and vrows_fires predicates disagree"
10446 .into(),
10447 );
10448 }
10449 if vrows_fires {
10450 let Some(SwigluClamp::Pre(limit)) = lim_exp else {
10451 return Err(
10452 "verify-rows MoE arm fired without a live PRE clamp: the predicate and \
10453 the dispatch disagree"
10454 .into(),
10455 );
10456 };
10457 let bases = slab_bases.expect("vrows_fires carries slab_bases.is_some()");
10458 let sel = match sel_dev.as_ref() {
10459 Some((si, sw)) => VrowsSel::Dev(si, sw),
10460 None => VrowsSel::Host(&sel_all, &w_all),
10461 };
10462 Self::moe_vrows_pairs_q8(
10463 e,
10464 m,
10465 z,
10466 sel,
10467 il,
10468 bases,
10469 t,
10470 n_embd,
10471 n_ff_exp,
10472 n_used,
10473 limit,
10474 &mut moe_out,
10475 )?;
10476 if memra_reference::hidden_trace::enabled() {
10477 memra_reference::hidden_trace::emit_last_row(
10478 "routed",
10479 il as i64,
10480 t,
10481 n_embd,
10482 &e.dtoh(&moe_out)?,
10483 );
10484 }
10485 Self::moe_shexp_add(e, m, z, zq8, t, cfg, lim_shexp, &mut moe_out)?;
10486 return Ok(moe_out);
10487 }
10488
10489 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;
10497 let mut scratch_u: Option<CudaSlice<u8>> = None;
10498 let mut scratch_d: Option<CudaSlice<u8>> = None;
10499 let page_window = moe_page_prefetch_window();
10507
10508 for tok in 0..t {
10511 let sel = &sel_all[tok * n_used..(tok + 1) * n_used];
10512 let w = &w_all[tok * n_used..(tok + 1) * n_used];
10513 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd); let mut tok_q8: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
10515
10516 let no_macros = m.gate_exps.macros.is_none()
10530 && m.up_exps.macros.is_none()
10531 && m.down_exps.macros.is_none();
10532 if slab_fused_may_fire {
10542 let (pg, pu, pd) = slab_bases.unwrap();
10543 let mut gp = [0u64; 8];
10544 let mut up = [0u64; 8];
10545 let mut dp = [0u64; 8];
10546 for (j, &ex) in sel.iter().enumerate() {
10547 let ex = ex as usize;
10548 gp[j] = pg + (ex * m.gate_exps.expert_stride) as u64;
10549 up[j] = pu + (ex * m.up_exps.expert_stride) as u64;
10550 dp[j] = pd + (ex * m.down_exps.expert_stride) as u64;
10551 }
10552 let mut wv = [0f32; 8];
10553 wv[..n_used].copy_from_slice(w);
10554 if tok_q8.is_none() {
10555 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
10556 }
10557 let (zq, zd) = tok_q8.as_ref().unwrap();
10558 let act = e.moe_gate_up_silu8_q8(
10559 crate::WPtr8(gp),
10560 crate::WPtr8(up),
10561 zq,
10562 zd,
10563 n_embd,
10564 n_ff_exp,
10565 n_used,
10566 m.gate_exps.qtype,
10567 m.up_exps.qtype,
10568 m.gate_exps.row_bytes,
10569 m.up_exps.row_bytes,
10570 )?;
10571 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
10572 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
10573 e.moe_down8_fma_q8(
10574 crate::WPtr8(dp),
10575 crate::F32x8(wv),
10576 &aq2,
10577 &ad2,
10578 &mut dst,
10579 n_ff_exp,
10580 n_embd,
10581 n_used,
10582 m.down_exps.qtype,
10583 m.down_exps.row_bytes,
10584 )?;
10585 continue;
10586 }
10587 if fused_epi_slab_may_fire {
10592 let Some(SwigluClamp::Pre(limit)) = lim_exp else {
10593 return Err(
10594 "fused MoE epilogue (slab) fired without a live PRE clamp: the \
10595 predicate and the dispatch disagree"
10596 .into(),
10597 );
10598 };
10599 let (pg, pu, pd) = slab_bases.unwrap();
10600 let mut g = [0u64; 8];
10601 let mut u = [0u64; 8];
10602 let mut d = [0u64; 8];
10603 for (j, &ex) in sel.iter().enumerate() {
10604 let ex = ex as usize;
10605 g[j] = pg + (ex * m.gate_exps.expert_stride) as u64;
10606 u[j] = pu + (ex * m.up_exps.expert_stride) as u64;
10607 d[j] = pd + (ex * m.down_exps.expert_stride) as u64;
10608 }
10609 if tok_q8.is_none() {
10610 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
10611 }
10612 let (zq, zd) = tok_q8.as_ref().unwrap();
10613 Self::moe_fused_epi_launch(
10614 e,
10615 m,
10616 zq,
10617 zd,
10618 sel,
10619 w,
10620 g,
10621 u,
10622 d,
10623 &mut moe_out,
10624 tok,
10625 n_embd,
10626 n_ff_exp,
10627 n_used,
10628 limit,
10629 )?;
10630 continue;
10631 }
10632 if fused_epi_may_fire {
10637 let Some(SwigluClamp::Pre(limit)) = lim_exp else {
10638 return Err(
10639 "fused MoE epilogue fired without a live PRE clamp: the predicate and \
10640 the dispatch disagree"
10641 .into(),
10642 );
10643 };
10644 if tok_q8.is_none() {
10645 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
10646 }
10647 let (zq, zd) = tok_q8.as_ref().unwrap();
10648 if Self::moe_fused_epi_token_q8(
10649 e,
10650 m,
10651 il,
10652 max_block,
10653 zq,
10654 zd,
10655 sel,
10656 w,
10657 &mut moe_out,
10658 tok,
10659 n_embd,
10660 n_ff_exp,
10661 n_used,
10662 limit,
10663 )? {
10664 continue;
10665 }
10666 }
10667 if gdec_may_fire && moe_q8 && cfg.m3.is_none() && no_macros {
10668 if tok_q8.is_none() {
10669 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
10670 }
10671 let (zq, zd) = tok_q8.as_ref().unwrap();
10672 if Self::moe_gdec_token_q8(
10673 e,
10674 m,
10675 il,
10676 max_block,
10677 zq,
10678 zd,
10679 sel,
10680 w,
10681 &mut moe_out,
10682 tok,
10683 n_embd,
10684 n_ff_exp,
10685 n_used,
10686 )? {
10687 continue;
10688 }
10689 } else if gdec_may_fire
10690 && cfg.m3.is_none()
10691 && no_macros
10692 && Self::moe_gdec_token(
10693 e,
10694 m,
10695 il,
10696 max_block,
10697 &zt,
10698 sel,
10699 w,
10700 &mut moe_out,
10701 tok,
10702 n_embd,
10703 n_ff_exp,
10704 n_used,
10705 )?
10706 {
10707 continue;
10708 }
10709
10710 if gdec_may_fire || slab_fused_may_fire || fused_epi_may_fire || fused_epi_slab_may_fire
10716 {
10717 let mut row = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
10718 e.memset_zeros_view(&mut row)?;
10719 }
10720
10721 let mut cpu_mask = vec![false; sel.len()];
10727 let cpu_worker = if let Some(host_input) = cpu_input.as_ref() {
10728 let gpu_resident = if use_cache {
10729 e.with_moe_cache(max_block, |cache, _| {
10730 Ok(sel
10731 .iter()
10732 .map(|&expert| {
10733 let expert = expert as u16;
10734 [PROJ_GATE, PROJ_UP, PROJ_DOWN]
10735 .into_iter()
10736 .filter(|&projection| {
10737 cache
10738 .resident(BlockId::new(il, projection, expert))
10739 .is_some()
10740 })
10741 .count()
10742 })
10743 .collect::<Vec<_>>())
10744 })?
10745 } else {
10746 vec![0; sel.len()]
10747 };
10748 let mut cpu_selected = Vec::new();
10749 for (index, (&expert, &route_weight)) in sel.iter().zip(w).enumerate() {
10750 if gpu_resident[index] != 3 {
10751 cpu_mask[index] = true;
10752 crate::cpu_experts::record_incomplete_gpu_residency(gpu_resident[index]);
10753 let expert = expert as usize;
10754 cpu_selected.push((expert, route_weight));
10755 }
10756 }
10757 if crate::cpu_experts::predictor_enabled() {
10758 let row = &host_input[tok * n_embd..(tok + 1) * n_embd];
10762 crate::cpu_experts::predictor_submit(il, row);
10763 }
10764 if cpu_selected.is_empty() {
10765 None
10766 } else {
10767 let row = &host_input[tok * n_embd..(tok + 1) * n_embd];
10768 let job = crate::cpu_experts::prepare_job(m, il, &cpu_selected, row)
10769 .map_err(std::io::Error::other)?;
10770 Some(crate::cpu_experts::submit(job).map_err(std::io::Error::other)?)
10771 }
10772 } else {
10773 None
10774 };
10775
10776 let worker_window = worker_disk_prefetch
10777 .then(worker_prefetch_window)
10778 .unwrap_or(0);
10779 for (j, &ex) in sel.iter().enumerate() {
10780 if cpu_mask[j] {
10781 continue;
10782 }
10783 let ex = ex as usize;
10784 if let Some(d) = slab_local {
10791 let gl = m.gate_exps.expert_layout(ex);
10792 let ul = m.up_exps.expert_layout(ex);
10793 let dl = m.down_exps.expert_layout(ex);
10794 let (g0, u0, d0) = (
10795 ex * m.gate_exps.expert_stride,
10796 ex * m.up_exps.expert_stride,
10797 ex * m.down_exps.expert_stride,
10798 );
10799 let (gate, up) = if moe_q8 {
10800 if tok_q8.is_none() {
10801 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
10802 }
10803 let (zq, zd) = tok_q8.as_ref().unwrap();
10804 (
10805 e.qmatvec_expert_q8(
10806 &d.gate,
10807 g0..g0 + gl.len,
10808 zq,
10809 zd,
10810 1,
10811 m.gate_exps.in_f,
10812 m.gate_exps.out_f,
10813 gl.qtype,
10814 gl.row_bytes,
10815 )?,
10816 e.qmatvec_expert_q8(
10817 &d.up,
10818 u0..u0 + ul.len,
10819 zq,
10820 zd,
10821 1,
10822 m.up_exps.in_f,
10823 m.up_exps.out_f,
10824 ul.qtype,
10825 ul.row_bytes,
10826 )?,
10827 )
10828 } else {
10829 (
10830 m.qmatvec_view(
10831 e,
10832 &d.gate,
10833 g0..g0 + gl.len,
10834 &zt,
10835 1,
10836 m.gate_exps.in_f,
10837 m.gate_exps.out_f,
10838 gl.qtype,
10839 gl.row_bytes,
10840 )?,
10841 m.qmatvec_view(
10842 e,
10843 &d.up,
10844 u0..u0 + ul.len,
10845 &zt,
10846 1,
10847 m.up_exps.in_f,
10848 m.up_exps.out_f,
10849 ul.qtype,
10850 ul.row_bytes,
10851 )?,
10852 )
10853 };
10854 let mut act = e.uninit(n_ff_exp)?;
10855 Self::ffn_act_lim(
10856 e,
10857 cfg,
10858 &gate,
10859 &up,
10860 m.gate_exps.macro_scale(ex),
10861 m.up_exps.macro_scale(ex),
10862 lim_exp,
10863 &mut act,
10864 n_ff_exp,
10865 )?;
10866 let y = if moe_q8 {
10867 let (aq2, ad2) = e.quantize_q8_1(&act, 1, n_ff_exp)?;
10868 e.qmatvec_expert_q8(
10869 &d.down,
10870 d0..d0 + dl.len,
10871 &aq2,
10872 &ad2,
10873 1,
10874 m.down_exps.in_f,
10875 m.down_exps.out_f,
10876 dl.qtype,
10877 dl.row_bytes,
10878 )?
10879 } else {
10880 let actv = act.slice(0..n_ff_exp);
10881 m.qmatvec_view(
10882 e,
10883 &d.down,
10884 d0..d0 + dl.len,
10885 &actv,
10886 1,
10887 m.down_exps.in_f,
10888 m.down_exps.out_f,
10889 dl.qtype,
10890 dl.row_bytes,
10891 )?
10892 };
10893 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
10894 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
10895 continue;
10896 }
10897 for next in page_prefetch_positions(j, sel.len(), page_window) {
10898 Self::moe_prefetch_host_expert(sel[next] as usize, m);
10899 }
10900 let keep = [
10901 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_GATE, ex as u16),
10902 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_UP, ex as u16),
10903 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_DOWN, ex as u16),
10904 ];
10905 if worker_disk_prefetch && worker_window > 0 {
10906 for next in worker_prefetch_positions(j, sel.len(), worker_window) {
10907 Self::moe_prefetch_disk_expert(
10908 e,
10909 il,
10910 sel[next] as usize,
10911 m,
10912 max_block,
10913 &keep,
10914 )?;
10915 }
10916 } else if cache_dispatch
10917 && !cpu_hybrid
10918 && moe_prefetch_enabled()
10919 && j + 1 < sel.len()
10920 {
10921 let next = sel[j + 1] as usize;
10922 Self::moe_prefetch_expert(e, il, next, m, max_block, &keep)?;
10923 }
10924 let [gate_q8, up_q8, down_q8] = [moe_q8; 3];
10925 if cache_dispatch && (gate_q8 || up_q8 || down_q8) {
10926 if (gate_q8 || up_q8) && tok_q8.is_none() {
10929 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
10930 }
10931 let gate = if gate_q8 {
10932 let (zq, zd) = tok_q8.as_ref().unwrap();
10933 Self::moe_cached_gemm_q8(e, il, PROJ_GATE, ex, m, max_block, zq, zd)?
10934 } else {
10935 Self::moe_cached_gemm(e, il, PROJ_GATE, ex, m, max_block, &zt)?
10936 };
10937 let up = if up_q8 {
10938 let (zq, zd) = tok_q8.as_ref().unwrap();
10939 Self::moe_cached_gemm_q8(e, il, PROJ_UP, ex, m, max_block, zq, zd)?
10940 } else {
10941 Self::moe_cached_gemm(e, il, PROJ_UP, ex, m, max_block, &zt)?
10942 };
10943 let mut act = e.uninit(n_ff_exp)?;
10944 Self::ffn_act_lim(
10945 e,
10946 cfg,
10947 &gate,
10948 &up,
10949 m.gate_exps.macro_scale(ex),
10950 m.up_exps.macro_scale(ex),
10951 lim_exp,
10952 &mut act,
10953 n_ff_exp,
10954 )?;
10955 let y = if down_q8 {
10956 let (aq2, ad2) = e.quantize_q8_1(&act, 1, n_ff_exp)?;
10957 Self::moe_cached_gemm_q8(e, il, PROJ_DOWN, ex, m, max_block, &aq2, &ad2)?
10958 } else {
10959 let actv = act.slice(0..n_ff_exp);
10960 Self::moe_cached_gemm(e, il, PROJ_DOWN, ex, m, max_block, &actv)?
10961 };
10962 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
10963 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
10965 } else if cache_dispatch {
10966 let gate = Self::moe_cached_gemm(e, il, PROJ_GATE, ex, m, max_block, &zt)?;
10971 let up = Self::moe_cached_gemm(e, il, PROJ_UP, ex, m, max_block, &zt)?;
10972 let mut act = e.uninit(n_ff_exp)?; Self::ffn_act_lim(
10974 e,
10975 cfg,
10976 &gate,
10977 &up,
10978 m.gate_exps.macro_scale(ex),
10979 m.up_exps.macro_scale(ex),
10980 lim_exp,
10981 &mut act,
10982 n_ff_exp,
10983 )?;
10984 let actv = act.slice(0..n_ff_exp);
10985 let y = Self::moe_cached_gemm(e, il, PROJ_DOWN, ex, m, max_block, &actv)?;
10986 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
10987 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
10989 } else if cache_frozen {
10990 let gate = Self::moe_frozen_gemm(
10995 e,
10996 il,
10997 PROJ_GATE,
10998 ex,
10999 m,
11000 max_block,
11001 &zt,
11002 &mut scratch_g,
11003 g_len,
11004 )?;
11005 let up = Self::moe_frozen_gemm(
11006 e,
11007 il,
11008 PROJ_UP,
11009 ex,
11010 m,
11011 max_block,
11012 &zt,
11013 &mut scratch_u,
11014 u_len,
11015 )?;
11016 let mut act = e.uninit(n_ff_exp)?;
11017 Self::ffn_act_lim(
11018 e,
11019 cfg,
11020 &gate,
11021 &up,
11022 m.gate_exps.macro_scale(ex),
11023 m.up_exps.macro_scale(ex),
11024 lim_exp,
11025 &mut act,
11026 n_ff_exp,
11027 )?;
11028 let actv = act.slice(0..n_ff_exp);
11029 let y = Self::moe_frozen_gemm(
11030 e,
11031 il,
11032 PROJ_DOWN,
11033 ex,
11034 m,
11035 max_block,
11036 &actv,
11037 &mut scratch_d,
11038 d_len,
11039 )?;
11040 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
11041 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
11042 } else {
11043 if scratch_g.is_none() {
11047 scratch_g = Some(e.alloc_u8_uninit(g_len)?);
11048 scratch_u = Some(e.alloc_u8_uninit(u_len)?);
11049 scratch_d = Some(e.alloc_u8_uninit(d_len)?);
11050 }
11051 let (sg, su, sd) = (
11052 scratch_g.as_mut().unwrap(),
11053 scratch_u.as_mut().unwrap(),
11054 scratch_d.as_mut().unwrap(),
11055 );
11056 let gl = m.gate_exps.expert_layout(ex);
11057 let ul = m.up_exps.expert_layout(ex);
11058 let dl = m.down_exps.expert_layout(ex);
11059 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
11060 let gate = m.qmatvec_view(
11061 e,
11062 sg,
11063 0..gl.len,
11064 &zt,
11065 1,
11066 m.gate_exps.in_f,
11067 m.gate_exps.out_f,
11068 gl.qtype,
11069 gl.row_bytes,
11070 )?;
11071
11072 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
11073 let up = m.qmatvec_view(
11074 e,
11075 su,
11076 0..ul.len,
11077 &zt,
11078 1,
11079 m.up_exps.in_f,
11080 m.up_exps.out_f,
11081 ul.qtype,
11082 ul.row_bytes,
11083 )?;
11084
11085 let mut act = e.uninit(n_ff_exp)?; Self::ffn_act_lim(
11087 e,
11088 cfg,
11089 &gate,
11090 &up,
11091 m.gate_exps.macro_scale(ex),
11092 m.up_exps.macro_scale(ex),
11093 lim_exp,
11094 &mut act,
11095 n_ff_exp,
11096 )?;
11097
11098 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
11099 let actv = act.slice(0..n_ff_exp);
11100 let y = m.qmatvec_view(
11101 e,
11102 sd,
11103 0..dl.len,
11104 &actv,
11105 1,
11106 m.down_exps.in_f,
11107 m.down_exps.out_f,
11108 dl.qtype,
11109 dl.row_bytes,
11110 )?;
11111
11112 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
11113 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
11114 }
11115 }
11116 if let Some(worker) = cpu_worker {
11117 let cpu_output = worker.wait().map_err(std::io::Error::other)?;
11118 let cpu_output = e.htod(&cpu_output)?;
11119 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
11120 e.axpy_into(&cpu_output, 1.0, &mut dst, n_embd)?;
11121 }
11122 if cpu_hybrid && !cache_frozen && cpu_expert_profile_admit_enabled() {
11123 for (j, &ex) in sel.iter().enumerate() {
11124 if cpu_mask[j] {
11125 Self::moe_profile_admit_expert(e, il, ex as usize, m, max_block)?;
11126 }
11127 }
11128 }
11129 }
11130
11131 if memra_reference::hidden_trace::enabled() {
11132 memra_reference::hidden_trace::emit_last_row(
11133 "routed",
11134 il as i64,
11135 t,
11136 n_embd,
11137 &e.dtoh(&moe_out)?,
11138 );
11139 }
11140
11141 Self::moe_shexp_add(e, m, z, zq8, t, cfg, lim_shexp, &mut moe_out)?;
11142
11143 Ok(moe_out)
11144 }
11145
11146 #[allow(clippy::too_many_arguments)]
11166 fn moe_ffn_glm5_ep(
11167 e: &Engine,
11168 m: &MoeWeights,
11169 ep: &crate::glm5_tp::Glm5EpExps,
11170 z: &CudaSlice<f32>,
11171 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
11172 t: usize,
11173 cfg: &ModelConfig,
11174 il: u16,
11175 prefill: bool,
11176 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
11177 let moe = cfg
11178 .moe
11179 .as_ref()
11180 .ok_or("glm5 EP execution requires MoE model metadata")?;
11181 let n_embd = cfg.n_embd as usize;
11182 let n_expert = moe.expert_count as usize;
11183 let n_used = moe.expert_used_count as usize;
11184 let n_ff_exp = moe.expert_ff_length as usize;
11185 let sig = cfg
11186 .sigmoid_router()
11187 .ok_or("glm5 EP execution requires the sigmoid router")?;
11188 let lim_exp = cfg.clamp_exp_at(il as u32);
11189 let lim_shexp = cfg.clamp_shexp_at(il as u32);
11190 if ep.slabs.iter().map(|s| s.n_experts).sum::<usize>() != n_expert {
11191 return Err(format!(
11192 "glm5 EP slabs cover {:?} experts, model declares {n_expert}",
11193 ep.slabs.iter().map(|s| s.n_experts).collect::<Vec<_>>()
11194 )
11195 .into());
11196 }
11197 let rt = &ep.rt;
11198 let ranks = ep.ranks();
11199
11200 let logits = Self::moe_router_logits(e, m, z, t, cfg)?;
11202 let (sel_all, w_all) =
11203 Self::moe_route_sigmoid_cfg(e, &logits, t, n_expert, n_used, m, sig)?;
11204 crate::moesd::record_host_routes(il, n_expert, n_used, &sel_all)?;
11205 Self::trace_moe_routes(il, t, &sel_all, &w_all)?;
11208
11209 if prefill && t > MOE_DEV_MAX_T {
11215 static EPGP_ANNOUNCED: std::sync::atomic::AtomicU8 =
11216 std::sync::atomic::AtomicU8::new(0);
11217 let enabled = crate::ep_grouped_prime_on() && moe_grouped_prefill_enabled();
11218 let bit = 1u8 << u8::from(enabled);
11219 if EPGP_ANNOUNCED.fetch_or(bit, std::sync::atomic::Ordering::Relaxed) & bit == 0 {
11220 eprintln!(
11221 "[glm5-ep-grouped-prime] flag={} t={t} il={il} (announce printed in both \
11222 arms; engagement is the dispatch counter + per-layer execute line)",
11223 if enabled { "on" } else { "off" },
11224 );
11225 }
11226 if enabled
11227 && let Some(mut out) =
11228 Self::moe_ffn_glm5_ep_grouped_prime(e, m, ep, z, &sel_all, &w_all, t, cfg, il)?
11229 {
11230 Self::moe_shexp_add(e, m, z, zq8, t, cfg, lim_shexp, &mut out)?;
11231 return Ok(out);
11232 }
11233 }
11234
11235 if t > 1 && !prefill {
11243 static EP_VERIFY_MARKED: std::sync::atomic::AtomicBool =
11244 std::sync::atomic::AtomicBool::new(false);
11245 if !EP_VERIFY_MARKED.swap(true, std::sync::atomic::Ordering::Relaxed) {
11246 eprintln!(
11247 "[glm5-tp-ep] verify rows ride the SEQUENTIAL EP walk (t={t}): the \
11248 batched vrows MoE pair is preempted by EP; the EP-aware vrows arm is \
11249 the named lever performance_claim=false"
11250 );
11251 }
11252 }
11253 if crate::ep_diet_on() {
11256 let mut out = Self::moe_ffn_glm5_ep_diet(e, m, ep, z, &sel_all, &w_all, t, cfg, il)?;
11257 Self::moe_shexp_add(e, m, z, zq8, t, cfg, lim_shexp, &mut out)?;
11258 return Ok(out);
11259 }
11260
11261 let mut moe_out = e.zeros(t * n_embd)?;
11262 use crate::tp_transport::TpTransport as TpXport;
11263 let hop = ep.rt.hop(e);
11264 let z_host = match hop.transport {
11273 TpXport::HostCanonical => Some(crate::tp_transport::host_stage_block(
11274 &hop,
11275 0,
11276 z,
11277 t * n_embd,
11278 )?),
11279 TpXport::PeerPull => None,
11280 };
11281 let z_peer_bulks = match hop.transport {
11282 TpXport::PeerPull => Some(crate::tp_transport::fanout_f32(&hop, z, t * n_embd)?),
11283 TpXport::HostCanonical => None,
11284 };
11285 for tok in 0..t {
11286 let sel = &sel_all[tok * n_used..(tok + 1) * n_used];
11287 let w = &w_all[tok * n_used..(tok + 1) * n_used];
11288 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
11289 let z_peer_row_holders: Option<Vec<CudaSlice<f32>>> = match &z_host {
11293 Some(h) => {
11294 let mut rows = Vec::with_capacity(ranks - 1);
11295 for r in 1..ranks {
11296 rows.push(crate::tp_transport::host_row_to(
11297 &hop,
11298 r,
11299 &h[tok * n_embd..(tok + 1) * n_embd],
11300 )?);
11301 }
11302 Some(rows)
11303 }
11304 None => None,
11305 };
11306 for (j, &ex) in sel.iter().enumerate() {
11309 let ex = ex as usize;
11310 let owner = ep.owner(ex);
11311 if owner != 0 {
11312 crate::glm5_tp::GLM5_EP_PEER_SLOT_DISPATCHES
11314 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11315 if matches!(
11318 crate::glm5_tp::gate_red(),
11319 Ok(Some(crate::glm5_tp::GateRed::SkipPeerCombine))
11320 ) {
11321 continue;
11322 }
11323 }
11324 let zin_holder;
11325 let (dev, slab, zin) = if owner == 0 {
11326 (e, &ep.slabs[0], &zt)
11327 } else {
11328 zin_holder = match (&z_peer_row_holders, &z_peer_bulks) {
11329 (Some(rows), _) => rows[owner - 1].slice(0..n_embd),
11330 (None, Some(bulks)) => {
11331 bulks[owner - 1].slice(tok * n_embd..(tok + 1) * n_embd)
11332 }
11333 (None, None) => {
11334 return Err(
11335 "glm5 EP: neither transport arm staged the peer activation".into(),
11336 );
11337 }
11338 };
11339 (
11340 crate::glm5_tp::rank_engine(e, rt, owner),
11341 &ep.slabs[owner],
11342 &zin_holder,
11343 )
11344 };
11345 let local = ep.local_of[ex] as usize;
11349 let gl = m.gate_exps.expert_stride;
11350 let ul = m.up_exps.expert_stride;
11351 let dl = m.down_exps.expert_stride;
11352 let gate = dev.qmatvec_view(
11353 &slab.gate,
11354 local * gl..(local + 1) * gl,
11355 zin,
11356 1,
11357 m.gate_exps.in_f,
11358 m.gate_exps.out_f,
11359 m.gate_exps.qtype,
11360 m.gate_exps.row_bytes,
11361 )?;
11362 let up = dev.qmatvec_view(
11363 &slab.up,
11364 local * ul..(local + 1) * ul,
11365 zin,
11366 1,
11367 m.up_exps.in_f,
11368 m.up_exps.out_f,
11369 m.up_exps.qtype,
11370 m.up_exps.row_bytes,
11371 )?;
11372 let mut act = dev.uninit(n_ff_exp)?; Self::ffn_act_lim(
11374 dev,
11375 cfg,
11376 &gate,
11377 &up,
11378 m.gate_exps.macro_scale(ex),
11379 m.up_exps.macro_scale(ex),
11380 lim_exp,
11381 &mut act,
11382 n_ff_exp,
11383 )?;
11384 let actv = act.slice(0..n_ff_exp);
11385 let y = dev.qmatvec_view(
11386 &slab.down,
11387 local * dl..(local + 1) * dl,
11388 &actv,
11389 1,
11390 m.down_exps.in_f,
11391 m.down_exps.out_f,
11392 m.down_exps.qtype,
11393 m.down_exps.row_bytes,
11394 )?;
11395 let y_root = if owner == 0 {
11400 y
11401 } else {
11402 crate::tp_transport::return_row_to_root(&hop, owner, &y, n_embd)?
11403 };
11404 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
11405 e.axpy_into(
11406 &y_root,
11407 w[j] * m.down_exps.macro_scale(ex),
11408 &mut dst,
11409 n_embd,
11410 )?;
11411 }
11412 }
11413
11414 Self::moe_shexp_add(e, m, z, zq8, t, cfg, lim_shexp, &mut moe_out)?;
11415 Ok(moe_out)
11416 }
11417
11418 #[allow(clippy::too_many_arguments)]
11446 fn moe_ffn_glm5_ep_diet(
11447 e: &Engine,
11448 m: &MoeWeights,
11449 ep: &crate::glm5_tp::Glm5EpExps,
11450 z: &CudaSlice<f32>,
11451 sel_all: &[u32],
11452 w_all: &[f32],
11453 t: usize,
11454 cfg: &ModelConfig,
11455 il: u16,
11456 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
11457 use std::sync::atomic::Ordering;
11458 let moe = cfg
11459 .moe
11460 .as_ref()
11461 .ok_or("glm5 EP execution requires MoE model metadata")?;
11462 let n_embd = cfg.n_embd as usize;
11463 let n_used = moe.expert_used_count as usize;
11464 let n_ff_exp = moe.expert_ff_length as usize;
11465 let lim_exp = cfg.clamp_exp_at(il as u32);
11466 let rt = &ep.rt;
11467 let n_pairs = t * n_used;
11468 if sel_all.len() < n_pairs || w_all.len() < n_pairs || z.len() < t * n_embd {
11469 return Err("glm5 EP diet geometry".into());
11470 }
11471 let red_skip_peer = matches!(
11472 crate::glm5_tp::gate_red(),
11473 Ok(Some(crate::glm5_tp::GateRed::SkipPeerCombine))
11474 );
11475
11476 let ranks = ep.ranks();
11482 let mut per_rank = vec![0usize; ranks];
11483 for &s in sel_all.iter().take(n_pairs) {
11484 let ex = s as usize;
11485 if ex >= ep.owner_of.len() {
11486 return Err(format!("glm5 EP diet: selection {ex} outside the bank").into());
11487 }
11488 per_rank[ep.owner(ex)] += 1;
11489 }
11490 let mut base = vec![0usize; ranks];
11491 for r in 1..ranks {
11492 base[r] = base[r - 1] + per_rank[r - 1];
11493 }
11494 let mut ids = vec![0i32; n_pairs];
11495 {
11496 let mut k = vec![0usize; ranks];
11497 for (p, id) in ids.iter_mut().enumerate() {
11498 let r = ep.owner(sel_all[p] as usize);
11499 *id = (base[r] + k[r]) as i32;
11500 k[r] += 1;
11501 }
11502 }
11503
11504 crate::glm5_tp::GLM5_EP_DIET_DISPATCHES.fetch_add(1, Ordering::Relaxed);
11505 for r in 1..ranks {
11506 crate::glm5_tp::GLM5_EP_DIET_FANOUT_UPLOADS_AVOIDED.fetch_add(
11507 if per_rank[r] > 0 {
11508 (t - 1) as u64
11509 } else {
11510 t as u64
11511 },
11512 Ordering::Relaxed,
11513 );
11514 }
11515 static EP_DIET_MARKED: std::sync::atomic::AtomicBool =
11516 std::sync::atomic::AtomicBool::new(false);
11517 if !EP_DIET_MARKED.swap(true, Ordering::Relaxed) {
11518 eprintln!(
11519 "[glm5-ep-diet] engaged: bulk fan-out + compact peer staging + single \
11520 slot-ordered scatter combine; per-slot host round-trips removed \
11521 transport={} performance_claim=false",
11522 ep.rt.transport.name(),
11523 );
11524 }
11525
11526 let expert_row = |dev: &Engine,
11529 slab: &crate::glm5_tp::EpRankSlab,
11530 zin: &cudarc::driver::CudaView<f32>,
11531 ex: usize|
11532 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
11533 let local = ep.local_of[ex] as usize;
11534 let gl = m.gate_exps.expert_stride;
11535 let ul = m.up_exps.expert_stride;
11536 let dl = m.down_exps.expert_stride;
11537 let gate = dev.qmatvec_view(
11538 &slab.gate,
11539 local * gl..(local + 1) * gl,
11540 zin,
11541 1,
11542 m.gate_exps.in_f,
11543 m.gate_exps.out_f,
11544 m.gate_exps.qtype,
11545 m.gate_exps.row_bytes,
11546 )?;
11547 let up = dev.qmatvec_view(
11548 &slab.up,
11549 local * ul..(local + 1) * ul,
11550 zin,
11551 1,
11552 m.up_exps.in_f,
11553 m.up_exps.out_f,
11554 m.up_exps.qtype,
11555 m.up_exps.row_bytes,
11556 )?;
11557 let mut act = dev.uninit(n_ff_exp)?; Self::ffn_act_lim(
11559 dev,
11560 cfg,
11561 &gate,
11562 &up,
11563 m.gate_exps.macro_scale(ex),
11564 m.up_exps.macro_scale(ex),
11565 lim_exp,
11566 &mut act,
11567 n_ff_exp,
11568 )?;
11569 let actv = act.slice(0..n_ff_exp);
11570 dev.qmatvec_view(
11571 &slab.down,
11572 local * dl..(local + 1) * dl,
11573 &actv,
11574 1,
11575 m.down_exps.in_f,
11576 m.down_exps.out_f,
11577 m.down_exps.qtype,
11578 m.down_exps.row_bytes,
11579 )
11580 };
11581
11582 let hop = ep.rt.hop(e);
11586 let mut y_peer_blks: Vec<Option<CudaSlice<f32>>> = (0..ranks).map(|_| None).collect();
11587 for r in 1..ranks {
11588 if per_rank[r] == 0 {
11589 continue;
11590 }
11591 let dev = crate::glm5_tp::rank_engine(e, rt, r);
11592 let z_r = crate::tp_transport::fanout_f32_to(&hop, r, z, t * n_embd)?;
11595 let mut blk = if red_skip_peer {
11598 dev.zeros(per_rank[r] * n_embd)?
11599 } else {
11600 dev.uninit(per_rank[r] * n_embd)?
11601 };
11602 let mut k = 0usize;
11603 for tok in 0..t {
11604 let sel = &sel_all[tok * n_used..(tok + 1) * n_used];
11605 for &ex in sel.iter() {
11606 let ex = ex as usize;
11607 if ep.owner(ex) != r {
11608 continue;
11609 }
11610 crate::glm5_tp::GLM5_EP_PEER_SLOT_DISPATCHES.fetch_add(1, Ordering::Relaxed);
11612 crate::glm5_tp::GLM5_EP_DIET_PEER_ROUNDTRIPS_AVOIDED
11613 .fetch_add(1, Ordering::Relaxed);
11614 if red_skip_peer {
11615 k += 1;
11616 continue;
11617 }
11618 let zt_r = z_r.slice(tok * n_embd..(tok + 1) * n_embd);
11619 let y = expert_row(dev, &ep.slabs[r], &zt_r, ex)?;
11620 dev.copy_into(&mut blk, k * n_embd, &y, n_embd)?;
11621 k += 1;
11622 }
11623 }
11624 y_peer_blks[r] = Some(blk);
11625 }
11626
11627 let mut y_all = e.uninit(n_pairs * n_embd)?;
11630 {
11631 let mut k = 0usize;
11632 for tok in 0..t {
11633 let sel = &sel_all[tok * n_used..(tok + 1) * n_used];
11634 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
11635 for &ex in sel.iter() {
11636 let ex = ex as usize;
11637 if ep.owner(ex) != 0 {
11638 continue;
11639 }
11640 let y = expert_row(e, &ep.slabs[0], &zt, ex)?;
11641 e.copy_into(&mut y_all, k * n_embd, &y, n_embd)?;
11642 k += 1;
11643 }
11644 }
11645 }
11646
11647 for r in 1..ranks {
11652 if let Some(blk) = &y_peer_blks[r] {
11653 crate::tp_transport::return_block_to_root(
11654 &hop,
11655 r,
11656 blk,
11657 &mut y_all,
11658 base[r] * n_embd,
11659 per_rank[r] * n_embd,
11660 )?;
11661 crate::glm5_tp::GLM5_EP_DIET_BULK_RETURNS.fetch_add(1, Ordering::Relaxed);
11662 }
11663 }
11664
11665 let mut wd = vec![0f32; n_pairs];
11669 for ((&id, &w), &s) in ids.iter().zip(w_all.iter()).zip(sel_all.iter()) {
11670 wd[id as usize] = w * m.down_exps.macro_scale(s as usize);
11671 }
11672 let pw = e.htod(&wd)?;
11673 let toff: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
11674 let toff_d = e.htod_i32(&toff)?;
11675 let ids_d = e.htod_i32(&ids)?;
11676 let mut moe_out = e.uninit(t * n_embd)?; e.moe_pairs_scatter(&y_all, &pw, &toff_d, &ids_d, &mut moe_out, t, n_embd)?;
11678 Ok(moe_out)
11679 }
11680
11681 #[allow(clippy::too_many_arguments)]
11700 fn moe_ffn_glm5_ep_grouped_prime(
11701 e: &Engine,
11702 m: &MoeWeights,
11703 ep: &crate::glm5_tp::Glm5EpExps,
11704 z: &CudaSlice<f32>,
11705 sel_all: &[u32],
11706 w_all: &[f32],
11707 t: usize,
11708 cfg: &ModelConfig,
11709 il: u16,
11710 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
11711 use std::sync::atomic::Ordering;
11712 let moe = cfg
11713 .moe
11714 .as_ref()
11715 .ok_or("glm5 EP grouped prime requires MoE model metadata")?;
11716 let n_embd = cfg.n_embd as usize;
11717 let n_expert = moe.expert_count as usize;
11718 let n_used = moe.expert_used_count as usize;
11719 let n_ff_exp = moe.expert_ff_length as usize;
11720 if crate::moe_f16g_mode() == 0 || std::env::var("MEMRA_MOE_GATE").is_ok() {
11724 return Ok(None);
11725 }
11726 if !(f16g_proj_ok(m.gate_exps.qtype, n_embd)
11727 && f16g_proj_ok(m.up_exps.qtype, n_embd)
11728 && f16g_proj_ok(m.down_exps.qtype, n_ff_exp))
11729 {
11730 return Ok(None);
11731 }
11732 if n_expert > 512 || n_used == 0 || n_used > 8 {
11733 return Ok(None);
11734 }
11735 let lim_exp = cfg.clamp_exp_at(il as u32);
11736 if matches!(lim_exp, Some(SwigluClamp::Post(_))) {
11737 return Err(
11738 "EP grouped prime is qualified for the PRE-clamped SwiGLU form only; \
11739 a POST-clamp layer must ride the sequential arm"
11740 .into(),
11741 );
11742 }
11743 let n_pairs = t * n_used;
11744 if sel_all.len() < n_pairs || w_all.len() < n_pairs || z.len() < t * n_embd {
11745 return Err("EP grouped prime geometry".into());
11746 }
11747 let rt = &ep.rt;
11748 let red_skip_peer = matches!(
11749 crate::glm5_tp::gate_red(),
11750 Ok(Some(crate::glm5_tp::GateRed::SkipPeerCombine))
11751 );
11752
11753 let rank_pass = |dev: &Engine,
11758 rank: u8,
11759 ptr_row: &CudaSlice<u64>,
11760 z_dev: &CudaSlice<f32>|
11761 -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
11762 let mut buckets_l: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
11765 let mut local_tok = Vec::new(); let mut local_ex = Vec::new(); let mut local_wd = Vec::new(); let mut local_count_per_tok = vec![0i32; t];
11769 for p in 0..n_pairs {
11770 let ex = sel_all[p] as usize;
11771 if ex >= n_expert {
11772 return Err(format!("EP grouped prime selection {ex} >= {n_expert}").into());
11773 }
11774 if ep.owner(ex) != rank as usize {
11775 continue;
11776 }
11777 let l = local_tok.len() as i32;
11778 buckets_l[ex].push(l);
11779 let tok = p / n_used;
11780 local_tok.push(tok as i32);
11781 local_ex.push(ex);
11782 local_wd.push(w_all[p] * m.down_exps.macro_scale(ex));
11783 local_count_per_tok[tok] += 1;
11784 }
11785 let n_owned = local_tok.len();
11786 if n_owned == 0 {
11787 return Ok(None);
11788 }
11789 let mut ex_ids: Vec<i32> = Vec::new();
11790 let mut ex_off: Vec<i32> = vec![0];
11791 let mut ex_pairs: Vec<i32> = Vec::with_capacity(n_owned); let mut csr_tok: Vec<i32> = Vec::with_capacity(n_owned);
11793 for (e_id, b) in buckets_l.iter().enumerate() {
11794 if !b.is_empty() {
11795 ex_ids.push(e_id as i32);
11796 for &l in b {
11797 ex_pairs.push(l);
11798 csr_tok.push(local_tok[l as usize]);
11799 }
11800 ex_off.push(ex_pairs.len() as i32);
11801 }
11802 }
11803 let n_active = ex_ids.len();
11804 if n_active == 0 || n_active > 512 {
11805 return Err(format!("EP grouped prime n_active {n_active} outside 1..=512").into());
11806 }
11807
11808 let exi = dev.htod_i32(&ex_ids)?;
11809 let exo = dev.htod_i32(&ex_off)?;
11810 let exp_d = dev.htod_i32(&ex_pairs)?;
11811 let csr_tok_d = dev.htod_i32(&csr_tok)?;
11812
11813 let (z16, zs) = dev.moe_f16g_act(z_dev, Some(&csr_tok_d), n_embd, n_owned)?;
11815 let mut g = dev.moe_f16_grouped(
11816 ptr_row,
11817 0,
11818 n_expert,
11819 &exi,
11820 &ex_off,
11821 &exo,
11822 &z16,
11823 &zs,
11824 n_embd,
11825 n_ff_exp,
11826 n_active,
11827 n_owned,
11828 m.gate_exps.qtype,
11829 m.gate_exps.row_bytes,
11830 )?;
11831 if m.gate_exps.macros.is_some() {
11832 let mg: Vec<f32> = ex_pairs
11833 .iter()
11834 .map(|&l| m.gate_exps.macro_scale(local_ex[l as usize]))
11835 .collect();
11836 let mg_d = dev.htod(&mg)?;
11837 dev.scale_rows(&mut g, &mg_d, n_ff_exp, n_owned)?;
11838 }
11839 let mut u = dev.moe_f16_grouped(
11840 ptr_row,
11841 1,
11842 n_expert,
11843 &exi,
11844 &ex_off,
11845 &exo,
11846 &z16,
11847 &zs,
11848 n_embd,
11849 n_ff_exp,
11850 n_active,
11851 n_owned,
11852 m.up_exps.qtype,
11853 m.up_exps.row_bytes,
11854 )?;
11855 if m.up_exps.macros.is_some() {
11856 let mu: Vec<f32> = ex_pairs
11857 .iter()
11858 .map(|&l| m.up_exps.macro_scale(local_ex[l as usize]))
11859 .collect();
11860 let mu_d = dev.htod(&mu)?;
11861 dev.scale_rows(&mut u, &mu_d, n_ff_exp, n_owned)?;
11862 }
11863
11864 let act = match lim_exp {
11866 Some(SwigluClamp::Pre(limit)) => {
11867 let mut a = dev.uninit(n_owned * n_ff_exp)?;
11868 dev.swiglu_preclamped_mul_scaled(
11869 &g,
11870 &u,
11871 1.0,
11872 1.0,
11873 limit,
11874 &mut a,
11875 n_owned * n_ff_exp,
11876 )?;
11877 a
11878 }
11879 None => dev.moe_pairs_silu_mul(&g, &u, n_owned * n_ff_exp)?,
11880 Some(SwigluClamp::Post(_)) => unreachable!("refused before any launch"),
11881 };
11882
11883 let (a16, a_s) = dev.moe_f16g_act(&act, None, n_ff_exp, n_owned)?;
11885 let d_csr = dev.moe_f16_grouped(
11886 ptr_row,
11887 2,
11888 n_expert,
11889 &exi,
11890 &ex_off,
11891 &exo,
11892 &a16,
11893 &a_s,
11894 n_ff_exp,
11895 n_embd,
11896 n_active,
11897 n_owned,
11898 m.down_exps.qtype,
11899 m.down_exps.row_bytes,
11900 )?;
11901 let y_local = dev.rows_permute(&d_csr, &exp_d, n_owned, n_embd)?;
11902 let mut toff: Vec<i32> = Vec::with_capacity(t + 1);
11903 let mut acc = 0i32;
11904 toff.push(0);
11905 for &c in &local_count_per_tok {
11906 acc += c;
11907 toff.push(acc);
11908 }
11909 let tids: Vec<i32> = (0..n_owned as i32).collect();
11910 let pw = dev.htod(&local_wd)?;
11911 let toff_d = dev.htod_i32(&toff)?;
11912 let tids_d = dev.htod_i32(&tids)?;
11913 let mut partial = dev.uninit(t * n_embd)?; dev.moe_pairs_scatter(&y_local, &pw, &toff_d, &tids_d, &mut partial, t, n_embd)?;
11915 Ok(Some(partial))
11916 };
11917
11918 let ranks = ep.ranks();
11923 let n_peer_pairs = sel_all
11924 .iter()
11925 .take(n_pairs)
11926 .filter(|&&ex| ep.owner(ex as usize) != 0)
11927 .count() as u64;
11928 crate::glm5_tp::GLM5_EP_PEER_SLOT_DISPATCHES.fetch_add(n_peer_pairs, Ordering::Relaxed);
11929 let hop = ep.rt.hop(e);
11930 let mut peer_partials: Vec<Option<CudaSlice<f32>>> = (0..ranks).map(|_| None).collect();
11931 for r in 1..ranks {
11932 let rank_owns_pairs = sel_all
11933 .iter()
11934 .take(n_pairs)
11935 .any(|&ex| ep.owner(ex as usize) == r);
11936 if !rank_owns_pairs {
11937 continue;
11938 }
11939 let dev = crate::glm5_tp::rank_engine(e, rt, r);
11940 let z_r = crate::tp_transport::fanout_f32_to(&hop, r, z, t * n_embd)?;
11942 dev.bind_runtime_device(dev.ctx().ordinal() as i32)?;
11943 let res = rank_pass(dev, r as u8, &ep.ptr_rows[r], &z_r);
11944 e.bind_runtime_device(e.ctx().ordinal() as i32)?;
11945 peer_partials[r] = res?;
11946 }
11947 let root_partial = rank_pass(e, 0, &ep.ptr_rows[0], z)?;
11948
11949 let mut out = match root_partial {
11955 Some(p) => p,
11956 None => e.zeros(t * n_embd)?,
11957 };
11958 for r in 1..ranks {
11959 if let Some(pp) = &peer_partials[r]
11960 && !red_skip_peer
11961 {
11962 let pp_root = crate::tp_transport::return_row_to_root(&hop, r, pp, t * n_embd)?;
11963 let mut dst = out.slice_mut(0..t * n_embd);
11964 e.axpy_into(&pp_root, 1.0, &mut dst, t * n_embd)?;
11965 crate::glm5_tp::GLM5_EP_DIET_BULK_RETURNS.fetch_add(1, Ordering::Relaxed);
11966 }
11967 }
11968 crate::glm5_tp::GLM5_EP_GROUPED_PRIME_DISPATCHES.fetch_add(1, Ordering::Relaxed);
11969 static EPGP_LOGGED: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
11970 let layer_bit = 1u64 << (il as u64 % 64);
11971 if EPGP_LOGGED.fetch_or(layer_bit, Ordering::Relaxed) & layer_bit == 0 {
11972 eprintln!(
11973 "[glm5-ep-grouped-prime] execute layer={il} tokens={t} \
11974 provenance=ep-rank-slabs router=sigmoid-host-oracle epilogue=pre-clamped \
11975 combine=rank-partial-add transport={} performance_claim=false \
11976 (logged once per layer)",
11977 hop.transport.name(),
11978 );
11979 }
11980 Ok(Some(out))
11981 }
11982
11983 #[allow(clippy::too_many_arguments)]
11991 fn moe_shexp_add(
11992 e: &Engine,
11993 m: &MoeWeights,
11994 z: &CudaSlice<f32>,
11995 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
11996 t: usize,
11997 cfg: &ModelConfig,
11998 lim_shexp: Option<memra_gguf::config::SwigluClamp>,
11999 moe_out: &mut CudaSlice<f32>,
12000 ) -> Result<(), Box<dyn std::error::Error>> {
12001 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
12002 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
12003 {
12004 let n_embd = cfg.n_embd as usize;
12005 let n_ff_sh = gate_shexp.out_features(); let verify_t = t > 1 && t < PRIME_MIN_T;
12014 let (sg_gate, sg_up) = if t == 1 {
12015 shexp_gate_up_t1(e, gate_shexp, up_shexp, z, zq8)?
12016 } else if verify_t {
12017 (
12018 e.matmul_decode_exact(gate_shexp, z, t)?,
12019 e.matmul_decode_exact(up_shexp, z, t)?,
12020 )
12021 } else {
12022 (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?) };
12024 let mut sa = e.uninit(t * n_ff_sh)?; Self::ffn_act_lim(
12026 e,
12027 cfg,
12028 &sg_gate,
12029 &sg_up,
12030 1.0,
12031 1.0,
12032 lim_shexp,
12033 &mut sa,
12034 t * n_ff_sh,
12035 )?;
12036 let sh = if verify_t {
12037 e.matmul_decode_exact(down_shexp, &sa, t)?
12038 } else {
12039 e.matmul(down_shexp, &sa, t)?
12040 }; if m.gate_inp_shexp.is_none() && crate::htod_diet_on() {
12059 crate::HTOD_DIET_AVOIDED.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
12060 e.add_scaled_rows_ones(&sh, moe_out, n_embd, t)?;
12061 return Ok(());
12062 }
12063 let g = match &m.gate_inp_shexp {
12064 Some(gate_inp_shexp) => {
12065 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
12066 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
12067 } else {
12068 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
12069 let mut g = e.uninit(t)?; e.sigmoid(&gs, &mut g, t)?;
12071 g
12072 }
12073 }
12074 None => e.htod(&vec![1.0f32; t])?,
12075 };
12076 e.add_scaled_rows(&sh, &g, moe_out, n_embd, t)?;
12078 }
12079 Ok(())
12080 }
12081
12082 pub fn stage1_h2d_per_token(&self) -> u64 {
12085 use crate::hybrid::Ffn;
12086 let n_used = self
12087 .cfg
12088 .moe
12089 .as_ref()
12090 .map(|m| m.expert_used_count as u64)
12091 .unwrap_or(0);
12092 let mut bytes = 0u64;
12093 for l in self.layers.iter() {
12094 if let Ffn::Moe(m) = &l.ffn {
12095 bytes += n_used
12096 * (m.gate_exps.max_expert_bytes()
12097 + m.up_exps.max_expert_bytes()
12098 + m.down_exps.max_expert_bytes()) as u64;
12099 }
12100 }
12101 bytes
12102 }
12103
12104 pub(crate) fn max_moe_block(&self) -> usize {
12108 use crate::hybrid::Ffn;
12109 let mut mx = 0usize;
12110 let mut scan = |ffn: &Ffn| {
12111 if let Ffn::Moe(m) = ffn {
12112 mx = mx
12113 .max(m.gate_exps.max_expert_bytes())
12114 .max(m.up_exps.max_expert_bytes())
12115 .max(m.down_exps.max_expert_bytes());
12116 }
12117 };
12118 for l in self.layers.iter() {
12119 scan(&l.ffn);
12120 }
12121 if let Some(mtp) = self.mtp.as_ref() {
12122 scan(&mtp.ffn);
12123 }
12124 mx
12125 }
12126
12127 pub(crate) fn moe_cache_block_sizes(&self) -> Vec<usize> {
12130 use crate::hybrid::Ffn;
12131 let mut sizes = Vec::new();
12132 let mut scan = |ffn: &Ffn| {
12133 let Ffn::Moe(m) = ffn else { return };
12134 for ex in 0..m.gate_exps.n_expert {
12135 if m.active_experts.as_ref().is_some_and(|active| !active[ex]) {
12136 continue;
12137 }
12138 for exps in [&m.gate_exps, &m.up_exps, &m.down_exps] {
12139 let len = exps.expert_layout(ex).len;
12140 if len > 0 {
12141 sizes.push(len);
12142 }
12143 }
12144 }
12145 };
12146 for layer in &self.layers {
12147 scan(&layer.ffn);
12148 }
12149 if let Some(mtp) = &self.mtp {
12150 scan(&mtp.ffn);
12151 }
12152 sizes
12153 }
12154
12155 pub fn save_cpu_expert_residency_profile(
12161 &self,
12162 e: &Engine,
12163 path: &std::path::Path,
12164 ) -> Result<(), Box<dyn std::error::Error>> {
12165 let Some(ids) = e.export_moe_residency() else {
12166 return Err("no MoE residency cache to persist".into());
12167 };
12168 let mut body = format!(
12169 "memra-freeze-profile v1 max_block={} blocks={}\n",
12170 self.max_moe_block(),
12171 ids.len()
12172 );
12173 for (layer, proj, ex) in &ids {
12174 body.push_str(&format!("{layer} {proj} {ex}\n"));
12175 }
12176 let tmp = path.with_extension("tmp");
12177 std::fs::write(&tmp, body)?;
12178 std::fs::rename(&tmp, path)?;
12179 println!(
12180 "[moe-cache] freeze profile saved: {} blocks -> {}",
12181 ids.len(),
12182 path.display()
12183 );
12184 Ok(())
12185 }
12186
12187 pub fn restore_cpu_expert_residency_profile(
12191 &self,
12192 e: &Engine,
12193 path: &std::path::Path,
12194 ) -> Result<bool, Box<dyn std::error::Error>> {
12195 use crate::hybrid::Ffn;
12196 use crate::moe_cache::BlockId;
12197 let Ok(content) = std::fs::read_to_string(path) else {
12198 return Ok(false);
12199 };
12200 let mut lines = content.lines();
12201 let Some(header) = lines.next() else {
12202 return Ok(false);
12203 };
12204 let expected = format!("memra-freeze-profile v1 max_block={}", self.max_moe_block());
12205 if !header.starts_with(&expected) {
12206 println!(
12207 "[moe-cache] freeze profile ignored (geometry mismatch): {}",
12208 path.display()
12209 );
12210 return Ok(false);
12211 }
12212 let mut by_layer: std::collections::HashMap<u16, Vec<BlockId>> =
12213 std::collections::HashMap::new();
12214 for line in lines {
12215 let mut fields = line.split_whitespace();
12216 let (Some(layer), Some(proj), Some(ex)) = (fields.next(), fields.next(), fields.next())
12217 else {
12218 continue;
12219 };
12220 let (Ok(layer), Ok(proj), Ok(ex)) =
12221 (layer.parse::<u16>(), proj.parse::<u8>(), ex.parse::<u16>())
12222 else {
12223 continue;
12224 };
12225 by_layer
12226 .entry(layer)
12227 .or_default()
12228 .push(BlockId::new(layer, proj, ex));
12229 }
12230 let requested: usize = by_layer.values().map(Vec::len).sum();
12231 if requested == 0 {
12232 return Ok(false);
12233 }
12234 let max_block = self.max_moe_block();
12235 let mut restaged = 0usize;
12236 let mut stage_layer =
12237 |layer_index: u16, ffn: &Ffn| -> Result<(), Box<dyn std::error::Error>> {
12238 let Ffn::Moe(m) = ffn else { return Ok(()) };
12239 let Some(ids) = by_layer.get(&layer_index) else {
12240 return Ok(());
12241 };
12242 e.with_moe_cache(max_block, |cache, eng| {
12243 for id in ids {
12244 if cache.restage_block(*id, m, eng)? {
12245 restaged += 1;
12246 }
12247 }
12248 Ok(())
12249 })
12250 };
12251 for (index, layer) in self.layers.iter().enumerate() {
12252 stage_layer(index as u16, &layer.ffn)?;
12253 }
12254 if let Some(mtp) = self.mtp.as_ref() {
12255 stage_layer(u16::MAX, &mtp.ffn)?;
12256 }
12257 e.freeze_moe_cache();
12258 println!(
12259 "[moe-cache] freeze profile restored: {restaged}/{requested} blocks restaged from {}",
12260 path.display()
12261 );
12262 Ok(true)
12263 }
12264
12265 pub fn freeze_cpu_expert_residency(
12267 &self,
12268 e: &Engine,
12269 ) -> Result<(), Box<dyn std::error::Error>> {
12270 e.freeze_moe_cache();
12271 Ok(())
12272 }
12273
12274 pub fn ffn_act(
12282 e: &Engine,
12283 cfg: &ModelConfig,
12284 gate: &CudaSlice<f32>,
12285 up: &CudaSlice<f32>,
12286 act: &mut CudaSlice<f32>,
12287 n: usize,
12288 ) -> Result<(), Box<dyn std::error::Error>> {
12289 Self::ffn_act_scaled(e, cfg, gate, up, 1.0, 1.0, act, n)
12290 }
12291
12292 #[allow(clippy::too_many_arguments)]
12296 pub(crate) fn ffn_act_scaled(
12297 e: &Engine,
12298 cfg: &ModelConfig,
12299 gate: &CudaSlice<f32>,
12300 up: &CudaSlice<f32>,
12301 gs: f32,
12302 us: f32,
12303 act: &mut CudaSlice<f32>,
12304 n: usize,
12305 ) -> Result<(), Box<dyn std::error::Error>> {
12306 Self::ffn_act_lim(e, cfg, gate, up, gs, us, None, act, n)
12307 }
12308
12309 #[allow(clippy::too_many_arguments)]
12320 pub(crate) fn ffn_act_lim(
12321 e: &Engine,
12322 cfg: &ModelConfig,
12323 gate: &CudaSlice<f32>,
12324 up: &CudaSlice<f32>,
12325 gs: f32,
12326 us: f32,
12327 limit: Option<SwigluClamp>,
12328 act: &mut CudaSlice<f32>,
12329 n: usize,
12330 ) -> Result<(), Box<dyn std::error::Error>> {
12331 if let Some(m3) = cfg.m3.as_ref() {
12332 debug_assert!(
12333 limit.is_none(),
12334 "m3 swigluoai and the step35/glm5_next clamps are different archs"
12335 );
12336 return e.swigluoai_mul_scaled(
12337 gate,
12338 up,
12339 gs,
12340 us,
12341 m3.swiglu_alpha,
12342 m3.swiglu_limit,
12343 act,
12344 n,
12345 );
12346 }
12347 match limit {
12348 Some(SwigluClamp::Post(l)) => {
12349 return e.swiglu_clamped_mul_scaled(gate, up, gs, us, l, act, n);
12350 }
12351 Some(SwigluClamp::Pre(l)) => {
12352 return e.swiglu_preclamped_mul_scaled(gate, up, gs, us, l, act, n);
12353 }
12354 None => {}
12355 }
12356 if gs == 1.0 && us == 1.0 {
12357 return e.silu_mul(gate, up, act, n);
12358 }
12359 e.silu_mul_scaled(gate, up, gs, us, act, n)
12360 }
12361
12362 fn fused_post_limit(lim: Option<SwigluClamp>) -> Result<Option<f32>, ()> {
12368 match lim {
12369 None => Ok(None),
12370 Some(SwigluClamp::Post(l)) => Ok(Some(l)),
12371 Some(SwigluClamp::Pre(_)) => Err(()),
12372 }
12373 }
12374
12375 fn moe_route(
12381 e: &Engine,
12382 logits: &CudaSlice<f32>,
12383 t: usize,
12384 n_expert: usize,
12385 n_used: usize,
12386 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
12387 Self::moe_route_cfg(e, logits, t, n_expert, n_used, None)
12388 }
12389
12390 #[allow(clippy::too_many_arguments)]
12398 fn moe_route_sigmoid_cfg(
12399 e: &Engine,
12400 logits: &CudaSlice<f32>,
12401 t: usize,
12402 n_expert: usize,
12403 n_used: usize,
12404 m: &MoeWeights,
12405 (sf, route_norm): (f32, bool),
12406 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
12407 if sigmoid_router_enabled() {
12408 return e.moe_router_sigmoid_topk_host(
12409 logits,
12410 t,
12411 n_expert,
12412 n_used,
12413 m.active_count(),
12414 &m.exp_probs_b_dev,
12415 &m.active_experts_dev,
12416 sf,
12417 route_norm,
12418 );
12419 }
12420 let lg = e.dtoh(logits)?;
12421 Self::moe_route_sigmoid_host(
12422 &lg,
12423 t,
12424 n_expert,
12425 n_used,
12426 m.exp_probs_b.as_deref(),
12427 sf,
12428 route_norm,
12429 m.active_experts.as_deref(),
12430 )
12431 }
12432
12433 #[allow(clippy::excessive_precision)] fn moe_route_cfg(
12437 e: &Engine,
12438 logits: &CudaSlice<f32>,
12439 t: usize,
12440 n_expert: usize,
12441 n_used: usize,
12442 active: Option<&[bool]>,
12443 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
12444 if active.is_none() && !matches!(std::env::var("MEMRA_FUSED_ROUTER").as_deref(), Ok("0")) {
12447 return e.moe_router_topk_host(logits, t, n_expert, n_used);
12448 }
12449 let lg = e.dtoh(logits)?; let mut sel = vec![0u32; t * n_used];
12452 let mut w_out = vec![0f32; t * n_used];
12453 for tok in 0..t {
12454 let row = &lg[tok * n_expert..(tok + 1) * n_expert];
12455 let maxl = row
12457 .iter()
12458 .enumerate()
12459 .filter(|(i, _)| active.is_none_or(|mask| mask[*i]))
12460 .map(|(_, &x)| x)
12461 .fold(f32::NEG_INFINITY, f32::max);
12462 let mut probs = vec![0f32; n_expert];
12463 let mut den = 0f32;
12464 for i in 0..n_expert {
12465 if active.is_some_and(|mask| !mask[i]) {
12466 continue;
12467 }
12468 let x = (row[i] - maxl).exp();
12469 probs[i] = x;
12470 den += x;
12471 }
12472 for p in probs.iter_mut() {
12473 *p /= den;
12474 }
12475 let mut idx: Vec<usize> = (0..n_expert)
12477 .filter(|&i| active.is_none_or(|mask| mask[i]))
12478 .collect();
12479 idx.sort_by(|&a, &b| probs[b].total_cmp(&probs[a]).then(a.cmp(&b)));
12480 let sl = &idx[..n_used];
12481 let mut wv: Vec<f32> = sl.iter().map(|&i| probs[i]).collect();
12482 let mut ws: f32 = wv.iter().sum();
12483 ws = ws.max(6.103515625e-5_f32); for x in wv.iter_mut() {
12485 *x /= ws;
12486 }
12487 for j in 0..n_used {
12488 sel[tok * n_used + j] = sl[j] as u32;
12489 w_out[tok * n_used + j] = wv[j];
12490 }
12491 }
12492 Ok((sel, w_out))
12493 }
12494
12495 #[allow(clippy::too_many_arguments)]
12496 #[allow(clippy::type_complexity)] fn moe_route_sigmoid_with_input(
12498 e: &Engine,
12499 logits: &CudaSlice<f32>,
12500 input: &CudaSlice<f32>,
12501 t: usize,
12502 in_features: usize,
12503 n_expert: usize,
12504 n_used: usize,
12505 bias: Option<&[f32]>,
12506 (sf, route_norm): (f32, bool),
12507 active: Option<&[bool]>,
12508 ) -> Result<(Vec<u32>, Vec<f32>, Vec<f32>), Box<dyn std::error::Error>> {
12509 let logit_values =
12510 active_matrix_values(logits.len(), t, n_expert, "sigmoid router logits")?;
12511 let input_values =
12512 active_matrix_values(input.len(), t, in_features, "sigmoid router input")?;
12513 let (lg, input) = e.dtoh_pair_views(
12514 &logits.slice(0..logit_values),
12515 &input.slice(0..input_values),
12516 )?;
12517 let (sel, w) =
12518 Self::moe_route_sigmoid_host(&lg, t, n_expert, n_used, bias, sf, route_norm, active)?;
12519 Ok((sel, w, input))
12520 }
12521
12522 pub fn start_moe_prefetch_predictor(
12527 &self,
12528 e: &Engine,
12529 cfg: &ModelConfig,
12530 ) -> Result<(), Box<dyn std::error::Error>> {
12531 use crate::hybrid::Ffn;
12532 let Some(sig) = cfg.sigmoid_router() else {
12533 return Err("prefetch predictor requires a sigmoid-router arch".into());
12534 };
12535 let resident: std::collections::HashSet<(u16, u8, u16)> = e
12536 .export_moe_residency()
12537 .ok_or("prefetch predictor needs the frozen MoE residency cache")?
12538 .into_iter()
12539 .collect();
12540 let mut layers = Vec::new();
12541 for (index, layer) in self.layers.iter().enumerate() {
12542 let Ffn::Moe(m) = &layer.ffn else { continue };
12543 let crate::model::GpuTensor::Float { data, .. } = &m.gate_inp else {
12544 continue;
12545 };
12546 let router = e.dtoh(data)?;
12547 let n_expert = m.gate_exps.n_expert;
12548 let n_embd = m.gate_exps.in_f;
12549 if router.len() != n_embd * n_expert {
12550 continue;
12551 }
12552 let build = |exps: &crate::model::HostExps| {
12553 (0..n_expert)
12554 .map(|expert| crate::cpu_experts::predictor_projection(exps, expert))
12555 .collect::<Vec<_>>()
12556 };
12557 layers.push((
12558 index as u16,
12559 crate::cpu_experts::PredictLayerInit {
12560 router,
12561 bias: m.exp_probs_b.clone(),
12562 active: m.active_experts.clone(),
12563 n_embd,
12564 n_used: cfg
12565 .moe
12566 .as_ref()
12567 .map(|moe| moe.expert_used_count as usize)
12568 .ok_or("prefetch predictor requires MoE config")?,
12569 sig,
12570 weights_n_expert: n_expert,
12571 gate: build(&m.gate_exps),
12572 up: build(&m.up_exps),
12573 down: build(&m.down_exps),
12574 },
12575 ));
12576 }
12577 crate::cpu_experts::start_prefetch_predictor(layers, resident).map_err(|error| error.into())
12578 }
12579
12580 #[allow(clippy::too_many_arguments)]
12583 pub fn moe_route_sigmoid_host_public(
12584 logits: &[f32],
12585 t: usize,
12586 n_expert: usize,
12587 n_used: usize,
12588 bias: Option<&[f32]>,
12589 sf: f32,
12590 route_norm: bool,
12591 active: Option<&[bool]>,
12592 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
12593 Self::moe_route_sigmoid_host(logits, t, n_expert, n_used, bias, sf, route_norm, active)
12594 }
12595
12596 #[allow(clippy::too_many_arguments)]
12597 fn moe_route_sigmoid_host(
12598 lg: &[f32],
12599 t: usize,
12600 n_expert: usize,
12601 n_used: usize,
12602 bias: Option<&[f32]>,
12603 sf: f32,
12604 route_norm: bool,
12605 active: Option<&[bool]>,
12606 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
12607 let active_count = active
12608 .map(|mask| mask.iter().filter(|&&enabled| enabled).count())
12609 .unwrap_or(n_expert);
12610 crate::sigrouter_contract::validate_active_count(n_used, active_count)?;
12611 if lg.len() != t * n_expert {
12612 return Err(format!(
12613 "sigmoid router logits length mismatch: got {}, expected {}",
12614 lg.len(),
12615 t * n_expert,
12616 )
12617 .into());
12618 }
12619 let mut sel = vec![0u32; t * n_used];
12620 let mut w_out = vec![0f32; t * n_used];
12621 for tok in 0..t {
12622 let row = &lg[tok * n_expert..(tok + 1) * n_expert];
12623 let scores: Vec<f32> = row.iter().map(|&x| 1.0 / (1.0 + (-x).exp())).collect();
12624 let selsc: Vec<f32> = match bias {
12626 Some(b) => scores.iter().zip(b).map(|(s, bb)| s + bb).collect(),
12627 None => scores.clone(),
12628 };
12629 let mut idx: Vec<usize> = (0..n_expert)
12630 .filter(|&i| active.is_none_or(|mask| mask[i]))
12631 .collect();
12632 idx.sort_by(|&a, &b| selsc[b].total_cmp(&selsc[a]).then(a.cmp(&b)));
12633 let sl = &idx[..n_used];
12634 let mut wv: Vec<f32> = sl.iter().map(|&i| scores[i]).collect();
12635 if route_norm {
12636 let ws: f32 = wv.iter().sum::<f32>().max(1e-20);
12637 for x in wv.iter_mut() {
12638 *x = *x / ws * sf;
12639 }
12640 } else {
12641 for x in wv.iter_mut() {
12642 *x *= sf;
12643 }
12644 }
12645 for j in 0..n_used {
12646 sel[tok * n_used + j] = sl[j] as u32;
12647 w_out[tok * n_used + j] = wv[j];
12648 }
12649 }
12650 Ok((sel, w_out))
12651 }
12652
12653 #[allow(clippy::too_many_arguments)]
12657 fn moe_ffn_sigmoid_dev(
12658 e: &Engine,
12659 m: &MoeWeights,
12660 z: &CudaSlice<f32>,
12661 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
12662 logits: &CudaSlice<f32>,
12663 t: usize,
12664 cfg: &ModelConfig,
12665 il: u16,
12666 (scaling_factor, route_norm): (f32, bool),
12667 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12668 let moe = cfg.moe.as_ref().unwrap();
12669 let n_embd = cfg.n_embd as usize;
12670 let n_expert = moe.expert_count as usize;
12671 let n_used = moe.expert_used_count as usize;
12672 let n_ff_exp = moe.expert_ff_length as usize;
12673 let dev = m.dev_exps.as_ref().unwrap();
12674 debug_assert_eq!(dev.dev, e.ctx().ordinal());
12675 debug_assert!(m.has_uniform_expert_layout());
12676 debug_assert!(!m.has_macros);
12677
12678 let (sel_d, w_d) = e.moe_router_sigmoid_topk(
12679 logits,
12680 t,
12681 n_expert,
12682 n_used,
12683 m.active_count(),
12684 &m.exp_probs_b_dev,
12685 &m.active_experts_dev,
12686 scaling_factor,
12687 route_norm,
12688 )?;
12689 crate::moesd::record_device_routes(e, il, n_expert, n_used, &sel_d)?;
12690 if let Some(fp8) = dev.fp8_blk.as_ref() {
12691 debug_assert_eq!(m.gate_exps.qtype, crate::QT_F8_E4M3_BLK);
12692 debug_assert_eq!(m.up_exps.qtype, crate::QT_F8_E4M3_BLK);
12693 debug_assert_eq!(m.down_exps.qtype, crate::QT_F8_E4M3_BLK);
12694 debug_assert_eq!(fp8.gate.rows, m.gate_exps.out_f.div_ceil(128));
12695 debug_assert_eq!(fp8.up.rows, m.up_exps.out_f.div_ceil(128));
12696 debug_assert_eq!(fp8.down.rows, m.down_exps.out_f.div_ceil(128));
12697
12698 let selected = e.dtoh_i32(&sel_d)?;
12705 let route_weights = e.dtoh(&w_d)?;
12706 let mut moe_out = e.zeros(t * n_embd)?;
12707 for tok in 0..t {
12708 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
12709 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
12710 for j in 0..n_used {
12711 let pair = tok * n_used + j;
12712 let expert = selected[pair] as usize;
12713 let gate = Self::moe_resident_fp8_e4m3(
12714 e,
12715 &m.gate_exps,
12716 &dev.gate,
12717 &fp8.gate,
12718 expert,
12719 &zt,
12720 1,
12721 )?;
12722 let up = Self::moe_resident_fp8_e4m3(
12723 e, &m.up_exps, &dev.up, &fp8.up, expert, &zt, 1,
12724 )?;
12725 let mut act = e.uninit(n_ff_exp)?;
12726 Self::ffn_act_lim(
12727 e,
12728 cfg,
12729 &gate,
12730 &up,
12731 1.0,
12732 1.0,
12733 cfg.clamp_exp_at(il as u32),
12734 &mut act,
12735 n_ff_exp,
12736 )?;
12737 let act = act.slice(0..n_ff_exp);
12738 let down = Self::moe_resident_fp8_e4m3(
12739 e,
12740 &m.down_exps,
12741 &dev.down,
12742 &fp8.down,
12743 expert,
12744 &act,
12745 1,
12746 )?;
12747 e.axpy_into(&down, route_weights[pair], &mut dst, n_embd)?;
12748 }
12749 }
12750 if std::env::var("MEMRA_SIG_ROUTER_DISPATCH_TRACE").as_deref() == Ok("1") {
12751 eprintln!(
12752 "[sigrouter-dev] layer={il} tokens={t} experts={n_expert} used={n_used} \
12753 native=fp8blk-w8a8-e4m3-reference clamp={}",
12754 cfg.clamp_exp_at(il as u32).is_some(),
12755 );
12756 }
12757 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut moe_out)?;
12758 return Ok(moe_out);
12759 }
12760 let (gate_row_bytes, up_row_bytes) = if dev.gu_il {
12761 let combined = m.gate_exps.row_bytes + m.up_exps.row_bytes;
12762 (combined, combined)
12763 } else {
12764 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
12765 };
12766 let (zq, zd) = match (t, zq8) {
12767 (1, Some((q, d))) => (q.clone(), d.clone()),
12768 _ => e.quantize_q8_1(z, t, n_embd)?,
12769 };
12770 let n_pairs = t * n_used;
12771 let mut moe_out = if cfg.clamp_exp_at(il as u32).is_some() {
12772 let pair_tok: Vec<i32> = (0..n_pairs).map(|pair| (pair / n_used) as i32).collect();
12776 let pair_tok_d = e.htod_i32(&pair_tok)?;
12777 let gate = e.moe_pairs_matvec_q8(
12778 &dev.ptr_row,
12779 0,
12780 &pair_tok_d,
12781 &sel_d,
12782 &zq,
12783 &zd,
12784 n_embd,
12785 n_ff_exp,
12786 n_expert,
12787 n_pairs,
12788 m.gate_exps.qtype,
12789 gate_row_bytes,
12790 )?;
12791 let up = e.moe_pairs_matvec_q8(
12792 &dev.ptr_row,
12793 1,
12794 &pair_tok_d,
12795 &sel_d,
12796 &zq,
12797 &zd,
12798 n_embd,
12799 n_ff_exp,
12800 n_expert,
12801 n_pairs,
12802 m.up_exps.qtype,
12803 up_row_bytes,
12804 )?;
12805 let mut act = e.uninit(n_pairs * n_ff_exp)?;
12806 Self::ffn_act_lim(
12807 e,
12808 cfg,
12809 &gate,
12810 &up,
12811 1.0,
12812 1.0,
12813 cfg.clamp_exp_at(il as u32),
12814 &mut act,
12815 n_pairs * n_ff_exp,
12816 )?;
12817 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
12818 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
12819 let pair_self_d = e.htod_i32(&pair_self)?;
12820 let down = e.moe_pairs_matvec_q8(
12821 &dev.ptr_row,
12822 2,
12823 &pair_self_d,
12824 &sel_d,
12825 &aq2,
12826 &ad2,
12827 n_ff_exp,
12828 n_embd,
12829 n_expert,
12830 n_pairs,
12831 m.down_exps.qtype,
12832 m.down_exps.row_bytes,
12833 )?;
12834 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
12835 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
12836 let tok_off_d = e.htod_i32(&tok_off)?;
12837 let tok_ids_d = e.htod_i32(&tok_ids)?;
12838 let mut output = e.uninit(t * n_embd)?;
12839 e.moe_pairs_scatter(&down, &w_d, &tok_off_d, &tok_ids_d, &mut output, t, n_embd)?;
12840 output
12841 } else {
12842 let act = e.moe_gate_up_silu8_dev_q8_rows(
12843 &dev.ptr_row,
12844 &sel_d,
12845 &zq,
12846 &zd,
12847 t,
12848 n_embd,
12849 n_ff_exp,
12850 n_used,
12851 n_expert,
12852 m.gate_exps.qtype,
12853 m.up_exps.qtype,
12854 gate_row_bytes,
12855 up_row_bytes,
12856 &m.dev_macros,
12857 )?;
12858 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
12859 let mut output = e.uninit(t * n_embd)?;
12860 e.moe_down8_fma_dev_q8_rows_g(
12861 &dev.ptr_row,
12862 &sel_d,
12863 &w_d,
12864 &aq2,
12865 &ad2,
12866 &mut output,
12867 t,
12868 n_ff_exp,
12869 n_embd,
12870 n_used,
12871 n_expert,
12872 m.down_exps.qtype,
12873 m.down_exps.row_bytes,
12874 )?;
12875 output
12876 };
12877
12878 if std::env::var("MEMRA_SIG_ROUTER_DISPATCH_TRACE").as_deref() == Ok("1") {
12879 eprintln!(
12880 "[sigrouter-dev] layer={il} tokens={t} experts={n_expert} used={n_used} clamp={} gu_il={}",
12881 cfg.clamp_exp_at(il as u32).is_some(),
12882 dev.gu_il,
12883 );
12884 }
12885 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut moe_out)?;
12886 Ok(moe_out)
12887 }
12888
12889 #[allow(clippy::too_many_arguments)]
12890 fn moe_resident_fp8_e4m3(
12891 e: &Engine,
12892 exps: &crate::model::HostExps,
12893 bytes: &CudaSlice<u8>,
12894 scales: &crate::hybrid::DevExpertFp8ProjectionScales,
12895 expert: usize,
12896 x: &cudarc::driver::CudaView<f32>,
12897 m: usize,
12898 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12899 let layout = exps.expert_layout(expert);
12900 debug_assert_eq!(layout.qtype, crate::QT_F8_E4M3_BLK);
12901 debug_assert_eq!(scales.rows * scales.cols, scales.expert_stride);
12902 let byte_start = expert * exps.expert_stride;
12903 let scale_start = expert * scales.expert_stride;
12904 let weight = bytes.slice(byte_start..byte_start + layout.len);
12905 let scale = scales
12906 .scales
12907 .slice(scale_start..scale_start + scales.expert_stride);
12908 e.qmatvec_mmq_fp8_blk_view(&weight, &scale, x, m, exps.in_f, exps.out_f)
12909 }
12910
12911 fn moe_ffn_pairs(
12920 e: &Engine,
12921 m: &MoeWeights,
12922 z: &CudaSlice<f32>,
12923 logits: &CudaSlice<f32>,
12924 t: usize,
12925 cfg: &ModelConfig,
12926 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12927 let moe = cfg.moe.as_ref().unwrap();
12928 let n_embd = cfg.n_embd as usize;
12929 let n_expert = moe.expert_count as usize;
12930 let n_used = moe.expert_used_count as usize;
12931 let n_ff_exp = moe.expert_ff_length as usize;
12932 debug_assert!(
12937 !cfg.swiglu_clamped_anywhere(),
12938 "moe_ffn_pairs has no per-layer clamp: fused epilogues are plain SiLU"
12939 );
12940 let dev = m.dev_exps.as_ref().unwrap();
12941 let (rbg_d, rbu_d) = if dev.gu_il {
12943 let sxx = m.gate_exps.row_bytes + m.up_exps.row_bytes;
12944 (sxx, sxx)
12945 } else {
12946 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
12947 };
12948
12949 let (sel_all, w_all) = Self::moe_route(e, logits, t, n_expert, n_used)?;
12950 let n_pairs = t * n_used;
12951 let pair_tok: Vec<i32> = (0..n_pairs).map(|p| (p / n_used) as i32).collect();
12954 let pair_ex: Vec<i32> = sel_all.iter().map(|&x| x as i32).collect();
12955 let pair_w: Vec<f32> = w_all.clone();
12956 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
12957 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
12958 let pt = e.htod_i32(&pair_tok)?;
12959 let px = e.htod_i32(&pair_ex)?;
12960 let pw = e.htod(&pair_w)?;
12961 let toff = e.htod_i32(&tok_off)?;
12962 let tids = e.htod_i32(&tok_ids)?;
12963
12964 let mut by_ex: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
12968 for p in 0..n_pairs {
12969 by_ex[pair_ex[p] as usize].push(p as i32);
12970 }
12971 let mut ex_ids: Vec<i32> = Vec::new();
12972 let mut ex_off: Vec<i32> = vec![0];
12973 let mut ex_pairs: Vec<i32> = Vec::with_capacity(n_pairs);
12974 for (ex, list) in by_ex.iter().enumerate() {
12975 if list.is_empty() {
12976 continue;
12977 }
12978 ex_ids.push(ex as i32);
12979 ex_pairs.extend_from_slice(list);
12980 ex_off.push(ex_pairs.len() as i32);
12981 }
12982 let n_active = ex_ids.len();
12983 let exi = e.htod_i32(&ex_ids)?;
12984 let exo = e.htod_i32(&ex_off)?;
12985 let exp_d = e.htod_i32(&ex_pairs)?;
12986 let _ = &px; static MMA_T: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
13007 let mma_t = *MMA_T.get_or_init(|| {
13008 std::env::var("MEMRA_MOE_MMA_T")
13009 .ok()
13010 .and_then(|v| v.parse().ok())
13011 .unwrap_or(16)
13012 });
13013 let use_mma = std::env::var("MEMRA_MOE_MMA")
13014 .map(|v| v != "0")
13015 .unwrap_or(true)
13016 && t >= mma_t
13017 && q8_expert_dec_supported(m.gate_exps.qtype)
13018 && q8_expert_dec_supported(m.up_exps.qtype)
13019 && q8_expert_dec_supported(m.down_exps.qtype)
13020 && n_embd.is_multiple_of(256)
13021 && n_ff_exp.is_multiple_of(256);
13022 let mma_capable = q8_expert_dec_supported(m.gate_exps.qtype)
13038 && q8_expert_dec_supported(m.up_exps.qtype)
13039 && q8_expert_dec_supported(m.down_exps.qtype)
13040 && n_embd.is_multiple_of(256)
13041 && n_ff_exp.is_multiple_of(256);
13042 let f16g_mode = crate::moe_f16g_mode();
13043 let f16g = f16g_mode != 0
13044 && t >= mma_t
13045 && (f16g_mode != 3 || !mma_capable)
13046 && f16g_proj_ok(m.gate_exps.qtype, n_embd)
13047 && f16g_proj_ok(m.up_exps.qtype, n_embd)
13048 && f16g_proj_ok(m.down_exps.qtype, n_ff_exp);
13049 if use_mma || f16g {
13050 let y_down = if f16g {
13058 let csr_tok: Vec<i32> = ex_pairs.iter().map(|&p| p / n_used as i32).collect();
13062 let csr_tok_d = e.htod_i32(&csr_tok)?;
13063 let (z_f16, z_s) = e.moe_f16g_act(z, Some(&csr_tok_d), n_embd, n_pairs)?;
13064 let g_csr = e.moe_f16_grouped(
13065 &dev.ptr_row,
13066 0,
13067 n_expert,
13068 &exi,
13069 &ex_off,
13070 &exo,
13071 &z_f16,
13072 &z_s,
13073 n_embd,
13074 n_ff_exp,
13075 n_active,
13076 n_pairs,
13077 m.gate_exps.qtype,
13078 rbg_d,
13079 )?;
13080 let u_csr = e.moe_f16_grouped(
13081 &dev.ptr_row,
13082 1,
13083 n_expert,
13084 &exi,
13085 &ex_off,
13086 &exo,
13087 &z_f16,
13088 &z_s,
13089 n_embd,
13090 n_ff_exp,
13091 n_active,
13092 n_pairs,
13093 m.up_exps.qtype,
13094 rbu_d,
13095 )?;
13096 let act_csr = e.moe_pairs_silu_mul(&g_csr, &u_csr, n_pairs * n_ff_exp)?;
13097 let (a_f16, a_s) = e.moe_f16g_act(&act_csr, None, n_ff_exp, n_pairs)?;
13098 let d_csr = e.moe_f16_grouped(
13099 &dev.ptr_row,
13100 2,
13101 n_expert,
13102 &exi,
13103 &ex_off,
13104 &exo,
13105 &a_f16,
13106 &a_s,
13107 n_ff_exp,
13108 n_embd,
13109 n_active,
13110 n_pairs,
13111 m.down_exps.qtype,
13112 m.down_exps.row_bytes,
13113 )?;
13114 e.rows_permute(&d_csr, &exp_d, n_pairs, n_embd)?
13115 } else {
13116 let z_scr = e.mmq_iq_quantize_act(z, n_embd, t)?;
13118 let gate = e.mmq_iq_experts(
13119 &dev.ptr_row,
13120 0,
13121 n_expert,
13122 &exi,
13123 &exo,
13124 &exp_d,
13125 &pt,
13126 &z_scr,
13127 n_embd,
13128 n_ff_exp,
13129 n_active,
13130 n_pairs,
13131 t,
13132 m.gate_exps.qtype,
13133 rbg_d,
13134 )?;
13135 let up = e.mmq_iq_experts(
13136 &dev.ptr_row,
13137 1,
13138 n_expert,
13139 &exi,
13140 &exo,
13141 &exp_d,
13142 &pt,
13143 &z_scr,
13144 n_embd,
13145 n_ff_exp,
13146 n_active,
13147 n_pairs,
13148 t,
13149 m.up_exps.qtype,
13150 rbu_d,
13151 )?;
13152 let a_scr = if crate::moe_fuse_actq_on() {
13158 e.mmq_iq_fused_act_quant(&gate, &up, n_ff_exp, n_pairs, 0)?
13159 } else {
13160 let act = e.moe_pairs_silu_mul(&gate, &up, n_pairs * n_ff_exp)?;
13161 e.mmq_iq_quantize_act(&act, n_ff_exp, n_pairs)?
13162 };
13163 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
13164 let pself = e.htod_i32(&pair_self)?;
13165 e.mmq_iq_experts(
13166 &dev.ptr_row,
13167 2,
13168 n_expert,
13169 &exi,
13170 &exo,
13171 &exp_d,
13172 &pself,
13173 &a_scr,
13174 n_ff_exp,
13175 n_embd,
13176 n_active,
13177 n_pairs,
13178 n_pairs,
13179 m.down_exps.qtype,
13180 m.down_exps.row_bytes,
13181 )?
13182 };
13183 let mut moe_out = e.uninit(t * n_embd)?;
13184 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
13185 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
13186 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
13187 {
13188 let n_ff_sh = gate_shexp.out_features();
13189 let sg_gate = e.matmul(gate_shexp, z, t)?;
13190 let sg_up = e.matmul(up_shexp, z, t)?;
13191 let mut sa = e.uninit(t * n_ff_sh)?;
13192 Self::ffn_act(e, cfg, &sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
13193 let sh = e.matmul(down_shexp, &sa, t)?;
13194 let g = match &m.gate_inp_shexp {
13200 Some(gate_inp_shexp) if crate::router_prefill_exact_on() => {
13201 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
13202 }
13203 Some(gate_inp_shexp) => {
13204 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
13205 let mut g = e.uninit(t)?;
13206 e.sigmoid(&gs, &mut g, t)?;
13207 g
13208 }
13209 None => e.htod(&vec![1.0f32; t])?,
13210 };
13211 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
13212 }
13213 return Ok(moe_out);
13214 }
13215
13216 let dec = std::env::var("MEMRA_MOE_DEC")
13219 .map(|v| v != "0")
13220 .unwrap_or(true);
13221 let matvec = |proj,
13222 exi: &_,
13223 exo: &_,
13224 exp_d: &_,
13225 pt: &_,
13226 aq: &_,
13227 ad: &_,
13228 inf,
13229 outf,
13230 qtype,
13231 rb|
13232 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
13233 let dec = dec && q8_expert_dec_supported(qtype);
13235 if dec {
13236 e.moe_pairs_matvec_q8_dec(
13237 &dev.ptr_row,
13238 proj,
13239 exi,
13240 exo,
13241 exp_d,
13242 pt,
13243 aq,
13244 ad,
13245 inf,
13246 outf,
13247 n_expert,
13248 n_active,
13249 n_pairs,
13250 qtype,
13251 rb,
13252 )
13253 } else {
13254 e.moe_pairs_matvec_q8_em(
13255 &dev.ptr_row,
13256 proj,
13257 exi,
13258 exo,
13259 exp_d,
13260 pt,
13261 aq,
13262 ad,
13263 inf,
13264 outf,
13265 n_expert,
13266 n_active,
13267 n_pairs,
13268 qtype,
13269 rb,
13270 )
13271 }
13272 };
13273 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
13274 let gate = matvec(
13275 0,
13276 &exi,
13277 &exo,
13278 &exp_d,
13279 &pt,
13280 &zq,
13281 &zd,
13282 n_embd,
13283 n_ff_exp,
13284 m.gate_exps.qtype,
13285 rbg_d,
13286 )?;
13287 let up = matvec(
13288 1,
13289 &exi,
13290 &exo,
13291 &exp_d,
13292 &pt,
13293 &zq,
13294 &zd,
13295 n_embd,
13296 n_ff_exp,
13297 m.up_exps.qtype,
13298 rbu_d,
13299 )?;
13300 let act = e.moe_pairs_silu_mul(&gate, &up, n_pairs * n_ff_exp)?;
13301 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
13302 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
13304 let pself = e.htod_i32(&pair_self)?;
13305 let y_down = matvec(
13306 2,
13307 &exi,
13308 &exo,
13309 &exp_d,
13310 &pself,
13311 &aq2,
13312 &ad2,
13313 n_ff_exp,
13314 n_embd,
13315 m.down_exps.qtype,
13316 m.down_exps.row_bytes,
13317 )?;
13318 let mut moe_out = e.uninit(t * n_embd)?; e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
13320
13321 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
13325 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
13326 {
13327 let n_ff_sh = gate_shexp.out_features();
13328 let step_exact = true;
13332 let verify_t = step_exact && t > 1 && t < PRIME_MIN_T;
13333 let (sg_gate, sg_up) = if step_exact && t == 1 {
13334 shexp_gate_up_t1(e, gate_shexp, up_shexp, z, None)?
13335 } else if verify_t {
13336 let mut fused = None;
13337 if crate::spec::spec_fused_t()
13338 && (2..=4).contains(&t)
13339 && e.uses_q8_1_fast(gate_shexp)
13340 && e.uses_q8_1_fast(up_shexp)
13341 {
13342 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
13343 fused = e.matmul_q8_fused2_t(gate_shexp, up_shexp, &zq, &zd, t)?;
13344 }
13345 match fused {
13346 Some(pair) => pair,
13347 None => (
13348 e.matmul_decode_exact(gate_shexp, z, t)?,
13349 e.matmul_decode_exact(up_shexp, z, t)?,
13350 ),
13351 }
13352 } else {
13353 (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?)
13354 };
13355 let mut sa = e.uninit(t * n_ff_sh)?;
13356 e.silu_mul(&sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
13357 let sh = if verify_t {
13358 e.matmul_decode_exact(down_shexp, &sa, t)?
13359 } else {
13360 e.matmul(down_shexp, &sa, t)?
13361 };
13362 let g = match &m.gate_inp_shexp {
13367 Some(gate_inp_shexp) if crate::router_prefill_exact_on() => {
13368 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
13369 }
13370 Some(gate_inp_shexp) => {
13371 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
13372 let mut g = e.uninit(t)?;
13373 e.sigmoid(&gs, &mut g, t)?;
13374 g
13375 }
13376 None => e.htod(&vec![1.0f32; t])?,
13377 };
13378 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
13379 }
13380 Ok(moe_out)
13381 }
13382
13383 #[allow(clippy::too_many_arguments)]
13385 #[allow(clippy::too_many_arguments)]
13386 fn moe_ffn_dev(
13387 e: &Engine,
13388 m: &MoeWeights,
13389 z: &CudaSlice<f32>,
13390 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
13391 logits: &CudaSlice<f32>,
13392 t: usize,
13393 cfg: &ModelConfig,
13394 il: u16,
13395 max_block: usize,
13396 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
13397 let moe = cfg.moe.as_ref().unwrap();
13398 let n_embd = cfg.n_embd as usize;
13399 let n_expert = moe.expert_count as usize;
13400 let n_used = moe.expert_used_count as usize;
13401 let n_ff_exp = moe.expert_ff_length as usize;
13402 debug_assert!(
13406 cfg.sigmoid_router().is_none(),
13407 "moe_ffn_dev routes SOFTMAX: a sigmoid-router arch would pick wrong experts"
13408 );
13409 debug_assert!(
13410 !cfg.swiglu_clamped_at(il as u32),
13411 "moe_ffn_dev's fused epilogue is plain SiLU: no clamped form"
13412 );
13413
13414 let (sel_d, mut w_d) = e.moe_router_topk(logits, t, n_expert, n_used)?;
13416 if m.has_macros {
13419 e.moe_w_scale_by_expert(&mut w_d, &sel_d, &m.dev_macros, n_expert, t * n_used)?;
13420 }
13421
13422 let mut moe_out = e.uninit(t * n_embd)?;
13424
13425 if let Some(dev) = m.dev_exps.as_ref() {
13428 let (rbg_d, rbu_d) = if dev.gu_il {
13431 let sxx = m.gate_exps.row_bytes + m.up_exps.row_bytes;
13432 (sxx, sxx)
13433 } else {
13434 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
13435 };
13436 let q8 = moe_q8_enabled_for_model(cfg, m);
13437 let rows_arm = q8
13446 && t > 1
13447 && crate::spec::spec_m2()
13448 && n_ff_exp == 512
13449 && n_used <= 8
13450 && std::env::var("MEMRA_MOE_DEVQ8_GU")
13451 .map(|v| v.is_empty() || v == "v")
13452 .unwrap_or(true)
13453 && std::env::var("MEMRA_MOE_DEVQ8_DOWN")
13454 .map(|v| v.is_empty() || v == "w8h2v")
13455 .unwrap_or(true);
13456 let csr_mode = std::env::var("MEMRA_MOE_CSR")
13465 .ok()
13466 .and_then(|v| v.parse::<i32>().ok())
13467 .unwrap_or(1);
13468 let csr_nvfp4_probe = std::env::var("MEMRA_MOE_CSR_NVFP4").as_deref() == Ok("1");
13485 let csr_qt = |qt: i32| {
13486 qt == crate::QT_IQ4_XS
13487 || qt == crate::QT_IQ3_S
13488 || (csr_nvfp4_probe && qt == crate::QT_NVFP4)
13489 };
13490 let csr_t_max = if csr_nvfp4_probe { MOE_DEV_MAX_T } else { 10 };
13491 let csr_uniform = m.gate_exps.qtype == m.up_exps.qtype;
13492 let csr_arm = rows_arm
13493 && csr_mode > 0
13494 && t <= csr_t_max
13495 && csr_uniform
13496 && csr_qt(m.gate_exps.qtype)
13497 && csr_qt(m.up_exps.qtype)
13498 && csr_qt(m.down_exps.qtype);
13499 if csr_arm {
13500 if csr_mode == 2 {
13501 static ENGAGED: std::sync::Once = std::sync::Once::new();
13502 ENGAGED.call_once(|| eprintln!("[memra] moe CSR byte-compare mode ON (t={t})"));
13503 }
13504 let n_pairs = t * n_used;
13505 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
13506 let act = e.moe_gate_up_silu8_dev_q8_csr(
13507 &dev.ptr_row,
13508 &sel_d,
13509 &zq,
13510 &zd,
13511 n_pairs,
13512 n_embd,
13513 n_ff_exp,
13514 n_used,
13515 n_expert,
13516 m.gate_exps.qtype,
13517 m.up_exps.qtype,
13518 rbg_d,
13519 rbu_d,
13520 )?;
13521 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
13522 e.moe_down8_fma_dev_q8_rows(
13526 &dev.ptr_row,
13527 &sel_d,
13528 &w_d,
13529 &aq2,
13530 &ad2,
13531 &mut moe_out,
13532 t,
13533 n_ff_exp,
13534 n_embd,
13535 n_used,
13536 n_expert,
13537 m.down_exps.qtype,
13538 m.down_exps.row_bytes,
13539 )?;
13540 if csr_mode == 2 {
13541 let act_r = e.moe_gate_up_silu8_dev_q8_rows(
13543 &dev.ptr_row,
13544 &sel_d,
13545 &zq,
13546 &zd,
13547 t,
13548 n_embd,
13549 n_ff_exp,
13550 n_used,
13551 n_expert,
13552 m.gate_exps.qtype,
13553 m.up_exps.qtype,
13554 rbg_d,
13555 rbu_d,
13556 &m.dev_macros,
13557 )?;
13558 let mut out_r = e.uninit(t * n_embd)?;
13559 let (aq2r, ad2r) = e.quantize_q8_1(&act_r, n_pairs, n_ff_exp)?;
13560 e.moe_down8_fma_dev_q8_rows(
13561 &dev.ptr_row,
13562 &sel_d,
13563 &w_d,
13564 &aq2r,
13565 &ad2r,
13566 &mut out_r,
13567 t,
13568 n_ff_exp,
13569 n_embd,
13570 n_used,
13571 n_expert,
13572 m.down_exps.qtype,
13573 m.down_exps.row_bytes,
13574 )?;
13575 let (a1, a2) = (e.dtoh(&act)?, e.dtoh(&act_r)?);
13576 let (o1, o2) = (e.dtoh(&moe_out)?, e.dtoh(&out_r)?);
13577 let ba = a1
13578 .iter()
13579 .zip(&a2)
13580 .filter(|(x, y)| x.to_bits() != y.to_bits())
13581 .count();
13582 let bo = o1
13583 .iter()
13584 .zip(&o2)
13585 .filter(|(x, y)| x.to_bits() != y.to_bits())
13586 .count();
13587 if ba + bo > 0 {
13588 eprintln!(
13589 "[csr-check] il={il} t={t} ACT diffs={ba}/{} OUT diffs={bo}/{}",
13590 a1.len(),
13591 o1.len()
13592 );
13593 let sel_h = e.dtoh_i32(&sel_d)?;
13595 let mut shown = 0;
13596 for (i, (x, y)) in a1.iter().zip(&a2).enumerate() {
13597 if x.to_bits() != y.to_bits() && shown < 4 {
13598 let (p, o) = (i / n_ff_exp, i % n_ff_exp);
13599 let ex = sel_h[p];
13600 let npx = sel_h.iter().filter(|&&v| v == ex).count();
13601 eprintln!(
13602 " ACT p={p} ex={ex} np={npx} o={o} csr={x:e} rows={y:e}"
13603 );
13604 shown += 1;
13605 }
13606 }
13607 std::process::exit(3);
13608 }
13609 }
13610 } else if rows_arm {
13611 if std::env::var("MEMRA_MOE_OVERLAP").as_deref() == Ok("1") {
13614 use std::sync::atomic::{AtomicU64, Ordering};
13615 static PAIRS: AtomicU64 = AtomicU64::new(0);
13616 static UNIQ: AtomicU64 = AtomicU64::new(0);
13617 static CALLS: AtomicU64 = AtomicU64::new(0);
13618 let sel_h = e.dtoh_i32(&sel_d)?;
13619 let mut u: Vec<i32> = sel_h.clone();
13620 u.sort_unstable();
13621 u.dedup();
13622 PAIRS.fetch_add(sel_h.len() as u64, Ordering::Relaxed);
13623 UNIQ.fetch_add(u.len() as u64, Ordering::Relaxed);
13624 let c = CALLS.fetch_add(1, Ordering::Relaxed) + 1;
13625 if c.is_multiple_of(480) {
13626 let p = PAIRS.load(Ordering::Relaxed);
13627 let q = UNIQ.load(Ordering::Relaxed);
13628 eprintln!(
13629 "[overlap] calls={c} pairs={p} unique={q} ratio={:.3} (t={t})",
13630 q as f64 / p as f64
13631 );
13632 }
13633 }
13634 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
13635 let act = e.moe_gate_up_silu8_dev_q8_rows(
13636 &dev.ptr_row,
13637 &sel_d,
13638 &zq,
13639 &zd,
13640 t,
13641 n_embd,
13642 n_ff_exp,
13643 n_used,
13644 n_expert,
13645 m.gate_exps.qtype,
13646 m.up_exps.qtype,
13647 rbg_d,
13648 rbu_d,
13649 &m.dev_macros,
13650 )?;
13651 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
13652 e.moe_down8_fma_dev_q8_rows(
13653 &dev.ptr_row,
13654 &sel_d,
13655 &w_d,
13656 &aq2,
13657 &ad2,
13658 &mut moe_out,
13659 t,
13660 n_ff_exp,
13661 n_embd,
13662 n_used,
13663 n_expert,
13664 m.down_exps.qtype,
13665 m.down_exps.row_bytes,
13666 )?;
13667 } else {
13668 for tok in 0..t {
13669 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
13670 let selt = sel_d.slice(tok * n_used..(tok + 1) * n_used);
13671 let wt = w_d.slice(tok * n_used..(tok + 1) * n_used);
13672 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
13673 if q8 {
13674 let (zq, zd) = match (t, zq8) {
13675 (1, Some((q, d))) => (q.clone(), d.clone()),
13676 _ => e.quantize_q8_1_view(&zt, 1, n_embd)?,
13677 };
13678 let act = e.moe_gate_up_silu8_dev_q8(
13679 &dev.ptr_row,
13680 &selt,
13681 &zq,
13682 &zd,
13683 n_embd,
13684 n_ff_exp,
13685 n_used,
13686 n_expert,
13687 m.gate_exps.qtype,
13688 m.up_exps.qtype,
13689 rbg_d,
13690 rbu_d,
13691 &m.dev_macros,
13692 )?;
13693 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
13694 e.moe_down8_fma_dev_q8(
13695 &dev.ptr_row,
13696 &selt,
13697 &wt,
13698 &aq2,
13699 &ad2,
13700 &mut dst,
13701 n_ff_exp,
13702 n_embd,
13703 n_used,
13704 n_expert,
13705 m.down_exps.qtype,
13706 m.down_exps.row_bytes,
13707 )?;
13708 } else {
13709 let act = e.moe_gate_up_silu8_dev(
13710 &dev.ptr_row,
13711 &selt,
13712 &zt,
13713 n_embd,
13714 n_ff_exp,
13715 n_used,
13716 n_expert,
13717 m.gate_exps.qtype,
13718 m.up_exps.qtype,
13719 rbg_d,
13720 rbu_d,
13721 &m.dev_macros,
13722 )?;
13723 e.moe_down8_fma_dev(
13724 &dev.ptr_row,
13725 &selt,
13726 &wt,
13727 &act,
13728 &mut dst,
13729 n_ff_exp,
13730 n_embd,
13731 n_used,
13732 n_expert,
13733 m.down_exps.qtype,
13734 m.down_exps.row_bytes,
13735 )?;
13736 }
13737 }
13738 }
13739 } else {
13740 let q8 = moe_q8_enabled_for_model(cfg, m);
13747 e.with_moe_cache(max_block, |c, eng| {
13748 let row = c
13749 .layer_dev_row(il, n_expert, eng)?
13750 .ok_or("moe_ffn_dev: layer row vanished under the lock")?;
13751 for tok in 0..t {
13752 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
13753 let selt = sel_d.slice(tok * n_used..(tok + 1) * n_used);
13754 let wt = w_d.slice(tok * n_used..(tok + 1) * n_used);
13755 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
13756 if q8 {
13757 let (zq, zd) = match (t, zq8) {
13758 (1, Some((q, d))) => (q.clone(), d.clone()),
13759 _ => eng.quantize_q8_1_view(&zt, 1, n_embd)?,
13760 };
13761 let act = eng.moe_gate_up_silu8_dev_q8(
13762 row,
13763 &selt,
13764 &zq,
13765 &zd,
13766 n_embd,
13767 n_ff_exp,
13768 n_used,
13769 n_expert,
13770 m.gate_exps.qtype,
13771 m.up_exps.qtype,
13772 m.gate_exps.row_bytes,
13773 m.up_exps.row_bytes,
13774 &m.dev_macros,
13775 )?;
13776 let (aq2, ad2) = eng.quantize_q8_1(&act, n_used, n_ff_exp)?;
13777 eng.moe_down8_fma_dev_q8(
13778 row,
13779 &selt,
13780 &wt,
13781 &aq2,
13782 &ad2,
13783 &mut dst,
13784 n_ff_exp,
13785 n_embd,
13786 n_used,
13787 n_expert,
13788 m.down_exps.qtype,
13789 m.down_exps.row_bytes,
13790 )?;
13791 } else {
13792 let act = eng.moe_gate_up_silu8_dev(
13793 row,
13794 &selt,
13795 &zt,
13796 n_embd,
13797 n_ff_exp,
13798 n_used,
13799 n_expert,
13800 m.gate_exps.qtype,
13801 m.up_exps.qtype,
13802 m.gate_exps.row_bytes,
13803 m.up_exps.row_bytes,
13804 &m.dev_macros,
13805 )?;
13806 eng.moe_down8_fma_dev(
13807 row,
13808 &selt,
13809 &wt,
13810 &act,
13811 &mut dst,
13812 n_ff_exp,
13813 n_embd,
13814 n_used,
13815 n_expert,
13816 m.down_exps.qtype,
13817 m.down_exps.row_bytes,
13818 )?;
13819 }
13820 }
13821 c.hits += (t * 3 * n_used) as u64;
13823 Ok(())
13824 })?;
13825 }
13826
13827 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
13832 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
13833 {
13834 let n_ff_sh = gate_shexp.out_features();
13835 let verify_t = t > 1 && t < PRIME_MIN_T;
13838 let (sg_gate, sg_up) = if t == 1 {
13839 shexp_gate_up_t1(e, gate_shexp, up_shexp, z, zq8)?
13840 } else if verify_t {
13841 let mut fused = None;
13845 if crate::spec::spec_fused_t()
13846 && (2..=4).contains(&t)
13847 && e.uses_q8_1_fast(gate_shexp)
13848 && e.uses_q8_1_fast(up_shexp)
13849 {
13850 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
13851 fused = e.matmul_q8_fused2_t(gate_shexp, up_shexp, &zq, &zd, t)?;
13852 }
13853 match fused {
13854 Some(pair) => pair,
13855 None => (
13856 e.matmul_decode_exact(gate_shexp, z, t)?,
13857 e.matmul_decode_exact(up_shexp, z, t)?,
13858 ),
13859 }
13860 } else {
13861 (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?)
13862 };
13863 let mut sa = e.uninit(t * n_ff_sh)?; e.silu_mul(&sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
13865 let sh = if verify_t {
13866 e.matmul_decode_exact(down_shexp, &sa, t)?
13867 } else {
13868 e.matmul(down_shexp, &sa, t)?
13869 };
13870 let g = match &m.gate_inp_shexp {
13874 Some(gate_inp_shexp) => {
13875 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
13878 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
13879 } else {
13880 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
13881 let mut g = e.uninit(t)?;
13882 e.sigmoid(&gs, &mut g, t)?;
13883 g
13884 }
13885 }
13886 None => e.htod(&vec![1.0f32; t])?,
13887 };
13888 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
13889 }
13890
13891 Ok(moe_out)
13892 }
13893
13894 #[allow(clippy::too_many_arguments)]
13904 #[allow(clippy::too_many_arguments)]
13907 fn moe_gdec_token_q8(
13908 e: &Engine,
13909 m: &MoeWeights,
13910 il: u16,
13911 max_block: usize,
13912 zq: &CudaSlice<i8>,
13913 zd: &CudaSlice<f32>,
13914 sel: &[u32],
13915 w: &[f32],
13916 moe_out: &mut CudaSlice<f32>,
13917 tok: usize,
13918 n_embd: usize,
13919 n_ff_exp: usize,
13920 n_used: usize,
13921 ) -> Result<bool, Box<dyn std::error::Error>> {
13922 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
13923 use cudarc::driver::DevicePtr;
13924 let ptrs = e.with_moe_cache(max_block, |c, eng| {
13925 let mut g = [0u64; 8];
13926 let mut u = [0u64; 8];
13927 let mut d = [0u64; 8];
13928 for (j, &ex) in sel.iter().enumerate() {
13929 let ex = ex as u16;
13930 let (Some(sg), Some(su), Some(sd)) = (
13931 c.resident(BlockId::new(il, PROJ_GATE, ex)),
13932 c.resident(BlockId::new(il, PROJ_UP, ex)),
13933 c.resident(BlockId::new(il, PROJ_DOWN, ex)),
13934 ) else {
13935 return Ok(None);
13936 };
13937 let __s = eng.stream();
13938 let (pg, _e0) = c.slot(sg).device_ptr(&__s);
13939 let (pu, _e1) = c.slot(su).device_ptr(&__s);
13940 let (pd, _e2) = c.slot(sd).device_ptr(&__s);
13941 g[j] = pg;
13942 u[j] = pu;
13943 d[j] = pd;
13944 }
13945 if cpu_expert_profile_admit_enabled() && !c.is_frozen() {
13946 for &ex in sel {
13947 let ex = ex as u16;
13948 for proj in [PROJ_GATE, PROJ_UP, PROJ_DOWN] {
13949 c.note_profile_hit(BlockId::new(il, proj, ex));
13950 }
13951 }
13952 }
13953 c.hits += (3 * n_used) as u64;
13954 Ok(Some((g, u, d)))
13955 })?;
13956 let Some((g, u, d)) = ptrs else {
13957 return Ok(false);
13958 };
13959 let mut wv = [0f32; 8];
13960 wv[..n_used].copy_from_slice(w);
13961 let act = e.moe_gate_up_silu8_q8(
13962 crate::WPtr8(g),
13963 crate::WPtr8(u),
13964 zq,
13965 zd,
13966 n_embd,
13967 n_ff_exp,
13968 n_used,
13969 m.gate_exps.qtype,
13970 m.up_exps.qtype,
13971 m.gate_exps.row_bytes,
13972 m.up_exps.row_bytes,
13973 )?;
13974 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
13976 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
13977 e.moe_down8_fma_q8(
13978 crate::WPtr8(d),
13979 crate::F32x8(wv),
13980 &aq2,
13981 &ad2,
13982 &mut dst,
13983 n_ff_exp,
13984 n_embd,
13985 n_used,
13986 m.down_exps.qtype,
13987 m.down_exps.row_bytes,
13988 )?;
13989 Ok(true)
13990 }
13991
13992 #[allow(clippy::too_many_arguments)]
14022 fn moe_fused_epi_token_q8(
14023 e: &Engine,
14024 m: &MoeWeights,
14025 il: u16,
14026 max_block: usize,
14027 zq: &CudaSlice<i8>,
14028 zd: &CudaSlice<f32>,
14029 sel: &[u32],
14030 w: &[f32],
14031 moe_out: &mut CudaSlice<f32>,
14032 tok: usize,
14033 n_embd: usize,
14034 n_ff_exp: usize,
14035 n_used: usize,
14036 limit: f32,
14037 ) -> Result<bool, Box<dyn std::error::Error>> {
14038 use crate::moe_cache::{BlockId, DispatchSlot, PROJ_DOWN, PROJ_GATE, PROJ_UP};
14039 use cudarc::driver::DevicePtr;
14040 debug_assert!(
14041 limit > 1e-6,
14042 "the fused epilogue's kernel collapses every gate to silu(0) at limit 0"
14043 );
14044 debug_assert_eq!(sel.len(), n_used);
14045 debug_assert_eq!(w.len(), n_used);
14046
14047 let ptrs = e.with_moe_cache(max_block, |c, eng| {
14048 if c.n_slots() < 3 * n_used {
14051 return Ok(None);
14052 }
14053 for &ex in sel.iter() {
14056 let ex_usize = ex as usize;
14057 for (proj, exps) in [
14058 (PROJ_GATE, &m.gate_exps),
14059 (PROJ_UP, &m.up_exps),
14060 (PROJ_DOWN, &m.down_exps),
14061 ] {
14062 let id = BlockId::new(il, proj, ex as u16);
14063 let DispatchSlot::Resident(_) =
14064 c.dispatch_source(id, exps.expert_source(ex_usize), eng)?;
14065 }
14066 }
14067 let mut g = [0u64; 8];
14069 let mut u = [0u64; 8];
14070 let mut d = [0u64; 8];
14071 for (j, &ex) in sel.iter().enumerate() {
14072 let ex = ex as u16;
14073 let (Some(sg), Some(su), Some(sd)) = (
14074 c.resident(BlockId::new(il, PROJ_GATE, ex)),
14075 c.resident(BlockId::new(il, PROJ_UP, ex)),
14076 c.resident(BlockId::new(il, PROJ_DOWN, ex)),
14077 ) else {
14078 return Ok(None);
14079 };
14080 let __s = eng.stream();
14081 let (pg, _e0) = c.slot(sg).device_ptr(&__s);
14082 let (pu, _e1) = c.slot(su).device_ptr(&__s);
14083 let (pd, _e2) = c.slot(sd).device_ptr(&__s);
14084 g[j] = pg;
14085 u[j] = pu;
14086 d[j] = pd;
14087 }
14088 Ok(Some((g, u, d)))
14089 })?;
14090 let Some((g, u, d)) = ptrs else {
14091 return Ok(false);
14092 };
14093
14094 Self::moe_fused_epi_launch(
14095 e, m, zq, zd, sel, w, g, u, d, moe_out, tok, n_embd, n_ff_exp, n_used, limit,
14096 )?;
14097 Ok(true)
14098 }
14099
14100 #[allow(clippy::too_many_arguments)]
14111 fn moe_fused_epi_launch(
14112 e: &Engine,
14113 m: &MoeWeights,
14114 zq: &CudaSlice<i8>,
14115 zd: &CudaSlice<f32>,
14116 sel: &[u32],
14117 w: &[f32],
14118 g: [u64; 8],
14119 u: [u64; 8],
14120 d: [u64; 8],
14121 moe_out: &mut CudaSlice<f32>,
14122 tok: usize,
14123 n_embd: usize,
14124 n_ff_exp: usize,
14125 n_used: usize,
14126 limit: f32,
14127 ) -> Result<(), Box<dyn std::error::Error>> {
14128 let mut gs = [0f32; 8];
14129 let mut us = [0f32; 8];
14130 let mut wv = [0f32; 8];
14131 for (j, &ex) in sel.iter().enumerate() {
14132 let ex = ex as usize;
14133 gs[j] = m.gate_exps.macro_scale(ex);
14134 us[j] = m.up_exps.macro_scale(ex);
14135 wv[j] = w[j] * m.down_exps.macro_scale(ex);
14136 }
14137 let act = e.moe_gate_up_preclamp8_q8(
14138 crate::WPtr8(g),
14139 crate::WPtr8(u),
14140 zq,
14141 zd,
14142 crate::F32x8(gs),
14143 crate::F32x8(us),
14144 limit,
14145 n_embd,
14146 n_ff_exp,
14147 n_used,
14148 m.gate_exps.qtype,
14149 m.up_exps.qtype,
14150 m.gate_exps.row_bytes,
14151 m.up_exps.row_bytes,
14152 )?;
14153 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
14155 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
14156 e.moe_down8_fma_q8(
14157 crate::WPtr8(d),
14158 crate::F32x8(wv),
14159 &aq2,
14160 &ad2,
14161 &mut dst,
14162 n_ff_exp,
14163 n_embd,
14164 n_used,
14165 m.down_exps.qtype,
14166 m.down_exps.row_bytes,
14167 )?;
14168 crate::MOE_FUSED_EPI_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
14169 Ok(())
14170 }
14171
14172 #[allow(clippy::too_many_arguments)]
14179 fn moe_vrows_pairs_q8(
14181 e: &Engine,
14182 m: &MoeWeights,
14183 z: &CudaSlice<f32>,
14184 sel: VrowsSel<'_>,
14185 il: u16,
14186 (pg, pu, pd): (u64, u64, u64),
14187 t: usize,
14188 n_embd: usize,
14189 n_ff_exp: usize,
14190 n_used: usize,
14191 limit: f32,
14192 moe_out: &mut CudaSlice<f32>,
14193 ) -> Result<(), Box<dyn std::error::Error>> {
14194 let n_pairs = t * n_used;
14195 let order_on = crate::moe_vrows_dedup_order_on() && !crate::moe_vrows_pack_on();
14203 let n_planes = if order_on { 4 } else { 3 };
14204 let mut ptrs_d = e.vws_uninit_u64(n_planes * n_pairs)?;
14210 let mut scl_d = e.vws_uninit(3 * n_pairs)?;
14211 match sel {
14217 VrowsSel::Host(sel_all, w_all) => {
14218 debug_assert_eq!(sel_all.len(), n_pairs);
14219 debug_assert_eq!(w_all.len(), n_pairs);
14220 let mut ptrs = vec![0u64; n_planes * n_pairs];
14221 let mut scl = vec![0f32; 3 * n_pairs];
14222 for (p, (&ex, &w)) in sel_all.iter().zip(w_all).enumerate() {
14223 let ex = ex as usize;
14224 ptrs[p] = pg + (ex * m.gate_exps.expert_stride) as u64;
14225 ptrs[n_pairs + p] = pu + (ex * m.up_exps.expert_stride) as u64;
14226 ptrs[2 * n_pairs + p] = pd + (ex * m.down_exps.expert_stride) as u64;
14227 scl[p] = m.gate_exps.macro_scale(ex);
14228 scl[n_pairs + p] = m.up_exps.macro_scale(ex);
14229 scl[2 * n_pairs + p] = w * m.down_exps.macro_scale(ex);
14232 }
14233 if order_on {
14234 ptrs[3 * n_pairs..].copy_from_slice(&crate::vrows_expert_major_order(sel_all));
14236 let (visits, distinct) = crate::vrows_overlap_counts(sel_all);
14239 crate::MOE_VROWS_SLAB_READS_AVOIDED
14240 .fetch_add(visits - distinct, std::sync::atomic::Ordering::Relaxed);
14241 }
14242 e.htod_u64_into(&ptrs, &mut ptrs_d)?;
14243 e.htod_f32_into(&scl, &mut scl_d)?;
14244 if crate::moe_vrows_dedup_stat_on() {
14249 let (visits, distinct) = crate::vrows_overlap_counts(sel_all);
14250 debug_assert_eq!(visits, n_pairs as u64);
14251 crate::MOE_VROWS_PAIR_VISITS
14252 .fetch_add(visits, std::sync::atomic::Ordering::Relaxed);
14253 crate::MOE_VROWS_PAIR_DISTINCT
14254 .fetch_add(distinct, std::sync::atomic::Ordering::Relaxed);
14255 crate::moe_vrows_dedup_report();
14256 }
14257 }
14258 VrowsSel::Dev(sel_d, selw_d) => {
14259 let macros = match (
14260 m.gate_exps.macros.as_deref(),
14261 m.up_exps.macros.as_deref(),
14262 m.down_exps.macros.as_deref(),
14263 ) {
14264 (Some(g), Some(u), Some(d)) => Some((g, u, d)),
14265 (None, None, None) => None,
14269 _ => {
14270 return Err("vrows device tables: expert macro planes are not uniform \
14271 across gate/up/down"
14272 .into());
14273 }
14274 };
14275 e.moe_vrows_tables_from_sel(
14276 sel_d,
14277 selw_d,
14278 il,
14279 macros,
14280 (pg, pu, pd),
14281 (
14282 m.gate_exps.expert_stride,
14283 m.up_exps.expert_stride,
14284 m.down_exps.expert_stride,
14285 ),
14286 n_pairs,
14287 &mut ptrs_d,
14288 &mut scl_d,
14289 )?;
14290 if order_on {
14291 e.moe_vrows_order_from_sel(sel_d, n_pairs, &mut ptrs_d)?;
14295 }
14296 if crate::MOE_VROWS_DEV_TABLES_DISPATCHES
14297 .fetch_add(1, std::sync::atomic::Ordering::Relaxed)
14298 == 0
14299 {
14300 eprintln!(
14301 "[moe-vrows-dev-tables] engaged: pointer/scale tables built on device \
14302 from the router's own sel/w; the per-layer pinned readback and its \
14303 cuStreamSynchronize are skipped (MEMRA_MOE_VROWS_DEV_TABLES=1)"
14304 );
14305 }
14306 }
14307 }
14308 let (mut zq, mut zd) = (
14311 e.vws_uninit_i8(t * n_embd)?,
14312 e.vws_uninit(t * (n_embd / 32))?,
14313 );
14314 e.quantize_q8_1_into(z, t, n_embd, &mut zq, &mut zd)?;
14315 let act = e.moe_gate_up_preclamp8_q8_rows(
14316 &ptrs_d,
14317 &scl_d,
14318 &zq,
14319 &zd,
14320 limit,
14321 n_embd,
14322 n_ff_exp,
14323 n_used,
14324 n_pairs,
14325 m.gate_exps.qtype,
14326 m.up_exps.qtype,
14327 m.gate_exps.row_bytes,
14328 m.up_exps.row_bytes,
14329 )?;
14330 let (mut aq2, mut ad2) = (
14332 e.vws_uninit_i8(n_pairs * n_ff_exp)?,
14333 e.vws_uninit(n_pairs * (n_ff_exp / 32))?,
14334 );
14335 e.quantize_q8_1_into(&act, n_pairs, n_ff_exp, &mut aq2, &mut ad2)?;
14336 e.moe_down8_fma_q8_rows(
14337 &ptrs_d,
14338 &scl_d,
14339 &aq2,
14340 &ad2,
14341 moe_out,
14342 n_ff_exp,
14343 n_embd,
14344 n_used,
14345 n_pairs,
14346 m.down_exps.qtype,
14347 m.down_exps.row_bytes,
14348 )?;
14349 e.vws_recycle_u64(ptrs_d);
14352 e.vws_recycle(scl_d);
14353 e.vws_recycle_i8(zq);
14354 e.vws_recycle(zd);
14355 e.vws_recycle(act);
14356 e.vws_recycle_i8(aq2);
14357 e.vws_recycle(ad2);
14358 if crate::MOE_VROWS_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed) == 0 {
14359 eprintln!(
14360 "[glm5-vrows] verify MoE batched across rows: pairs={n_pairs} (t={t} x \
14361 {n_used}), one gate/up+preclamp launch + one down/FMA launch per layer-call \
14362 (rides MEMRA_GLM5_VERIFY_BATCH)"
14363 );
14364 }
14365 Ok(())
14366 }
14367
14368 fn moe_ffn_grouped_prefill_sigmoid(
14406 e: &Engine,
14407 m: &MoeWeights,
14408 z: &CudaSlice<f32>,
14409 t: usize,
14410 cfg: &ModelConfig,
14411 il: u16,
14412 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
14413 let Some(dev) = m
14418 .dev_exps
14419 .as_ref()
14420 .filter(|d| moe_slab_enabled() && d.dev == e.ctx().ordinal())
14421 else {
14422 return Ok(None);
14423 };
14424 if crate::moe_f16g_mode() == 0 {
14425 return Ok(None);
14426 }
14427 if std::env::var("MEMRA_MOE_GATE").is_ok() {
14431 return Ok(None);
14432 }
14433 let moe = cfg
14434 .moe
14435 .as_ref()
14436 .ok_or("grouped sigmoid prefill requires MoE model metadata")?;
14437 let n_embd = cfg.n_embd as usize;
14438 let n_expert = moe.expert_count as usize;
14439 let n_used = moe.expert_used_count as usize;
14440 let n_ff_exp = moe.expert_ff_length as usize;
14441 if !(f16g_proj_ok(m.gate_exps.qtype, n_embd)
14442 && f16g_proj_ok(m.up_exps.qtype, n_embd)
14443 && f16g_proj_ok(m.down_exps.qtype, n_ff_exp))
14444 {
14445 return Ok(None);
14446 }
14447 if n_expert > 512 || n_used == 0 || n_used > 8 {
14450 return Ok(None);
14451 }
14452 let sigmoid = cfg
14453 .sigmoid_router()
14454 .ok_or("grouped sigmoid prefill requires the sigmoid router")?;
14455 let lim_exp = cfg.clamp_exp_at(il as u32);
14459 if matches!(lim_exp, Some(SwigluClamp::Post(_))) {
14460 return Err(
14461 "grouped sigmoid prefill is qualified for the PRE-clamped SwiGLU form only; \
14462 a POST-clamp layer must ride the sequential arm"
14463 .into(),
14464 );
14465 }
14466
14467 let logits = Self::moe_router_logits(e, m, z, t, cfg)?;
14472 Self::trace_sigmoid_router_logits(e, il, t, n_expert, n_used, &logits, m, sigmoid)?;
14473 let (sel_all, w_all) =
14474 Self::moe_route_sigmoid_cfg(e, &logits, t, n_expert, n_used, m, sigmoid)?;
14475 crate::moesd::record_host_routes(il, n_expert, n_used, &sel_all)?;
14476 Self::trace_moe_routes(il, t, &sel_all, &w_all)?;
14477 Self::trace_moe_input(e, il, t, n_embd, z)?;
14478
14479 let mprof = std::env::var("MEMRA_PRIME_PROF").as_deref() == Ok("1");
14480 let mut mt = std::time::Instant::now();
14481 let mut phase = |on: bool| -> f64 {
14482 if on {
14483 let _ = e.stream().synchronize();
14484 let v = mt.elapsed().as_secs_f64() * 1e3;
14485 mt = std::time::Instant::now();
14486 v
14487 } else {
14488 0.0
14489 }
14490 };
14491 let d_router = phase(mprof);
14492
14493 let n_pairs = t * n_used;
14495 if sel_all.len() < n_pairs || w_all.len() < n_pairs || z.len() < t * n_embd {
14496 return Err("grouped sigmoid prefill geometry".into());
14497 }
14498 let mut buckets: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
14499 for (p, &s_id) in sel_all.iter().take(n_pairs).enumerate() {
14500 let s_id = s_id as usize;
14501 if s_id >= n_expert {
14502 return Err(format!("grouped prefill selection {s_id} >= {n_expert}").into());
14503 }
14504 buckets[s_id].push(p as i32);
14505 }
14506 let mut ex_ids: Vec<i32> = Vec::new();
14507 let mut ex_off: Vec<i32> = vec![0];
14508 let mut ex_pairs: Vec<i32> = Vec::with_capacity(n_pairs);
14509 for (e_id, b) in buckets.iter().enumerate() {
14510 if !b.is_empty() {
14511 ex_ids.push(e_id as i32);
14512 ex_pairs.extend_from_slice(b);
14513 ex_off.push(ex_pairs.len() as i32);
14514 }
14515 }
14516 let n_active = ex_ids.len();
14517 if n_active == 0 || n_active > 512 {
14518 return Err(format!("grouped prefill n_active {n_active} outside 1..=512").into());
14519 }
14520 let csr_tok: Vec<i32> = ex_pairs.iter().map(|&p| p / n_used as i32).collect();
14521
14522 let wd: Vec<f32> = (0..n_pairs)
14526 .map(|p| w_all[p] * m.down_exps.macro_scale(sel_all[p] as usize))
14527 .collect();
14528
14529 let (rbg_d, rbu_d) = if dev.gu_il {
14531 let sxx = m.gate_exps.row_bytes + m.up_exps.row_bytes;
14532 (sxx, sxx)
14533 } else {
14534 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
14535 };
14536
14537 let exi = e.htod_i32(&ex_ids)?;
14538 let exo = e.htod_i32(&ex_off)?;
14539 let exp_d = e.htod_i32(&ex_pairs)?;
14540 let csr_tok_d = e.htod_i32(&csr_tok)?;
14541 let pw = e.htod(&wd)?;
14542
14543 let (z16, zs) = e.moe_f16g_act(z, Some(&csr_tok_d), n_embd, n_pairs)?;
14545 let mut g = e.moe_f16_grouped(
14546 &dev.ptr_row,
14547 0,
14548 n_expert,
14549 &exi,
14550 &ex_off,
14551 &exo,
14552 &z16,
14553 &zs,
14554 n_embd,
14555 n_ff_exp,
14556 n_active,
14557 n_pairs,
14558 m.gate_exps.qtype,
14559 rbg_d,
14560 )?;
14561 if m.gate_exps.macros.is_some() {
14562 let mg: Vec<f32> = ex_pairs
14563 .iter()
14564 .map(|&p| m.gate_exps.macro_scale(sel_all[p as usize] as usize))
14565 .collect();
14566 let mg_d = e.htod(&mg)?;
14567 e.scale_rows(&mut g, &mg_d, n_ff_exp, n_pairs)?;
14568 }
14569 let mut u = e.moe_f16_grouped(
14570 &dev.ptr_row,
14571 1,
14572 n_expert,
14573 &exi,
14574 &ex_off,
14575 &exo,
14576 &z16,
14577 &zs,
14578 n_embd,
14579 n_ff_exp,
14580 n_active,
14581 n_pairs,
14582 m.up_exps.qtype,
14583 rbu_d,
14584 )?;
14585 if m.up_exps.macros.is_some() {
14586 let mu: Vec<f32> = ex_pairs
14587 .iter()
14588 .map(|&p| m.up_exps.macro_scale(sel_all[p as usize] as usize))
14589 .collect();
14590 let mu_d = e.htod(&mu)?;
14591 e.scale_rows(&mut u, &mu_d, n_ff_exp, n_pairs)?;
14592 }
14593
14594 let act = match lim_exp {
14597 Some(SwigluClamp::Pre(limit)) => {
14598 let mut a = e.uninit(n_pairs * n_ff_exp)?;
14599 e.swiglu_preclamped_mul_scaled(
14602 &g,
14603 &u,
14604 1.0,
14605 1.0,
14606 limit,
14607 &mut a,
14608 n_pairs * n_ff_exp,
14609 )?;
14610 a
14611 }
14612 None => e.moe_pairs_silu_mul(&g, &u, n_pairs * n_ff_exp)?,
14613 Some(SwigluClamp::Post(_)) => unreachable!("refused before any launch"),
14614 };
14615 let d_gemm_gu = phase(mprof);
14616
14617 let (a16, a_s) = e.moe_f16g_act(&act, None, n_ff_exp, n_pairs)?;
14619 let d_csr = e.moe_f16_grouped(
14620 &dev.ptr_row,
14621 2,
14622 n_expert,
14623 &exi,
14624 &ex_off,
14625 &exo,
14626 &a16,
14627 &a_s,
14628 n_ff_exp,
14629 n_embd,
14630 n_active,
14631 n_pairs,
14632 m.down_exps.qtype,
14633 m.down_exps.row_bytes,
14634 )?;
14635 let y_pair = e.rows_permute(&d_csr, &exp_d, n_pairs, n_embd)?;
14636 let toff: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
14637 let tids: Vec<i32> = (0..n_pairs as i32).collect();
14638 let toff_d = e.htod_i32(&toff)?;
14639 let tids_d = e.htod_i32(&tids)?;
14640 let mut moe_out = e.uninit(t * n_embd)?;
14643 e.moe_pairs_scatter(&y_pair, &pw, &toff_d, &tids_d, &mut moe_out, t, n_embd)?;
14644 let d_down = phase(mprof);
14645
14646 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut moe_out)?;
14649 if mprof {
14650 let d_shared = phase(true);
14651 eprintln!(
14652 "[moe-grouped-prefill-prof] il={il} t={t} router={d_router:.1}ms \
14653 gemm_gu={d_gemm_gu:.1}ms down_scatter={d_down:.1}ms shared={d_shared:.1}ms"
14654 );
14655 }
14656
14657 crate::MOE_GROUPED_PREFILL_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
14658 static GPF_LOGGED: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
14661 let layer_bit = 1u64 << (il as u64 % 64);
14662 if GPF_LOGGED.fetch_or(layer_bit, std::sync::atomic::Ordering::Relaxed) & layer_bit == 0 {
14663 eprintln!(
14664 "[moe-grouped-prefill] execute layer={il} tokens={t} n_active={n_active} \
14665 provenance=resident-slab router=sigmoid-host-oracle epilogue=pre-clamped \
14666 macro_fold=gate-up-rows+down-weight performance_claim=false \
14667 (logged once per layer)"
14668 );
14669 }
14670 Ok(Some(moe_out))
14671 }
14672
14673 #[allow(clippy::too_many_arguments)] fn moe_gdec_token(
14675 e: &Engine,
14676 m: &MoeWeights,
14677 il: u16,
14678 max_block: usize,
14679 zt: &cudarc::driver::CudaView<f32>,
14680 sel: &[u32],
14681 w: &[f32],
14682 moe_out: &mut CudaSlice<f32>,
14683 tok: usize,
14684 n_embd: usize,
14685 n_ff_exp: usize,
14686 n_used: usize,
14687 ) -> Result<bool, Box<dyn std::error::Error>> {
14688 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
14689 use cudarc::driver::DevicePtr;
14690 let ptrs = e.with_moe_cache(max_block, |c, eng| {
14692 let mut g = [0u64; 8];
14693 let mut u = [0u64; 8];
14694 let mut d = [0u64; 8];
14695 for (j, &ex) in sel.iter().enumerate() {
14696 let ex = ex as u16;
14697 let (Some(sg), Some(su), Some(sd)) = (
14698 c.resident(BlockId::new(il, PROJ_GATE, ex)),
14699 c.resident(BlockId::new(il, PROJ_UP, ex)),
14700 c.resident(BlockId::new(il, PROJ_DOWN, ex)),
14701 ) else {
14702 return Ok(None);
14703 };
14704 let __s = eng.stream();
14705 let (pg, _e0) = c.slot(sg).device_ptr(&__s);
14706 let (pu, _e1) = c.slot(su).device_ptr(&__s);
14707 let (pd, _e2) = c.slot(sd).device_ptr(&__s);
14708 g[j] = pg;
14709 u[j] = pu;
14710 d[j] = pd;
14711 }
14712 if cpu_expert_profile_admit_enabled() && !c.is_frozen() {
14713 for &ex in sel {
14714 let ex = ex as u16;
14715 for proj in [PROJ_GATE, PROJ_UP, PROJ_DOWN] {
14716 c.note_profile_hit(BlockId::new(il, proj, ex));
14717 }
14718 }
14719 }
14720 c.hits += (3 * n_used) as u64; Ok(Some((g, u, d)))
14722 })?;
14723 let Some((g, u, d)) = ptrs else {
14724 return Ok(false);
14725 };
14726 let mut wv = [0f32; 8];
14727 wv[..n_used].copy_from_slice(w);
14728 let act = e.moe_gate_up_silu8(
14730 crate::WPtr8(g),
14731 crate::WPtr8(u),
14732 zt,
14733 n_embd,
14734 n_ff_exp,
14735 n_used,
14736 m.gate_exps.qtype,
14737 m.up_exps.qtype,
14738 m.gate_exps.row_bytes,
14739 m.up_exps.row_bytes,
14740 )?;
14741 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
14742 e.moe_down8_fma_into(
14743 crate::WPtr8(d),
14744 crate::F32x8(wv),
14745 &act,
14746 &mut dst,
14747 n_ff_exp,
14748 n_embd,
14749 n_used,
14750 m.down_exps.qtype,
14751 m.down_exps.row_bytes,
14752 )?;
14753 Ok(true)
14754 }
14755
14756 #[allow(clippy::too_many_arguments)] fn moe_cached_gemm_q8(
14762 e: &Engine,
14763 il: u16,
14764 proj: u8,
14765 ex: usize,
14766 m: &MoeWeights,
14767 max_block: usize,
14768 aq: &CudaSlice<i8>,
14769 ad: &CudaSlice<f32>,
14770 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
14771 use crate::moe_cache::{BlockId, DispatchSlot, PROJ_GATE, PROJ_UP};
14772 let exps = match proj {
14773 PROJ_GATE => &m.gate_exps,
14774 PROJ_UP => &m.up_exps,
14775 _ => &m.down_exps,
14776 };
14777 let layout = exps.expert_layout(ex);
14778 let id = BlockId::new(il, proj, ex as u16);
14779 let source = exps.expert_source(ex);
14780 e.with_moe_cache(max_block, |c, eng| {
14781 let slot = c.dispatch_source(id, source, eng)?;
14782 let DispatchSlot::Resident(sl) = slot;
14783 let buf = c.slot(sl);
14784 eng.qmatvec_expert_q8(
14785 buf,
14786 0..layout.len,
14787 aq,
14788 ad,
14789 1,
14790 exps.in_f,
14791 exps.out_f,
14792 layout.qtype,
14793 layout.row_bytes,
14794 )
14795 })
14796 }
14797
14798 fn moe_cached_gemm(
14799 e: &Engine,
14800 il: u16,
14801 proj: u8,
14802 ex: usize,
14803 m: &MoeWeights,
14804 max_block: usize,
14805 x: &cudarc::driver::CudaView<f32>,
14806 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
14807 use crate::moe_cache::{BlockId, DispatchSlot, PROJ_GATE, PROJ_UP};
14808 let exps = match proj {
14809 PROJ_GATE => &m.gate_exps,
14810 PROJ_UP => &m.up_exps,
14811 _ => &m.down_exps,
14812 };
14813 let layout = exps.expert_layout(ex);
14814 let id = BlockId::new(il, proj, ex as u16);
14815 let source = exps.expert_source(ex);
14816 e.with_moe_cache(max_block, |c, eng| {
14818 let slot = c.dispatch_source(id, source, eng)?;
14819 let DispatchSlot::Resident(sl) = slot;
14822 let buf = c.slot(sl);
14823 m.qmatvec_view(
14824 eng,
14825 buf,
14826 0..layout.len,
14827 x,
14828 1,
14829 exps.in_f,
14830 exps.out_f,
14831 layout.qtype,
14832 layout.row_bytes,
14833 )
14834 })
14835 }
14836
14837 fn moe_profile_admit_expert(
14841 e: &Engine,
14842 il: u16,
14843 ex: usize,
14844 m: &MoeWeights,
14845 max_block: usize,
14846 ) -> Result<(), Box<dyn std::error::Error>> {
14847 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
14848 e.with_moe_cache(max_block, |cache, eng| {
14849 for (proj, exps) in [
14850 (PROJ_GATE, &m.gate_exps),
14851 (PROJ_UP, &m.up_exps),
14852 (PROJ_DOWN, &m.down_exps),
14853 ] {
14854 let id = BlockId::new(il, proj, ex as u16);
14855 let _ = cache.dispatch_source(id, exps.expert_source(ex), eng)?;
14856 }
14857 Ok(())
14858 })
14859 }
14860
14861 #[allow(clippy::too_many_arguments)]
14864 fn moe_frozen_gemm(
14865 e: &Engine,
14866 il: u16,
14867 proj: u8,
14868 ex: usize,
14869 m: &MoeWeights,
14870 max_block: usize,
14871 x: &cudarc::driver::CudaView<f32>,
14872 scratch: &mut Option<CudaSlice<u8>>,
14873 scratch_len: usize,
14874 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
14875 use crate::moe_cache::{BlockId, PROJ_GATE, PROJ_UP};
14876 let exps = match proj {
14877 PROJ_GATE => &m.gate_exps,
14878 PROJ_UP => &m.up_exps,
14879 _ => &m.down_exps,
14880 };
14881 let layout = exps.expert_layout(ex);
14882 let id = BlockId::new(il, proj, ex as u16);
14883 if let Some(output) = e.with_moe_cache(max_block, |cache, eng| {
14884 let Some(slot) = cache.resident(id) else {
14885 return Ok(None);
14886 };
14887 let buf = cache.slot(slot);
14888 Ok(Some(m.qmatvec_view(
14889 eng,
14890 buf,
14891 0..layout.len,
14892 x,
14893 1,
14894 exps.in_f,
14895 exps.out_f,
14896 layout.qtype,
14897 layout.row_bytes,
14898 )?))
14899 })? {
14900 return Ok(output);
14901 }
14902 if scratch.is_none() {
14903 *scratch = Some(e.alloc_u8_uninit(scratch_len)?);
14904 }
14905 let scratch = scratch.as_mut().unwrap();
14906 e.stage_expert(exps.expert_bytes(ex), scratch, 0)?;
14907 m.qmatvec_view(
14908 e,
14909 scratch,
14910 0..layout.len,
14911 x,
14912 1,
14913 exps.in_f,
14914 exps.out_f,
14915 layout.qtype,
14916 layout.row_bytes,
14917 )
14918 }
14919
14920 fn moe_prefetch_expert(
14921 e: &Engine,
14922 il: u16,
14923 ex: usize,
14924 m: &MoeWeights,
14925 max_block: usize,
14926 keep: &[crate::moe_cache::BlockId],
14927 ) -> Result<(), Box<dyn std::error::Error>> {
14928 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
14929 e.with_moe_cache(max_block, |c, eng| {
14930 for (proj, exps) in [
14931 (PROJ_GATE, &m.gate_exps),
14932 (PROJ_UP, &m.up_exps),
14933 (PROJ_DOWN, &m.down_exps),
14934 ] {
14935 let id = BlockId::new(il, proj, ex as u16);
14936 let _ = c.prefetch_source(id, exps.expert_source(ex), keep, eng)?;
14937 }
14938 Ok(())
14939 })
14940 }
14941
14942 fn moe_prefetch_disk_expert(
14945 e: &Engine,
14946 il: u16,
14947 ex: usize,
14948 m: &MoeWeights,
14949 max_block: usize,
14950 keep: &[crate::moe_cache::BlockId],
14951 ) -> Result<(), Box<dyn std::error::Error>> {
14952 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
14953 e.with_moe_cache(max_block, |c, eng| {
14954 for (proj, exps) in [
14955 (PROJ_GATE, &m.gate_exps),
14956 (PROJ_UP, &m.up_exps),
14957 (PROJ_DOWN, &m.down_exps),
14958 ] {
14959 let source = exps.expert_source(ex);
14960 if let crate::model::ExpertSource::Disk { .. } = &source {
14961 let id = BlockId::new(il, proj, ex as u16);
14962 let _ = c.prefetch_source(id, source, keep, eng)?;
14963 }
14964 }
14965 Ok(())
14966 })
14967 }
14968
14969 #[inline]
14970 fn moe_prefetch_host_expert(ex: usize, m: &MoeWeights) {
14971 let _ = m.gate_exps.prefetch_expert_pages(ex);
14972 let _ = m.up_exps.prefetch_expert_pages(ex);
14973 let _ = m.down_exps.prefetch_expert_pages(ex);
14974 }
14975}
14976
14977impl HybridModel {
14994 #[allow(clippy::too_many_arguments)]
14998 fn moe_ffn_grouped_resident_q8(
14999 e: &Engine,
15000 m: &MoeWeights,
15001 z: &CudaSlice<f32>,
15002 t: usize,
15003 cfg: &ModelConfig,
15004 il: u16,
15005 sel_all: &[u32],
15006 w_all: &[f32],
15007 table: &CudaSlice<u64>,
15008 gu_il: bool,
15009 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
15010 let moe = cfg.moe.as_ref().unwrap();
15011 let n_embd = cfg.n_embd as usize;
15012 let n_expert = moe.expert_count as usize;
15013 let n_used = moe.expert_used_count as usize;
15014 let n_ff_exp = moe.expert_ff_length as usize;
15015 let n_pairs = t * n_used;
15016 debug_assert_eq!(sel_all.len(), n_pairs);
15017 debug_assert_eq!(w_all.len(), n_pairs);
15018 debug_assert!(
15019 m.gate_exps.macros.is_none()
15020 && m.up_exps.macros.is_none()
15021 && m.down_exps.macros.is_none(),
15022 "resident grouped q8 does not fold per-expert macro scales",
15023 );
15024
15025 if !cfg.swiglu_clamped_at(il as u32) {
15031 let sel: Vec<i32> = sel_all.iter().map(|&expert| expert as i32).collect();
15032 let sel_d = e.htod_i32(&sel)?;
15033 let w_d = e.htod(w_all)?;
15034 let (gate_row_bytes, up_row_bytes) = if gu_il {
15035 let combined = m.gate_exps.row_bytes + m.up_exps.row_bytes;
15036 (combined, combined)
15037 } else {
15038 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
15039 };
15040 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
15041 let act = e.moe_gate_up_silu8_dev_q8_rows(
15042 table,
15043 &sel_d,
15044 &zq,
15045 &zd,
15046 t,
15047 n_embd,
15048 n_ff_exp,
15049 n_used,
15050 n_expert,
15051 m.gate_exps.qtype,
15052 m.up_exps.qtype,
15053 gate_row_bytes,
15054 up_row_bytes,
15055 &m.dev_macros,
15056 )?;
15057 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
15058 let mut moe_out = e.uninit(t * n_embd)?;
15059 e.moe_down8_fma_dev_q8_rows_g(
15060 table,
15061 &sel_d,
15062 &w_d,
15063 &aq2,
15064 &ad2,
15065 &mut moe_out,
15066 t,
15067 n_ff_exp,
15068 n_embd,
15069 n_used,
15070 n_expert,
15071 m.down_exps.qtype,
15072 m.down_exps.row_bytes,
15073 )?;
15074
15075 if std::env::var("MEMRA_MOE_STATS").is_ok() {
15076 let mut counts = vec![0usize; n_expert];
15077 for &expert in sel_all {
15078 counts[expert as usize] += 1;
15079 }
15080 let mut sizes: Vec<usize> =
15081 counts.into_iter().filter(|&count| count != 0).collect();
15082 sizes.sort_unstable();
15083 let mean = sizes.iter().sum::<usize>() as f64 / sizes.len().max(1) as f64;
15084 println!(
15085 "moe-grouped il={il} t={t} dispatch=resident-q8-rows active={}/{} \
15086 m_e: min={} median={} mean={mean:.1} max={}",
15087 sizes.len(),
15088 n_expert,
15089 sizes.first().copied().unwrap_or(0),
15090 sizes.get(sizes.len() / 2).copied().unwrap_or(0),
15091 sizes.last().copied().unwrap_or(0),
15092 );
15093 }
15094 return Ok(moe_out);
15095 }
15096
15097 let pair_tok: Vec<i32> = (0..n_pairs).map(|pair| (pair / n_used) as i32).collect();
15101 let pair_ex: Vec<i32> = sel_all.iter().map(|&expert| expert as i32).collect();
15102 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
15103 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
15104
15105 let mut by_expert: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
15106 for (pair, &expert) in pair_ex.iter().enumerate() {
15107 by_expert[expert as usize].push(pair as i32);
15108 }
15109
15110 let pair_tok_d = e.htod_i32(&pair_tok)?;
15111 let pair_ex_d = e.htod_i32(&pair_ex)?;
15112 let pair_w_d = e.htod(w_all)?;
15113 let tok_off_d = e.htod_i32(&tok_off)?;
15114 let tok_ids_d = e.htod_i32(&tok_ids)?;
15115
15116 let matvec = |proj: i32,
15117 pair_rows: &CudaSlice<i32>,
15118 aq: &CudaSlice<i8>,
15119 ad: &CudaSlice<f32>,
15120 in_f: usize,
15121 out_f: usize,
15122 qtype: i32,
15123 row_bytes: usize|
15124 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
15125 e.moe_pairs_matvec_q8(
15126 table, proj, pair_rows, &pair_ex_d, aq, ad, in_f, out_f, n_expert, n_pairs, qtype,
15127 row_bytes,
15128 )
15129 };
15130
15131 let (gate_row_bytes, up_row_bytes) = if gu_il {
15132 let combined = m.gate_exps.row_bytes + m.up_exps.row_bytes;
15133 (combined, combined)
15134 } else {
15135 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
15136 };
15137 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
15138 let gate = matvec(
15139 0,
15140 &pair_tok_d,
15141 &zq,
15142 &zd,
15143 n_embd,
15144 n_ff_exp,
15145 m.gate_exps.qtype,
15146 gate_row_bytes,
15147 )?;
15148 let up = matvec(
15149 1,
15150 &pair_tok_d,
15151 &zq,
15152 &zd,
15153 n_embd,
15154 n_ff_exp,
15155 m.up_exps.qtype,
15156 up_row_bytes,
15157 )?;
15158 let mut act = e.uninit(n_pairs * n_ff_exp)?;
15159 Self::ffn_act_lim(
15160 e,
15161 cfg,
15162 &gate,
15163 &up,
15164 1.0,
15165 1.0,
15166 cfg.clamp_exp_at(il as u32),
15167 &mut act,
15168 n_pairs * n_ff_exp,
15169 )?;
15170 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
15171 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
15172 let pair_self_d = e.htod_i32(&pair_self)?;
15173 let down = matvec(
15174 2,
15175 &pair_self_d,
15176 &aq2,
15177 &ad2,
15178 n_ff_exp,
15179 n_embd,
15180 m.down_exps.qtype,
15181 m.down_exps.row_bytes,
15182 )?;
15183 let mut moe_out = e.uninit(t * n_embd)?;
15184 e.moe_pairs_scatter(
15185 &down,
15186 &pair_w_d,
15187 &tok_off_d,
15188 &tok_ids_d,
15189 &mut moe_out,
15190 t,
15191 n_embd,
15192 )?;
15193
15194 if std::env::var("MEMRA_MOE_STATS").is_ok() {
15195 let mut sizes: Vec<usize> = by_expert
15196 .iter()
15197 .filter_map(|pairs| (!pairs.is_empty()).then_some(pairs.len()))
15198 .collect();
15199 sizes.sort_unstable();
15200 let mean = sizes.iter().sum::<usize>() as f64 / sizes.len().max(1) as f64;
15201 println!(
15202 "moe-grouped il={il} t={t} dispatch=resident-q8-clamped-pairs active={}/{} \
15203 m_e: min={} median={} mean={mean:.1} max={}",
15204 sizes.len(),
15205 n_expert,
15206 sizes.first().copied().unwrap_or(0),
15207 sizes.get(sizes.len() / 2).copied().unwrap_or(0),
15208 sizes.last().copied().unwrap_or(0),
15209 );
15210 }
15211 Ok(moe_out)
15212 }
15213
15214 #[allow(clippy::too_many_arguments)]
15220 #[allow(clippy::map_entry)] fn shexp_split_matvec(
15222 e: &Engine,
15223 rank1: &Engine,
15224 wg: &CudaSlice<u8>,
15225 wu: &CudaSlice<u8>,
15226 wd: &CudaSlice<u8>,
15227 z: &CudaSlice<f32>,
15228 lim: Option<SwigluClamp>,
15229 cfg: &ModelConfig,
15230 il: u16,
15231 n_embd: usize,
15232 n_ff_sh: usize,
15233 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
15234 use cudarc::driver::DevicePtr;
15235 if !n_ff_sh.is_multiple_of(2) || !n_embd.is_multiple_of(2) {
15236 return Ok(None);
15237 }
15238 let hf = n_ff_sh / 2;
15239 let nd = n_embd / 2;
15240 struct Rep {
15241 wg1: CudaSlice<u8>,
15242 wu1: CudaSlice<u8>,
15243 wd1: CudaSlice<u8>,
15244 }
15245 struct SplitWs {
15246 pin_dev: usize,
15247 gate0: CudaSlice<f32>,
15249 up0: CudaSlice<f32>,
15250 act: CudaSlice<f32>,
15251 sh_buf: CudaSlice<f32>,
15252 ev_z: cudarc::driver::CudaEvent,
15253 ev_act0: cudarc::driver::CudaEvent,
15254 z1: CudaSlice<f32>,
15256 g1: CudaSlice<f32>,
15257 u1: CudaSlice<f32>,
15258 a1h: CudaSlice<f32>,
15259 act1: CudaSlice<f32>,
15260 y1: CudaSlice<f32>,
15261 ev_act1: cudarc::driver::CudaEvent,
15262 ev_y1: cudarc::driver::CudaEvent,
15263 raw_act_e: u64,
15264 raw_sh_e: u64,
15265 raw_z1: u64,
15266 raw_a1h: u64,
15267 raw_act1: u64,
15268 raw_y1: u64,
15269 }
15270 static WS: std::sync::Mutex<Option<SplitWs>> = std::sync::Mutex::new(None);
15271 static REPS: std::sync::Mutex<Option<std::collections::HashMap<u64, Rep>>> =
15272 std::sync::Mutex::new(None);
15273 let mut guard = WS.lock().map_err(|_| "shexp split lock is poisoned")?;
15274 let mut reps_guard = REPS.lock().map_err(|_| "shexp reps lock is poisoned")?;
15275 let reps = reps_guard.get_or_insert_with(std::collections::HashMap::new);
15276 let pins = e.ctx().ordinal();
15277 if guard.as_ref().is_none_or(|w| w.pin_dev != pins) {
15278 let (gate0, up0, act, sh_buf, ev_z, ev_act0) = {
15279 let _m = e.gpu.enter_main()?;
15280 (
15281 e.htod(&vec![0.0f32; hf])?,
15282 e.htod(&vec![0.0f32; hf])?,
15283 e.htod(&vec![0.0f32; n_ff_sh])?,
15284 e.htod(&vec![0.0f32; n_embd])?,
15285 e.ctx().new_event(None)?,
15286 e.ctx().new_event(None)?,
15287 )
15288 };
15289 let (z1, g1, u1, a1h, act1, y1, ev_act1, ev_y1) = {
15290 let _r = rank1.gpu.enter_main()?;
15291 (
15292 rank1.htod(&vec![0.0f32; n_embd])?,
15293 rank1.htod(&vec![0.0f32; hf])?,
15294 rank1.htod(&vec![0.0f32; hf])?,
15295 rank1.htod(&vec![0.0f32; hf])?,
15296 rank1.htod(&vec![0.0f32; n_ff_sh])?,
15297 rank1.htod(&vec![0.0f32; nd])?,
15298 rank1.ctx().new_event(None)?,
15299 rank1.ctx().new_event(None)?,
15300 )
15301 };
15302 let (raw_act_e, raw_sh_e) = {
15303 let _m = e.gpu.enter_main()?;
15304 let stream = e.stream();
15305 let (a, _g0) = act.device_ptr(&stream);
15306 let (b, _g1) = sh_buf.device_ptr(&stream);
15307 (a, b)
15308 };
15309 let (raw_z1, raw_a1h, raw_act1, raw_y1) = {
15310 let _r = rank1.gpu.enter_main()?;
15311 let rs = rank1.stream();
15312 let (a, _g0) = z1.device_ptr(&rs);
15313 let (b, _g1) = a1h.device_ptr(&rs);
15314 let (c, _g2) = act1.device_ptr(&rs);
15315 let (d, _g3) = y1.device_ptr(&rs);
15316 (a, b, c, d)
15317 };
15318 *guard = Some(SplitWs {
15319 pin_dev: pins,
15320 gate0,
15321 up0,
15322 act,
15323 sh_buf,
15324 ev_z,
15325 ev_act0,
15326 z1,
15327 g1,
15328 u1,
15329 a1h,
15330 act1,
15331 y1,
15332 ev_act1,
15333 ev_y1,
15334 raw_act_e,
15335 raw_sh_e,
15336 raw_z1,
15337 raw_a1h,
15338 raw_act1,
15339 raw_y1,
15340 });
15341 }
15342 let ws = guard.as_mut().expect("armed above");
15343 let wg_pin = {
15344 let _m = e.gpu.enter_main()?;
15345 let stream = e.stream();
15346 let (p, _g) = wg.device_ptr(&stream);
15347 p
15348 };
15349 if !reps.contains_key(&wg_pin) {
15350 let up = |src: &CudaSlice<u8>,
15352 off_bytes: usize,
15353 len: usize|
15354 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
15355 use cudarc::driver::sys;
15356 let sptr = {
15357 let _m = e.gpu.enter_main()?;
15358 let stream = e.stream();
15359 let (p, _g) = src.device_ptr(&stream);
15360 p + off_bytes as u64
15361 };
15362 let dst = {
15363 let _r = rank1.gpu.enter_main()?;
15364 rank1.alloc_u8_uninit(len)?
15365 };
15366 let dptr = {
15367 let _r = rank1.gpu.enter_main()?;
15368 let rs = rank1.stream();
15369 let (p, _g) = dst.device_ptr(&rs);
15370 p
15371 };
15372 let _r = rank1.gpu.enter_main()?;
15373 let r = unsafe {
15374 sys::cuMemcpyAsync(
15375 dptr as sys::CUdeviceptr,
15376 sptr as sys::CUdeviceptr,
15377 len,
15378 rank1.stream().cu_stream() as sys::CUstream,
15379 )
15380 };
15381 if r != sys::CUresult::CUDA_SUCCESS {
15382 return Err(format!("shexp split replica upload: {r:?}").into());
15383 }
15384 rank1.stream().synchronize()?;
15385 Ok(dst)
15386 };
15387 let wg1 = up(wg, hf * n_embd * 2, hf * n_embd * 2)?;
15388 let wu1 = up(wu, hf * n_embd * 2, hf * n_embd * 2)?;
15389 let wd1 = up(wd, nd * n_ff_sh * 2, nd * n_ff_sh * 2)?;
15390 reps.insert(wg_pin, Rep { wg1, wu1, wd1 });
15391 }
15392 let _ = il;
15393 let raw_z = {
15395 let _m = e.gpu.enter_main()?;
15396 let stream = e.stream();
15397 let (p, _g) = z.device_ptr(&stream);
15398 ws.ev_z.record(&stream)?;
15399 p
15400 };
15401 {
15403 let rep = reps.get(&wg_pin).expect("uploaded above");
15404 let _r = rank1.gpu.enter_main()?;
15405 rank1.stream().wait(&ws.ev_z)?;
15406 crate::tp::raw_copy_bytes(ws.raw_z1, raw_z, n_embd * 4, rank1)?;
15407 let SplitWs {
15408 z1, g1, u1, a1h, ..
15409 } = &mut *ws;
15410 rank1.matvec_bf16_dual_into(&rep.wg1, &rep.wu1, z1, g1, u1, n_embd, hf)?;
15411 Self::ffn_act_lim(rank1, cfg, g1, u1, 1.0, 1.0, lim, a1h, hf)?;
15412 crate::tp::raw_copy_bytes(ws.raw_act1 + (hf * 4) as u64, ws.raw_a1h, hf * 4, rank1)?;
15414 crate::tp::raw_copy_bytes(ws.raw_act_e + (hf * 4) as u64, ws.raw_a1h, hf * 4, rank1)?;
15415 ws.ev_act1.record(&rank1.stream())?;
15416 }
15417 {
15419 let _m = e.gpu.enter_main()?;
15420 let SplitWs {
15421 gate0, up0, act, ..
15422 } = &mut *ws;
15423 let wg_lo = wg.slice(0..hf * n_embd * 2);
15424 let wu_lo = wu.slice(0..hf * n_embd * 2);
15425 e.matvec_bf16_dual_view_into(&wg_lo, &wu_lo, z, gate0, up0, n_embd, hf)?;
15426 Self::ffn_act_lim(e, cfg, gate0, up0, 1.0, 1.0, lim, act, hf)?;
15427 ws.ev_act0.record(&e.stream())?;
15428 }
15429 {
15431 let rep = reps.get(&wg_pin).expect("uploaded above");
15432 let _r = rank1.gpu.enter_main()?;
15433 rank1.stream().wait(&ws.ev_act0)?;
15434 crate::tp::raw_copy_bytes(ws.raw_act1, ws.raw_act_e, hf * 4, rank1)?;
15435 let SplitWs { act1, y1, .. } = &mut *ws;
15436 rank1.matvec_bf16_into(&rep.wd1, act1, y1, n_ff_sh, nd)?;
15437 crate::tp::raw_copy_bytes(ws.raw_sh_e + (nd * 4) as u64, ws.raw_y1, nd * 4, rank1)?;
15438 ws.ev_y1.record(&rank1.stream())?;
15439 }
15440 {
15442 let _m = e.gpu.enter_main()?;
15443 e.stream().wait(&ws.ev_act1)?;
15444 let SplitWs { act, sh_buf, .. } = &mut *ws;
15445 let wd_lo = wd.slice(0..nd * n_ff_sh * 2);
15446 e.matvec_bf16_view_into(&wd_lo, act, sh_buf, n_ff_sh, nd)?;
15447 e.stream().wait(&ws.ev_y1)?;
15448 let mut sh = e.uninit(n_embd)?;
15449 {
15450 let mut dst = sh.slice_mut(0..n_embd);
15451 e.stream()
15452 .memcpy_dtod(&ws.sh_buf.slice(0..n_embd), &mut dst)?;
15453 }
15454 Ok(Some(sh))
15455 }
15456 }
15457
15458 fn shexp_overlap_issue(
15465 e: &Engine,
15466 m: &MoeWeights,
15467 z: &CudaSlice<f32>,
15468 cfg: &ModelConfig,
15469 il: u16,
15470 n_embd: usize,
15471 ) -> Result<bool, Box<dyn std::error::Error>> {
15472 if cfg.m3.is_some() || m.gate_inp_shexp.is_some() {
15473 return Ok(false);
15474 }
15475 let (
15476 Some(crate::model::GpuTensor::FloatBf16 { data: wg, .. }),
15477 Some(crate::model::GpuTensor::FloatBf16 { data: wu, .. }),
15478 Some(crate::model::GpuTensor::FloatBf16 { data: wd, .. }),
15479 ) = (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
15480 else {
15481 return Ok(false);
15482 };
15483 let n_ff_sh = m
15484 .gate_shexp
15485 .as_ref()
15486 .expect("matched Some above")
15487 .out_features();
15488 let Ok(lim) = Self::fused_post_limit(cfg.clamp_shexp_at(il as u32)) else {
15491 return Ok(false);
15492 };
15493 let mut guard = SHEXP_OV_WS
15494 .lock()
15495 .map_err(|_| "shexp overlap workspace lock is poisoned")?;
15496 let pins = (e.ctx().ordinal(), n_embd, n_ff_sh);
15497 if guard
15498 .as_ref()
15499 .is_none_or(|(d, ne, nf, ..)| (*d, *ne, *nf) != pins)
15500 {
15501 *guard = Some((
15502 pins.0,
15503 pins.1,
15504 pins.2,
15505 e.uninit(n_ff_sh)?,
15506 e.uninit(n_embd)?,
15507 ));
15508 }
15509 let (_, _, _, act, sh) = guard.as_mut().expect("armed above");
15510 e.matvec_bf16_dual_silu_into(wg, wu, z, act, n_embd, n_ff_sh, lim)?;
15511 e.matvec_bf16_into(wd, act, sh, n_ff_sh, n_embd)?;
15512 drop(guard);
15513 Ok(true)
15514 }
15515
15516 #[allow(clippy::too_many_arguments)]
15522 #[allow(clippy::map_entry)] fn shexp_dev1_issue(
15524 e: &Engine,
15525 rank1: &Engine,
15526 m: &MoeWeights,
15527 z: &CudaSlice<f32>,
15528 cfg: &ModelConfig,
15529 il: u16,
15530 n_embd: usize,
15531 ) -> Result<bool, Box<dyn std::error::Error>> {
15532 use cudarc::driver::DevicePtr;
15533 if cfg.m3.is_some() || m.gate_inp_shexp.is_some() {
15534 return Ok(false);
15535 }
15536 let (
15537 Some(crate::model::GpuTensor::FloatBf16 { data: wg, .. }),
15538 Some(crate::model::GpuTensor::FloatBf16 { data: wu, .. }),
15539 Some(crate::model::GpuTensor::FloatBf16 { data: wd, .. }),
15540 ) = (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
15541 else {
15542 return Ok(false);
15543 };
15544 let n_ff_sh = m
15545 .gate_shexp
15546 .as_ref()
15547 .expect("matched Some above")
15548 .out_features();
15549 let Ok(lim) = Self::fused_post_limit(cfg.clamp_shexp_at(il as u32)) else {
15552 return Ok(false);
15553 };
15554 let mut ws_guard = SHEXP_D1_WS
15556 .lock()
15557 .map_err(|_| "shexp dev1 workspace lock is poisoned")?;
15558 if ws_guard
15559 .as_ref()
15560 .is_none_or(|(k, ..)| *k != (n_embd, n_ff_sh))
15561 {
15562 let (act1, z1, ev_done) = {
15563 let _r1 = rank1.gpu.enter_main()?;
15564 (
15565 rank1.htod(&vec![0.0f32; n_ff_sh])?,
15566 rank1.htod(&vec![0.0f32; n_embd])?,
15567 rank1.ctx().new_event(None)?,
15568 )
15569 };
15570 let (sh_root, ev_z) = {
15571 let _main = e.gpu.enter_main()?;
15572 (e.htod(&vec![0.0f32; n_embd])?, e.ctx().new_event(None)?)
15573 };
15574 *ws_guard = Some(((n_embd, n_ff_sh), act1, z1, sh_root, ev_z, ev_done));
15575 }
15576 let mut reps_guard = SHEXP_D1_REPS
15578 .lock()
15579 .map_err(|_| "shexp dev1 replica lock is poisoned")?;
15580 let reps = reps_guard.get_or_insert_with(Default::default);
15581 if !reps.contains_key(&il) {
15582 let (wg1, wu1, wd1) = {
15583 let _r1 = rank1.gpu.enter_main()?;
15584 (
15585 rank1.alloc_u8_uninit(n_ff_sh * n_embd * 2)?,
15586 rank1.alloc_u8_uninit(n_ff_sh * n_embd * 2)?,
15587 rank1.alloc_u8_uninit(n_embd * n_ff_sh * 2)?,
15588 )
15589 };
15590 for (src, dst) in [(wg, &wg1), (wu, &wu1), (wd, &wd1)] {
15591 let s_ptr = {
15592 let _main = e.gpu.enter_main()?;
15593 let stream = e.stream();
15594 let (p, _g) = src.device_ptr(&stream);
15595 p
15596 };
15597 let d_ptr = {
15598 let _r1 = rank1.gpu.enter_main()?;
15599 let stream = rank1.stream();
15600 let (p, _g) = dst.device_ptr(&stream);
15601 p
15602 };
15603 let _r1 = rank1.gpu.enter_main()?;
15604 crate::tp::raw_copy_bytes(d_ptr, s_ptr, src.len(), rank1)?;
15605 }
15606 {
15607 let _r1 = rank1.gpu.enter_main()?;
15608 rank1.stream().synchronize()?;
15609 }
15610 reps.insert(il, (wg1, wu1, wd1));
15611 }
15612 let (wg1, wu1, wd1) = reps.get(&il).expect("armed above");
15613 let (_, act1, z1, sh_root, ev_z, ev_done) = ws_guard.as_mut().expect("armed above");
15614 let (raw_z, raw_sh) = {
15617 let _main = e.gpu.enter_main()?;
15618 let stream = e.stream();
15619 let (a, _g0) = z.device_ptr(&stream);
15620 let (b, _g1) = sh_root.device_ptr(&stream);
15621 ev_z.record(&stream)?;
15622 (a, b)
15623 };
15624 {
15625 let _r1 = rank1.gpu.enter_main()?;
15626 rank1.stream().wait(ev_z)?;
15627 let raw_z1 = {
15628 let stream = rank1.stream();
15629 let (p, _g) = z1.device_ptr(&stream);
15630 p
15631 };
15632 crate::tp::raw_copy_bytes(raw_z1, raw_z, n_embd * 4, rank1)?;
15633 rank1.matvec_bf16_dual_silu_into(wg1, wu1, z1, act1, n_embd, n_ff_sh, lim)?;
15634 rank1.matvec_bf16_raw_out(wd1, act1, raw_sh, n_ff_sh, n_embd)?;
15638 ev_done.record(&rank1.stream())?;
15639 }
15640 Ok(true)
15641 }
15642
15643 fn shexp_dev1_apply(
15645 e: &Engine,
15646 output: &mut CudaSlice<f32>,
15647 n_embd: usize,
15648 ) -> Result<(), Box<dyn std::error::Error>> {
15649 let guard = SHEXP_D1_WS
15650 .lock()
15651 .map_err(|_| "shexp dev1 workspace lock is poisoned")?;
15652 let (pin, _, _, sh_root, _, ev_done) =
15653 guard.as_ref().ok_or("shexp dev1 apply without issue")?;
15654 if pin.0 != n_embd {
15655 return Err("shexp dev1 width drifted".into());
15656 }
15657 let _main = e.gpu.enter_main()?;
15658 e.stream().wait(ev_done)?;
15659 static ONES_D1: std::sync::Mutex<Option<(usize, CudaSlice<f32>)>> =
15660 std::sync::Mutex::new(None);
15661 let mut og = ONES_D1.lock().map_err(|_| "ones lock is poisoned")?;
15662 if og.as_ref().is_none_or(|(d, _)| *d != e.ctx().ordinal()) {
15663 *og = Some((e.ctx().ordinal(), e.htod(&[1.0f32])?));
15664 }
15665 let ones = &og.as_ref().expect("armed above").1;
15666 e.add_scaled_rows(sh_root, ones, output, n_embd, 1)?;
15667 Ok(())
15668 }
15669
15670 fn shexp_overlap_tail_ptrs(
15674 e: &Engine,
15675 m: &MoeWeights,
15676 cfg: &ModelConfig,
15677 n_embd: usize,
15678 ) -> Result<Option<(u64, u64)>, Box<dyn std::error::Error>> {
15679 use cudarc::driver::DevicePtr;
15680 if cfg.m3.is_some() || m.gate_inp_shexp.is_some() {
15681 return Ok(None);
15682 }
15683 let (
15684 Some(crate::model::GpuTensor::FloatBf16 { .. }),
15685 Some(crate::model::GpuTensor::FloatBf16 { .. }),
15686 Some(crate::model::GpuTensor::FloatBf16 { .. }),
15687 ) = (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
15688 else {
15689 return Ok(None);
15690 };
15691 let n_ff_sh = m
15692 .gate_shexp
15693 .as_ref()
15694 .expect("matched Some above")
15695 .out_features();
15696 let mut guard = SHEXP_OV_WS
15697 .lock()
15698 .map_err(|_| "shexp overlap workspace lock is poisoned")?;
15699 let pins = (e.ctx().ordinal(), n_embd, n_ff_sh);
15700 if guard
15701 .as_ref()
15702 .is_none_or(|(d, ne, nf, ..)| (*d, *ne, *nf) != pins)
15703 {
15704 *guard = Some((
15705 pins.0,
15706 pins.1,
15707 pins.2,
15708 e.uninit(n_ff_sh)?,
15709 e.uninit(n_embd)?,
15710 ));
15711 }
15712 let sh_raw = {
15713 let (_, _, _, _, sh) = guard.as_ref().expect("armed above");
15714 let stream = e.stream();
15715 let (p, _g) = sh.device_ptr(&stream);
15716 p
15717 };
15718 static ONES_T3: std::sync::Mutex<Option<(usize, CudaSlice<f32>)>> =
15719 std::sync::Mutex::new(None);
15720 let mut og = ONES_T3.lock().map_err(|_| "ones lock is poisoned")?;
15721 if og.as_ref().is_none_or(|(d, _)| *d != e.ctx().ordinal()) {
15722 *og = Some((e.ctx().ordinal(), e.htod(&[1.0f32])?));
15723 }
15724 let ones_raw = {
15725 let stream = e.stream();
15726 let (p, _g) = og.as_ref().expect("armed above").1.device_ptr(&stream);
15727 p
15728 };
15729 Ok(Some((sh_raw, ones_raw)))
15730 }
15731
15732 fn shexp_overlap_apply(
15735 e: &Engine,
15736 output: &mut CudaSlice<f32>,
15737 n_embd: usize,
15738 ) -> Result<(), Box<dyn std::error::Error>> {
15739 let guard = SHEXP_OV_WS
15740 .lock()
15741 .map_err(|_| "shexp overlap workspace lock is poisoned")?;
15742 let (_, ne, _, _, sh) = guard.as_ref().ok_or("shexp overlap apply without issue")?;
15743 if *ne != n_embd {
15744 return Err("shexp overlap width drifted".into());
15745 }
15746 static ONES_OV: std::sync::Mutex<Option<(usize, CudaSlice<f32>)>> =
15747 std::sync::Mutex::new(None);
15748 let mut og = ONES_OV.lock().map_err(|_| "ones lock is poisoned")?;
15749 if og.as_ref().is_none_or(|(d, _)| *d != e.ctx().ordinal()) {
15750 *og = Some((e.ctx().ordinal(), e.htod(&[1.0f32])?));
15751 }
15752 let ones = &og.as_ref().expect("armed above").1;
15753 e.add_scaled_rows(sh, ones, output, n_embd, 1)?;
15754 Ok(())
15755 }
15756
15757 fn moe_ffn_grouped_add_shared(
15758 e: &Engine,
15759 m: &MoeWeights,
15760 z: &CudaSlice<f32>,
15761 t: usize,
15762 cfg: &ModelConfig,
15763 il: u16,
15764 moe_out: &mut CudaSlice<f32>,
15765 ) -> Result<(), Box<dyn std::error::Error>> {
15766 static SHEXP_NS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
15769 static SHEXP_CALLS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
15770 let shexp_timing = std::env::var("MEMRA_STEP_TP_TIMING").as_deref() == Ok("1");
15771 let shexp_started = shexp_timing.then(std::time::Instant::now);
15772 let result = Self::moe_ffn_grouped_add_shared_inner(e, m, z, t, cfg, il, moe_out);
15773 if let Some(started) = shexp_started {
15774 use std::sync::atomic::Ordering;
15775 e.stream().synchronize()?;
15776 let ns = SHEXP_NS.fetch_add(started.elapsed().as_nanos() as u64, Ordering::Relaxed)
15777 + started.elapsed().as_nanos() as u64;
15778 let calls = SHEXP_CALLS.fetch_add(1, Ordering::Relaxed) + 1;
15779 if calls.is_multiple_of(430) {
15780 eprintln!(
15781 "[moe-shexp-timing] calls={calls} total_ms={:.1} avg_us={:.1}",
15782 ns as f64 / 1.0e6,
15783 ns as f64 / calls as f64 / 1.0e3,
15784 );
15785 }
15786 }
15787 result
15788 }
15789
15790 #[allow(clippy::too_many_arguments)]
15791 fn moe_ffn_grouped_add_shared_inner(
15792 e: &Engine,
15793 m: &MoeWeights,
15794 z: &CudaSlice<f32>,
15795 t: usize,
15796 cfg: &ModelConfig,
15797 il: u16,
15798 moe_out: &mut CudaSlice<f32>,
15799 ) -> Result<(), Box<dyn std::error::Error>> {
15800 let n_embd = cfg.n_embd as usize;
15801 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
15802 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
15803 {
15804 let n_ff_sh = gate_shexp.out_features();
15805 let lim = cfg.clamp_shexp_at(il as u32);
15806 let fused = t == 1
15813 && lim.is_none()
15814 && cfg.m3.is_none()
15815 && e.uses_q8_1_fast(gate_shexp)
15816 && e.uses_q8_1_fast(up_shexp);
15817 let canonical_w4a16_rows =
15818 t <= 32 && m.step_ep.as_ref().is_some_and(|ep| ep.nvfp4_device_routes);
15819 let bf16_dual = if (t == 1 || canonical_w4a16_rows)
15824 && crate::Engine::bf16_mmv_on()
15825 && n_embd.is_multiple_of(8)
15826 {
15827 match (gate_shexp, up_shexp) {
15828 (
15829 crate::model::GpuTensor::FloatBf16 { data: wg, .. },
15830 crate::model::GpuTensor::FloatBf16 { data: wu, .. },
15831 ) => Some((wg, wu)),
15832 _ => None,
15833 }
15834 } else {
15835 None
15836 };
15837 let sh = if let Some((wg, wu)) = bf16_dual {
15838 type SharedExpertWorkspace = (
15842 usize,
15843 usize,
15844 usize,
15845 CudaSlice<f32>,
15846 CudaSlice<f32>,
15847 CudaSlice<f32>,
15848 CudaSlice<f32>,
15849 );
15850 static SHEXP_WS: std::sync::Mutex<
15851 Option<std::collections::HashMap<usize, SharedExpertWorkspace>>,
15852 > = std::sync::Mutex::new(None);
15853 let down_bf16 = match down_shexp {
15854 crate::model::GpuTensor::FloatBf16 { data, .. } => Some(data),
15855 _ => None,
15856 };
15857 let mut guard = SHEXP_WS
15858 .lock()
15859 .map_err(|_| "shexp workspace lock is poisoned")?;
15860 let capacity = if canonical_w4a16_rows { 32 } else { 1 };
15861 let device = e.ctx().ordinal();
15862 let workspaces = guard.get_or_insert_with(Default::default);
15863 if workspaces
15864 .get(&device)
15865 .is_none_or(|(ne, nf, cap, ..)| (*ne, *nf, *cap) != (n_embd, n_ff_sh, capacity))
15866 {
15867 workspaces.insert(
15868 device,
15869 (
15870 n_embd,
15871 n_ff_sh,
15872 capacity,
15873 e.uninit(capacity * n_ff_sh)?,
15874 e.uninit(capacity * n_ff_sh)?,
15875 e.uninit(capacity * n_ff_sh)?,
15876 e.uninit(capacity * n_embd)?,
15877 ),
15878 );
15879 }
15880 {
15883 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15884 let split_on = *ON
15885 .get_or_init(|| std::env::var("MEMRA_SHEXP_SPLIT").as_deref() == Ok("1"));
15886 if split_on
15887 && t == 1
15888 && let (Some(wd), Some(rank1)) = (
15889 match down_shexp {
15890 crate::model::GpuTensor::FloatBf16 { data, .. } => Some(data),
15891 _ => None,
15892 },
15893 m.step_tp.as_ref().and_then(|st| st.runtime.rank_engine(1)),
15894 )
15895 && let Some(sh) = Self::shexp_split_matvec(
15896 e, rank1, wg, wu, wd, z, lim, cfg, il, n_embd, n_ff_sh,
15897 )?
15898 {
15899 drop(guard);
15900 let gate = match &m.gate_inp_shexp {
15901 Some(gate_inp_shexp) => {
15902 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
15903 }
15904 None => e.htod(&vec![1.0f32; t])?,
15905 };
15906 e.add_scaled_rows(&sh, &gate, moe_out, n_embd, t)?;
15907 return Ok(());
15908 }
15909 }
15910 let (_, _, _, gate, up, act, sh_buf) = workspaces
15911 .get_mut(&device)
15912 .expect("shexp workspace initialized above");
15913 if let (true, Ok(lim_post)) = (cfg.m3.is_none(), Self::fused_post_limit(lim)) {
15916 if canonical_w4a16_rows {
15919 e.matvec_bf16_dual_silu_rows_into(
15920 wg, wu, z, act, n_embd, n_ff_sh, lim_post, t,
15921 )?;
15922 } else {
15923 e.matvec_bf16_dual_silu_into(wg, wu, z, act, n_embd, n_ff_sh, lim_post)?;
15924 }
15925 let _ = (&gate, &up);
15926 } else {
15927 e.matvec_bf16_dual_into(wg, wu, z, gate, up, n_embd, n_ff_sh)?;
15928 Self::ffn_act_lim(e, cfg, gate, up, 1.0, 1.0, lim, act, n_ff_sh)?;
15929 }
15930 if let Some(down) = down_bf16 {
15931 if canonical_w4a16_rows {
15932 e.matvec_bf16_rows_into(down, act, sh_buf, n_ff_sh, n_embd, t)?;
15933 let mut sh = e.uninit(t * n_embd)?;
15934 {
15935 let mut dst = sh.slice_mut(0..t * n_embd);
15936 e.stream()
15937 .memcpy_dtod(&sh_buf.slice(0..t * n_embd), &mut dst)?;
15938 }
15939 sh
15940 } else {
15941 static FUSE_DA: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15948 let fuse_da = *FUSE_DA.get_or_init(|| {
15949 std::env::var("MEMRA_FUSE_DOWN_ADDSCALE").as_deref() != Ok("0")
15950 });
15951 if fuse_da && m.gate_inp_shexp.is_none() {
15952 static ONES1: std::sync::Mutex<Option<(usize, CudaSlice<f32>)>> =
15953 std::sync::Mutex::new(None);
15954 let mut og = ONES1.lock().map_err(|_| "shexp ones lock is poisoned")?;
15955 if og.as_ref().is_none_or(|(d, _)| *d != e.ctx().ordinal()) {
15956 *og = Some((e.ctx().ordinal(), e.htod(&[1.0f32])?));
15957 }
15958 let ones = &og.as_ref().expect("armed above").1;
15959 e.matvec_bf16_down_addscale_into(
15960 down, act, ones, moe_out, n_ff_sh, n_embd,
15961 )?;
15962 return Ok(());
15963 }
15964 e.matvec_bf16_into(down, act, sh_buf, n_ff_sh, n_embd)?;
15965 let sh = e.uninit(n_embd)?;
15966 let mut sh = sh;
15968 {
15969 let mut dst = sh.slice_mut(0..n_embd);
15970 e.stream().memcpy_dtod(&sh_buf.slice(0..n_embd), &mut dst)?;
15971 }
15972 sh
15973 }
15974 } else {
15975 e.matmul(down_shexp, act, t)?
15976 }
15977 } else if fused {
15978 let (zq, zd) = e.quantize_q8_1(z, 1, n_embd)?;
15979 let pair = match e.matmul_pre_dual_noscale(gate_shexp, up_shexp, &zq, &zd, 1)? {
15980 Some((gate, up)) => Some((gate, up)),
15981 None => {
15982 match (
15983 e.matmul_pre_noscale(gate_shexp, &zq, &zd, 1)?,
15984 e.matmul_pre_noscale(up_shexp, &zq, &zd, 1)?,
15985 ) {
15986 (Some(gate), Some(up)) => Some((gate, up)),
15987 _ => None,
15988 }
15989 }
15990 };
15991 match pair {
15992 Some(((gate, gs), (up, us))) => {
15993 if e.uses_q8_1_fast(down_shexp) {
15994 let (aq, ad) = e.silu_mul_scaled_q8_1(&gate, &up, gs, us, n_ff_sh)?;
15995 e.matmul_pre(down_shexp, &aq, &ad, &gate, 1)?
15996 } else {
15997 let mut act = e.uninit(n_ff_sh)?;
15998 e.silu_mul_scaled(&gate, &up, gs, us, &mut act, n_ff_sh)?;
15999 e.matmul(down_shexp, &act, 1)?
16000 }
16001 }
16002 None => {
16003 let gate = e.matmul_pre(gate_shexp, &zq, &zd, z, 1)?;
16004 let up = e.matmul_pre(up_shexp, &zq, &zd, z, 1)?;
16005 let mut act = e.uninit(n_ff_sh)?;
16006 Self::ffn_act(e, cfg, &gate, &up, &mut act, n_ff_sh)?;
16007 e.matmul(down_shexp, &act, 1)?
16008 }
16009 }
16010 } else {
16011 let sg_gate = e.matmul(gate_shexp, z, t)?;
16012 let sg_up = e.matmul(up_shexp, z, t)?;
16013 let mut sa = e.uninit(t * n_ff_sh)?;
16014 Self::ffn_act_lim(
16015 e,
16016 cfg,
16017 &sg_gate,
16018 &sg_up,
16019 1.0,
16020 1.0,
16021 lim,
16022 &mut sa,
16023 t * n_ff_sh,
16024 )?;
16025 e.matmul(down_shexp, &sa, t)?
16026 };
16027 let gate = match &m.gate_inp_shexp {
16028 Some(gate_inp_shexp) => {
16029 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
16030 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
16031 } else {
16032 let raw = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
16033 let mut gate = e.uninit(t)?;
16034 e.sigmoid(&raw, &mut gate, t)?;
16035 gate
16036 }
16037 }
16038 None if t == 1 => {
16043 static ONES: std::sync::Mutex<
16044 Option<std::collections::HashMap<usize, CudaSlice<f32>>>,
16045 > = std::sync::Mutex::new(None);
16046 let mut guard = ONES.lock().map_err(|_| "shexp ones lock is poisoned")?;
16047 let device = e.ctx().ordinal();
16048 let rows = guard.get_or_insert_with(Default::default);
16049 use std::collections::hash_map::Entry;
16052 let ones = match rows.entry(device) {
16053 Entry::Occupied(occupied) => occupied.into_mut(),
16054 Entry::Vacant(vacant) => vacant.insert(e.htod(&[1.0f32])?),
16055 };
16056 e.add_scaled_rows(&sh, ones, moe_out, n_embd, t)?;
16057 return Ok(());
16058 }
16059 None => e.htod(&vec![1.0f32; t])?,
16060 };
16061 e.add_scaled_rows(&sh, &gate, moe_out, n_embd, t)?;
16062 }
16063 Ok(())
16064 }
16065
16066 pub(crate) fn moe_ffn_grouped(
16069 e: &Engine,
16070 m: &MoeWeights,
16071 z: &CudaSlice<f32>,
16072 t: usize,
16073 cfg: &ModelConfig,
16074 il: u16,
16075 max_block: usize,
16076 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
16077 let moe = cfg.moe.as_ref().unwrap();
16078 let n_embd = cfg.n_embd as usize;
16079 let n_expert = moe.expert_count as usize;
16080 let n_used = moe.expert_used_count as usize;
16081 let n_ff_exp = moe.expert_ff_length as usize;
16082 let lim_exp = cfg.clamp_exp_at(il as u32);
16084
16085 let logits = Self::moe_router_logits(e, m, z, t, cfg)?;
16089 if let Some(sig) = cfg.sigmoid_router() {
16090 Self::trace_sigmoid_router_logits(e, il, t, n_expert, n_used, &logits, m, sig)?;
16091 }
16092 let (sel_all, w_all) = if let Some(sig) = cfg.sigmoid_router() {
16093 Self::moe_route_sigmoid_cfg(e, &logits, t, n_expert, n_used, m, sig)?
16094 } else {
16095 Self::moe_route_cfg(e, &logits, t, n_expert, n_used, m.active_experts.as_deref())?
16096 };
16097 crate::moesd::record_host_routes(il, n_expert, n_used, &sel_all)?;
16098 Self::trace_moe_routes(il, t, &sel_all, &w_all)?;
16099 Self::trace_moe_input(e, il, t, n_embd, z)?;
16100
16101 let no_exp_macros = m.gate_exps.macros.is_none()
16106 && m.up_exps.macros.is_none()
16107 && m.down_exps.macros.is_none();
16108 let resident_q8 = m.dev_exps.as_ref().filter(|dev| {
16109 m.has_uniform_expert_layout()
16110 && no_exp_macros
16111 && moe_q8_enabled_for_model(cfg, m)
16112 && moe_slab_enabled()
16113 && dev.dev == e.ctx().ordinal()
16114 });
16115 if let Some(dev) = resident_q8 {
16116 let mut moe_out = Self::moe_ffn_grouped_resident_q8(
16117 e,
16118 m,
16119 z,
16120 t,
16121 cfg,
16122 il,
16123 &sel_all,
16124 &w_all,
16125 &dev.ptr_row,
16126 dev.gu_il,
16127 )?;
16128 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut moe_out)?;
16129 return Ok(moe_out);
16130 }
16131
16132 struct ExpertGroup {
16136 tok_indices: Vec<i32>, slot_indices: Vec<i32>, weights: Vec<f32>, }
16140 let mut groups: Vec<ExpertGroup> = (0..n_expert)
16141 .map(|_| ExpertGroup {
16142 tok_indices: Vec::new(),
16143 slot_indices: Vec::new(),
16144 weights: Vec::new(),
16145 })
16146 .collect();
16147
16148 for tok in 0..t {
16149 for j in 0..n_used {
16150 let ex = sel_all[tok * n_used + j] as usize;
16151 let w = w_all[tok * n_used + j];
16152 groups[ex].tok_indices.push(tok as i32);
16153 groups[ex].slot_indices.push(j as i32);
16154 groups[ex].weights.push(w);
16155 }
16156 }
16157
16158 let mut slot_buf = e.zeros(t * n_used * n_embd)?;
16161 let mut wbuf = e.zeros(t * n_used)?; let g_len = m.gate_exps.max_expert_bytes();
16165 let u_len = m.up_exps.max_expert_bytes();
16166 let d_len = m.down_exps.max_expert_bytes();
16167 let moe_q8 = moe_q8_enabled_for_model(cfg, m);
16168 let slab_local = m
16171 .dev_exps
16172 .as_ref()
16173 .filter(|dev| !dev.gu_il && moe_slab_enabled() && dev.dev == e.ctx().ordinal());
16174 let use_cache =
16175 slab_local.is_none() && Engine::moe_cache_enabled() && !e.moe_cache_frozen();
16176 let grouped_q8 = moe_q8 && (slab_local.is_some() || use_cache);
16179
16180 let (mut scratch_g, mut scratch_u, mut scratch_d) = if slab_local.is_none() && !use_cache {
16182 (
16183 Some(e.alloc_u8(g_len)?),
16184 Some(e.alloc_u8(u_len)?),
16185 Some(e.alloc_u8(d_len)?),
16186 )
16187 } else {
16188 (None, None, None)
16189 };
16190
16191 let mut order: Vec<usize> = (0..n_expert)
16202 .filter(|&ex| !groups[ex].tok_indices.is_empty())
16203 .collect();
16204 order.sort_by(|&a, &b| {
16205 groups[b]
16206 .tok_indices
16207 .len()
16208 .cmp(&groups[a].tok_indices.len())
16209 .then(a.cmp(&b))
16210 });
16211 let mut m_dist: Vec<usize> = Vec::new(); let page_window = moe_page_prefetch_window();
16213 let worker_disk_prefetch = use_cache && crate::spill_pread::worker_enabled();
16214 if worker_disk_prefetch
16215 && let Some(first) = grouped_worker_prefetch_position(order.len(), None)
16216 {
16217 Self::moe_prefetch_disk_expert(e, il, order[first], m, max_block, &[])?;
16218 }
16219 for (order_pos, &ex) in order.iter().enumerate() {
16220 for next in page_prefetch_positions(order_pos, order.len(), page_window) {
16221 Self::moe_prefetch_host_expert(order[next], m);
16222 }
16223 if worker_disk_prefetch
16224 && let Some(next) = grouped_worker_prefetch_position(order.len(), Some(order_pos))
16225 {
16226 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
16227 let keep = [
16228 BlockId::new(il, PROJ_GATE, ex as u16),
16229 BlockId::new(il, PROJ_UP, ex as u16),
16230 BlockId::new(il, PROJ_DOWN, ex as u16),
16231 ];
16232 Self::moe_prefetch_disk_expert(e, il, order[next], m, max_block, &keep)?;
16233 }
16234 let grp = &groups[ex];
16235 let m_e = grp.tok_indices.len();
16236 m_dist.push(m_e);
16237 let gl = m.gate_exps.expert_layout(ex);
16238 let ul = m.up_exps.expert_layout(ex);
16239 let dl = m.down_exps.expert_layout(ex);
16240
16241 let tok_idx_d = e.htod_i32(&grp.tok_indices)?;
16245 let slot_idx_d = e.htod_i32(&grp.slot_indices)?;
16246 let dmac = m.down_exps.macro_scale(ex);
16247 let weight_d = if dmac == 1.0 {
16248 e.htod(&grp.weights)?
16249 } else {
16250 let scaled: Vec<f32> = grp.weights.iter().map(|&w| w * dmac).collect();
16251 e.htod(&scaled)?
16252 };
16253
16254 let mut gathered = e.zeros(m_e * n_embd)?;
16256 e.gather_rows(z, &tok_idx_d, &mut gathered, n_embd, m_e)?;
16257 let gv = gathered.slice(0..m_e * n_embd);
16258
16259 let y = if let Some(dev) = slab_local {
16262 let gate_start = ex * m.gate_exps.expert_stride;
16263 let up_start = ex * m.up_exps.expert_stride;
16264 let down_start = ex * m.down_exps.expert_stride;
16265 if grouped_q8 {
16266 let (zq, zd) = e.quantize_q8_1(&gathered, m_e, n_embd)?;
16267 let gate = e.qmatvec_expert_q8(
16268 &dev.gate,
16269 gate_start..gate_start + gl.len,
16270 &zq,
16271 &zd,
16272 m_e,
16273 m.gate_exps.in_f,
16274 m.gate_exps.out_f,
16275 gl.qtype,
16276 gl.row_bytes,
16277 )?;
16278 let up = e.qmatvec_expert_q8(
16279 &dev.up,
16280 up_start..up_start + ul.len,
16281 &zq,
16282 &zd,
16283 m_e,
16284 m.up_exps.in_f,
16285 m.up_exps.out_f,
16286 ul.qtype,
16287 ul.row_bytes,
16288 )?;
16289 let mut act = e.uninit(m_e * n_ff_exp)?;
16290 Self::ffn_act_lim(
16291 e,
16292 cfg,
16293 &gate,
16294 &up,
16295 m.gate_exps.macro_scale(ex),
16296 m.up_exps.macro_scale(ex),
16297 lim_exp,
16298 &mut act,
16299 m_e * n_ff_exp,
16300 )?;
16301 let (aq2, ad2) = e.quantize_q8_1(&act, m_e, n_ff_exp)?;
16302 e.qmatvec_expert_q8(
16303 &dev.down,
16304 down_start..down_start + dl.len,
16305 &aq2,
16306 &ad2,
16307 m_e,
16308 m.down_exps.in_f,
16309 m.down_exps.out_f,
16310 dl.qtype,
16311 dl.row_bytes,
16312 )?
16313 } else {
16314 let gate = m.qmatvec_view(
16315 e,
16316 &dev.gate,
16317 gate_start..gate_start + gl.len,
16318 &gv,
16319 m_e,
16320 m.gate_exps.in_f,
16321 m.gate_exps.out_f,
16322 gl.qtype,
16323 gl.row_bytes,
16324 )?;
16325 let up = m.qmatvec_view(
16326 e,
16327 &dev.up,
16328 up_start..up_start + ul.len,
16329 &gv,
16330 m_e,
16331 m.up_exps.in_f,
16332 m.up_exps.out_f,
16333 ul.qtype,
16334 ul.row_bytes,
16335 )?;
16336 let mut act = e.uninit(m_e * n_ff_exp)?;
16337 Self::ffn_act_lim(
16338 e,
16339 cfg,
16340 &gate,
16341 &up,
16342 m.gate_exps.macro_scale(ex),
16343 m.up_exps.macro_scale(ex),
16344 lim_exp,
16345 &mut act,
16346 m_e * n_ff_exp,
16347 )?;
16348 let actv = act.slice(0..m_e * n_ff_exp);
16349 m.qmatvec_view(
16350 e,
16351 &dev.down,
16352 down_start..down_start + dl.len,
16353 &actv,
16354 m_e,
16355 m.down_exps.in_f,
16356 m.down_exps.out_f,
16357 dl.qtype,
16358 dl.row_bytes,
16359 )?
16360 }
16361 } else if use_cache {
16362 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
16363 if grouped_q8 {
16364 let (zq, zd) = e.quantize_q8_1(&gathered, m_e, n_embd)?;
16365 let gate = e.with_moe_cache(max_block, |cache, eng| {
16366 let id = BlockId::new(il, PROJ_GATE, ex as u16);
16367 let slot = cache.dispatch_source(id, m.gate_exps.expert_source(ex), eng)?;
16368 eng.qmatvec_expert_q8(
16369 cache.buf(slot),
16370 0..gl.len,
16371 &zq,
16372 &zd,
16373 m_e,
16374 m.gate_exps.in_f,
16375 m.gate_exps.out_f,
16376 gl.qtype,
16377 gl.row_bytes,
16378 )
16379 })?;
16380 let up = e.with_moe_cache(max_block, |cache, eng| {
16381 let id = BlockId::new(il, PROJ_UP, ex as u16);
16382 let slot = cache.dispatch_source(id, m.up_exps.expert_source(ex), eng)?;
16383 eng.qmatvec_expert_q8(
16384 cache.buf(slot),
16385 0..ul.len,
16386 &zq,
16387 &zd,
16388 m_e,
16389 m.up_exps.in_f,
16390 m.up_exps.out_f,
16391 ul.qtype,
16392 ul.row_bytes,
16393 )
16394 })?;
16395 let mut act = e.uninit(m_e * n_ff_exp)?;
16396 Self::ffn_act_lim(
16397 e,
16398 cfg,
16399 &gate,
16400 &up,
16401 m.gate_exps.macro_scale(ex),
16402 m.up_exps.macro_scale(ex),
16403 lim_exp,
16404 &mut act,
16405 m_e * n_ff_exp,
16406 )?;
16407 let (aq2, ad2) = e.quantize_q8_1(&act, m_e, n_ff_exp)?;
16408 e.with_moe_cache(max_block, |cache, eng| {
16409 let id = BlockId::new(il, PROJ_DOWN, ex as u16);
16410 let slot = cache.dispatch_source(id, m.down_exps.expert_source(ex), eng)?;
16411 eng.qmatvec_expert_q8(
16412 cache.buf(slot),
16413 0..dl.len,
16414 &aq2,
16415 &ad2,
16416 m_e,
16417 m.down_exps.in_f,
16418 m.down_exps.out_f,
16419 dl.qtype,
16420 dl.row_bytes,
16421 )
16422 })?
16423 } else {
16424 let gate = e.with_moe_cache(max_block, |cache, eng| {
16425 let id = BlockId::new(il, PROJ_GATE, ex as u16);
16426 let slot = cache.dispatch_source(id, m.gate_exps.expert_source(ex), eng)?;
16427 m.qmatvec_view(
16428 eng,
16429 cache.buf(slot),
16430 0..gl.len,
16431 &gv,
16432 m_e,
16433 m.gate_exps.in_f,
16434 m.gate_exps.out_f,
16435 gl.qtype,
16436 gl.row_bytes,
16437 )
16438 })?;
16439 let up = e.with_moe_cache(max_block, |cache, eng| {
16440 let id = BlockId::new(il, PROJ_UP, ex as u16);
16441 let slot = cache.dispatch_source(id, m.up_exps.expert_source(ex), eng)?;
16442 m.qmatvec_view(
16443 eng,
16444 cache.buf(slot),
16445 0..ul.len,
16446 &gv,
16447 m_e,
16448 m.up_exps.in_f,
16449 m.up_exps.out_f,
16450 ul.qtype,
16451 ul.row_bytes,
16452 )
16453 })?;
16454 let mut act = e.uninit(m_e * n_ff_exp)?;
16455 Self::ffn_act_lim(
16456 e,
16457 cfg,
16458 &gate,
16459 &up,
16460 m.gate_exps.macro_scale(ex),
16461 m.up_exps.macro_scale(ex),
16462 lim_exp,
16463 &mut act,
16464 m_e * n_ff_exp,
16465 )?;
16466 let actv = act.slice(0..m_e * n_ff_exp);
16467 e.with_moe_cache(max_block, |cache, eng| {
16468 let id = BlockId::new(il, PROJ_DOWN, ex as u16);
16469 let slot = cache.dispatch_source(id, m.down_exps.expert_source(ex), eng)?;
16470 m.qmatvec_view(
16471 eng,
16472 cache.buf(slot),
16473 0..dl.len,
16474 &actv,
16475 m_e,
16476 m.down_exps.in_f,
16477 m.down_exps.out_f,
16478 dl.qtype,
16479 dl.row_bytes,
16480 )
16481 })?
16482 }
16483 } else {
16484 let sg = scratch_g.as_mut().unwrap();
16485 let su = scratch_u.as_mut().unwrap();
16486 let sd = scratch_d.as_mut().unwrap();
16487 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
16488 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
16489 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
16490 if grouped_q8 {
16491 let (zq, zd) = e.quantize_q8_1(&gathered, m_e, n_embd)?;
16492 let gate = e.qmatvec_expert_q8(
16493 sg,
16494 0..gl.len,
16495 &zq,
16496 &zd,
16497 m_e,
16498 m.gate_exps.in_f,
16499 m.gate_exps.out_f,
16500 gl.qtype,
16501 gl.row_bytes,
16502 )?;
16503 let up = e.qmatvec_expert_q8(
16504 su,
16505 0..ul.len,
16506 &zq,
16507 &zd,
16508 m_e,
16509 m.up_exps.in_f,
16510 m.up_exps.out_f,
16511 ul.qtype,
16512 ul.row_bytes,
16513 )?;
16514 let mut act = e.uninit(m_e * n_ff_exp)?;
16515 Self::ffn_act_lim(
16516 e,
16517 cfg,
16518 &gate,
16519 &up,
16520 m.gate_exps.macro_scale(ex),
16521 m.up_exps.macro_scale(ex),
16522 lim_exp,
16523 &mut act,
16524 m_e * n_ff_exp,
16525 )?;
16526 let (aq2, ad2) = e.quantize_q8_1(&act, m_e, n_ff_exp)?;
16527 e.qmatvec_expert_q8(
16528 sd,
16529 0..dl.len,
16530 &aq2,
16531 &ad2,
16532 m_e,
16533 m.down_exps.in_f,
16534 m.down_exps.out_f,
16535 dl.qtype,
16536 dl.row_bytes,
16537 )?
16538 } else {
16539 let gate = m.qmatvec_view(
16540 e,
16541 sg,
16542 0..gl.len,
16543 &gv,
16544 m_e,
16545 m.gate_exps.in_f,
16546 m.gate_exps.out_f,
16547 gl.qtype,
16548 gl.row_bytes,
16549 )?;
16550 let up = m.qmatvec_view(
16551 e,
16552 su,
16553 0..ul.len,
16554 &gv,
16555 m_e,
16556 m.up_exps.in_f,
16557 m.up_exps.out_f,
16558 ul.qtype,
16559 ul.row_bytes,
16560 )?;
16561 let mut act = e.uninit(m_e * n_ff_exp)?;
16562 Self::ffn_act_lim(
16563 e,
16564 cfg,
16565 &gate,
16566 &up,
16567 m.gate_exps.macro_scale(ex),
16568 m.up_exps.macro_scale(ex),
16569 lim_exp,
16570 &mut act,
16571 m_e * n_ff_exp,
16572 )?;
16573 let actv = act.slice(0..m_e * n_ff_exp);
16574 m.qmatvec_view(
16575 e,
16576 sd,
16577 0..dl.len,
16578 &actv,
16579 m_e,
16580 m.down_exps.in_f,
16581 m.down_exps.out_f,
16582 dl.qtype,
16583 dl.row_bytes,
16584 )?
16585 }
16586 };
16587
16588 e.scatter_slot(
16590 &y,
16591 &tok_idx_d,
16592 &slot_idx_d,
16593 &weight_d,
16594 &mut slot_buf,
16595 &mut wbuf,
16596 n_embd,
16597 n_used,
16598 m_e,
16599 )?;
16600 }
16601
16602 let mut moe_out = e.zeros(t * n_embd)?;
16604 e.reduce_slots(&slot_buf, &wbuf, &mut moe_out, n_embd, n_used, t)?;
16605
16606 if std::env::var("MEMRA_MOE_STATS").is_ok() && !m_dist.is_empty() {
16608 m_dist.sort_unstable();
16609 let active = m_dist.len();
16610 let mean = m_dist.iter().sum::<usize>() as f64 / active as f64;
16611 let median = m_dist[active / 2];
16612 let max_m = *m_dist.last().unwrap();
16613 let min_m = m_dist[0];
16614 let above16 = m_dist.iter().filter(|&&x| x >= 16).count();
16615 println!(
16616 "moe-grouped il={il} t={t} active={active}/{n_expert} \
16617 m_e: min={min_m} median={median} mean={mean:.1} max={max_m} \
16618 above_gemm_threshold(>=16)={above16}/{active}"
16619 );
16620 }
16621
16622 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut moe_out)?;
16623 Ok(moe_out)
16624 }
16625
16626 pub(crate) fn moe_ffn_lockstep(
16633 &self,
16634 e: &Engine,
16635 m: &MoeWeights,
16636 zbatch: &CudaSlice<f32>,
16637 mrows: usize,
16638 il: u16,
16639 max_block: usize,
16640 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
16641 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
16642 let cfg = &self.cfg;
16643 let moe = cfg.moe.as_ref().unwrap();
16644 let n_embd = cfg.n_embd as usize;
16645 let n_expert = moe.expert_count as usize;
16646 let n_used = moe.expert_used_count as usize;
16647 let n_ff_exp = moe.expert_ff_length as usize;
16648 let lim_exp = cfg.clamp_exp_at(il as u32);
16650 let lim_shexp = cfg.clamp_shexp_at(il as u32);
16651
16652 let logits = e.matmul(&m.gate_inp, zbatch, mrows)?;
16653 if let Some(sig) = cfg.sigmoid_router() {
16654 Self::trace_sigmoid_router_logits(e, il, mrows, n_expert, n_used, &logits, m, sig)?;
16655 }
16656 let (sel_all, w_all) = if let Some(sig) = cfg.sigmoid_router() {
16657 Self::moe_route_sigmoid_cfg(e, &logits, mrows, n_expert, n_used, m, sig)?
16658 } else {
16659 Self::moe_route_cfg(
16660 e,
16661 &logits,
16662 mrows,
16663 n_expert,
16664 n_used,
16665 m.active_experts.as_deref(),
16666 )?
16667 };
16668 Self::trace_moe_routes(il, mrows, &sel_all, &w_all)?;
16669
16670 let resident_expert: Vec<bool> = e.with_moe_cache(max_block, |c, _| {
16672 Ok((0..n_expert)
16673 .map(|ex| {
16674 [PROJ_GATE, PROJ_UP, PROJ_DOWN]
16675 .into_iter()
16676 .all(|p| c.resident(BlockId::new(il, p, ex as u16)).is_some())
16677 })
16678 .collect())
16679 })?;
16680
16681 struct Group {
16682 rows: Vec<i32>,
16683 slots: Vec<i32>,
16684 weights: Vec<f32>,
16685 }
16686 let mut groups: std::collections::HashMap<usize, Group> = Default::default();
16687 let mut cpu_rows: Vec<Vec<(usize, f32)>> = vec![Vec::new(); mrows];
16688 let mut cpu_by_expert: std::collections::HashMap<usize, Vec<(usize, f32)>> =
16689 Default::default();
16690 for row in 0..mrows {
16691 for j in 0..n_used {
16692 let ex = sel_all[row * n_used + j] as usize;
16693 let w = w_all[row * n_used + j];
16694 if resident_expert[ex] {
16695 let group = groups.entry(ex).or_insert_with(|| Group {
16696 rows: Vec::new(),
16697 slots: Vec::new(),
16698 weights: Vec::new(),
16699 });
16700 group.rows.push(row as i32);
16701 group.slots.push(j as i32);
16702 group.weights.push(w);
16703 } else {
16704 crate::cpu_experts::record_incomplete_gpu_residency(0);
16705 cpu_rows[row].push((ex, w));
16706 cpu_by_expert.entry(ex).or_default().push((row, w));
16707 }
16708 }
16709 }
16710
16711 let host_rows = e.dtoh(zbatch)?;
16717 let rows_ok = crate::cpu_experts::rows_supported();
16718 enum CpuPart {
16719 Single { row: usize },
16720 Rows { rows: Vec<usize> },
16721 }
16722 let mut tickets: Vec<(CpuPart, crate::cpu_experts::CpuExpertTicket)> = Vec::new();
16723 let mut rows_served: std::collections::HashSet<(usize, usize)> = Default::default();
16724 if rows_ok {
16725 let mut shared: Vec<(usize, Vec<(usize, f32)>)> = cpu_by_expert
16726 .into_iter()
16727 .filter(|(_, rows)| rows.len() >= 2)
16728 .collect();
16729 shared.sort_by_key(|(ex, _)| *ex);
16730 for (ex, mut row_weights) in shared {
16731 row_weights.sort_by_key(|(row, _)| *row);
16732 let inputs: Vec<(&[f32], f32)> = row_weights
16733 .iter()
16734 .map(|&(row, w)| (&host_rows[row * n_embd..(row + 1) * n_embd], w))
16735 .collect();
16736 let job = crate::cpu_experts::prepare_rows_job(m, ex, &inputs)
16737 .map_err(std::io::Error::other)?;
16738 for &(row, _) in &row_weights {
16739 rows_served.insert((row, ex));
16740 }
16741 tickets.push((
16742 CpuPart::Rows {
16743 rows: row_weights.iter().map(|&(row, _)| row).collect(),
16744 },
16745 crate::cpu_experts::submit_rows(job).map_err(std::io::Error::other)?,
16746 ));
16747 }
16748 }
16749 for (row, selected) in cpu_rows.iter().enumerate() {
16750 let leftover: Vec<(usize, f32)> = selected
16751 .iter()
16752 .copied()
16753 .filter(|&(ex, _)| !rows_served.contains(&(row, ex)))
16754 .collect();
16755 if leftover.is_empty() {
16756 continue;
16757 }
16758 let host_row = &host_rows[row * n_embd..(row + 1) * n_embd];
16759 let job = crate::cpu_experts::prepare_job(m, il, &leftover, host_row)
16760 .map_err(std::io::Error::other)?;
16761 tickets.push((
16762 CpuPart::Single { row },
16763 crate::cpu_experts::submit(job).map_err(std::io::Error::other)?,
16764 ));
16765 }
16766
16767 let mut slot_buf = e.zeros(mrows * n_used * n_embd)?;
16768 let mut wbuf = e.zeros(mrows * n_used)?;
16769 let mut order: Vec<usize> = groups.keys().copied().collect();
16770 order.sort_by(|&a, &b| {
16771 groups[&b]
16772 .rows
16773 .len()
16774 .cmp(&groups[&a].rows.len())
16775 .then(a.cmp(&b))
16776 });
16777 for &ex in &order {
16778 let group = &groups[&ex];
16779 let m_e = group.rows.len();
16780 let gl = m.gate_exps.expert_layout(ex);
16781 let ul = m.up_exps.expert_layout(ex);
16782 let dl = m.down_exps.expert_layout(ex);
16783 let row_idx_d = e.htod_i32(&group.rows)?;
16784 let slot_idx_d = e.htod_i32(&group.slots)?;
16785 let dmac = m.down_exps.macro_scale(ex);
16786 let weight_d = if dmac == 1.0 {
16787 e.htod(&group.weights)?
16788 } else {
16789 let scaled: Vec<f32> = group.weights.iter().map(|&w| w * dmac).collect();
16790 e.htod(&scaled)?
16791 };
16792 let mut gathered = e.zeros(m_e * n_embd)?;
16793 e.gather_rows(zbatch, &row_idx_d, &mut gathered, n_embd, m_e)?;
16794 let gv = gathered.slice(0..m_e * n_embd);
16795 let gate = e.with_moe_cache(max_block, |c, eng| {
16796 let slot = c
16797 .resident(BlockId::new(il, PROJ_GATE, ex as u16))
16798 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
16799 m.qmatvec_view(
16800 eng,
16801 c.buf(crate::moe_cache::DispatchSlot::Resident(slot)),
16802 0..gl.len,
16803 &gv,
16804 m_e,
16805 m.gate_exps.in_f,
16806 m.gate_exps.out_f,
16807 gl.qtype,
16808 gl.row_bytes,
16809 )
16810 })?;
16811 let up = e.with_moe_cache(max_block, |c, eng| {
16812 let slot = c
16813 .resident(BlockId::new(il, PROJ_UP, ex as u16))
16814 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
16815 m.qmatvec_view(
16816 eng,
16817 c.buf(crate::moe_cache::DispatchSlot::Resident(slot)),
16818 0..ul.len,
16819 &gv,
16820 m_e,
16821 m.up_exps.in_f,
16822 m.up_exps.out_f,
16823 ul.qtype,
16824 ul.row_bytes,
16825 )
16826 })?;
16827 let mut act = e.zeros(m_e * n_ff_exp)?;
16828 Self::ffn_act_lim(
16829 e,
16830 cfg,
16831 &gate,
16832 &up,
16833 m.gate_exps.macro_scale(ex),
16834 m.up_exps.macro_scale(ex),
16835 lim_exp,
16836 &mut act,
16837 m_e * n_ff_exp,
16838 )?;
16839 let actv = act.slice(0..m_e * n_ff_exp);
16840 let y = e.with_moe_cache(max_block, |c, eng| {
16841 let slot = c
16842 .resident(BlockId::new(il, PROJ_DOWN, ex as u16))
16843 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
16844 m.qmatvec_view(
16845 eng,
16846 c.buf(crate::moe_cache::DispatchSlot::Resident(slot)),
16847 0..dl.len,
16848 &actv,
16849 m_e,
16850 m.down_exps.in_f,
16851 m.down_exps.out_f,
16852 dl.qtype,
16853 dl.row_bytes,
16854 )
16855 })?;
16856 e.scatter_slot(
16857 &y,
16858 &row_idx_d,
16859 &slot_idx_d,
16860 &weight_d,
16861 &mut slot_buf,
16862 &mut wbuf,
16863 n_embd,
16864 n_used,
16865 m_e,
16866 )?;
16867 }
16868 let mut moe_out = e.zeros(mrows * n_embd)?;
16869 e.reduce_slots(&slot_buf, &wbuf, &mut moe_out, n_embd, n_used, mrows)?;
16870
16871 let mut row_sums: Vec<Option<Vec<f32>>> = vec![None; mrows];
16873 for (part, ticket) in tickets {
16874 let cpu_output = ticket.wait().map_err(std::io::Error::other)?;
16875 let mut add_row = |row: usize, chunk: &[f32]| {
16876 let sum = row_sums[row].get_or_insert_with(|| vec![0.0f32; n_embd]);
16877 for (accumulator, value) in sum.iter_mut().zip(chunk) {
16878 *accumulator += value;
16879 }
16880 };
16881 match part {
16882 CpuPart::Single { row } => add_row(row, &cpu_output),
16883 CpuPart::Rows { rows } => {
16884 for (slot, row) in rows.into_iter().enumerate() {
16885 add_row(row, &cpu_output[slot * n_embd..(slot + 1) * n_embd]);
16886 }
16887 }
16888 }
16889 }
16890 for (row, sum) in row_sums.into_iter().enumerate() {
16891 let Some(sum) = sum else { continue };
16892 let cpu_output = e.htod(&sum)?;
16893 let mut dst = moe_out.slice_mut(row * n_embd..(row + 1) * n_embd);
16894 e.axpy_into(&cpu_output, 1.0, &mut dst, n_embd)?;
16895 }
16896
16897 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
16898 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
16899 {
16900 let n_ff_sh = gate_shexp.out_features();
16901 let sg_gate = e.matmul(gate_shexp, zbatch, mrows)?;
16902 let sg_up = e.matmul(up_shexp, zbatch, mrows)?;
16903 let mut sa = e.zeros(mrows * n_ff_sh)?;
16904 Self::ffn_act_lim(
16905 e,
16906 cfg,
16907 &sg_gate,
16908 &sg_up,
16909 1.0,
16910 1.0,
16911 lim_shexp,
16912 &mut sa,
16913 mrows * n_ff_sh,
16914 )?;
16915 let sh = e.matmul(down_shexp, &sa, mrows)?;
16916 let g = match &m.gate_inp_shexp {
16919 Some(gate_inp_shexp) => {
16920 e.sigmoid_dot_rows(zbatch, gate_inp_shexp.float_data(), n_embd, mrows)?
16921 }
16922 None => e.htod(&vec![1.0f32; mrows])?,
16923 };
16924 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, mrows)?;
16925 }
16926
16927 Ok(moe_out)
16928 }
16929}
16930
16931impl HybridModel {
16937 pub(crate) fn gemma4_rope_dims(&self, il: usize) -> usize {
16951 let g = self
16952 .cfg
16953 .gemma4
16954 .as_ref()
16955 .expect("gemma4_rope_dims on a non-gemma4 config");
16956 if g.swa_pattern[il] {
16957 g.rope_dims_swa as usize
16958 } else {
16959 g.rope_dims_global as usize
16960 }
16961 }
16962
16963 pub(crate) fn gemma4_geom(&self, il: usize) -> (usize, usize, usize, f32, f32, bool) {
16964 let g = self.cfg.gemma4.as_ref().unwrap();
16965 let swa = g.swa_pattern[il];
16966 let hd = if swa {
16967 g.key_length_swa
16968 } else {
16969 g.key_length_global
16970 } as usize;
16971 (
16975 hd,
16976 g.head_count_kv[il] as usize,
16977 self.cfg.n_head as usize,
16978 if swa {
16979 g.rope_base_swa
16980 } else {
16981 g.rope_base_global
16982 },
16983 1.0,
16984 swa,
16985 )
16986 }
16987
16988 pub(crate) fn gemma4_suppress(
16992 &self,
16993 e: &Engine,
16994 ld: &mut CudaSlice<f32>,
16995 t: usize,
16996 ) -> Result<(), Box<dyn std::error::Error>> {
16997 if let Some((ids, n)) = self.gemma4_aux.as_ref().and_then(|a| a.suppress_d.as_ref()) {
16998 #[cfg(debug_assertions)]
17003 crate::debug_assert_tensor_stream_device(
17004 ids,
17005 &e.stream(),
17006 "gemma4_suppress.suppress_d",
17007 );
17008 e.mask_ids_rows(ld, ids, *n, self.output.out_features(), t)?;
17009 }
17010 Ok(())
17011 }
17012
17013 #[allow(clippy::too_many_arguments)]
17018 fn gemma_fa_one_program() -> bool {
17027 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
17028 *ON.get_or_init(|| std::env::var("MEMRA_GEMMA_FA_ONE_PROGRAM").as_deref() == Ok("1"))
17029 }
17030
17031 #[allow(clippy::too_many_arguments)] fn gemma4_attn_prime(
17033 &self,
17034 e: &Engine,
17035 fa: &crate::hybrid::FullAttnLayer,
17036 il: usize,
17037 h: &CudaSlice<f32>,
17038 pos_d: &CudaSlice<i32>,
17039 t: usize,
17040 cache: Option<&mut Cache>,
17041 island: Option<&CudaSlice<i32>>,
17042 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
17043 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
17044 let eps = self.cfg.rms_eps;
17045 let aux = self.gemma4_aux.as_ref().unwrap();
17046 let ones = aux.ones(e);
17047 #[cfg(debug_assertions)]
17048 crate::debug_assert_tensor_stream_device(ones, &e.stream(), "gemma4_attn_prime.ones");
17049
17050 e.mmq_act_begin();
17053 let q0 = e.matmul(&fa.wq, h, t)?; if il == 0 && std::env::var("MEMRA_G4_PRIME_TRACE").as_deref() == Ok("1") {
17055 let v = e.dtoh(&q0)?;
17056 let nan = v.iter().filter(|x| x.is_nan()).count();
17057 let amax = v.iter().fold(0f32, |a, x| a.max(x.abs()));
17058 eprintln!(
17059 "[g4-prime-trace] L0 q0: nan={nan}/{} amax={amax:.3}",
17060 v.len()
17061 );
17062 }
17063 let k0 = e.matmul(&fa.wk, h, t)?; let v0 = if swa {
17067 e.matmul(&fa.wv, h, t)?
17068 } else {
17069 e.clone_dtod(&k0)?
17070 };
17071 if il == 0 && std::env::var("MEMRA_G4_PRIME_TRACE").as_deref() == Ok("1") {
17072 for (tag, buf) in [("k0", &k0), ("v0", &v0)] {
17073 let v = e.dtoh(buf)?;
17074 let nan = v.iter().filter(|x| x.is_nan()).count();
17075 let amax = v.iter().fold(0f32, |a, x| a.max(x.abs()));
17076 eprintln!(
17077 "[g4-prime-trace] L0 {tag}: nan={nan}/{} amax={amax:.3}",
17078 v.len()
17079 );
17080 }
17081 }
17082
17083 let mut q = e.uninit(t * nh * hd)?;
17084 let mut k = e.uninit(t * nkv * hd)?;
17085 let mut v = e.uninit(t * nkv * hd)?;
17087 static EMIT: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
17091 let emit = island.is_none()
17094 && t >= 16
17095 && crate::Engine::qkvnorm_w_on_prefill(nh * t + 2 * nkv * t, hd)
17096 && *EMIT.get_or_init(|| {
17097 std::env::var("MEMRA_FA_EMIT")
17098 .map(|s| s != "0")
17099 .unwrap_or(true)
17100 });
17101 let mut qb = e.alloc_uninit::<u8>(if emit { t * nh * hd * 2 } else { 1 })?;
17102 let mut kb = e.alloc_uninit::<u8>(if emit { t * nkv * hd * 2 } else { 1 })?;
17103 let mut vb = e.alloc_uninit::<u8>(if emit { t * nkv * hd * 2 } else { 1 })?;
17104 let v_f16 = emit
17107 && crate::fa_f16pv_on()
17108 && match hd {
17109 512 => true,
17110 256 => swa && crate::faw_hp_on() && nh % 2 == 0 && (nh / nkv) % 2 == 0,
17111 _ => false,
17112 };
17113 if emit {
17114 e.rms_norm_qkv_w4b(
17115 &q0,
17116 &k0,
17117 &v0,
17118 fa.q_norm.float_data(),
17119 fa.k_norm.float_data(),
17120 ones,
17121 &mut q,
17122 &mut k,
17123 &mut v,
17124 &mut vb,
17125 hd,
17126 nh * t,
17127 nkv * t,
17128 eps,
17129 v_f16,
17130 )?;
17131 } else {
17132 e.rms_norm_qkv(
17133 &q0,
17134 &k0,
17135 &v0,
17136 fa.q_norm.float_data(),
17137 fa.k_norm.float_data(),
17138 ones,
17139 &mut q,
17140 &mut k,
17141 &mut v,
17142 hd,
17143 nh * t,
17144 nkv * t,
17145 eps,
17146 )?;
17147 }
17148
17149 let ff = if swa {
17150 None
17151 } else {
17152 Some(
17153 aux.rope_freqs(e)
17154 .expect("gemma4 global rope needs rope_freqs.weight"),
17155 )
17156 };
17157 #[cfg(debug_assertions)]
17158 if let Some(ff) = ff {
17159 crate::debug_assert_tensor_stream_device(
17160 ff,
17161 &e.stream(),
17162 "gemma4_attn_prime.rope_freqs",
17163 );
17164 }
17165 if emit {
17166 e.rope_neox2_bf16e(
17167 &mut q, &mut k, &mut qb, &mut kb, pos_d, hd, hd, nh, nkv, t, base, 1.0, ff,
17168 )?;
17169 } else {
17170 e.rope_neox2(&mut q, &mut k, pos_d, hd, hd, nh, nkv, t, base, 1.0, ff)?;
17171 }
17172
17173 if let Some(cache) = cache {
17174 let kvl = cache.kv[il].as_mut().unwrap();
17175 assert_eq!(kvl.len, 0, "gemma4 prime is fresh-prompt only (v0)");
17176 e.append_kv_quantized_rows(
17177 &k,
17178 &v,
17179 &mut kvl.k,
17180 &mut kvl.v,
17181 kvl.len,
17182 t,
17183 kvl.kv_dim_k,
17184 kvl.kv_dim_v,
17185 kvl.k_tok_bytes,
17186 kvl.v_tok_bytes,
17187 (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on()),
17188 )?;
17189 kvl.len += t;
17190 }
17191 let mut attn = e.zeros(t * nh * hd)?;
17192 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
17196 if let Some(span) = island {
17197 let w = if swa && t > win { win } else { 0 };
17202 e.sdpa_naive_island(&q, &k, &v, &mut attn, span, hd, nh, nkv, t, t, scale, w)?;
17203 } else if swa && (t > win || Self::gemma_fa_one_program()) {
17204 if hd == 256 && std::env::var("MEMRA_NOFA").is_err() {
17205 if emit {
17206 e.fa_prefill_w_pre(
17207 &qb, &kb, &vb, &mut attn, hd, nh, nkv, t, t, scale, true, win, v_f16,
17208 )?;
17209 } else {
17210 e.fa_prefill_w(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
17211 }
17212 } else {
17213 e.sdpa_naive_w(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
17214 }
17215 } else if hd == 256 && std::env::var("MEMRA_NOFA").is_err() {
17216 e.fa_prefill(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
17217 } else if hd == 512 && std::env::var("MEMRA_NOFA").is_err() {
17218 if emit {
17219 e.fa_prefill_hd512_pre(
17220 &qb, &kb, &vb, &mut attn, hd, nh, nkv, t, t, scale, true, v_f16,
17221 )?;
17222 } else {
17223 e.fa_prefill_hd512(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
17224 }
17225 } else {
17226 e.sdpa_naive(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
17227 }
17228 e.matmul(&fa.wo, &attn, t)
17229 }
17230
17231 fn gemma4_attn(
17233 &self,
17234 e: &Engine,
17235 fa: &crate::hybrid::FullAttnLayer,
17236 il: usize,
17237 h: &CudaSlice<f32>,
17238 pos_d: &CudaSlice<i32>,
17239 t: usize,
17240 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
17241 self.gemma4_attn_prime(e, fa, il, h, pos_d, t, None, None)
17242 }
17243
17244 fn gemma4_moe_q8(
17249 &self,
17250 e: &Engine,
17251 m: &crate::hybrid::MoeWeights,
17252 bits: &crate::hybrid::Gemma4MoeBits,
17253 mq: &(CudaSlice<i8>, CudaSlice<f32>),
17254 router_in: &CudaSlice<f32>,
17255 t: usize,
17256 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
17257 let cfg = &self.cfg;
17258 let moe = cfg.moe.as_ref().unwrap();
17259 let n_embd = cfg.n_embd as usize;
17260 let n_expert = moe.expert_count as usize;
17261 let n_used = moe.expert_used_count as usize;
17262 let n_ff_exp = moe.expert_ff_length as usize;
17263 let logits = if crate::router_kernel_on() {
17267 e.router_gemv(m.gate_inp.float_data(), router_in, n_embd, n_expert, t)?
17268 } else {
17269 e.matmul(&m.gate_inp, router_in, t)?
17270 };
17271 let dev = m.dev_exps.as_ref().unwrap();
17272 let (sel_d, w_d) =
17273 e.moe_router_topk_scaled(&logits, t, n_expert, n_used, &bits.per_expert_scale_d)?;
17274 let (zq, zd) = mq;
17275 if t == 1 {
17276 let selv = sel_d.slice(0..n_used);
17277 let wv = w_d.slice(0..n_used);
17278 let act = e.moe_gate_up_gelu8_dev_q8(
17279 &dev.ptr_row,
17280 &selv,
17281 zq,
17282 zd,
17283 n_embd,
17284 n_ff_exp,
17285 n_used,
17286 n_expert,
17287 m.gate_exps.qtype,
17288 m.up_exps.qtype,
17289 m.gate_exps.row_bytes,
17290 m.up_exps.row_bytes,
17291 )?;
17292 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
17293 let mut moe_out = e.uninit(n_embd)?;
17294 e.moe_down8_fma_dev_q8(
17295 &dev.ptr_row,
17296 &selv,
17297 &wv,
17298 &aq2,
17299 &ad2,
17300 &mut moe_out.slice_mut(0..n_embd),
17301 n_ff_exp,
17302 n_embd,
17303 n_used,
17304 n_expert,
17305 m.down_exps.qtype,
17306 m.down_exps.row_bytes,
17307 )?;
17308 return Ok(moe_out);
17309 }
17310 let csr = t <= 10 && std::env::var("MEMRA_GEMMA_CSR").as_deref() != Ok("0");
17311 let act = if csr {
17312 e.moe_gate_up_gelu8_dev_q8_csr(
17313 &dev.ptr_row,
17314 &sel_d,
17315 zq,
17316 zd,
17317 t * n_used,
17318 n_embd,
17319 n_ff_exp,
17320 n_used,
17321 n_expert,
17322 m.gate_exps.qtype,
17323 m.up_exps.qtype,
17324 m.gate_exps.row_bytes,
17325 m.up_exps.row_bytes,
17326 )?
17327 } else {
17328 e.moe_gate_up_gelu8_dev_q8_rows(
17329 &dev.ptr_row,
17330 &sel_d,
17331 zq,
17332 zd,
17333 t,
17334 n_embd,
17335 n_ff_exp,
17336 n_used,
17337 n_expert,
17338 m.gate_exps.qtype,
17339 m.up_exps.qtype,
17340 m.gate_exps.row_bytes,
17341 m.up_exps.row_bytes,
17342 )?
17343 };
17344 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
17345 let mut moe_out = e.uninit(t * n_embd)?;
17346 e.moe_down8_fma_dev_q8_rows_g(
17349 &dev.ptr_row,
17350 &sel_d,
17351 &w_d,
17352 &aq2,
17353 &ad2,
17354 &mut moe_out,
17355 t,
17356 n_ff_exp,
17357 n_embd,
17358 n_used,
17359 n_expert,
17360 m.down_exps.qtype,
17361 m.down_exps.row_bytes,
17362 )?;
17363 Ok(moe_out)
17364 }
17365
17366 fn gemma4_moe(
17370 &self,
17371 e: &Engine,
17372 m: &crate::hybrid::MoeWeights,
17373 bits: &crate::hybrid::Gemma4MoeBits,
17374 moe_in: &CudaSlice<f32>,
17375 router_in: &CudaSlice<f32>,
17376 t: usize,
17377 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
17378 let cfg = &self.cfg;
17379 let moe = cfg.moe.as_ref().unwrap();
17380 let n_embd = cfg.n_embd as usize;
17381 let n_expert = moe.expert_count as usize;
17382 let n_used = moe.expert_used_count as usize;
17383 let n_ff_exp = moe.expert_ff_length as usize;
17384
17385 let logits = if t < PRIME_MIN_T && crate::router_kernel_on() {
17389 e.router_gemv(m.gate_inp.float_data(), router_in, n_embd, n_expert, t)?
17390 } else {
17391 e.matmul(&m.gate_inp, router_in, t)?
17392 };
17393
17394 if t < PRIME_MIN_T
17399 && m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
17400 && expert_dp4a_supported(m.gate_exps.qtype)
17401 && expert_dp4a_supported(m.up_exps.qtype)
17402 && expert_dp4a_supported(m.down_exps.qtype)
17403 && std::env::var("MEMRA_GEMMA_MOE_FAST").as_deref() != Ok("0")
17404 {
17405 let dev = m.dev_exps.as_ref().unwrap();
17406 let (sel_d, w_d) =
17407 e.moe_router_topk_scaled(&logits, t, n_expert, n_used, &bits.per_expert_scale_d)?;
17408 if t == 1 {
17409 let (zq, zd) = e.quantize_q8_1(moe_in, 1, n_embd)?;
17410 let selv = sel_d.slice(0..n_used);
17411 let wv = w_d.slice(0..n_used);
17412 let act = e.moe_gate_up_gelu8_dev_q8(
17413 &dev.ptr_row,
17414 &selv,
17415 &zq,
17416 &zd,
17417 n_embd,
17418 n_ff_exp,
17419 n_used,
17420 n_expert,
17421 m.gate_exps.qtype,
17422 m.up_exps.qtype,
17423 m.gate_exps.row_bytes,
17424 m.up_exps.row_bytes,
17425 )?;
17426 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
17427 let mut moe_out = e.uninit(n_embd)?;
17428 e.moe_down8_fma_dev_q8(
17429 &dev.ptr_row,
17430 &selv,
17431 &wv,
17432 &aq2,
17433 &ad2,
17434 &mut moe_out.slice_mut(0..n_embd),
17435 n_ff_exp,
17436 n_embd,
17437 n_used,
17438 n_expert,
17439 m.down_exps.qtype,
17440 m.down_exps.row_bytes,
17441 )?;
17442 return Ok(moe_out);
17443 }
17444 let (zq, zd) = e.quantize_q8_1(moe_in, t, n_embd)?;
17449 let csr = t <= 10 && std::env::var("MEMRA_GEMMA_CSR").as_deref() != Ok("0");
17450 let act = if csr {
17451 e.moe_gate_up_gelu8_dev_q8_csr(
17452 &dev.ptr_row,
17453 &sel_d,
17454 &zq,
17455 &zd,
17456 t * n_used,
17457 n_embd,
17458 n_ff_exp,
17459 n_used,
17460 n_expert,
17461 m.gate_exps.qtype,
17462 m.up_exps.qtype,
17463 m.gate_exps.row_bytes,
17464 m.up_exps.row_bytes,
17465 )?
17466 } else {
17467 e.moe_gate_up_gelu8_dev_q8_rows(
17468 &dev.ptr_row,
17469 &sel_d,
17470 &zq,
17471 &zd,
17472 t,
17473 n_embd,
17474 n_ff_exp,
17475 n_used,
17476 n_expert,
17477 m.gate_exps.qtype,
17478 m.up_exps.qtype,
17479 m.gate_exps.row_bytes,
17480 m.up_exps.row_bytes,
17481 )?
17482 };
17483 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
17484 let mut moe_out = e.uninit(t * n_embd)?;
17485 e.moe_down8_fma_dev_q8_rows_g(
17486 &dev.ptr_row,
17487 &sel_d,
17488 &w_d,
17489 &aq2,
17490 &ad2,
17491 &mut moe_out,
17492 t,
17493 n_ff_exp,
17494 n_embd,
17495 n_used,
17496 n_expert,
17497 m.down_exps.qtype,
17498 m.down_exps.row_bytes,
17499 )?;
17500 return Ok(moe_out);
17501 }
17502
17503 let (sel_all, mut w_all) = Self::moe_route(e, &logits, t, n_expert, n_used)?;
17504 for (i, &sx) in sel_all.iter().enumerate() {
17505 w_all[i] *= bits.per_expert_scale[sx as usize];
17506 }
17507
17508 if t >= PRIME_MIN_T
17512 && m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
17513 && expert_dp4a_supported(m.gate_exps.qtype)
17514 && expert_dp4a_supported(m.up_exps.qtype)
17515 && expert_dp4a_supported(m.down_exps.qtype)
17516 && std::env::var("MEMRA_GEMMA_MOE_PAIRS").as_deref() != Ok("0")
17517 {
17518 let dev = m.dev_exps.as_ref().unwrap();
17519 let n_pairs = t * n_used;
17520 let pair_ex: Vec<i32> = sel_all.iter().map(|&x| x as i32).collect();
17521 let pair_tok: Vec<i32> = (0..n_pairs).map(|p| (p / n_used) as i32).collect();
17522 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
17523 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
17524 let pt = e.htod_i32(&pair_tok)?;
17525 let pw = e.htod(&w_all)?;
17526 let toff = e.htod_i32(&tok_off)?;
17527 let tids = e.htod_i32(&tok_ids)?;
17528 let mut by_ex: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
17529 for p in 0..n_pairs {
17530 by_ex[pair_ex[p] as usize].push(p as i32);
17531 }
17532 let mut ex_ids: Vec<i32> = Vec::new();
17533 let mut ex_off: Vec<i32> = vec![0];
17534 let mut ex_pairs: Vec<i32> = Vec::with_capacity(n_pairs);
17535 for (ex, list) in by_ex.iter().enumerate() {
17536 if list.is_empty() {
17537 continue;
17538 }
17539 ex_ids.push(ex as i32);
17540 ex_pairs.extend_from_slice(list);
17541 ex_off.push(ex_pairs.len() as i32);
17542 }
17543 let n_active = ex_ids.len();
17544 let exi = e.htod_i32(&ex_ids)?;
17545 let exo = e.htod_i32(&ex_off)?;
17546 let exp_d = e.htod_i32(&ex_pairs)?;
17547 if crate::moe_f16g_gemma_on()
17555 && f16g_proj_ok(m.gate_exps.qtype, n_embd)
17556 && f16g_proj_ok(m.up_exps.qtype, n_embd)
17557 && f16g_proj_ok(m.down_exps.qtype, n_ff_exp)
17558 {
17559 let csr_tok: Vec<i32> = ex_pairs.iter().map(|&p| p / n_used as i32).collect();
17560 let csr_tok_d = e.htod_i32(&csr_tok)?;
17561 let (z_f16, z_s) = e.moe_f16g_act(moe_in, Some(&csr_tok_d), n_embd, n_pairs)?;
17562 let g_csr = e.moe_f16_grouped(
17563 &dev.ptr_row,
17564 0,
17565 n_expert,
17566 &exi,
17567 &ex_off,
17568 &exo,
17569 &z_f16,
17570 &z_s,
17571 n_embd,
17572 n_ff_exp,
17573 n_active,
17574 n_pairs,
17575 m.gate_exps.qtype,
17576 m.gate_exps.row_bytes,
17577 )?;
17578 let u_csr = e.moe_f16_grouped(
17579 &dev.ptr_row,
17580 1,
17581 n_expert,
17582 &exi,
17583 &ex_off,
17584 &exo,
17585 &z_f16,
17586 &z_s,
17587 n_embd,
17588 n_ff_exp,
17589 n_active,
17590 n_pairs,
17591 m.up_exps.qtype,
17592 m.up_exps.row_bytes,
17593 )?;
17594 let act_csr = e.moe_pairs_gelu_mul(&g_csr, &u_csr, n_pairs * n_ff_exp)?;
17595 let (a_f16, a_s) = e.moe_f16g_act(&act_csr, None, n_ff_exp, n_pairs)?;
17596 let d_csr = e.moe_f16_grouped(
17597 &dev.ptr_row,
17598 2,
17599 n_expert,
17600 &exi,
17601 &ex_off,
17602 &exo,
17603 &a_f16,
17604 &a_s,
17605 n_ff_exp,
17606 n_embd,
17607 n_active,
17608 n_pairs,
17609 m.down_exps.qtype,
17610 m.down_exps.row_bytes,
17611 )?;
17612 let y_down = e.rows_permute(&d_csr, &exp_d, n_pairs, n_embd)?;
17613 let mut moe_out = e.uninit(t * n_embd)?;
17614 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
17615 if std::env::var("MEMRA_F16G_DEBUG").is_ok() {
17616 let scan = |v: &[f32]| v.iter().filter(|x| !x.is_finite()).count();
17617 let (yd, mo) = (e.dtoh(&y_down)?, e.dtoh(&moe_out)?);
17618 eprintln!(
17619 "[f16g-debug] post-permute bad={} post-scatter bad={}",
17620 scan(&yd),
17621 scan(&mo)
17622 );
17623 }
17624 return Ok(moe_out);
17625 }
17626 let mma = n_embd.is_multiple_of(256)
17629 && std::env::var("MEMRA_GEMMA_MOE_MMA").as_deref() != Ok("0");
17630 let (gate, up) = if mma {
17631 let z_scr = e.mmq_iq_quantize_act(moe_in, n_embd, t)?;
17632 (
17633 e.mmq_iq_experts(
17634 &dev.ptr_row,
17635 0,
17636 n_expert,
17637 &exi,
17638 &exo,
17639 &exp_d,
17640 &pt,
17641 &z_scr,
17642 n_embd,
17643 n_ff_exp,
17644 n_active,
17645 n_pairs,
17646 t,
17647 m.gate_exps.qtype,
17648 m.gate_exps.row_bytes,
17649 )?,
17650 e.mmq_iq_experts(
17651 &dev.ptr_row,
17652 1,
17653 n_expert,
17654 &exi,
17655 &exo,
17656 &exp_d,
17657 &pt,
17658 &z_scr,
17659 n_embd,
17660 n_ff_exp,
17661 n_active,
17662 n_pairs,
17663 t,
17664 m.up_exps.qtype,
17665 m.up_exps.row_bytes,
17666 )?,
17667 )
17668 } else {
17669 let (zq, zd) = e.quantize_q8_1(moe_in, t, n_embd)?;
17670 (
17671 e.moe_pairs_matvec_q8_dec(
17672 &dev.ptr_row,
17673 0,
17674 &exi,
17675 &exo,
17676 &exp_d,
17677 &pt,
17678 &zq,
17679 &zd,
17680 n_embd,
17681 n_ff_exp,
17682 n_expert,
17683 n_active,
17684 n_pairs,
17685 m.gate_exps.qtype,
17686 m.gate_exps.row_bytes,
17687 )?,
17688 e.moe_pairs_matvec_q8_dec(
17689 &dev.ptr_row,
17690 1,
17691 &exi,
17692 &exo,
17693 &exp_d,
17694 &pt,
17695 &zq,
17696 &zd,
17697 n_embd,
17698 n_ff_exp,
17699 n_expert,
17700 n_active,
17701 n_pairs,
17702 m.up_exps.qtype,
17703 m.up_exps.row_bytes,
17704 )?,
17705 )
17706 };
17707 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
17708 let pself = e.htod_i32(&pair_self)?;
17709 let y_down = if mma {
17721 let in_pad = n_ff_exp.div_ceil(256) * 256;
17722 let a_scr = if crate::moe_fuse_actq_on() {
17723 e.mmq_iq_fused_act_quant(&gate, &up, n_ff_exp, n_pairs, 1)?
17724 } else {
17725 let act = e.moe_pairs_gelu_mul(&gate, &up, n_pairs * n_ff_exp)?;
17726 e.mmq_iq_quantize_act(&act, n_ff_exp, n_pairs)?
17727 };
17728 e.mmq_iq_experts(
17729 &dev.ptr_row,
17730 2,
17731 n_expert,
17732 &exi,
17733 &exo,
17734 &exp_d,
17735 &pself,
17736 &a_scr,
17737 in_pad,
17738 n_embd,
17739 n_active,
17740 n_pairs,
17741 n_pairs,
17742 m.down_exps.qtype,
17743 m.down_exps.row_bytes,
17744 )?
17745 } else {
17746 let act = e.moe_pairs_gelu_mul(&gate, &up, n_pairs * n_ff_exp)?;
17747 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
17748 e.moe_pairs_matvec_q8_dec(
17749 &dev.ptr_row,
17750 2,
17751 &exi,
17752 &exo,
17753 &exp_d,
17754 &pself,
17755 &aq2,
17756 &ad2,
17757 n_ff_exp,
17758 n_embd,
17759 n_expert,
17760 n_active,
17761 n_pairs,
17762 m.down_exps.qtype,
17763 m.down_exps.row_bytes,
17764 )?
17765 };
17766 let mut moe_out = e.uninit(t * n_embd)?;
17767 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
17768 return Ok(moe_out);
17769 }
17770
17771 let g_len = m.gate_exps.expert_stride;
17772 let u_len = m.up_exps.expert_stride;
17773 let d_len = m.down_exps.expert_stride;
17774 let dev = m.dev_exps.as_ref().filter(|d| !d.gu_il);
17778 let (mut sg, mut su, mut sd) = if dev.is_some() {
17779 (None, None, None)
17780 } else {
17781 (
17782 Some(e.alloc_u8_uninit(g_len)?),
17783 Some(e.alloc_u8_uninit(u_len)?),
17784 Some(e.alloc_u8_uninit(d_len)?),
17785 )
17786 };
17787 let mut moe_out = e.zeros(t * n_embd)?;
17788 for tok in 0..t {
17789 let sel = &sel_all[tok * n_used..(tok + 1) * n_used];
17790 let w = &w_all[tok * n_used..(tok + 1) * n_used];
17791 let zt = moe_in.slice(tok * n_embd..(tok + 1) * n_embd);
17792 for (j, &ex) in sel.iter().enumerate() {
17793 let ex = ex as usize;
17794 let gate = match dev {
17795 Some(d) => m.qmatvec_view(
17796 e,
17797 &d.gate,
17798 ex * g_len..(ex + 1) * g_len,
17799 &zt,
17800 1,
17801 m.gate_exps.in_f,
17802 m.gate_exps.out_f,
17803 m.gate_exps.qtype,
17804 m.gate_exps.row_bytes,
17805 )?,
17806 None => {
17807 let sg = sg.as_mut().unwrap();
17808 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
17809 m.qmatvec_view(
17810 e,
17811 sg,
17812 0..g_len,
17813 &zt,
17814 1,
17815 m.gate_exps.in_f,
17816 m.gate_exps.out_f,
17817 m.gate_exps.qtype,
17818 m.gate_exps.row_bytes,
17819 )?
17820 }
17821 };
17822 let up = match dev {
17823 Some(d) => m.qmatvec_view(
17824 e,
17825 &d.up,
17826 ex * u_len..(ex + 1) * u_len,
17827 &zt,
17828 1,
17829 m.up_exps.in_f,
17830 m.up_exps.out_f,
17831 m.up_exps.qtype,
17832 m.up_exps.row_bytes,
17833 )?,
17834 None => {
17835 let su = su.as_mut().unwrap();
17836 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
17837 m.qmatvec_view(
17838 e,
17839 su,
17840 0..u_len,
17841 &zt,
17842 1,
17843 m.up_exps.in_f,
17844 m.up_exps.out_f,
17845 m.up_exps.qtype,
17846 m.up_exps.row_bytes,
17847 )?
17848 }
17849 };
17850 let mut act = e.uninit(n_ff_exp)?;
17851 e.gelu_tanh_mul(&gate, &up, &mut act, n_ff_exp)?;
17852 let actv = act.slice(0..n_ff_exp);
17853 let y = match dev {
17854 Some(d) => m.qmatvec_view(
17855 e,
17856 &d.down,
17857 ex * d_len..(ex + 1) * d_len,
17858 &actv,
17859 1,
17860 m.down_exps.in_f,
17861 m.down_exps.out_f,
17862 m.down_exps.qtype,
17863 m.down_exps.row_bytes,
17864 )?,
17865 None => {
17866 let sd = sd.as_mut().unwrap();
17867 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
17868 m.qmatvec_view(
17869 e,
17870 sd,
17871 0..d_len,
17872 &actv,
17873 1,
17874 m.down_exps.in_f,
17875 m.down_exps.out_f,
17876 m.down_exps.qtype,
17877 m.down_exps.row_bytes,
17878 )?
17879 }
17880 };
17881 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
17882 e.axpy_into(&y, w[j], &mut dst, n_embd)?;
17883 }
17884 }
17885 Ok(moe_out)
17886 }
17887
17888 fn gemma4_layer(
17890 &self,
17891 e: &Engine,
17892 il: usize,
17893 layer: &crate::hybrid::HybridLayer,
17894 x: &CudaSlice<f32>,
17895 pos_d: &CudaSlice<i32>,
17896 t: usize,
17897 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
17898 let n_embd = self.cfg.n_embd as usize;
17899 let eps = self.cfg.rms_eps;
17900
17901 let mut h = e.zeros(t * n_embd)?;
17902 e.rms_norm(x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
17903 let Mixer::Full(fa) = &layer.mixer else {
17904 panic!("gemma4 layer {il} not full-attn")
17905 };
17906 let o = self.gemma4_attn(e, fa, il, &h, pos_d, t)?;
17907 let mut cur = e.zeros(t * n_embd)?;
17909 e.rms_norm(
17910 &o,
17911 layer.post_attn_norm.float_data(),
17912 &mut cur,
17913 n_embd,
17914 t,
17915 eps,
17916 )?;
17917 self.gemma4_layer_tail_add(e, layer, &cur, x, t)
17918 }
17919
17920 fn gemma4_layer_tail_add(
17924 &self,
17925 e: &Engine,
17926 layer: &crate::hybrid::HybridLayer,
17927 cur: &CudaSlice<f32>,
17928 x: &CudaSlice<f32>,
17929 t: usize,
17930 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
17931 Ok(self.gemma4_layer_tail_add_n(e, layer, cur, x, t, None)?.0)
17932 }
17933
17934 #[allow(clippy::type_complexity)] fn gemma4_layer_tail_add_n(
17938 &self,
17939 e: &Engine,
17940 layer: &crate::hybrid::HybridLayer,
17941 cur: &CudaSlice<f32>,
17942 x: &CudaSlice<f32>,
17943 t: usize,
17944 next_norm: Option<&CudaSlice<f32>>,
17945 ) -> Result<(CudaSlice<f32>, Option<CudaSlice<f32>>), Box<dyn std::error::Error>> {
17946 let n_embd = self.cfg.n_embd as usize;
17947 let bits = layer.gemma4.as_ref().unwrap();
17948 let (sn, attn_out) = self.gemma4_layer_tail_core(e, layer, cur, x, t)?;
17949 let mut xn = e.uninit(t * n_embd)?;
17950 match next_norm {
17951 Some(w) => {
17952 let mut hn = e.uninit(t * n_embd)?;
17953 e.add_scale_rms_norm(
17954 &sn,
17955 &attn_out,
17956 bits.layer_scale,
17957 w,
17958 &mut xn,
17959 &mut hn,
17960 n_embd,
17961 t,
17962 self.cfg.rms_eps,
17963 )?;
17964 Ok((xn, Some(hn)))
17965 }
17966 None => {
17967 e.add_scale(&sn, &attn_out, bits.layer_scale, &mut xn, t * n_embd)?;
17968 Ok((xn, None))
17969 }
17970 }
17971 }
17972
17973 fn gemma4_layer_tail_core(
17976 &self,
17977 e: &Engine,
17978 layer: &crate::hybrid::HybridLayer,
17979 cur: &CudaSlice<f32>,
17980 x: &CudaSlice<f32>,
17981 t: usize,
17982 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
17983 self.gemma4_layer_tail_core_pn(e, layer, cur, x, t, None, false)
17984 }
17985
17986 #[allow(clippy::too_many_arguments)] fn gemma4_layer_tail_core_pn(
17994 &self,
17995 e: &Engine,
17996 layer: &crate::hybrid::HybridLayer,
17997 cur: &CudaSlice<f32>,
17998 x: &CudaSlice<f32>,
17999 t: usize,
18000 pre_norm: Option<&CudaSlice<f32>>,
18001 defer_post_norm: bool,
18002 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18003 let n_embd = self.cfg.n_embd as usize;
18004 let eps = self.cfg.rms_eps;
18005 let bits = layer.gemma4.as_ref().unwrap();
18006
18007 let Some(mbits) = bits.moe_bits.as_ref() else {
18010 let crate::hybrid::Ffn::Dense {
18011 ffn_gate,
18012 ffn_up,
18013 ffn_down,
18014 } = &layer.ffn
18015 else {
18016 panic!("gemma4 dense layer without Dense ffn")
18017 };
18018 let mut attn_out = e.uninit(t * n_embd)?;
18019 let mut zsh = e.uninit(t * n_embd)?;
18020 let mut zpair: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
18023 match pre_norm {
18024 Some(wa) if t == 1 => {
18025 zpair = Some(e.rms_pre_add_rms_norm_q8z(
18026 cur,
18027 wa,
18028 x,
18029 bits.ffn_norm.float_data(),
18030 &mut attn_out,
18031 &mut zsh,
18032 n_embd,
18033 t,
18034 eps,
18035 )?);
18036 }
18037 Some(wa) => e.rms_pre_add_rms_norm(
18038 cur,
18039 wa,
18040 x,
18041 bits.ffn_norm.float_data(),
18042 &mut attn_out,
18043 &mut zsh,
18044 n_embd,
18045 t,
18046 eps,
18047 )?,
18048 None => e.add_rms_norm(
18049 cur,
18050 x,
18051 bits.ffn_norm.float_data(),
18052 &mut attn_out,
18053 &mut zsh,
18054 n_embd,
18055 t,
18056 eps,
18057 )?,
18058 }
18059 let n_ff = ffn_gate.out_features();
18060 let (gate, up) = if t == 1 {
18066 let (zq, zd) = match zpair {
18067 Some(p) => p,
18068 None => e.quantize_q8_1(&zsh, 1, n_embd)?,
18069 };
18070 match e.matmul_q4_fused2(ffn_gate, ffn_up, &zq, &zd)? {
18071 Some(p) => p,
18072 None => match e.matmul_nvfp4_fused2(ffn_gate, ffn_up, &zq, &zd, 1)? {
18074 Some(p) => p,
18075 None => (
18076 e.matmul_pre(ffn_gate, &zq, &zd, &zsh, 1)?,
18077 e.matmul_pre(ffn_up, &zq, &zd, &zsh, 1)?,
18078 ),
18079 },
18080 }
18081 } else {
18082 static F2B: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
18087 let f2b = *F2B.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
18088 let fused = if f2b {
18089 let (zq, zd) = e.quantize_q8_1(&zsh, t, n_embd)?;
18090 e.matmul_q4_fused2_batched(ffn_gate, ffn_up, &zq, &zd, t)?
18091 } else {
18092 None
18093 };
18094 match fused {
18095 Some(p) => p,
18096 None => {
18097 e.mmq_act_begin();
18099 (e.matmul(ffn_gate, &zsh, t)?, e.matmul(ffn_up, &zsh, t)?)
18100 }
18101 }
18102 };
18103 let mut act = e.uninit(t * n_ff)?;
18104 let f0 = if e.uses_q8_1_fast(ffn_down) {
18107 let upv = e.view(&up, t * n_ff);
18108 let up_all = upv.slice(0..t * n_ff);
18109 let (aq, ad) = e.gelu_tanh_mul_q8_1(&gate, &up_all, &mut act, n_ff, t)?;
18110 e.matmul_pre(ffn_down, &aq, &ad, &act, t)?
18111 } else {
18112 e.gelu_tanh_mul(&gate, &up, &mut act, t * n_ff)?;
18113 e.matmul(ffn_down, &act, t)?
18114 };
18115 if defer_post_norm {
18116 return Ok((f0, attn_out));
18117 }
18118 let mut sn = e.uninit(t * n_embd)?;
18119 e.rms_norm(
18120 &f0,
18121 bits.post_ffw_norm.float_data(),
18122 &mut sn,
18123 n_embd,
18124 t,
18125 eps,
18126 )?;
18127 return Ok((sn, attn_out));
18128 };
18129
18130 assert!(pre_norm.is_none(), "pre-norm fold is dense-entry only");
18131 let mut attn_out = e.uninit(t * n_embd)?;
18136 let mut router_in = e.uninit(t * n_embd)?;
18137 let fast_moe = match &layer.ffn {
18138 crate::hybrid::Ffn::Moe(m) => {
18139 m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
18140 && expert_dp4a_supported(m.gate_exps.qtype)
18141 && expert_dp4a_supported(m.up_exps.qtype)
18142 && expert_dp4a_supported(m.down_exps.qtype)
18143 && std::env::var("MEMRA_GEMMA_MOE_FAST").as_deref() != Ok("0")
18144 }
18145 _ => false,
18146 };
18147 let q8z = t < PRIME_MIN_T && fast_moe;
18148 let (zsh_f32, zsh_q8, moe_q8) = if q8z {
18149 let (z0, m2) = e.add_rms_norm3_q8z(
18150 cur,
18151 x,
18152 bits.ffn_norm.float_data(),
18153 &mbits.router_scale_pre,
18154 mbits.pre_ffw_norm_2.float_data(),
18155 &mut attn_out,
18156 &mut router_in,
18157 n_embd,
18158 t,
18159 eps,
18160 )?;
18161 (None, Some(z0), Some(m2))
18162 } else {
18163 let mut zsh = e.uninit(t * n_embd)?;
18164 let mut moe_in = e.uninit(t * n_embd)?;
18165 e.add_rms_norm3(
18166 cur,
18167 x,
18168 bits.ffn_norm.float_data(),
18169 &mbits.router_scale_pre,
18170 mbits.pre_ffw_norm_2.float_data(),
18171 &mut attn_out,
18172 &mut zsh,
18173 &mut router_in,
18174 &mut moe_in,
18175 n_embd,
18176 t,
18177 eps,
18178 )?;
18179 (Some((zsh, moe_in)), None, None)
18180 };
18181 let attn_out2 = attn_out;
18182 #[allow(unused_variables)]
18183 let attn_out = &attn_out2;
18184 let n_ff = mbits.shared_gate.out_features();
18185 let (gate, up) = if let Some((zq, zd)) = zsh_q8.as_ref() {
18186 if t == 1 {
18187 match e.matmul_q4_fused2(&mbits.shared_gate, &mbits.shared_up, zq, zd)? {
18188 Some(p) => p,
18189 None => match e.matmul_nvfp4_fused2(
18190 &mbits.shared_gate,
18191 &mbits.shared_up,
18192 zq,
18193 zd,
18194 1,
18195 )? {
18196 Some(p) => p,
18197 None => {
18198 let h0 = e.zeros(0)?;
18199 (
18200 e.matmul_pre(&mbits.shared_gate, zq, zd, &h0, 1)?,
18201 e.matmul_pre(&mbits.shared_up, zq, zd, &h0, 1)?,
18202 )
18203 }
18204 },
18205 }
18206 } else {
18207 let h0 = e.zeros(0)?;
18209 (
18210 e.matmul_pre(&mbits.shared_gate, zq, zd, &h0, t)?,
18211 e.matmul_pre(&mbits.shared_up, zq, zd, &h0, t)?,
18212 )
18213 }
18214 } else {
18215 let (zsh, _) = zsh_f32.as_ref().unwrap();
18216 (
18217 e.matmul(&mbits.shared_gate, zsh, t)?,
18218 e.matmul(&mbits.shared_up, zsh, t)?,
18219 )
18220 };
18221 let mut act = e.uninit(t * n_ff)?;
18222 e.gelu_tanh_mul(&gate, &up, &mut act, t * n_ff)?;
18223 let mlp0 = e.matmul(&mbits.shared_down, &act, t)?;
18224 let crate::hybrid::Ffn::Moe(m) = &layer.ffn else {
18225 panic!("gemma4 layer not MoE")
18226 };
18227 let moe0 = match (&moe_q8, &zsh_f32) {
18228 (Some(mq), _) => self.gemma4_moe_q8(e, m, mbits, mq, &router_in, t)?,
18229 (None, Some((_, moe_in))) => self.gemma4_moe(e, m, mbits, moe_in, &router_in, t)?,
18230 _ => unreachable!(),
18231 };
18232 let mut mlp = e.uninit(t * n_embd)?;
18234 let mut moe = e.uninit(t * n_embd)?;
18235 e.rms_norm2x(
18236 &mlp0,
18237 &moe0,
18238 mbits.post_ffw_norm_1.float_data(),
18239 mbits.post_ffw_norm_2.float_data(),
18240 &mut mlp,
18241 &mut moe,
18242 n_embd,
18243 t,
18244 eps,
18245 )?;
18246
18247 let mut sum = e.uninit(t * n_embd)?;
18250 let mut sn = e.uninit(t * n_embd)?;
18251 e.add_rms_norm(
18252 &mlp,
18253 &moe,
18254 bits.post_ffw_norm.float_data(),
18255 &mut sum,
18256 &mut sn,
18257 n_embd,
18258 t,
18259 eps,
18260 )?;
18261 Ok((sn, attn_out2))
18262 }
18263
18264 #[allow(clippy::type_complexity)] pub(crate) fn gemma4_layer_tail_add_nq_pn(
18275 &self,
18276 e: &Engine,
18277 layer: &crate::hybrid::HybridLayer,
18278 o: &CudaSlice<f32>,
18279 x: &CudaSlice<f32>,
18280 t: usize,
18281 next_norm: Option<&CudaSlice<f32>>,
18282 ) -> Result<(CudaSlice<f32>, Option<(CudaSlice<i8>, CudaSlice<f32>)>), Box<dyn std::error::Error>>
18283 {
18284 let n_embd = self.cfg.n_embd as usize;
18285 let eps = self.cfg.rms_eps;
18286 let bits = layer.gemma4.as_ref().unwrap();
18287 if Engine::g4_pnfold_on() && matches!(layer.ffn, crate::hybrid::Ffn::Dense { .. }) {
18288 let (f0, attn_out) = self.gemma4_layer_tail_core_pn(
18289 e,
18290 layer,
18291 o,
18292 x,
18293 t,
18294 Some(layer.post_attn_norm.float_data()),
18295 true,
18296 )?;
18297 let mut xn = e.uninit(t * n_embd)?;
18298 return match next_norm {
18299 Some(w) => {
18300 let pair = e.rms_pre_add_scale_rms_norm_q8_1(
18301 &f0,
18302 bits.post_ffw_norm.float_data(),
18303 &attn_out,
18304 bits.layer_scale,
18305 w,
18306 &mut xn,
18307 n_embd,
18308 t,
18309 eps,
18310 )?;
18311 Ok((xn, Some(pair)))
18312 }
18313 None => {
18314 let mut sn = e.uninit(t * n_embd)?;
18315 e.rms_norm(
18316 &f0,
18317 bits.post_ffw_norm.float_data(),
18318 &mut sn,
18319 n_embd,
18320 t,
18321 eps,
18322 )?;
18323 e.add_scale(&sn, &attn_out, bits.layer_scale, &mut xn, t * n_embd)?;
18324 Ok((xn, None))
18325 }
18326 };
18327 }
18328 let mut cur = e.uninit(t * n_embd)?;
18329 e.rms_norm(
18330 o,
18331 layer.post_attn_norm.float_data(),
18332 &mut cur,
18333 n_embd,
18334 t,
18335 eps,
18336 )?;
18337 self.gemma4_layer_tail_add_nq(e, layer, &cur, x, t, next_norm)
18338 }
18339
18340 #[allow(clippy::type_complexity)] pub(crate) fn gemma4_layer_tail_add_nq(
18342 &self,
18343 e: &Engine,
18344 layer: &crate::hybrid::HybridLayer,
18345 cur: &CudaSlice<f32>,
18346 x: &CudaSlice<f32>,
18347 t: usize,
18348 next_norm: Option<&CudaSlice<f32>>,
18349 ) -> Result<(CudaSlice<f32>, Option<(CudaSlice<i8>, CudaSlice<f32>)>), Box<dyn std::error::Error>>
18350 {
18351 let n_embd = self.cfg.n_embd as usize;
18352 let bits = layer.gemma4.as_ref().unwrap();
18353 let (sn, attn_out) = self.gemma4_layer_tail_core(e, layer, cur, x, t)?;
18354 let mut xn = e.uninit(t * n_embd)?;
18355 match next_norm {
18356 Some(w) => {
18357 let pair = e.add_scale_rms_norm_q8_1(
18358 &sn,
18359 &attn_out,
18360 bits.layer_scale,
18361 w,
18362 &mut xn,
18363 n_embd,
18364 t,
18365 self.cfg.rms_eps,
18366 )?;
18367 Ok((xn, Some(pair)))
18368 }
18369 None => {
18370 e.add_scale(&sn, &attn_out, bits.layer_scale, &mut xn, t * n_embd)?;
18371 Ok((xn, None))
18372 }
18373 }
18374 }
18375
18376 fn gemma4_forward(
18379 &self,
18380 e: &Engine,
18381 tokens: &[u32],
18382 last_only: bool,
18383 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
18384 if self.is_gemma4_e4b() {
18387 return self.gemma4_e4b_forward(e, tokens, last_only);
18388 }
18389 let n_embd = self.cfg.n_embd as usize;
18390 let t = tokens.len();
18391 let pos: Vec<i32> = (0..t as i32).collect();
18392 let pos_d = e.htod_i32(&pos)?;
18393
18394 let mut x = self.embed(e, tokens)?;
18395 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
18396 let probe = std::env::var("MEMRA_GEMMA_PROBE").is_ok();
18399 let stat =
18400 |e: &Engine, x: &CudaSlice<f32>, tag: &str| -> Result<(), Box<dyn std::error::Error>> {
18401 let h = e.dtoh(x)?;
18402 let bad = h.iter().filter(|v| !v.is_finite()).count();
18403 let mx = h
18404 .iter()
18405 .filter(|v| v.is_finite())
18406 .fold(0.0f32, |m, v| m.max(v.abs()));
18407 eprintln!(
18408 "[gemma-probe] {tag}: tok0_first3={:?} bad={bad} max={mx:.3e}",
18409 &h[..3]
18410 );
18411 Ok(())
18412 };
18413 if probe {
18414 stat(e, &x, "embed")?;
18415 }
18416 for (il, layer) in self.layers.iter().enumerate() {
18417 x = self.gemma4_layer(e, il, layer, &x, &pos_d, t)?;
18418 if probe {
18419 stat(e, &x, &format!("L{il}"))?;
18420 }
18421 }
18422 let mut hn = e.zeros(t * n_embd)?;
18423 e.rms_norm(
18424 &x,
18425 self.output_norm.float_data(),
18426 &mut hn,
18427 n_embd,
18428 t,
18429 self.cfg.rms_eps,
18430 )?;
18431 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
18432 let n_vocab = self.output.out_features();
18433 let logits = if last_only {
18434 let hv = e.view(&hn, t * n_embd);
18435 let last_row = hv.slice((t - 1) * n_embd..t * n_embd);
18436 let mut hlast = e.zeros(n_embd)?;
18437 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
18438 let mut ld = e.matmul(&self.output, &hlast, 1)?;
18439 e.softcap(&mut ld, cap, n_vocab)?;
18440 self.gemma4_suppress(e, &mut ld, 1)?;
18441 e.dtoh(&ld)?
18442 } else {
18443 let mut ld = e.matmul(&self.output, &hn, t)?;
18444 e.softcap(&mut ld, cap, t * n_vocab)?;
18445 self.gemma4_suppress(e, &mut ld, t)?;
18446 e.dtoh(&ld)?
18447 };
18448 Ok(logits)
18449 }
18450
18451 #[allow(clippy::type_complexity)] pub(crate) fn gemma4_prime(
18457 &self,
18458 e: &Engine,
18459 tokens: &[u32],
18460 cache: &mut Cache,
18461 overlay: Option<&crate::vision::EmbedOverlay>,
18462 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
18463 if cache.pos != 0 {
18468 return Err(
18469 "gemma4 prime v0 is fresh-prompt only (no continuation/chunked prime) \
18470 — prime the full prompt in one call or decode tokenwise"
18471 .into(),
18472 );
18473 }
18474 let n_embd = self.cfg.n_embd as usize;
18475 let eps = self.cfg.rms_eps;
18476 let t = tokens.len();
18477 let pos: Vec<i32> = (0..t as i32).collect();
18478 let pos_d = e.htod_i32(&pos)?;
18479 let mut x = self.embed(e, tokens)?;
18480 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
18481 let island: Option<CudaSlice<i32>> = match overlay {
18488 Some(ov) => {
18489 ov.require_resident(e)?;
18493 let mut span_id = vec![-1i32; t];
18494 for (i, &(pos, row_off, n_rows)) in ov.spans.iter().enumerate() {
18495 if pos + n_rows > t {
18496 return Err(format!(
18497 "gemma4 overlay span {i} [{pos}, {}) exceeds the prompt ({t})",
18498 pos + n_rows
18499 )
18500 .into());
18501 }
18502 let view = ov.rows.slice(row_off * n_embd..(row_off + n_rows) * n_embd);
18503 e.copy_view_into(&mut x, pos * n_embd, &view, n_rows * n_embd)?;
18504 for s in span_id.iter_mut().skip(pos).take(n_rows) {
18505 *s = i as i32;
18506 }
18507 }
18508 if std::env::var("MEMRA_GV_FORCE_CAUSAL").as_deref() == Ok("1") {
18512 None
18513 } else {
18514 Some(e.htod_i32(&span_id)?)
18515 }
18516 }
18517 None => None,
18518 };
18519 for (il, layer) in self.layers.iter().enumerate() {
18520 let mut h = e.zeros(t * n_embd)?;
18521 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
18522 let Mixer::Full(fa) = &layer.mixer else {
18523 panic!("gemma4 layer not full-attn")
18524 };
18525 let trace = il == 0 && std::env::var("MEMRA_G4_PRIME_TRACE").as_deref() == Ok("1");
18526 if trace {
18527 let v = e.dtoh(&h)?;
18528 let nan = v.iter().filter(|x| x.is_nan()).count();
18529 eprintln!("[g4-prime-trace] L0 post-attn_norm: nan={nan}/{}", v.len());
18530 }
18531 let o =
18532 self.gemma4_attn_prime(e, fa, il, &h, &pos_d, t, Some(cache), island.as_ref())?;
18533 if trace {
18534 let v = e.dtoh(&o)?;
18535 let nan = v.iter().filter(|x| x.is_nan()).count();
18536 eprintln!("[g4-prime-trace] L0 post-attn: nan={nan}/{}", v.len());
18537 }
18538 let mut cur = e.zeros(t * n_embd)?;
18539 e.rms_norm(
18540 &o,
18541 layer.post_attn_norm.float_data(),
18542 &mut cur,
18543 n_embd,
18544 t,
18545 eps,
18546 )?;
18547 x = self.gemma4_layer_tail_add(e, layer, &cur, &x, t)?;
18548 self.dflash_tap(e, cache, il, &x, t)?;
18549 if std::env::var("MEMRA_G4_PRIME_TRACE").as_deref() == Ok("1") {
18551 let h = e.dtoh(&x)?;
18552 let nan = h.iter().filter(|v| v.is_nan()).count();
18553 let amax = h.iter().fold(0f32, |a, v| a.max(v.abs()));
18554 eprintln!(
18555 "[g4-prime-trace] layer {il}: nan={nan}/{} amax={amax:.3}",
18556 h.len()
18557 );
18558 if nan > 0 {
18559 return Err(format!("g4-prime-trace: first NaN at layer {il}").into());
18560 }
18561 }
18562 }
18563 cache.pos += t;
18564 let hiddens = e.clone_dtod(&x)?;
18565 let xv = e.view(&x, t * n_embd);
18566 let last_row = xv.slice((t - 1) * n_embd..t * n_embd);
18567 let mut h_seed = e.zeros(n_embd)?;
18568 e.copy_view_into(&mut h_seed, 0, &last_row, n_embd)?;
18569 let mut hn = e.uninit(n_embd)?;
18570 e.rms_norm(
18571 &h_seed,
18572 self.output_norm.float_data(),
18573 &mut hn,
18574 n_embd,
18575 1,
18576 eps,
18577 )?;
18578 let mut ld = e.matmul(&self.output, &hn, 1)?;
18579 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
18580 e.softcap(&mut ld, cap, self.output.out_features())?;
18581 self.gemma4_suppress(e, &mut ld, 1)?;
18582 let logits = e.dtoh(&ld)?;
18583 Ok((logits, h_seed, hiddens))
18584 }
18585
18586 #[allow(clippy::too_many_arguments)] fn gemma4_decode_attn(
18592 &self,
18593 e: &Engine,
18594 fa: &crate::hybrid::FullAttnLayer,
18595 il: usize,
18596 hq: &CudaSlice<i8>,
18597 hdq: &CudaSlice<f32>,
18598 pos_d: &CudaSlice<i32>,
18599 cache: &mut Cache,
18600 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
18601 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
18602 let eps = self.cfg.rms_eps;
18603 let aux = self.gemma4_aux.as_ref().unwrap();
18604 let ones = aux.ones(e);
18605 #[cfg(debug_assertions)]
18606 crate::debug_assert_tensor_stream_device(ones, &e.stream(), "gemma4_decode_attn.ones");
18607 let (hq, hdq) = (hq, hdq);
18608 let h0 = e.zeros(0)?;
18609 let h = &h0;
18610 let (q0, k0, v0) = if swa {
18611 match e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)? {
18612 Some(t3) => t3,
18613 None => match e.matmul_nvfp4_fused2(&fa.wq, &fa.wk, hq, hdq, 1)? {
18616 Some((q0, k0)) => {
18617 let v0 = e.matmul_pre(&fa.wv, hq, hdq, h, 1)?;
18618 (q0, k0, v0)
18619 }
18620 None => (
18621 e.matmul_pre(&fa.wq, hq, hdq, h, 1)?,
18622 e.matmul_pre(&fa.wk, hq, hdq, h, 1)?,
18623 e.matmul_pre(&fa.wv, hq, hdq, h, 1)?,
18624 ),
18625 },
18626 }
18627 } else {
18628 let (q0, k0) = match e.matmul_q4_fused2(&fa.wq, &fa.wk, hq, hdq)? {
18629 Some(p) => p,
18630 None => match e.matmul_nvfp4_fused2(&fa.wq, &fa.wk, hq, hdq, 1)? {
18631 Some(p) => p,
18632 None => (
18633 e.matmul_pre(&fa.wq, hq, hdq, h, 1)?,
18634 e.matmul_pre(&fa.wk, hq, hdq, h, 1)?,
18635 ),
18636 },
18637 };
18638 let v0 = e.clone_dtod(&k0)?;
18639 (q0, k0, v0)
18640 };
18641 let mut q = e.uninit(nh * hd)?;
18642 let mut k = e.uninit(nkv * hd)?;
18643 let mut v = e.uninit(nkv * hd)?;
18644 let ff = if swa {
18647 None
18648 } else {
18649 Some(
18650 aux.rope_freqs(e)
18651 .expect("gemma4 global rope needs rope_freqs.weight"),
18652 )
18653 };
18654 #[cfg(debug_assertions)]
18655 if let Some(ff) = ff {
18656 crate::debug_assert_tensor_stream_device(
18657 ff,
18658 &e.stream(),
18659 "gemma4_decode_attn.rope_freqs",
18660 );
18661 }
18662 let kvl = cache.kv[il].as_mut().unwrap();
18663 let kv_fp8 = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
18664 if crate::Engine::qkv_append_on() {
18665 e.rms_norm_qkv_rope_append(
18669 &q0,
18670 &k0,
18671 &v0,
18672 fa.q_norm.float_data(),
18673 fa.k_norm.float_data(),
18674 ones,
18675 &mut q,
18676 &mut k,
18677 &mut v,
18678 hd,
18679 self.gemma4_rope_dims(il),
18680 nh,
18681 nkv,
18682 pos_d,
18683 nh,
18684 nkv,
18685 base,
18686 1.0,
18687 ff,
18688 eps,
18689 &mut kvl.k,
18690 &mut kvl.v,
18691 kvl.len,
18692 kvl.k_tok_bytes,
18693 kvl.v_tok_bytes,
18694 kv_fp8,
18695 )?;
18696 } else {
18697 e.rms_norm_qkv_rope(
18698 &q0,
18699 &k0,
18700 &v0,
18701 fa.q_norm.float_data(),
18702 fa.k_norm.float_data(),
18703 ones,
18704 &mut q,
18705 &mut k,
18706 &mut v,
18707 hd,
18708 self.gemma4_rope_dims(il),
18709 nh,
18710 nkv,
18711 pos_d,
18712 nh,
18713 nkv,
18714 base,
18715 1.0,
18716 ff,
18717 eps,
18718 )?;
18719 e.append_kv_quantized(
18720 &k,
18721 &v,
18722 &mut kvl.k,
18723 &mut kvl.v,
18724 kvl.len,
18725 kvl.kv_dim_k,
18726 kvl.kv_dim_v,
18727 kvl.k_tok_bytes,
18728 kvl.v_tok_bytes,
18729 kv_fp8,
18730 )?;
18731 }
18732 kvl.len += 1;
18733 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
18737 let mut attn = e.uninit(nh * hd)?;
18738 if !swa
18740 && hd == 512
18741 && kvl.len >= crate::fa512_min_tkv()
18742 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
18743 {
18744 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
18745 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
18746 let base = kvl.len as i32;
18748 e.i32_set_k(&mut kvl.len_d, base)?;
18749 e.fa_decode_rows(
18750 &q,
18751 &kp,
18752 &vp,
18753 &mut attn,
18754 hd,
18755 nh,
18756 nkv,
18757 kvl.len - 1,
18758 1,
18759 scale,
18760 kvl.k_tok_bytes,
18761 kvl.v_tok_bytes,
18762 Some((&kvl.len_d, -1)),
18763 false,
18764 false,
18765 None,
18766 )?;
18767 return e.matmul(&fa.wo, &attn, 1);
18768 }
18769 if swa
18771 && kvl.len > win
18772 && hd == 256
18773 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
18774 {
18775 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
18776 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
18777 let base = kvl.len as i32;
18778 e.i32_set_k(&mut kvl.len_d, base)?;
18779 e.fa_decode_rows_w(
18780 &q,
18781 &kp,
18782 &vp,
18783 &mut attn,
18784 hd,
18785 nh,
18786 nkv,
18787 &kvl.len_d,
18788 -1,
18789 1,
18790 scale,
18791 win,
18792 kvl.k_tok_bytes,
18793 kvl.v_tok_bytes,
18794 None,
18795 )?;
18796 return e.matmul(&fa.wo, &attn, 1);
18797 }
18798 let (off_tok, t_kv) = if swa && kvl.len > win {
18799 (kvl.len - win, win)
18800 } else {
18801 (0, kvl.len)
18802 };
18803 let k_view = e.view_u8_range(
18804 &kvl.k,
18805 off_tok * kvl.k_tok_bytes,
18806 (off_tok + t_kv) * kvl.k_tok_bytes,
18807 );
18808 let v_view = e.view_u8_range(
18809 &kvl.v,
18810 off_tok * kvl.v_tok_bytes,
18811 (off_tok + t_kv) * kvl.v_tok_bytes,
18812 );
18813 e.fa_decode_kvmod(
18814 &q,
18815 &k_view,
18816 &v_view,
18817 &mut attn,
18818 hd,
18819 nh,
18820 nkv,
18821 t_kv,
18822 scale,
18823 kvl.k_tok_bytes,
18824 kvl.v_tok_bytes,
18825 swa && crate::Engine::wkv_on(),
18826 )?;
18827 e.matmul(&fa.wo, &attn, 1)
18828 }
18829
18830 #[allow(clippy::too_many_arguments)]
18837 pub fn gemma4_decode_step_dc(
18838 &self,
18839 e: &Engine,
18840 token_d: &CudaSlice<u32>,
18841 pos_d: &mut CudaSlice<i32>,
18842 embd_gpu: &CudaSlice<u8>,
18843 embd_qt: i32,
18844 embd_rb: usize,
18845 cache: &mut Cache,
18846 n_vocab: usize,
18847 cap_bucket_max: Option<(usize, usize)>,
18848 ) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
18849 let mut tok_out = e.stream().alloc_zeros::<u32>(1)?;
18850 self.gemma4_decode_step_dc_into(
18851 e,
18852 token_d,
18853 pos_d,
18854 embd_gpu,
18855 embd_qt,
18856 embd_rb,
18857 cache,
18858 n_vocab,
18859 cap_bucket_max,
18860 &mut tok_out,
18861 )?;
18862 Ok(tok_out)
18863 }
18864
18865 #[allow(clippy::too_many_arguments)]
18868 pub fn gemma4_decode_step_dc_into(
18869 &self,
18870 e: &Engine,
18871 token_d: &CudaSlice<u32>,
18872 pos_d: &mut CudaSlice<i32>,
18873 embd_gpu: &CudaSlice<u8>,
18874 embd_qt: i32,
18875 embd_rb: usize,
18876 cache: &mut Cache,
18877 n_vocab: usize,
18878 cap_bucket_max: Option<(usize, usize)>,
18879 tok_out: &mut CudaSlice<u32>,
18880 ) -> Result<(), Box<dyn std::error::Error>> {
18881 let n_embd = self.cfg.n_embd as usize;
18882 let eps = self.cfg.rms_eps;
18883 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
18884 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
18885 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
18886 let n_layers = self.layers.len();
18887 for (il, layer) in self.layers.iter().enumerate() {
18888 let (hq, hdq) = match h_carry.take() {
18889 Some(p) => p,
18890 None => {
18891 e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, 1, eps)?
18892 }
18893 };
18894 let Mixer::Full(fa) = &layer.mixer else {
18895 panic!("gemma4 layer {il} not full-attn")
18896 };
18897 let o =
18898 self.gemma4_decode_attn_dc(e, fa, il, &hq, &hdq, pos_d, cache, cap_bucket_max)?;
18899 let next_norm = if il + 1 < n_layers {
18900 Some(self.layers[il + 1].attn_norm.float_data())
18901 } else {
18902 None
18903 };
18904 let (xn, hn) = self.gemma4_layer_tail_add_nq_pn(e, layer, &o, &x, 1, next_norm)?;
18905 x = xn;
18906 h_carry = hn;
18907 }
18908 let mut hn = e.uninit(n_embd)?;
18909 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
18910 let mut logits = e.matmul(&self.output, &hn, 1)?;
18911 self.gemma4_suppress(e, &mut logits, 1)?; e.argmax_token_device_into(&logits, tok_out, n_vocab)?;
18913 e.inc_seqlen(pos_d)?;
18914 if cap_bucket_max.is_none() {
18915 cache.pos += 1;
18916 }
18917 Ok(())
18918 }
18919
18920 pub fn g4_dc_slots(&self, e: &Engine) -> Result<G4DcSlots, Box<dyn std::error::Error>> {
18927 let n_embd = self.cfg.n_embd as usize;
18928 let n_vocab = self.output.out_features();
18929 let n_layers = self.layers.len();
18930 let (mut qmax, mut kvmax, mut ffmax) = (0usize, 0usize, 0usize);
18931 for il in 0..n_layers {
18932 let (hd, nkv, nh, _b, _s, _w) = self.gemma4_geom(il);
18933 qmax = qmax.max(nh * hd);
18934 kvmax = kvmax.max(nkv * hd);
18935 if let crate::hybrid::Ffn::Dense { ffn_gate, .. } = &self.layers[il].ffn {
18936 ffmax = ffmax.max(ffn_gate.out_features());
18937 }
18938 }
18939 Ok(G4DcSlots {
18940 x: e.uninit(n_embd)?,
18941 xn: e.uninit(n_embd)?,
18942 cur: e.uninit(n_embd)?,
18943 hq: e.alloc_i8_uninit(n_embd)?,
18944 hd_: e.uninit(n_embd / 32)?,
18945 q0: e.uninit(qmax)?,
18946 k0: e.uninit(kvmax)?,
18947 v0: e.uninit(kvmax)?,
18948 q: e.uninit(qmax)?,
18949 k: e.uninit(kvmax)?,
18950 v: e.uninit(kvmax)?,
18951 attn: e.uninit(qmax)?,
18952 o: e.uninit(n_embd)?,
18953 attn_out: e.uninit(n_embd)?,
18954 zsh: e.uninit(n_embd)?,
18955 zq: e.alloc_i8_uninit(n_embd.max(qmax))?,
18958 zd: e.uninit(n_embd.max(qmax) / 32)?,
18959 gate: e.uninit(ffmax)?,
18960 up: e.uninit(ffmax)?,
18961 act: e.uninit(ffmax)?,
18962 actq: e.alloc_i8_uninit(ffmax)?,
18963 actd: e.uninit(ffmax / 32)?,
18964 f0: e.uninit(n_embd)?,
18965 sn: e.uninit(n_embd)?,
18966 hn: e.uninit(n_embd)?,
18967 logits: e.uninit(n_vocab)?,
18968 })
18969 }
18970
18971 fn g4_matvec_m1_into(
18974 &self,
18975 e: &Engine,
18976 w: &crate::model::GpuTensor,
18977 aq: &CudaSlice<i8>,
18978 ad: &CudaSlice<f32>,
18979 y: &mut CudaSlice<f32>,
18980 ) -> Result<(), Box<dyn std::error::Error>> {
18981 use crate::model::GpuTensor;
18982 let (bytes, qtype, row_bytes, scale, rp) = match w {
18983 GpuTensor::Quant {
18984 bytes,
18985 qtype,
18986 row_bytes,
18987 scale,
18988 rp,
18989 ..
18990 } => (bytes, *qtype, *row_bytes, *scale, *rp),
18991 _ => return Err("g4_matvec_m1_into: non-quant tensor".into()),
18992 };
18993 let (mbytes, mrp) = match w {
18994 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
18995 _ => (bytes, rp),
18996 };
18997 e.qmatvec_mmvq_into(
18998 mbytes,
18999 aq,
19000 ad,
19001 1,
19002 w.in_features(),
19003 w.out_features(),
19004 qtype,
19005 row_bytes,
19006 scale,
19007 mrp,
19008 y,
19009 )
19010 }
19011
19012 #[allow(clippy::too_many_arguments)]
19016 pub fn gemma4_decode_step_dc_slotted(
19017 &self,
19018 e: &Engine,
19019 token_d: &CudaSlice<u32>,
19020 pos_d: &mut CudaSlice<i32>,
19021 embd_gpu: &CudaSlice<u8>,
19022 embd_qt: i32,
19023 embd_rb: usize,
19024 cache: &mut Cache,
19025 n_vocab: usize,
19026 cap_bucket_max: Option<(usize, usize)>,
19027 sl: &mut G4DcSlots,
19028 tok_out: &mut CudaSlice<u32>,
19029 ring: Option<(&mut CudaSlice<u32>, usize)>,
19030 ) -> Result<(), Box<dyn std::error::Error>> {
19031 let n_embd = self.cfg.n_embd as usize;
19032 let eps = self.cfg.rms_eps;
19033 e.embed_gather_device_into(embd_gpu, token_d, &mut sl.x, n_embd, embd_qt, embd_rb)?;
19034 e.scale_inplace(&mut sl.x, (n_embd as f32).sqrt(), n_embd)?;
19035 let n_layers = self.layers.len();
19036 let mut has_carry = false;
19037 for il in 0..n_layers {
19038 if !has_carry {
19039 e.rms_norm_q8_1_into(
19040 &sl.x,
19041 self.layers[il].attn_norm.float_data(),
19042 n_embd,
19043 1,
19044 eps,
19045 &mut sl.hq,
19046 &mut sl.hd_,
19047 )?;
19048 }
19049 has_carry = true;
19050 let layer = &self.layers[il];
19051 let Mixer::Full(fa) = &layer.mixer else {
19052 panic!("gemma4 layer {il} not full-attn")
19053 };
19054 self.gemma4_decode_attn_dc_slotted(e, fa, il, pos_d, cache, cap_bucket_max, sl)?;
19055 if !Engine::g4_pnfold_on() {
19058 e.rms_norm(
19059 &sl.o,
19060 layer.post_attn_norm.float_data(),
19061 &mut sl.cur,
19062 n_embd,
19063 1,
19064 eps,
19065 )?;
19066 }
19067 let next_norm = if il + 1 < n_layers {
19068 Some(self.layers[il + 1].attn_norm.float_data())
19069 } else {
19070 None
19071 };
19072 self.gemma4_layer_tail_slotted(e, layer, next_norm, sl)?;
19073 std::mem::swap(&mut sl.x, &mut sl.xn);
19074 }
19075 e.rms_norm(
19076 &sl.x,
19077 self.output_norm.float_data(),
19078 &mut sl.hn,
19079 n_embd,
19080 1,
19081 eps,
19082 )?;
19083 e.quantize_q8_1_into(&sl.hn, 1, n_embd, &mut sl.zq, &mut sl.zd)?;
19084 {
19086 let (zq, zd) = (&sl.zq, &sl.zd);
19087 let zq = unsafe { &*(zq as *const CudaSlice<i8>) };
19088 let zd = unsafe { &*(zd as *const CudaSlice<f32>) };
19089 self.g4_matvec_m1_into(e, &self.output, zq, zd, &mut sl.logits)?;
19090 }
19091 self.gemma4_suppress(e, &mut sl.logits, 1)?;
19092 e.argmax_token_device_into(&sl.logits, tok_out, n_vocab)?;
19093 if let Some((ring, base)) = ring {
19094 e.plain_tok_ring(tok_out, pos_d, base, ring)?;
19098 }
19099 e.inc_seqlen(pos_d)?;
19100 if cap_bucket_max.is_none() {
19101 cache.pos += 1;
19102 }
19103 Ok(())
19104 }
19105
19106 #[allow(clippy::too_many_arguments)]
19108 fn gemma4_decode_attn_dc_slotted(
19109 &self,
19110 e: &Engine,
19111 fa: &crate::hybrid::FullAttnLayer,
19112 il: usize,
19113 pos_d: &CudaSlice<i32>,
19114 cache: &mut Cache,
19115 cap_bucket_max: Option<(usize, usize)>,
19116 sl: &mut G4DcSlots,
19117 ) -> Result<(), Box<dyn std::error::Error>> {
19118 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
19119 let eps = self.cfg.rms_eps;
19120 let aux = self.gemma4_aux.as_ref().unwrap();
19121 let ones = aux.ones(e);
19122 #[cfg(debug_assertions)]
19123 crate::debug_assert_tensor_stream_device(
19124 ones,
19125 &e.stream(),
19126 "gemma4_decode_attn_dc_slotted.ones",
19127 );
19128 {
19129 let hq = unsafe { &*(&sl.hq as *const CudaSlice<i8>) };
19130 let hdq = unsafe { &*(&sl.hd_ as *const CudaSlice<f32>) };
19131 if swa {
19132 if !e.matmul_q4_fused3_into(
19133 &fa.wq, &fa.wk, &fa.wv, hq, hdq, &mut sl.q0, &mut sl.k0, &mut sl.v0,
19134 )? {
19135 if e.matmul_nvfp4_fused2_into(&fa.wq, &fa.wk, hq, hdq, &mut sl.q0, &mut sl.k0)?
19139 {
19140 self.g4_matvec_m1_into(e, &fa.wv, hq, hdq, &mut sl.v0)?;
19141 } else {
19142 return Err("slotted step: fused3 unavailable (non-uniform trunk)".into());
19143 }
19144 }
19145 } else {
19146 if !e.matmul_q4_fused2_into(&fa.wq, &fa.wk, hq, hdq, &mut sl.q0, &mut sl.k0)?
19147 && !e
19148 .matmul_nvfp4_fused2_into(&fa.wq, &fa.wk, hq, hdq, &mut sl.q0, &mut sl.k0)?
19149 {
19150 return Err("slotted step: fused2 unavailable".into());
19151 }
19152 let k0r = unsafe { &*(&sl.k0 as *const CudaSlice<f32>) };
19153 e.copy_into(&mut sl.v0, 0, k0r, nkv * hd)?;
19154 }
19155 }
19156 let ff = if swa {
19159 None
19160 } else {
19161 Some(
19162 aux.rope_freqs(e)
19163 .expect("gemma4 global rope needs rope_freqs.weight"),
19164 )
19165 };
19166 #[cfg(debug_assertions)]
19167 if let Some(ff) = ff {
19168 crate::debug_assert_tensor_stream_device(
19169 ff,
19170 &e.stream(),
19171 "gemma4_decode_attn_dc_slotted.rope_freqs",
19172 );
19173 }
19174 let kvl = cache.kv[il].as_mut().unwrap();
19175 let kv_fp8 = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
19176 if crate::Engine::qkv_append_on() {
19177 e.rms_norm_qkv_rope_append_dc(
19179 &sl.q0,
19180 &sl.k0,
19181 &sl.v0,
19182 fa.q_norm.float_data(),
19183 fa.k_norm.float_data(),
19184 ones,
19185 &mut sl.q,
19186 &mut sl.k,
19187 &mut sl.v,
19188 hd,
19189 self.gemma4_rope_dims(il),
19190 nh,
19191 nkv,
19192 pos_d,
19193 nh,
19194 nkv,
19195 base,
19196 1.0,
19197 ff,
19198 eps,
19199 &mut kvl.k,
19200 &mut kvl.v,
19201 &kvl.len_d,
19202 kvl.k_tok_bytes,
19203 kvl.v_tok_bytes,
19204 kv_fp8,
19205 )?;
19206 } else {
19207 e.rms_norm_qkv_rope(
19208 &sl.q0,
19209 &sl.k0,
19210 &sl.v0,
19211 fa.q_norm.float_data(),
19212 fa.k_norm.float_data(),
19213 ones,
19214 &mut sl.q,
19215 &mut sl.k,
19216 &mut sl.v,
19217 hd,
19218 self.gemma4_rope_dims(il),
19219 nh,
19220 nkv,
19221 pos_d,
19222 nh,
19223 nkv,
19224 base,
19225 1.0,
19226 ff,
19227 eps,
19228 )?;
19229 e.append_kv_quantized_dc(
19230 &sl.k,
19231 &sl.v,
19232 &mut kvl.k,
19233 &mut kvl.v,
19234 &kvl.len_d,
19235 kvl.kv_dim_k,
19236 kvl.kv_dim_v,
19237 kvl.k_tok_bytes,
19238 kvl.v_tok_bytes,
19239 kv_fp8,
19240 )?;
19241 }
19242 e.inc_seqlen(&mut kvl.len_d)?;
19243 let (b_swa, b_glob) = cap_bucket_max.expect("slotted step is capture-only");
19244 let k_view = e.view_u8(&kvl.k, kvl.k.len());
19245 let v_view = e.view_u8(&kvl.v, kvl.v.len());
19246 let rows_on = std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0");
19247 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
19248 let mut fa_q8 = false;
19252 if !swa && hd == 512 && b_glob >= crate::fa512_min_tkv() && rows_on {
19253 e.fa_decode_rows(
19254 &sl.q,
19255 &k_view,
19256 &v_view,
19257 &mut sl.attn,
19258 hd,
19259 nh,
19260 nkv,
19261 b_glob - 1,
19262 1,
19263 scale,
19264 kvl.k_tok_bytes,
19265 kvl.v_tok_bytes,
19266 Some((&kvl.len_d, -1)),
19267 false,
19268 false,
19269 Some((&mut sl.zq, &mut sl.zd)),
19270 )?;
19271 fa_q8 = true;
19272 } else if swa && b_swa > win && hd == 256 && rows_on {
19273 e.fa_decode_rows_w(
19274 &sl.q,
19275 &k_view,
19276 &v_view,
19277 &mut sl.attn,
19278 hd,
19279 nh,
19280 nkv,
19281 &kvl.len_d,
19282 -1,
19283 1,
19284 scale,
19285 win,
19286 kvl.k_tok_bytes,
19287 kvl.v_tok_bytes,
19288 Some((&mut sl.zq, &mut sl.zd)),
19289 )?;
19290 fa_q8 = true;
19291 } else {
19292 let b = if swa { b_swa } else { b_glob };
19293 e.fa_decode_dc(
19294 &sl.q,
19295 &k_view,
19296 &v_view,
19297 &mut sl.attn,
19298 hd,
19299 nh,
19300 nkv,
19301 &kvl.len_d,
19302 b,
19303 scale,
19304 kvl.k_tok_bytes,
19305 kvl.v_tok_bytes,
19306 swa && crate::Engine::wkv_on(),
19307 )?;
19308 }
19309 if !fa_q8 {
19310 let aq = unsafe { &*(&sl.attn as *const CudaSlice<f32>) };
19311 e.quantize_q8_1_into(aq, 1, nh * hd, &mut sl.zq, &mut sl.zd)?;
19312 }
19313 {
19314 let zq = unsafe { &*(&sl.zq as *const CudaSlice<i8>) };
19315 let zd = unsafe { &*(&sl.zd as *const CudaSlice<f32>) };
19316 self.g4_matvec_m1_into(e, &fa.wo, zq, zd, &mut sl.o)?;
19317 }
19318 Ok(())
19319 }
19320
19321 fn gemma4_layer_tail_slotted(
19324 &self,
19325 e: &Engine,
19326 layer: &crate::hybrid::HybridLayer,
19327 next_norm: Option<&CudaSlice<f32>>,
19328 sl: &mut G4DcSlots,
19329 ) -> Result<(), Box<dyn std::error::Error>> {
19330 let n_embd = self.cfg.n_embd as usize;
19331 let eps = self.cfg.rms_eps;
19332 let bits = layer.gemma4.as_ref().unwrap();
19333 let crate::hybrid::Ffn::Dense {
19334 ffn_gate,
19335 ffn_up,
19336 ffn_down,
19337 } = &layer.ffn
19338 else {
19339 return Err("slotted tail: dense ffn only".into());
19340 };
19341 let pnfold = Engine::g4_pnfold_on();
19342 if pnfold {
19343 let or = unsafe { &*(&sl.o as *const CudaSlice<f32>) };
19346 let xr = unsafe { &*(&sl.x as *const CudaSlice<f32>) };
19347 e.rms_pre_add_rms_norm_q8z_into(
19348 or,
19349 layer.post_attn_norm.float_data(),
19350 xr,
19351 bits.ffn_norm.float_data(),
19352 &mut sl.attn_out,
19353 &mut sl.zsh,
19354 n_embd,
19355 1,
19356 eps,
19357 &mut sl.zq,
19358 &mut sl.zd,
19359 )?;
19360 } else {
19361 e.add_rms_norm(
19362 &sl.cur,
19363 &sl.x,
19364 bits.ffn_norm.float_data(),
19365 &mut sl.attn_out,
19366 &mut sl.zsh,
19367 n_embd,
19368 1,
19369 eps,
19370 )?;
19371 }
19372 let n_ff = ffn_gate.out_features();
19373 if !pnfold {
19374 let zshr = unsafe { &*(&sl.zsh as *const CudaSlice<f32>) };
19375 e.quantize_q8_1_into(zshr, 1, n_embd, &mut sl.zq, &mut sl.zd)?;
19376 }
19377 {
19378 let zq = unsafe { &*(&sl.zq as *const CudaSlice<i8>) };
19379 let zd = unsafe { &*(&sl.zd as *const CudaSlice<f32>) };
19380 if !e.matmul_q4_fused2_into(ffn_gate, ffn_up, zq, zd, &mut sl.gate, &mut sl.up)?
19381 && !e.matmul_nvfp4_fused2_into(
19382 ffn_gate,
19383 ffn_up,
19384 zq,
19385 zd,
19386 &mut sl.gate,
19387 &mut sl.up,
19388 )?
19389 {
19390 return Err("slotted tail: ffn fused2 unavailable".into());
19391 }
19392 }
19393 debug_assert!(e.uses_q8_1_fast(ffn_down));
19394 {
19395 let upr = unsafe { &*(&sl.up as *const CudaSlice<f32>) };
19396 let upv = e.view(upr, n_ff);
19397 let up_all = upv.slice(0..n_ff);
19398 let gr = unsafe { &*(&sl.gate as *const CudaSlice<f32>) };
19399 e.gelu_tanh_mul_q8_1_into(
19400 gr,
19401 &up_all,
19402 &mut sl.act,
19403 n_ff,
19404 1,
19405 &mut sl.actq,
19406 &mut sl.actd,
19407 )?;
19408 }
19409 {
19410 let aq = unsafe { &*(&sl.actq as *const CudaSlice<i8>) };
19411 let ad = unsafe { &*(&sl.actd as *const CudaSlice<f32>) };
19412 self.g4_matvec_m1_into(e, ffn_down, aq, ad, &mut sl.f0)?;
19413 }
19414 if pnfold {
19415 if let Some(w) = next_norm {
19418 let f0r = unsafe { &*(&sl.f0 as *const CudaSlice<f32>) };
19419 let aor = unsafe { &*(&sl.attn_out as *const CudaSlice<f32>) };
19420 e.rms_pre_add_scale_rms_norm_q8_1_into(
19421 f0r,
19422 bits.post_ffw_norm.float_data(),
19423 aor,
19424 bits.layer_scale,
19425 w,
19426 &mut sl.xn,
19427 n_embd,
19428 1,
19429 eps,
19430 &mut sl.hq,
19431 &mut sl.hd_,
19432 )?;
19433 return Ok(());
19434 }
19435 }
19436 e.rms_norm(
19437 &sl.f0,
19438 bits.post_ffw_norm.float_data(),
19439 &mut sl.sn,
19440 n_embd,
19441 1,
19442 eps,
19443 )?;
19444 match next_norm {
19445 Some(w) => {
19446 e.add_scale_rms_norm_q8_1_into(
19447 &sl.sn,
19448 &sl.attn_out,
19449 bits.layer_scale,
19450 w,
19451 &mut sl.xn,
19452 n_embd,
19453 1,
19454 eps,
19455 &mut sl.hq,
19456 &mut sl.hd_,
19457 )?;
19458 }
19459 None => {
19460 e.add_scale(&sl.sn, &sl.attn_out, bits.layer_scale, &mut sl.xn, n_embd)?;
19461 }
19462 }
19463 Ok(())
19464 }
19465
19466 #[allow(clippy::too_many_arguments)]
19468 fn gemma4_decode_attn_dc(
19469 &self,
19470 e: &Engine,
19471 fa: &crate::hybrid::FullAttnLayer,
19472 il: usize,
19473 hq: &CudaSlice<i8>,
19474 hdq: &CudaSlice<f32>,
19475 pos_d: &CudaSlice<i32>,
19476 cache: &mut Cache,
19477 cap_bucket_max: Option<(usize, usize)>,
19478 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
19479 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
19480 let eps = self.cfg.rms_eps;
19481 let aux = self.gemma4_aux.as_ref().unwrap();
19482 let ones = aux.ones(e);
19483 #[cfg(debug_assertions)]
19484 crate::debug_assert_tensor_stream_device(ones, &e.stream(), "gemma4_decode_attn_dc.ones");
19485 let (q0, k0, v0) = if swa {
19486 match e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)? {
19487 Some(t3) => t3,
19488 None => match e.matmul_nvfp4_fused2(&fa.wq, &fa.wk, hq, hdq, 1)? {
19490 Some((q0, k0)) => {
19491 let h0 = e.zeros(0)?;
19492 let v0 = e.matmul_pre(&fa.wv, hq, hdq, &h0, 1)?;
19493 (q0, k0, v0)
19494 }
19495 None => {
19496 let h0 = e.zeros(0)?;
19497 (
19498 e.matmul_pre(&fa.wq, hq, hdq, &h0, 1)?,
19499 e.matmul_pre(&fa.wk, hq, hdq, &h0, 1)?,
19500 e.matmul_pre(&fa.wv, hq, hdq, &h0, 1)?,
19501 )
19502 }
19503 },
19504 }
19505 } else {
19506 let (q0, k0) = match e.matmul_q4_fused2(&fa.wq, &fa.wk, hq, hdq)? {
19507 Some(p) => p,
19508 None => match e.matmul_nvfp4_fused2(&fa.wq, &fa.wk, hq, hdq, 1)? {
19509 Some(p) => p,
19510 None => {
19511 let h0 = e.zeros(0)?;
19512 (
19513 e.matmul_pre(&fa.wq, hq, hdq, &h0, 1)?,
19514 e.matmul_pre(&fa.wk, hq, hdq, &h0, 1)?,
19515 )
19516 }
19517 },
19518 };
19519 let v0 = e.clone_dtod(&k0)?;
19520 (q0, k0, v0)
19521 };
19522 let mut q = e.uninit(nh * hd)?;
19523 let mut k = e.uninit(nkv * hd)?;
19524 let mut v = e.uninit(nkv * hd)?;
19525 let ff = if swa {
19527 None
19528 } else {
19529 Some(
19530 aux.rope_freqs(e)
19531 .expect("gemma4 global rope needs rope_freqs.weight"),
19532 )
19533 };
19534 #[cfg(debug_assertions)]
19535 if let Some(ff) = ff {
19536 crate::debug_assert_tensor_stream_device(
19537 ff,
19538 &e.stream(),
19539 "gemma4_decode_attn_dc.rope_freqs",
19540 );
19541 }
19542 let kvl = cache.kv[il].as_mut().unwrap();
19543 let kv_fp8 = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
19544 if crate::Engine::qkv_append_on() {
19545 e.rms_norm_qkv_rope_append_dc(
19547 &q0,
19548 &k0,
19549 &v0,
19550 fa.q_norm.float_data(),
19551 fa.k_norm.float_data(),
19552 ones,
19553 &mut q,
19554 &mut k,
19555 &mut v,
19556 hd,
19557 self.gemma4_rope_dims(il),
19558 nh,
19559 nkv,
19560 pos_d,
19561 nh,
19562 nkv,
19563 base,
19564 1.0,
19565 ff,
19566 eps,
19567 &mut kvl.k,
19568 &mut kvl.v,
19569 &kvl.len_d,
19570 kvl.k_tok_bytes,
19571 kvl.v_tok_bytes,
19572 kv_fp8,
19573 )?;
19574 } else {
19575 e.rms_norm_qkv_rope(
19576 &q0,
19577 &k0,
19578 &v0,
19579 fa.q_norm.float_data(),
19580 fa.k_norm.float_data(),
19581 ones,
19582 &mut q,
19583 &mut k,
19584 &mut v,
19585 hd,
19586 self.gemma4_rope_dims(il),
19587 nh,
19588 nkv,
19589 pos_d,
19590 nh,
19591 nkv,
19592 base,
19593 1.0,
19594 ff,
19595 eps,
19596 )?;
19597 e.append_kv_quantized_dc(
19598 &k,
19599 &v,
19600 &mut kvl.k,
19601 &mut kvl.v,
19602 &kvl.len_d,
19603 kvl.kv_dim_k,
19604 kvl.kv_dim_v,
19605 kvl.k_tok_bytes,
19606 kvl.v_tok_bytes,
19607 kv_fp8,
19608 )?;
19609 }
19610 e.inc_seqlen(&mut kvl.len_d)?;
19611 let mut attn = e.uninit(nh * hd)?;
19612 let mut fa_q8: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
19615 match cap_bucket_max {
19620 None => {
19621 kvl.len += 1;
19625 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
19626 if !swa
19627 && hd == 512
19628 && kvl.len >= crate::fa512_min_tkv()
19629 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
19630 {
19631 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
19634 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
19635 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
19636 e.fa_decode_rows(
19637 &q,
19638 &kp,
19639 &vp,
19640 &mut attn,
19641 hd,
19642 nh,
19643 nkv,
19644 kvl.len - 1,
19645 1,
19646 scale,
19647 kvl.k_tok_bytes,
19648 kvl.v_tok_bytes,
19649 Some((&kvl.len_d, -1)),
19650 false,
19651 false,
19652 Some((&mut aq8, &mut ad8)),
19653 )?;
19654 fa_q8 = Some((aq8, ad8));
19655 } else if swa
19656 && kvl.len > win
19657 && hd == 256
19658 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
19659 {
19660 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
19662 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
19663 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
19664 e.fa_decode_rows_w(
19665 &q,
19666 &kp,
19667 &vp,
19668 &mut attn,
19669 hd,
19670 nh,
19671 nkv,
19672 &kvl.len_d,
19673 -1,
19674 1,
19675 scale,
19676 win,
19677 kvl.k_tok_bytes,
19678 kvl.v_tok_bytes,
19679 Some((&mut aq8, &mut ad8)),
19680 )?;
19681 fa_q8 = Some((aq8, ad8));
19682 } else {
19683 let (off_tok, t_kv) = if swa && kvl.len > win {
19684 (kvl.len - win, win)
19685 } else {
19686 (0, kvl.len)
19687 };
19688 let k_view = e.view_u8_range(
19689 &kvl.k,
19690 off_tok * kvl.k_tok_bytes,
19691 (off_tok + t_kv) * kvl.k_tok_bytes,
19692 );
19693 let v_view = e.view_u8_range(
19694 &kvl.v,
19695 off_tok * kvl.v_tok_bytes,
19696 (off_tok + t_kv) * kvl.v_tok_bytes,
19697 );
19698 e.fa_decode_kvmod(
19699 &q,
19700 &k_view,
19701 &v_view,
19702 &mut attn,
19703 hd,
19704 nh,
19705 nkv,
19706 t_kv,
19707 scale,
19708 kvl.k_tok_bytes,
19709 kvl.v_tok_bytes,
19710 swa && crate::Engine::wkv_on(),
19711 )?;
19712 }
19713 }
19714 Some((b_swa, b_glob)) => {
19715 let k_view = e.view_u8(&kvl.k, kvl.k.len());
19721 let v_view = e.view_u8(&kvl.v, kvl.v.len());
19722 let rows_on = std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0");
19723 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
19724 if !swa && hd == 512 && b_glob >= crate::fa512_min_tkv() && rows_on {
19725 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
19726 e.fa_decode_rows(
19727 &q,
19728 &k_view,
19729 &v_view,
19730 &mut attn,
19731 hd,
19732 nh,
19733 nkv,
19734 b_glob - 1,
19735 1,
19736 scale,
19737 kvl.k_tok_bytes,
19738 kvl.v_tok_bytes,
19739 Some((&kvl.len_d, -1)),
19740 false,
19741 false,
19742 Some((&mut aq8, &mut ad8)),
19743 )?;
19744 fa_q8 = Some((aq8, ad8));
19745 } else if swa && b_swa > win && hd == 256 && rows_on {
19746 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
19747 e.fa_decode_rows_w(
19748 &q,
19749 &k_view,
19750 &v_view,
19751 &mut attn,
19752 hd,
19753 nh,
19754 nkv,
19755 &kvl.len_d,
19756 -1,
19757 1,
19758 scale,
19759 win,
19760 kvl.k_tok_bytes,
19761 kvl.v_tok_bytes,
19762 Some((&mut aq8, &mut ad8)),
19763 )?;
19764 fa_q8 = Some((aq8, ad8));
19765 } else {
19766 let b = if swa { b_swa } else { b_glob };
19767 e.fa_decode_dc(
19768 &q,
19769 &k_view,
19770 &v_view,
19771 &mut attn,
19772 hd,
19773 nh,
19774 nkv,
19775 &kvl.len_d,
19776 b,
19777 scale,
19778 kvl.k_tok_bytes,
19779 kvl.v_tok_bytes,
19780 swa && crate::Engine::wkv_on(),
19781 )?;
19782 }
19783 }
19784 }
19785 if let Some((aq8, ad8)) = fa_q8 {
19788 let mut y = e.uninit(fa.wo.out_features())?;
19789 self.g4_matvec_m1_into(e, &fa.wo, &aq8, &ad8, &mut y)?;
19790 return Ok(y);
19791 }
19792 e.matmul(&fa.wo, &attn, 1)
19793 }
19794
19795 #[allow(clippy::too_many_arguments)]
19800 #[allow(clippy::map_entry)] pub fn gemma4_generate_graph(
19803 &self,
19804 e: &Engine,
19805 prompt_pos: usize,
19806 first_token: u32,
19807 cache: &mut Cache,
19808 max_new: usize,
19809 eos: &[u32],
19810 mut on_token: impl FnMut(u32) -> bool,
19811 ) -> Result<(Vec<u32>, crate::decode::StopReason), Box<dyn std::error::Error>> {
19812 if self.is_gemma4_e4b() {
19813 return Err(
19814 "E4B graph serving is unwired (HANDOVER-E4B.md) — dc-eager is the serving arm"
19815 .into(),
19816 );
19817 }
19818 use crate::decode::StopReason;
19819 let n_vocab = self.output.out_features();
19820 let n_embd = self.cfg.n_embd as usize;
19821 let embd_gpu = self
19822 .embd_gpu
19823 .get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload"));
19824 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
19825 for kvl in cache.kv.iter_mut().flatten() {
19826 e.set_i32_one(&mut kvl.len_d, kvl.len as i32)?;
19827 }
19828 let mut token_d = e.stream().clone_htod(&[first_token])?;
19829 let mut pos_d = e.htod_i32(&[prompt_pos as i32])?;
19830 let g4 = self.cfg.gemma4.as_ref().unwrap();
19831 let (hd_s, hd_g) = (g4.key_length_swa as usize, g4.key_length_global as usize);
19832 let nkv_s = g4
19834 .head_count_kv
19835 .iter()
19836 .zip(g4.swa_pattern.iter())
19837 .find(|p| *p.1)
19838 .map(|p| *p.0 as usize)
19839 .unwrap_or(8);
19840 let nkv_g = g4
19841 .head_count_kv
19842 .iter()
19843 .zip(g4.swa_pattern.iter())
19844 .find(|p| !*p.1)
19845 .map(|p| *p.0 as usize)
19846 .unwrap_or(2);
19847 #[allow(clippy::type_complexity)]
19848 let mut graphs: std::collections::HashMap<
19850 ((bool, usize), (bool, usize), bool, bool),
19851 (
19852 cudarc::driver::CudaGraph,
19853 Vec<Box<dyn std::any::Any + Send>>,
19854 ),
19855 > = Default::default();
19856 let mut slots = self.g4_dc_slots(e)?;
19859 const RING: usize = 64;
19862 const DRAIN: usize = 1;
19868 let mut ring = e.stream().alloc_zeros::<u32>(RING)?;
19869 let ring_base = prompt_pos;
19870 let mut out = Vec::with_capacity(max_new);
19871 let mut reason = StopReason::MaxNew;
19872 let mut next = first_token;
19873 let mut captures = 0usize;
19874 for _ in 0..max_new {
19875 out.push(next);
19876 if eos.contains(&next) {
19877 reason = StopReason::Eos;
19878 break;
19879 }
19880 if !on_token(next) {
19881 reason = StopReason::Callback;
19882 break;
19883 }
19884 let t_kv = cache.pos + 1;
19885 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
19893 let f512 = crate::fa512_min_tkv();
19894 let key_s = if t_kv > win {
19895 (true, usize::MAX)
19896 } else {
19897 e.fa_bucket_key(t_kv, hd_s, nkv_s, crate::Engine::wkv_on())
19898 };
19899 let (key_g, rung_end) = if t_kv >= f512 {
19900 let end = (t_kv + 1).next_power_of_two().max(f512 * 2);
19903 ((true, end), end)
19904 } else {
19905 (e.fa_bucket_key(t_kv, hd_g, nkv_g, false), t_kv)
19906 };
19907 let key = (key_s, key_g, t_kv >= f512, t_kv > win);
19908 if !graphs.contains_key(&key) {
19909 let bucket_max = (t_kv, rung_end);
19910 let snap = cache.snapshot(e)?;
19912 let pos_save = e.dtoh_i32_one(&pos_d)?;
19913 let len_save: Vec<Option<i32>> = cache
19914 .kv
19915 .iter()
19916 .map(|k| k.as_ref().map(|kvl| e.dtoh_i32_one(&kvl.len_d).unwrap()))
19917 .collect();
19918 let tok_save = e.dtoh_u32_one(&token_d)?;
19919 let graph = {
19924 let tok_ref = &mut token_d;
19925 let pos_ref = &mut pos_d;
19926 let cache_ref = &mut *cache;
19927 let slots_ref = &mut slots;
19928 let ring_ref = &mut ring;
19929 e.capture_graph_retained_flags(
19930 cudarc::driver::sys::CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_USE_NODE_PRIORITY,
19931 |e| {
19932 let tok_in = unsafe { &*(tok_ref as *const CudaSlice<u32>) };
19934 let sl = unsafe { &mut *(slots_ref as *mut G4DcSlots) };
19935 let rg = unsafe { &mut *(ring_ref as *mut CudaSlice<u32>) };
19936 self.gemma4_decode_step_dc_slotted(e, tok_in, pos_ref, embd_gpu, qt, rb,
19937 cache_ref, n_vocab, Some(bucket_max),
19938 sl, tok_ref, Some((rg, ring_base)))
19939 })?
19940 };
19941 cache.rollback(e, &snap, 0)?;
19942 e.set_i32_one(&mut pos_d, pos_save)?;
19943 for (il, ls) in len_save.iter().enumerate() {
19944 if let (Some(kvl), Some(v)) = (cache.kv[il].as_mut(), ls) {
19945 e.set_i32_one(&mut kvl.len_d, *v)?;
19946 }
19947 }
19948 e.set_u32_one(&mut token_d, tok_save)?;
19949 if std::env::var("MEMRA_GRAPH_CENSUS").as_deref() == Ok("1")
19950 && let Ok(c) = crate::graph_update::node_census(&graph.0)
19951 {
19952 eprintln!("[graph-census] {c:?}");
19953 }
19954 graphs.insert(key, graph);
19955 captures += 1;
19956 }
19957 let mut chunk = 1usize;
19962 let drain_cap: usize = std::env::var("MEMRA_GRAPH_DRAIN")
19963 .ok()
19964 .and_then(|v| v.parse().ok())
19965 .unwrap_or(DRAIN);
19966 while chunk < drain_cap && out.len() + chunk < max_new {
19967 let t_next = cache.pos + 1 + chunk;
19968 let key_s2 = if t_next > win {
19969 (true, usize::MAX)
19970 } else {
19971 e.fa_bucket_key(t_next, hd_s, nkv_s, crate::Engine::wkv_on())
19972 };
19973 let key_g2 = if t_next >= f512 {
19974 (true, (t_next + 1).next_power_of_two().max(f512 * 2))
19975 } else {
19976 e.fa_bucket_key(t_next, hd_g, nkv_g, false)
19977 };
19978 if (key_s2, key_g2, t_next >= f512, t_next > win) != key {
19979 break;
19980 }
19981 chunk += 1;
19982 }
19983 let g = &graphs.get(&key).unwrap().0;
19984 for _ in 0..chunk {
19985 g.launch()?;
19986 }
19987 e.stream().synchronize()?;
19988 let ringh = e.dtoh_u32(&ring)?;
19989 for j in 0..chunk {
19990 let pos_j = cache.pos + j;
19991 let tok_j = ringh[(pos_j - ring_base) % RING];
19992 cache.pos += 0; if j + 1 == chunk {
19994 next = tok_j;
19995 } else {
19996 out.push(tok_j);
19997 if eos.contains(&tok_j) || !on_token(tok_j) {
19998 reason = if eos.contains(&tok_j) {
19999 StopReason::Eos
20000 } else {
20001 StopReason::Callback
20002 };
20003 let keep = cache.pos + j + 1;
20005 e.set_i32_one(&mut pos_d, keep as i32)?;
20006 for kvl in cache.kv.iter_mut().filter_map(|k| k.as_mut()) {
20007 e.set_i32_one(&mut kvl.len_d, keep as i32)?;
20008 kvl.len = keep;
20009 }
20010 cache.pos = keep;
20011 if std::env::var("MEMRA_GRAPH_STATS").is_ok() {
20012 eprintln!("[gemma-graph] captures={captures} buckets={}", graphs.len());
20013 }
20014 return Ok((out, reason));
20015 }
20016 }
20017 }
20018 cache.pos += chunk;
20019 for kvl in cache.kv.iter_mut().filter_map(|k| k.as_mut()) {
20020 kvl.len += chunk;
20021 }
20022 }
20023 if std::env::var("MEMRA_GRAPH_STATS").is_ok() {
20024 eprintln!("[gemma-graph] captures={captures} buckets={}", graphs.len());
20025 }
20026 Ok((out, reason))
20027 }
20028
20029 pub(crate) fn gemma4_decode_step_t(
20035 &self,
20036 e: &Engine,
20037 tokens: &[u32],
20038 pos0: usize,
20039 cache: &mut Cache,
20040 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
20041 Ok(self.gemma4_decode_step_t_h(e, tokens, pos0, cache)?.0)
20042 }
20043
20044 pub(crate) fn gemma4_decode_step_t_am(
20048 &self,
20049 e: &Engine,
20050 tokens: &[u32],
20051 pos0: usize,
20052 cache: &mut Cache,
20053 ) -> Result<(Vec<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
20054 let (ld, hn) = self.gemma4_verify_trunk(e, tokens, pos0, cache, None)?;
20055 let t = tokens.len();
20056 let n_vocab = self.output.out_features();
20057 let mut toks = e.stream().alloc_zeros::<u32>(t)?;
20058 for i in 0..t {
20059 e.argmax_token_device_col(&ld, i, n_vocab, &mut toks, i)?;
20060 }
20061 Ok((e.dtoh_u32(&toks)?, hn))
20062 }
20063
20064 pub(crate) fn gemma4_decode_step_t_am_dev(
20067 &self,
20068 e: &Engine,
20069 tok_d: &CudaSlice<u32>,
20070 t: usize,
20071 pos0: usize,
20072 cache: &mut Cache,
20073 ) -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
20074 let (ld, hn) = self.gemma4_verify_trunk(e, &vec![0u32; t], pos0, cache, Some(tok_d))?;
20075 let n_vocab = self.output.out_features();
20076 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
20077 for i in 0..t {
20078 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
20079 }
20080 Ok((vam, hn))
20081 }
20082
20083 pub(crate) fn gemma4_decode_step_t_h(
20086 &self,
20087 e: &Engine,
20088 tokens: &[u32],
20089 pos0: usize,
20090 cache: &mut Cache,
20091 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
20092 let (mut ld, hn) = self.gemma4_verify_trunk(e, tokens, pos0, cache, None)?;
20093 let t = tokens.len();
20094 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
20095 e.softcap(&mut ld, cap, t * self.output.out_features())?;
20096 Ok((e.dtoh(&ld)?, hn))
20097 }
20098
20099 pub(crate) fn verify_stream_scratch(
20102 &self,
20103 e: &Engine,
20104 cap: usize,
20105 ) -> Result<VerifyStreamScratch, Box<dyn std::error::Error>> {
20106 Ok(VerifyStreamScratch {
20107 pos_d: e.htod_i32(&vec![0i32; cap])?,
20108 row_ctrs: (0..cap)
20109 .map(|_| e.htod_i32(&[0]))
20110 .collect::<Result<_, _>>()?,
20111 })
20112 }
20113
20114 #[allow(clippy::too_many_arguments)] pub(crate) fn gemma4_verify_t_am_stream(
20123 &self,
20124 e: &Engine,
20125 tok_d: &CudaSlice<u32>,
20126 t: usize,
20127 ctr: &CudaSlice<i32>,
20128 hint: usize,
20129 cache: &mut Cache,
20130 scr: &mut VerifyStreamScratch,
20131 ) -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
20132 let n_embd = self.cfg.n_embd as usize;
20133 let eps = self.cfg.rms_eps;
20134 assert!(t <= scr.row_ctrs.len() && t <= 64);
20135 e.i32_iota_from(ctr, &mut scr.pos_d, t)?;
20136 for i in 0..t {
20137 e.i32_copy_add(ctr, &mut scr.row_ctrs[i], (i + 1) as i32)?;
20138 }
20139 let (pos_d, row_ctrs) = (&scr.pos_d, &scr.row_ctrs);
20140 let embd_gpu = self
20141 .embd_gpu
20142 .get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload"));
20143 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
20144 let mut x = e.embed_gather_device_td(embd_gpu, tok_d, t, n_embd, qt, rb)?;
20145 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
20146 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
20147 let n_layers = self.layers.len();
20148 for (il, layer) in self.layers.iter().enumerate() {
20149 let (hq, hdq) = match h_carry.take() {
20150 Some(p) => p,
20151 None => {
20152 e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, t, eps)?
20153 }
20154 };
20155 let Mixer::Full(fa) = &layer.mixer else {
20156 panic!("gemma4 layer {il} not full-attn")
20157 };
20158 let o = self
20159 .gemma4_verify_attn_stream(e, fa, il, &hq, &hdq, pos_d, t, cache, hint, row_ctrs)?;
20160 let next_norm = if il + 1 < n_layers {
20161 Some(self.layers[il + 1].attn_norm.float_data())
20162 } else {
20163 None
20164 };
20165 let (xn, hn) = self.gemma4_layer_tail_add_nq_pn(e, layer, &o, &x, t, next_norm)?;
20166 x = xn;
20167 h_carry = hn;
20168 self.dflash_tap(e, cache, il, &x, t)?;
20169 }
20170 let mut hn = e.uninit(t * n_embd)?;
20171 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
20172 let ld = e.matmul(&self.output, &hn, t)?;
20173 let n_vocab = self.output.out_features();
20174 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
20175 for i in 0..t {
20176 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
20177 }
20178 Ok((vam, hn))
20179 }
20180
20181 pub(crate) fn dflash_tap(
20188 &self,
20189 e: &Engine,
20190 cache: &mut Cache,
20191 il: usize,
20192 x: &CudaSlice<f32>,
20193 t: usize,
20194 ) -> Result<(), Box<dyn std::error::Error>> {
20195 let Some(taps) = cache.dflash_taps.as_mut() else {
20196 return Ok(());
20197 };
20198 let Some(slot) = taps.layer_ids.iter().position(|&l| l == il) else {
20199 return Ok(());
20200 };
20201 let h = taps.hidden;
20202 let n_taps = taps.layer_ids.len();
20203 let base = taps.base;
20204 debug_assert!(
20205 base + t <= taps.t,
20206 "tap window {base}+{t} exceeds sink {}",
20207 taps.t
20208 );
20209 let xv = e.view(x, t * h);
20210 for r in 0..t {
20211 let row = xv.slice(r * h..(r + 1) * h);
20212 e.copy_view_into(&mut taps.buf, (base + r) * n_taps * h + slot * h, &row, h)?;
20213 }
20214 Ok(())
20215 }
20216
20217 fn gemma4_verify_trunk(
20218 &self,
20219 e: &Engine,
20220 tokens: &[u32],
20221 pos0: usize,
20222 cache: &mut Cache,
20223 tok_dev: Option<&CudaSlice<u32>>,
20224 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
20225 let n_embd = self.cfg.n_embd as usize;
20226 let eps = self.cfg.rms_eps;
20227 let t = tokens.len();
20228 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
20229 let pos_d = e.htod_i32(&pos)?;
20230 let mut x = match tok_dev {
20231 Some(td) => {
20232 let embd_gpu = self
20233 .embd_gpu
20234 .get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload"));
20235 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
20236 e.embed_gather_device_td(embd_gpu, td, t, n_embd, qt, rb)?
20237 }
20238 None => e.htod(&self.embd.gather(n_embd, tokens))?,
20239 };
20240 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
20241 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
20242 let n_layers = self.layers.len();
20243 for (il, layer) in self.layers.iter().enumerate() {
20244 let (hq, hdq) = match h_carry.take() {
20245 Some(p) => p,
20246 None => {
20247 e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, t, eps)?
20248 }
20249 };
20250 let Mixer::Full(fa) = &layer.mixer else {
20251 panic!("gemma4 layer {il} not full-attn")
20252 };
20253 let o = self.gemma4_verify_attn(e, fa, il, &hq, &hdq, &pos_d, t, cache)?;
20254 let next_norm = if il + 1 < n_layers {
20255 Some(self.layers[il + 1].attn_norm.float_data())
20256 } else {
20257 None
20258 };
20259 let (xn, hn) = self.gemma4_layer_tail_add_nq_pn(e, layer, &o, &x, t, next_norm)?;
20260 x = xn;
20261 h_carry = hn;
20262 self.dflash_tap(e, cache, il, &x, t)?;
20263 }
20264 let mut hn = e.uninit(t * n_embd)?;
20265 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
20266 let mut ld = e.matmul(&self.output, &hn, t)?;
20267 self.gemma4_suppress(e, &mut ld, t)?; cache.pos += t;
20269 Ok((ld, hn))
20270 }
20271
20272 #[allow(clippy::too_many_arguments)]
20280 fn gemma4_verify_attn_stream(
20281 &self,
20282 e: &Engine,
20283 fa: &crate::hybrid::FullAttnLayer,
20284 il: usize,
20285 hq: &CudaSlice<i8>,
20286 hdq: &CudaSlice<f32>,
20287 pos_d: &CudaSlice<i32>,
20288 t: usize,
20289 cache: &mut Cache,
20290 hint: usize,
20291 row_ctrs: &[CudaSlice<i32>],
20292 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
20293 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
20294 let eps = self.cfg.rms_eps;
20295 let aux = self.gemma4_aux.as_ref().unwrap();
20296 let ones = aux.ones(e);
20297 #[cfg(debug_assertions)]
20298 crate::debug_assert_tensor_stream_device(
20299 ones,
20300 &e.stream(),
20301 "gemma4_verify_attn_stream.ones",
20302 );
20303 let h0 = e.zeros(0)?;
20304 let h = &h0;
20305 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
20308 let f2b = *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
20309 let fused_qkv = if f2b {
20310 if swa {
20311 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
20312 .map(|(a, b, c)| (a, b, Some(c)))
20313 } else {
20314 e.matmul_q4_fused2_batched(&fa.wq, &fa.wk, hq, hdq, t)?
20315 .map(|(a, b)| (a, b, None))
20316 }
20317 } else {
20318 None
20319 };
20320 let (q0, k0, v0) = match fused_qkv {
20321 Some((a, b, cv)) => {
20322 let v = match cv {
20323 Some(c) => c,
20324 None => e.clone_dtod(&b)?,
20325 };
20326 (a, b, v)
20327 }
20328 None => {
20329 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
20330 let k0 = e.matmul_pre(&fa.wk, hq, hdq, h, t)?;
20331 let v0 = if swa {
20332 e.matmul_pre(&fa.wv, hq, hdq, h, t)?
20333 } else {
20334 e.clone_dtod(&k0)?
20335 };
20336 (q0, k0, v0)
20337 }
20338 };
20339 let mut q = e.uninit(t * nh * hd)?;
20340 let mut k = e.uninit(t * nkv * hd)?;
20341 let mut v = e.uninit(t * nkv * hd)?;
20342 let ff = if swa {
20345 None
20346 } else {
20347 Some(
20348 aux.rope_freqs(e)
20349 .expect("gemma4 global rope needs rope_freqs.weight"),
20350 )
20351 };
20352 #[cfg(debug_assertions)]
20353 if let Some(ff) = ff {
20354 crate::debug_assert_tensor_stream_device(
20355 ff,
20356 &e.stream(),
20357 "gemma4_verify_attn_stream.rope_freqs",
20358 );
20359 }
20360 e.rms_norm_qkv_rope(
20361 &q0,
20362 &k0,
20363 &v0,
20364 fa.q_norm.float_data(),
20365 fa.k_norm.float_data(),
20366 ones,
20367 &mut q,
20368 &mut k,
20369 &mut v,
20370 hd,
20371 self.gemma4_rope_dims(il),
20372 nh * t,
20373 nkv * t,
20374 pos_d,
20375 nh,
20376 nkv,
20377 base,
20378 1.0,
20379 ff,
20380 eps,
20381 )?;
20382 let kvl = cache.kv[il].as_mut().unwrap();
20383 e.append_kv_quantized_rows_dc(
20385 &k,
20386 &v,
20387 &mut kvl.k,
20388 &mut kvl.v,
20389 &kvl.len_d,
20390 t,
20391 kvl.kv_dim_k,
20392 kvl.kv_dim_v,
20393 kvl.k_tok_bytes,
20394 kvl.v_tok_bytes,
20395 (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on()),
20396 )?;
20397 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
20400 let mut attn = e.uninit(t * nh * hd)?;
20401 let k_view = e.view_u8(&kvl.k, kvl.k.len());
20402 let v_view = e.view_u8(&kvl.v, kvl.v.len());
20403 if swa && hint + 1 >= win {
20406 e.fa_decode_rows_w(
20409 &q,
20410 &k_view,
20411 &v_view,
20412 &mut attn,
20413 hd,
20414 nh,
20415 nkv,
20416 &kvl.len_d,
20417 0,
20418 t,
20419 scale,
20420 win,
20421 kvl.k_tok_bytes,
20422 kvl.v_tok_bytes,
20423 None,
20424 )?;
20425 } else if hd == 512 && hint + t < crate::fa512_min_tkv() {
20426 let bucket = (hint + t + 2)
20439 .next_power_of_two()
20440 .min(crate::fa512_min_tkv().saturating_sub(1));
20441 let qv = e.view(&q, t * nh * hd);
20442 #[allow(clippy::needless_range_loop)]
20443 for i in 0..t {
20445 let q_row = qv.slice(i * nh * hd..(i + 1) * nh * hd);
20446 let mut q_one = e.uninit(nh * hd)?;
20447 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
20448 let mut a_one = e.uninit(nh * hd)?;
20449 e.fa_decode_dc(
20450 &q_one,
20451 &k_view,
20452 &v_view,
20453 &mut a_one,
20454 hd,
20455 nh,
20456 nkv,
20457 &row_ctrs[i],
20458 bucket,
20459 scale,
20460 kvl.k_tok_bytes,
20461 kvl.v_tok_bytes,
20462 false,
20463 )?;
20464 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
20465 }
20466 } else if hd == 512 {
20467 e.fa_decode_rows(
20470 &q,
20471 &k_view,
20472 &v_view,
20473 &mut attn,
20474 hd,
20475 nh,
20476 nkv,
20477 hint,
20478 t,
20479 scale,
20480 kvl.k_tok_bytes,
20481 kvl.v_tok_bytes,
20482 Some((&kvl.len_d, 0)),
20483 false,
20484 false,
20485 None,
20486 )?;
20487 } else {
20488 e.fa_decode_rows_dc(
20490 &q,
20491 &k_view,
20492 &v_view,
20493 &mut attn,
20494 hd,
20495 nh,
20496 nkv,
20497 &kvl.len_d,
20498 hint + t,
20499 t,
20500 scale,
20501 kvl.k_tok_bytes,
20502 kvl.v_tok_bytes,
20503 0,
20504 swa && crate::Engine::wkv_on(),
20505 )?;
20506 }
20507 e.matmul(&fa.wo, &attn, t)
20508 }
20509
20510 #[allow(clippy::too_many_arguments)] fn gemma4_verify_attn(
20512 &self,
20513 e: &Engine,
20514 fa: &crate::hybrid::FullAttnLayer,
20515 il: usize,
20516 hq: &CudaSlice<i8>,
20517 hdq: &CudaSlice<f32>,
20518 pos_d: &CudaSlice<i32>,
20519 t: usize,
20520 cache: &mut Cache,
20521 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
20522 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
20523 let eps = self.cfg.rms_eps;
20524 let aux = self.gemma4_aux.as_ref().unwrap();
20525 let ones = aux.ones(e);
20526 #[cfg(debug_assertions)]
20527 crate::debug_assert_tensor_stream_device(ones, &e.stream(), "gemma4_verify_attn.ones");
20528 let n_embd = self.cfg.n_embd as usize;
20529 let _ = n_embd;
20530
20531 let h0 = e.zeros(0)?;
20532 let h = &h0;
20533 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
20536 let f2b = *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
20537 let fused_qkv = if f2b {
20538 if swa {
20539 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
20540 .map(|(a, b, c)| (a, b, Some(c)))
20541 } else {
20542 e.matmul_q4_fused2_batched(&fa.wq, &fa.wk, hq, hdq, t)?
20543 .map(|(a, b)| (a, b, None))
20544 }
20545 } else {
20546 None
20547 };
20548 let (q0, k0, v0) = match fused_qkv {
20549 Some((a, b, cv)) => {
20550 let v = match cv {
20551 Some(c) => c,
20552 None => e.clone_dtod(&b)?,
20553 };
20554 (a, b, v)
20555 }
20556 None => {
20557 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
20558 let k0 = e.matmul_pre(&fa.wk, hq, hdq, h, t)?;
20559 let v0 = if swa {
20560 e.matmul_pre(&fa.wv, hq, hdq, h, t)?
20561 } else {
20562 e.clone_dtod(&k0)?
20563 };
20564 (q0, k0, v0)
20565 }
20566 };
20567 let mut q = e.uninit(t * nh * hd)?;
20568 let mut k = e.uninit(t * nkv * hd)?;
20569 let mut v = e.uninit(t * nkv * hd)?;
20570 let ff = if swa {
20573 None
20574 } else {
20575 Some(
20576 aux.rope_freqs(e)
20577 .expect("gemma4 global rope needs rope_freqs.weight"),
20578 )
20579 };
20580 #[cfg(debug_assertions)]
20581 if let Some(ff) = ff {
20582 crate::debug_assert_tensor_stream_device(
20583 ff,
20584 &e.stream(),
20585 "gemma4_verify_attn.rope_freqs",
20586 );
20587 }
20588 e.rms_norm_qkv_rope(
20589 &q0,
20590 &k0,
20591 &v0,
20592 fa.q_norm.float_data(),
20593 fa.k_norm.float_data(),
20594 ones,
20595 &mut q,
20596 &mut k,
20597 &mut v,
20598 hd,
20599 self.gemma4_rope_dims(il),
20600 nh * t,
20601 nkv * t,
20602 pos_d,
20603 nh,
20604 nkv,
20605 base,
20606 1.0,
20607 ff,
20608 eps,
20609 )?;
20610 let kvl = cache.kv[il].as_mut().unwrap();
20611 let base_len = kvl.len;
20612 e.append_kv_quantized_rows(
20613 &k,
20614 &v,
20615 &mut kvl.k,
20616 &mut kvl.v,
20617 base_len,
20618 t,
20619 kvl.kv_dim_k,
20620 kvl.kv_dim_v,
20621 kvl.k_tok_bytes,
20622 kvl.v_tok_bytes,
20623 (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on()),
20624 )?;
20625 kvl.len += t;
20626 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
20627 let mut attn = e.uninit(t * nh * hd)?;
20628 let rows_ok = (hd == 256 && base_len + 1 >= crate::fa_vec_min_tkv())
20631 || (hd == 512 && !swa && base_len + 1 >= crate::fa512_min_tkv());
20634 if rows_ok && (!swa || base_len + t <= win) {
20635 let k_view = e.view_u8(&kvl.k, (base_len + t) * kvl.k_tok_bytes);
20636 let v_view = e.view_u8(&kvl.v, (base_len + t) * kvl.v_tok_bytes);
20637 if hd == 512 {
20638 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
20640 e.fa_decode_rows(
20641 &q,
20642 &k_view,
20643 &v_view,
20644 &mut attn,
20645 hd,
20646 nh,
20647 nkv,
20648 base_len,
20649 t,
20650 scale,
20651 kvl.k_tok_bytes,
20652 kvl.v_tok_bytes,
20653 Some((&kvl.len_d, 0)),
20654 false,
20655 swa && crate::Engine::wkv_on(),
20656 None,
20657 )?;
20658 } else {
20659 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
20663 e.fa_decode_rows_dc(
20664 &q,
20665 &k_view,
20666 &v_view,
20667 &mut attn,
20668 hd,
20669 nh,
20670 nkv,
20671 &kvl.len_d,
20672 base_len + t,
20673 t,
20674 scale,
20675 kvl.k_tok_bytes,
20676 kvl.v_tok_bytes,
20677 0,
20678 swa && crate::Engine::wkv_on(),
20679 )?;
20680 }
20681 return e.matmul(&fa.wo, &attn, t);
20682 }
20683 if hd == 256
20691 && swa
20692 && base_len + 1 >= win
20693 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
20694 {
20695 let k_view = e.view_u8(&kvl.k, (base_len + t) * kvl.k_tok_bytes);
20696 let v_view = e.view_u8(&kvl.v, (base_len + t) * kvl.v_tok_bytes);
20697 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
20698 e.fa_decode_rows_w(
20699 &q,
20700 &k_view,
20701 &v_view,
20702 &mut attn,
20703 hd,
20704 nh,
20705 nkv,
20706 &kvl.len_d,
20707 0,
20708 t,
20709 scale,
20710 win,
20711 kvl.k_tok_bytes,
20712 kvl.v_tok_bytes,
20713 None,
20714 )?;
20715 return e.matmul(&fa.wo, &attn, t);
20716 }
20717 for i in 0..t {
20718 let avail = base_len + i + 1;
20719 let (off_tok, t_kv) = if swa && avail > win {
20720 (avail - win, win)
20721 } else {
20722 (0, avail)
20723 };
20724 let k_view = e.view_u8_range(
20725 &kvl.k,
20726 off_tok * kvl.k_tok_bytes,
20727 (off_tok + t_kv) * kvl.k_tok_bytes,
20728 );
20729 let v_view = e.view_u8_range(
20730 &kvl.v,
20731 off_tok * kvl.v_tok_bytes,
20732 (off_tok + t_kv) * kvl.v_tok_bytes,
20733 );
20734 let qi = e.view(&q, t * nh * hd);
20735 let q_row = qi.slice(i * nh * hd..(i + 1) * nh * hd);
20736 let mut q_one = e.uninit(nh * hd)?;
20737 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
20738 let mut a_one = e.uninit(nh * hd)?;
20739 if swa
20743 && avail > win
20744 && hd == 256
20745 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
20746 {
20747 let kp = e.view_u8(&kvl.k, avail * kvl.k_tok_bytes);
20748 let vp = e.view_u8(&kvl.v, avail * kvl.v_tok_bytes);
20749 e.i32_set_k(&mut kvl.len_d, (avail - 1) as i32)?;
20750 e.fa_decode_rows_w(
20751 &q_one,
20752 &kp,
20753 &vp,
20754 &mut a_one,
20755 hd,
20756 nh,
20757 nkv,
20758 &kvl.len_d,
20759 0,
20760 1,
20761 scale,
20762 win,
20763 kvl.k_tok_bytes,
20764 kvl.v_tok_bytes,
20765 None,
20766 )?;
20767 } else if !swa
20768 && hd == 512
20769 && avail >= crate::fa512_min_tkv()
20770 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
20771 {
20772 let kp = e.view_u8(&kvl.k, avail * kvl.k_tok_bytes);
20773 let vp = e.view_u8(&kvl.v, avail * kvl.v_tok_bytes);
20774 e.i32_set_k(&mut kvl.len_d, (avail - 1) as i32)?;
20775 e.fa_decode_rows(
20776 &q_one,
20777 &kp,
20778 &vp,
20779 &mut a_one,
20780 hd,
20781 nh,
20782 nkv,
20783 avail - 1,
20784 1,
20785 scale,
20786 kvl.k_tok_bytes,
20787 kvl.v_tok_bytes,
20788 Some((&kvl.len_d, 0)),
20789 false,
20790 false,
20791 None,
20792 )?;
20793 } else {
20794 e.fa_decode_kvmod(
20795 &q_one,
20796 &k_view,
20797 &v_view,
20798 &mut a_one,
20799 hd,
20800 nh,
20801 nkv,
20802 t_kv,
20803 scale,
20804 kvl.k_tok_bytes,
20805 kvl.v_tok_bytes,
20806 swa && crate::Engine::wkv_on(),
20807 )?;
20808 }
20809 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
20810 }
20811 e.matmul(&fa.wo, &attn, t)
20812 }
20813
20814 pub(crate) fn gemma4_decode_step_h(
20817 &self,
20818 e: &Engine,
20819 token: u32,
20820 cache: &mut Cache,
20821 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
20822 if let Some(split) = crate::pp::pp2_split(self.layers.len()) {
20827 let rt = crate::pp::Pp2Rt::get(e)?;
20828 let _walk = rt.acquire_walk("gemma4_decode_step_h_pp2")?;
20829 return self.gemma4_decode_step_h_pp2(e, token, cache, split);
20830 }
20831 if crate::pp::pp_cuts(self.layers.len()).is_some() {
20832 crate::pp::warn_unwired_once("gemma4 eager decode (N>2)");
20833 }
20834 let n_embd = self.cfg.n_embd as usize;
20835 let eps = self.cfg.rms_eps;
20836 let pos_d = e.htod_i32(&[cache.pos as i32])?;
20837 let mut x = e.htod(&self.embd.gather(n_embd, &[token]))?;
20838 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
20839 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
20842 let n_layers = self.layers.len();
20843 for (il, layer) in self.layers.iter().enumerate() {
20844 let (hq, hdq) = match h_carry.take() {
20845 Some(p) => p,
20846 None => {
20847 e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, 1, eps)?
20848 }
20849 };
20850 let Mixer::Full(fa) = &layer.mixer else {
20851 panic!("gemma4 layer {il} not full-attn")
20852 };
20853 let o = self.gemma4_decode_attn(e, fa, il, &hq, &hdq, &pos_d, cache)?;
20854 let next_norm = if il + 1 < n_layers {
20855 Some(self.layers[il + 1].attn_norm.float_data())
20856 } else {
20857 None
20858 };
20859 let (xn, hn) = self.gemma4_layer_tail_add_nq_pn(e, layer, &o, &x, 1, next_norm)?;
20860 x = xn;
20861 h_carry = hn;
20862 }
20863 let mut hn = e.uninit(n_embd)?;
20864 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
20865 let h_seed = e.clone_dtod(&x)?;
20866 let mut ld = e.matmul(&self.output, &hn, 1)?;
20867 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
20868 e.softcap(&mut ld, cap, self.output.out_features())?; self.gemma4_suppress(e, &mut ld, 1)?;
20870 let logits = e.dtoh(&ld)?;
20871 cache.pos += 1;
20872 Ok((logits, h_seed))
20873 }
20874
20875 fn gemma4_decode_layers(
20883 &self,
20884 e: &Engine,
20885 mut x: CudaSlice<f32>,
20886 lo: usize,
20887 hi: usize,
20888 pos_d: &CudaSlice<i32>,
20889 cache: &mut Cache,
20890 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
20891 let n_embd = self.cfg.n_embd as usize;
20892 let eps = self.cfg.rms_eps;
20893 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
20894 for il in lo..hi {
20895 let layer = &self.layers[il];
20896 let (hq, hdq) = match h_carry.take() {
20897 Some(p) => p,
20898 None => {
20900 e.rms_norm_q8_1(&x, self.layers[il].attn_norm.float_data(), n_embd, 1, eps)?
20901 }
20902 };
20903 let Mixer::Full(fa) = &layer.mixer else {
20904 panic!("gemma4 layer {il} not full-attn")
20905 };
20906 let o = self.gemma4_decode_attn(e, fa, il, &hq, &hdq, pos_d, cache)?;
20907 let next_norm = if il + 1 < hi {
20908 Some(self.layers[il + 1].attn_norm.float_data())
20909 } else {
20910 None
20911 };
20912 let (xn, hn) = self.gemma4_layer_tail_add_nq_pn(e, layer, &o, &x, 1, next_norm)?;
20913 x = xn;
20914 h_carry = hn;
20915 }
20916 Ok(x)
20917 }
20918
20919 fn gemma4_decode_step_h_pp2(
20927 &self,
20928 e: &Engine,
20929 token: u32,
20930 cache: &mut Cache,
20931 split: usize,
20932 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
20933 if crate::pp::pp2_streams_off() {
20934 return self.gemma4_decode_step_h_pp2_samestream(e, token, cache, split);
20935 }
20936 let rt = crate::pp::Pp2Rt::get(e)?;
20937 let e0 = rt.engine(0, e);
20938 let e1 = rt.engine(1, e);
20939 let n_embd = self.cfg.n_embd as usize;
20940 let eps = self.cfg.rms_eps;
20941 let pos = cache.pos as i32;
20942
20943 let slot = {
20945 let _st0 = rt.enter(0);
20946 let pos_d = e0.htod_i32(&[pos])?;
20947 #[cfg(debug_assertions)]
20948 crate::debug_assert_tensor_stream_device(
20949 &pos_d,
20950 &e0.stream(),
20951 "gemma4_decode_step_h_pp2.stage0.pos_d",
20952 );
20953 let mut x = e0.htod(&self.embd.gather(n_embd, &[token]))?;
20954 e0.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
20955 let x = self.gemma4_decode_layers(e0, x, 0, split, &pos_d, cache)?;
20956 rt.tx(0, &x, n_embd)?
20957 };
20958
20959 let _st1 = rt.enter(1);
20961 let pos_d = e1.htod_i32(&[pos])?;
20962 #[cfg(debug_assertions)]
20963 crate::debug_assert_tensor_stream_device(
20964 &pos_d,
20965 &e1.stream(),
20966 "gemma4_decode_step_h_pp2.stage1.pos_d",
20967 );
20968 let x = rt.rx(0, slot, n_embd)?;
20969 let x = self.gemma4_decode_layers(e1, x, split, self.layers.len(), &pos_d, cache)?;
20970
20971 let mut hn = e1.uninit(n_embd)?;
20972 e1.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
20973 let h_seed = e1.clone_dtod(&x)?;
20974 let mut ld = e1.matmul(&self.output, &hn, 1)?;
20975 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
20976 e1.softcap(&mut ld, cap, self.output.out_features())?;
20977 self.gemma4_suppress(e1, &mut ld, 1)?;
20978 let logits = e1.dtoh(&ld)?;
20979 cache.pos += 1;
20980 Ok((logits, h_seed))
20981 }
20982
20983 fn gemma4_decode_step_h_pp2_samestream(
20986 &self,
20987 e: &Engine,
20988 token: u32,
20989 cache: &mut Cache,
20990 split: usize,
20991 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
20992 let n_embd = self.cfg.n_embd as usize;
20993 let eps = self.cfg.rms_eps;
20994 let pos_d = e.htod_i32(&[cache.pos as i32])?;
20995
20996 let mut x = e.htod(&self.embd.gather(n_embd, &[token]))?;
20998 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
20999 let x = self.gemma4_decode_layers(e, x, 0, split, &pos_d, cache)?;
21000
21001 let boundary_tx = e.clone_dtod(&x)?;
21003 let boundary_rx = e.clone_dtod(&boundary_tx)?;
21004
21005 let x =
21007 self.gemma4_decode_layers(e, boundary_rx, split, self.layers.len(), &pos_d, cache)?;
21008
21009 let mut hn = e.uninit(n_embd)?;
21010 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
21011 let h_seed = e.clone_dtod(&x)?;
21012 let mut ld = e.matmul(&self.output, &hn, 1)?;
21013 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
21014 e.softcap(&mut ld, cap, self.output.out_features())?;
21015 self.gemma4_suppress(e, &mut ld, 1)?;
21016 let logits = e.dtoh(&ld)?;
21017 cache.pos += 1;
21018 Ok((logits, h_seed))
21019 }
21020}
21021
21022impl HybridModel {
21041 pub(crate) fn step35_geom(&self, il: usize) -> memra_gguf::config::LayerGeometry {
21044 let geometry = self
21045 .cfg
21046 .layer_geometry(il as u32)
21047 .unwrap_or_else(|| panic!("step35 layer {il} has no geometry-table row"));
21048 debug_assert_eq!(
21049 geometry.attention_gate,
21050 memra_gguf::config::AttentionGateKind::SeparateHead
21051 );
21052 geometry
21053 }
21054
21055 #[allow(clippy::too_many_arguments)]
21115 fn step35_attn_pre_wo(
21116 &self,
21117 e: &Engine,
21118 fa: &FullAttnLayer,
21119 mut g3: Vec<CudaSlice<f32>>,
21120 hg: Option<&CudaSlice<f32>>,
21121 gt_pre: Option<&CudaSlice<f32>>,
21122 pos_d: &CudaSlice<i32>,
21123 t: usize,
21124 cache: Option<&mut Cache>,
21125 il: usize,
21126 seq_end: usize,
21127 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
21128 let geometry = self.cfg.full_attention_geometry_at(il as u32);
21129 let hd = geometry.head_dim_k as usize;
21130 let nkv = geometry.n_head_kv as usize;
21131 let nh = geometry.n_head as usize;
21132 let rbase = geometry.rope_base;
21133 let scale = geometry.attention_scale();
21134 let swa = geometry.window.is_some();
21135 let eps = self.cfg.rms_eps;
21136 let win = geometry.window.unwrap_or(0) as usize;
21137 let n_rot = geometry.n_rot as usize;
21138
21139 let v = g3.pop().unwrap();
21140 let k0 = g3.pop().unwrap();
21141 let q0 = g3.pop().unwrap();
21142
21143 let mut q = e.uninit(t * nh * hd)?;
21147 e.rms_norm(&q0, fa.q_norm.float_data(), &mut q, hd, nh * t, eps)?;
21148 let mut k = e.uninit(t * nkv * hd)?;
21149 e.rms_norm(&k0, fa.k_norm.float_data(), &mut k, hd, nkv * t, eps)?;
21150 let ff = if geometry.rope_factors {
21151 self.step35_aux.as_ref().and_then(|a| a.rope_freqs(e))
21152 } else {
21153 None
21154 };
21155 #[cfg(debug_assertions)]
21156 if let Some(ff) = ff {
21157 crate::debug_assert_tensor_stream_device(
21158 ff,
21159 &e.stream(),
21160 "step35_attn_pre_wo.rope_freqs",
21161 );
21162 }
21163 e.rope_neox2(&mut q, &mut k, pos_d, hd, n_rot, nh, nkv, t, rbase, 1.0, ff)?;
21164
21165 let mut attn = e.uninit(t * nh * hd)?;
21166 match cache {
21167 Some(cache) => {
21168 let base_len = cache.kv[il].as_ref().unwrap().len;
21169 let legacy_tkv = std::env::var("MEMRA_STEP35_SWA_TKV").as_deref() == Ok("1");
21171 let legacy_calllocal = std::env::var("MEMRA_PRIME_CALLLOCAL").as_deref() == Ok("1");
21172 let off = if swa {
21173 let raw = base_len.saturating_sub(win - 1);
21174 if legacy_tkv || legacy_calllocal {
21175 raw
21176 } else {
21177 raw & !31usize
21178 }
21179 } else {
21180 0
21181 };
21182 {
21183 let kvl = cache.kv[il].as_mut().unwrap();
21184 assert!(kvl.len + t <= cache.max_ctx, "step35 prime: KV overflow");
21185 let write_row = e.prepare_kv_append(kvl, off, t)?;
21186 e.append_kv_quantized_rows(
21187 &k,
21188 &v,
21189 &mut kvl.k,
21190 &mut kvl.v,
21191 write_row,
21192 t,
21193 kvl.kv_dim_k,
21194 kvl.kv_dim_v,
21195 kvl.k_tok_bytes,
21196 kvl.v_tok_bytes,
21197 crate::Engine::kv_fp8_on(),
21198 )?;
21199 kvl.len += t;
21200 let new_len = kvl.len as i32;
21201 e.set_i32_one(&mut kvl.len_d, new_len)?;
21202 }
21203 let kvl = cache.kv[il].as_ref().unwrap();
21204 let t_kv = base_len + t - off;
21227 let physical = kvl.physical_rows(off, off + t_kv)?;
21228 let k_view = e.view_u8_range(
21229 &kvl.k,
21230 physical.start * kvl.k_tok_bytes,
21231 physical.end * kvl.k_tok_bytes,
21232 );
21233 let v_view = e.view_u8_range(
21234 &kvl.v,
21235 physical.start * kvl.v_tok_bytes,
21236 physical.end * kvl.v_tok_bytes,
21237 );
21238 let swa_naive = if legacy_tkv {
21250 t_kv > win
21251 } else {
21252 seq_end > win
21253 };
21254 if swa && swa_naive {
21255 if std::env::var("MEMRA_STEP35_SWA_FA").as_deref() == Ok("0") {
21268 e.sdpa_naive_w_quantized_view(
21269 &q,
21270 &k_view,
21271 &v_view,
21272 &mut attn,
21273 hd,
21274 nh,
21275 nkv,
21276 t,
21277 t_kv,
21278 scale,
21279 true,
21280 win,
21281 kvl.k_tok_bytes,
21282 kvl.v_tok_bytes,
21283 )?;
21284 } else {
21285 e.fa_prefill_view_ws_w_hd128(
21286 &q,
21287 &k_view,
21288 &v_view,
21289 &mut attn,
21290 hd,
21291 nh,
21292 nkv,
21293 t,
21294 t_kv,
21295 scale,
21296 true,
21297 win,
21298 kvl.k_tok_bytes,
21299 kvl.v_tok_bytes,
21300 )?;
21301 }
21302 } else if std::env::var("MEMRA_NOFA").is_ok() {
21303 e.sdpa_naive_quantized_view(
21304 &q,
21305 &k_view,
21306 &v_view,
21307 &mut attn,
21308 hd,
21309 nh,
21310 nkv,
21311 t,
21312 t_kv,
21313 scale,
21314 true,
21315 kvl.k_tok_bytes,
21316 kvl.v_tok_bytes,
21317 )?;
21318 } else {
21319 e.fa_prefill_view_ws(
21324 &q,
21325 &k_view,
21326 &v_view,
21327 &mut attn,
21328 hd,
21329 nh,
21330 nkv,
21331 t,
21332 t_kv,
21333 scale,
21334 true,
21335 kvl.k_tok_bytes,
21336 kvl.v_tok_bytes,
21337 crate::Engine::kv_fp8_on(),
21338 )?;
21339 }
21340 }
21341 None => {
21342 debug_assert_eq!(
21347 seq_end, t,
21348 "step35 cacheless prefill is monolithic (seq_end == t)"
21349 );
21350 if swa && seq_end > win {
21351 e.sdpa_naive_w(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
21352 } else if std::env::var("MEMRA_NOFA").is_ok() {
21353 e.sdpa_naive(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
21354 } else {
21355 e.fa_prefill(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
21356 }
21357 }
21358 }
21359
21360 let gw = fa
21363 .attn_gate
21364 .as_ref()
21365 .ok_or("step35 layer is missing attn_gate.weight (head-wise attention gate)")?;
21366 let gt_owned = if gt_pre.is_none() {
21367 Some(e.matmul(
21368 gw,
21369 hg.ok_or("step35 attention needs hg when gt_pre is absent")?,
21370 t,
21371 )?)
21372 } else {
21373 None
21374 };
21375 let gt = gt_pre.or(gt_owned.as_ref()).unwrap();
21376 let mut ag = e.uninit(t * nh * hd)?;
21377 e.attn_head_gate(&attn, gt, &mut ag, None, hd, nh, t)?;
21378 Ok(ag)
21379 }
21380
21381 pub(crate) fn step35_attn(
21384 &self,
21385 e: &Engine,
21386 fa: &FullAttnLayer,
21387 h: &CudaSlice<f32>,
21388 pos_d: &CudaSlice<i32>,
21389 t: usize,
21390 il: usize,
21391 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
21392 let g3 = match self.full_attn_tp_qkv(e, fa, h, t)? {
21393 Some(g3) => g3,
21394 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
21395 };
21396 let ag = self.step35_attn_pre_wo(e, fa, g3, Some(h), None, pos_d, t, None, il, t)?;
21398 self.full_attn_o(e, fa, &ag, t)
21399 }
21400
21401 #[allow(clippy::too_many_arguments)]
21408 pub(crate) fn step35_attn_prime(
21409 &self,
21410 e: &Engine,
21411 fa: &FullAttnLayer,
21412 h: &CudaSlice<f32>,
21413 hx: Option<&CudaSlice<u8>>,
21414 pos_d: &CudaSlice<i32>,
21415 t: usize,
21416 cache: &mut Cache,
21417 il: usize,
21418 seq_end: usize,
21419 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
21420 if step_tp_prefill_enabled()? && fa.step_tp_qkv.is_some() {
21421 if hx.is_some() {
21422 return Err(
21423 "rank-local Step prefill preserves BF16 activations and refuses the q8_1 \
21424 pre-quantized prime path"
21425 .into(),
21426 );
21427 }
21428 return self.step35_tp_prefill_attn_resident(e, fa, il, h, pos_d, t, cache, seq_end);
21429 }
21430 let g3 = if fa.step_tp_qkv.is_some() {
21431 if hx.is_some() {
21432 return Err(
21433 "Step Q/K/V TP preserves BF16 activations and refuses the q8_1 \
21434 pre-quantized prime path"
21435 .into(),
21436 );
21437 }
21438 self.full_attn_tp_qkv(e, fa, h, t)?
21439 .expect("Step Q/K/V TP disappeared after the presence check")
21440 } else {
21441 match hx {
21442 Some(xh) => e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], h, xh, t)?,
21443 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
21444 }
21445 };
21446 let ag =
21447 self.step35_attn_pre_wo(e, fa, g3, Some(h), None, pos_d, t, Some(cache), il, seq_end)?;
21448 self.full_attn_o(e, fa, &ag, t)
21449 }
21450
21451 fn ensure_step_tp_kv_cache(
21452 &self,
21453 e: &Engine,
21454 fa: &FullAttnLayer,
21455 il: usize,
21456 cache: &mut Cache,
21457 ) -> Result<bool, Box<dyn std::error::Error>> {
21458 let tp = fa
21459 .step_tp_qkv
21460 .as_ref()
21461 .ok_or("Step TP cache hydration lost its resident projections")?;
21462 let geometry = self.cfg.full_attention_geometry_at(il as u32);
21463 let window = geometry.window.map(|window| window as usize);
21464 let ranks = tp.runtime.devices().len();
21465 let head_dim = geometry.head_dim_k as usize;
21466 let kv_heads = geometry.n_head_kv as usize;
21467 let max_ctx = cache.max_ctx;
21468
21469 if cache.tp_kv[il].is_some() {
21470 return Ok(false);
21471 }
21472 let local = cache.kv[il]
21473 .as_ref()
21474 .ok_or_else(|| format!("Step TP layer {il} has no owning-stage KV cache"))?;
21475 if local.kv_dim_k != kv_heads * head_dim || local.kv_dim_v != kv_heads * head_dim {
21476 return Err(format!(
21477 "Step TP layer {il} local KV geometry k={} v={} != {}",
21478 local.kv_dim_k,
21479 local.kv_dim_v,
21480 kv_heads * head_dim
21481 )
21482 .into());
21483 }
21484 let resident_start = window
21485 .map(|window| local.len.saturating_sub(window.saturating_sub(1)) & !31usize)
21486 .unwrap_or(0);
21487 let resident_rows = local.len - resident_start;
21488 let physical = local.physical_rows(resident_start, local.len)?;
21489 let k_rows = if resident_rows == 0 {
21490 Vec::new()
21491 } else {
21492 e.dtoh_u8_view(&e.view_u8_range(
21493 &local.k,
21494 physical.start * local.k_tok_bytes,
21495 physical.end * local.k_tok_bytes,
21496 ))?
21497 };
21498 let v_rows = if resident_rows == 0 {
21499 Vec::new()
21500 } else {
21501 e.dtoh_u8_view(&e.view_u8_range(
21502 &local.v,
21503 physical.start * local.v_tok_bytes,
21504 physical.end * local.v_tok_bytes,
21505 ))?
21506 };
21507 let mut distributed = match window {
21508 Some(window) => tp.runtime.allocate_tp_swa_kv_cache(
21509 kv_heads * head_dim,
21510 kv_heads * head_dim,
21511 max_ctx,
21512 window,
21513 )?,
21514 None => tp.runtime.allocate_tp_kv_cache(
21515 kv_heads * head_dim,
21516 kv_heads * head_dim,
21517 max_ctx,
21518 )?,
21519 };
21520 if distributed.k_tok_bytes() * ranks != local.k_tok_bytes
21521 || distributed.v_tok_bytes() * ranks != local.v_tok_bytes
21522 {
21523 return Err(format!(
21524 "Step TP layer {il} distributed/local KV token bytes disagree: \
21525 k={}x{ranks}/{} v={}x{ranks}/{}",
21526 distributed.k_tok_bytes(),
21527 local.k_tok_bytes,
21528 distributed.v_tok_bytes(),
21529 local.v_tok_bytes,
21530 )
21531 .into());
21532 }
21533 tp.runtime.hydrate_tp_kv_cache_from(
21534 &mut distributed,
21535 local.len,
21536 resident_start,
21537 &k_rows,
21538 &v_rows,
21539 )?;
21540 cache.tp_kv[il] = Some(distributed);
21541 Ok(true)
21542 }
21543
21544 #[allow(clippy::too_many_arguments)]
21545 #[allow(clippy::manual_is_multiple_of)] fn step35_tp_prefill_attn_resident(
21547 &self,
21548 e: &Engine,
21549 fa: &FullAttnLayer,
21550 il: usize,
21551 h: &CudaSlice<f32>,
21552 pos_d: &CudaSlice<i32>,
21553 tokens: usize,
21554 cache: &mut Cache,
21555 seq_end: usize,
21556 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
21557 let tp = fa
21558 .step_tp_qkv
21559 .as_ref()
21560 .ok_or("Step TP prefill lost its resident projections")?;
21561 let attention = tp
21562 .attention
21563 .as_ref()
21564 .ok_or("Step TP prefill lost its resident attention auxiliaries")?;
21565 let ranks = tp.runtime.devices().len();
21566 if !step_tp_prefill_shape(
21567 true,
21568 tokens,
21569 ranks,
21570 tp.runtime.native_p2p(),
21571 true,
21572 crate::Engine::kv_fp8_on(),
21573 ) {
21574 return Err(format!(
21575 "rank-local Step prefill requires tokens>={PRIME_MIN_T}, TP2/TP4 native P2P, \
21576 rank-local attention, and q8_0/q5_1 KV; got tokens={tokens} ranks={ranks} \
21577 native_p2p={} fp8_kv={}",
21578 tp.runtime.native_p2p(),
21579 crate::Engine::kv_fp8_on(),
21580 )
21581 .into());
21582 }
21583 for seam in [
21584 "MEMRA_STEP35_SWA_TKV",
21585 "MEMRA_PRIME_CALLLOCAL",
21586 "MEMRA_PRIME_F32CHUNK0",
21587 ] {
21588 if std::env::var(seam).as_deref() == Ok("1") {
21589 return Err(format!(
21590 "rank-local Step prefill has not qualified the legacy seam {seam}=1"
21591 )
21592 .into());
21593 }
21594 }
21595
21596 let geometry = self.cfg.full_attention_geometry_at(il as u32);
21597 let window = geometry.window.map(|window| window as usize);
21598 let head_dim = geometry.head_dim_k as usize;
21599 let heads = geometry.n_head as usize;
21600 let kv_heads = geometry.n_head_kv as usize;
21601 if heads % ranks != 0 || kv_heads % ranks != 0 {
21602 return Err(format!(
21603 "Step attention heads q={heads} kv={kv_heads} are not divisible by TP={ranks}"
21604 )
21605 .into());
21606 }
21607 let local_heads = heads / ranks;
21608 let local_kv_heads = kv_heads / ranks;
21609 let local_kv_dim = local_kv_heads * head_dim;
21610 let hidden = self.cfg.n_embd as usize;
21611 let expected_input = tokens
21612 .checked_mul(hidden)
21613 .ok_or("Step TP prefill input size overflow")?;
21614 if h.len() < expected_input {
21615 return Err(format!(
21616 "Step TP prefill input {} is shorter than {tokens}x{hidden}",
21617 h.len()
21618 )
21619 .into());
21620 }
21621 let positions = e.dtoh_i32(pos_d)?;
21622 if positions.len() != tokens {
21623 return Err(format!(
21624 "rank-local Step prefill positions {} != tokens {tokens}",
21625 positions.len()
21626 )
21627 .into());
21628 }
21629
21630 let mut active_input = e.uninit(expected_input)?;
21631 e.copy_view_into(
21632 &mut active_input,
21633 0,
21634 &h.slice(0..expected_input),
21635 expected_input,
21636 )?;
21637 let mut input = tp.runtime.allocate_replicated_device_rows(tokens, hidden)?;
21638 e.stream().synchronize()?;
21643 tp.runtime
21644 .refresh_replicated_device_rows_from_root(&mut input, &active_input)?;
21645 let q_raw = tp
21646 .runtime
21647 .bf16_column_parallel_resident_replicated_device_shards(&tp.q, &input)?;
21648 let k_raw = tp
21649 .runtime
21650 .bf16_column_parallel_resident_replicated_device_shards(&tp.k, &input)?;
21651 let v_raw = tp
21652 .runtime
21653 .bf16_column_parallel_resident_replicated_device_shards(&tp.v, &input)?;
21654 let mut q = Vec::with_capacity(ranks);
21655 let mut k = Vec::with_capacity(ranks);
21656 for rank in 0..ranks {
21657 let engine = tp
21658 .runtime
21659 .rank_engine(rank)
21660 .ok_or_else(|| format!("Step TP layer {il} has no engine for rank {rank}"))?;
21661 let _main = engine.gpu.enter_main()?;
21662 let mut q_rank = engine.uninit(tokens * local_heads * head_dim)?;
21663 engine.rms_norm(
21664 &q_raw[rank],
21665 &attention.q_norm[rank],
21666 &mut q_rank,
21667 head_dim,
21668 tokens * local_heads,
21669 self.cfg.rms_eps,
21670 )?;
21671 let mut k_rank = engine.uninit(tokens * local_kv_dim)?;
21672 engine.rms_norm(
21673 &k_raw[rank],
21674 &attention.k_norm[rank],
21675 &mut k_rank,
21676 head_dim,
21677 tokens * local_kv_heads,
21678 self.cfg.rms_eps,
21679 )?;
21680 let position = engine.htod_i32(&positions)?;
21681 let rope_freqs = if geometry.rope_factors {
21682 self.step35_aux
21683 .as_ref()
21684 .and_then(|aux| aux.rope_freqs(engine))
21685 } else {
21686 None
21687 };
21688 engine.rope_neox2(
21689 &mut q_rank,
21690 &mut k_rank,
21691 &position,
21692 head_dim,
21693 geometry.n_rot as usize,
21694 local_heads,
21695 local_kv_heads,
21696 tokens,
21697 geometry.rope_base,
21698 1.0,
21699 rope_freqs,
21700 )?;
21701 q.push(q_rank);
21702 k.push(k_rank);
21703 }
21704
21705 let gate_weight = fa
21706 .attn_gate
21707 .as_ref()
21708 .ok_or("step35 layer is missing attn_gate.weight")?;
21709 let gate = e.dtoh(&e.matmul(gate_weight, h, tokens)?)?;
21710 if gate.len() != tokens * heads {
21711 return Err(format!(
21712 "Step TP layer {il} gate output {} != {tokens}x{heads}",
21713 gate.len()
21714 )
21715 .into());
21716 }
21717
21718 let hydrated = self.ensure_step_tp_kv_cache(e, fa, il, cache)?;
21719 let base_len = cache.kv[il]
21720 .as_ref()
21721 .ok_or_else(|| format!("Step TP layer {il} lost its local KV cache"))?
21722 .len;
21723 let distributed = cache.tp_kv[il]
21724 .as_ref()
21725 .ok_or_else(|| format!("Step TP layer {il} lost its distributed KV cache"))?;
21726 if distributed.committed_len() != base_len || distributed.staged_len() != base_len {
21727 return Err(format!(
21728 "Step TP layer {il} cache lengths diverged before prefill: \
21729 local={base_len} distributed={}/{}",
21730 distributed.committed_len(),
21731 distributed.staged_len()
21732 )
21733 .into());
21734 }
21735 let target_len = base_len
21736 .checked_add(tokens)
21737 .ok_or("Step TP prefill cache length overflow")?;
21738 if target_len > cache.max_ctx {
21739 return Err(format!(
21740 "Step TP layer {il} prefill exceeds cache: {base_len}+{tokens}>{}",
21741 cache.max_ctx
21742 )
21743 .into());
21744 }
21745 if seq_end < target_len {
21746 return Err(format!(
21747 "Step TP layer {il} request end {seq_end} precedes chunk end {target_len}"
21748 )
21749 .into());
21750 }
21751
21752 let transaction = cache.tp_kv[il]
21753 .as_mut()
21754 .expect("distributed cache checked above")
21755 .begin_transaction()?;
21756 if let Err(error) = tp.runtime.append_tp_kv_transaction(
21757 cache.tp_kv[il]
21758 .as_mut()
21759 .expect("distributed cache checked above"),
21760 transaction,
21761 &k,
21762 &v_raw,
21763 tokens,
21764 ) {
21765 let _ = tp.runtime.rollback_tp_kv_transaction(
21766 cache.tp_kv[il]
21767 .as_mut()
21768 .expect("distributed cache checked above"),
21769 transaction,
21770 );
21771 return Err(error);
21772 }
21773
21774 let staged = (|| -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
21775 let distributed = cache.tp_kv[il]
21776 .as_ref()
21777 .expect("distributed cache checked above");
21778 let staged_len = distributed.staged_len();
21779 let view_start = window
21780 .map(|window| base_len.saturating_sub(window.saturating_sub(1)) & !31usize)
21781 .unwrap_or(0);
21782 let physical = distributed.physical_range(view_start, staged_len)?;
21783 let t_kv = staged_len - view_start;
21784 let swa_naive = window.is_some_and(|window| seq_end > window);
21785 let mut gated = Vec::with_capacity(ranks);
21786 #[allow(clippy::needless_range_loop)]
21787 for rank in 0..ranks {
21789 let engine = tp
21790 .runtime
21791 .rank_engine(rank)
21792 .ok_or_else(|| format!("Step TP layer {il} has no engine for rank {rank}"))?;
21793 let _main = engine.gpu.enter_main()?;
21794 let rank_cache = distributed
21795 .rank(rank)
21796 .ok_or_else(|| format!("Step TP layer {il} has no KV cache rank {rank}"))?;
21797 let k_view = engine.view_u8_range(
21798 rank_cache.k(),
21799 physical.start * distributed.k_tok_bytes(),
21800 physical.end * distributed.k_tok_bytes(),
21801 );
21802 let v_view = engine.view_u8_range(
21803 rank_cache.v(),
21804 physical.start * distributed.v_tok_bytes(),
21805 physical.end * distributed.v_tok_bytes(),
21806 );
21807 let mut attention_out = engine.uninit(tokens * local_heads * head_dim)?;
21808 if swa_naive {
21809 let window = window.expect("SWA predicate requires a window");
21810 if std::env::var("MEMRA_STEP35_SWA_FA").as_deref() == Ok("0") {
21811 engine.sdpa_naive_w_quantized_view(
21812 &q[rank],
21813 &k_view,
21814 &v_view,
21815 &mut attention_out,
21816 head_dim,
21817 local_heads,
21818 local_kv_heads,
21819 tokens,
21820 t_kv,
21821 geometry.attention_scale(),
21822 true,
21823 window,
21824 distributed.k_tok_bytes(),
21825 distributed.v_tok_bytes(),
21826 )?;
21827 } else {
21828 engine.fa_prefill_view_ws_w_hd128(
21829 &q[rank],
21830 &k_view,
21831 &v_view,
21832 &mut attention_out,
21833 head_dim,
21834 local_heads,
21835 local_kv_heads,
21836 tokens,
21837 t_kv,
21838 geometry.attention_scale(),
21839 true,
21840 window,
21841 distributed.k_tok_bytes(),
21842 distributed.v_tok_bytes(),
21843 )?;
21844 }
21845 } else if std::env::var("MEMRA_NOFA").is_ok() {
21846 engine.sdpa_naive_quantized_view(
21847 &q[rank],
21848 &k_view,
21849 &v_view,
21850 &mut attention_out,
21851 head_dim,
21852 local_heads,
21853 local_kv_heads,
21854 tokens,
21855 t_kv,
21856 geometry.attention_scale(),
21857 true,
21858 distributed.k_tok_bytes(),
21859 distributed.v_tok_bytes(),
21860 )?;
21861 } else {
21862 engine.fa_prefill_view_ws(
21863 &q[rank],
21864 &k_view,
21865 &v_view,
21866 &mut attention_out,
21867 head_dim,
21868 local_heads,
21869 local_kv_heads,
21870 tokens,
21871 t_kv,
21872 geometry.attention_scale(),
21873 true,
21874 distributed.k_tok_bytes(),
21875 distributed.v_tok_bytes(),
21876 false,
21877 )?;
21878 }
21879
21880 let gate_start = rank * local_heads;
21881 let mut gate_rank = Vec::with_capacity(tokens * local_heads);
21882 for token in 0..tokens {
21883 let start = token * heads + gate_start;
21884 gate_rank.extend_from_slice(&gate[start..start + local_heads]);
21885 }
21886 let gate_rank = engine.htod(&gate_rank)?;
21887 let mut gated_rank = engine.uninit(tokens * local_heads * head_dim)?;
21888 engine.attn_head_gate(
21889 &attention_out,
21890 &gate_rank,
21891 &mut gated_rank,
21892 None,
21893 head_dim,
21894 local_heads,
21895 tokens,
21896 )?;
21897 gated.push(gated_rank);
21898 }
21899 for rank in 1..ranks {
21900 let engine = tp
21901 .runtime
21902 .rank_engine(rank)
21903 .ok_or_else(|| format!("Step TP layer {il} has no engine for rank {rank}"))?;
21904 let _main = engine.gpu.enter_main()?;
21905 engine.stream().synchronize()?;
21906 }
21907
21908 let (output, k_shadow, v_shadow) = if tp.runtime.bulk_p2p() {
21909 let output = tp
21910 .runtime
21911 .step_bf16_row_parallel_resident_root_device(&tp.o, &gated, tokens)?;
21912 let k_shadow =
21913 tp.runtime
21914 .gather_native_column_shards_device(&k, tokens, local_kv_dim)?;
21915 let v_shadow =
21916 tp.runtime
21917 .gather_native_column_shards_device(&v_raw, tokens, local_kv_dim)?;
21918 let root = tp
21919 .runtime
21920 .rank_engine(0)
21921 .ok_or("Step TP prefill lost its root engine")?;
21922 let _main = root.gpu.enter_main()?;
21923 root.stream().synchronize()?;
21924 (output, k_shadow, v_shadow)
21925 } else {
21926 let attention = tp.runtime.gather_native_column_shards(
21927 &gated,
21928 tokens,
21929 local_heads * head_dim,
21930 )?;
21931 let output = tp
21932 .runtime
21933 .step_bf16_row_parallel_resident_native(&tp.o, &attention, tokens)?;
21934 let k_shadow = tp
21935 .runtime
21936 .gather_native_column_shards(&k, tokens, local_kv_dim)?;
21937 let v_shadow =
21938 tp.runtime
21939 .gather_native_column_shards(&v_raw, tokens, local_kv_dim)?;
21940 (e.htod(&output)?, e.htod(&k_shadow)?, e.htod(&v_shadow)?)
21941 };
21942 let local = cache.kv[il]
21943 .as_mut()
21944 .ok_or_else(|| format!("Step TP layer {il} lost its local KV cache"))?;
21945 if local.len != base_len {
21946 return Err(format!(
21947 "Step TP layer {il} local cache changed during prefill: \
21948 len={} base={base_len}",
21949 local.len
21950 )
21951 .into());
21952 }
21953 let retain_from = window
21954 .map(|window| {
21955 let staged_retain = staged_len.saturating_sub(window) & !31usize;
21956 let rollback_retain =
21957 base_len.saturating_sub(window.saturating_sub(1)) & !31usize;
21958 staged_retain.min(rollback_retain)
21959 })
21960 .unwrap_or(0);
21961 let write_row = e.prepare_kv_append(local, retain_from, tokens)?;
21962 e.append_kv_quantized_rows(
21963 &k_shadow,
21964 &v_shadow,
21965 &mut local.k,
21966 &mut local.v,
21967 write_row,
21968 tokens,
21969 local.kv_dim_k,
21970 local.kv_dim_v,
21971 local.k_tok_bytes,
21972 local.v_tok_bytes,
21973 false,
21974 )?;
21975 local.len = staged_len;
21976 e.set_i32_one(&mut local.len_d, staged_len as i32)?;
21977 Ok(output)
21978 })();
21979
21980 let output = match staged {
21981 Ok(output) => output,
21982 Err(error) => {
21983 let _ = tp.runtime.rollback_tp_kv_transaction(
21984 cache.tp_kv[il]
21985 .as_mut()
21986 .expect("distributed cache checked above"),
21987 transaction,
21988 );
21989 if let Some(local) = cache.kv[il].as_mut() {
21990 local.len = base_len;
21991 let _ = e.set_i32_one(&mut local.len_d, base_len as i32);
21992 }
21993 return Err(error);
21994 }
21995 };
21996 if let Err(error) = tp.runtime.commit_tp_kv_transaction(
21997 cache.tp_kv[il]
21998 .as_mut()
21999 .expect("distributed cache checked above"),
22000 transaction,
22001 tokens,
22002 ) {
22003 let _ = tp.runtime.rollback_tp_kv_transaction(
22004 cache.tp_kv[il]
22005 .as_mut()
22006 .expect("distributed cache checked above"),
22007 transaction,
22008 );
22009 let local = cache.kv[il].as_mut().expect("local cache checked above");
22010 local.len = base_len;
22011 e.set_i32_one(&mut local.len_d, base_len as i32)?;
22012 return Err(error);
22013 }
22014
22015 let committed = cache.tp_kv[il]
22016 .as_ref()
22017 .expect("distributed cache checked above")
22018 .committed_len();
22019 let local_len = cache.kv[il]
22020 .as_ref()
22021 .expect("local cache checked above")
22022 .len;
22023 if committed != local_len {
22024 return Err(format!(
22025 "Step TP layer {il} committed cache length {committed} != local shadow {local_len}"
22026 )
22027 .into());
22028 }
22029 eprintln!(
22030 "[step-tp-prefill-attn] execute layer={} devices={:?} tokens={tokens} \
22031 qkv_tensor_parallel=true qk_norm_rank_local=true rope_rank_local=true \
22032 kv_cache_distributed=true kv_cache_hydrated={} attention_tensor_parallel=true \
22033 attention_scope={} input_path=root-device-replicated gate_tensor_parallel=false \
22034 gate_shards=host-canonical o_tensor_parallel=true local_cache_shadow=true \
22035 cache_commit=chunk transport={} native_p2p=true bulk_p2p={} \
22036 output={} performance_claim=false",
22037 tp.layer,
22038 tp.devices,
22039 hydrated,
22040 if window.is_some() {
22041 "rank-local-swa-ring"
22042 } else {
22043 "rank-local-global"
22044 },
22045 tp.runtime.transport_label(),
22046 tp.runtime.bulk_p2p(),
22047 if tp.runtime.bulk_p2p() {
22048 "root-device"
22049 } else {
22050 "root-readback"
22051 },
22052 );
22053 Ok(output)
22054 }
22055
22056 fn step35_tp_decode_attn_resident(
22057 &self,
22058 e: &Engine,
22059 fa: &FullAttnLayer,
22060 il: usize,
22061 h: &CudaSlice<f32>,
22062 pos_d: &CudaSlice<i32>,
22063 cache: &mut Cache,
22064 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
22065 static ATTN_NS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
22069 static ATTN_CALLS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
22070 let timing = std::env::var("MEMRA_STEP_TP_TIMING").as_deref() == Ok("1");
22071 let started = timing.then(std::time::Instant::now);
22072 let result = if crate::tp::step_tp_decode_v2_enabled()? {
22073 self.step35_tp_decode_attn_resident_v2(e, fa, il, h, pos_d, cache)
22074 } else {
22075 self.step35_tp_decode_attn_resident_inner(e, fa, il, h, pos_d, cache)
22076 };
22077 if let Some(started) = started {
22078 use std::sync::atomic::Ordering;
22079 let ns = ATTN_NS.fetch_add(started.elapsed().as_nanos() as u64, Ordering::Relaxed)
22080 + started.elapsed().as_nanos() as u64;
22081 let calls = ATTN_CALLS.fetch_add(1, Ordering::Relaxed) + 1;
22082 if calls.is_multiple_of(430) {
22083 eprintln!(
22084 "[step-tp-attn-timing] calls={calls} total_ms={:.1} avg_us={:.1}",
22085 ns as f64 / 1.0e6,
22086 ns as f64 / calls as f64 / 1.0e3,
22087 );
22088 }
22089 }
22090 result
22091 }
22092
22093 #[allow(clippy::too_many_arguments)]
22094 fn step35_tp_decode_attn_resident_inner(
22095 &self,
22096 e: &Engine,
22097 fa: &FullAttnLayer,
22098 il: usize,
22099 h: &CudaSlice<f32>,
22100 pos_d: &CudaSlice<i32>,
22101 cache: &mut Cache,
22102 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
22103 static T_POS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
22108 static T_QKV: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
22109 static T_NORMROPE: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
22110 static T_GATE: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
22111 static T_APPEND: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
22112 static T_ATTN: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
22113 static T_OPROJ: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
22114 static T_SHADOW: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
22115 static T_PHASE_CALLS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
22116 let timing = std::env::var("MEMRA_STEP_TP_TIMING").as_deref() == Ok("1");
22117 #[allow(clippy::manual_is_multiple_of)] fn lap(
22119 runtime: &crate::tp::TpE4m3HostBounce,
22120 e: &Engine,
22121 timer: &std::sync::atomic::AtomicU64,
22122 started: &mut Option<std::time::Instant>,
22123 ) -> Result<(), Box<dyn std::error::Error>> {
22124 let Some(start) = started.as_mut() else {
22125 return Ok(());
22126 };
22127 for rank in 0..runtime.devices().len() {
22128 if let Some(engine) = runtime.rank_engine(rank) {
22129 let _main = engine.gpu.enter_main()?;
22130 engine.stream().synchronize()?;
22131 }
22132 }
22133 e.stream().synchronize()?;
22134 timer.fetch_add(
22135 start.elapsed().as_nanos() as u64,
22136 std::sync::atomic::Ordering::Relaxed,
22137 );
22138 *start = std::time::Instant::now();
22139 Ok(())
22140 }
22141 let tp = fa
22142 .step_tp_qkv
22143 .as_ref()
22144 .ok_or("Step TP decode lost its resident projections")?;
22145 let attention = tp
22146 .attention
22147 .as_ref()
22148 .ok_or("Step TP decode lost its resident attention auxiliaries")?;
22149 if !tp.runtime.native_p2p() {
22150 return Err("rank-local Step attention requires native P2P".into());
22151 }
22152 if crate::Engine::kv_fp8_on() {
22153 return Err("rank-local Step attention has not qualified the FP8 KV cache".into());
22154 }
22155
22156 let geometry = self.step35_geom(il);
22157 let window = geometry.window.map(|window| window as usize);
22158 let ranks = tp.runtime.devices().len();
22159 let head_dim = geometry.head_dim_k as usize;
22160 let heads = geometry.n_head as usize;
22161 let kv_heads = geometry.n_head_kv as usize;
22162 if !heads.is_multiple_of(ranks) || !kv_heads.is_multiple_of(ranks) {
22163 return Err(format!(
22164 "Step attention heads q={heads} kv={kv_heads} are not divisible by TP={ranks}"
22165 )
22166 .into());
22167 }
22168 let local_heads = heads / ranks;
22169 let local_kv_heads = kv_heads / ranks;
22170 let local_kv_dim = local_kv_heads * head_dim;
22171 let max_ctx = cache.max_ctx;
22172
22173 let hydrated = self.ensure_step_tp_kv_cache(e, fa, il, cache)?;
22174
22175 let base_len = cache.kv[il]
22176 .as_ref()
22177 .ok_or_else(|| format!("Step TP layer {il} lost its local KV cache"))?
22178 .len;
22179 let distributed = cache.tp_kv[il]
22180 .as_ref()
22181 .ok_or_else(|| format!("Step TP layer {il} lost its distributed KV cache"))?;
22182 if distributed.committed_len() != base_len || distributed.staged_len() != base_len {
22183 return Err(format!(
22184 "Step TP layer {il} cache lengths diverged before decode: \
22185 local={base_len} distributed={}/{}",
22186 distributed.committed_len(),
22187 distributed.staged_len()
22188 )
22189 .into());
22190 }
22191
22192 let mut lap_start = timing.then(std::time::Instant::now);
22193 let positions = e.dtoh_i32(pos_d)?;
22194 if positions.len() != 1 {
22195 return Err(format!(
22196 "rank-local Step decode requires one position, got {}",
22197 positions.len()
22198 )
22199 .into());
22200 }
22201 lap(&tp.runtime, e, &T_POS, &mut lap_start)?;
22202 let (q_raw, k_raw, v_raw, input_path) = if let Some(decode_input) =
22203 attention.decode_input.as_ref()
22204 {
22205 let mut decode_input = decode_input
22206 .lock()
22207 .map_err(|_| "Step TP replicated decode input lock is poisoned")?;
22208 e.stream().synchronize()?;
22212 tp.runtime
22213 .refresh_replicated_device_rows_from_root(&mut decode_input, h)?;
22214 let q_raw = tp
22215 .runtime
22216 .bf16_column_parallel_resident_replicated_device_shards(&tp.q, &decode_input)?;
22217 let k_raw = tp
22218 .runtime
22219 .bf16_column_parallel_resident_replicated_device_shards(&tp.k, &decode_input)?;
22220 let v_raw = tp
22221 .runtime
22222 .bf16_column_parallel_resident_replicated_device_shards(&tp.v, &decode_input)?;
22223 (q_raw, k_raw, v_raw, "root-device-replicated")
22224 } else {
22225 let activation = e.dtoh(h)?;
22226 let q_raw =
22227 tp.runtime
22228 .bf16_column_parallel_resident_device_shards(&tp.q, &activation, 1)?;
22229 let k_raw =
22230 tp.runtime
22231 .bf16_column_parallel_resident_device_shards(&tp.k, &activation, 1)?;
22232 let v_raw =
22233 tp.runtime
22234 .bf16_column_parallel_resident_device_shards(&tp.v, &activation, 1)?;
22235 (q_raw, k_raw, v_raw, "host-replicated")
22236 };
22237 lap(&tp.runtime, e, &T_QKV, &mut lap_start)?;
22238 let mut q = Vec::with_capacity(ranks);
22239 let mut k = Vec::with_capacity(ranks);
22240 for rank in 0..ranks {
22241 let engine = tp
22242 .runtime
22243 .rank_engine(rank)
22244 .ok_or_else(|| format!("Step TP layer {il} has no engine for rank {rank}"))?;
22245 let _main = engine.gpu.enter_main()?;
22246 let mut q_rank = engine.uninit(local_heads * head_dim)?;
22247 engine.rms_norm(
22248 &q_raw[rank],
22249 &attention.q_norm[rank],
22250 &mut q_rank,
22251 head_dim,
22252 local_heads,
22253 self.cfg.rms_eps,
22254 )?;
22255 let mut k_rank = engine.uninit(local_kv_dim)?;
22256 engine.rms_norm(
22257 &k_raw[rank],
22258 &attention.k_norm[rank],
22259 &mut k_rank,
22260 head_dim,
22261 local_kv_heads,
22262 self.cfg.rms_eps,
22263 )?;
22264 let position = engine.htod_i32(&positions)?;
22265 let rope_freqs = if geometry.rope_factors {
22266 self.step35_aux
22267 .as_ref()
22268 .and_then(|aux| aux.rope_freqs(engine))
22269 } else {
22270 None
22271 };
22272 engine.rope_neox2(
22273 &mut q_rank,
22274 &mut k_rank,
22275 &position,
22276 head_dim,
22277 geometry.n_rot as usize,
22278 local_heads,
22279 local_kv_heads,
22280 1,
22281 geometry.rope_base,
22282 1.0,
22283 rope_freqs,
22284 )?;
22285 q.push(q_rank);
22286 k.push(k_rank);
22287 }
22288 lap(&tp.runtime, e, &T_NORMROPE, &mut lap_start)?;
22289
22290 let gate_weight = fa
22291 .attn_gate
22292 .as_ref()
22293 .ok_or("step35 layer is missing attn_gate.weight")?;
22294 let gate = e.matmul(gate_weight, h, 1)?;
22295 let gate = e.dtoh(&gate)?;
22296 if gate.len() != heads {
22297 return Err(format!("Step TP layer {il} gate output {} != {heads}", gate.len()).into());
22298 }
22299 lap(&tp.runtime, e, &T_GATE, &mut lap_start)?;
22300
22301 let transaction = cache.tp_kv[il]
22302 .as_mut()
22303 .expect("distributed cache checked above")
22304 .begin_transaction()?;
22305 if let Err(error) = tp.runtime.append_tp_kv_transaction(
22306 cache.tp_kv[il]
22307 .as_mut()
22308 .expect("distributed cache checked above"),
22309 transaction,
22310 &k,
22311 &v_raw,
22312 1,
22313 ) {
22314 let _ = tp.runtime.rollback_tp_kv_transaction(
22315 cache.tp_kv[il]
22316 .as_mut()
22317 .expect("distributed cache checked above"),
22318 transaction,
22319 );
22320 return Err(error);
22321 }
22322 lap(&tp.runtime, e, &T_APPEND, &mut lap_start)?;
22323
22324 let staged = (|| -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
22325 let distributed = cache.tp_kv[il]
22326 .as_ref()
22327 .expect("distributed cache checked above");
22328 let staged_len = distributed.staged_len();
22329 let view_start = window
22330 .map(|window| staged_len.saturating_sub(window))
22331 .unwrap_or(0);
22332 let physical = distributed.physical_range(view_start, staged_len)?;
22333 let t_kv = staged_len - view_start;
22334 let mut gated = Vec::with_capacity(ranks);
22335 #[allow(clippy::needless_range_loop)]
22336 for rank in 0..ranks {
22338 let engine = tp
22339 .runtime
22340 .rank_engine(rank)
22341 .ok_or_else(|| format!("Step TP layer {il} has no engine for rank {rank}"))?;
22342 let _main = engine.gpu.enter_main()?;
22343 let rank_cache = distributed
22344 .rank(rank)
22345 .ok_or_else(|| format!("Step TP layer {il} has no KV cache rank {rank}"))?;
22346 let k_view = engine.view_u8_range(
22347 rank_cache.k(),
22348 physical.start * distributed.k_tok_bytes(),
22349 physical.end * distributed.k_tok_bytes(),
22350 );
22351 let v_view = engine.view_u8_range(
22352 rank_cache.v(),
22353 physical.start * distributed.v_tok_bytes(),
22354 physical.end * distributed.v_tok_bytes(),
22355 );
22356 let mut attention_out = engine.uninit(local_heads * head_dim)?;
22357 engine.fa_decode_kvmod(
22358 &q[rank],
22359 &k_view,
22360 &v_view,
22361 &mut attention_out,
22362 head_dim,
22363 local_heads,
22364 local_kv_heads,
22365 t_kv,
22366 geometry.attention_scale(),
22367 distributed.k_tok_bytes(),
22368 distributed.v_tok_bytes(),
22369 false,
22370 )?;
22371 let gate_start = rank * local_heads;
22372 let gate_rank = engine.htod(&gate[gate_start..gate_start + local_heads])?;
22373 let mut gated_rank = engine.uninit(local_heads * head_dim)?;
22374 engine.attn_head_gate(
22375 &attention_out,
22376 &gate_rank,
22377 &mut gated_rank,
22378 None,
22379 head_dim,
22380 local_heads,
22381 1,
22382 )?;
22383 gated.push(gated_rank);
22384 }
22385 lap(&tp.runtime, e, &T_ATTN, &mut lap_start)?;
22386
22387 let gathered =
22388 tp.runtime
22389 .gather_native_column_shards(&gated, 1, local_heads * head_dim)?;
22390 let output = tp
22391 .runtime
22392 .step_bf16_row_parallel_resident_native(&tp.o, &gathered, 1)?;
22393 let output = e.htod(&output)?;
22394 lap(&tp.runtime, e, &T_OPROJ, &mut lap_start)?;
22395
22396 let k_shadow = tp
22397 .runtime
22398 .gather_native_column_shards(&k, 1, local_kv_dim)?;
22399 let v_shadow = tp
22400 .runtime
22401 .gather_native_column_shards(&v_raw, 1, local_kv_dim)?;
22402 let k_shadow = e.htod(&k_shadow)?;
22403 let v_shadow = e.htod(&v_shadow)?;
22404 let local = cache.kv[il]
22405 .as_mut()
22406 .ok_or_else(|| format!("Step TP layer {il} lost its local KV cache"))?;
22407 if local.len != base_len || base_len + 1 > max_ctx {
22408 return Err(format!(
22409 "Step TP layer {il} local cache changed during decode: \
22410 len={} base={base_len} max={max_ctx}",
22411 local.len
22412 )
22413 .into());
22414 }
22415 let retain_from = window
22416 .map(|window| {
22417 let staged_retain = (base_len + 1).saturating_sub(window) & !31usize;
22418 let rollback_retain =
22419 base_len.saturating_sub(window.saturating_sub(1)) & !31usize;
22420 staged_retain.min(rollback_retain)
22421 })
22422 .unwrap_or(0);
22423 let write_row = e.prepare_kv_append(local, retain_from, 1)?;
22424 e.append_kv_quantized(
22425 &k_shadow,
22426 &v_shadow,
22427 &mut local.k,
22428 &mut local.v,
22429 write_row,
22430 local.kv_dim_k,
22431 local.kv_dim_v,
22432 local.k_tok_bytes,
22433 local.v_tok_bytes,
22434 false,
22435 )?;
22436 local.len = base_len + 1;
22437 e.set_i32_one(&mut local.len_d, local.len as i32)?;
22438 Ok(output)
22439 })();
22440
22441 let output = match staged {
22442 Ok(output) => output,
22443 Err(error) => {
22444 let _ = tp.runtime.rollback_tp_kv_transaction(
22445 cache.tp_kv[il]
22446 .as_mut()
22447 .expect("distributed cache checked above"),
22448 transaction,
22449 );
22450 if let Some(local) = cache.kv[il].as_mut() {
22451 local.len = base_len;
22452 let _ = e.set_i32_one(&mut local.len_d, base_len as i32);
22453 }
22454 return Err(error);
22455 }
22456 };
22457 if let Err(error) = tp.runtime.commit_tp_kv_transaction(
22458 cache.tp_kv[il]
22459 .as_mut()
22460 .expect("distributed cache checked above"),
22461 transaction,
22462 1,
22463 ) {
22464 let _ = tp.runtime.rollback_tp_kv_transaction(
22465 cache.tp_kv[il]
22466 .as_mut()
22467 .expect("distributed cache checked above"),
22468 transaction,
22469 );
22470 let local = cache.kv[il].as_mut().expect("local cache checked above");
22471 local.len = base_len;
22472 e.set_i32_one(&mut local.len_d, base_len as i32)?;
22473 return Err(error);
22474 }
22475
22476 let committed = cache.tp_kv[il]
22477 .as_ref()
22478 .expect("distributed cache checked above")
22479 .committed_len();
22480 let local_len = cache.kv[il]
22481 .as_ref()
22482 .expect("local cache checked above")
22483 .len;
22484 if committed != local_len {
22485 return Err(format!(
22486 "Step TP layer {il} committed cache length {committed} != local shadow {local_len}"
22487 )
22488 .into());
22489 }
22490 lap(&tp.runtime, e, &T_SHADOW, &mut lap_start)?;
22491 if timing {
22492 use std::sync::atomic::Ordering;
22493 let calls = T_PHASE_CALLS.fetch_add(1, Ordering::Relaxed) + 1;
22494 if calls.is_multiple_of(430) {
22495 let avg = |t: &std::sync::atomic::AtomicU64| {
22496 t.load(Ordering::Relaxed) as f64 / calls as f64 / 1.0e3
22497 };
22498 eprintln!(
22499 "[step-tp-attn-phase] calls={calls} avg_us pos={:.1} qkv={:.1} \
22500 normrope={:.1} gate={:.1} append={:.1} attn={:.1} oproj={:.1} shadow={:.1}",
22501 avg(&T_POS),
22502 avg(&T_QKV),
22503 avg(&T_NORMROPE),
22504 avg(&T_GATE),
22505 avg(&T_APPEND),
22506 avg(&T_ATTN),
22507 avg(&T_OPROJ),
22508 avg(&T_SHADOW),
22509 );
22510 }
22511 }
22512 eprintln!(
22513 "[step-tp-attn] execute layer={} devices={:?} tokens=1 \
22514 qkv_tensor_parallel=true qk_norm_rank_local=true rope_rank_local=true \
22515 kv_cache_distributed=true kv_cache_hydrated={} attention_tensor_parallel=true \
22516 attention_scope={} input_path={} kv_physical_rows={} \
22517 gate_tensor_parallel=false gate_shards=host-canonical o_tensor_parallel=true \
22518 local_cache_shadow=true cache_commit=immediate transport={} native_p2p=true \
22519 bulk_p2p={} output=root-readback performance_claim=false",
22520 tp.layer,
22521 tp.devices,
22522 hydrated,
22523 if window.is_some() {
22524 "rank-local-swa-ring"
22525 } else {
22526 "rank-local-global"
22527 },
22528 input_path,
22529 cache.tp_kv[il]
22530 .as_ref()
22531 .expect("distributed cache checked above")
22532 .physical_capacity(),
22533 tp.runtime.transport_label(),
22534 tp.runtime.bulk_p2p(),
22535 );
22536 Ok(output)
22537 }
22538
22539 #[allow(clippy::too_many_arguments)]
22546 pub(crate) fn step35_verify_qkv_precompute(
22551 &self,
22552 e: &Engine,
22553 il: usize,
22554 h_t: &CudaSlice<f32>,
22555 t: usize,
22556 ) -> Result<bool, Box<dyn std::error::Error>> {
22557 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
22558 return Ok(false);
22559 };
22560 let Some(tp) = fa.step_tp_qkv.as_ref() else {
22561 return Ok(false);
22562 };
22563 let Some(attention) = tp.attention.as_ref() else {
22564 return Ok(false);
22565 };
22566 if !tp.runtime.native_p2p() || !crate::tp::step_tp_qkv_fused_enabled()? {
22567 return Ok(false);
22568 }
22569 let geometry = self.step35_geom(il);
22570 let heads = geometry.n_head as usize;
22571 let ws_index = tp
22572 .runtime
22573 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
22574 let gate_shards = attention
22575 .gate_shards_bf16
22576 .as_deref()
22577 .map(crate::tp::StepTpGateShards::Bf16);
22578 tp.runtime.decode_v2_input_qkv_tcol(
22579 ws_index,
22580 e,
22581 h_t,
22582 t,
22583 &tp.q,
22584 &tp.k,
22585 &tp.v,
22586 gate_shards,
22587 )?;
22588 Ok(true)
22589 }
22590
22591 pub(crate) fn step35_verify_oproj_tcol(
22596 &self,
22597 e: &Engine,
22598 il: usize,
22599 t: usize,
22600 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
22601 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
22602 return Err("tcol o_proj join expects full attention".into());
22603 };
22604 let tp = fa
22605 .step_tp_qkv
22606 .as_ref()
22607 .ok_or("tcol o_proj join lost its resident projections")?;
22608 let heads = self.step35_geom(il).n_head as usize;
22609 let ws_index = tp
22610 .runtime
22611 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
22612 tp.runtime.decode_v2_oproj_tcol(ws_index, e, &tp.o, t)
22613 }
22614
22615 #[allow(dead_code)] pub(crate) fn step35_spec_fa2_precheck(
22623 &self,
22624 cache: &Cache,
22625 il: usize,
22626 pos0: usize,
22627 ) -> Result<bool, Box<dyn std::error::Error>> {
22628 fn nope(clause: &str, il: usize, pos0: usize) -> bool {
22631 static DBG: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
22632 static SEEN: std::sync::Mutex<Vec<&'static str>> = std::sync::Mutex::new(Vec::new());
22633 if *DBG.get_or_init(|| std::env::var("MEMRA_SPEC_FA2_DEBUG").as_deref() == Ok("1")) {
22634 let mut seen = SEEN.lock().unwrap();
22635 if !seen.contains(&clause) {
22636 seen.push(Box::leak(clause.to_string().into_boxed_str()));
22638 eprintln!("[spec-fa2] precheck FAIL clause={clause} il={il} pos0={pos0}");
22639 }
22640 }
22641 false
22642 }
22643 static ONLY: std::sync::OnceLock<Option<usize>> = std::sync::OnceLock::new();
22645 if let Some(only) =
22646 ONLY.get_or_init(|| std::env::var("MEMRA_SPEC_FA2_LAYER").ok()?.parse().ok())
22647 && *only != il
22648 {
22649 return Ok(false);
22650 }
22651 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
22652 return Ok(nope("mixer", il, pos0));
22653 };
22654 let Some(tp) = fa.step_tp_qkv.as_ref() else {
22655 return Ok(nope("step_tp", il, pos0));
22656 };
22657 let Some(attention) = tp.attention.as_ref() else {
22658 return Ok(nope("attention", il, pos0));
22659 };
22660 if !tp.runtime.native_p2p()
22661 || crate::Engine::kv_fp8_on()
22662 || !crate::tp::step_tp_dcw_enabled()?
22663 || !crate::tp::step_tp_qkv_fused_enabled()?
22664 || (attention.gate_shards.is_none() && attention.gate_shards_bf16.is_none())
22665 {
22666 return Ok(nope("runtime-doors", il, pos0));
22667 }
22668 let geometry = self.step35_geom(il);
22669 let head_dim = geometry.head_dim_k as usize;
22670 if head_dim > 256 || !head_dim.is_multiple_of(32) || !crate::fa_v3_on() {
22671 return Ok(nope("fa-class", il, pos0));
22672 }
22673 let Some(distributed) = cache.tp_kv[il].as_ref() else {
22674 return Ok(nope("tp-kv", il, pos0));
22675 };
22676 if distributed.staged_len() != pos0 {
22677 return Ok(nope("staged-len", il, pos0));
22678 }
22679 let (_, would_rebase) = distributed.peek_append_ring(2)?;
22682 if would_rebase {
22683 return Ok(nope("rebase", il, pos0));
22684 }
22685 let window = geometry.window.map(|w| w as usize);
22686 if let Some(w) = window
22694 && pos0 + 2 > w
22695 {
22696 return Ok(nope("swa-capped", il, pos0));
22697 }
22698 let (t0, t1) = (pos0 + 1, pos0 + 2);
22702 if t0 < 96 {
22703 return Ok(nope("dcw-floor", il, pos0));
22704 }
22705 if std::env::var("MEMRA_NO_FA_VEC").is_ok() || t0 < crate::fa_vec_min_tkv() {
22706 return Ok(nope("vec-floor", il, pos0));
22707 }
22708 let ranks = tp.runtime.devices().len();
22714 let local_kv_heads = (geometry.n_head_kv as usize / ranks).max(1);
22715 let sp0 = crate::fa_split_keys_pub(t0, local_kv_heads);
22716 let sp1 = crate::fa_split_keys_pub(t1, local_kv_heads);
22717 if sp0 != sp1 {
22718 return Ok(nope("partition-sp", il, pos0));
22719 }
22720 let (ns0, ns1) = (t0.div_ceil(sp0), t1.div_ceil(sp1));
22721 if ns0 != ns1 {
22722 return Ok(nope("partition-ns", il, pos0));
22723 }
22724 if t0.div_ceil(ns0) != t1.div_ceil(ns1) {
22725 return Ok(nope("partition-per", il, pos0));
22726 }
22727 Ok(true)
22728 }
22729
22730 pub(crate) fn step35_fa_rows_precheck(
22735 &self,
22736 cache: &Cache,
22737 il: usize,
22738 pos0: usize,
22739 t: usize,
22740 ) -> Result<bool, Box<dyn std::error::Error>> {
22741 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
22742 return Ok(false);
22743 };
22744 let Some(tp) = fa.step_tp_qkv.as_ref() else {
22745 return Ok(false);
22746 };
22747 let Some(attention) = tp.attention.as_ref() else {
22748 return Ok(false);
22749 };
22750 if !tp.runtime.native_p2p()
22751 || crate::Engine::kv_fp8_on()
22752 || !crate::tp::step_tp_dcw_enabled()?
22753 || !crate::tp::step_tp_qkv_fused_enabled()?
22754 || (attention.gate_shards.is_none() && attention.gate_shards_bf16.is_none())
22755 {
22756 return Ok(false);
22757 }
22758 let geometry = self.step35_geom(il);
22759 let head_dim = geometry.head_dim_k as usize;
22760 if head_dim > 256 || !head_dim.is_multiple_of(32) || !crate::fa_v3_on() {
22761 return Ok(false);
22762 }
22763 if crate::fa_sm_count() < 128
22764 || std::env::var("MEMRA_FA_SPLIT").is_ok()
22765 || std::env::var("MEMRA_FA_SP_SHORT").is_ok()
22766 || std::env::var("MEMRA_FA_SP16").is_ok()
22767 || std::env::var("MEMRA_NO_FA_VEC").is_ok()
22768 {
22769 return Ok(false);
22770 }
22771 let Some(distributed) = cache.tp_kv[il].as_ref() else {
22772 return Ok(false);
22773 };
22774 if distributed.staged_len() != pos0 {
22775 return Ok(false);
22776 }
22777 let (_, would_rebase) = distributed.peek_append_ring(t)?;
22778 if would_rebase {
22779 return Ok(false);
22780 }
22781 let window = geometry.window.map(|w| w as usize);
22784 let t0 = window.map(|w| (pos0 + 1).min(w)).unwrap_or(pos0 + 1);
22785 if t0 < 96 || t0 < crate::fa_vec_min_tkv() {
22786 return Ok(false);
22787 }
22788 Ok(true)
22789 }
22790
22791 pub(crate) fn step35_verify_fa_rows_join(
22795 &self,
22796 e: &Engine,
22797 il: usize,
22798 cache: &Cache,
22799 pos0: usize,
22800 t: usize,
22801 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
22802 use cudarc::driver::DevicePtr;
22803 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
22804 return Err("fa rows join expects full attention".into());
22805 };
22806 let tp = fa
22807 .step_tp_qkv
22808 .as_ref()
22809 .ok_or("fa rows join lost its resident projections")?;
22810 let geometry = self.step35_geom(il);
22811 let heads = geometry.n_head as usize;
22812 let head_dim = geometry.head_dim_k as usize;
22813 let window = geometry.window.map(|w| w as usize);
22814 let distributed = cache.tp_kv[il]
22815 .as_ref()
22816 .ok_or("fa rows join lost its distributed KV cache")?;
22817 let (k_tok_bytes, v_tok_bytes) = (distributed.k_tok_bytes(), distributed.v_tok_bytes());
22818 let ladder = |t_kv: usize| -> usize {
22820 if t_kv <= 2048 {
22821 16
22822 } else if t_kv <= 16384 {
22823 64
22824 } else {
22825 128
22826 }
22827 };
22828 let mut max_ns = 1usize;
22829 for r in 0..t {
22830 let t_kv = window
22831 .map(|w| (pos0 + r + 1).min(w))
22832 .unwrap_or(pos0 + r + 1);
22833 max_ns = max_ns.max(t_kv.div_ceil(ladder(t_kv)));
22834 }
22835 let ranks = tp.runtime.devices().len();
22840 let mut tables = Vec::with_capacity(ranks);
22841 for rank in 0..ranks {
22842 let engine = tp
22843 .runtime
22844 .rank_engine(rank)
22845 .ok_or("fa rows join lost a rank engine")?;
22846 let rank_cache = distributed
22847 .rank(rank)
22848 .ok_or("fa rows join lost a KV cache rank")?;
22849 let _main = engine.gpu.enter_main()?;
22850 let s = engine.stream();
22851 let (kp, _g0) = rank_cache.k().device_ptr(&s);
22852 let (vp, _g1) = rank_cache.v().device_ptr(&s);
22853 let (lp, _g2) = rank_cache.len_d().device_ptr(&s);
22854 let bp = match rank_cache.base_d() {
22855 Some(b) => {
22856 let (p, _g) = b.device_ptr(&s);
22857 p
22858 }
22859 None => 0u64,
22860 };
22861 let mut host = Vec::with_capacity(t * 6);
22862 for r in 0..t {
22863 host.extend_from_slice(&[kp, vp, lp, bp, 0u64, (t - 1 - r) as u64]);
22864 }
22865 tables.push(engine.stream().clone_htod(&host)?);
22866 }
22867 let tabs: Vec<&CudaSlice<u64>> = tables.iter().collect();
22868 let ws_index = tp
22869 .runtime
22870 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
22871 tp.runtime.decode_v2_fa_rows_join(
22872 ws_index,
22873 e,
22874 &tp.o,
22875 &tabs,
22876 t,
22877 head_dim,
22878 window.unwrap_or(0),
22879 max_ns,
22880 geometry.attention_scale(),
22881 k_tok_bytes,
22882 v_tok_bytes,
22883 )
22884 }
22885
22886 pub(crate) fn step35_batch_fa_rows_precheck(
22890 &self,
22891 caches: &[&mut Cache],
22892 row_to_cache: impl Fn(usize) -> usize,
22893 positions: &[i32],
22894 il: usize,
22895 ) -> Result<bool, Box<dyn std::error::Error>> {
22896 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
22897 return Ok(false);
22898 };
22899 let Some(tp) = fa.step_tp_qkv.as_ref() else {
22900 return Ok(false);
22901 };
22902 let Some(attention) = tp.attention.as_ref() else {
22903 return Ok(false);
22904 };
22905 if !tp.runtime.native_p2p()
22906 || crate::Engine::kv_fp8_on()
22907 || !crate::tp::step_tp_dcw_enabled()?
22908 || !crate::tp::step_tp_qkv_fused_enabled()?
22909 || (attention.gate_shards.is_none() && attention.gate_shards_bf16.is_none())
22910 {
22911 return Ok(false);
22912 }
22913 let geometry = self.step35_geom(il);
22914 let head_dim = geometry.head_dim_k as usize;
22915 if head_dim > 256 || !head_dim.is_multiple_of(32) || !crate::fa_v3_on() {
22916 return Ok(false);
22917 }
22918 if crate::fa_sm_count() < 128
22919 || std::env::var("MEMRA_FA_SPLIT").is_ok()
22920 || std::env::var("MEMRA_FA_SP_SHORT").is_ok()
22921 || std::env::var("MEMRA_FA_SP16").is_ok()
22922 || std::env::var("MEMRA_NO_FA_VEC").is_ok()
22923 {
22924 return Ok(false);
22925 }
22926 let window = geometry.window.map(|w| w as usize);
22927 for (r, &pos) in positions.iter().enumerate() {
22928 let cache = &caches[row_to_cache(r)];
22929 let Some(distributed) = cache.tp_kv[il].as_ref() else {
22930 return Ok(false);
22931 };
22932 if distributed.staged_len() != pos as usize {
22933 return Ok(false);
22934 }
22935 if distributed.peek_append_ring(1)?.1 {
22936 return Ok(false);
22937 }
22938 let t0 = window
22939 .map(|w| (pos as usize + 1).min(w))
22940 .unwrap_or(pos as usize + 1);
22941 if t0 < 96 || t0 < crate::fa_vec_min_tkv() {
22942 return Ok(false);
22943 }
22944 }
22945 Ok(true)
22946 }
22947
22948 pub(crate) fn step35_verify_rope_fa_pass(
22954 &self,
22955 e: &Engine,
22956 il: usize,
22957 cache: &Cache,
22958 pos0: usize,
22959 t: usize,
22960 stage_pos: bool,
22961 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
22962 use cudarc::driver::DevicePtr;
22963 if !crate::tp::fuse_rope_append_on() {
22964 return Ok(None);
22965 }
22966 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
22967 return Ok(None);
22968 };
22969 let Some(tp) = fa.step_tp_qkv.as_ref() else {
22970 return Ok(None);
22971 };
22972 let Some(attention) = tp.attention.as_ref() else {
22973 return Ok(None);
22974 };
22975 let geometry = self.step35_geom(il);
22976 let head_dim = geometry.head_dim_k as usize;
22977 if head_dim != 128 {
22978 return Ok(None);
22979 }
22980 let heads = geometry.n_head as usize;
22981 let window = geometry.window.map(|w| w as usize);
22982 let ranks = tp.runtime.devices().len();
22983 let Some(distributed) = cache.tp_kv[il].as_ref() else {
22984 return Ok(None);
22985 };
22986 if distributed.kv_dim_k() != distributed.kv_dim_v() {
22987 return Ok(None);
22988 }
22989 {
22990 let rank0 = distributed.rank(0).ok_or("verify rope pass lost rank 0")?;
22991 if rank0.base_d().is_none()
22992 && distributed.staged_len() + t > distributed.physical_capacity()
22993 {
22994 return Ok(None);
22995 }
22996 }
22997 let mut rope_freqs = Vec::with_capacity(ranks);
22998 for rank in 0..ranks {
22999 let engine = tp
23000 .runtime
23001 .rank_engine(rank)
23002 .ok_or("verify rope pass lost a rank engine")?;
23003 rope_freqs.push(if geometry.rope_factors {
23004 match self
23005 .step35_aux
23006 .as_ref()
23007 .and_then(|aux| aux.rope_freqs(engine))
23008 {
23009 Some(f) => Some(f),
23010 None => return Ok(None),
23011 }
23012 } else {
23013 None
23014 });
23015 }
23016 let ladder = |t_kv: usize| -> usize {
23017 if t_kv <= 2048 {
23018 16
23019 } else if t_kv <= 16384 {
23020 64
23021 } else {
23022 128
23023 }
23024 };
23025 let (k_tok_bytes, v_tok_bytes) = (distributed.k_tok_bytes(), distributed.v_tok_bytes());
23026 let mut max_ns = 1usize;
23027 let mut positions = Vec::with_capacity(t);
23028 for r in 0..t {
23029 positions.push((pos0 + r) as i32);
23030 let t_kv = window
23031 .map(|w| (pos0 + r + 1).min(w))
23032 .unwrap_or(pos0 + r + 1);
23033 max_ns = max_ns.max(t_kv.div_ceil(ladder(t_kv)));
23034 }
23035 let mut session_parts: Vec<Vec<[u64; 4]>> = vec![Vec::with_capacity(t); ranks];
23036 let mut tab_keys = vec![0u64; ranks];
23037 for rank in 0..ranks {
23038 let engine = tp
23039 .runtime
23040 .rank_engine(rank)
23041 .ok_or("verify rope pass lost a rank engine")?;
23042 let rank_cache = distributed
23043 .rank(rank)
23044 .ok_or("verify rope pass lost a KV cache rank")?;
23045 let _main = engine.gpu.enter_main()?;
23046 let s = engine.stream();
23047 let (kp, _g0) = rank_cache.k().device_ptr(&s);
23048 let (vp, _g1) = rank_cache.v().device_ptr(&s);
23049 let (lp, _g2) = rank_cache.len_d().device_ptr(&s);
23050 let bp = match rank_cache.base_d() {
23051 Some(b) => {
23052 let (p, _g) = b.device_ptr(&s);
23053 p
23054 }
23055 None => 0u64,
23056 };
23057 tab_keys[rank] = kp
23058 .rotate_left(17)
23059 .wrapping_add(bp)
23060 .wrapping_add((il as u64) << 32)
23061 .wrapping_add(t as u64)
23062 .wrapping_add(1 << 63);
23063 for _r in 0..t {
23064 session_parts[rank].push([kp, vp, lp, bp]);
23065 }
23066 }
23067 let ws_index = tp
23068 .runtime
23069 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
23070 tp.runtime
23071 .decode_v2_rope_fa_rows(
23072 ws_index,
23073 e,
23074 &tp.o,
23075 &session_parts,
23076 &tab_keys,
23077 &positions,
23078 stage_pos,
23079 true,
23080 &attention.q_norm,
23081 &attention.k_norm,
23082 &rope_freqs,
23083 t,
23084 head_dim,
23085 geometry.n_rot as usize,
23086 window.unwrap_or(0),
23087 max_ns,
23088 geometry.attention_scale(),
23089 k_tok_bytes,
23090 v_tok_bytes,
23091 self.cfg.rms_eps,
23092 geometry.rope_base,
23093 )
23094 .map(Some)
23095 }
23096
23097 #[allow(clippy::too_many_arguments)]
23102 pub(crate) fn step35_batch_rope_fa_pass(
23103 &self,
23104 e: &Engine,
23105 il: usize,
23106 caches: &[&mut Cache],
23107 row_to_cache: impl Fn(usize) -> usize,
23108 positions: &[i32],
23109 t: usize,
23110 stage_pos: bool,
23111 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
23112 use cudarc::driver::DevicePtr;
23113 if !crate::tp::fuse_rope_append_on() {
23114 return Ok(None);
23115 }
23116 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
23117 return Ok(None);
23118 };
23119 let Some(tp) = fa.step_tp_qkv.as_ref() else {
23120 return Ok(None);
23121 };
23122 let Some(attention) = tp.attention.as_ref() else {
23123 return Ok(None);
23124 };
23125 let geometry = self.step35_geom(il);
23126 let head_dim = geometry.head_dim_k as usize;
23127 if head_dim != 128 {
23128 return Ok(None);
23129 }
23130 let heads = geometry.n_head as usize;
23131 let window = geometry.window.map(|w| w as usize);
23132 let ranks = tp.runtime.devices().len();
23133 for r in 0..t {
23136 let cache = &caches[row_to_cache(r)];
23137 let Some(distributed) = cache.tp_kv[il].as_ref() else {
23138 return Ok(None);
23139 };
23140 if distributed.kv_dim_k() != distributed.kv_dim_v() {
23141 return Ok(None);
23142 }
23143 let rank0 = distributed.rank(0).ok_or("rope fa pass lost rank 0")?;
23144 if rank0.base_d().is_none()
23145 && distributed.staged_len() + t > distributed.physical_capacity()
23146 {
23147 return Ok(None);
23148 }
23149 }
23150 let mut rope_freqs = Vec::with_capacity(ranks);
23151 for rank in 0..ranks {
23152 let engine = tp
23153 .runtime
23154 .rank_engine(rank)
23155 .ok_or("rope fa pass lost a rank engine")?;
23156 rope_freqs.push(if geometry.rope_factors {
23157 match self
23158 .step35_aux
23159 .as_ref()
23160 .and_then(|aux| aux.rope_freqs(engine))
23161 {
23162 Some(f) => Some(f),
23163 None => return Ok(None),
23164 }
23165 } else {
23166 None
23167 });
23168 }
23169 let ladder = |t_kv: usize| -> usize {
23170 if t_kv <= 2048 {
23171 16
23172 } else if t_kv <= 16384 {
23173 64
23174 } else {
23175 128
23176 }
23177 };
23178 let (mut max_ns, mut k_tok_bytes, mut v_tok_bytes) = (1usize, 0usize, 0usize);
23179 let mut session_parts: Vec<Vec<[u64; 4]>> = vec![Vec::with_capacity(t); ranks];
23180 let mut tab_keys = vec![0u64; ranks];
23181 for (r, &pos) in positions.iter().enumerate().take(t) {
23182 let cache = &caches[row_to_cache(r)];
23183 let distributed = cache.tp_kv[il]
23184 .as_ref()
23185 .ok_or("rope fa pass lost a distributed KV cache")?;
23186 (k_tok_bytes, v_tok_bytes) = (distributed.k_tok_bytes(), distributed.v_tok_bytes());
23187 let t_kv = window
23188 .map(|w| (pos as usize + 1).min(w))
23189 .unwrap_or(pos as usize + 1);
23190 max_ns = max_ns.max(t_kv.div_ceil(ladder(t_kv)));
23191 for rank in 0..ranks {
23192 let engine = tp
23193 .runtime
23194 .rank_engine(rank)
23195 .ok_or("rope fa pass lost a rank engine")?;
23196 let rank_cache = distributed
23197 .rank(rank)
23198 .ok_or("rope fa pass lost a KV cache rank")?;
23199 let _main = engine.gpu.enter_main()?;
23200 let s = engine.stream();
23201 let (kp, _g0) = rank_cache.k().device_ptr(&s);
23202 let (vp, _g1) = rank_cache.v().device_ptr(&s);
23203 let (lp, _g2) = rank_cache.len_d().device_ptr(&s);
23204 let bp = match rank_cache.base_d() {
23205 Some(b) => {
23206 let (p, _g) = b.device_ptr(&s);
23207 p
23208 }
23209 None => 0u64,
23210 };
23211 tab_keys[rank] = tab_keys[rank]
23212 .rotate_left(9)
23213 .wrapping_add(kp)
23214 .wrapping_add(bp)
23215 .wrapping_add(il as u64);
23216 session_parts[rank].push([kp, vp, lp, bp]);
23217 }
23218 }
23219 let ws_index = tp
23220 .runtime
23221 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
23222 tp.runtime
23223 .decode_v2_rope_fa_rows(
23224 ws_index,
23225 e,
23226 &tp.o,
23227 &session_parts,
23228 &tab_keys,
23229 positions,
23230 stage_pos,
23231 false,
23232 &attention.q_norm,
23233 &attention.k_norm,
23234 &rope_freqs,
23235 t,
23236 head_dim,
23237 geometry.n_rot as usize,
23238 window.unwrap_or(0),
23239 max_ns,
23240 geometry.attention_scale(),
23241 k_tok_bytes,
23242 v_tok_bytes,
23243 self.cfg.rms_eps,
23244 geometry.rope_base,
23245 )
23246 .map(Some)
23247 }
23248
23249 #[allow(clippy::too_many_arguments)]
23253 pub(crate) fn step35_batch_fa_rows_join(
23254 &self,
23255 e: &Engine,
23256 il: usize,
23257 caches: &[&mut Cache],
23258 row_to_cache: impl Fn(usize) -> usize,
23259 positions: &[i32],
23260 t: usize,
23261 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
23262 use cudarc::driver::DevicePtr;
23263 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
23264 return Err("batch fa rows join expects full attention".into());
23265 };
23266 let tp = fa
23267 .step_tp_qkv
23268 .as_ref()
23269 .ok_or("batch fa rows join lost its resident projections")?;
23270 let geometry = self.step35_geom(il);
23271 let heads = geometry.n_head as usize;
23272 let head_dim = geometry.head_dim_k as usize;
23273 let window = geometry.window.map(|w| w as usize);
23274 let ladder = |t_kv: usize| -> usize {
23275 if t_kv <= 2048 {
23276 16
23277 } else if t_kv <= 16384 {
23278 64
23279 } else {
23280 128
23281 }
23282 };
23283 let (mut max_ns, mut k_tok_bytes, mut v_tok_bytes) = (1usize, 0usize, 0usize);
23284 for (r, &pos) in positions.iter().enumerate() {
23285 let cache = &caches[row_to_cache(r)];
23286 let distributed = cache.tp_kv[il]
23287 .as_ref()
23288 .ok_or("batch fa rows join lost a distributed KV cache")?;
23289 (k_tok_bytes, v_tok_bytes) = (distributed.k_tok_bytes(), distributed.v_tok_bytes());
23290 let t_kv = window
23291 .map(|w| (pos as usize + 1).min(w))
23292 .unwrap_or(pos as usize + 1);
23293 max_ns = max_ns.max(t_kv.div_ceil(ladder(t_kv)));
23294 }
23295 let ranks = tp.runtime.devices().len();
23299 let mut tables = Vec::with_capacity(ranks);
23300 for rank in 0..ranks {
23301 let engine = tp
23302 .runtime
23303 .rank_engine(rank)
23304 .ok_or("batch fa rows join lost a rank engine")?;
23305 let _main = engine.gpu.enter_main()?;
23306 let s = engine.stream();
23307 let mut host = Vec::with_capacity(t * 6);
23308 for r in 0..t {
23309 let cache = &caches[row_to_cache(r)];
23310 let distributed = cache.tp_kv[il]
23311 .as_ref()
23312 .ok_or("batch fa rows join lost a distributed KV cache")?;
23313 let rank_cache = distributed
23314 .rank(rank)
23315 .ok_or("batch fa rows join lost a KV cache rank")?;
23316 let (kp, _g0) = rank_cache.k().device_ptr(&s);
23317 let (vp, _g1) = rank_cache.v().device_ptr(&s);
23318 let (lp, _g2) = rank_cache.len_d().device_ptr(&s);
23319 let bp = match rank_cache.base_d() {
23320 Some(b) => {
23321 let (p, _g) = b.device_ptr(&s);
23322 p
23323 }
23324 None => 0u64,
23325 };
23326 host.extend_from_slice(&[kp, vp, lp, bp, 0u64, 0u64]);
23327 }
23328 tables.push(engine.stream().clone_htod(&host)?);
23329 }
23330 let tabs: Vec<&CudaSlice<u64>> = tables.iter().collect();
23331 let ws_index = tp
23332 .runtime
23333 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
23334 tp.runtime.decode_v2_fa_rows_join(
23335 ws_index,
23336 e,
23337 &tp.o,
23338 &tabs,
23339 t,
23340 head_dim,
23341 window.unwrap_or(0),
23342 max_ns,
23343 geometry.attention_scale(),
23344 k_tok_bytes,
23345 v_tok_bytes,
23346 )
23347 }
23348
23349 #[allow(dead_code)] pub(crate) fn step35_verify_spec_fa2_join(
23354 &self,
23355 e: &Engine,
23356 il: usize,
23357 cache: &Cache,
23358 pos0: usize,
23359 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
23360 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
23361 return Err("spec fa2 join expects full attention".into());
23362 };
23363 let tp = fa
23364 .step_tp_qkv
23365 .as_ref()
23366 .ok_or("spec fa2 join lost its resident projections")?;
23367 let geometry = self.step35_geom(il);
23368 let heads = geometry.n_head as usize;
23369 let head_dim = geometry.head_dim_k as usize;
23370 let window = geometry.window.map(|w| w as usize);
23371 let bucket = window.map(|w| (pos0 + 2).min(w)).unwrap_or(pos0 + 2);
23374 let distributed = cache.tp_kv[il]
23375 .as_ref()
23376 .ok_or("spec fa2 join lost its distributed KV cache")?;
23377 let ws_index = tp
23378 .runtime
23379 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
23380 tp.runtime.decode_v2_spec_fa2_join(
23381 ws_index,
23382 e,
23383 &tp.o,
23384 distributed,
23385 head_dim,
23386 window.unwrap_or(0),
23387 bucket,
23388 geometry.attention_scale(),
23389 )
23390 }
23391
23392 pub(crate) fn step35_tp_decode_attn_resident_v2(
23393 &self,
23394 e: &Engine,
23395 fa: &FullAttnLayer,
23396 il: usize,
23397 h: &CudaSlice<f32>,
23398 pos_d: &CudaSlice<i32>,
23399 cache: &mut Cache,
23400 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
23401 let tp = fa
23402 .step_tp_qkv
23403 .as_ref()
23404 .ok_or("Step TP decode lost its resident projections")?;
23405 let attention = tp
23406 .attention
23407 .as_ref()
23408 .ok_or("Step TP decode lost its resident attention auxiliaries")?;
23409 if !tp.runtime.native_p2p() {
23410 return Err("rank-local Step attention requires native P2P".into());
23411 }
23412 if crate::Engine::kv_fp8_on() {
23413 return Err("rank-local Step attention has not qualified the FP8 KV cache".into());
23414 }
23415
23416 let geometry = self.cfg.full_attention_geometry_at(il as u32);
23417 let window = geometry.window.map(|window| window as usize);
23418 let ranks = tp.runtime.devices().len();
23419 let head_dim = geometry.head_dim_k as usize;
23420 let heads = geometry.n_head as usize;
23421 let kv_heads = geometry.n_head_kv as usize;
23422 if !heads.is_multiple_of(ranks) || !kv_heads.is_multiple_of(ranks) {
23423 return Err(format!(
23424 "Step attention heads q={heads} kv={kv_heads} are not divisible by TP={ranks}"
23425 )
23426 .into());
23427 }
23428 let local_heads = heads / ranks;
23429 let local_kv_heads = kv_heads / ranks;
23430 let max_ctx = cache.max_ctx;
23431
23432 let hydrated = self.ensure_step_tp_kv_cache(e, fa, il, cache)?;
23433
23434 let base_len = cache.kv[il]
23435 .as_ref()
23436 .ok_or_else(|| format!("Step TP layer {il} lost its local KV cache"))?
23437 .len;
23438 {
23439 let distributed = cache.tp_kv[il]
23440 .as_ref()
23441 .ok_or_else(|| format!("Step TP layer {il} lost its distributed KV cache"))?;
23442 if distributed.committed_len() != base_len || distributed.staged_len() != base_len {
23443 return Err(format!(
23444 "Step TP layer {il} cache lengths diverged before decode: \
23445 local={base_len} distributed={}/{}",
23446 distributed.committed_len(),
23447 distributed.staged_len()
23448 )
23449 .into());
23450 }
23451 }
23452 if pos_d.len() != 1 {
23453 return Err(format!(
23454 "rank-local Step decode requires one position, got {}",
23455 pos_d.len()
23456 )
23457 .into());
23458 }
23459
23460 let decode_input = attention
23461 .decode_input
23462 .as_ref()
23463 .ok_or("Step TP decode v2 requires the replicated decode input")?;
23464 let mut decode_input = decode_input
23465 .lock()
23466 .map_err(|_| "Step TP replicated decode input lock is poisoned")?;
23467
23468 let has_gate = fa.attn_gate.is_some();
23469 let use_gate_shards = has_gate
23474 && (attention.gate_shards.is_some() || attention.gate_shards_bf16.is_some())
23475 && crate::tp::step_tp_qkv_fused_enabled()?;
23476 let gate_raw = if !has_gate || use_gate_shards {
23477 None
23478 } else {
23479 let gate_weight = fa
23480 .attn_gate
23481 .as_ref()
23482 .ok_or("step35 layer is missing attn_gate.weight")?;
23483 let gate_raw = e.matmul(gate_weight, h, 1)?;
23484 if gate_raw.len() != heads {
23485 return Err(format!(
23486 "Step TP layer {il} gate output {} != {heads}",
23487 gate_raw.len()
23488 )
23489 .into());
23490 }
23491 Some(gate_raw)
23492 };
23493
23494 let ws_index = tp
23495 .runtime
23496 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
23497 let mut ws_guard = tp
23498 .runtime
23499 .decode_v2_workspace()
23500 .lock()
23501 .map_err(|_| "Step TP decode v2 workspace lock is poisoned")?;
23502 let ws = ws_guard
23503 .get_mut(ws_index)
23504 .ok_or("Step TP decode v2 workspace missing after ensure")?;
23505
23506 let mut rope_freqs = Vec::with_capacity(ranks);
23507 for rank in 0..ranks {
23508 let engine = tp
23509 .runtime
23510 .rank_engine(rank)
23511 .ok_or_else(|| format!("Step TP layer {il} has no engine for rank {rank}"))?;
23512 rope_freqs.push(if geometry.rope_factors {
23513 self.step35_aux
23514 .as_ref()
23515 .and_then(|aux| aux.rope_freqs(engine))
23516 } else {
23517 None
23518 });
23519 }
23520 let staged_next = base_len + 1;
23527 let t_kv_eff = window
23528 .map(|window| staged_next.min(window))
23529 .unwrap_or(staged_next);
23530 let dcw = crate::tp::step_tp_dcw_enabled()?
23531 && (use_gate_shards || (!has_gate && crate::tp::step_tp_qkv_fused_enabled()?))
23532 && t_kv_eff >= 96
23533 && {
23534 let (write_row, would_rebase) = cache.tp_kv[il]
23535 .as_ref()
23536 .expect("distributed cache checked above")
23537 .peek_append_ring(1)?;
23538 if !would_rebase {
23539 let base = (base_len - write_row) as i32;
23541 let distributed = cache.tp_kv[il]
23542 .as_mut()
23543 .expect("distributed cache checked above");
23544 for rank in 0..ranks {
23545 let engine = tp.runtime.rank_engine(rank).ok_or_else(|| {
23546 format!("Step TP layer {il} has no engine for rank {rank}")
23547 })?;
23548 let _main = engine.gpu.enter_main()?;
23549 let rank_cache = distributed.rank_mut(rank).ok_or_else(|| {
23550 format!("Step TP layer {il} has no KV cache rank {rank}")
23551 })?;
23552 if rank_cache.base_d().is_none() {
23553 rank_cache.arm_base_d(engine.htod_i32(&[base])?);
23554 }
23555 }
23556 }
23557 !would_rebase
23558 };
23559 let fuse_rope = dcw
23560 && crate::tp::fuse_rope_append_on()
23561 && head_dim == 128
23562 && cache.tp_kv[il]
23563 .as_ref()
23564 .map(|d| d.kv_dim_k() == d.kv_dim_v() && d.kv_dim_k() == local_kv_heads * head_dim)
23565 .unwrap_or(false);
23566
23567 let tcol_col = crate::tp::take_verify_tcol();
23568 let fa2_col = crate::tp::take_spec_fa2_defer();
23575 tp.runtime.decode_v2_input_qkv(
23576 ws,
23577 e,
23578 h,
23579 pos_d,
23580 gate_raw.as_ref(),
23581 if !use_gate_shards {
23582 None
23583 } else if let Some(shards) = attention.gate_shards.as_deref() {
23584 Some(crate::tp::StepTpGateShards::F32(shards))
23585 } else {
23586 attention
23587 .gate_shards_bf16
23588 .as_deref()
23589 .map(crate::tp::StepTpGateShards::Bf16)
23590 },
23591 &mut decode_input,
23592 &tp.q,
23593 &tp.k,
23594 &tp.v,
23595 &attention.q_norm,
23596 &attention.k_norm,
23597 head_dim,
23598 geometry.n_rot as usize,
23599 geometry.rope_base,
23600 &rope_freqs,
23601 self.cfg.rms_eps,
23602 has_gate,
23603 fuse_rope,
23604 tcol_col,
23605 )?;
23606
23607 let transaction = cache.tp_kv[il]
23608 .as_mut()
23609 .expect("distributed cache checked above")
23610 .begin_transaction()?;
23611 let append_result = tp.runtime.append_tp_kv_transaction_inner(
23612 cache.tp_kv[il]
23613 .as_mut()
23614 .expect("distributed cache checked above"),
23615 transaction,
23616 &ws.k,
23617 &ws.v_raw,
23618 1,
23619 dcw,
23620 );
23621 if let Err(error) = append_result {
23622 let _ = tp.runtime.rollback_tp_kv_transaction(
23623 cache.tp_kv[il]
23624 .as_mut()
23625 .expect("distributed cache checked above"),
23626 transaction,
23627 );
23628 return Err(error);
23629 }
23630
23631 let staged = (|| -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
23632 let (staged_len, physical, k_tok_bytes_c, v_tok_bytes_c, capacity) = {
23635 let distributed = cache.tp_kv[il]
23636 .as_ref()
23637 .expect("distributed cache checked above");
23638 let staged_len = distributed.staged_len();
23639 let view_start = window
23640 .map(|window| staged_len.saturating_sub(window))
23641 .unwrap_or(0);
23642 (
23643 staged_len,
23644 distributed.physical_range(view_start, staged_len)?,
23645 distributed.k_tok_bytes(),
23646 distributed.v_tok_bytes(),
23647 distributed.physical_capacity(),
23648 )
23649 };
23650 let view_start = window
23651 .map(|window| staged_len.saturating_sub(window))
23652 .unwrap_or(0);
23653 let t_kv = staged_len - view_start;
23654 for rank in 0..ranks {
23655 let engine = tp
23656 .runtime
23657 .rank_engine(rank)
23658 .ok_or_else(|| format!("Step TP layer {il} has no engine for rank {rank}"))?;
23659 let _main = engine.gpu.enter_main()?;
23660 if dcw {
23661 {
23665 let distributed_mut = cache.tp_kv[il]
23666 .as_mut()
23667 .expect("distributed cache checked above");
23668 let (kv_dim_k, kv_dim_v) =
23669 (distributed_mut.kv_dim_k(), distributed_mut.kv_dim_v());
23670 let (k_tok_bytes, v_tok_bytes) =
23671 (distributed_mut.k_tok_bytes(), distributed_mut.v_tok_bytes());
23672 let rank_cache = distributed_mut.rank_mut(rank).ok_or_else(|| {
23673 format!("Step TP layer {il} has no KV cache rank {rank}")
23674 })?;
23675 let (k_plane, v_plane, len_d, base_d) =
23676 rank_cache.planes_and_counters_mut();
23677 if fuse_rope {
23678 let same_dev = engine.ctx().ordinal() == e.ctx().ordinal();
23681 let crate::tp::StepTpDecodeV2Ws {
23682 q_raw,
23683 k_raw,
23684 v_raw,
23685 q,
23686 k,
23687 pos,
23688 pos_stage,
23689 fuse_ctr,
23690 ..
23691 } = &mut *ws;
23692 let pos_ref: &CudaSlice<i32> = if same_dev {
23696 pos_stage
23697 .as_ref()
23698 .ok_or("step TP decode v2 pos stage not armed")?
23699 } else {
23700 &pos[rank]
23701 };
23702 engine.qk_norm_rope_append_inc_dcw(
23703 &q_raw[rank],
23704 &k_raw[rank],
23705 &v_raw[rank],
23706 &attention.q_norm[rank],
23707 &attention.k_norm[rank],
23708 &mut q[rank],
23709 &mut k[rank],
23710 pos_ref,
23711 k_plane,
23712 v_plane,
23713 len_d,
23714 base_d,
23715 &mut fuse_ctr[rank],
23716 kv_dim_k,
23717 kv_dim_v,
23718 k_tok_bytes,
23719 v_tok_bytes,
23720 head_dim,
23721 geometry.n_rot as usize,
23722 local_heads,
23723 local_kv_heads,
23724 self.cfg.rms_eps,
23725 geometry.rope_base,
23726 1.0,
23727 rope_freqs[rank],
23728 )?;
23729 } else {
23730 engine.append_kv_quantized_dcw(
23731 &ws.k[rank],
23732 &ws.v_raw[rank],
23733 k_plane,
23734 v_plane,
23735 len_d,
23736 base_d,
23737 kv_dim_k,
23738 kv_dim_v,
23739 k_tok_bytes,
23740 v_tok_bytes,
23741 )?;
23742 }
23743 if !fuse_rope {
23744 let rank_cache = distributed_mut.rank_mut(rank).ok_or_else(|| {
23745 format!("Step TP layer {il} has no KV cache rank {rank}")
23746 })?;
23747 engine.inc_i32(rank_cache.len_d_mut())?;
23748 }
23749 }
23750 let distributed = cache.tp_kv[il]
23751 .as_ref()
23752 .expect("distributed cache checked above");
23753 let rank_cache = distributed
23754 .rank(rank)
23755 .ok_or_else(|| format!("Step TP layer {il} has no KV cache rank {rank}"))?;
23756 let k_ring = engine.view_u8_range(rank_cache.k(), 0, capacity * k_tok_bytes_c);
23757 let v_ring = engine.view_u8_range(rank_cache.v(), 0, capacity * v_tok_bytes_c);
23758 if fa2_col.is_some() {
23759 continue;
23762 }
23763 {
23764 let crate::tp::StepTpDecodeV2Ws { q, gate, gated, .. } = &mut *ws;
23767 engine.fa_decode_dcw(
23768 &q[rank],
23769 &k_ring,
23770 &v_ring,
23771 &mut gated[rank],
23772 head_dim,
23773 local_heads,
23774 local_kv_heads,
23775 rank_cache.len_d(),
23776 rank_cache.base_d(),
23777 window.unwrap_or(0),
23778 t_kv,
23779 geometry.attention_scale(),
23780 k_tok_bytes_c,
23781 v_tok_bytes_c,
23782 has_gate.then_some(&gate[rank]),
23783 )?;
23784 }
23785 continue;
23786 }
23787 let distributed = cache.tp_kv[il]
23788 .as_ref()
23789 .expect("distributed cache checked above");
23790 let rank_cache = distributed
23791 .rank(rank)
23792 .ok_or_else(|| format!("Step TP layer {il} has no KV cache rank {rank}"))?;
23793 let k_view = engine.view_u8_range(
23794 rank_cache.k(),
23795 physical.start * k_tok_bytes_c,
23796 physical.end * k_tok_bytes_c,
23797 );
23798 let v_view = engine.view_u8_range(
23799 rank_cache.v(),
23800 physical.start * v_tok_bytes_c,
23801 physical.end * v_tok_bytes_c,
23802 );
23803 if has_gate {
23804 engine.fa_decode_kvmod(
23805 &ws.q[rank],
23806 &k_view,
23807 &v_view,
23808 &mut ws.attn_out[rank],
23809 head_dim,
23810 local_heads,
23811 local_kv_heads,
23812 t_kv,
23813 geometry.attention_scale(),
23814 k_tok_bytes_c,
23815 v_tok_bytes_c,
23816 false,
23817 )?;
23818 engine.attn_head_gate(
23819 &ws.attn_out[rank],
23820 &ws.gate[rank],
23821 &mut ws.gated[rank],
23822 None,
23823 head_dim,
23824 local_heads,
23825 1,
23826 )?;
23827 } else {
23828 engine.fa_decode_kvmod(
23829 &ws.q[rank],
23830 &k_view,
23831 &v_view,
23832 &mut ws.gated[rank],
23833 head_dim,
23834 local_heads,
23835 local_kv_heads,
23836 t_kv,
23837 geometry.attention_scale(),
23838 k_tok_bytes_c,
23839 v_tok_bytes_c,
23840 false,
23841 )?;
23842 }
23843 }
23844
23845 let output = if let Some(col) = fa2_col.filter(|_| dcw) {
23852 tp.runtime.decode_v2_stash_fa2(ws, e, col)?;
23856 crate::tp::set_spec_fa2_stashed();
23857 e.uninit(ws.o_out)?
23858 } else if let Some(col) = crate::tp::take_tcol_oproj_defer() {
23859 if tp.runtime.decode_v2_oproj_tcol_eligible(ws, &tp.o) {
23860 tp.runtime.decode_v2_stash_gated(ws, e, col)?;
23861 crate::tp::set_tcol_oproj_stashed();
23862 e.uninit(ws.o_out)?
23863 } else {
23864 tp.runtime.decode_v2_finish(ws, e, &tp.o)?
23865 }
23866 } else {
23867 tp.runtime.decode_v2_finish(ws, e, &tp.o)?
23868 };
23869
23870 let local = cache.kv[il]
23874 .as_mut()
23875 .ok_or_else(|| format!("Step TP layer {il} lost its local KV cache"))?;
23876 if local.len != base_len || base_len + 1 > max_ctx {
23877 return Err(format!(
23878 "Step TP layer {il} local cache changed during decode: \
23879 len={} base={base_len} max={max_ctx}",
23880 local.len
23881 )
23882 .into());
23883 }
23884 if crate::tp::no_local_shadow_on() {
23885 local.len = base_len + 1;
23888 if !crate::tp::len_mirror_lazy_on() {
23892 e.set_i32_one(&mut local.len_d, local.len as i32)?;
23893 }
23894 } else {
23895 let retain_from = window
23896 .map(|window| {
23897 let staged_retain = (base_len + 1).saturating_sub(window) & !31usize;
23898 let rollback_retain =
23899 base_len.saturating_sub(window.saturating_sub(1)) & !31usize;
23900 staged_retain.min(rollback_retain)
23901 })
23902 .unwrap_or(0);
23903 let write_row = e.prepare_kv_append(local, retain_from, 1)?;
23904 e.append_kv_quantized(
23905 &ws.k_shadow,
23906 &ws.v_shadow,
23907 &mut local.k,
23908 &mut local.v,
23909 write_row,
23910 local.kv_dim_k,
23911 local.kv_dim_v,
23912 local.k_tok_bytes,
23913 local.v_tok_bytes,
23914 false,
23915 )?;
23916 local.len = base_len + 1;
23917 e.set_i32_one(&mut local.len_d, local.len as i32)?;
23918 }
23919 Ok(output)
23920 })();
23921
23922 let output = match staged {
23923 Ok(output) => output,
23924 Err(error) => {
23925 let _ = tp.runtime.rollback_tp_kv_transaction(
23926 cache.tp_kv[il]
23927 .as_mut()
23928 .expect("distributed cache checked above"),
23929 transaction,
23930 );
23931 if let Some(local) = cache.kv[il].as_mut() {
23932 local.len = base_len;
23933 let _ = e.set_i32_one(&mut local.len_d, base_len as i32);
23934 }
23935 return Err(error);
23936 }
23937 };
23938 let lazy_commit = fuse_rope && crate::tp::len_mirror_lazy_on();
23943 if lazy_commit {
23944 if let Err(error) = tp.runtime.commit_tp_kv_transaction_external(
23945 cache.tp_kv[il]
23946 .as_mut()
23947 .expect("distributed cache checked above"),
23948 transaction,
23949 1,
23950 ) {
23951 let _ = tp.runtime.rollback_tp_kv_transaction(
23952 cache.tp_kv[il]
23953 .as_mut()
23954 .expect("distributed cache checked above"),
23955 transaction,
23956 );
23957 let local = cache.kv[il].as_mut().expect("local cache checked above");
23958 local.len = base_len;
23959 e.set_i32_one(&mut local.len_d, base_len as i32)?;
23960 return Err(error);
23961 }
23962 } else if let Err(error) = tp.runtime.commit_tp_kv_transaction(
23963 cache.tp_kv[il]
23964 .as_mut()
23965 .expect("distributed cache checked above"),
23966 transaction,
23967 1,
23968 ) {
23969 let _ = tp.runtime.rollback_tp_kv_transaction(
23970 cache.tp_kv[il]
23971 .as_mut()
23972 .expect("distributed cache checked above"),
23973 transaction,
23974 );
23975 let local = cache.kv[il].as_mut().expect("local cache checked above");
23976 local.len = base_len;
23977 e.set_i32_one(&mut local.len_d, base_len as i32)?;
23978 return Err(error);
23979 }
23980
23981 let committed = cache.tp_kv[il]
23982 .as_ref()
23983 .expect("distributed cache checked above")
23984 .committed_len();
23985 let local_len = cache.kv[il]
23986 .as_ref()
23987 .expect("local cache checked above")
23988 .len;
23989 if committed != local_len {
23990 return Err(format!(
23991 "Step TP layer {il} committed cache length {committed} != local shadow {local_len}"
23992 )
23993 .into());
23994 }
23995 static V2_LOGGED: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false);
23996 if !V2_LOGGED.swap(true, std::sync::atomic::Ordering::Relaxed) {
23997 eprintln!(
23998 "[step-tp-attn-v2] execute layer={} devices={:?} tokens=1 driver=v2 \
23999 qkv_tensor_parallel=true qk_norm_rank_local=true rope_rank_local=true \
24000 kv_cache_distributed=true kv_cache_hydrated={hydrated} \
24001 attention_tensor_parallel=true attention_scope={} \
24002 input_path=root-device-replicated gate={} gate_tensor_parallel={} \
24003 gate_shards={} o_tensor_parallel=true o_reduce=root-device \
24004 local_cache_shadow=true cache_commit=immediate transport={} native_p2p=true \
24005 bulk_p2p={} workspace=persistent ordering=evented output=e-device \
24006 performance_claim=false (logged once; every decode layer runs this driver)",
24007 tp.layer,
24008 tp.devices,
24009 if window.is_some() {
24010 "rank-local-swa-ring"
24011 } else {
24012 "rank-local-global"
24013 },
24014 has_gate,
24015 use_gate_shards,
24016 if use_gate_shards {
24017 "device-staged"
24018 } else if has_gate {
24019 "root-staged"
24020 } else {
24021 "none"
24022 },
24023 tp.runtime.transport_label(),
24024 tp.runtime.bulk_p2p(),
24025 );
24026 }
24027 Ok(output)
24028 }
24029
24030 #[allow(clippy::too_many_arguments)]
24040 pub(crate) fn step35_decode_attn(
24041 &self,
24042 e: &Engine,
24043 fa: &FullAttnLayer,
24044 il: usize,
24045 h: &CudaSlice<f32>,
24046 pre_q: Option<(&CudaSlice<i8>, &CudaSlice<f32>)>,
24047 pos_d: &CudaSlice<i32>,
24048 cache: &mut Cache,
24049 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
24050 if fa
24051 .step_tp_qkv
24052 .as_ref()
24053 .is_some_and(|tp| tp.attention.is_some())
24054 {
24055 if pre_q.is_some() {
24056 return Err(
24057 "rank-local Step attention preserves BF16 activations and refuses the q8_1 \
24058 pre-quantized decode path"
24059 .into(),
24060 );
24061 }
24062 return self.step35_tp_decode_attn_resident(e, fa, il, h, pos_d, cache);
24063 }
24064
24065 let geometry = self.step35_geom(il);
24066 let hd = geometry.head_dim_k as usize;
24067 let nkv = geometry.n_head_kv as usize;
24068 let nh = geometry.n_head as usize;
24069 let rbase = geometry.rope_base;
24070 let scale = geometry.attention_scale();
24071 let swa = geometry.window.is_some();
24072 let eps = self.cfg.rms_eps;
24073 let win = geometry.window.unwrap_or(0) as usize;
24074 let n_rot = geometry.n_rot as usize;
24075 let n_embd = self.cfg.n_embd as usize;
24076 let gw = fa
24077 .attn_gate
24078 .as_ref()
24079 .ok_or("step35 layer is missing attn_gate.weight (head-wise attention gate)")?;
24080
24081 let tp_qkv = if fa.step_tp_qkv.is_some() {
24082 if pre_q.is_some() {
24083 return Err(
24084 "Step Q/K/V TP preserves BF16 activations and refuses the q8_1 \
24085 pre-quantized decode path"
24086 .into(),
24087 );
24088 }
24089 self.full_attn_tp_qkv(e, fa, h, 1)?
24090 } else {
24091 None
24092 };
24093
24094 let (q0, k0, v0, gt) = match tp_qkv {
24095 Some(mut g3) => {
24096 let v = g3.pop().unwrap();
24097 let k = g3.pop().unwrap();
24098 let q = g3.pop().unwrap();
24099 let gt = e.matmul(gw, h, 1)?;
24100 (q, k, v, gt)
24101 }
24102 None => match pre_q {
24103 Some((hq, hdq)) => {
24104 debug_assert!(
24105 e.uses_q8_1_fast(gw),
24106 "step35 pre-quantized decode requires attn_gate on the q8_1 fast path \
24107 (h is a zero-length placeholder here) — see mixer_in_q8_1_fast"
24108 );
24109 let (a, b, c) = match e.matmul_q8_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)? {
24110 Some(t3) => t3,
24111 None => (
24112 e.matmul_pre(&fa.wq, hq, hdq, h, 1)?,
24113 e.matmul_pre(&fa.wk, hq, hdq, h, 1)?,
24114 e.matmul_pre(&fa.wv, hq, hdq, h, 1)?,
24115 ),
24116 };
24117 let gt = e.matmul_pre(gw, hq, hdq, h, 1)?;
24118 (a, b, c, gt)
24119 }
24120 None => {
24121 if e.uses_q8_1_fast(&fa.wq)
24122 && e.uses_q8_1_fast(&fa.wk)
24123 && e.uses_q8_1_fast(&fa.wv)
24124 && e.uses_q8_1_fast(gw)
24125 {
24126 let (hq, hdq) = e.quantize_q8_1(h, 1, n_embd)?;
24127 let (a, b, c) =
24128 match e.matmul_q8_fused3(&fa.wq, &fa.wk, &fa.wv, &hq, &hdq)? {
24129 Some(t3) => t3,
24130 None => (
24131 e.matmul_pre(&fa.wq, &hq, &hdq, h, 1)?,
24132 e.matmul_pre(&fa.wk, &hq, &hdq, h, 1)?,
24133 e.matmul_pre(&fa.wv, &hq, &hdq, h, 1)?,
24134 ),
24135 };
24136 let gt = e.matmul_pre(gw, &hq, &hdq, h, 1)?;
24137 (a, b, c, gt)
24138 } else {
24139 (
24140 e.matmul(&fa.wq, h, 1)?,
24141 e.matmul(&fa.wk, h, 1)?,
24142 e.matmul(&fa.wv, h, 1)?,
24143 e.matmul(gw, h, 1)?,
24144 )
24145 }
24146 }
24147 },
24148 };
24149
24150 let mut q = e.uninit(nh * hd)?;
24151 e.rms_norm(&q0, fa.q_norm.float_data(), &mut q, hd, nh, eps)?;
24152 let mut k = e.uninit(nkv * hd)?;
24153 e.rms_norm(&k0, fa.k_norm.float_data(), &mut k, hd, nkv, eps)?;
24154 let ff = if swa {
24155 None
24156 } else {
24157 self.step35_aux.as_ref().and_then(|a| a.rope_freqs(e))
24158 };
24159 #[cfg(debug_assertions)]
24160 if let Some(ff) = ff {
24161 crate::debug_assert_tensor_stream_device(
24162 ff,
24163 &e.stream(),
24164 "step35_decode_attn.rope_freqs",
24165 );
24166 }
24167 e.rope_neox2(&mut q, &mut k, pos_d, hd, n_rot, nh, nkv, 1, rbase, 1.0, ff)?;
24168
24169 if std::env::var("MEMRA_NOFA").is_ok() {
24170 return Err(
24171 "MEMRA_NOFA (naive f32 SDPA) is incompatible with the quantized KV \
24172 cache; unset MEMRA_NOFA to use fa_decode"
24173 .into(),
24174 );
24175 }
24176 let kvl = cache.kv[il].as_mut().unwrap();
24177 let next_len = kvl.len + 1;
24178 let (off, t_kv) = if swa && next_len > win {
24179 (next_len - win, win)
24180 } else {
24181 (0, next_len)
24182 };
24183 let write_row = e.prepare_kv_append(kvl, off & !31usize, 1)?;
24184 e.append_kv_quantized(
24185 &k,
24186 &v0,
24187 &mut kvl.k,
24188 &mut kvl.v,
24189 write_row,
24190 kvl.kv_dim_k,
24191 kvl.kv_dim_v,
24192 kvl.k_tok_bytes,
24193 kvl.v_tok_bytes,
24194 crate::Engine::kv_fp8_on(),
24195 )?;
24196 kvl.len = next_len;
24197 let physical = kvl.physical_rows(off, off + t_kv)?;
24198 let k_view = e.view_u8_range(
24199 &kvl.k,
24200 physical.start * kvl.k_tok_bytes,
24201 physical.end * kvl.k_tok_bytes,
24202 );
24203 let v_view = e.view_u8_range(
24204 &kvl.v,
24205 physical.start * kvl.v_tok_bytes,
24206 physical.end * kvl.v_tok_bytes,
24207 );
24208 let mut attn = e.uninit(nh * hd)?;
24209 e.fa_decode_kvmod(
24210 &q,
24211 &k_view,
24212 &v_view,
24213 &mut attn,
24214 hd,
24215 nh,
24216 nkv,
24217 t_kv,
24218 scale,
24219 kvl.k_tok_bytes,
24220 kvl.v_tok_bytes,
24221 crate::Engine::kv_fp8_on(),
24222 )?;
24223
24224 let mut ag = e.uninit(nh * hd)?;
24225 e.attn_head_gate(&attn, >, &mut ag, None, hd, nh, 1)?;
24226 self.full_attn_o(e, fa, &ag, 1)
24227 }
24228}
24229
24230impl HybridModel {
24239 pub fn is_gemma4_e4b(&self) -> bool {
24240 self.gemma4_aux.as_ref().is_some_and(|a| a.e4b.is_some())
24241 }
24242
24243 fn gemma4_e4b_geom(&self, il: usize) -> (usize, usize, usize, f32, f32, bool) {
24247 let g = self.cfg.gemma4.as_ref().unwrap();
24248 let swa = g.swa_pattern[il];
24249 let hd = if swa {
24250 g.key_length_swa
24251 } else {
24252 g.key_length_global
24253 } as usize;
24254 let Mixer::Full(fa) = &self.layers[il].mixer else {
24255 panic!("e4b layer {il} not full-attn")
24256 };
24257 let nh = fa.wq.out_features() / hd;
24258 let nkv = fa.wk.out_features() / hd;
24259 (
24260 hd,
24261 nkv,
24262 nh,
24263 if swa {
24264 g.rope_base_swa
24265 } else {
24266 g.rope_base_global
24267 },
24268 1.0,
24269 swa,
24270 )
24271 }
24272
24273 fn gemma4_e4b_kv_target(&self, il: usize) -> Option<usize> {
24275 self.layers[il]
24276 .gemma4
24277 .as_ref()
24278 .and_then(|b| b.e4b.as_ref())
24279 .and_then(|e4| e4.kv_share.map(|t| t as usize))
24280 }
24281
24282 fn gemma4_e4b_inp_pl(
24287 &self,
24288 e: &Engine,
24289 tokens: &[u32],
24290 x_scaled: &CudaSlice<f32>,
24291 t: usize,
24292 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
24293 let tok_d = e.stream().clone_htod(&tokens.to_vec())?;
24294 self.gemma4_e4b_inp_pl_dev(e, &tok_d, x_scaled, t)
24295 }
24296
24297 fn gemma4_e4b_inp_pl_dev(
24299 &self,
24300 e: &Engine,
24301 tok_d: &CudaSlice<u32>,
24302 x_scaled: &CudaSlice<f32>,
24303 t: usize,
24304 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
24305 let aux = self.gemma4_aux.as_ref().unwrap();
24306 let m = aux.e4b.as_ref().unwrap();
24307 let n_embd = self.cfg.n_embd as usize;
24308 let n_layer = self.layers.len();
24309 let width = m.n_epl * n_layer;
24310 let tbl = m.tok_tbl_gpu.get_or_init(|| {
24311 e.upload_u8(&m.tok_embd_bytes)
24312 .expect("e4b per-layer token table upload")
24313 });
24314 let mut a =
24315 e.embed_gather_device_td(tbl, tok_d, t, width, m.tok_embd_qt, m.tok_embd_row_bytes)?;
24316 e.scale_inplace(&mut a, (m.n_epl as f32).sqrt(), t * width)?;
24317 let mut p = e.matmul(&m.model_proj, x_scaled, t)?;
24318 e.scale_inplace(&mut p, 1.0 / (n_embd as f32).sqrt(), t * width)?;
24319 let mut pn = e.uninit(t * width)?;
24320 e.rms_norm(
24321 &p,
24322 m.proj_norm.float_data(),
24323 &mut pn,
24324 m.n_epl,
24325 t * n_layer,
24326 self.cfg.rms_eps,
24327 )?;
24328 let mut out = e.uninit(t * width)?;
24329 e.add_scale(&a, &pn, 1.0 / 2f32.sqrt(), &mut out, t * width)?;
24330 Ok(out)
24331 }
24332
24333 #[allow(clippy::too_many_arguments)]
24338 fn gemma4_e4b_attn(
24339 &self,
24340 e: &Engine,
24341 il: usize,
24342 hq: &CudaSlice<i8>,
24343 hdq: &CudaSlice<f32>,
24344 pos_d: &CudaSlice<i32>,
24345 t: usize,
24346 cache: &mut Cache,
24347 dc_bucket: Option<usize>,
24348 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
24349 let (hd, nkv, nh, base, scale, swa) = self.gemma4_e4b_geom(il);
24350 let eps = self.cfg.rms_eps;
24351 let aux = self.gemma4_aux.as_ref().unwrap();
24352 let ones = aux.ones(e);
24353 #[cfg(debug_assertions)]
24354 crate::debug_assert_tensor_stream_device(ones, &e.stream(), "gemma4_e4b_attn.ones");
24355 let Mixer::Full(fa) = &self.layers[il].mixer else {
24356 unreachable!()
24357 };
24358 let h0 = e.zeros(0)?;
24362 let h = &h0;
24363
24364 let ff = if swa {
24365 None
24366 } else {
24367 Some(
24368 aux.rope_freqs(e)
24369 .expect("e4b global rope needs rope_freqs.weight"),
24370 )
24371 };
24372 #[cfg(debug_assertions)]
24373 if let Some(ff) = ff {
24374 crate::debug_assert_tensor_stream_device(ff, &e.stream(), "gemma4_e4b_attn.rope_freqs");
24375 }
24376 let share = self.gemma4_e4b_kv_target(il);
24377 let mut kv_f32: Option<(CudaSlice<f32>, CudaSlice<f32>)> = None;
24379 let mut q;
24380 if let Some(_tgt) = share {
24381 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
24382 q = e.uninit(t * nh * hd)?;
24383 let mut kdummy = e.uninit(1)?;
24386 let mut vdummy = e.uninit(1)?;
24387 e.rms_norm_qkv_rope(
24388 &q0,
24389 &q0,
24390 &q0,
24391 fa.q_norm.float_data(),
24392 fa.q_norm.float_data(),
24393 ones,
24394 &mut q,
24395 &mut kdummy,
24396 &mut vdummy,
24397 hd,
24398 self.gemma4_rope_dims(il),
24399 nh * t,
24400 0,
24401 pos_d,
24402 nh,
24403 1,
24404 base,
24405 1.0,
24406 ff,
24407 eps,
24408 )?;
24409 } else {
24410 let e4bits = self.layers[il].gemma4.as_ref().and_then(|g| g.e4b.as_ref());
24414 let cat = e4bits.and_then(|e4| e4.qkv_cat.as_ref());
24415 q = e.uninit(t * nh * hd)?;
24416 let mut k = e.uninit(t * nkv * hd)?;
24417 let mut v = e.uninit(t * nkv * hd)?;
24418 if t == 1 && cat.is_some() {
24419 #[allow(clippy::unnecessary_unwrap)]
24420 let qkv0 = e.matmul_pre(cat.unwrap(), hq, hdq, h, 1)?;
24422 e.rms_norm_qkv_rope_cat(
24423 &qkv0,
24424 fa.q_norm.float_data(),
24425 fa.k_norm.float_data(),
24426 ones,
24427 &mut q,
24428 &mut k,
24429 &mut v,
24430 hd,
24431 self.gemma4_rope_dims(il),
24432 nh,
24433 nkv,
24434 pos_d,
24435 nh,
24436 nkv,
24437 base,
24438 1.0,
24439 ff,
24440 eps,
24441 )?;
24442 } else {
24443 let (q0, k0, v0) = match if t == 1 {
24444 e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)?
24445 } else {
24446 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
24449 if *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0")) {
24450 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
24451 } else {
24452 None
24453 }
24454 } {
24455 Some(triple) => triple,
24456 None => (
24457 e.matmul_pre(&fa.wq, hq, hdq, h, t)?,
24458 e.matmul_pre(&fa.wk, hq, hdq, h, t)?,
24459 e.matmul_pre(&fa.wv, hq, hdq, h, t)?,
24460 ), };
24462 e.rms_norm_qkv_rope(
24465 &q0,
24466 &k0,
24467 &v0,
24468 fa.q_norm.float_data(),
24469 fa.k_norm.float_data(),
24470 ones,
24471 &mut q,
24472 &mut k,
24473 &mut v,
24474 hd,
24475 self.gemma4_rope_dims(il),
24476 nh * t,
24477 nkv * t,
24478 pos_d,
24479 nh,
24480 nkv,
24481 base,
24482 1.0,
24483 ff,
24484 eps,
24485 )?;
24486 }
24487 let kvl = cache.kv[il].as_mut().unwrap();
24488 let cls = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
24492 if dc_bucket.is_some() {
24493 debug_assert!(t == 1);
24498 e.append_kv_quantized_row_dc_inc(
24500 &k,
24501 &v,
24502 &mut kvl.k,
24503 &mut kvl.v,
24504 &mut kvl.len_d,
24505 kvl.kv_dim_k,
24506 kvl.kv_dim_v,
24507 kvl.k_tok_bytes,
24508 kvl.v_tok_bytes,
24509 cls,
24510 )?;
24511 } else {
24512 e.append_kv_quantized_rows(
24513 &k,
24514 &v,
24515 &mut kvl.k,
24516 &mut kvl.v,
24517 kvl.len,
24518 t,
24519 kvl.kv_dim_k,
24520 kvl.kv_dim_v,
24521 kvl.k_tok_bytes,
24522 kvl.v_tok_bytes,
24523 cls,
24524 )?;
24525 kvl.len += t;
24526 }
24527 kv_f32 = Some((k, v));
24528 }
24529 let kvl_idx = share.unwrap_or(il);
24532 let kvl = cache.kv[kvl_idx].as_ref().unwrap();
24533 let base_len = kvl.len - t; let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
24535 let mut attn = e.uninit(t * nh * hd)?;
24536 if t > 1 && base_len == 0 && std::env::var("MEMRA_NOFA").is_err() {
24548 if let Some((kf, vf)) = &kv_f32 {
24549 if hd == 256 && t <= win {
24550 e.fa_prefill(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale, true)?;
24551 return e.matmul(&fa.wo, &attn, t);
24552 }
24553 if hd == 256 && swa && t > win {
24554 e.fa_prefill_w(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
24555 return e.matmul(&fa.wo, &attn, t);
24556 }
24557 if hd == 512 && !swa {
24558 e.fa_prefill_hd512(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale, true)?;
24559 return e.matmul(&fa.wo, &attn, t);
24560 }
24561 } else if share.is_some() {
24562 let g = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
24563 let k_view = e.view_u8(&kvl.k, kvl.k.len());
24564 let v_view = e.view_u8(&kvl.v, kvl.v.len());
24565 if hd == 256 && (!swa || t <= win) {
24566 e.fa_prefill_view(
24568 &q,
24569 &k_view,
24570 &v_view,
24571 &mut attn,
24572 hd,
24573 nh,
24574 nkv,
24575 t,
24576 t,
24577 scale,
24578 true,
24579 kvl.k_tok_bytes,
24580 kvl.v_tok_bytes,
24581 g,
24582 )?;
24583 return e.matmul(&fa.wo, &attn, t);
24584 }
24585 let kv_dim = nkv * hd;
24588 let mut kf = e.uninit(t * kv_dim)?;
24589 let mut vf = e.uninit(t * kv_dim)?;
24590 e.fa_dequant_kv_view_f32(
24591 &k_view,
24592 &v_view,
24593 &mut kf,
24594 &mut vf,
24595 kv_dim,
24596 kv_dim,
24597 t,
24598 kvl.k_tok_bytes,
24599 kvl.v_tok_bytes,
24600 g,
24601 )?;
24602 if hd == 512 {
24603 e.fa_prefill_hd512(&q, &kf, &vf, &mut attn, hd, nh, nkv, t, t, scale, true)?;
24604 } else {
24605 e.fa_prefill_w(&q, &kf, &vf, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
24606 }
24607 return e.matmul(&fa.wo, &attn, t);
24608 }
24609 }
24610 if let Some(bucket) = dc_bucket {
24611 assert!(t == 1);
24616 let bucket = if hd == 512 && win <= crate::fa512_min_tkv() {
24622 bucket.min(crate::fa512_min_tkv().saturating_sub(1))
24623 } else {
24624 bucket
24625 };
24626 let k_view = e.view_u8(&kvl.k, kvl.k.len());
24627 let v_view = e.view_u8(&kvl.v, kvl.v.len());
24628 let g = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
24629 if crate::Engine::wpf_level() >= 1 {
24637 e.prefetch_weight_l2(&fa.wo)?;
24638 }
24639 if e.uses_q8_1_fast(&fa.wo) {
24642 let mut oq = e.alloc_i8_uninit(nh * hd)?;
24643 let mut od = e.zeros(nh * hd / 32)?;
24644 e.fa_decode_dc_q8(
24645 &q,
24646 &k_view,
24647 &v_view,
24648 &mut attn,
24649 hd,
24650 nh,
24651 nkv,
24652 &kvl.len_d,
24653 bucket,
24654 scale,
24655 kvl.k_tok_bytes,
24656 kvl.v_tok_bytes,
24657 g,
24658 Some((&mut oq, &mut od)),
24659 )?;
24660 return e.matmul_pre(&fa.wo, &oq, &od, &attn, t);
24661 }
24662 e.fa_decode_dc(
24663 &q,
24664 &k_view,
24665 &v_view,
24666 &mut attn,
24667 hd,
24668 nh,
24669 nkv,
24670 &kvl.len_d,
24671 bucket,
24672 scale,
24673 kvl.k_tok_bytes,
24674 kvl.v_tok_bytes,
24675 g,
24676 )?;
24677 return e.matmul(&fa.wo, &attn, t);
24678 }
24679 for i in 0..t {
24680 let avail = base_len + i + 1;
24681 let (off_tok, t_kv) = if swa && avail > win {
24682 (avail - win, win)
24683 } else {
24684 (0, avail)
24685 };
24686 let k_view = e.view_u8_range(
24687 &kvl.k,
24688 off_tok * kvl.k_tok_bytes,
24689 (off_tok + t_kv) * kvl.k_tok_bytes,
24690 );
24691 let v_view = e.view_u8_range(
24692 &kvl.v,
24693 off_tok * kvl.v_tok_bytes,
24694 (off_tok + t_kv) * kvl.v_tok_bytes,
24695 );
24696 let qv = e.view(&q, t * nh * hd);
24697 let q_row = qv.slice(i * nh * hd..(i + 1) * nh * hd);
24698 let mut q_one = e.uninit(nh * hd)?;
24699 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
24700 let mut a_one = e.uninit(nh * hd)?;
24701 e.fa_decode_kvmod(
24705 &q_one,
24706 &k_view,
24707 &v_view,
24708 &mut a_one,
24709 hd,
24710 nh,
24711 nkv,
24712 t_kv,
24713 scale,
24714 kvl.k_tok_bytes,
24715 kvl.v_tok_bytes,
24716 (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on()),
24717 )?;
24718 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
24719 }
24720 e.matmul(&fa.wo, &attn, t)
24721 }
24722
24723 fn gemma4_e4b_trunk(
24728 &self,
24729 e: &Engine,
24730 tokens: &[u32],
24731 pos0: usize,
24732 cache: &mut Cache,
24733 head_last: bool,
24734 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
24735 let n_embd = self.cfg.n_embd as usize;
24736 let t = tokens.len();
24737 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
24738 let pos_d = e.htod_i32(&pos)?;
24739 let mut x = e.htod(&self.embd.gather(n_embd, tokens))?;
24740 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
24741 let inp_pl = self.gemma4_e4b_inp_pl(e, tokens, &x, t)?;
24742 self.gemma4_e4b_trunk_core(e, x, inp_pl, &pos_d, t, cache, None, true, head_last)
24743 }
24744
24745 #[allow(clippy::too_many_arguments)] fn gemma4_e4b_trunk_core(
24750 &self,
24751 e: &Engine,
24752 x_in: CudaSlice<f32>,
24753 inp_pl: CudaSlice<f32>,
24754 pos_d: &CudaSlice<i32>,
24755 t: usize,
24756 cache: &mut Cache,
24757 dc_bucket: Option<usize>,
24758 cap_logits: bool,
24759 head_last: bool,
24760 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
24761 let n_embd = self.cfg.n_embd as usize;
24762 let eps = self.cfg.rms_eps;
24763 let n_layer = self.layers.len();
24764 let mut x = x_in;
24765 let aux_e4b = self.gemma4_aux.as_ref().unwrap().e4b.as_ref().unwrap();
24766 let n_epl = aux_e4b.n_epl;
24767
24768 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
24774 for il in 0..n_layer {
24775 let layer = &self.layers[il];
24776 let (hq, hdq) = match h_carry.take() {
24777 Some(p) => p,
24778 None => e.rms_norm_q8_1(&x, layer.attn_norm.float_data(), n_embd, t, eps)?,
24779 };
24780 let o = self.gemma4_e4b_attn(e, il, &hq, &hdq, pos_d, t, cache, dc_bucket)?;
24781 let bits = layer.gemma4.as_ref().unwrap();
24784 let e4b = bits.e4b.as_ref().expect("e4b layer bits");
24785 let fuse_exit = e.uses_q8_1_fast(&e4b.inp_gate);
24796 let (sn, attn_out) = self.gemma4_layer_tail_core_pn(
24797 e,
24798 layer,
24799 &o,
24800 &x,
24801 t,
24802 Some(layer.post_attn_norm.float_data()),
24803 fuse_exit,
24804 )?;
24805 let mut resid = e.uninit(t * n_embd)?;
24806 let g = if fuse_exit {
24812 let (rq, rd) = e.rms_pre_add_q8_1(
24814 &sn,
24815 bits.post_ffw_norm.float_data(),
24816 &attn_out,
24817 &mut resid,
24818 n_embd,
24819 t,
24820 self.cfg.rms_eps,
24821 )?;
24822 e.matmul_pre(&e4b.inp_gate, &rq, &rd, &resid, t)?
24823 } else {
24824 e.add(&sn, &attn_out, &mut resid, t * n_embd)?;
24825 e.matmul(&e4b.inp_gate, &resid, t)?
24826 };
24827 let mut act = e.uninit(t * n_epl)?;
24828 let y = if t == 1 && e.uses_q8_1_fast(&e4b.proj) {
24829 let ipv = e.view(&inp_pl, n_epl * n_layer);
24830 let row = ipv.slice(il * n_epl..(il + 1) * n_epl);
24831 let (aq, ad) = e.gelu_tanh_mul_q8_1(&g, &row, &mut act, n_epl, 1)?;
24832 e.matmul_pre(&e4b.proj, &aq, &ad, &act, t)?
24833 } else {
24834 let mut inp_this = e.uninit(t * n_epl)?;
24835 e.copy_rows_strided(
24836 &inp_pl,
24837 &mut inp_this,
24838 n_epl,
24839 t,
24840 n_epl * n_layer,
24841 il * n_epl,
24842 )?;
24843 e.gelu_tanh_mul(&g, &inp_this, &mut act, t * n_epl)?;
24844 e.matmul(&e4b.proj, &act, t)?
24845 };
24846 let next_norm = if il + 1 < n_layer {
24849 self.layers[il + 1].attn_norm.float_data()
24850 } else {
24851 self.output_norm.float_data()
24852 };
24853 let mut xn = e.uninit(t * n_embd)?;
24854 let pair = e.rms_pre_add_scale_rms_norm_q8_1(
24855 &y,
24856 e4b.post_norm.float_data(),
24857 &resid,
24858 bits.layer_scale,
24859 next_norm,
24860 &mut xn,
24861 n_embd,
24862 t,
24863 eps,
24864 )?;
24865 h_carry = Some(pair);
24866 x = xn;
24867 }
24868 let (oq, odq) = h_carry.take().unwrap();
24872 let h0 = e.zeros(0)?;
24873 let hm = if head_last { 1 } else { t };
24874 let (hq, hd) = if head_last && t > 1 {
24875 let mut q1 = e.uninit_i8(n_embd)?;
24876 e.dtod_copy_view_i8(&oq.slice((t - 1) * n_embd..t * n_embd), &mut q1)?;
24877 let nb = n_embd / 32;
24878 let mut d1 = e.uninit(nb)?;
24879 e.dtod_copy_view(&odq.slice((t - 1) * nb..t * nb), &mut d1)?;
24880 (q1, d1)
24881 } else {
24882 (oq, odq)
24883 };
24884 let mut ld = e.matmul_pre(&self.output, &hq, &hd, &h0, hm)?;
24885 if cap_logits {
24889 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
24890 e.softcap(&mut ld, cap, hm * self.output.out_features())?;
24891 }
24892 self.gemma4_suppress(e, &mut ld, hm)?; Ok((ld, x))
24894 }
24895
24896 pub fn gemma4_e4b_decode_step_t_am_dev(
24903 &self,
24904 e: &Engine,
24905 tok_d: &CudaSlice<u32>,
24906 t: usize,
24907 pos0: usize,
24908 cache: &mut Cache,
24909 ) -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
24910 let n_embd = self.cfg.n_embd as usize;
24911 let eps = self.cfg.rms_eps;
24912 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
24913 let pos_d = e.htod_i32(&pos)?;
24914 let embd_gpu = self
24915 .embd_gpu
24916 .get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload"));
24917 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
24918 let mut x = e.embed_gather_device_td(embd_gpu, tok_d, t, n_embd, qt, rb)?;
24919 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
24920 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, tok_d, &x, t)?;
24921 let (ld, xp) =
24922 self.gemma4_e4b_trunk_core(e, x, inp_pl, &pos_d, t, cache, None, true, false)?;
24923 let n_vocab = self.output.out_features();
24926 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
24927 for i in 0..t {
24928 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
24929 }
24930 let mut hn = e.uninit(t * n_embd)?;
24931 e.rms_norm(&xp, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
24932 cache.pos += t;
24933 Ok((vam, hn))
24934 }
24935
24936 pub(crate) fn gemma4_e4b_decode_step_t_h(
24939 &self,
24940 e: &Engine,
24941 tokens: &[u32],
24942 pos0: usize,
24943 cache: &mut Cache,
24944 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
24945 let n_embd = self.cfg.n_embd as usize;
24946 let eps = self.cfg.rms_eps;
24947 let t = tokens.len();
24948 let (ld, xp) = self.gemma4_e4b_trunk(e, tokens, pos0, cache, false)?;
24949 let mut hn = e.uninit(t * n_embd)?;
24950 e.rms_norm(&xp, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
24951 cache.pos += t;
24952 Ok((e.dtoh(&ld)?, hn))
24953 }
24954
24955 #[allow(clippy::too_many_arguments)] pub fn gemma4_e4b_decode_step_dcg(
24962 &self,
24963 e: &Engine,
24964 token_d: &mut CudaSlice<u32>,
24965 pos_d: &mut CudaSlice<i32>,
24966 embd_gpu: &CudaSlice<u8>,
24967 embd_qt: i32,
24968 embd_rb: usize,
24969 cache: &mut Cache,
24970 n_vocab: usize,
24971 bucket: usize,
24972 ) -> Result<(), Box<dyn std::error::Error>> {
24973 let n_embd = self.cfg.n_embd as usize;
24974 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
24975 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
24976 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, token_d, &x, 1)?;
24977 let (ld, _x) =
24978 self.gemma4_e4b_trunk_core(e, x, inp_pl, pos_d, 1, cache, Some(bucket), false, false)?;
24979 e.argmax_token_device_into(&ld, token_d, n_vocab)?;
24980 e.inc_seqlen(pos_d)?;
24981 Ok(())
24982 }
24983
24984 #[allow(clippy::too_many_arguments)]
24992 pub fn gemma4_e4b_decode_step_dc(
24993 &self,
24994 e: &Engine,
24995 token_d: &CudaSlice<u32>,
24996 pos_d: &mut CudaSlice<i32>,
24997 embd_gpu: &CudaSlice<u8>,
24998 embd_qt: i32,
24999 embd_rb: usize,
25000 cache: &mut Cache,
25001 n_vocab: usize,
25002 ) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
25003 let n_embd = self.cfg.n_embd as usize;
25004 let eps = self.cfg.rms_eps;
25005 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
25006 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
25007 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, token_d, &x, 1)?;
25008 let (ld, _x) =
25009 self.gemma4_e4b_trunk_core(e, x, inp_pl, pos_d, 1, cache, None, false, false)?;
25010 let mut tok_out = e.stream().alloc_zeros::<u32>(1)?;
25011 e.argmax_token_device_into(&ld, &mut tok_out, n_vocab)?;
25012 e.inc_seqlen(pos_d)?;
25013 cache.pos += 1;
25014 let _ = eps;
25015 Ok(tok_out)
25016 }
25017
25018 pub(crate) fn gemma4_e4b_decode_step_h(
25021 &self,
25022 e: &Engine,
25023 token: u32,
25024 cache: &mut Cache,
25025 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
25026 let (ld, x) = self.gemma4_e4b_trunk(e, &[token], cache.pos, cache, false)?;
25027 let logits = e.dtoh(&ld)?;
25028 cache.pos += 1;
25029 Ok((logits, x))
25030 }
25031
25032 #[allow(clippy::type_complexity)] pub(crate) fn gemma4_e4b_prime(
25037 &self,
25038 e: &Engine,
25039 tokens: &[u32],
25040 cache: &mut Cache,
25041 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
25042 if cache.pos != 0 {
25045 return Err(
25046 "e4b prime is fresh-prompt only (v0) — prime the full prompt in one \
25047 call or decode tokenwise"
25048 .into(),
25049 );
25050 }
25051 let n_embd = self.cfg.n_embd as usize;
25052 let t = tokens.len();
25053 let (ld, x) = self.gemma4_e4b_trunk(e, tokens, 0, cache, true)?;
25054 cache.pos += t;
25055 let last = e.dtoh(&ld)?; let xv = e.view(&x, t * n_embd);
25057 let row = xv.slice((t - 1) * n_embd..t * n_embd);
25058 let mut h_seed = e.uninit(n_embd)?;
25059 e.copy_view_into(&mut h_seed, 0, &row, n_embd)?;
25060 Ok((last, h_seed, x))
25061 }
25062
25063 pub(crate) fn gemma4_e4b_forward(
25065 &self,
25066 e: &Engine,
25067 tokens: &[u32],
25068 last_only: bool,
25069 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
25070 let mut cache = Cache::new(e, &self.cfg, tokens.len() + 8)?;
25071 let (ld, _x) = self.gemma4_e4b_trunk(e, tokens, 0, &mut cache, last_only)?;
25072 e.dtoh(&ld) }
25074}
25075
25076#[cfg(test)]
25077mod prime_chunk_schedule_tests {
25078 use super::{
25079 CUDA_GRID_YZ_MAX, PRIME_CHUNK_LAUNCH_CAP, PRIME_MIN_T, PRIME_PIPE_MIN_CHUNK, PrimePpSignal,
25080 PrimePpStageChannels, PrimePpWaveCredits, PrimePpWaveSlot, active_matrix_values,
25081 align_prime_ranges_to_gdn, dynamic_prime_chunk_ranges, explicit_prime_chunk,
25082 fixed_prime_chunk_ranges, fixed_prime_chunk_ranges_for_ring, move_prime_cache_layers,
25083 parse_step_ep_grouped_prefill, parse_step_tp_prefill, prime_cache_stage_for_layer,
25084 recv_prime_pp_signal, restore_prime_cache_layers, step_grouped_decode_shape,
25085 step_grouped_prefill_shape, step_tp_prefill_shape, validate_step_prime_batch_modes,
25086 };
25087
25088 fn sizes(ranges: &[(usize, usize)]) -> Vec<usize> {
25089 ranges.iter().map(|(start, end)| end - start).collect()
25090 }
25091
25092 #[allow(clippy::manual_clamp)] fn auto_chunk(t: usize) -> usize {
25094 t.div_ceil(8).max(PRIME_PIPE_MIN_CHUNK).min(4096)
25095 }
25096
25097 #[test]
25102 #[allow(clippy::assertions_on_constants)]
25105 fn monolithic_prime_chunk_caps_at_the_cuda_launch_wall() {
25106 assert_eq!(explicit_prime_chunk(0, false), PRIME_CHUNK_LAUNCH_CAP);
25109 assert_eq!(explicit_prime_chunk(100_000, false), PRIME_CHUNK_LAUNCH_CAP);
25110 assert_eq!(explicit_prime_chunk(4096, false), 4096);
25111 assert_eq!(
25112 explicit_prime_chunk(PRIME_CHUNK_LAUNCH_CAP, false),
25113 PRIME_CHUNK_LAUNCH_CAP
25114 );
25115 assert_eq!(
25117 explicit_prime_chunk(0, true),
25118 crate::cache::PRIME_CHUNK_MAX_TOKENS
25119 );
25120 assert_eq!(
25121 explicit_prime_chunk(100_000, true),
25122 crate::cache::PRIME_CHUNK_MAX_TOKENS
25123 );
25124 assert_eq!(explicit_prime_chunk(512, true), 512);
25125 assert!(PRIME_CHUNK_LAUNCH_CAP + PRIME_MIN_T - 1 <= CUDA_GRID_YZ_MAX);
25127 }
25128
25129 #[test]
25130 fn capped_monolithic_ranges_are_identical_below_the_wall_and_legal_above() {
25131 let chunk = explicit_prime_chunk(0, false);
25132 for t in [
25136 PRIME_MIN_T,
25137 4096,
25138 61_000,
25139 64_984,
25140 PRIME_CHUNK_LAUNCH_CAP,
25141 PRIME_CHUNK_LAUNCH_CAP + 1,
25142 CUDA_GRID_YZ_MAX,
25143 ] {
25144 assert_eq!(
25145 fixed_prime_chunk_ranges_for_ring(t, chunk, false),
25146 vec![(0, t)],
25147 "t={t} must stay a single monolithic range"
25148 );
25149 }
25150 for t in [65_536, 65_643, 66_045, 79_717, 82_440, 262_144] {
25154 let ranges = fixed_prime_chunk_ranges_for_ring(t, chunk, false);
25155 assert!(ranges.len() >= 2, "t={t} must chunk");
25156 let mut cursor = 0usize;
25157 for &(start, end) in &ranges {
25158 assert_eq!(start, cursor, "t={t}: ranges must be contiguous");
25159 assert!(
25160 end - start <= CUDA_GRID_YZ_MAX,
25161 "t={t}: range width {} exceeds the CUDA grid.y limit",
25162 end - start
25163 );
25164 assert!(
25165 end - start >= PRIME_MIN_T,
25166 "t={t}: range width {} below PRIME_MIN_T",
25167 end - start
25168 );
25169 cursor = end;
25170 }
25171 assert_eq!(cursor, t, "t={t}: ranges must cover the prompt");
25172 }
25173 for t in (CUDA_GRID_YZ_MAX - 64)..=(CUDA_GRID_YZ_MAX + 2 * PRIME_MIN_T + 64) {
25175 for &(start, end) in &fixed_prime_chunk_ranges_for_ring(t, chunk, false) {
25176 assert!(
25177 end - start <= CUDA_GRID_YZ_MAX,
25178 "t={t} width {}",
25179 end - start
25180 );
25181 }
25182 }
25183 }
25184
25185 #[test]
25186 fn ppn_prime_cache_partition_moves_and_restores_every_layer() {
25187 let round_trip = |fence: &[usize], layers: usize| {
25188 let original: Vec<Option<usize>> = (0..layers).map(Some).collect();
25189 let mut parent = original.clone();
25190 let mut stages: Vec<Vec<Option<usize>>> =
25191 (0..fence.len() - 1).map(|_| vec![None; layers]).collect();
25192
25193 move_prime_cache_layers(&mut parent, &mut stages, fence);
25194 assert!(parent.iter().all(Option::is_none));
25195 for layer in 0..layers {
25196 let owner = prime_cache_stage_for_layer(fence, layer);
25197 for (stage, values) in stages.iter().enumerate() {
25198 assert_eq!(values[layer], (stage == owner).then_some(layer));
25199 }
25200 }
25201
25202 restore_prime_cache_layers(&mut parent, &mut stages, fence);
25203 assert_eq!(parent, original);
25204 assert!(stages.iter().flatten().all(Option::is_none));
25205 };
25206
25207 round_trip(&[0, 5, 8], 10);
25209 round_trip(&[0, 2, 5, 8], 10);
25210 round_trip(&[0, 1, 3, 6, 8], 10);
25211 }
25212
25213 #[test]
25214 fn ppn_prime_wave_credit_requires_the_exact_oldest_wave_and_slot() {
25215 let mut credits = PrimePpWaveCredits::default();
25216 let wave0 = PrimePpWaveSlot { wave: 0, slot: 1 };
25217 let wave1 = PrimePpWaveSlot { wave: 1, slot: 0 };
25218 credits.record_send(wave0).unwrap();
25219 assert_eq!(credits.release_required(), None);
25220 credits.record_send(wave1).unwrap();
25221 assert_eq!(credits.release_required(), Some(wave0));
25222
25223 assert!(
25224 credits
25225 .record_release(PrimePpWaveSlot { wave: 0, slot: 0 })
25226 .unwrap_err()
25227 .contains("does not match oldest pending")
25228 );
25229 assert_eq!(credits.release_required(), Some(wave0));
25230 credits.record_release(wave0).unwrap();
25231 credits
25232 .record_send(PrimePpWaveSlot { wave: 2, slot: 1 })
25233 .unwrap();
25234 assert!(
25235 credits
25236 .record_send(PrimePpWaveSlot { wave: 4, slot: 0 })
25237 .unwrap_err()
25238 .contains("while wave 3 was next")
25239 );
25240 assert!(
25241 credits
25242 .record_send(PrimePpWaveSlot { wave: 3, slot: 1 })
25243 .unwrap_err()
25244 .contains("reused slot 1")
25245 );
25246 }
25247
25248 #[test]
25249 fn ppn_prime_wave_signal_reports_order_error_injected_error_and_closure() {
25250 let expected = PrimePpWaveSlot { wave: 2, slot: 1 };
25251
25252 let (sender, receiver) = std::sync::mpsc::channel();
25253 sender.send(PrimePpSignal::Slot(expected)).unwrap();
25254 assert_eq!(
25255 recv_prime_pp_signal(&receiver, expected, true, "test").unwrap(),
25256 expected
25257 );
25258
25259 let (sender, receiver) = std::sync::mpsc::channel();
25260 sender
25261 .send(PrimePpSignal::Slot(PrimePpWaveSlot { wave: 3, slot: 1 }))
25262 .unwrap();
25263 assert!(
25264 recv_prime_pp_signal(&receiver, expected, true, "test")
25265 .unwrap_err()
25266 .contains("expected wave/slot")
25267 );
25268
25269 let (sender, receiver) = std::sync::mpsc::channel();
25270 sender
25271 .send(PrimePpSignal::Error("injected stage failure".into()))
25272 .unwrap();
25273 assert_eq!(
25274 recv_prime_pp_signal(&receiver, expected, true, "test").unwrap_err(),
25275 "injected stage failure"
25276 );
25277
25278 let (upstream_sender, upstream_receiver) = std::sync::mpsc::channel();
25279 let (outgoing_sender, outgoing_receiver) = std::sync::mpsc::channel();
25280 let (_release_sender, released_downstream) = std::sync::mpsc::channel();
25281 PrimePpStageChannels {
25282 incoming: None,
25283 release_upstream: Some(upstream_sender),
25284 outgoing: outgoing_sender,
25285 released_downstream,
25286 }
25287 .notify_failure("injected worker error");
25288 assert_eq!(
25289 recv_prime_pp_signal(&upstream_receiver, expected, false, "test").unwrap_err(),
25290 "injected worker error"
25291 );
25292 assert_eq!(
25293 recv_prime_pp_signal(&outgoing_receiver, expected, false, "test").unwrap_err(),
25294 "injected worker error"
25295 );
25296
25297 let (sender, receiver) = std::sync::mpsc::channel::<PrimePpSignal>();
25298 drop(sender);
25299 assert!(
25300 recv_prime_pp_signal(&receiver, expected, true, "test")
25301 .unwrap_err()
25302 .contains("channel closed while waiting for wave 2")
25303 );
25304 }
25305
25306 #[test]
25311 fn auto_prime_ranges_align_to_the_gdn_grid() {
25312 let c = 32usize; let assert_covers = |ranges: &[(usize, usize)], t: usize| {
25314 assert_eq!(ranges.first().map(|&(s, _)| s), Some(0));
25315 assert_eq!(ranges.last().map(|&(_, e)| e), Some(t));
25316 for w in ranges.windows(2) {
25317 assert_eq!(w[0].1, w[1].0, "ranges must stay contiguous");
25318 }
25319 assert!(ranges.iter().all(|&(s, e)| e > s), "no empty range");
25320 };
25321
25322 let t = 9510usize;
25325 let fill = auto_chunk(t);
25326 let fixed = fixed_prime_chunk_ranges(t, fill);
25327 assert!(
25328 fixed[..fixed.len() - 1].iter().any(|&(_, e)| e % c != 0),
25329 "broken arm vanished: fixed auto boundaries all landed on-grid"
25330 );
25331 let dynamic = dynamic_prime_chunk_ranges(t, fill, &fixed);
25332 assert!(
25333 dynamic[..dynamic.len() - 1]
25334 .iter()
25335 .any(|&(_, e)| e % c != 0),
25336 "broken arm vanished: dynamic auto boundaries all landed on-grid"
25337 );
25338
25339 for ranges in [&fixed, &dynamic] {
25340 let aligned = align_prime_ranges_to_gdn(ranges, t, c);
25341 assert_covers(&aligned, t);
25342 for &(_, e) in &aligned[..aligned.len() - 1] {
25343 assert_eq!(e % c, 0, "internal boundary {e} off the {c}-grid");
25344 }
25345 for (&(_, a), &(_, b)) in aligned.iter().zip(ranges.iter()) {
25347 assert!(a <= b && b - a < c);
25348 }
25349 }
25350
25351 let tight = vec![(0usize, 33usize), (33, 40), (40, 200)];
25354 let aligned = align_prime_ranges_to_gdn(&tight, 200, c);
25355 assert_covers(&aligned, 200);
25356 assert_eq!(aligned, vec![(0, 32), (32, 200)]);
25357
25358 assert_eq!(align_prime_ranges_to_gdn(&[(0, 200)], 200, c), [(0, 200)]);
25360 assert_eq!(align_prime_ranges_to_gdn(&tight, 200, 0), tight.as_slice());
25361 let on_grid = vec![(0usize, 128usize), (128, 256), (256, 300)];
25362 assert_eq!(
25363 align_prime_ranges_to_gdn(&on_grid, 300, c),
25364 on_grid.as_slice()
25365 );
25366 }
25367
25368 #[test]
25369 fn active_matrix_prefix_scopes_reused_prime_slabs() {
25370 assert_eq!(
25371 active_matrix_values(40 * 4096, 29, 4096, "activation").unwrap(),
25372 29 * 4096
25373 );
25374 assert_eq!(
25375 active_matrix_values(29 * 4096, 29, 4096, "activation").unwrap(),
25376 29 * 4096
25377 );
25378 assert_eq!(
25379 active_matrix_values(29 * 4096, 24, 4096, "activation").unwrap(),
25380 24 * 4096
25381 );
25382 assert!(active_matrix_values(28 * 4096, 29, 4096, "activation").is_err());
25383 assert!(active_matrix_values(usize::MAX, usize::MAX, 2, "activation").is_err());
25384 }
25385
25386 #[test]
25387 fn step_tp_prefill_batch_refuses_before_scheduler_fallback() {
25388 assert!(validate_step_prime_batch_modes(false, false).is_ok());
25389
25390 let grouped_without_tp = validate_step_prime_batch_modes(false, true).unwrap_err();
25391 assert!(grouped_without_tp.contains("requires MEMRA_STEP_TP_PREFILL=1"));
25392
25393 for grouped in [false, true] {
25394 let err = validate_step_prime_batch_modes(true, grouped).unwrap_err();
25395 assert!(err.contains("did not clear the live-server performance gate"));
25396 assert!(err.contains("per-session grouped prefill"));
25397 }
25398 }
25399
25400 #[test]
25401 fn step_grouped_path_is_eager_single_token_only() {
25402 assert!(step_grouped_decode_shape(false, 1));
25403 assert!(!step_grouped_decode_shape(true, 1));
25404 assert!(!step_grouped_decode_shape(false, 2));
25405 assert!(!step_grouped_decode_shape(true, 2));
25406 }
25407
25408 #[test]
25409 fn step_grouped_prefill_door_is_strict_and_capacity_bounded() {
25410 assert!(!parse_step_ep_grouped_prefill(None).unwrap());
25411 assert!(!parse_step_ep_grouped_prefill(Some("")).unwrap());
25412 assert!(!parse_step_ep_grouped_prefill(Some("0")).unwrap());
25413 assert!(parse_step_ep_grouped_prefill(Some("1")).unwrap());
25414 assert!(parse_step_ep_grouped_prefill(Some("true")).is_err());
25415 assert!(parse_step_ep_grouped_prefill(Some("2")).is_err());
25416
25417 assert!(step_grouped_prefill_shape(true, true, PRIME_MIN_T));
25418 assert!(step_grouped_prefill_shape(
25419 true,
25420 true,
25421 crate::cache::PRIME_CHUNK_MAX_TOKENS,
25422 ));
25423 assert!(!step_grouped_prefill_shape(true, true, PRIME_MIN_T - 1,));
25424 assert!(!step_grouped_prefill_shape(
25425 true,
25426 true,
25427 crate::cache::PRIME_CHUNK_MAX_TOKENS + 1,
25428 ));
25429 assert!(!step_grouped_prefill_shape(false, true, PRIME_MIN_T));
25430 assert!(!step_grouped_prefill_shape(true, false, PRIME_MIN_T));
25431 }
25432
25433 #[test]
25434 fn step_tp_prefill_door_is_strict_and_default_off() {
25435 assert!(!parse_step_tp_prefill(None).unwrap());
25436 assert!(!parse_step_tp_prefill(Some("")).unwrap());
25437 assert!(!parse_step_tp_prefill(Some("0")).unwrap());
25438 assert!(parse_step_tp_prefill(Some("1")).unwrap());
25439 assert!(parse_step_tp_prefill(Some("true")).is_err());
25440 assert!(parse_step_tp_prefill(Some("2")).is_err());
25441 }
25442
25443 #[test]
25444 fn step_tp_prefill_requires_a_qualified_even_rank_shape() {
25445 assert!(step_tp_prefill_shape(
25446 true,
25447 PRIME_MIN_T,
25448 4,
25449 true,
25450 true,
25451 false,
25452 ));
25453 assert!(!step_tp_prefill_shape(
25454 false,
25455 PRIME_MIN_T,
25456 4,
25457 true,
25458 true,
25459 false,
25460 ));
25461 assert!(!step_tp_prefill_shape(
25462 true,
25463 PRIME_MIN_T - 1,
25464 4,
25465 true,
25466 true,
25467 false,
25468 ));
25469 assert!(step_tp_prefill_shape(
25471 true,
25472 PRIME_MIN_T,
25473 2,
25474 true,
25475 true,
25476 false
25477 ));
25478 assert!(!step_tp_prefill_shape(
25479 true,
25480 PRIME_MIN_T,
25481 1,
25482 true,
25483 true,
25484 false
25485 ));
25486 assert!(!step_tp_prefill_shape(
25487 true,
25488 PRIME_MIN_T,
25489 3,
25490 true,
25491 true,
25492 false
25493 ));
25494 assert!(!step_tp_prefill_shape(
25495 true,
25496 PRIME_MIN_T,
25497 4,
25498 false,
25499 true,
25500 false,
25501 ));
25502 assert!(!step_tp_prefill_shape(
25503 true,
25504 PRIME_MIN_T,
25505 4,
25506 true,
25507 false,
25508 false,
25509 ));
25510 assert!(!step_tp_prefill_shape(
25511 true,
25512 PRIME_MIN_T,
25513 4,
25514 true,
25515 true,
25516 true,
25517 ));
25518 }
25519
25520 #[test]
25521 fn fixed_schedule_retains_measured_geometry() {
25522 assert_eq!(
25523 sizes(&fixed_prime_chunk_ranges(461, 128)),
25524 vec![128, 128, 128, 77]
25525 );
25526 assert_eq!(
25527 sizes(&fixed_prime_chunk_ranges(1833, 230)),
25528 vec![230, 230, 230, 230, 230, 230, 230, 223]
25529 );
25530 assert_eq!(sizes(&fixed_prime_chunk_ranges(4096, 512)), vec![512; 8]);
25531 let capped = sizes(&fixed_prime_chunk_ranges_for_ring(8200, 4096, true));
25532 assert_eq!(capped, vec![4096, 4088, 16]);
25533 assert!(capped.iter().all(|&rows| rows <= 4096));
25534 assert_eq!(
25535 sizes(&fixed_prime_chunk_ranges_for_ring(4100, 4096, false)),
25536 vec![4100],
25537 "flag-off schedule remains byte-for-byte the legacy monolithic tail",
25538 );
25539 }
25540
25541 #[test]
25542 fn dynamic_schedule_matches_registered_shapes() {
25543 let cases = [
25544 (461, vec![64, 141, 132, 124]),
25545 (1833, vec![115, 269, 260, 252, 244, 237, 231, 225]),
25546 (4096, vec![256, 602, 580, 563, 545, 531, 516, 503]),
25547 ];
25548 for (t, expected) in cases {
25549 let chunk = auto_chunk(t);
25550 let fixed = fixed_prime_chunk_ranges(t, chunk);
25551 assert_eq!(
25552 sizes(&dynamic_prime_chunk_ranges(t, chunk, &fixed)),
25553 expected
25554 );
25555 }
25556 }
25557
25558 #[test]
25559 fn dynamic_schedule_covers_exactly_and_shrinks_after_fill() {
25560 for t in 256..=8192 {
25561 let chunk = auto_chunk(t);
25562 let fixed = fixed_prime_chunk_ranges(t, chunk);
25563 let dynamic = dynamic_prime_chunk_ranges(t, chunk, &fixed);
25564 assert_eq!(dynamic.len(), fixed.len(), "T={t}");
25565 assert_eq!(dynamic.first().unwrap().0, 0, "T={t}");
25566 assert_eq!(dynamic.last().unwrap().1, t, "T={t}");
25567 for pair in dynamic.windows(2) {
25568 assert_eq!(pair[0].1, pair[1].0, "T={t}");
25569 }
25570 assert!(
25571 dynamic
25572 .iter()
25573 .all(|(start, end)| end - start >= PRIME_MIN_T),
25574 "T={t} sizes={:?}",
25575 sizes(&dynamic)
25576 );
25577 if dynamic.len() >= 3 {
25578 let chunk_sizes = sizes(&dynamic);
25579 assert!(
25580 chunk_sizes[0] < chunk_sizes[1],
25581 "T={t} sizes={chunk_sizes:?}"
25582 );
25583 assert!(
25584 chunk_sizes[1..].windows(2).all(|pair| pair[0] >= pair[1]),
25585 "T={t} sizes={chunk_sizes:?}"
25586 );
25587 }
25588 }
25589 }
25590}
25591
25592#[cfg(test)]
25593mod page_prefetch_tests {
25594 use super::{
25595 grouped_worker_prefetch_position, page_prefetch_positions,
25596 page_prefetch_window_from_values, worker_prefetch_positions,
25597 };
25598
25599 #[test]
25600 fn page_prefetch_window_keeps_existing_opt_in_default() {
25601 assert_eq!(page_prefetch_window_from_values(false, None), 0);
25602 assert_eq!(page_prefetch_window_from_values(false, Some("8")), 0);
25603 assert_eq!(page_prefetch_window_from_values(true, None), 1);
25604 assert_eq!(page_prefetch_window_from_values(true, Some("bad")), 1);
25605 assert_eq!(page_prefetch_window_from_values(true, Some("0")), 0);
25606 assert_eq!(page_prefetch_window_from_values(true, Some("8")), 8);
25607 }
25608
25609 #[test]
25610 fn rolling_page_prefetch_advises_each_future_expert_once() {
25611 let advised: Vec<_> = (0..7)
25612 .flat_map(|position| page_prefetch_positions(position, 7, 3))
25613 .collect();
25614 assert_eq!(advised, vec![1, 2, 3, 4, 5, 6]);
25615
25616 let one_ahead: Vec<_> = (0..4)
25617 .flat_map(|position| page_prefetch_positions(position, 4, 1))
25618 .collect();
25619 assert_eq!(one_ahead, vec![1, 2, 3]);
25620 assert!(page_prefetch_positions(0, 4, 0).is_empty());
25621 }
25622
25623 #[test]
25624 fn grouped_worker_prefetch_primes_first_then_each_known_next_once() {
25625 assert_eq!(grouped_worker_prefetch_position(0, None), None);
25626 let positions: Vec<_> = std::iter::once(grouped_worker_prefetch_position(4, None).unwrap())
25627 .chain(
25628 (0..4).filter_map(|position| grouped_worker_prefetch_position(4, Some(position))),
25629 )
25630 .collect();
25631 assert_eq!(positions, vec![0, 1, 2, 3]);
25632 assert_eq!(grouped_worker_prefetch_position(1, Some(0)), None);
25633 }
25634
25635 #[test]
25636 fn rolling_worker_prefetch_primes_current_and_each_future_expert_once() {
25637 let queued: Vec<_> = (0..8)
25638 .flat_map(|position| worker_prefetch_positions(position, 8, 5))
25639 .collect();
25640 assert_eq!(queued, (0..8).collect::<Vec<_>>());
25641
25642 let one_at_a_time: Vec<_> = (0..4)
25643 .flat_map(|position| worker_prefetch_positions(position, 4, 1))
25644 .collect();
25645 assert_eq!(one_at_a_time, vec![0, 1, 2, 3]);
25646 assert!(worker_prefetch_positions(0, 4, 0).is_empty());
25647 }
25648}
25649
25650pub struct G4DcSlots {
25651 x: CudaSlice<f32>,
25652 xn: CudaSlice<f32>,
25653 cur: CudaSlice<f32>,
25654 hq: CudaSlice<i8>,
25655 hd_: CudaSlice<f32>,
25656 q0: CudaSlice<f32>,
25657 k0: CudaSlice<f32>,
25658 v0: CudaSlice<f32>,
25659 q: CudaSlice<f32>,
25660 k: CudaSlice<f32>,
25661 v: CudaSlice<f32>,
25662 attn: CudaSlice<f32>,
25663 o: CudaSlice<f32>,
25664 attn_out: CudaSlice<f32>,
25665 zsh: CudaSlice<f32>,
25666 zq: CudaSlice<i8>,
25667 zd: CudaSlice<f32>,
25668 gate: CudaSlice<f32>,
25669 up: CudaSlice<f32>,
25670 act: CudaSlice<f32>,
25671 actq: CudaSlice<i8>,
25672 actd: CudaSlice<f32>,
25673 f0: CudaSlice<f32>,
25674 sn: CudaSlice<f32>,
25675 hn: CudaSlice<f32>,
25676 logits: CudaSlice<f32>,
25677}
25678
25679pub struct Step35TokenGraphState {
25684 pub graphs: Vec<(usize, crate::tp::TokenGraph)>,
25686 pub token_d: cudarc::driver::CudaSlice<u32>,
25687 pub pos_d: cudarc::driver::CudaSlice<i32>,
25688 pub logits_stage: cudarc::driver::CudaSlice<f32>,
25689 pub x: cudarc::driver::CudaSlice<f32>,
25694 pub x1: cudarc::driver::CudaSlice<f32>,
25695 pub mixed_stage: cudarc::driver::CudaSlice<f32>,
25696 pub sh_stage: cudarc::driver::CudaSlice<f32>,
25697 pub k_shadow_stage: cudarc::driver::CudaSlice<f32>,
25698 pub v_shadow_stage: cudarc::driver::CudaSlice<f32>,
25699 pub router_logits: cudarc::driver::CudaSlice<f32>,
25702 pub shexp_gate: cudarc::driver::CudaSlice<f32>,
25703 pub shexp_up: cudarc::driver::CudaSlice<f32>,
25704 pub shexp_act: cudarc::driver::CudaSlice<f32>,
25705 pub gate_sig: cudarc::driver::CudaSlice<f32>,
25706 pub dense_z: cudarc::driver::CudaSlice<f32>,
25707 pub dense_gate: cudarc::driver::CudaSlice<f32>,
25708 pub dense_up: cudarc::driver::CudaSlice<f32>,
25709 pub dense_act: cudarc::driver::CudaSlice<f32>,
25710 pub hn: cudarc::driver::CudaSlice<f32>,
25711 pub probe_mixed: cudarc::driver::CudaSlice<f32>,
25714 pub probe_x: cudarc::driver::CudaSlice<f32>,
25715 pub token_hist: cudarc::driver::CudaSlice<u32>,
25718 pub hist_idx: cudarc::driver::CudaSlice<i32>,
25719}
25720
25721impl HybridModel {
25722 #[allow(clippy::type_complexity)] pub(crate) fn step35_token_graph_step(
25734 &self,
25735 e: &Engine,
25736 token: u32,
25737 cache: &mut Cache,
25738 ) -> Result<Option<(Vec<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
25739 if !self.uses_sliding_gated_moe_program()
25740 || !crate::tp::step_tp_graph_enabled()?
25741 || !crate::tp::step_tp_dcw_enabled()?
25742 || !crate::tp::step_tp_qkv_fused_enabled()?
25743 || !crate::tp::step_tp_dev_router_enabled()?
25744 || !crate::tp::step_nvfp4_dev_routes_enabled()?
25745 {
25746 return Ok(None);
25747 }
25748 if !crate::spec::graph_launch_headroom_ok(e) {
25754 static NOTED: std::sync::Once = std::sync::Once::new();
25755 NOTED.call_once(|| crate::spec::graph_replay_suspended_note("step-tp-token"));
25756 return Ok(None);
25757 }
25758 let n_embd = self.cfg.n_embd as usize;
25759 let n_vocab = self.cfg.n_vocab as usize;
25760 let n_layers = self.layers.len();
25761 let pos = cache.pos;
25762 let staged_next = pos + 1;
25763 if staged_next < 96 {
25764 return Ok(None); }
25766
25767 for il in 0..n_layers {
25770 let Some(tp_kv) = cache.tp_kv[il].as_ref() else {
25771 return Ok(None); };
25773 if tp_kv.peek_append_ring(1)?.1 {
25774 return Ok(None);
25775 }
25776 }
25777
25778 let (fa_vec, n_splits) = e.fa_geom_eager(staged_next, 128, 8, false);
25781 if !fa_vec {
25782 return Ok(None);
25783 }
25784 let sp = crate::fa_split_keys(staged_next, 8);
25785 let bucket_max = (n_splits * sp).max(staged_next);
25786
25787 let mut state_guard = self
25788 .step35_token_graph
25789 .lock()
25790 .map_err(|_| "step35 token graph lock is poisoned")?;
25791 if state_guard.is_none() {
25792 let _main = e.gpu.enter_main()?;
25793 let n_expert = self
25794 .cfg
25795 .moe
25796 .as_ref()
25797 .map(|m| m.expert_count as usize)
25798 .unwrap_or(0);
25799 let n_ff_sh = self
25800 .layers
25801 .iter()
25802 .find_map(|l| match &l.ffn {
25803 crate::hybrid::Ffn::Moe(m) => m.gate_shexp.as_ref().map(|g| g.out_features()),
25804 _ => None,
25805 })
25806 .unwrap_or(0);
25807 let n_ff_dense = self
25808 .layers
25809 .iter()
25810 .find_map(|l| match &l.ffn {
25811 crate::hybrid::Ffn::Dense { ffn_gate, .. } => Some(ffn_gate.out_features()),
25812 _ => None,
25813 })
25814 .unwrap_or(0);
25815 *state_guard = Some(Step35TokenGraphState {
25816 graphs: Vec::new(),
25817 token_d: e.stream().clone_htod(&[0u32])?,
25818 pos_d: e.htod_i32(&[pos as i32])?,
25819 logits_stage: e.htod(&vec![0.0f32; n_vocab])?,
25820 x: e.htod(&vec![0.0f32; n_embd])?,
25821 x1: e.htod(&vec![0.0f32; n_embd])?,
25822 mixed_stage: e.htod(&vec![0.0f32; n_embd])?,
25823 sh_stage: e.htod(&vec![0.0f32; n_embd])?,
25824 k_shadow_stage: e.htod(&vec![0.0f32; 2048])?,
25825 v_shadow_stage: e.htod(&vec![0.0f32; 2048])?,
25826 router_logits: e.htod(&vec![0.0f32; n_expert.max(1)])?,
25827 shexp_gate: e.htod(&vec![0.0f32; n_ff_sh.max(1)])?,
25828 shexp_up: e.htod(&vec![0.0f32; n_ff_sh.max(1)])?,
25829 shexp_act: e.htod(&vec![0.0f32; n_ff_sh.max(1)])?,
25830 gate_sig: e.htod(&[1.0f32; 1])?,
25831 dense_z: e.htod(&vec![0.0f32; n_embd])?,
25832 dense_gate: e.htod(&vec![0.0f32; n_ff_dense.max(1)])?,
25833 dense_up: e.htod(&vec![0.0f32; n_ff_dense.max(1)])?,
25834 dense_act: e.htod(&vec![0.0f32; n_ff_dense.max(1)])?,
25835 hn: e.htod(&vec![0.0f32; n_embd])?,
25836 probe_mixed: e.htod(&vec![0.0f32; n_embd])?,
25837 probe_x: e.htod(&vec![0.0f32; n_embd])?,
25838 token_hist: e.stream().clone_htod(&[0u32; 16])?,
25839 hist_idx: e.htod_i32(&[0])?,
25840 });
25841 }
25842 let state = state_guard.as_mut().expect("state armed above");
25843 {
25847 let _main = e.gpu.enter_main()?;
25848 let Step35TokenGraphState {
25849 logits_stage,
25850 token_d,
25851 ..
25852 } = &mut *state;
25853 e.argmax_token_device_into(logits_stage, token_d, n_vocab)?;
25854 }
25855
25856 if state.graphs.is_empty() {
25861 self.step35_token_graph_build(e, cache, state, bucket_max)?;
25864 }
25865 {
25866 let (b, g) = state.graphs.first_mut().expect("graph built above");
25867 if *b != bucket_max {
25868 g.retarget_bucket(bucket_max)?;
25869 *b = bucket_max;
25870 }
25871 }
25872 let graph = state
25873 .graphs
25874 .first()
25875 .map(|(_, g)| g)
25876 .expect("graph built above");
25877
25878 let tg_timing = std::env::var("MEMRA_STEP_TP_TIMING").as_deref() == Ok("1");
25879 let t_fence = tg_timing.then(std::time::Instant::now);
25880 {
25885 let fa0 = match &self.layers[0].mixer {
25886 Mixer::Full(fa) => fa,
25887 _ => return Err("step35 token graph expects full-attention layers".into()),
25888 };
25889 let tp0 = fa0
25890 .step_tp_qkv
25891 .as_ref()
25892 .ok_or("step35 token graph lost its TP state")?;
25893 for rank in 0..tp0.runtime.devices().len() {
25894 let engine = tp0
25895 .runtime
25896 .rank_engine(rank)
25897 .ok_or("step35 token graph lost a rank engine")?;
25898 let _main = engine.gpu.enter_main()?;
25899 engine.stream().synchronize()?;
25900 }
25901 }
25902
25903 {
25905 let _main = e.gpu.enter_main()?;
25906 e.set_u32_one(&mut state.token_d, token)?;
25907 e.set_i32_one(&mut state.pos_d, pos as i32)?;
25908 }
25909 let t_launch = tg_timing.then(std::time::Instant::now);
25910 graph.launch(e)?;
25911 let t_book = tg_timing.then(std::time::Instant::now);
25912 for il in 0..n_layers {
25917 let tp_kv = cache.tp_kv[il].as_mut().expect("eligibility checked above");
25918 let transaction = tp_kv.begin_transaction()?;
25919 let fa = match &self.layers[il].mixer {
25920 Mixer::Full(fa) => fa,
25921 _ => return Err("step35 token graph expects full-attention layers".into()),
25922 };
25923 let tp = fa
25924 .step_tp_qkv
25925 .as_ref()
25926 .ok_or("step35 token graph lost its TP state")?;
25927 let empty: [CudaSlice<f32>; 0] = [];
25930 tp.runtime.append_tp_kv_transaction_inner(
25931 tp_kv,
25932 transaction,
25933 &empty,
25934 &empty,
25935 1,
25936 true,
25937 )?;
25938 tp.runtime
25939 .commit_tp_kv_transaction_external(tp_kv, transaction, 1)?;
25940 if let Some(local) = cache.kv[il].as_mut() {
25942 local.len = pos + 1;
25943 let _main = e.gpu.enter_main()?;
25944 e.set_i32_one(&mut local.len_d, (pos + 1) as i32)?;
25945 }
25946 }
25947 cache.pos = pos + 1;
25948 let t_sync = tg_timing.then(std::time::Instant::now);
25949 let (logits, h_seed) = {
25950 let _main = e.gpu.enter_main()?;
25951 e.stream().synchronize()?;
25952 (e.dtoh(&state.logits_stage)?, e.clone_dtod(&state.x)?)
25953 };
25954 if let (Some(f), Some(l), Some(b), Some(sy)) = (t_fence, t_launch, t_book, t_sync) {
25955 use std::sync::atomic::{AtomicU64, Ordering};
25956 static NS: [AtomicU64; 5] = [
25957 AtomicU64::new(0),
25958 AtomicU64::new(0),
25959 AtomicU64::new(0),
25960 AtomicU64::new(0),
25961 AtomicU64::new(0),
25962 ];
25963 static CALLS: AtomicU64 = AtomicU64::new(0);
25964 let now = std::time::Instant::now();
25965 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;
25971 if calls.is_multiple_of(100) {
25972 let avg = |i: usize| NS[i].load(Ordering::Relaxed) as f64 / calls as f64 / 1e3;
25973 eprintln!(
25974 "[tg-timing] calls={calls} fence_us={:.0} launch_us={:.0} book_us={:.0} \
25975 syncdtoh_us={:.0} total_us={:.0}",
25976 avg(0),
25977 avg(1),
25978 avg(2),
25979 avg(3),
25980 avg(4)
25981 );
25982 }
25983 }
25984 if std::env::var("MEMRA_TG_PROBE_LAYER").is_ok() {
25986 use std::io::Write;
25987 let (pm, px) = {
25988 let _main = e.gpu.enter_main()?;
25989 (e.dtoh(&state.probe_mixed)?, e.dtoh(&state.probe_x)?)
25990 };
25991 for (path, data) in [
25992 ("/root/tg-probe-mixed.bin", &pm),
25993 ("/root/tg-probe-x.bin", &px),
25994 ] {
25995 let mut fo = std::fs::OpenOptions::new()
25996 .create(true)
25997 .append(true)
25998 .open(path)?;
25999 for v in data {
26000 fo.write_all(&v.to_le_bytes())?;
26001 }
26002 }
26003 }
26004 if let Ok(path) = std::env::var("MEMRA_DUMP_HN") {
26007 let hh = {
26008 let _main = e.gpu.enter_main()?;
26009 e.dtoh(&state.hn)?
26010 };
26011 use std::io::Write;
26012 let mut fo = std::fs::OpenOptions::new()
26013 .create(true)
26014 .append(true)
26015 .open(path)?;
26016 for v in &hh {
26017 fo.write_all(&v.to_le_bytes())?;
26018 }
26019 }
26020 if std::env::var("MEMRA_STEP_TP_GRAPH_DEBUG").as_deref() == Ok("1") {
26023 for il in [0usize, 1, 44] {
26024 let tp_kv = cache.tp_kv[il].as_ref().expect("eligibility checked above");
26025 let host_len = tp_kv.staged_len();
26026 let fa = match &self.layers[il].mixer {
26027 Mixer::Full(fa) => fa,
26028 _ => continue,
26029 };
26030 let tp = fa
26031 .step_tp_qkv
26032 .as_ref()
26033 .ok_or("step35 token graph lost its TP state")?;
26034 for rank in 0..tp.runtime.devices().len() {
26035 let engine = tp
26036 .runtime
26037 .rank_engine(rank)
26038 .ok_or("step35 token graph lost a rank engine")?;
26039 let rank_cache = tp_kv.rank(rank).ok_or("debug rank cache missing")?;
26040 let _main = engine.gpu.enter_main()?;
26041 engine.stream().synchronize()?;
26042 let len_d = engine.dtoh_i32_one(rank_cache.len_d())?;
26043 let base_d = match rank_cache.base_d() {
26044 Some(b) => engine.dtoh_i32_one(b)?,
26045 None => -1,
26046 };
26047 eprintln!(
26048 "[graph-debug] pos={pos} il={il} rank={rank} host_len={host_len} \
26049 len_d={len_d} base_d={base_d}"
26050 );
26051 }
26052 }
26053 }
26054 Ok(Some((logits, h_seed)))
26055 }
26056
26057 pub(crate) fn head_split_matvec(
26063 &self,
26064 e: &Engine,
26065 hn: &CudaSlice<f32>,
26066 ) -> Result<Option<Vec<f32>>, Box<dyn std::error::Error>> {
26067 if self.head_split_fill_device(e, hn)?.is_none() {
26068 return Ok(None);
26069 }
26070 let guard = HEAD_SPLIT_WS
26071 .lock()
26072 .map_err(|_| "head split lock is poisoned")?;
26073 let ws = guard.as_ref().expect("filled above");
26074 let _main = e.gpu.enter_main()?;
26075 Ok(Some(e.dtoh(&ws.logits_e)?))
26076 }
26077
26078 fn head_split_fill_device(
26082 &self,
26083 e: &Engine,
26084 hn: &CudaSlice<f32>,
26085 ) -> Result<Option<()>, Box<dyn std::error::Error>> {
26086 use cudarc::driver::DevicePtr;
26087 let crate::model::GpuTensor::FloatBf16 { data: head, .. } = &self.output else {
26088 return Ok(None);
26089 };
26090 let Some(rank1) = self.layers.first().and_then(|l| match &l.mixer {
26091 Mixer::Full(fa) => fa
26092 .step_tp_qkv
26093 .as_ref()
26094 .and_then(|tp| tp.runtime.rank_engine(1)),
26095 _ => None,
26096 }) else {
26097 return Ok(None);
26098 };
26099 let n_embd = self.cfg.n_embd as usize;
26100 let n_vocab = self.cfg.n_vocab as usize;
26101 let half = n_vocab / 2;
26102 let mut guard = HEAD_SPLIT_WS
26103 .lock()
26104 .map_err(|_| "head split lock is poisoned")?;
26105 let pin = {
26106 let _main = e.gpu.enter_main()?;
26107 let stream = e.stream();
26108 let (ptr, _g) = head.device_ptr(&stream);
26109 ptr
26110 };
26111 if guard.as_ref().is_none_or(|ws| ws.pin != pin) {
26112 let hi_rows = n_vocab - half;
26114 let (w1, hn1, y1, ev_done) = {
26115 let _r1 = rank1.gpu.enter_main()?;
26116 (
26117 rank1.alloc_u8_uninit(hi_rows * n_embd * 2)?,
26118 rank1.htod(&vec![0.0f32; n_embd])?,
26119 rank1.htod(&vec![0.0f32; hi_rows])?,
26120 rank1.ctx().new_event(None)?,
26121 )
26122 };
26123 {
26124 use cudarc::driver::sys;
26125 let src = pin + (half * n_embd * 2) as u64;
26126 let dst = {
26127 let _r1 = rank1.gpu.enter_main()?;
26128 let rstream = rank1.stream();
26129 let (d, _g) = w1.device_ptr(&rstream);
26130 d
26131 };
26132 let _r1 = rank1.gpu.enter_main()?;
26133 let r = unsafe {
26134 sys::cuMemcpyAsync(
26135 dst as sys::CUdeviceptr,
26136 src as sys::CUdeviceptr,
26137 hi_rows * n_embd * 2,
26138 rank1.stream().cu_stream() as sys::CUstream,
26139 )
26140 };
26141 if r != sys::CUresult::CUDA_SUCCESS {
26142 return Err(format!("head split replica upload: {r:?}").into());
26143 }
26144 rank1.stream().synchronize()?;
26145 }
26146 let (logits_e, ev_hn) = {
26147 let _main = e.gpu.enter_main()?;
26148 (e.htod(&vec![0.0f32; n_vocab])?, e.ctx().new_event(None)?)
26149 };
26150 let (raw_hn1, raw_y1) = {
26151 let _r1 = rank1.gpu.enter_main()?;
26152 let rstream = rank1.stream();
26153 let (a, _g0) = hn1.device_ptr(&rstream);
26154 let (b, _g1) = y1.device_ptr(&rstream);
26155 (a, b)
26156 };
26157 let raw_logits_hi = {
26158 let _main = e.gpu.enter_main()?;
26159 let stream = e.stream();
26160 let (l, _g) = logits_e.device_ptr(&stream);
26161 l + (half * 4) as u64
26162 };
26163 *guard = Some(HeadSplit {
26164 pin,
26165 w1,
26166 hn1,
26167 y1,
26168 logits_e,
26169 ev_hn,
26170 ev_done,
26171 raw_hn1,
26172 raw_y1,
26173 raw_logits_hi,
26174 samp: None,
26175 });
26176 }
26177 let ws = guard.as_mut().expect("armed above");
26178 let hi_rows = n_vocab - half;
26179 let raw_hn = {
26181 let _main = e.gpu.enter_main()?;
26182 let stream = e.stream();
26183 let (h, _g) = hn.device_ptr(&stream);
26184 ws.ev_hn.record(&stream)?;
26185 h
26186 };
26187 {
26188 let _r1 = rank1.gpu.enter_main()?;
26189 rank1.stream().wait(&ws.ev_hn)?;
26190 crate::tp::raw_copy_bytes(ws.raw_hn1, raw_hn, n_embd * 4, rank1)?;
26191 let HeadSplit { w1, hn1, y1, .. } = &mut *ws;
26192 rank1.matvec_bf16_into(w1, hn1, y1, n_embd, hi_rows)?;
26193 crate::tp::raw_copy_bytes(ws.raw_logits_hi, ws.raw_y1, hi_rows * 4, rank1)?;
26194 ws.ev_done.record(&rank1.stream())?;
26195 }
26196 {
26197 let _main = e.gpu.enter_main()?;
26198 let head_lo = head.slice(0..half * n_embd * 2);
26199 let HeadSplit { logits_e, .. } = &mut *ws;
26200 e.matvec_bf16_view_into(&head_lo, hn, logits_e, n_embd, half)?;
26202 e.stream().wait(&ws.ev_done)?;
26203 Ok(Some(()))
26204 }
26205 }
26206
26207 pub(crate) fn head_split_argmax_device(
26212 &self,
26213 e: &Engine,
26214 hn: &CudaSlice<f32>,
26215 token_d: &mut CudaSlice<u32>,
26216 ) -> Result<bool, Box<dyn std::error::Error>> {
26217 if self.head_split_fill_device(e, hn)?.is_none() {
26218 return Ok(false);
26219 }
26220 let n_vocab = self.cfg.n_vocab as usize;
26221 let guard = HEAD_SPLIT_WS
26222 .lock()
26223 .map_err(|_| "head split lock is poisoned")?;
26224 let ws = guard.as_ref().expect("filled above");
26225 let _main = e.gpu.enter_main()?;
26226 e.argmax_token_device_into(&ws.logits_e, token_d, n_vocab)?;
26227 Ok(true)
26228 }
26229
26230 pub(crate) fn head_split_sample_device(
26236 &self,
26237 e: &Engine,
26238 hn: &CudaSlice<f32>,
26239 token_d: &mut CudaSlice<u32>,
26240 samp: &crate::decode_batch::DevSamp,
26241 ctr: u32,
26242 ) -> Result<bool, Box<dyn std::error::Error>> {
26243 if self.head_split_fill_device(e, hn)?.is_none() {
26244 return Ok(false);
26245 }
26246 let n_vocab = self.cfg.n_vocab as usize;
26247 let guard = HEAD_SPLIT_WS
26248 .lock()
26249 .map_err(|_| "head split lock is poisoned")?;
26250 let mut guard = guard;
26251 let ws = guard.as_mut().expect("filled above");
26252 let _main = e.gpu.enter_main()?;
26253 if ws.samp.is_none() {
26254 ws.samp = Some(SampScratch {
26255 pb: e.zeros(n_vocab)?,
26256 th: e.zeros(1)?,
26257 z: e.zeros(1)?,
26258 mx: e.zeros(1)?,
26259 rows: e.htod_i32(&[0i32])?,
26260 });
26261 }
26262 let filtered = samp.top_k > 0 || samp.top_p < 1.0 || samp.min_p > 0.0;
26263 let HeadSplit {
26264 logits_e,
26265 samp: scratch,
26266 ..
26267 } = &mut *ws;
26268 let sc = scratch.as_mut().expect("armed above");
26269 if filtered {
26270 e.filter_stats(
26271 logits_e, n_vocab, &sc.rows, &mut sc.th, &mut sc.z, &mut sc.mx, n_vocab, 1,
26272 samp.temp, samp.top_k, samp.top_p, samp.min_p,
26273 )?;
26274 let SampScratch { pb, th, mx, .. } = sc;
26275 e.gumbel_perturb_filtered_col(
26276 logits_e, 0, pb, n_vocab, samp.seed, ctr, samp.temp, mx, th, 0,
26277 )?;
26278 } else {
26279 e.gumbel_perturb_col(logits_e, 0, &mut sc.pb, n_vocab, samp.seed, ctr, samp.temp)?;
26280 }
26281 e.argmax_token_device_col(&sc.pb, 0, n_vocab, token_d, 0)?;
26282 Ok(true)
26283 }
26284
26285 pub(crate) fn head_split_logits_dtoh(
26288 &self,
26289 e: &Engine,
26290 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
26291 let guard = HEAD_SPLIT_WS
26292 .lock()
26293 .map_err(|_| "head split lock is poisoned")?;
26294 let ws = guard.as_ref().ok_or("head split logits not armed")?;
26295 let _main = e.gpu.enter_main()?;
26296 e.dtoh(&ws.logits_e)
26297 }
26298
26299 #[allow(clippy::type_complexity)] pub fn step35_token_graph_chunk(
26309 &self,
26310 e: &Engine,
26311 token: u32,
26312 k_target: usize,
26313 cache: &mut Cache,
26314 ) -> Result<Option<(Vec<u32>, Vec<f32>)>, Box<dyn std::error::Error>> {
26315 if !self.uses_sliding_gated_moe_program()
26316 || !crate::tp::step_tp_graph_enabled()?
26317 || !crate::tp::step_tp_dcw_enabled()?
26318 || !crate::tp::step_tp_qkv_fused_enabled()?
26319 || !crate::tp::step_tp_dev_router_enabled()?
26320 || !crate::tp::step_nvfp4_dev_routes_enabled()?
26321 {
26322 return Ok(None);
26323 }
26324 if !crate::spec::graph_launch_headroom_ok(e) {
26327 static NOTED: std::sync::Once = std::sync::Once::new();
26328 NOTED.call_once(|| crate::spec::graph_replay_suspended_note("step-tp-token"));
26329 return Ok(None);
26330 }
26331 let n_layers = self.layers.len();
26332 let pos = cache.pos;
26333 let staged_next = pos + 1;
26334 if staged_next < 96 {
26335 return Ok(None);
26336 }
26337 let (fa_vec, n_splits) = e.fa_geom_eager(staged_next, 128, 8, false);
26340 if !fa_vec {
26341 return Ok(None);
26342 }
26343 let sp = crate::fa_split_keys(staged_next, 8);
26344 let bucket_max = (n_splits * sp).max(staged_next);
26345 let to_boundary = bucket_max.saturating_sub(staged_next) + 1;
26346 let mut k = k_target.min(to_boundary).min(16);
26347 if k < 2 {
26348 return Ok(None);
26349 }
26350 for il in 0..n_layers {
26352 let Some(tp_kv) = cache.tp_kv[il].as_ref() else {
26353 return Ok(None);
26354 };
26355 while k >= 2 && tp_kv.peek_append_ring(k)?.1 {
26356 k -= 1;
26357 }
26358 if k < 2 {
26359 return Ok(None);
26360 }
26361 }
26362
26363 let mut state_guard = self
26364 .step35_token_graph
26365 .lock()
26366 .map_err(|_| "step35 token graph lock is poisoned")?;
26367 let Some(state) = state_guard.as_mut() else {
26368 return Ok(None); };
26370 if state.graphs.is_empty() {
26371 return Ok(None);
26372 }
26373 {
26374 let (b, g) = state.graphs.first_mut().expect("checked above");
26375 if *b != bucket_max {
26376 g.retarget_bucket(bucket_max)?;
26377 *b = bucket_max;
26378 }
26379 }
26380 let graph = state.graphs.first().map(|(_, g)| g).expect("checked above");
26381
26382 {
26384 let fa0 = match &self.layers[0].mixer {
26385 Mixer::Full(fa) => fa,
26386 _ => return Err("step35 token graph expects full-attention layers".into()),
26387 };
26388 let tp0 = fa0
26389 .step_tp_qkv
26390 .as_ref()
26391 .ok_or("step35 token graph lost its TP state")?;
26392 for rank in 0..tp0.runtime.devices().len() {
26393 let engine = tp0
26394 .runtime
26395 .rank_engine(rank)
26396 .ok_or("step35 token graph lost a rank engine")?;
26397 let _main = engine.gpu.enter_main()?;
26398 engine.stream().synchronize()?;
26399 }
26400 }
26401
26402 {
26405 let _main = e.gpu.enter_main()?;
26406 e.set_u32_one(&mut state.token_d, token)?;
26407 e.set_i32_one(&mut state.pos_d, pos as i32)?;
26408 e.set_i32_one(&mut state.hist_idx, 0)?;
26409 }
26410 for _ in 0..k {
26411 graph.launch(e)?;
26412 }
26413 for il in 0..n_layers {
26415 let tp_kv = cache.tp_kv[il].as_mut().expect("eligibility checked above");
26416 let transaction = tp_kv.begin_transaction()?;
26417 let fa = match &self.layers[il].mixer {
26418 Mixer::Full(fa) => fa,
26419 _ => return Err("step35 token graph expects full-attention layers".into()),
26420 };
26421 let tp = fa
26422 .step_tp_qkv
26423 .as_ref()
26424 .ok_or("step35 token graph lost its TP state")?;
26425 let empty: [CudaSlice<f32>; 0] = [];
26426 tp.runtime.append_tp_kv_transaction_inner(
26427 tp_kv,
26428 transaction,
26429 &empty,
26430 &empty,
26431 k,
26432 true,
26433 )?;
26434 tp.runtime
26435 .commit_tp_kv_transaction_external(tp_kv, transaction, k)?;
26436 if let Some(local) = cache.kv[il].as_mut() {
26437 local.len = pos + k;
26438 let _main = e.gpu.enter_main()?;
26439 e.set_i32_one(&mut local.len_d, (pos + k) as i32)?;
26440 }
26441 }
26442 cache.pos = pos + k;
26443 let (hist, logits) = {
26444 let _main = e.gpu.enter_main()?;
26445 e.stream().synchronize()?;
26446 (e.dtoh_u32(&state.token_hist)?, e.dtoh(&state.logits_stage)?)
26447 };
26448 Ok(Some((hist[..k].to_vec(), logits)))
26449 }
26450}
26451
26452impl HybridModel {
26453 #[allow(clippy::too_many_arguments)]
26458 fn step35_token_graph_build(
26459 &self,
26460 e: &Engine,
26461 cache: &mut Cache,
26462 state: &mut Step35TokenGraphState,
26463 bucket_max: usize,
26464 ) -> Result<(), Box<dyn std::error::Error>> {
26465 use cudarc::driver::DevicePtr;
26466 let n_embd = self.cfg.n_embd as usize;
26467 let eps = self.cfg.rms_eps;
26468 let n_layers = self.layers.len();
26469 let started = std::time::Instant::now();
26470 if !crate::router_kernel_on() {
26471 return Err(
26472 "step35 token graph requires the router kernel (MEMRA_ROUTER_KERNEL=0)".into(),
26473 );
26474 }
26475 if !Engine::bf16_mmv_on() || !n_embd.is_multiple_of(8) {
26476 return Err("step35 token graph requires MEMRA_BF16_MMV bf16-resident matvecs".into());
26477 }
26478
26479 let embd_gpu = self
26481 .embd_gpu_try(e)
26482 .ok_or("step35 token graph could not upload the device embed table")?;
26483 let embd_qtype = match self.embd.ggml_type {
26484 memra_gguf::GgmlType::BF16 => crate::QT_BF16,
26485 memra_gguf::GgmlType::Q8_0 => crate::QT_Q8_0,
26486 other => return Err(format!("token graph embed dtype {other:?} unhandled").into()),
26487 };
26488 let embd_row_bytes = self.embd.raw.len() / self.cfg.n_vocab as usize;
26489
26490 let (p_mixed, p_kshadow, p_vshadow) = {
26492 let _main = e.gpu.enter_main()?;
26493 let stream = e.stream();
26494 let (a, _g) = state.mixed_stage.device_ptr(&stream);
26495 let (b, _g) = state.k_shadow_stage.device_ptr(&stream);
26496 let (c, _g) = state.v_shadow_stage.device_ptr(&stream);
26497 (a, b, c)
26498 };
26499
26500 crate::tp::token_graph_build_begin()?;
26501 let mut group_id: u32 = 0;
26502 for il in 0..n_layers {
26503 let layer = &self.layers[il];
26504 let fa = match &layer.mixer {
26505 Mixer::Full(fa) => fa,
26506 _ => return Err("step35 token graph expects full-attention layers".into()),
26507 };
26508 let tp = fa
26509 .step_tp_qkv
26510 .as_ref()
26511 .ok_or("step35 token graph lost its TP state")?;
26512 let attention = tp
26513 .attention
26514 .as_ref()
26515 .ok_or("step35 token graph lost its attention aux")?;
26516 let geometry = self.step35_geom(il);
26517 let window = geometry.window.map(|w| w as usize);
26518 let head_dim = geometry.head_dim_k as usize;
26519 let heads = geometry.n_head as usize;
26520 let kv_heads = geometry.n_head_kv as usize;
26521 let ranks = tp.runtime.devices().len();
26522 let local_heads = heads / ranks;
26523 let local_kv_heads = kv_heads / ranks;
26524 let layer_bucket = window.map(|w| bucket_max.min(w)).unwrap_or(bucket_max);
26525 let use_gate_shards =
26526 attention.gate_shards.is_some() || attention.gate_shards_bf16.is_some();
26527 if !use_gate_shards {
26528 return Err("step35 token graph requires the fused gate shards".into());
26529 }
26530
26531 let ws_index = tp
26532 .runtime
26533 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
26534 let ws_mutex = tp.runtime.decode_v2_workspace();
26535 let mut ws_guard = ws_mutex
26536 .lock()
26537 .map_err(|_| "step TP decode v2 workspace lock is poisoned")?;
26538 let ws = ws_guard
26539 .get_mut(ws_index)
26540 .ok_or("step TP decode v2 workspace missing after ensure")?;
26541 tp.runtime
26542 .decode_v2_arm_token_mirrors(ws, p_mixed, (p_kshadow, p_vshadow))?;
26543 let mut rope_freqs = Vec::with_capacity(ranks);
26544 for rank in 0..ranks {
26545 let engine = tp
26546 .runtime
26547 .rank_engine(rank)
26548 .ok_or("step35 token graph lost a rank engine")?;
26549 rope_freqs.push(if geometry.rope_factors {
26550 self.step35_aux
26551 .as_ref()
26552 .and_then(|aux| aux.rope_freqs(engine))
26553 } else {
26554 None
26555 });
26556 }
26557 let gate_shards_arg = if let Some(shards) = attention.gate_shards.as_deref() {
26558 Some(crate::tp::StepTpGateShards::F32(shards))
26559 } else {
26560 attention
26561 .gate_shards_bf16
26562 .as_deref()
26563 .map(crate::tp::StepTpGateShards::Bf16)
26564 };
26565
26566 let decode_input = attention
26568 .decode_input
26569 .as_ref()
26570 .ok_or("step35 token graph requires the replicated decode input")?;
26571 let mut decode_input = decode_input
26572 .lock()
26573 .map_err(|_| "replicated decode input lock is poisoned")?;
26574 if ws.h_stage.is_none() {
26576 return Err(
26577 "step35 token graph requires the stage flow armed (run eager dcw first)".into(),
26578 );
26579 }
26580 {
26581 let state_x = &mut state.x;
26582 let token_d = &state.token_d;
26583 let pos_d = &state.pos_d;
26584 crate::tp::graph_section(e, None, || {
26585 let _main = e.gpu.enter_main()?;
26586 if il == 0 {
26587 e.embed_gather_device_into(
26588 embd_gpu,
26589 token_d,
26590 state_x,
26591 n_embd,
26592 embd_qtype,
26593 embd_row_bytes,
26594 )?;
26595 }
26596 {
26597 let h_stage = ws.h_stage.as_mut().expect("stage armed checked above");
26598 e.rms_norm(
26599 state_x,
26600 layer.attn_norm.float_data(),
26601 h_stage,
26602 n_embd,
26603 1,
26604 eps,
26605 )?;
26606 }
26607 {
26608 let pos_stage = ws.pos_stage.as_mut().expect("stage armed above");
26609 let mut dst = pos_stage.slice_mut(0..1);
26610 e.stream().memcpy_dtod(&pos_d.slice(0..1), &mut dst)?;
26611 }
26612 Ok(())
26613 })?;
26614 }
26615
26616 group_id += 1;
26618 for rank in 0..ranks {
26619 let engine = tp
26620 .runtime
26621 .rank_engine(rank)
26622 .ok_or("step35 token graph lost a rank engine")?;
26623 {
26624 let ceiling = window
26629 .map(|w| cache.max_ctx.min(w))
26630 .unwrap_or(cache.max_ctx);
26631 let _main = engine.gpu.enter_main()?;
26632 engine.fa_dcw_pool_ensure(
26633 head_dim,
26634 local_heads,
26635 local_kv_heads,
26636 ceiling.min(2048),
26637 )?;
26638 engine.fa_dcw_pool_ensure(head_dim, local_heads, local_kv_heads, ceiling)?;
26639 engine.fa_dcw_pool_ensure(
26640 head_dim,
26641 local_heads,
26642 local_kv_heads,
26643 layer_bucket,
26644 )?;
26645 }
26646 let runtime = &tp.runtime;
26647 let q_norm = &attention.q_norm;
26648 let k_norm = &attention.k_norm;
26649 let gate_ref = gate_shards_arg.as_ref();
26650 crate::tp::graph_section(engine, Some(group_id), || {
26651 runtime.decode_v2_input_qkv_rank(
26652 ws,
26653 &state.pos_d,
26654 &mut decode_input,
26655 &tp.q,
26656 &tp.k,
26657 &tp.v,
26658 q_norm,
26659 k_norm,
26660 head_dim,
26661 geometry.n_rot as usize,
26662 geometry.rope_base,
26663 &rope_freqs,
26664 eps,
26665 gate_ref,
26666 true,
26667 true,
26668 false,
26669 rank,
26670 None,
26671 )?;
26672 let distributed = cache.tp_kv[il]
26675 .as_mut()
26676 .ok_or("step35 token graph lost a TP cache")?;
26677 let (kv_dim_k, kv_dim_v) = (distributed.kv_dim_k(), distributed.kv_dim_v());
26678 let (ktb, vtb) = (distributed.k_tok_bytes(), distributed.v_tok_bytes());
26679 let capacity = distributed.physical_capacity();
26680 {
26681 let rank_cache = distributed
26682 .rank_mut(rank)
26683 .ok_or("step35 token graph lost a rank cache")?;
26684 let (k_plane, v_plane, len_d, base_d) =
26685 rank_cache.planes_and_counters_mut();
26686 engine.append_kv_quantized_dcw(
26687 &ws.k[rank],
26688 &ws.v_raw[rank],
26689 k_plane,
26690 v_plane,
26691 len_d,
26692 base_d,
26693 kv_dim_k,
26694 kv_dim_v,
26695 ktb,
26696 vtb,
26697 )?;
26698 }
26699 {
26700 let rank_cache = distributed
26701 .rank_mut(rank)
26702 .ok_or("step35 token graph lost a rank cache")?;
26703 engine.inc_i32(rank_cache.len_d_mut())?;
26704 }
26705 let rank_cache = distributed
26706 .rank(rank)
26707 .ok_or("step35 token graph lost a rank cache")?;
26708 let k_ring = engine.view_u8_range(rank_cache.k(), 0, capacity * ktb);
26709 let v_ring = engine.view_u8_range(rank_cache.v(), 0, capacity * vtb);
26710 engine.fa_decode_dcw(
26715 &ws.q[rank],
26716 &k_ring,
26717 &v_ring,
26718 &mut ws.attn_out[rank],
26719 head_dim,
26720 local_heads,
26721 local_kv_heads,
26722 rank_cache.len_d(),
26723 rank_cache.base_d(),
26724 window.unwrap_or(0),
26725 layer_bucket,
26726 geometry.attention_scale(),
26727 ktb,
26728 vtb,
26729 None,
26730 )?;
26731 engine.attn_head_gate(
26732 &ws.attn_out[rank],
26733 &ws.gate[rank],
26734 &mut ws.gated[rank],
26735 None,
26736 head_dim,
26737 local_heads,
26738 1,
26739 )?;
26740 runtime.decode_v2_finish_rank_partial(ws, &tp.o, true, rank)?;
26741 Ok(())
26742 })?;
26743 }
26744
26745 {
26747 let root = tp
26748 .runtime
26749 .rank_engine(0)
26750 .ok_or("step35 token graph lost the root engine")?;
26751 let runtime = &tp.runtime;
26752 crate::tp::graph_section(root, None, || runtime.decode_v2_finish_root_fused(ws))?;
26753 }
26754 drop(ws_guard);
26755 drop(decode_input);
26756
26757 let probe_layer: Option<usize> = std::env::var("MEMRA_TG_PROBE_LAYER")
26758 .ok()
26759 .and_then(|v| v.parse().ok());
26760 if probe_layer == Some(il) {
26761 let Step35TokenGraphState {
26762 mixed_stage,
26763 probe_mixed,
26764 ..
26765 } = &mut *state;
26766 crate::tp::graph_section(e, None, || {
26767 let _main = e.gpu.enter_main()?;
26768 let mut dst = probe_mixed.slice_mut(0..n_embd);
26769 e.stream()
26770 .memcpy_dtod(&mixed_stage.slice(0..n_embd), &mut dst)?;
26771 Ok(())
26772 })?;
26773 }
26774
26775 match &layer.ffn {
26777 crate::hybrid::Ffn::Dense {
26778 ffn_gate,
26779 ffn_up,
26780 ffn_down,
26781 } => {
26782 let n_ff = ffn_gate.out_features();
26783 let lim = self.cfg.clamp_shexp_at(il as u32);
26784 if lim.is_some() {
26788 return Err("step35 token graph dense FFN with clamp unsupported".into());
26789 }
26790 let (wg_d, wu_d, wd_d) = match (ffn_gate, ffn_up, ffn_down) {
26791 (
26792 crate::model::GpuTensor::FloatBf16 { data: wg, .. },
26793 crate::model::GpuTensor::FloatBf16 { data: wu, .. },
26794 crate::model::GpuTensor::FloatBf16 { data: wd, .. },
26795 ) => (wg, wu, wd),
26796 _ => {
26797 return Err(
26798 "step35 token graph dense FFN requires bf16-resident weights"
26799 .into(),
26800 );
26801 }
26802 };
26803 crate::tp::graph_section(e, None, || {
26804 let _main = e.gpu.enter_main()?;
26805 let Step35TokenGraphState {
26806 x,
26807 x1,
26808 mixed_stage,
26809 dense_z,
26810 dense_gate,
26811 dense_up,
26812 dense_act,
26813 sh_stage,
26814 ..
26815 } = &mut *state;
26816 e.add_rms_norm(
26817 x,
26818 mixed_stage,
26819 layer.post_attn_norm.float_data(),
26820 x1,
26821 dense_z,
26822 n_embd,
26823 1,
26824 eps,
26825 )?;
26826 e.matvec_bf16_into(wg_d, dense_z, dense_gate, n_embd, n_ff)?;
26830 e.matvec_bf16_into(wu_d, dense_z, dense_up, n_embd, n_ff)?;
26831 Self::ffn_act_lim(
26832 e, &self.cfg, dense_gate, dense_up, 1.0, 1.0, lim, dense_act, n_ff,
26833 )?;
26834 e.matvec_bf16_into(wd_d, dense_act, sh_stage, n_ff, n_embd)?;
26835 e.add(x1, sh_stage, x, n_embd)?;
26836 Ok(())
26837 })?;
26838 }
26839 crate::hybrid::Ffn::Moe(m) => {
26840 let moe = self
26841 .cfg
26842 .moe
26843 .as_ref()
26844 .ok_or("step35 token graph needs moe cfg")?;
26845 let n_expert = moe.expert_count as usize;
26846 let n_used = moe.expert_used_count as usize;
26847 let sigmoid = self
26848 .cfg
26849 .sigmoid_router()
26850 .ok_or("step35 token graph needs the sigmoid router")?;
26851 let step_tp = m
26852 .step_tp
26853 .as_ref()
26854 .ok_or("step35 token graph needs TP experts")?;
26855 let bank = match &step_tp.experts {
26856 crate::hybrid::StepTpExpertBank::Nvfp4(bank) => bank,
26857 _ => return Err("step35 token graph needs the NVFP4 bank".into()),
26858 };
26859 let routes_ws_mutex = bank.device_workspace_handle();
26860 let mut routes_guard = routes_ws_mutex
26861 .lock()
26862 .map_err(|_| "routes workspace lock is poisoned")?;
26863 let routes_ws = routes_guard
26864 .as_mut()
26865 .ok_or("step35 token graph requires the routes workspace warmed")?;
26866 routes_ws.arm_stages(e, bank.input_width, n_used)?;
26867 step_tp.runtime.routes_arm_raw(bank, routes_ws)?;
26868 let p_z = {
26869 let root = step_tp
26870 .runtime
26871 .rank_engine(0)
26872 .ok_or("routes root engine missing")?;
26873 let _main = root.gpu.enter_main()?;
26874 let stream = root.stream();
26875 let in_stage = routes_ws
26876 .in_stage_handle()
26877 .ok_or("routes in stage not armed")?;
26878 let (a, _g) = in_stage.device_ptr(&stream);
26879 a
26880 };
26881 let local_out = bank.expert_width / ranks;
26882
26883 crate::tp::graph_section(e, None, || {
26885 let _main = e.gpu.enter_main()?;
26886 {
26887 let in_stage = routes_ws
26888 .in_stage_mut()
26889 .ok_or("routes in stage not armed")?;
26890 let Step35TokenGraphState {
26891 x, x1, mixed_stage, ..
26892 } = &mut *state;
26893 e.add_rms_norm(
26894 x,
26895 mixed_stage,
26896 layer.post_attn_norm.float_data(),
26897 x1,
26898 in_stage,
26899 n_embd,
26900 1,
26901 eps,
26902 )?;
26903 }
26904 {
26905 let z_ref = routes_ws
26906 .in_stage_handle()
26907 .ok_or("routes in stage not armed")?;
26908 e.router_gemv_into(
26909 m.gate_inp.float_data(),
26910 z_ref,
26911 &mut state.router_logits,
26912 n_embd,
26913 n_expert,
26914 1,
26915 )?;
26916 }
26917 let (sel_e, w_e) = routes_ws
26918 .dev_route_e_mut()
26919 .ok_or("routes staging not armed")?;
26920 e.moe_router_sigmoid_topk_into(
26921 &state.router_logits,
26922 1,
26923 n_expert,
26924 n_used,
26925 m.active_count(),
26926 &m.exp_probs_b_dev,
26927 &m.active_experts_dev,
26928 sigmoid.0,
26929 sigmoid.1,
26930 sel_e,
26931 w_e,
26932 )?;
26933 Ok(())
26934 })?;
26935
26936 group_id += 1;
26938 for rank in 0..ranks {
26939 let engine = step_tp
26940 .runtime
26941 .rank_engine(rank)
26942 .ok_or("routes rank engine missing")?;
26943 let runtime = &step_tp.runtime;
26944 crate::tp::graph_section(engine, Some(group_id), || {
26945 runtime.routes_rank_section(
26946 bank,
26947 routes_ws,
26948 p_z,
26949 local_out,
26950 n_used,
26951 step_tp.activation_limit,
26952 rank,
26953 )
26954 })?;
26955 }
26956
26957 {
26959 let root = step_tp
26960 .runtime
26961 .rank_engine(0)
26962 .ok_or("routes root engine missing")?;
26963 let runtime = &step_tp.runtime;
26964 crate::tp::graph_section(root, None, || {
26965 runtime.routes_root_section(bank, routes_ws)
26966 })?;
26967 }
26968
26969 let lim_sh = self.cfg.clamp_shexp_at(il as u32);
26973 let (wg_sh, wu_sh, wd_sh) = match (&m.gate_shexp, &m.up_shexp, &m.down_shexp) {
26974 (
26975 Some(crate::model::GpuTensor::FloatBf16 { data: wg, .. }),
26976 Some(crate::model::GpuTensor::FloatBf16 { data: wu, .. }),
26977 Some(crate::model::GpuTensor::FloatBf16 { data: wd, .. }),
26978 ) => (wg, wu, wd),
26979 _ => {
26980 return Err(
26981 "step35 token graph shexp requires bf16-resident weights".into()
26982 );
26983 }
26984 };
26985 let n_ff_sh = m
26986 .gate_shexp
26987 .as_ref()
26988 .expect("matched Some above")
26989 .out_features();
26990 let gate_inp_shexp = m.gate_inp_shexp.as_ref();
26993 crate::tp::graph_section(e, None, || {
26994 let _main = e.gpu.enter_main()?;
26995 let (z_ref, out_stage) = routes_ws
26996 .in_and_out_stages_mut()
26997 .ok_or("routes stages not armed")?;
26998 let Step35TokenGraphState {
26999 x,
27000 x1,
27001 sh_stage,
27002 shexp_gate,
27003 shexp_up,
27004 shexp_act,
27005 gate_sig,
27006 ..
27007 } = &mut *state;
27008 e.matvec_bf16_dual_into(
27009 wg_sh, wu_sh, z_ref, shexp_gate, shexp_up, n_embd, n_ff_sh,
27010 )?;
27011 Self::ffn_act_lim(
27012 e, &self.cfg, shexp_gate, shexp_up, 1.0, 1.0, lim_sh, shexp_act,
27013 n_ff_sh,
27014 )?;
27015 e.matvec_bf16_into(wd_sh, shexp_act, sh_stage, n_ff_sh, n_embd)?;
27016 if let Some(gate_w) = gate_inp_shexp {
27017 e.sigmoid_dot_rows_into(
27018 z_ref,
27019 gate_w.float_data(),
27020 gate_sig,
27021 n_embd,
27022 1,
27023 )?;
27024 }
27025 e.add_scaled_rows(sh_stage, gate_sig, out_stage, n_embd, 1)?;
27026 e.add(x1, out_stage, x, n_embd)?;
27027 Ok(())
27028 })?;
27029 }
27030 }
27031 if probe_layer == Some(il) {
27032 let Step35TokenGraphState { x, probe_x, .. } = &mut *state;
27033 crate::tp::graph_section(e, None, || {
27034 let _main = e.gpu.enter_main()?;
27035 let mut dst = probe_x.slice_mut(0..n_embd);
27036 e.stream().memcpy_dtod(&x.slice(0..n_embd), &mut dst)?;
27037 Ok(())
27038 })?;
27039 }
27040 }
27041
27042 let head = match &self.output {
27044 crate::model::GpuTensor::FloatBf16 { data, .. } => data,
27045 _ => return Err("step35 token graph head requires the bf16-resident output".into()),
27046 };
27047 crate::tp::graph_section(e, None, || {
27048 let _main = e.gpu.enter_main()?;
27049 let Step35TokenGraphState {
27050 x,
27051 hn,
27052 logits_stage,
27053 token_d,
27054 pos_d,
27055 token_hist,
27056 hist_idx,
27057 ..
27058 } = &mut *state;
27059 e.rms_norm(x, self.output_norm.float_data(), hn, n_embd, 1, eps)?;
27060 e.matvec_bf16_into(head, hn, logits_stage, n_embd, self.cfg.n_vocab as usize)?;
27061 e.argmax_token_device_into(logits_stage, token_d, self.cfg.n_vocab as usize)?;
27067 e.u32_hist_append(token_d, token_hist, hist_idx)?;
27068 e.inc_i32(pos_d)?;
27069 Ok(())
27070 })?;
27071
27072 let graph = crate::tp::token_graph_build_finish()?;
27073 state.graphs.push((bucket_max, graph));
27074 eprintln!(
27075 "[step35-token-graph] built bucket={bucket_max} layers={n_layers} \
27076 build_ms={:.0} performance_claim=false",
27077 started.elapsed().as_secs_f64() * 1e3
27078 );
27079 Ok(())
27080 }
27081}