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 glm5_decode_graph: None,
283 })
284 })
285 .collect();
286 Self {
287 parent,
288 fence: fence.to_vec(),
289 stages,
290 committed: false,
291 }
292 }
293
294 fn pp2_parts(&mut self) -> (&mut Cache, &mut Cache) {
295 assert_eq!(self.stages.len(), 2);
296 let (stage0, stage1) = self.stages.split_at_mut(1);
297 (
298 stage0[0]
299 .get_mut()
300 .unwrap_or_else(|poisoned| poisoned.into_inner()),
301 stage1[0]
302 .get_mut()
303 .unwrap_or_else(|poisoned| poisoned.into_inner()),
304 )
305 }
306
307 fn stages(&self) -> &[std::sync::Mutex<Cache>] {
308 &self.stages
309 }
310
311 fn commit(&mut self) {
312 self.committed = true;
313 }
314}
315
316impl Drop for PrimeCacheStages<'_> {
317 fn drop(&mut self) {
318 let n = self.parent.kv.len();
319 for i in 0..n {
320 let stage = prime_cache_stage_for_layer(&self.fence, i);
321 let source = self.stages[stage]
322 .get_mut()
323 .unwrap_or_else(|poisoned| poisoned.into_inner());
324 debug_assert!(self.parent.kv[i].is_none());
325 debug_assert!(self.parent.tp_kv[i].is_none());
326 debug_assert!(self.parent.recur[i].is_none());
327 debug_assert!(self.parent.latent[i].is_none());
328 self.parent.kv[i] = source.kv[i].take();
329 self.parent.tp_kv[i] = source.tp_kv[i].take();
330 self.parent.recur[i] = source.recur[i].take();
331 self.parent.latent[i] = source.latent[i].take();
332 self.parent.glm5_tp_recur[i] = source.glm5_tp_recur[i].take();
333 self.parent.glm5_tp_latent_peer[i] = source.glm5_tp_latent_peer[i].take();
334 }
335 self.parent.pos = self
336 .stages
337 .iter_mut()
338 .map(|stage| {
339 stage
340 .get_mut()
341 .unwrap_or_else(|poisoned| poisoned.into_inner())
342 .pos
343 })
344 .min()
345 .unwrap_or(self.parent.pos);
346 if !self.committed {
347 self.parent.mark_tainted();
348 }
349 }
350}
351
352struct CacheTaintGuard {
356 caches: Vec<*mut Cache>,
357 committed: bool,
358}
359
360impl CacheTaintGuard {
361 fn arm(caches: &mut [&mut Cache]) -> Self {
362 Self {
363 caches: caches
364 .iter_mut()
365 .map(|cache| *cache as *mut Cache)
366 .collect(),
367 committed: false,
368 }
369 }
370
371 fn commit(&mut self) {
372 self.committed = true;
373 }
374}
375
376impl Drop for CacheTaintGuard {
377 fn drop(&mut self) {
378 if self.committed {
379 return;
380 }
381 for cache in &self.caches {
382 unsafe { (&mut **cache).mark_tainted() };
385 }
386 }
387}
388
389#[derive(Debug, Clone, Copy, PartialEq, Eq)]
390struct PrimePpWaveSlot {
391 wave: usize,
392 slot: usize,
393}
394
395#[derive(Debug)]
396enum PrimePpSignal {
397 Slot(PrimePpWaveSlot),
398 Error(String),
399}
400
401#[derive(Default)]
402struct PrimePpWaveCredits {
403 next_wave: usize,
404 pending: std::collections::VecDeque<PrimePpWaveSlot>,
405}
406
407impl PrimePpWaveCredits {
408 fn release_required(&self) -> Option<PrimePpWaveSlot> {
409 (self.pending.len() == 2).then(|| self.pending[0])
410 }
411
412 fn record_release(&mut self, released: PrimePpWaveSlot) -> Result<(), String> {
413 let expected =
414 self.pending.front().copied().ok_or_else(|| {
415 "prime PP received a slot release with no pending wave".to_string()
416 })?;
417 if released != expected {
418 return Err(format!(
419 "prime PP slot release {:?} does not match oldest pending {:?}",
420 released, expected
421 ));
422 }
423 self.pending.pop_front();
424 Ok(())
425 }
426
427 fn record_send(&mut self, sent: PrimePpWaveSlot) -> Result<(), String> {
428 if sent.wave != self.next_wave {
429 return Err(format!(
430 "prime PP sent wave {} while wave {} was next",
431 sent.wave, self.next_wave
432 ));
433 }
434 if sent.slot >= 2 {
435 return Err(format!(
436 "prime PP boundary returned invalid slot {}",
437 sent.slot
438 ));
439 }
440 if self.pending.iter().any(|pending| pending.slot == sent.slot) {
441 return Err(format!(
442 "prime PP reused slot {} before its exact-wave release",
443 sent.slot
444 ));
445 }
446 self.pending.push_back(sent);
447 self.next_wave += 1;
448 Ok(())
449 }
450}
451
452fn recv_prime_pp_signal(
453 receiver: &std::sync::mpsc::Receiver<PrimePpSignal>,
454 expected: PrimePpWaveSlot,
455 exact_slot: bool,
456 label: &str,
457) -> Result<PrimePpWaveSlot, String> {
458 match receiver.recv() {
459 Ok(PrimePpSignal::Error(error)) => Err(error),
460 Ok(PrimePpSignal::Slot(received))
461 if received.wave == expected.wave
462 && (!exact_slot || received.slot == expected.slot) =>
463 {
464 if received.slot >= 2 {
465 Err(format!(
466 "{label}: wave {} carried invalid slot {}",
467 received.wave, received.slot
468 ))
469 } else {
470 Ok(received)
471 }
472 }
473 Ok(PrimePpSignal::Slot(received)) => Err(format!(
474 "{label}: expected wave/slot {:?}, received {:?}",
475 expected, received
476 )),
477 Err(_) => Err(format!(
478 "{label}: channel closed while waiting for wave {}",
479 expected.wave
480 )),
481 }
482}
483
484fn send_prime_pp_signal(
485 sender: &std::sync::mpsc::Sender<PrimePpSignal>,
486 signal: PrimePpSignal,
487 label: &str,
488) -> Result<(), String> {
489 sender
490 .send(signal)
491 .map_err(|_| format!("{label}: channel closed"))
492}
493
494struct PrimePpWave<'a> {
495 start: usize,
496 end: usize,
497 tokens: &'a [u32],
498}
499
500struct PrimePpStageChannels {
501 incoming: Option<std::sync::mpsc::Receiver<PrimePpSignal>>,
502 release_upstream: Option<std::sync::mpsc::Sender<PrimePpSignal>>,
503 outgoing: std::sync::mpsc::Sender<PrimePpSignal>,
504 released_downstream: std::sync::mpsc::Receiver<PrimePpSignal>,
505}
506
507impl PrimePpStageChannels {
508 fn notify_failure(&self, error: &str) {
509 if let Some(upstream) = &self.release_upstream {
510 let _ = upstream.send(PrimePpSignal::Error(error.to_string()));
511 }
512 let _ = self.outgoing.send(PrimePpSignal::Error(error.to_string()));
513 }
514}
515
516pub static MLA_SEG_WS_DISPATCHES: std::sync::atomic::AtomicU64 =
540 std::sync::atomic::AtomicU64::new(0);
541
542pub fn mla_seg_ws_dispatches() -> u64 {
544 MLA_SEG_WS_DISPATCHES.load(std::sync::atomic::Ordering::Relaxed)
545}
546
547pub(crate) struct MlaSegWs {
548 pub q_nope: CudaSlice<f32>,
549 pub q_pe: CudaSlice<f32>,
550 pub q_an: CudaSlice<f32>,
551 pub c_kv_n: CudaSlice<f32>,
552 pub k_pe: CudaSlice<f32>,
553 pub sig: (usize, usize, usize, usize, usize),
555}
556
557impl MlaSegWs {
558 pub(crate) fn new(
560 e: &Engine,
561 nh: usize,
562 dn: usize,
563 dr: usize,
564 r: usize,
565 q_lora: usize,
566 ) -> Result<Self, Box<dyn std::error::Error>> {
567 Ok(Self {
568 q_nope: e.uninit(nh * dn)?,
569 q_pe: e.uninit((nh * dr).max(1))?,
570 q_an: e.uninit(q_lora)?,
571 c_kv_n: e.uninit(r)?,
572 k_pe: e.uninit(dr.max(1))?,
573 sig: (nh, dn, dr, r, q_lora),
574 })
575 }
576}
577
578pub(crate) struct MlaMidIn<'a> {
581 pub q_an: &'a CudaSlice<f32>,
582 pub c_kv_n: &'a CudaSlice<f32>,
583 pub k_pe: &'a CudaSlice<f32>,
584}
585
586pub(crate) struct MlaPreOut {
589 pub q_nope: CudaSlice<f32>,
590 pub q_pe: CudaSlice<f32>,
591 pub q_an: CudaSlice<f32>,
592 pub c_kv_n: CudaSlice<f32>,
593 pub k_pe: CudaSlice<f32>,
594}
595
596pub struct IndexerPlanes<'a> {
597 pub state: &'a mut CudaSlice<f32>,
598 pub pool_keys: &'a mut Option<CudaSlice<f32>>,
599 pub ready: &'a mut usize,
600 pub state_ring_rows: usize,
605 pub capacity_tokens: usize,
609}
610
611pub(crate) struct AttnPre {
613 pub q: cudarc::driver::CudaSlice<f32>,
614 pub k: cudarc::driver::CudaSlice<f32>,
615 pub v: cudarc::driver::CudaSlice<f32>,
616 pub gate: Option<cudarc::driver::CudaSlice<f32>>,
617}
618
619pub(crate) struct GdnPrep {
621 pub hk: usize,
622 pub q_l2: cudarc::driver::CudaSlice<f32>,
623 pub k_l2: cudarc::driver::CudaSlice<f32>,
624 pub v_g: cudarc::driver::CudaSlice<f32>,
625 pub beta: cudarc::driver::CudaSlice<f32>,
626 pub g_log: cudarc::driver::CudaSlice<f32>,
627 pub kb16: Option<cudarc::driver::CudaSlice<u8>>,
628 pub qb16: Option<cudarc::driver::CudaSlice<u8>>,
629}
630
631pub(crate) struct VerifyStreamScratch {
633 pub pos_d: CudaSlice<i32>,
634 pub row_ctrs: Vec<CudaSlice<i32>>,
635}
636use crate::hybrid::{FullAttnLayer, HybridModel, LinearAttnLayer, Mixer, MoeWeights};
637
638struct MoeInputTraceWriter {
639 dir: std::path::PathBuf,
640 index: std::fs::File,
641 payloads: std::collections::HashMap<u16, (std::fs::File, u64)>,
642}
643
644static MOE_INPUT_TRACE_WRITER: std::sync::OnceLock<std::sync::Mutex<Option<MoeInputTraceWriter>>> =
645 std::sync::OnceLock::new();
646
647fn gdec_enabled() -> bool {
650 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
651 *E.get_or_init(|| {
652 std::env::var("MEMRA_MOE_GDEC")
653 .map(|v| v != "0")
654 .unwrap_or(true)
655 })
656}
657
658fn moe_slab_enabled() -> bool {
669 std::env::var("MEMRA_MOE_SLAB").as_deref() != Ok("0")
670}
671
672fn moe_fused_epi_enabled() -> bool {
685 std::env::var("MEMRA_MOE_FUSED_EPI")
686 .map(|v| v != "0")
687 .unwrap_or(false)
688}
689
690pub(crate) static HC_WS_FORCE_PLAIN: std::sync::atomic::AtomicBool =
706 std::sync::atomic::AtomicBool::new(false);
707
708fn hyper_decode_ws_on() -> bool {
709 !HC_WS_FORCE_PLAIN.load(std::sync::atomic::Ordering::Relaxed)
710 && std::env::var("MEMRA_HC_DECODE_WS").as_deref() == Ok("1")
711}
712
713pub static HC_DECODE_WS_DISPATCHES: std::sync::atomic::AtomicU64 =
716 std::sync::atomic::AtomicU64::new(0);
717
718pub static GLM5_Q8_FUSE_DISPATCHES: std::sync::atomic::AtomicU64 =
724 std::sync::atomic::AtomicU64::new(0);
725
726fn mla_tc_prefill_enabled() -> bool {
749 std::env::var("MEMRA_MLA_TC_PREFILL")
750 .map(|v| v != "0")
751 .unwrap_or(true)
752}
753
754pub(crate) enum HcTapArm<'a, 's: 'a> {
822 FromCache,
824 Shared(&'a std::sync::Mutex<&'s mut crate::cache::HcTapSink>, usize),
826}
827
828fn hyper_pipe_decline_once(reason: &str) {
833 static DECLINED: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false);
834 if !DECLINED.swap(true, std::sync::atomic::Ordering::Relaxed) {
835 eprintln!(
836 "[hyper-prime-pipe] DECLINED: {reason} -> the serial per-chunk stage walk serves \
837 this prime (MEMRA_B200_PRIME_V2 arm 1 is unaffected; logged once per process)"
838 );
839 }
840}
841
842pub fn b200_prime_v2_on() -> bool {
843 std::env::var("MEMRA_B200_PRIME_V2").as_deref() == Ok("1")
844}
845
846pub static HYPER_PRIME_NATURAL_SCHEDULES: std::sync::atomic::AtomicU64 =
853 std::sync::atomic::AtomicU64::new(0);
854
855pub static HYPER_PRIME_PIPELINED_CHUNKS: std::sync::atomic::AtomicU64 =
856 std::sync::atomic::AtomicU64::new(0);
857
858fn moe_grouped_enabled(_cfg: &ModelConfig, _prefill: bool) -> bool {
862 std::env::var("MEMRA_MOE_GROUPED")
863 .map(|value| value != "0")
864 .unwrap_or(false)
865}
866
867fn moe_grouped_prefill_enabled() -> bool {
888 std::env::var("MEMRA_MOE_GROUPED_PREFILL")
889 .map(|v| v != "0")
890 .unwrap_or(true)
891}
892
893fn moe_prefetch_enabled() -> bool {
896 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
897 *E.get_or_init(|| {
898 std::env::var("MEMRA_MOE_PREFETCH").as_deref() == Ok("1")
899 || crate::spill_pread::worker_enabled()
900 })
901}
902
903fn moe_page_prefetch_window() -> usize {
908 static W: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
909 *W.get_or_init(|| {
910 page_prefetch_window_from_values(
911 std::env::var("MEMRA_MOE_PAGE_PREFETCH").as_deref() == Ok("1"),
912 std::env::var("MEMRA_MOE_PAGE_PREFETCH_WINDOW")
913 .ok()
914 .as_deref(),
915 )
916 })
917}
918
919fn page_prefetch_window_from_values(enabled: bool, raw_window: Option<&str>) -> usize {
920 if !enabled {
921 return 0;
922 }
923 raw_window.and_then(|value| value.parse().ok()).unwrap_or(1)
924}
925
926fn page_prefetch_positions(position: usize, len: usize, window: usize) -> std::ops::Range<usize> {
930 if window == 0 || position >= len {
931 return len..len;
932 }
933 let (start, count) = if position == 0 {
934 (1, window)
935 } else {
936 (position.saturating_add(window), 1)
937 };
938 let start = start.min(len);
939 start..start.saturating_add(count).min(len)
940}
941
942fn grouped_worker_prefetch_position(order_len: usize, current: Option<usize>) -> Option<usize> {
945 let position = current.map_or(0, |position| position.saturating_add(1));
946 (position < order_len).then_some(position)
947}
948
949fn worker_prefetch_window() -> usize {
954 static WINDOW: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
955 *WINDOW.get_or_init(|| {
956 let automatic = crate::spill_pread::configured_depth().saturating_sub(1) / 3;
957 std::env::var("MEMRA_SPILL_WORKER_EXPERT_WINDOW")
958 .ok()
959 .and_then(|value| value.parse::<usize>().ok())
960 .unwrap_or(automatic.max(1))
961 })
962}
963
964fn worker_prefetch_positions(position: usize, len: usize, window: usize) -> std::ops::Range<usize> {
968 if window == 0 || position >= len {
969 return len..len;
970 }
971 let (start, count) = if position == 0 {
972 (0, window)
973 } else {
974 (position.saturating_add(window).saturating_sub(1), 1)
975 };
976 let start = start.min(len);
977 start..start.saturating_add(count).min(len)
978}
979
980fn moe_dev_enabled() -> bool {
985 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
986 *E.get_or_init(|| {
987 std::env::var("MEMRA_MOE_DEV")
988 .map(|v| v != "0")
989 .unwrap_or(true)
990 && !matches!(std::env::var("MEMRA_FUSED_ROUTER").as_deref(), Ok("0"))
991 })
992}
993
994enum VrowsSel<'a> {
1003 Host(&'a [u32], &'a [f32]),
1004 Dev(&'a CudaSlice<i32>, &'a CudaSlice<f32>),
1005}
1006
1007pub(crate) fn sigmoid_router_enabled() -> bool {
1008 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1009 *E.get_or_init(|| {
1010 std::env::var("MEMRA_SIG_ROUTER")
1011 .map(|v| v != "0")
1012 .unwrap_or(true)
1013 })
1014}
1015
1016fn moe_q8_enabled() -> bool {
1021 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1022 *E.get_or_init(|| {
1023 std::env::var("MEMRA_MOE_Q8")
1024 .map(|v| v != "0")
1025 .unwrap_or(true)
1026 })
1027}
1028
1029fn expert_dp4a_supported(qt: i32) -> bool {
1032 qt == crate::QT_Q4_0
1033 || qt == crate::QT_IQ3_S
1034 || qt == crate::QT_IQ4_XS
1035 || qt == crate::QT_Q3_K
1036 || qt == crate::QT_Q4_K
1037 || qt == crate::QT_Q6_K
1038}
1039
1040fn q8_expert_supported(qt: i32) -> bool {
1041 static KQ: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1047 let kq = *KQ.get_or_init(|| {
1048 std::env::var("MEMRA_MOE_Q8_KQ")
1049 .map(|v| v != "0")
1050 .unwrap_or(true)
1051 });
1052 let nvfp4_q8 = std::env::var("MEMRA_MOE_Q8_NVFP4")
1059 .map(|v| v != "0")
1060 .unwrap_or(true);
1061 qt == crate::QT_IQ3_S
1062 || qt == crate::QT_IQ4_XS
1063 || (nvfp4_q8 && qt == crate::QT_NVFP4)
1064 || (kq && (qt == crate::QT_Q3_K || qt == crate::QT_Q4_K || qt == crate::QT_Q6_K))
1065}
1066
1067fn q8_expert_supported_for_model(cfg: &ModelConfig, qt: i32) -> bool {
1071 let weight_only_nvfp4 = cfg.hy3.as_ref().is_some_and(|hy3| hy3.weight_only_nvfp4);
1072 q8_expert_supported(qt) && !(weight_only_nvfp4 && qt == crate::QT_NVFP4)
1073}
1074
1075fn moe_q8_enabled_for_model(cfg: &ModelConfig, m: &MoeWeights) -> bool {
1076 m.has_uniform_expert_layout()
1077 && moe_q8_enabled()
1078 && q8_expert_supported_for_model(cfg, m.gate_exps.qtype)
1079 && q8_expert_supported_for_model(cfg, m.up_exps.qtype)
1080 && q8_expert_supported_for_model(cfg, m.down_exps.qtype)
1081}
1082
1083#[cfg(test)]
1084mod w4a16_dispatch_tests {
1085 use super::q8_expert_supported_for_model;
1086 use memra_gguf::config::{HfConfig, ModelConfig};
1087
1088 #[test]
1089 fn hy3_w4a16_never_admits_q8_activations() {
1090 let hf = HfConfig::parse(
1091 r#"{"model_type":"hy_v3","num_hidden_layers":2,"hidden_size":8,
1092 "num_attention_heads":2,"intermediate_size":16,"vocab_size":32,
1093 "max_position_embeddings":32,
1094 "quantization_config":{"quant_method":"modelopt","quant_algo":"W4A16_NVFP4"}}"#,
1095 );
1096 let cfg = ModelConfig::from_hf(&hf);
1097 assert!(!q8_expert_supported_for_model(&cfg, crate::QT_NVFP4));
1098 assert!(q8_expert_supported_for_model(&cfg, crate::QT_IQ4_XS));
1099 }
1100}
1101
1102fn q8_expert_dec_supported(qt: i32) -> bool {
1105 qt == crate::QT_IQ3_S || qt == crate::QT_IQ4_XS || qt == crate::QT_Q4_0
1106}
1107
1108fn f16g_proj_ok(qt: i32, in_f: usize) -> bool {
1114 match qt {
1115 crate::QT_Q4_0 => in_f.is_multiple_of(32),
1116 crate::QT_IQ4_XS | crate::QT_IQ3_S | crate::QT_Q3_K | crate::QT_Q4_K | crate::QT_Q6_K => {
1117 in_f.is_multiple_of(256)
1118 }
1119 crate::QT_NVFP4 | crate::QT_NVFP4_V2 => in_f.is_multiple_of(64),
1124 _ => false,
1125 }
1126}
1127
1128fn moe_prewarm_enabled() -> bool {
1131 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1132 *E.get_or_init(|| {
1133 std::env::var("MEMRA_MOE_PREWARM")
1134 .map(|v| v != "0")
1135 .unwrap_or(true)
1136 })
1137}
1138
1139fn cpu_expert_profile_admit_enabled() -> bool {
1143 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1144 *E.get_or_init(|| std::env::var("MEMRA_CPU_EXPERT_FREEZE_PROFILE_ADMIT").as_deref() == Ok("1"))
1145}
1146
1147pub const PRIME_MIN_T: usize = 16;
1151
1152pub const CUDA_GRID_YZ_MAX: usize = 65_535;
1165
1166pub const PRIME_CHUNK_LAUNCH_CAP: usize = CUDA_GRID_YZ_MAX - (PRIME_MIN_T - 1);
1174
1175fn explicit_prime_chunk(parsed: usize, ring_on: bool) -> usize {
1182 if ring_on {
1183 if parsed == 0 {
1184 crate::cache::PRIME_CHUNK_MAX_TOKENS
1185 } else {
1186 parsed.min(crate::cache::PRIME_CHUNK_MAX_TOKENS)
1187 }
1188 } else if parsed == 0 {
1189 PRIME_CHUNK_LAUNCH_CAP
1190 } else {
1191 parsed.min(PRIME_CHUNK_LAUNCH_CAP)
1192 }
1193}
1194
1195fn step_gemm_prime_suffix_on() -> bool {
1225 std::env::var("MEMRA_STEP_GEMM_PRIME_SUFFIX").as_deref() != Ok("0")
1226}
1227
1228const MOE_DEV_MAX_T: usize = 16;
1237const PRIME_PIPE_MICROBATCHES: usize = 8;
1238const PRIME_PIPE_MIN_CHUNK: usize = 128;
1239const PRIME_PIPE_EDGE_MIN_CHUNK: usize = 64;
1240const PRIME_PIPE_LINEAR_WORK: usize = 8;
1241
1242fn prime_pp2_auto_geometry(n_layers: usize) -> bool {
1243 crate::pp::prime_pp_on()
1244 && !crate::pp::pp2_streams_off()
1245 && crate::pp::pp_cuts(n_layers).is_some_and(|cuts| cuts.len() == 3)
1246}
1247
1248fn prime_ppn_wave_auto_geometry(n_layers: usize) -> bool {
1249 crate::pp::prime_pp_on()
1250 && !crate::pp::pp2_streams_off()
1251 && crate::pp::pp_wave_on() == Ok(true)
1252 && crate::pp::pp_cuts(n_layers).is_some_and(|cuts| matches!(cuts.len(), 4 | 5))
1253}
1254
1255fn prime_pipeline_auto_geometry(n_layers: usize) -> bool {
1256 prime_pp2_auto_geometry(n_layers) || prime_ppn_wave_auto_geometry(n_layers)
1257}
1258
1259pub fn prime_chunk_tokens(t: usize, n_layers: usize) -> usize {
1266 if let Ok(value) = std::env::var("MEMRA_PRIME_CHUNK") {
1267 let parsed = value
1268 .parse::<usize>()
1269 .unwrap_or(crate::cache::PRIME_CHUNK_MAX_TOKENS);
1270 return explicit_prime_chunk(parsed, crate::cache::swa_ring_on());
1271 }
1272 let chunk = crate::cache::PRIME_CHUNK_MAX_TOKENS;
1273 if prime_pipeline_auto_geometry(n_layers) && t >= 2 * PRIME_PIPE_MIN_CHUNK {
1274 chunk.min(
1275 t.div_ceil(PRIME_PIPE_MICROBATCHES)
1276 .max(PRIME_PIPE_MIN_CHUNK),
1277 )
1278 } else {
1279 chunk
1280 }
1281}
1282
1283fn fixed_prime_chunk_ranges(t: usize, chunk: usize) -> Vec<(usize, usize)> {
1284 fixed_prime_chunk_ranges_for_ring(t, chunk, crate::cache::swa_ring_on())
1285}
1286
1287fn fixed_prime_chunk_ranges_for_ring(t: usize, chunk: usize, ring_on: bool) -> Vec<(usize, usize)> {
1288 if chunk == 0 || t <= chunk {
1289 return vec![(0, t)];
1290 }
1291 let mut ranges = Vec::with_capacity(t.div_ceil(chunk));
1292 let mut start = 0usize;
1293 while start < t {
1294 let mut end = (start + chunk).min(t);
1295 if t - end > 0 && t - end < PRIME_MIN_T {
1296 if ring_on {
1297 let shifted = t - PRIME_MIN_T;
1298 end = if shifted > start { shifted } else { t };
1299 } else {
1300 end = t;
1301 }
1302 }
1303 ranges.push((start, end));
1304 start = end;
1305 }
1306 ranges
1307}
1308
1309fn prime_chunk_work(prefix: usize, total: usize) -> u128 {
1310 let prefix = prefix as u128;
1311 prefix * (prefix + (PRIME_PIPE_LINEAR_WORK as u128) * (total as u128))
1312}
1313
1314fn dynamic_prime_chunk_ranges(
1315 t: usize,
1316 fixed_chunk: usize,
1317 fixed: &[(usize, usize)],
1318) -> Vec<(usize, usize)> {
1319 let n = fixed.len();
1320 if n < 3 {
1321 return fixed.to_vec();
1322 }
1323
1324 let max_first = t - (n - 1) * PRIME_MIN_T;
1325 let first = fixed_chunk
1326 .div_ceil(2)
1327 .max(PRIME_PIPE_EDGE_MIN_CHUNK)
1328 .min(max_first);
1329 let mut ranges = Vec::with_capacity(n);
1330 ranges.push((0, first));
1331
1332 let first_work = prime_chunk_work(first, t);
1333 let work_span = prime_chunk_work(t, t) - first_work;
1334 let denominator = (n - 1) as u128;
1335 let mut previous = first;
1336 for boundary in 1..n - 1 {
1337 let target = first_work * denominator + work_span * (boundary as u128);
1338 let remaining = n - 1 - boundary;
1339 let mut low = previous + PRIME_MIN_T;
1340 let mut high = t - remaining * PRIME_MIN_T;
1341 while low < high {
1342 let mid = low + (high - low) / 2;
1343 if prime_chunk_work(mid, t) * denominator >= target {
1344 high = mid;
1345 } else {
1346 low = mid + 1;
1347 }
1348 }
1349 ranges.push((previous, low));
1350 previous = low;
1351 }
1352 ranges.push((previous, t));
1353 ranges
1354}
1355
1356pub fn prime_chunk_ranges(t: usize, n_layers: usize, gdn_grid: bool) -> Vec<(usize, usize)> {
1366 let explicit_chunk = std::env::var_os("MEMRA_PRIME_CHUNK").is_some();
1367 let chunk = prime_chunk_tokens(t, n_layers);
1368 let fixed = fixed_prime_chunk_ranges(t, chunk);
1369 let dynamic = match std::env::var("MEMRA_PRIME_CHUNK_SCHED") {
1370 Ok(value) => value == "dynamic",
1371 Err(_) => true,
1372 };
1373 if explicit_chunk {
1374 return fixed;
1375 }
1376 let ranges = if !dynamic || !prime_pipeline_auto_geometry(n_layers) {
1377 fixed
1378 } else {
1379 dynamic_prime_chunk_ranges(t, chunk, &fixed)
1380 };
1381 if gdn_grid && std::env::var("MEMRA_PRIME_GRID_ALIGN").as_deref() != Ok("0") {
1385 align_prime_ranges_to_gdn(&ranges, t, Engine::gdn_chunk_size())
1386 } else {
1387 ranges
1388 }
1389}
1390
1391pub fn hyper_prime_ranges(t: usize, n_layers: usize, gdn_grid: bool) -> Vec<(usize, usize)> {
1447 if b200_prime_v2_on() && std::env::var_os("MEMRA_PRIME_CHUNK").is_none() {
1454 return hyper_prime_ranges_natural(t, gdn_grid);
1455 }
1456 prime_chunk_ranges(t, n_layers, gdn_grid)
1457}
1458
1459fn hyper_prime_ranges_natural(t: usize, gdn_grid: bool) -> Vec<(usize, usize)> {
1466 let ranges = fixed_prime_chunk_ranges(t, crate::cache::PRIME_CHUNK_MAX_TOKENS);
1467 let ranges = if gdn_grid && std::env::var("MEMRA_PRIME_GRID_ALIGN").as_deref() != Ok("0") {
1468 align_prime_ranges_to_gdn(&ranges, t, Engine::gdn_chunk_size())
1469 } else {
1470 ranges
1471 };
1472 HYPER_PRIME_NATURAL_SCHEDULES.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
1473 static SAID: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false);
1478 if !SAID.swap(true, std::sync::atomic::Ordering::Relaxed) {
1479 eprintln!(
1480 "[prime-v2] arm1 engaged t={t} chunk={} chunks={} widths={:?} (natural chunk \
1481 replaces the PRIME_PIPE_MICROBATCHES geometry; MEMRA_B200_PRIME_V2, logged once \
1482 per process)",
1483 crate::cache::PRIME_CHUNK_MAX_TOKENS,
1484 ranges.len(),
1485 ranges.iter().map(|&(a, b)| b - a).collect::<Vec<_>>(),
1486 );
1487 }
1488 ranges
1489}
1490
1491#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1508pub struct PrimeWorkspaceShape {
1509 pub call_row_bytes: usize,
1511 pub prompt_row_bytes: usize,
1513 pub n_layers: usize,
1515}
1516
1517impl PrimeWorkspaceShape {
1518 pub fn admission_bytes_with_call_rows(&self, prompt_rows: usize, call_rows: usize) -> usize {
1521 let rows = prompt_rows.min(call_rows.max(1));
1522 self.call_row_bytes
1523 .saturating_mul(rows)
1524 .saturating_add(self.prompt_row_bytes.saturating_mul(prompt_rows))
1525 }
1526
1527 pub fn admission_bytes(&self, prompt_rows: usize) -> usize {
1530 self.admission_bytes_with_call_rows(
1531 prompt_rows,
1532 prime_chunk_tokens(prompt_rows, self.n_layers),
1533 )
1534 }
1535}
1536
1537pub fn hyper_prime_call_rows(t: usize, n_layers: usize, gdn_grid: bool) -> usize {
1538 hyper_prime_ranges(t, n_layers, gdn_grid)
1539 .iter()
1540 .map(|&(start, end)| end - start)
1541 .max()
1542 .unwrap_or(0)
1543}
1544
1545#[derive(Debug, Clone, Copy)]
1557pub struct HyperPrimeWorkspaceShape {
1558 pub chunk_token_bytes: usize,
1568 pub prompt_bytes_per_token: usize,
1572 pub kpool_score_pool: usize,
1577 pub n_layers: usize,
1580 pub gdn_grid: bool,
1582}
1583
1584impl HyperPrimeWorkspaceShape {
1585 pub fn admission_bytes(&self, prompt_rows: usize) -> usize {
1596 let rows = hyper_prime_call_rows(prompt_rows, self.n_layers, self.gdn_grid);
1597 let chunk = self.chunk_token_bytes.saturating_mul(rows);
1598 let score = prompt_rows
1599 .checked_div(self.kpool_score_pool)
1600 .map(|pools| rows.saturating_mul(pools).saturating_mul(size_of::<f32>()))
1601 .unwrap_or(0);
1602 chunk
1603 .saturating_add(score)
1604 .saturating_add(self.prompt_bytes_per_token.saturating_mul(prompt_rows))
1605 }
1606}
1607
1608pub fn align_prime_ranges_to_gdn(
1628 ranges: &[(usize, usize)],
1629 t: usize,
1630 c: usize,
1631) -> Vec<(usize, usize)> {
1632 if c == 0 || ranges.len() < 2 {
1633 return ranges.to_vec();
1634 }
1635 let mut out: Vec<(usize, usize)> = Vec::with_capacity(ranges.len());
1636 let mut start = 0usize;
1637 for (i, &(_, end)) in ranges.iter().enumerate() {
1638 let e = if i + 1 == ranges.len() {
1639 t
1640 } else {
1641 end / c * c
1642 };
1643 if e > start {
1644 out.push((start, e));
1645 start = e;
1646 } }
1648 debug_assert_eq!(out.last().map(|&(_, e)| e), Some(t));
1649 out
1650}
1651
1652struct HeadSplit {
1653 pin: u64,
1654 w1: CudaSlice<u8>,
1655 hn1: CudaSlice<f32>,
1656 y1: CudaSlice<f32>,
1657 logits_e: CudaSlice<f32>,
1658 ev_hn: cudarc::driver::CudaEvent,
1659 ev_done: cudarc::driver::CudaEvent,
1660 raw_hn1: u64,
1661 raw_y1: u64,
1662 raw_logits_hi: u64,
1663 samp: Option<SampScratch>,
1668}
1669
1670struct SampScratch {
1671 pb: CudaSlice<f32>,
1672 th: CudaSlice<f32>,
1673 z: CudaSlice<f32>,
1674 mx: CudaSlice<f32>,
1675 rows: CudaSlice<i32>,
1676}
1677static HEAD_SPLIT_WS: std::sync::Mutex<Option<HeadSplit>> = std::sync::Mutex::new(None);
1679
1680#[allow(clippy::type_complexity)]
1684static DEV1_ROUTER_REPS: std::sync::Mutex<
1685 Option<(
1686 std::collections::HashMap<u16, (CudaSlice<f32>, CudaSlice<f32>, CudaSlice<u8>)>,
1687 Option<CudaSlice<f32>>,
1688 )>,
1689> = std::sync::Mutex::new(None);
1690
1691#[allow(clippy::type_complexity)]
1695static SHEXP_D1_REPS: std::sync::Mutex<
1696 Option<std::collections::HashMap<u16, (CudaSlice<u8>, CudaSlice<u8>, CudaSlice<u8>)>>,
1697> = std::sync::Mutex::new(None);
1698#[allow(clippy::type_complexity)]
1699static SHEXP_D1_WS: std::sync::Mutex<
1700 Option<(
1701 (usize, usize),
1702 CudaSlice<f32>,
1703 CudaSlice<f32>,
1704 CudaSlice<f32>,
1705 cudarc::driver::CudaEvent,
1706 cudarc::driver::CudaEvent,
1707 )>,
1708> = std::sync::Mutex::new(None);
1709
1710#[allow(clippy::type_complexity)] static SHEXP_OV_WS: std::sync::Mutex<
1713 Option<(usize, usize, usize, CudaSlice<f32>, CudaSlice<f32>)>,
1714> = std::sync::Mutex::new(None);
1715
1716impl HybridModel {
1717 pub fn gdn_prime_grid_on(&self) -> bool {
1723 Engine::gdn_chunked_enabled()
1724 && self
1725 .layers
1726 .iter()
1727 .any(|l| matches!(l.mixer, crate::hybrid::Mixer::Linear(_)))
1728 }
1729
1730 fn full_attn_tp_device_resident(e: &Engine, tp: &crate::hybrid::StepTpQkv) -> bool {
1735 tp.runtime.native_p2p() && tp.runtime.root_shares_ctx(e)
1736 }
1737
1738 pub(crate) fn full_attn_tp_qkv(
1739 &self,
1740 e: &Engine,
1741 fa: &FullAttnLayer,
1742 h: &CudaSlice<f32>,
1743 t: usize,
1744 ) -> Result<Option<Vec<CudaSlice<f32>>>, Box<dyn std::error::Error>> {
1745 let Some(tp) = fa.step_tp_qkv.as_ref() else {
1746 return Ok(None);
1747 };
1748 let values = active_matrix_values(
1749 h.len(),
1750 t,
1751 self.cfg.n_embd as usize,
1752 "Step TP QKV activation",
1753 )?;
1754 if Self::full_attn_tp_device_resident(e, tp) {
1763 e.stream().synchronize()?;
1766 let q = tp
1767 .runtime
1768 .bf16_column_parallel_resident_native_device(&tp.q, h, t)?;
1769 let k = tp
1770 .runtime
1771 .bf16_column_parallel_resident_native_device(&tp.k, h, t)?;
1772 let v = tp
1773 .runtime
1774 .bf16_column_parallel_resident_native_device(&tp.v, h, t)?;
1775 Self::full_attn_tp_log_once(tp, "qkv", "device-resident");
1776 return Ok(Some(vec![q, k, v]));
1777 }
1778 let host = e.dtoh_view(&h.slice(0..values))?;
1779 let q = if tp.runtime.native_p2p() {
1780 tp.runtime
1781 .bf16_column_parallel_resident_native(&tp.q, &host, t)?
1782 } else {
1783 tp.runtime
1784 .bf16_column_parallel_resident(&tp.q, &host, t)?
1785 .gathered
1786 };
1787 let k = if tp.runtime.native_p2p() {
1788 tp.runtime
1789 .bf16_column_parallel_resident_native(&tp.k, &host, t)?
1790 } else {
1791 tp.runtime
1792 .bf16_column_parallel_resident(&tp.k, &host, t)?
1793 .gathered
1794 };
1795 let v = if tp.runtime.native_p2p() {
1796 tp.runtime
1797 .bf16_column_parallel_resident_native(&tp.v, &host, t)?
1798 } else {
1799 tp.runtime
1800 .bf16_column_parallel_resident(&tp.v, &host, t)?
1801 .gathered
1802 };
1803 Self::full_attn_tp_log_once(tp, "qkv", "host-canonical");
1804 Ok(Some(vec![e.htod(&q)?, e.htod(&k)?, e.htod(&v)?]))
1805 }
1806
1807 fn full_attn_tp_log_once(tp: &crate::hybrid::StepTpQkv, proj: &str, activation: &'static str) {
1811 use std::sync::atomic::{AtomicBool, Ordering};
1812 static LOGGED: [AtomicBool; 4] = [
1813 AtomicBool::new(false),
1814 AtomicBool::new(false),
1815 AtomicBool::new(false),
1816 AtomicBool::new(false),
1817 ];
1818 let idx = 2 * usize::from(proj == "o") + usize::from(activation == "device-resident");
1819 if LOGGED[idx].swap(true, Ordering::Relaxed) {
1820 return;
1821 }
1822 eprintln!(
1823 "[step-tp-{proj}] execute layer={} devices={:?} projections={proj} \
1824 tensor_parallel=true attention_local=true kv_local=true transport={} \
1825 native_p2p={} bulk_p2p={} activation={activation} \
1826 output={} performance_claim=false (logged once per transport)",
1827 tp.layer,
1828 tp.devices,
1829 tp.runtime.transport_label(),
1830 tp.runtime.native_p2p(),
1831 tp.runtime.bulk_p2p(),
1832 if activation == "device-resident" {
1833 "root-resident"
1834 } else {
1835 "root-readback"
1836 },
1837 );
1838 }
1839
1840 pub(crate) fn full_attn_tp_o(
1841 &self,
1842 e: &Engine,
1843 fa: &FullAttnLayer,
1844 activation: &CudaSlice<f32>,
1845 tokens: usize,
1846 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
1847 let Some(tp) = fa.step_tp_qkv.as_ref() else {
1848 return Ok(None);
1849 };
1850 if Self::full_attn_tp_device_resident(e, tp) {
1854 e.stream().synchronize()?; let output = tp
1856 .runtime
1857 .step_bf16_row_parallel_resident_native_device(&tp.o, activation, tokens)?;
1858 Self::full_attn_tp_log_once(tp, "o", "device-resident");
1859 return Ok(Some(output));
1860 }
1861 let host = e.dtoh(activation)?;
1862 let output = if tp.runtime.native_p2p() {
1863 tp.runtime
1864 .step_bf16_row_parallel_resident_native(&tp.o, &host, tokens)?
1865 } else {
1866 tp.runtime
1867 .step_bf16_row_parallel_resident(&tp.o, &host, tokens)?
1868 };
1869 Self::full_attn_tp_log_once(tp, "o", "host-canonical");
1870 Ok(Some(e.htod(&output)?))
1871 }
1872
1873 fn full_attn_o(
1874 &self,
1875 e: &Engine,
1876 fa: &FullAttnLayer,
1877 activation: &CudaSlice<f32>,
1878 tokens: usize,
1879 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1880 match self.full_attn_tp_o(e, fa, activation, tokens)? {
1881 Some(output) => Ok(output),
1882 None => e.matmul(&fa.wo, activation, tokens),
1883 }
1884 }
1885
1886 fn prime_trace_path() -> Option<&'static str> {
1891 static P: std::sync::OnceLock<Option<String>> = std::sync::OnceLock::new();
1892 P.get_or_init(|| std::env::var("MEMRA_PRIME_TRACE").ok())
1893 .as_deref()
1894 }
1895
1896 fn prime_anatomy_on() -> bool {
1902 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1903 *E.get_or_init(|| std::env::var("MEMRA_PRIME_ANATOMY").as_deref() == Ok("1"))
1904 }
1905
1906 fn prime_anatomy_slots() -> &'static [std::sync::atomic::AtomicU64; 5] {
1907 static S: [std::sync::atomic::AtomicU64; 5] = [
1908 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), ];
1914 &S
1915 }
1916
1917 pub(crate) fn refuse_hyper(&self, path: &str) -> Result<(), Box<dyn std::error::Error>> {
1924 if let Some(topology) = self.hyper.as_ref() {
1925 return Err(format!(
1926 "{path} runs a serial residual, but this model's ModelPlan declares \
1927 ResidualTopology::HyperConnections{{ streams: {}, collapse: {:?} }}. Refusing: \
1928 that path would compute a different model. Converted paths: forward, \
1929 forward_last, prime_cache, decode_step, and the batched serving chain \
1930 decode_step_batch / _sampled / _lean / _masked.",
1931 topology.streams, topology.collapse
1932 )
1933 .into());
1934 }
1935 Ok(())
1936 }
1937
1938 #[allow(clippy::too_many_arguments)] pub(crate) fn hyper_ffn_branch(
1949 &self,
1950 e: &Engine,
1951 layer: &crate::hybrid::HybridLayer,
1952 z: &CudaSlice<f32>,
1953 t: usize,
1954 il: usize,
1955 prefill: bool,
1956 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
1957 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1958 match &layer.ffn {
1959 crate::hybrid::Ffn::Dense {
1960 ffn_gate,
1961 ffn_up,
1962 ffn_down,
1963 } => {
1964 let n_ff = ffn_gate.out_features();
1965 let mut g2 = e.matmul_group(&[ffn_gate, ffn_up], z, t)?;
1966 let up = g2.pop().unwrap();
1967 let gate = g2.pop().unwrap();
1968 let mut act = e.uninit(t * n_ff)?;
1969 Self::ffn_act_lim(
1971 e,
1972 &self.cfg,
1973 &gate,
1974 &up,
1975 1.0,
1976 1.0,
1977 self.cfg.clamp_shexp_at(il as u32),
1978 &mut act,
1979 t * n_ff,
1980 )?;
1981 e.matmul(ffn_down, &act, t)
1982 }
1983 crate::hybrid::Ffn::Moe(m) => {
1984 if prefill {
1985 self.moe_ffn_il_prefill(e, m, z, t, il as u16)
1986 } else {
1987 self.moe_ffn_il_zq8(e, m, z, zq8, t, il as u16)
1988 }
1989 }
1990 }
1991 }
1992
1993 fn forward_hyper(
1998 &self,
1999 e: &Engine,
2000 tokens: &[u32],
2001 last_only: bool,
2002 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
2003 let topology = *self
2004 .hyper
2005 .as_ref()
2006 .ok_or("forward_hyper on a model with no HyperConnections topology")?;
2007 if let Some(fence) = crate::pp::pp_cuts(self.layers.len()) {
2013 if !self.rewrite_allowed(memra_gguf::execution_manifest::RewriteSurface::Pipeline) {
2014 return Err("pipeline rewrite is not qualified for this ModelPlan".into());
2015 }
2016 return self.forward_hyper_ppn(e, tokens, last_only, &topology, &fence);
2017 }
2018 let n_embd = self.cfg.n_embd as usize;
2019 let t = tokens.len();
2020 let eps = self.cfg.rms_eps;
2021 let pos: Vec<i32> = (0..t as i32).collect();
2022 let pos_d = e.htod_i32(&pos)?;
2023
2024 let embedded = self.embed(e, tokens)?;
2025 let mut x = crate::hyper::expand(e, &topology, &embedded, t, n_embd)?;
2026 let trace = memra_reference::hidden_trace::enabled();
2027 if trace {
2028 memra_reference::hidden_trace::emit_tokens(tokens);
2029 let streams = x.len() / (t * n_embd);
2030 memra_reference::hidden_trace::emit_last_row(
2031 "expand",
2032 -1,
2033 t,
2034 streams * n_embd,
2035 &e.dtoh(&x)?,
2036 );
2037 }
2038
2039 x = self.hyper_range_forward(e, &topology, x, 0, self.layers.len(), &pos_d, t, trace)?;
2040
2041 self.hyper_head_logits(e, &topology, &x, t, n_embd, eps, last_only)
2045 }
2046
2047 #[allow(clippy::type_complexity)] fn prime_cache_hyper(
2059 &self,
2060 e: &Engine,
2061 tokens: &[u32],
2062 cache: &mut Cache,
2063 queued_after: usize,
2064 overlay: Option<&crate::vision::EmbedOverlay>,
2065 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
2066 let topology = *self
2067 .hyper
2068 .as_ref()
2069 .ok_or("prime_cache_hyper on a model with no HyperConnections topology")?;
2070 if let Some(fence) = crate::pp::pp_cuts(self.layers.len()) {
2093 if !self.rewrite_allowed(memra_gguf::execution_manifest::RewriteSurface::Pipeline) {
2094 return Err("pipeline rewrite is not qualified for this ModelPlan".into());
2095 }
2096 let n_embd = self.cfg.n_embd as usize;
2105 let t = tokens.len();
2106 if cache.pos + t > cache.max_ctx {
2107 return Err("prime_cache: prompt exceeds cache max_ctx".into());
2108 }
2109 let ranges = hyper_prime_ranges(t, self.layers.len(), self.gdn_prime_grid_on());
2110 if ranges.len() == 1 {
2115 self.glm5_taps_range_begin(cache, 0);
2116 let out = self.prime_cache_hyper_ppn(
2117 e,
2118 tokens,
2119 cache,
2120 queued_after,
2121 &topology,
2122 &fence,
2123 overlay,
2124 )?;
2125 self.glm5_taps_range_done(e, cache, 0, t)?;
2126 return Ok(out);
2127 }
2128 if b200_prime_v2_on() && !crate::pp::pp2_streams_off() && fence.len() == 3 {
2133 if !crate::pp::prime_pipe_on() {
2134 hyper_pipe_decline_once("MEMRA_PRIME_PIPE=0 (the operator rollback seam)");
2135 } else if overlay.is_some() {
2136 hyper_pipe_decline_once(
2137 "a vision embedding overlay is present: the splice is a stage-0 \
2138 embedding-intake transform whose gate ran on the serial ppN body",
2139 );
2140 } else {
2141 let seq_end = cache.pos + t + queued_after;
2149 let mut sink = cache.hc_taps.take();
2150 let out = {
2151 let lock = sink.as_mut().map(std::sync::Mutex::new);
2152 self.prime_cache_hyper_pp2_pipelined(
2153 e,
2154 tokens,
2155 cache,
2156 seq_end,
2157 &topology,
2158 &ranges,
2159 &fence,
2160 lock.as_ref(),
2161 )
2162 };
2163 if let Some(mut s) = sink {
2167 if let Some(&(last, _)) = ranges.last() {
2170 s.base = cache.pos.saturating_sub(t).saturating_add(last);
2171 }
2172 cache.hc_taps = Some(s);
2173 }
2174 return out;
2175 }
2176 }
2177 let mut hiddens = e.uninit(t * n_embd)?;
2178 let mut last: Option<(Vec<f32>, CudaSlice<f32>)> = None;
2179 for &(start, end) in &ranges {
2180 let ov = overlay.and_then(|o| o.window(start, end - start));
2181 self.glm5_taps_range_begin(cache, start);
2182 let (l, hs, x) = self.prime_cache_hyper_ppn(
2183 e,
2184 &tokens[start..end],
2185 cache,
2186 queued_after + (t - end),
2187 &topology,
2188 &fence,
2189 ov.as_ref(),
2190 )?;
2191 self.glm5_taps_range_done(e, cache, start, end)?;
2192 e.copy_into(&mut hiddens, start * n_embd, &x, (end - start) * n_embd)?;
2193 last = Some((l, hs));
2194 crate::progress::note_prime_rows(end - start);
2198 }
2199 let (logits, h_seed) =
2200 last.expect("hyper_prime_ranges never returns an empty schedule");
2201 return Ok((logits, h_seed, hiddens));
2202 }
2203 let n_embd = self.cfg.n_embd as usize;
2204 let t = tokens.len();
2205 if cache.pos + t > cache.max_ctx {
2206 return Err("prime_cache: prompt exceeds cache max_ctx".into());
2207 }
2208 let seq_end = cache.pos + t + queued_after;
2212 let ranges = hyper_prime_ranges(t, self.layers.len(), self.gdn_prime_grid_on());
2213 if ranges.len() == 1 {
2215 self.glm5_taps_range_begin(cache, 0);
2216 let out = self.prime_chunk_hyper(e, tokens, cache, seq_end, 0, overlay)?;
2217 self.glm5_taps_range_done(e, cache, 0, t)?;
2218 return Ok(out);
2219 }
2220 let mut hiddens = e.uninit(t * n_embd)?;
2221 let mut last: Option<(Vec<f32>, CudaSlice<f32>)> = None;
2222 for &(start, end) in &ranges {
2223 self.glm5_taps_range_begin(cache, start);
2224 let (l, hs, x) =
2225 self.prime_chunk_hyper(e, &tokens[start..end], cache, seq_end, start, overlay)?;
2226 self.glm5_taps_range_done(e, cache, start, end)?;
2227 e.copy_into(&mut hiddens, start * n_embd, &x, (end - start) * n_embd)?;
2228 last = Some((l, hs));
2229 crate::progress::note_prime_rows(end - start);
2233 }
2234 let (logits, h_seed) = last.expect("hyper_prime_ranges never returns an empty schedule");
2235 Ok((logits, h_seed, hiddens))
2236 }
2237
2238 pub fn prime_workspace_shape(&self) -> Option<PrimeWorkspaceShape> {
2266 if self.hyper.is_some() {
2267 return None;
2268 }
2269 let h = self.cfg.n_embd as usize;
2270 let n_ff_max = self
2271 .layers
2272 .iter()
2273 .map(|l| match &l.ffn {
2274 crate::hybrid::Ffn::Dense { ffn_gate, .. } => ffn_gate.out_features(),
2275 _ => h,
2276 })
2277 .max()
2278 .unwrap_or(h)
2279 .max(h);
2280 let f32b = std::mem::size_of::<f32>();
2281 let mut call_row_bytes =
2284 (7 + 2) * h * f32b + (2 + 1) * h * 2 + 3 * n_ff_max * f32b + n_ff_max * 2;
2285 if let Some(moe) = self.cfg.moe.as_ref() {
2292 let u = moe.expert_used_count as usize;
2293 let f = moe.expert_ff_length as usize;
2294 call_row_bytes += u * (10 * h + 14 * f);
2295 }
2296 Some(PrimeWorkspaceShape {
2297 call_row_bytes,
2298 prompt_row_bytes: h * f32b,
2299 n_layers: self.layers.len(),
2300 })
2301 }
2302
2303 pub fn hyper_prime_workspace_shape(&self) -> Option<HyperPrimeWorkspaceShape> {
2304 let topology = self.hyper.as_ref()?;
2305 let h = self.cfg.n_embd as usize;
2306 let s = topology.streams;
2307 let f32b = std::mem::size_of::<f32>();
2308 let mut chunk_token_bytes = 4 * s * h * f32b + 4 * h * f32b + 2 * h * f32b;
2310 if let Some(moe) = self.cfg.moe.as_ref() {
2311 let u = moe.expert_used_count as usize;
2312 let f = moe.expert_ff_length as usize;
2313 chunk_token_bytes += u * (10 * h + 14 * f) + 4 * h;
2314 }
2315 let mut kpool_score_pool = 0;
2316 if let Some(glm5) = self.cfg.glm5.as_ref() {
2317 let heads = self.cfg.n_head as usize;
2318 chunk_token_bytes +=
2319 heads * (glm5.qk_head_dim as usize + glm5.v_head_dim as usize) * f32b;
2320 if glm5.index_kpool > 0 {
2321 chunk_token_bytes += (glm5.index_topk as usize / glm5.index_kpool as usize + 1)
2322 * std::mem::size_of::<i32>();
2323 kpool_score_pool = glm5.index_kpool as usize;
2324 }
2325 }
2326 Some(HyperPrimeWorkspaceShape {
2327 chunk_token_bytes,
2328 prompt_bytes_per_token: h * f32b,
2329 kpool_score_pool,
2330 n_layers: self.layers.len(),
2331 gdn_grid: self.gdn_prime_grid_on(),
2332 })
2333 }
2334
2335 pub(crate) fn decode_step_hyper(
2340 &self,
2341 e: &Engine,
2342 token: u32,
2343 cache: &mut Cache,
2344 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
2345 let topology = *self
2346 .hyper
2347 .as_ref()
2348 .ok_or("decode_step_hyper on a model with no HyperConnections topology")?;
2349 if let Some(fence) = crate::pp::pp_cuts(self.layers.len()) {
2353 if !self.rewrite_allowed(memra_gguf::execution_manifest::RewriteSurface::Pipeline) {
2354 return Err("pipeline rewrite is not qualified for this ModelPlan".into());
2355 }
2356 return self.decode_step_hyper_ppn(e, token, cache, &topology, &fence);
2357 }
2358 let n_embd = self.cfg.n_embd as usize;
2359 let eps = self.cfg.rms_eps;
2360 let pos = cache.pos;
2361 let pos_d = e.htod_i32(&[pos as i32])?;
2362
2363 let embedded = e.htod(&self.embd.try_gather(n_embd, &[token])?)?;
2364 let mut x = crate::hyper::expand(e, &topology, &embedded, 1, n_embd)?;
2365
2366 x = self.hyper_range_decode(e, &topology, x, 0, self.layers.len(), &pos_d, pos, cache)?;
2367
2368 self.hyper_decode_tail(e, &topology, &x, n_embd, eps, cache)
2370 }
2371
2372 #[allow(clippy::too_many_arguments)]
2379 fn hyper_range_forward(
2380 &self,
2381 e: &Engine,
2382 topology: &crate::hyper::HyperTopology,
2383 mut x: CudaSlice<f32>,
2384 lo: usize,
2385 hi: usize,
2386 pos_d: &CudaSlice<i32>,
2387 t: usize,
2388 trace: bool,
2389 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2390 let n_embd = self.cfg.n_embd as usize;
2391 let eps = self.cfg.rms_eps;
2392 for il in lo..hi {
2393 let layer = &self.layers[il];
2394 let hyper = layer.hyper.as_ref().ok_or_else(|| {
2395 format!("layer {il} carries no hyper-connection weights under an hc plan")
2396 })?;
2397
2398 let (y, mix) = crate::hyper::pre(e, topology, &hyper.attn, &x, t, n_embd)?;
2399 let mut h = e.uninit(t * n_embd)?;
2400 e.rms_norm(&y, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
2401 let mixed = match &layer.mixer {
2402 Mixer::Full(fa) => self.full_attn(e, fa, &h, pos_d, t, il)?,
2403 Mixer::Linear(la) => self.linear_attn(e, la, &h, t)?,
2404 Mixer::Mla(mla) => self.mla_attn(e, mla, &h, pos_d, t, il)?,
2405 Mixer::Kda(la) => crate::kda::kda_attn(e, la, &h, t, eps)?,
2406 };
2407 x = crate::hyper::post(e, topology, &mixed, &x, &mix, t, n_embd)?;
2408 if trace {
2409 let index = il as i64;
2410 memra_reference::hidden_trace::emit_last_row(
2411 "mixer",
2412 index,
2413 t,
2414 n_embd,
2415 &e.dtoh(&mixed)?,
2416 );
2417 let streams = x.len() / (t * n_embd);
2418 memra_reference::hidden_trace::emit_last_row(
2419 "attn",
2420 index,
2421 t,
2422 streams * n_embd,
2423 &e.dtoh(&x)?,
2424 );
2425 }
2426
2427 let (y, mix) = crate::hyper::pre(e, topology, &hyper.mlp, &x, t, n_embd)?;
2428 let mut z = e.uninit(t * n_embd)?;
2429 e.rms_norm(
2430 &y,
2431 layer.post_attn_norm.float_data(),
2432 &mut z,
2433 n_embd,
2434 t,
2435 eps,
2436 )?;
2437 let ffn_out = self.hyper_ffn_branch(e, layer, &z, t, il, true, None)?;
2438 x = crate::hyper::post(e, topology, &ffn_out, &x, &mix, t, n_embd)?;
2439 if trace {
2440 let index = il as i64;
2441 memra_reference::hidden_trace::emit_last_row(
2442 "ffn",
2443 index,
2444 t,
2445 n_embd,
2446 &e.dtoh(&ffn_out)?,
2447 );
2448 let streams = x.len() / (t * n_embd);
2449 memra_reference::hidden_trace::emit_last_row(
2450 "layer",
2451 index,
2452 t,
2453 streams * n_embd,
2454 &e.dtoh(&x)?,
2455 );
2456 }
2457 }
2458 Ok(x)
2459 }
2460
2461 #[allow(clippy::too_many_arguments)]
2466 #[allow(clippy::too_many_arguments)] fn hyper_range_prime(
2468 &self,
2469 e: &Engine,
2470 topology: &crate::hyper::HyperTopology,
2471 mut x: CudaSlice<f32>,
2472 lo: usize,
2473 hi: usize,
2474 pos_d: &CudaSlice<i32>,
2475 t: usize,
2476 cache: &mut Cache,
2477 seq_end: usize,
2478 taps: HcTapArm<'_, '_>,
2479 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2480 let n_embd = self.cfg.n_embd as usize;
2481 let eps = self.cfg.rms_eps;
2482 let prof = std::env::var("MEMRA_PRIME_PROF").as_deref() == Ok("1");
2494 let mut ph = [0f64; 4]; let mut pt = std::time::Instant::now();
2496 let mark = |e: &Engine, acc: usize, pt: &mut std::time::Instant, ph: &mut [f64; 4]| {
2497 if prof {
2498 let _ = e.stream().synchronize();
2499 ph[acc] += pt.elapsed().as_secs_f64() * 1e3;
2500 *pt = std::time::Instant::now();
2501 }
2502 };
2503 for il in lo..hi {
2504 let layer = &self.layers[il];
2505 let hyper = layer.hyper.as_ref().ok_or_else(|| {
2506 format!("layer {il} carries no hyper-connection weights under an hc plan")
2507 })?;
2508
2509 let (y, mix) = crate::hyper::pre(e, topology, &hyper.attn, &x, t, n_embd)?;
2510 let mut h = e.uninit(t * n_embd)?;
2511 e.rms_norm(&y, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
2512 mark(e, 0, &mut pt, &mut ph);
2513 let mixed = match &layer.mixer {
2514 Mixer::Full(fa) => {
2515 self.full_attn_prime(e, fa, &h, None, pos_d, t, cache, il, seq_end)?
2516 }
2517 Mixer::Linear(la) => self.linear_attn_prime(e, la, &h, None, t, cache, il)?,
2518 Mixer::Mla(mla) if mla.tp.is_some() => {
2519 self.mla_tp_attn_cached(e, mla, &h, pos_d, t, il, cache, false)?
2520 }
2521 Mixer::Mla(mla) => self.mla_attn_cached(e, mla, &h, pos_d, t, il, cache)?,
2522 Mixer::Kda(la) if la.tp.is_some() => crate::glm5_tp::kda_tp_cached(
2523 e,
2524 la,
2525 &h,
2526 t,
2527 eps,
2528 cache,
2529 il,
2530 crate::kda::ConvArm::Prefill,
2531 )?,
2532 Mixer::Kda(la) => crate::kda::kda_prime_cached(e, la, &h, t, eps, cache, il)?,
2533 };
2534 mark(e, 1, &mut pt, &mut ph);
2535 x = crate::hyper::post(e, topology, &mixed, &x, &mix, t, n_embd)?;
2536
2537 let (y, mix) = crate::hyper::pre(e, topology, &hyper.mlp, &x, t, n_embd)?;
2538 let mut z = e.uninit(t * n_embd)?;
2539 e.rms_norm(
2540 &y,
2541 layer.post_attn_norm.float_data(),
2542 &mut z,
2543 n_embd,
2544 t,
2545 eps,
2546 )?;
2547 mark(e, 2, &mut pt, &mut ph);
2548 let ffn_out = self.hyper_ffn_branch(e, layer, &z, t, il, true, None)?;
2549 x = crate::hyper::post(e, topology, &ffn_out, &x, &mix, t, n_embd)?;
2550 match &taps {
2553 HcTapArm::FromCache => self.glm5_hc_tap(e, cache, topology, il, &x, t)?,
2554 HcTapArm::Shared(sink, base) => {
2555 let mut guard = sink
2558 .lock()
2559 .map_err(|_| "hc tap sink lock poisoned by a failed stage walk")?;
2560 self.glm5_hc_tap_into(e, &mut guard, *base, topology, il, &x, t)?;
2561 }
2562 }
2563 mark(e, 3, &mut pt, &mut ph);
2564 }
2565 if prof {
2566 eprintln!(
2570 "[prime-prof] walk=hyper t={t} layers={} attn_glue={:.0}ms mixer={:.0}ms \
2571 ffn_glue={:.0}ms ffn={:.0}ms",
2572 hi - lo,
2573 ph[0],
2574 ph[1],
2575 ph[2],
2576 ph[3]
2577 );
2578 }
2579 Ok(x)
2580 }
2581
2582 #[allow(clippy::too_many_arguments)]
2584 fn hyper_range_decode(
2585 &self,
2586 e: &Engine,
2587 topology: &crate::hyper::HyperTopology,
2588 x: CudaSlice<f32>,
2589 lo: usize,
2590 hi: usize,
2591 pos_d: &CudaSlice<i32>,
2592 pos: usize,
2593 cache: &mut Cache,
2594 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2595 if crate::glm5_decode_graph_on() && self.glm5_decode_graph_ready(e, cache, lo, hi) {
2602 return self.hyper_range_decode_graphed(e, topology, x, lo, hi, pos_d, pos, cache);
2603 }
2604 if crate::glm5_graph_trace_on() {
2608 return self.hyper_range_decode_eager_traced(e, topology, x, lo, hi, pos_d, pos, cache);
2609 }
2610 self.hyper_range_decode_eager(e, topology, x, lo, hi, pos_d, pos, cache)
2611 }
2612
2613 #[allow(clippy::too_many_arguments)]
2616 pub(crate) fn hyper_range_decode_eager(
2617 &self,
2618 e: &Engine,
2619 topology: &crate::hyper::HyperTopology,
2620 mut x: CudaSlice<f32>,
2621 lo: usize,
2622 hi: usize,
2623 pos_d: &CudaSlice<i32>,
2624 pos: usize,
2625 cache: &mut Cache,
2626 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2627 if hyper_decode_ws_on() {
2633 return self.hyper_range_decode_ws(e, topology, x, lo, hi, pos_d, pos, cache);
2634 }
2635 let n_embd = self.cfg.n_embd as usize;
2636 let eps = self.cfg.rms_eps;
2637 for il in lo..hi {
2638 let layer = &self.layers[il];
2639 let hyper = layer.hyper.as_ref().ok_or_else(|| {
2640 format!("layer {il} carries no hyper-connection weights under an hc plan")
2641 })?;
2642
2643 let (y, mix) = crate::hyper::pre(e, topology, &hyper.attn, &x, 1, n_embd)?;
2644 let mut h = e.uninit(n_embd)?;
2645 let attn_q8 = if crate::glm5_q8_fuse_attn_on()
2651 && matches!(&layer.mixer, Mixer::Kda(la) if la.tp.is_none())
2652 {
2653 let pair =
2654 e.rms_norm_zq8_f32(&y, layer.attn_norm.float_data(), &mut h, n_embd, 1, eps)?;
2655 if crate::GLM5_Q8_FUSE_ATTN_DISPATCHES
2656 .fetch_add(1, std::sync::atomic::Ordering::Relaxed)
2657 == 0
2658 {
2659 eprintln!(
2660 "[glm5-q8-fuse-attn] engaged (rms_norm_zq8_f32 -> kda6, hyper_range_decode)"
2661 );
2662 }
2663 Some(pair)
2664 } else {
2665 e.rms_norm(&y, layer.attn_norm.float_data(), &mut h, n_embd, 1, eps)?;
2666 None
2667 };
2668 let mixed = match &layer.mixer {
2669 Mixer::Full(fa) => self.full_attn_decode(e, fa, &h, pos_d, pos, cache, il)?,
2670 Mixer::Linear(la) => self.linear_attn_decode(e, la, &h, cache, il)?,
2671 Mixer::Mla(mla) if mla.tp.is_some() => {
2672 self.mla_tp_attn_cached(e, mla, &h, pos_d, 1, il, cache, false)?
2673 }
2674 Mixer::Mla(mla) => self.mla_attn_cached(e, mla, &h, pos_d, 1, il, cache)?,
2675 Mixer::Kda(la) if la.tp.is_some() => crate::glm5_tp::kda_tp_cached(
2676 e,
2677 la,
2678 &h,
2679 1,
2680 eps,
2681 cache,
2682 il,
2683 crate::kda::ConvArm::Decode,
2684 )?,
2685 Mixer::Kda(la) => crate::kda::kda_decode_cached_q8(
2686 e,
2687 la,
2688 &h,
2689 attn_q8.as_ref().map(|(q, d)| (q, d)),
2690 eps,
2691 cache,
2692 il,
2693 )?,
2694 };
2695 x = crate::hyper::post(e, topology, &mixed, &x, &mix, 1, n_embd)?;
2696
2697 let (y, mix) = crate::hyper::pre(e, topology, &hyper.mlp, &x, 1, n_embd)?;
2698 let mut z = e.uninit(n_embd)?;
2699 let zq8 = if crate::glm5_q8_fuse_on() {
2707 let pair = e.rms_norm_zq8_f32(
2708 &y,
2709 layer.post_attn_norm.float_data(),
2710 &mut z,
2711 n_embd,
2712 1,
2713 eps,
2714 )?;
2715 if GLM5_Q8_FUSE_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed) == 0 {
2716 eprintln!("[glm5-q8-fuse] engaged (rms_norm_zq8_f32, hyper_range_decode)");
2717 }
2718 Some(pair)
2719 } else {
2720 e.rms_norm(
2721 &y,
2722 layer.post_attn_norm.float_data(),
2723 &mut z,
2724 n_embd,
2725 1,
2726 eps,
2727 )?;
2728 None
2729 };
2730 let ffn_out = self.hyper_ffn_branch(e, layer, &z, 1, il, false, zq8.as_ref())?;
2731 x = crate::hyper::post(e, topology, &ffn_out, &x, &mix, 1, n_embd)?;
2732 }
2733 Ok(x)
2734 }
2735
2736 #[allow(clippy::too_many_arguments)]
2743 fn hyper_range_decode_ws(
2744 &self,
2745 e: &Engine,
2746 topology: &crate::hyper::HyperTopology,
2747 x: CudaSlice<f32>,
2748 lo: usize,
2749 hi: usize,
2750 pos_d: &CudaSlice<i32>,
2751 pos: usize,
2752 cache: &mut Cache,
2753 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2754 let n_embd = self.cfg.n_embd as usize;
2755 let mut ws = match e.hyper_ws_take() {
2756 Some(ws) if ws.matches(topology, n_embd) => ws,
2757 _ => crate::hyper::HyperDecodeWs::new(e, topology, n_embd)?,
2758 };
2759 if HC_DECODE_WS_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed) == 0 {
2760 eprintln!(
2761 "[hc-decode-ws] engaged streams={} hidden={n_embd} (persistent hc-glue \
2762 workspace, per-engine pool; MEMRA_HC_DECODE_WS=1)",
2763 topology.streams
2764 );
2765 }
2766 let mut x = x;
2767 let out = self
2768 .hyper_range_decode_ws_body(e, topology, &mut x, lo, hi, pos_d, pos, cache, &mut ws)
2769 .map(|()| x);
2770 e.hyper_ws_put(ws);
2771 out
2772 }
2773
2774 #[allow(clippy::too_many_arguments)]
2779 pub(crate) fn hyper_range_decode_ws_body(
2780 &self,
2781 e: &Engine,
2782 topology: &crate::hyper::HyperTopology,
2783 x: &mut CudaSlice<f32>,
2784 lo: usize,
2785 hi: usize,
2786 pos_d: &CudaSlice<i32>,
2787 pos: usize,
2788 cache: &mut Cache,
2789 ws: &mut crate::hyper::HyperDecodeWs,
2790 ) -> Result<(), Box<dyn std::error::Error>> {
2791 let n_embd = self.cfg.n_embd as usize;
2792 let eps = self.cfg.rms_eps;
2793 for il in lo..hi {
2794 let layer = &self.layers[il];
2795 let hyper = layer.hyper.as_ref().ok_or_else(|| {
2796 format!("layer {il} carries no hyper-connection weights under an hc plan")
2797 })?;
2798
2799 crate::hyper::pre_t1_ws(e, topology, &hyper.attn, x, ws, n_embd)?;
2800 let attn_q8 = if crate::glm5_q8_fuse_attn_on()
2802 && matches!(&layer.mixer, Mixer::Kda(la) if la.tp.is_none())
2803 {
2804 let pair = e.rms_norm_zq8_f32(
2805 &ws.y,
2806 layer.attn_norm.float_data(),
2807 &mut ws.h,
2808 n_embd,
2809 1,
2810 eps,
2811 )?;
2812 if crate::GLM5_Q8_FUSE_ATTN_DISPATCHES
2813 .fetch_add(1, std::sync::atomic::Ordering::Relaxed)
2814 == 0
2815 {
2816 eprintln!(
2817 "[glm5-q8-fuse-attn] engaged (rms_norm_zq8_f32 -> kda6, hyper_range_decode_ws_body)"
2818 );
2819 }
2820 Some(pair)
2821 } else {
2822 e.rms_norm(
2823 &ws.y,
2824 layer.attn_norm.float_data(),
2825 &mut ws.h,
2826 n_embd,
2827 1,
2828 eps,
2829 )?;
2830 None
2831 };
2832 let mixed = match &layer.mixer {
2833 Mixer::Full(fa) => self.full_attn_decode(e, fa, &ws.h, pos_d, pos, cache, il)?,
2834 Mixer::Linear(la) => self.linear_attn_decode(e, la, &ws.h, cache, il)?,
2835 Mixer::Mla(mla) => self.mla_attn_cached(e, mla, &ws.h, pos_d, 1, il, cache)?,
2836 Mixer::Kda(la) => crate::kda::kda_decode_cached_q8(
2837 e,
2838 la,
2839 &ws.h,
2840 attn_q8.as_ref().map(|(q, d)| (q, d)),
2841 eps,
2842 cache,
2843 il,
2844 )?,
2845 };
2846 crate::hyper::post_t1_ws(e, topology, &mixed, x, ws, n_embd)?;
2847 std::mem::swap(x, &mut ws.xb);
2848
2849 crate::hyper::pre_t1_ws(e, topology, &hyper.mlp, x, ws, n_embd)?;
2850 let zq8 = if crate::glm5_q8_fuse_on() {
2853 let pair = e.rms_norm_zq8_f32(
2854 &ws.y,
2855 layer.post_attn_norm.float_data(),
2856 &mut ws.z,
2857 n_embd,
2858 1,
2859 eps,
2860 )?;
2861 if GLM5_Q8_FUSE_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed) == 0 {
2862 eprintln!(
2863 "[glm5-q8-fuse] engaged (rms_norm_zq8_f32, hyper_range_decode_ws_body)"
2864 );
2865 }
2866 Some(pair)
2867 } else {
2868 e.rms_norm(
2869 &ws.y,
2870 layer.post_attn_norm.float_data(),
2871 &mut ws.z,
2872 n_embd,
2873 1,
2874 eps,
2875 )?;
2876 None
2877 };
2878 let ffn_out = self.hyper_ffn_branch(e, layer, &ws.z, 1, il, false, zq8.as_ref())?;
2879 crate::hyper::post_t1_ws(e, topology, &ffn_out, x, ws, n_embd)?;
2880 std::mem::swap(x, &mut ws.xb);
2881 }
2882 Ok(())
2883 }
2884
2885 #[allow(clippy::too_many_arguments)]
2920 pub(crate) fn hyper_batch_range_decode(
2921 &self,
2922 e: &Engine,
2923 topology: &crate::hyper::HyperTopology,
2924 mut x: CudaSlice<f32>,
2925 lo: usize,
2926 hi: usize,
2927 pos_rows: &[CudaSlice<i32>],
2928 caches: &mut [&mut Cache],
2929 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2930 let b_n = caches.len();
2931 assert_eq!(
2932 pos_rows.len(),
2933 b_n,
2934 "hyper_batch_range_decode: pos_rows built for a different batch width"
2935 );
2936 if b_n == 1 && crate::hyper_batch_solo_on() {
2961 let pos = caches[0].pos;
2962 return self.hyper_range_decode(e, topology, x, lo, hi, &pos_rows[0], pos, caches[0]);
2963 }
2964 let n_embd = self.cfg.n_embd as usize;
2965 let eps = self.cfg.rms_eps;
2966 for il in lo..hi {
2967 let layer = &self.layers[il];
2968 let hyper = layer.hyper.as_ref().ok_or_else(|| {
2969 format!("layer {il} carries no hyper-connection weights under an hc plan")
2970 })?;
2971
2972 let (y, mix) = crate::hyper::pre_exact(e, topology, &hyper.attn, &x, b_n, n_embd)?;
2973 let mut h = e.uninit(b_n * n_embd)?;
2974 e.rms_norm(&y, layer.attn_norm.float_data(), &mut h, n_embd, b_n, eps)?;
2975 let mut mixed = e.uninit(b_n * n_embd)?;
2979 for bi in 0..b_n {
2980 let mut h_row = e.uninit(n_embd)?;
2981 e.dtod_copy_view(&h.slice(bi * n_embd..(bi + 1) * n_embd), &mut h_row)?;
2982 let cache: &mut Cache = &mut *caches[bi];
2983 let pos = cache.pos;
2984 let out_row = match &layer.mixer {
2985 Mixer::Full(fa) => {
2986 self.full_attn_decode(e, fa, &h_row, &pos_rows[bi], pos, cache, il)?
2987 }
2988 Mixer::Linear(la) => self.linear_attn_decode(e, la, &h_row, cache, il)?,
2989 Mixer::Mla(mla) => {
2990 self.mla_attn_cached(e, mla, &h_row, &pos_rows[bi], 1, il, cache)?
2991 }
2992 Mixer::Kda(la) => crate::kda::kda_decode_cached(e, la, &h_row, eps, cache, il)?,
2993 };
2994 e.copy_into(&mut mixed, bi * n_embd, &out_row, n_embd)?;
2995 }
2996 x = crate::hyper::post(e, topology, &mixed, &x, &mix, b_n, n_embd)?;
2997
2998 let (y, mix) = crate::hyper::pre_exact(e, topology, &hyper.mlp, &x, b_n, n_embd)?;
2999 let mut z = e.uninit(b_n * n_embd)?;
3000 e.rms_norm(
3001 &y,
3002 layer.post_attn_norm.float_data(),
3003 &mut z,
3004 n_embd,
3005 b_n,
3006 eps,
3007 )?;
3008 let ffn_out = self.hyper_ffn_branch_batch(e, layer, &z, b_n, il, false)?;
3009 x = crate::hyper::post(e, topology, &ffn_out, &x, &mix, b_n, n_embd)?;
3010 }
3011 Ok(x)
3012 }
3013
3014 pub(crate) fn hyper_ffn_branch_batch(
3034 &self,
3035 e: &Engine,
3036 layer: &crate::hybrid::HybridLayer,
3037 z: &CudaSlice<f32>,
3038 b_n: usize,
3039 il: usize,
3040 vrows: bool,
3041 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3042 let n_embd = self.cfg.n_embd as usize;
3043 match &layer.ffn {
3044 crate::hybrid::Ffn::Dense { .. } => {
3045 let mut out = e.uninit(b_n * n_embd)?;
3046 for bi in 0..b_n {
3047 let mut z_row = e.uninit(n_embd)?;
3048 e.dtod_copy_view(&z.slice(bi * n_embd..(bi + 1) * n_embd), &mut z_row)?;
3049 let row = self.hyper_ffn_branch(e, layer, &z_row, 1, il, false, None)?;
3050 e.copy_into(&mut out, bi * n_embd, &row, n_embd)?;
3051 }
3052 Ok(out)
3053 }
3054 crate::hybrid::Ffn::Moe(m) => {
3055 if vrows {
3056 self.moe_ffn_il_zq8_vrows(e, m, z, b_n, il as u16)
3057 } else {
3058 self.moe_ffn_il_zq8(e, m, z, None, b_n, il as u16)
3059 }
3060 }
3061 }
3062 }
3063
3064 fn forward_hyper_ppn(
3102 &self,
3103 e: &Engine,
3104 tokens: &[u32],
3105 last_only: bool,
3106 topology: &crate::hyper::HyperTopology,
3107 fence: &[usize],
3108 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
3109 let n_embd = self.cfg.n_embd as usize;
3110 let eps = self.cfg.rms_eps;
3111 let t = tokens.len();
3112 let width = topology.streams * n_embd;
3113 let trace = memra_reference::hidden_trace::enabled();
3114 if trace {
3115 memra_reference::hidden_trace::emit_tokens(tokens);
3116 }
3117 let pos: Vec<i32> = (0..t as i32).collect();
3118
3119 if crate::pp::pp2_streams_off() {
3120 let pos_d = e.htod_i32(&pos)?;
3124 let embedded = self.embed(e, tokens)?;
3125 let mut x = crate::hyper::expand(e, topology, &embedded, t, n_embd)?;
3126 x = self.hyper_range_forward(e, topology, x, fence[0], fence[1], &pos_d, t, trace)?;
3127 for s in 1..fence.len() - 1 {
3128 let boundary_tx = e.clone_dtod(&x)?;
3129 let boundary_rx = e.clone_dtod(&boundary_tx)?;
3130 x = self.hyper_range_forward(
3131 e,
3132 topology,
3133 boundary_rx,
3134 fence[s],
3135 fence[s + 1],
3136 &pos_d,
3137 t,
3138 trace,
3139 )?;
3140 }
3141 return self.hyper_head_logits(e, topology, &x, t, n_embd, eps, last_only);
3142 }
3143
3144 let rt = crate::pp::PpNRt::get(e)?;
3145 let n_st = fence.len() - 1;
3146 assert_eq!(
3147 rt.n_stages(),
3148 n_st,
3149 "PpNRt stage count {} != fence stages {n_st}",
3150 rt.n_stages()
3151 );
3152 rt.fence_stages_behind(&e.stream())?;
3155
3156 let mut slot = {
3157 let _st0 = rt.enter(0);
3158 let e0 = rt.engine(0, e);
3159 let pos_d = e0.htod_i32(&pos)?;
3162 let embedded = self.embed(e0, tokens)?;
3163 let x = crate::hyper::expand(e0, topology, &embedded, t, n_embd)?;
3164 let x =
3165 self.hyper_range_forward(e0, topology, x, fence[0], fence[1], &pos_d, t, trace)?;
3166 rt.tx(0, &x, t * width)?
3167 };
3168 for s in 1..n_st - 1 {
3169 let _st = rt.enter(s);
3170 let es = rt.engine(s, e);
3171 let pos_d = es.htod_i32(&pos)?;
3172 let x = rt.rx(s - 1, slot, t * width)?;
3173 let x = self.hyper_range_forward(
3174 es,
3175 topology,
3176 x,
3177 fence[s],
3178 fence[s + 1],
3179 &pos_d,
3180 t,
3181 trace,
3182 )?;
3183 slot = rt.tx(s, &x, t * width)?;
3184 }
3185 let _stl = rt.enter(n_st - 1);
3186 let el = rt.engine(n_st - 1, e);
3187 let pos_d = el.htod_i32(&pos)?;
3188 let x = rt.rx(n_st - 2, slot, t * width)?;
3189 let x = self.hyper_range_forward(
3190 el,
3191 topology,
3192 x,
3193 fence[n_st - 1],
3194 fence[n_st],
3195 &pos_d,
3196 t,
3197 trace,
3198 )?;
3199 self.hyper_head_logits(el, topology, &x, t, n_embd, eps, last_only)
3200 }
3201
3202 #[allow(clippy::too_many_arguments)]
3207 fn hyper_head_logits(
3208 &self,
3209 e: &Engine,
3210 topology: &crate::hyper::HyperTopology,
3211 x: &CudaSlice<f32>,
3212 t: usize,
3213 n_embd: usize,
3214 eps: f32,
3215 last_only: bool,
3216 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
3217 let collapsed =
3218 crate::hyper::collapse(e, topology, self.hyper_head.as_ref(), x, t, n_embd)?;
3219 if memra_reference::hidden_trace::enabled() {
3220 memra_reference::hidden_trace::emit_last_row(
3221 "collapse",
3222 -1,
3223 t,
3224 n_embd,
3225 &e.dtoh(&collapsed)?,
3226 );
3227 }
3228 let mut hn = e.uninit(t * n_embd)?;
3229 e.rms_norm(
3230 &collapsed,
3231 self.output_norm.float_data(),
3232 &mut hn,
3233 n_embd,
3234 t,
3235 eps,
3236 )?;
3237 let logits = if last_only {
3238 let last = e.view(&hn, t * n_embd);
3239 let last_row = last.slice((t - 1) * n_embd..t * n_embd);
3240 let mut hlast = e.uninit(n_embd)?;
3241 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
3242 e.matmul(&self.output, &hlast, 1)?
3243 } else {
3244 e.matmul(&self.output, &hn, t)?
3245 };
3246 e.dtoh(&logits)
3247 }
3248
3249 #[allow(clippy::type_complexity)] #[allow(clippy::too_many_arguments)]
3254 fn prime_cache_hyper_ppn(
3255 &self,
3256 e: &Engine,
3257 tokens: &[u32],
3258 cache: &mut Cache,
3259 queued_after: usize,
3260 topology: &crate::hyper::HyperTopology,
3261 fence: &[usize],
3262 overlay: Option<&crate::vision::EmbedOverlay>,
3263 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
3264 let n_embd = self.cfg.n_embd as usize;
3265 let eps = self.cfg.rms_eps;
3266 let t = tokens.len();
3267 let width = topology.streams * n_embd;
3268 if cache.pos + t > cache.max_ctx {
3269 return Err("prime_cache: prompt exceeds cache max_ctx".into());
3270 }
3271 let seq_end = cache.pos + t + queued_after;
3272 let pos: Vec<i32> = (cache.pos as i32..(cache.pos + t) as i32).collect();
3273 if let Some(sink) = cache.hc_taps.as_mut() {
3276 sink.base = cache.pos;
3277 }
3278
3279 if crate::pp::pp2_streams_off() {
3280 let pos_d = e.htod_i32(&pos)?;
3281 let mut embedded = self.embed(e, tokens)?;
3282 if let Some(ov) = overlay {
3283 ov.splice_into(e, &mut embedded, 0, t, n_embd)?;
3287 }
3288 let mut x = crate::hyper::expand(e, topology, &embedded, t, n_embd)?;
3289 x = self.hyper_range_prime(
3290 e,
3291 topology,
3292 x,
3293 fence[0],
3294 fence[1],
3295 &pos_d,
3296 t,
3297 cache,
3298 seq_end,
3299 HcTapArm::FromCache,
3300 )?;
3301 for s in 1..fence.len() - 1 {
3302 let boundary_tx = e.clone_dtod(&x)?;
3303 let boundary_rx = e.clone_dtod(&boundary_tx)?;
3304 x = self.hyper_range_prime(
3305 e,
3306 topology,
3307 boundary_rx,
3308 fence[s],
3309 fence[s + 1],
3310 &pos_d,
3311 t,
3312 cache,
3313 seq_end,
3314 HcTapArm::FromCache,
3315 )?;
3316 }
3317 return self.hyper_prime_tail(e, topology, &x, t, n_embd, eps, cache);
3318 }
3319
3320 {
3321 let rt = crate::pp::PpNRt::get(e)?;
3322 let n_st = fence.len() - 1;
3323 assert_eq!(
3324 rt.n_stages(),
3325 n_st,
3326 "PpNRt stage count {} != fence stages {n_st}",
3327 rt.n_stages()
3328 );
3329 let caller_stream = e.stream();
3330 rt.fence_stages_behind(&caller_stream)?;
3331 if let Some(ov) = overlay
3355 && !ov.resident_in(rt.engine(0, e))
3356 {
3357 return Err(format!(
3358 "vision embedding overlay rows are resident on dev{} but pp stage 0's \
3359 embedding intake runs on dev{}: the overlay must be published into the \
3360 intake engine's context (build it with EmbedOverlay::new_published; \
3361 MEMRA_VISION_OVERLAY_PUBLISH=0 pins the pre-publication program, whose \
3362 only vision-capable shape is MEMRA_PP_STREAMS=0)",
3363 ov.ctx().ordinal(),
3364 rt.engine(0, e).ctx().ordinal(),
3365 )
3366 .into());
3367 }
3368 let mut slot = {
3369 let _st0 = rt.enter(0);
3370 let e0 = rt.engine(0, e);
3371 let pos_d = e0.htod_i32(&pos)?;
3372 let mut embedded = self.embed(e0, tokens)?;
3373 if let Some(ov) = overlay {
3374 ov.splice_into(e0, &mut embedded, 0, t, n_embd)?;
3380 }
3381 let x = crate::hyper::expand(e0, topology, &embedded, t, n_embd)?;
3382 let x = self.hyper_range_prime(
3383 e0,
3384 topology,
3385 x,
3386 fence[0],
3387 fence[1],
3388 &pos_d,
3389 t,
3390 cache,
3391 seq_end,
3392 HcTapArm::FromCache,
3393 )?;
3394 rt.tx(0, &x, t * width)?
3395 };
3396 for s in 1..n_st - 1 {
3397 let _st = rt.enter(s);
3398 let es = rt.engine(s, e);
3399 let pos_d = es.htod_i32(&pos)?;
3400 let x = rt.rx(s - 1, slot, t * width)?;
3401 let x = self.hyper_range_prime(
3402 es,
3403 topology,
3404 x,
3405 fence[s],
3406 fence[s + 1],
3407 &pos_d,
3408 t,
3409 cache,
3410 seq_end,
3411 HcTapArm::FromCache,
3412 )?;
3413 slot = rt.tx(s, &x, t * width)?;
3414 }
3415 let out = {
3416 let _stl = rt.enter(n_st - 1);
3417 let el = rt.engine(n_st - 1, e);
3418 let pos_d = el.htod_i32(&pos)?;
3419 let x = rt.rx(n_st - 2, slot, t * width)?;
3420 let x = self.hyper_range_prime(
3421 el,
3422 topology,
3423 x,
3424 fence[n_st - 1],
3425 fence[n_st],
3426 &pos_d,
3427 t,
3428 cache,
3429 seq_end,
3430 HcTapArm::FromCache,
3431 )?;
3432 self.hyper_prime_tail(el, topology, &x, t, n_embd, eps, cache)?
3433 };
3434 rt.publish_all_to(&caller_stream)?;
3442 Ok(out)
3443 }
3444 }
3445
3446 #[allow(clippy::too_many_arguments)] #[allow(clippy::type_complexity)] fn prime_cache_hyper_pp2_pipelined(
3475 &self,
3476 e: &Engine,
3477 tokens: &[u32],
3478 cache: &mut Cache,
3479 seq_end: usize,
3480 topology: &crate::hyper::HyperTopology,
3481 ranges: &[(usize, usize)],
3482 fence: &[usize],
3483 taps: Option<&std::sync::Mutex<&mut crate::cache::HcTapSink>>,
3484 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
3485 debug_assert_eq!(fence.len(), 3);
3486 debug_assert!(ranges.len() >= 2);
3487 let arm = |base: usize| match taps {
3490 Some(sink) => HcTapArm::Shared(sink, base),
3491 None => HcTapArm::FromCache,
3492 };
3493 let rt = crate::pp::PpNRt::get(e)?;
3494 assert_eq!(
3495 rt.n_stages(),
3496 2,
3497 "the pipelined mHC prime requires exactly two PP stages"
3498 );
3499 let n_embd = self.cfg.n_embd as usize;
3500 let eps = self.cfg.rms_eps;
3501 let width = topology.streams * n_embd;
3502 let t = tokens.len();
3503 let initial_base = cache.pos;
3504 let caller_stream = e.stream();
3505
3506 rt.fence_stages_behind(&caller_stream)?;
3510 let max_payload = ranges.iter().map(|(s, x)| (x - s) * width).max().unwrap();
3511 rt.prepare_overlap_slots(0, max_payload)?;
3512
3513 static SAID: std::sync::atomic::AtomicU8 = std::sync::atomic::AtomicU8::new(0);
3518 let bit = 1u8 << u8::from(taps.is_some());
3519 if SAID.fetch_or(bit, std::sync::atomic::Ordering::Relaxed) & bit == 0 {
3520 eprintln!(
3521 "[prime-v2] arm2 pipelined stages=2 chunks={} overlap={} t={t} devices={:?} \
3522 taps={} (stage 0 of chunk k+1 overlaps stage 1 of chunk k on two host \
3523 threads; MEMRA_B200_PRIME_V2, logged once per process)",
3524 ranges.len(),
3525 ranges.len().saturating_sub(1),
3526 (0..2)
3527 .map(|s| rt.engine(s, e).ctx().ordinal())
3528 .collect::<Vec<_>>(),
3529 if taps.is_some() { "armed" } else { "none" },
3530 );
3531 }
3532
3533 let mut hiddens = e.uninit(t * n_embd)?;
3534 let mut last: Option<(Vec<f32>, CudaSlice<f32>)> = None;
3535 let mut stage_caches = PrimeCacheStages::new(cache, fence);
3536 let (cache0, cache1) = stage_caches.pp2_parts();
3537 let (first_start, first_end) = ranges[0];
3538 let mut slot = self.prime_hyper_pp2_stage0_enqueue(
3539 e,
3540 rt,
3541 topology,
3542 &tokens[first_start..first_end],
3543 cache0,
3544 seq_end,
3545 fence,
3546 initial_base + first_start,
3547 arm(initial_base + first_start),
3548 )?;
3549 cache0.pos = initial_base + first_end;
3550
3551 for (i, &(start, end)) in ranges.iter().enumerate() {
3552 let base = initial_base + start;
3553 debug_assert_eq!(
3554 cache1.pos, base,
3555 "stage 1 must drain chunks in original position order"
3556 );
3557 let (out, next_slot) = if let Some(&(next_start, next_end)) = ranges.get(i + 1) {
3558 let next_base = initial_base + next_start;
3559 debug_assert_eq!(
3560 cache0.pos, next_base,
3561 "stage 0 must issue chunks in original position order"
3562 );
3563 let cache0_stage = &mut *cache0;
3564 std::thread::scope(|scope| -> Result<_, Box<dyn std::error::Error>> {
3565 let stage0 = scope.spawn(move || -> Result<usize, String> {
3566 let next = self
3567 .prime_hyper_pp2_stage0_enqueue(
3568 e,
3569 rt,
3570 topology,
3571 &tokens[next_start..next_end],
3572 cache0_stage,
3573 seq_end,
3574 fence,
3575 next_base,
3576 arm(next_base),
3577 )
3578 .map_err(|err| err.to_string())?;
3579 cache0_stage.pos = initial_base + next_end;
3580 Ok(next)
3581 });
3582 let x = self.prime_hyper_pp2_stage1_enqueue(
3583 e,
3584 rt,
3585 topology,
3586 slot,
3587 end - start,
3588 cache1,
3589 seq_end,
3590 fence,
3591 base,
3592 arm(base),
3593 )?;
3594 let out = {
3595 rt.bind_stage(1)?;
3596 let _st1 = rt.enter(1);
3597 let e1 = rt.engine(1, e);
3598 self.hyper_prime_tail(e1, topology, &x, end - start, n_embd, eps, cache1)?
3599 };
3600 let next = match stage0.join() {
3601 Ok(result) => {
3602 result.map_err(|err| -> Box<dyn std::error::Error> { err.into() })?
3603 }
3604 Err(payload) => std::panic::resume_unwind(payload),
3605 };
3606 Ok((out, Some(next)))
3607 })?
3608 } else {
3609 let x = self.prime_hyper_pp2_stage1_enqueue(
3610 e,
3611 rt,
3612 topology,
3613 slot,
3614 end - start,
3615 cache1,
3616 seq_end,
3617 fence,
3618 base,
3619 arm(base),
3620 )?;
3621 let out = {
3622 rt.bind_stage(1)?;
3623 let _st1 = rt.enter(1);
3624 let e1 = rt.engine(1, e);
3625 self.hyper_prime_tail(e1, topology, &x, end - start, n_embd, eps, cache1)?
3626 };
3627 (out, None)
3628 };
3629
3630 rt.publish_to(1, &caller_stream)?;
3631 e.copy_into(&mut hiddens, start * n_embd, &out.2, (end - start) * n_embd)?;
3632 last = Some((out.0, out.1));
3633 crate::pp::PRIME_SPLIT_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
3634 crate::progress::note_prime_rows(end - start);
3638 HYPER_PRIME_PIPELINED_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
3639
3640 if let Some(next) = next_slot {
3641 rt.fence_stages_behind(&caller_stream)?;
3646 slot = next;
3647 }
3648 }
3649
3650 debug_assert_eq!(cache0.pos, initial_base + t);
3651 debug_assert_eq!(cache1.pos, initial_base + t);
3652 let (logits, h_seed) = last.expect("the pipelined mHC prime ran at least one chunk");
3653 stage_caches.commit();
3654 Ok((logits, h_seed, hiddens))
3655 }
3656
3657 #[allow(clippy::too_many_arguments)] fn prime_hyper_pp2_stage0_enqueue(
3663 &self,
3664 e: &Engine,
3665 rt: &crate::pp::PpNRt,
3666 topology: &crate::hyper::HyperTopology,
3667 tokens: &[u32],
3668 cache: &mut Cache,
3669 seq_end: usize,
3670 fence: &[usize],
3671 base: usize,
3672 taps: HcTapArm<'_, '_>,
3673 ) -> Result<usize, Box<dyn std::error::Error>> {
3674 let t = tokens.len();
3675 let n_embd = self.cfg.n_embd as usize;
3676 let width = topology.streams * n_embd;
3677 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
3678 rt.bind_stage(0)?;
3679 let _st0 = rt.enter(0);
3680 let e0 = rt.engine(0, e);
3681 let pos_d = e0.htod_i32(&pos)?;
3682 let embedded = self.embed(e0, tokens)?;
3683 let x = crate::hyper::expand(e0, topology, &embedded, t, n_embd)?;
3684 let _overlap = crate::pp::enter_prime_pipe_stage();
3685 let x = self.hyper_range_prime(
3686 e0, topology, x, fence[0], fence[1], &pos_d, t, cache, seq_end, taps,
3687 )?;
3688 rt.tx_pipelined(0, &x, t * width)
3689 }
3690
3691 #[allow(clippy::too_many_arguments)] fn prime_hyper_pp2_stage1_enqueue(
3696 &self,
3697 e: &Engine,
3698 rt: &crate::pp::PpNRt,
3699 topology: &crate::hyper::HyperTopology,
3700 slot: usize,
3701 t: usize,
3702 cache: &mut Cache,
3703 seq_end: usize,
3704 fence: &[usize],
3705 base: usize,
3706 taps: HcTapArm<'_, '_>,
3707 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3708 let n_embd = self.cfg.n_embd as usize;
3709 let width = topology.streams * n_embd;
3710 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
3711 rt.bind_stage(1)?;
3712 let _st1 = rt.enter(1);
3713 let e1 = rt.engine(1, e);
3714 let pos_d = e1.htod_i32(&pos)?;
3715 let x = rt.rx(0, slot, t * width)?;
3716 let _overlap = crate::pp::enter_prime_pipe_stage();
3717 self.hyper_range_prime(
3718 e1, topology, x, fence[1], fence[2], &pos_d, t, cache, seq_end, taps,
3719 )
3720 }
3721
3722 #[allow(clippy::type_complexity)] fn prime_chunk_hyper(
3728 &self,
3729 e: &Engine,
3730 tokens: &[u32],
3731 cache: &mut Cache,
3732 seq_end: usize,
3733 chunk_off: usize,
3734 overlay: Option<&crate::vision::EmbedOverlay>,
3735 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
3736 let topology = *self
3737 .hyper
3738 .as_ref()
3739 .ok_or("prime_chunk_hyper on a model with no HyperConnections topology")?;
3740 let n_embd = self.cfg.n_embd as usize;
3741 let t = tokens.len();
3742 let eps = self.cfg.rms_eps;
3743 let pos: Vec<i32> = (cache.pos as i32..(cache.pos + t) as i32).collect();
3744 let pos_d = e.htod_i32(&pos)?;
3745 if let Some(sink) = cache.hc_taps.as_mut() {
3748 sink.base = cache.pos;
3749 }
3750
3751 let mut embedded = self.embed(e, tokens)?;
3752 if let Some(ov) = overlay {
3753 ov.splice_into(e, &mut embedded, chunk_off, t, n_embd)?;
3759 }
3760 let mut x = crate::hyper::expand(e, &topology, &embedded, t, n_embd)?;
3761
3762 for (il, layer) in self.layers.iter().enumerate() {
3763 let hyper = layer.hyper.as_ref().ok_or_else(|| {
3764 format!("layer {il} carries no hyper-connection weights under an hc plan")
3765 })?;
3766
3767 let (y, mix) = crate::hyper::pre(e, &topology, &hyper.attn, &x, t, n_embd)?;
3768 let mut h = e.uninit(t * n_embd)?;
3769 e.rms_norm(&y, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
3770 let mixed = match &layer.mixer {
3771 Mixer::Full(fa) => {
3772 self.full_attn_prime(e, fa, &h, None, &pos_d, t, cache, il, seq_end)?
3773 }
3774 Mixer::Linear(la) => self.linear_attn_prime(e, la, &h, None, t, cache, il)?,
3775 Mixer::Mla(mla) if mla.tp.is_some() => {
3776 self.mla_tp_attn_cached(e, mla, &h, &pos_d, t, il, cache, false)?
3777 }
3778 Mixer::Mla(mla) => self.mla_attn_cached(e, mla, &h, &pos_d, t, il, cache)?,
3779 Mixer::Kda(la) if la.tp.is_some() => crate::glm5_tp::kda_tp_cached(
3780 e,
3781 la,
3782 &h,
3783 t,
3784 eps,
3785 cache,
3786 il,
3787 crate::kda::ConvArm::Prefill,
3788 )?,
3789 Mixer::Kda(la) => crate::kda::kda_prime_cached(e, la, &h, t, eps, cache, il)?,
3790 };
3791 x = crate::hyper::post(e, &topology, &mixed, &x, &mix, t, n_embd)?;
3792
3793 let (y, mix) = crate::hyper::pre(e, &topology, &hyper.mlp, &x, t, n_embd)?;
3794 let mut z = e.uninit(t * n_embd)?;
3795 e.rms_norm(
3796 &y,
3797 layer.post_attn_norm.float_data(),
3798 &mut z,
3799 n_embd,
3800 t,
3801 eps,
3802 )?;
3803 let ffn_out = self.hyper_ffn_branch(e, layer, &z, t, il, true, None)?;
3804 x = crate::hyper::post(e, &topology, &ffn_out, &x, &mix, t, n_embd)?;
3805 self.glm5_hc_tap(e, cache, &topology, il, &x, t)?;
3808 }
3809
3810 let hiddens =
3811 crate::hyper::collapse(e, &topology, self.hyper_head.as_ref(), &x, t, n_embd)?;
3812 let mut hn = e.uninit(t * n_embd)?;
3813 e.rms_norm(
3814 &hiddens,
3815 self.output_norm.float_data(),
3816 &mut hn,
3817 n_embd,
3818 t,
3819 eps,
3820 )?;
3821 let last = e.view(&hn, t * n_embd);
3822 let last_row = last.slice((t - 1) * n_embd..t * n_embd);
3823 let mut hlast = e.uninit(n_embd)?;
3824 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
3825 let logits = e.matmul(&self.output, &hlast, 1)?;
3826 let host = e.dtoh(&logits)?;
3827
3828 let stack = e.view(&hiddens, t * n_embd);
3831 let seed_row = stack.slice((t - 1) * n_embd..t * n_embd);
3832 let mut h_seed = e.uninit(n_embd)?;
3833 e.copy_view_into(&mut h_seed, 0, &seed_row, n_embd)?;
3834 cache.pos += t;
3835 Ok((host, h_seed, hiddens))
3836 }
3837
3838 #[allow(clippy::too_many_arguments)]
3842 #[allow(clippy::type_complexity)] fn hyper_prime_tail(
3844 &self,
3845 e: &Engine,
3846 topology: &crate::hyper::HyperTopology,
3847 x: &CudaSlice<f32>,
3848 t: usize,
3849 n_embd: usize,
3850 eps: f32,
3851 cache: &mut Cache,
3852 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
3853 let hiddens = crate::hyper::collapse(e, topology, self.hyper_head.as_ref(), x, t, n_embd)?;
3854 let mut hn = e.uninit(t * n_embd)?;
3855 e.rms_norm(
3856 &hiddens,
3857 self.output_norm.float_data(),
3858 &mut hn,
3859 n_embd,
3860 t,
3861 eps,
3862 )?;
3863 let last = e.view(&hn, t * n_embd);
3864 let last_row = last.slice((t - 1) * n_embd..t * n_embd);
3865 let mut hlast = e.uninit(n_embd)?;
3866 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
3867 let logits = e.matmul(&self.output, &hlast, 1)?;
3868 let host = e.dtoh(&logits)?;
3869 let stack = e.view(&hiddens, t * n_embd);
3870 let seed_row = stack.slice((t - 1) * n_embd..t * n_embd);
3871 let mut h_seed = e.uninit(n_embd)?;
3872 e.copy_view_into(&mut h_seed, 0, &seed_row, n_embd)?;
3873 cache.pos += t;
3874 Ok((host, h_seed, hiddens))
3875 }
3876
3877 fn decode_step_hyper_ppn(
3882 &self,
3883 e: &Engine,
3884 token: u32,
3885 cache: &mut Cache,
3886 topology: &crate::hyper::HyperTopology,
3887 fence: &[usize],
3888 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
3889 let n_embd = self.cfg.n_embd as usize;
3890 let eps = self.cfg.rms_eps;
3891 let pos = cache.pos;
3892 let width = topology.streams * n_embd;
3893
3894 if crate::pp::pp2_streams_off() {
3895 let pos_d = e.htod_i32(&[pos as i32])?;
3896 let embedded = e.htod(&self.embd.try_gather(n_embd, &[token])?)?;
3897 let mut x = crate::hyper::expand(e, topology, &embedded, 1, n_embd)?;
3898 x = self.hyper_range_decode(e, topology, x, fence[0], fence[1], &pos_d, pos, cache)?;
3899 for s in 1..fence.len() - 1 {
3900 let boundary_tx = e.clone_dtod(&x)?;
3901 let boundary_rx = e.clone_dtod(&boundary_tx)?;
3902 x = self.hyper_range_decode(
3903 e,
3904 topology,
3905 boundary_rx,
3906 fence[s],
3907 fence[s + 1],
3908 &pos_d,
3909 pos,
3910 cache,
3911 )?;
3912 }
3913 return self.hyper_decode_tail(e, topology, &x, n_embd, eps, cache);
3914 }
3915
3916 let rt = crate::pp::PpNRt::get(e)?;
3917 let n_st = fence.len() - 1;
3918 assert_eq!(
3919 rt.n_stages(),
3920 n_st,
3921 "PpNRt stage count {} != fence stages {n_st}",
3922 rt.n_stages()
3923 );
3924 rt.fence_stages_behind(&e.stream())?;
3925
3926 let mut slot = {
3927 let _st0 = rt.enter(0);
3928 let e0 = rt.engine(0, e);
3929 let pos_d = e0.htod_i32(&[pos as i32])?;
3930 let embedded = e0.htod(&self.embd.try_gather(n_embd, &[token])?)?;
3931 let x = crate::hyper::expand(e0, topology, &embedded, 1, n_embd)?;
3932 let x =
3933 self.hyper_range_decode(e0, topology, x, fence[0], fence[1], &pos_d, pos, cache)?;
3934 rt.tx(0, &x, width)?
3935 };
3936 for s in 1..n_st - 1 {
3937 let _st = rt.enter(s);
3938 let es = rt.engine(s, e);
3939 let pos_d = es.htod_i32(&[pos as i32])?;
3940 let x = rt.rx(s - 1, slot, width)?;
3941 let x = self.hyper_range_decode(
3942 es,
3943 topology,
3944 x,
3945 fence[s],
3946 fence[s + 1],
3947 &pos_d,
3948 pos,
3949 cache,
3950 )?;
3951 slot = rt.tx(s, &x, width)?;
3952 }
3953 let _stl = rt.enter(n_st - 1);
3954 let el = rt.engine(n_st - 1, e);
3955 let pos_d = el.htod_i32(&[pos as i32])?;
3956 let x = rt.rx(n_st - 2, slot, width)?;
3957 let x = self.hyper_range_decode(
3958 el,
3959 topology,
3960 x,
3961 fence[n_st - 1],
3962 fence[n_st],
3963 &pos_d,
3964 pos,
3965 cache,
3966 )?;
3967 self.hyper_decode_tail(el, topology, &x, n_embd, eps, cache)
3968 }
3969
3970 fn hyper_decode_tail(
3974 &self,
3975 e: &Engine,
3976 topology: &crate::hyper::HyperTopology,
3977 x: &CudaSlice<f32>,
3978 n_embd: usize,
3979 eps: f32,
3980 cache: &mut Cache,
3981 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
3982 let h_seed = crate::hyper::collapse(e, topology, self.hyper_head.as_ref(), x, 1, n_embd)?;
3983 let mut hn = e.uninit(n_embd)?;
3984 e.rms_norm(
3985 &h_seed,
3986 self.output_norm.float_data(),
3987 &mut hn,
3988 n_embd,
3989 1,
3990 eps,
3991 )?;
3992 let logits = e.matmul(&self.output, &hn, 1)?;
3993 let host = e.dtoh(&logits)?;
3994 cache.pos += 1;
3995 Ok((host, h_seed))
3996 }
3997
3998 pub fn forward(
4000 &self,
4001 e: &Engine,
4002 tokens: &[u32],
4003 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
4004 if self.hyper.is_some() {
4005 return self.forward_hyper(e, tokens, false);
4006 }
4007 if self.is_gemma4_e4b() {
4008 return self.gemma4_e4b_forward(e, tokens, false);
4009 }
4010 if self.uses_gemma_program() {
4011 return self.gemma4_forward(e, tokens, false);
4012 }
4013 let cfg = &self.cfg;
4014 let n_embd = cfg.n_embd as usize;
4015 let t = tokens.len();
4016 let eps = cfg.rms_eps;
4017 let pos: Vec<i32> = (0..t as i32).collect();
4018 let pos_d = e.htod_i32(&pos)?;
4019
4020 let mut x = self.embed(e, tokens)?; for (il, layer) in self.layers.iter().enumerate() {
4023 let mut h = e.uninit(t * n_embd)?;
4025 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
4026
4027 let mixed = match &layer.mixer {
4028 Mixer::Full(fa) => self.full_attn(e, fa, &h, &pos_d, t, il)?,
4029 Mixer::Linear(la) => self.linear_attn(e, la, &h, t)?,
4030 Mixer::Mla(mla) => self.mla_attn(e, mla, &h, &pos_d, t, il)?,
4031 Mixer::Kda(la) => crate::kda::kda_attn(e, la, &h, t, eps)?,
4032 };
4033
4034 let mut x1 = e.uninit(t * n_embd)?;
4036 e.add(&x, &mixed, &mut x1, t * n_embd)?;
4037
4038 let mut z = e.uninit(t * n_embd)?;
4040 e.rms_norm(
4041 &x1,
4042 layer.post_attn_norm.float_data(),
4043 &mut z,
4044 n_embd,
4045 t,
4046 eps,
4047 )?;
4048 let ffn_out = match &layer.ffn {
4049 crate::hybrid::Ffn::Dense {
4050 ffn_gate,
4051 ffn_up,
4052 ffn_down,
4053 } => {
4054 let n_ff = ffn_gate.out_features();
4055 let mut g2 = e.matmul_group(&[ffn_gate, ffn_up], &z, t)?;
4056 let up = g2.pop().unwrap();
4057 let gate = g2.pop().unwrap();
4058 let mut act = e.uninit(t * n_ff)?;
4059 Self::ffn_act_lim(
4064 e,
4065 &self.cfg,
4066 &gate,
4067 &up,
4068 1.0,
4069 1.0,
4070 self.cfg.clamp_shexp_at(il as u32),
4071 &mut act,
4072 t * n_ff,
4073 )?;
4074 e.matmul(ffn_down, &act, t)?
4075 }
4076 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il_prefill(e, m, &z, t, il as u16)?,
4077 };
4078 let mut x2 = e.uninit(t * n_embd)?;
4079 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
4080 x = x2;
4081 }
4082
4083 let mut hn = e.uninit(t * n_embd)?;
4084 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
4085 let logits = e.matmul(&self.output, &hn, t)?;
4086 e.dtoh(&logits)
4087 }
4088
4089 pub fn forward_last(
4095 &self,
4096 e: &Engine,
4097 tokens: &[u32],
4098 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
4099 if self.hyper.is_some() {
4100 return self.forward_hyper(e, tokens, true);
4101 }
4102 if self.uses_gemma_program() {
4103 return self.gemma4_forward(e, tokens, true);
4104 }
4105 let cfg = &self.cfg;
4106 let n_embd = cfg.n_embd as usize;
4107 let t = tokens.len();
4108 let eps = cfg.rms_eps;
4109 let pos: Vec<i32> = (0..t as i32).collect();
4110 let pos_d = e.htod_i32(&pos)?;
4111
4112 let mut x = self.embed(e, tokens)?; let probe = std::env::var("MEMRA_LAYER_PROBE").is_ok();
4116 let anat = Self::prime_anatomy_on();
4117 let mut anat_last = if anat {
4118 e.stream().synchronize()?;
4119 Some(std::time::Instant::now())
4120 } else {
4121 None
4122 };
4123 macro_rules! anat_mark {
4124 ($slot:expr) => {
4125 if let Some(ts) = anat_last.as_mut() {
4126 e.stream().synchronize()?;
4127 Self::prime_anatomy_slots()[$slot].fetch_add(
4128 ts.elapsed().as_nanos() as u64,
4129 std::sync::atomic::Ordering::Relaxed,
4130 );
4131 *ts = std::time::Instant::now();
4132 }
4133 };
4134 }
4135 for (il, layer) in self.layers.iter().enumerate() {
4136 let mut h = e.uninit(t * n_embd)?;
4137 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
4138 if probe {
4139 e.stream().synchronize()?;
4140 eprintln!("[probe] L{il} norm ok");
4141 }
4142 anat_mark!(4);
4143 let mixed = match &layer.mixer {
4144 Mixer::Full(fa) => {
4145 let y = self.full_attn(e, fa, &h, &pos_d, t, il)?;
4146 anat_mark!(0);
4147 y
4148 }
4149 Mixer::Linear(la) => {
4150 let y = self.linear_attn(e, la, &h, t)?;
4151 anat_mark!(1);
4152 y
4153 }
4154 Mixer::Mla(mla) => self.mla_attn(e, mla, &h, &pos_d, t, il)?,
4155 Mixer::Kda(la) => {
4156 let y = crate::kda::kda_attn(e, la, &h, t, eps)?;
4157 anat_mark!(1);
4159 y
4160 }
4161 };
4162 if probe {
4163 e.stream().synchronize()?;
4164 eprintln!("[probe] L{il} mixer ok");
4165 }
4166 let mut x1 = e.uninit(t * n_embd)?;
4167 e.add(&x, &mixed, &mut x1, t * n_embd)?;
4168 let mut z = e.uninit(t * n_embd)?;
4169 e.rms_norm(
4170 &x1,
4171 layer.post_attn_norm.float_data(),
4172 &mut z,
4173 n_embd,
4174 t,
4175 eps,
4176 )?;
4177 anat_mark!(4);
4178 let ffn_out = match &layer.ffn {
4179 crate::hybrid::Ffn::Dense {
4180 ffn_gate,
4181 ffn_up,
4182 ffn_down,
4183 } => {
4184 let n_ff = ffn_gate.out_features();
4185 let mut g2 = e.matmul_group(&[ffn_gate, ffn_up], &z, t)?;
4186 let up = g2.pop().unwrap();
4187 let gate = g2.pop().unwrap();
4188 let mut act = e.uninit(t * n_ff)?;
4189 Self::ffn_act_lim(
4191 e,
4192 &self.cfg,
4193 &gate,
4194 &up,
4195 1.0,
4196 1.0,
4197 self.cfg.clamp_shexp_at(il as u32),
4198 &mut act,
4199 t * n_ff,
4200 )?;
4201 let y = e.matmul(ffn_down, &act, t)?;
4202 anat_mark!(3);
4203 y
4204 }
4205 crate::hybrid::Ffn::Moe(m) => {
4206 let y = self.moe_ffn_il_prefill(e, m, &z, t, il as u16)?;
4207 anat_mark!(2);
4208 y
4209 }
4210 };
4211 if probe {
4212 e.stream().synchronize()?;
4213 eprintln!("[probe] L{il} ffn ok");
4214 }
4215 let mut x2 = e.uninit(t * n_embd)?;
4216 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
4217 x = x2;
4218 }
4219 if anat {
4220 let s = Self::prime_anatomy_slots();
4221 let ms = |i: usize| s[i].load(std::sync::atomic::Ordering::Relaxed) as f64 / 1.0e6;
4222 eprintln!(
4223 "[prime-anatomy] cumulative ms: attn_full={:.1} gdn_linear={:.1} moe={:.1} \
4224 dense={:.1} norms_adds={:.1} (t={t}, forward_last)",
4225 ms(0),
4226 ms(1),
4227 ms(2),
4228 ms(3),
4229 ms(4)
4230 );
4231 }
4232 let mut hn = e.uninit(t * n_embd)?;
4234 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
4235 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)?;
4238 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
4239 let logits = e.matmul(&self.output, &hlast, 1)?; e.dtoh(&logits)
4241 }
4242
4243 #[allow(clippy::type_complexity)] pub fn prime_cache(
4276 &self,
4277 e: &Engine,
4278 tokens: &[u32],
4279 cache: &mut Cache,
4280 queued_after: usize,
4281 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4282 self.prime_cache_overlaid(e, tokens, cache, queued_after, None)
4283 }
4284
4285 pub fn vision_intake_engine<'a>(
4300 &self,
4301 e: &'a Engine,
4302 ) -> Result<&'a Engine, Box<dyn std::error::Error>> {
4303 if self.hyper.is_some()
4304 && !crate::pp::pp2_streams_off()
4305 && crate::pp::pp_cuts(self.layers.len()).is_some()
4306 {
4307 let rt = crate::pp::PpNRt::get(e)?;
4308 return Ok(rt.engine(0, e));
4309 }
4310 Ok(e)
4311 }
4312
4313 #[allow(clippy::type_complexity)] pub fn prime_cache_overlaid(
4322 &self,
4323 e: &Engine,
4324 tokens: &[u32],
4325 cache: &mut Cache,
4326 queued_after: usize,
4327 overlay: Option<&crate::vision::EmbedOverlay>,
4328 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4329 let events_before = crate::progress::events();
4339 let out = self.prime_cache_overlaid_inner(e, tokens, cache, queued_after, overlay);
4340 if out.is_ok() && crate::progress::events() == events_before {
4341 crate::progress::note_prime_rows(tokens.len());
4342 }
4343 out
4344 }
4345
4346 #[allow(clippy::type_complexity)] fn prime_cache_overlaid_inner(
4348 &self,
4349 e: &Engine,
4350 tokens: &[u32],
4351 cache: &mut Cache,
4352 queued_after: usize,
4353 overlay: Option<&crate::vision::EmbedOverlay>,
4354 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4355 cache.ensure_usable("prime_cache")?;
4356 if self.hyper.is_some() {
4357 return self.prime_cache_hyper(e, tokens, cache, queued_after, overlay);
4361 }
4362 let _pp_walk =
4363 if crate::pp::pp_cuts(self.layers.len()).is_some() && !crate::pp::pp2_streams_off() {
4364 let rt = crate::pp::PpNRt::get(e)?;
4365 Some(rt.acquire_walk("prime_cache")?)
4366 } else {
4367 None
4368 };
4369 let n_embd = self.cfg.n_embd as usize;
4370 let t = tokens.len();
4371 let legacy_calllocal = std::env::var("MEMRA_PRIME_CALLLOCAL").as_deref() == Ok("1");
4400 let seq_end = if legacy_calllocal {
4401 cache.pos + t
4402 } else {
4403 cache.pos + t + queued_after
4404 };
4405 if overlay.is_none()
4410 && (cache.pos == 0 || step_gemm_prime_suffix_on())
4411 && t >= PRIME_MIN_T
4412 && crate::step_gemm_prime_on()
4413 && self.uses_sliding_gated_moe_program()
4414 {
4415 let n_embd = self.cfg.n_embd as usize;
4416 let base = cache.pos;
4417 let width = crate::cache::PRIME_CHUNK_MAX_TOKENS;
4418 let mut hiddens = e.uninit(t * n_embd)?;
4419 let mut last: Option<(Vec<f32>, CudaSlice<f32>)> = None;
4420 let mut start = 0usize;
4421 while start < t {
4422 let mut end = (start + width).min(t);
4425 if t - end > 0 && t - end < PRIME_MIN_T {
4426 end = t;
4427 }
4428 let mut out = self.step35_prime_cache_batch(
4429 e,
4430 &[&tokens[start..end]],
4431 &mut [cache],
4432 &[seq_end],
4433 )?;
4434 if out.len() != 1 {
4435 return Err("B=1 batched prime returned a non-singleton".into());
4436 }
4437 let (logits, h_seed, hidden) = out.remove(0);
4438 e.copy_into(
4439 &mut hiddens,
4440 start * n_embd,
4441 &hidden,
4442 (end - start) * n_embd,
4443 )?;
4444 last = Some((logits, h_seed));
4445 crate::progress::note_prime_rows(end - start);
4449 start = end;
4450 }
4451 let (logits, h_seed) = last.expect("prime produced no chunk");
4452 eprintln!(
4457 "[gemm-prime] ENGAGED t={t} base={base} seq_end={seq_end} chunks<={width} (GEMM trunk + grouped MoE)"
4458 );
4459 return Ok((logits, h_seed, hiddens));
4460 }
4461 if self.uses_sliding_gated_moe_program() {
4462 eprintln!(
4463 "[gemm-prime] WALK t={t} base={} seq_end={seq_end} (batched prime declined)",
4464 cache.pos
4465 );
4466 }
4467 if overlay.is_none()
4468 && let Some(out) = self.step35_prime_trows(e, tokens, cache)?
4469 {
4470 return Ok(out);
4471 }
4472 assert!(
4476 t >= PRIME_MIN_T,
4477 "prime_cache needs T >= {PRIME_MIN_T} (caller gates)"
4478 );
4479 assert!(
4480 cache.pos + t <= cache.max_ctx,
4481 "prime_cache: prompt exceeds cache max_ctx"
4482 );
4483
4484 if self.is_gemma4_e4b() || self.uses_gemma_program() {
4496 if self.is_gemma4_e4b() {
4497 if overlay.is_some() {
4498 return Err(
4499 "vision embedding overlay is unsupported on gemma4 E4B (PLE prime)".into(),
4500 );
4501 }
4502 return self.gemma4_e4b_prime(e, tokens, cache);
4503 }
4504 return self.gemma4_prime(e, tokens, cache, overlay);
4509 }
4510 if crate::pp::prime_pipe_on()
4511 && crate::pp::prime_pp_on()
4512 && !crate::pp::pp2_streams_off()
4513 && crate::pp::pp_cuts(self.layers.len())
4514 .is_some_and(|fence| matches!(fence.len(), 4 | 5))
4515 {
4516 crate::pp::pp_wave_on()
4517 .map_err(|reason| -> Box<dyn std::error::Error> { reason.into() })?;
4518 }
4519 let ranges = prime_chunk_ranges(t, self.layers.len(), self.gdn_prime_grid_on());
4520 if ranges.len() == 1 {
4558 return self.prime_chunk(e, tokens, cache, seq_end, 0, overlay);
4559 }
4560 if crate::pp::prime_pipe_on() && crate::pp::prime_pp_on() && !crate::pp::pp2_streams_off() {
4564 if let Some(fence) = crate::pp::pp_cuts(self.layers.len()).filter(|f| f.len() == 3) {
4565 if overlay.is_some() {
4566 return Err(
4567 "vision embedding overlay + pipelined PP prime unsupported (v1); \
4568 run the serial prime (single device or MEMRA_PRIME_PIPE=0)"
4569 .into(),
4570 );
4571 }
4572 if crate::pp::pp_multi_stream_same_device() {
4573 return Err(
4574 "prime chunk pipeline refused with 2 stage streams on one device — \
4575 that concurrent-stream placement remains quarantined by the deferred \
4576 pp flake record. Use one device per stage or MEMRA_PRIME_PIPE=0 for \
4577 the serial split."
4578 .into(),
4579 );
4580 }
4581 return self.prime_cache_pp2_pipelined(e, tokens, cache, seq_end, &ranges, &fence);
4582 }
4583 if let Some(fence) =
4584 crate::pp::pp_cuts(self.layers.len()).filter(|f| matches!(f.len(), 4 | 5))
4585 {
4586 let wave_on = crate::pp::pp_wave_on()
4587 .map_err(|reason| -> Box<dyn std::error::Error> { reason.into() })?;
4588 let stages = fence.len() - 1;
4589 if crate::pp::pp_wave_route_enabled(
4590 wave_on,
4591 crate::pp::pp2_overlap(),
4592 stages,
4593 ranges.len(),
4594 ) {
4595 if overlay.is_some() {
4596 return Err(
4597 "vision embedding overlay + pipelined PP prime unsupported; \
4598 run the serial prime (MEMRA_PP_WAVE=0 or MEMRA_PRIME_PIPE=0)"
4599 .into(),
4600 );
4601 }
4602 let rt = crate::pp::PpNRt::get(e)?;
4603 let double_slot = crate::pp::pp2_overlap();
4604 crate::pp::pp_wave_eligibility(
4605 stages,
4606 double_slot,
4607 rt.host_bounce_active(),
4608 rt.repeated_stage_device(),
4609 )
4610 .map_err(|reason| -> Box<dyn std::error::Error> { reason.into() })?;
4611 return self
4612 .prime_cache_ppn_pipelined(e, tokens, cache, seq_end, &ranges, &fence);
4613 }
4614 }
4615 }
4616 let mut hiddens = e.uninit(t * n_embd)?;
4617 let mut last: Option<(Vec<f32>, CudaSlice<f32>)> = None;
4618 for &(start, end) in &ranges {
4619 if let Some(taps) = cache.dflash_taps.as_mut() {
4621 taps.base = start;
4622 }
4623 let (l, hs, x) =
4624 self.prime_chunk(e, &tokens[start..end], cache, seq_end, start, overlay)?;
4625 e.copy_into(&mut hiddens, start * n_embd, &x, (end - start) * n_embd)?;
4626 last = Some((l, hs));
4627 crate::progress::note_prime_rows(end - start);
4631 }
4632 let (logits, h_seed) = last.unwrap();
4633 Ok((logits, h_seed, hiddens))
4634 }
4635
4636 #[allow(clippy::type_complexity)] fn prime_cache_pp2_pipelined(
4642 &self,
4643 e: &Engine,
4644 tokens: &[u32],
4645 cache: &mut Cache,
4646 seq_end: usize,
4647 ranges: &[(usize, usize)],
4648 fence: &[usize],
4649 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4650 debug_assert_eq!(fence.len(), 3);
4651 debug_assert!(ranges.len() >= 2);
4652 let rt = crate::pp::PpNRt::get(e)?;
4653 assert_eq!(
4654 rt.n_stages(),
4655 2,
4656 "prime pipeline requires exactly two PP stages"
4657 );
4658 let n_embd = self.cfg.n_embd as usize;
4659 let t = tokens.len();
4660 let initial_base = cache.pos;
4661 let caller_stream = e.stream();
4662
4663 rt.fence_stages_behind(&caller_stream)?;
4668 let max_payload = ranges.iter().map(|(s, e)| (e - s) * n_embd).max().unwrap();
4669 rt.prepare_overlap_slots(0, max_payload)?;
4670
4671 let mut hiddens = e.uninit(t * n_embd)?;
4672 let mut last: Option<(Vec<f32>, CudaSlice<f32>)> = None;
4673 let mut stage_caches = PrimeCacheStages::new(cache, fence);
4674 let (cache0, cache1) = stage_caches.pp2_parts();
4675 let (first_start, first_end) = ranges[0];
4676 let mut slot = self.prime_pp2_stage0_enqueue(
4677 e,
4678 rt,
4679 &tokens[first_start..first_end],
4680 cache0,
4681 seq_end,
4682 fence,
4683 initial_base + first_start,
4684 true,
4685 )?;
4686 cache0.pos = initial_base + first_end;
4687
4688 for (i, &(start, end)) in ranges.iter().enumerate() {
4689 let base = initial_base + start;
4690 debug_assert_eq!(
4691 cache1.pos, base,
4692 "stage 1 must drain chunks in original position order"
4693 );
4694 let (out, next_slot) = if let Some(&(next_start, next_end)) = ranges.get(i + 1) {
4695 let next_base = initial_base + next_start;
4696 debug_assert_eq!(
4697 cache0.pos, next_base,
4698 "stage 0 must issue chunks in original position order"
4699 );
4700 let cache0_stage = &mut *cache0;
4701 std::thread::scope(|scope| -> Result<_, Box<dyn std::error::Error>> {
4706 let stage0 = scope.spawn(move || -> Result<usize, String> {
4707 let next = self
4708 .prime_pp2_stage0_enqueue(
4709 e,
4710 rt,
4711 &tokens[next_start..next_end],
4712 cache0_stage,
4713 seq_end,
4714 fence,
4715 next_base,
4716 true,
4717 )
4718 .map_err(|err| err.to_string())?;
4719 cache0_stage.pos = initial_base + next_end;
4720 Ok(next)
4721 });
4722 let x = self.prime_pp2_stage1_enqueue(
4723 e,
4724 rt,
4725 slot,
4726 end - start,
4727 cache1,
4728 seq_end,
4729 fence,
4730 base,
4731 true,
4732 )?;
4733 let out = {
4734 rt.bind_stage(1)?;
4735 let _st1 = rt.enter(1);
4736 let e1 = rt.engine(1, e);
4737 self.prime_chunk_epilogue(e1, x, end - start, cache1)?
4738 };
4739 let next = match stage0.join() {
4740 Ok(result) => {
4741 result.map_err(|err| -> Box<dyn std::error::Error> { err.into() })?
4742 }
4743 Err(payload) => std::panic::resume_unwind(payload),
4744 };
4745 Ok((out, Some(next)))
4746 })?
4747 } else {
4748 let x = self.prime_pp2_stage1_enqueue(
4749 e,
4750 rt,
4751 slot,
4752 end - start,
4753 cache1,
4754 seq_end,
4755 fence,
4756 base,
4757 true,
4758 )?;
4759 let out = {
4760 rt.bind_stage(1)?;
4761 let _st1 = rt.enter(1);
4762 let e1 = rt.engine(1, e);
4763 self.prime_chunk_epilogue(e1, x, end - start, cache1)?
4764 };
4765 (out, None)
4766 };
4767
4768 rt.publish_to(1, &caller_stream)?;
4769 e.copy_into(&mut hiddens, start * n_embd, &out.2, (end - start) * n_embd)?;
4770 last = Some((out.0, out.1));
4771 crate::pp::PRIME_SPLIT_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
4772 crate::progress::note_prime_rows(end - start);
4776
4777 if let Some(next) = next_slot {
4778 rt.fence_stages_behind(&caller_stream)?;
4783 slot = next;
4784 }
4785 }
4786
4787 debug_assert_eq!(cache0.pos, initial_base + t);
4788 debug_assert_eq!(cache1.pos, initial_base + t);
4789 let (logits, h_seed) = last.unwrap();
4790 stage_caches.commit();
4791 Ok((logits, h_seed, hiddens))
4792 }
4793
4794 #[allow(clippy::type_complexity)] fn prime_cache_ppn_pipelined(
4804 &self,
4805 e: &Engine,
4806 tokens: &[u32],
4807 cache: &mut Cache,
4808 seq_end: usize,
4809 ranges: &[(usize, usize)],
4810 fence: &[usize],
4811 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
4812 let stages = fence.len().saturating_sub(1);
4813 debug_assert!((3..=4).contains(&stages));
4814 debug_assert!(ranges.len() >= 2);
4815 let rt = crate::pp::PpNRt::get(e)?;
4816 assert_eq!(
4817 rt.n_stages(),
4818 stages,
4819 "prime wavefront PpNRt/fence stage mismatch"
4820 );
4821 let n_embd = self.cfg.n_embd as usize;
4822 let initial_base = cache.pos;
4823 let caller_stream = e.stream();
4824 let primary_context = crate::pp::PrimaryContextRestore::new(e);
4825
4826 rt.fence_stages_behind(&caller_stream)?;
4827 let max_payload = ranges
4828 .iter()
4829 .map(|(start, end)| (end - start) * n_embd)
4830 .max()
4831 .unwrap_or(0);
4832 for boundary in 0..stages - 1 {
4833 rt.prepare_overlap_slots(boundary, max_payload)?;
4834 }
4835
4836 let mut stage_caches = PrimeCacheStages::new(cache, fence);
4837 let waves: Vec<_> = ranges
4838 .iter()
4839 .map(|&(start, end)| PrimePpWave {
4840 start,
4841 end,
4842 tokens: &tokens[start..end],
4843 })
4844 .collect();
4845 let mut forward_senders = Vec::with_capacity(stages - 1);
4846 let mut forward_receivers = Vec::with_capacity(stages - 1);
4847 let mut release_senders = Vec::with_capacity(stages - 1);
4848 let mut release_receivers = Vec::with_capacity(stages - 1);
4849 for _ in 0..stages - 1 {
4850 let (forward_sender, forward_receiver) = std::sync::mpsc::channel();
4851 let (release_sender, release_receiver) = std::sync::mpsc::channel();
4852 forward_senders.push(Some(forward_sender));
4853 forward_receivers.push(Some(forward_receiver));
4854 release_senders.push(Some(release_sender));
4855 release_receivers.push(Some(release_receiver));
4856 }
4857 let mut stage_channels = Vec::with_capacity(stages - 1);
4858 for stage in 0..stages - 1 {
4859 stage_channels.push(Some(PrimePpStageChannels {
4860 incoming: (stage > 0).then(|| forward_receivers[stage - 1].take().unwrap()),
4861 release_upstream: (stage > 0).then(|| release_senders[stage - 1].take().unwrap()),
4862 outgoing: forward_senders[stage].take().unwrap(),
4863 released_downstream: release_receivers[stage].take().unwrap(),
4864 }));
4865 }
4866 let head_incoming = forward_receivers[stages - 2].take().unwrap();
4867 let head_release = release_senders[stages - 2].take().unwrap();
4868 #[allow(clippy::type_complexity)]
4871 let mut results: Vec<Option<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>)>> =
4872 std::iter::repeat_with(|| None).take(waves.len()).collect();
4873 let walk_result = std::thread::scope(|scope| -> Result<(), Box<dyn std::error::Error>> {
4874 let waves_ref = &waves;
4875 let mut handles = Vec::with_capacity(stages - 1);
4876 #[allow(clippy::needless_range_loop)]
4880 for stage in 0..stages - 1 {
4881 let channels = stage_channels[stage].take().unwrap();
4882 let cache_state = &stage_caches.stages()[stage];
4883 handles.push(scope.spawn(move || -> Result<(), String> {
4884 let result = self.prime_ppn_wave_worker(
4885 e,
4886 rt,
4887 waves_ref,
4888 cache_state,
4889 channels.incoming.as_ref(),
4890 channels.release_upstream.as_ref(),
4891 &channels.outgoing,
4892 &channels.released_downstream,
4893 stage,
4894 seq_end,
4895 fence,
4896 initial_base,
4897 );
4898 if let Err(error) = &result {
4899 channels.notify_failure(&error.to_string());
4900 }
4901 result.map_err(|error| error.to_string())
4902 }));
4903 }
4904
4905 let mut head_panic = None;
4906 let head_result = match std::panic::catch_unwind(std::panic::AssertUnwindSafe(
4907 || -> Result<(), Box<dyn std::error::Error>> {
4908 let mut head_cache = stage_caches.stages()[stages - 1]
4909 .lock()
4910 .map_err(|_| "prime PP head cache lock poisoned")?;
4911 for (wave_index, wave) in waves_ref.iter().enumerate() {
4912 let incoming = recv_prime_pp_signal(
4913 &head_incoming,
4914 PrimePpWaveSlot {
4915 wave: wave_index,
4916 slot: 0,
4917 },
4918 false,
4919 "prime PP head input",
4920 )?;
4921 results[wave_index] = Some(self.prime_ppn_wave_final(
4922 e,
4923 rt,
4924 wave,
4925 &mut head_cache,
4926 incoming,
4927 &head_release,
4928 seq_end,
4929 fence,
4930 initial_base,
4931 )?);
4932 crate::progress::note_prime_rows(wave.end - wave.start);
4938 }
4939 Ok(())
4940 },
4941 )) {
4942 Ok(result) => result,
4943 Err(payload) => {
4944 head_panic = Some(payload);
4945 Err("prime PP head-stage host walker panicked".into())
4946 }
4947 };
4948 if let Err(error) = &head_result {
4949 let _ = head_release.send(PrimePpSignal::Error(error.to_string()));
4950 }
4951 let mut first_error = head_result.err().map(|error| error.to_string());
4952 let mut worker_panic = None;
4953 for handle in handles {
4954 match handle.join() {
4955 Ok(Ok(())) => {}
4956 Ok(Err(error)) => {
4957 first_error.get_or_insert(error);
4958 }
4959 Err(payload) => {
4960 if worker_panic.is_none() {
4961 worker_panic = Some(payload);
4962 }
4963 }
4964 }
4965 }
4966 if let Some(payload) = head_panic {
4967 std::panic::resume_unwind(payload);
4968 }
4969 if let Some(payload) = worker_panic {
4970 std::panic::resume_unwind(payload);
4971 }
4972 if let Some(error) = first_error {
4973 return Err(error.into());
4974 }
4975 Ok(())
4976 });
4977 let publish_result = if walk_result.is_ok() {
4978 Some(rt.publish_to(stages - 1, &caller_stream))
4979 } else {
4980 None
4981 };
4982 let restore_result = primary_context.restore();
4983 walk_result?;
4984 if let Some(result) = publish_result {
4985 result?;
4986 }
4987 restore_result?;
4988 let mut hiddens = e.uninit(tokens.len() * n_embd)?;
4989 let mut last: Option<(Vec<f32>, CudaSlice<f32>)> = None;
4990 for (wave_index, (wave, result)) in waves.iter().zip(results).enumerate() {
4991 debug_assert_eq!((wave.start, wave.end), ranges[wave_index]);
4992 let out = result.ok_or("prime PP wavefront completed without a head-stage result")?;
4993 e.copy_into(
4994 &mut hiddens,
4995 wave.start * n_embd,
4996 &out.2,
4997 (wave.end - wave.start) * n_embd,
4998 )?;
4999 last = Some((out.0, out.1));
5000 crate::pp::PRIME_SPLIT_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
5001 }
5006 stage_caches.commit();
5007 drop(stage_caches);
5008
5009 static LOGGED: std::sync::Once = std::sync::Once::new();
5010 LOGGED.call_once(|| {
5011 eprintln!(
5012 "[pp-wave] PP{stages} prime wavefront engaged: microchunks={} \
5013 (experimental, MEMRA_PP_WAVE=1)",
5014 ranges.len(),
5015 );
5016 });
5017 let (logits, h_seed) = last.expect("prime PP wavefront produced no microchunk");
5018 crate::pp::record_pp_wave_tick();
5019 Ok((logits, h_seed, hiddens))
5020 }
5021
5022 #[allow(clippy::too_many_arguments)]
5023 fn prime_ppn_wave_worker(
5024 &self,
5025 e: &Engine,
5026 rt: &crate::pp::PpNRt,
5027 waves: &[PrimePpWave<'_>],
5028 cache: &std::sync::Mutex<Cache>,
5029 incoming: Option<&std::sync::mpsc::Receiver<PrimePpSignal>>,
5030 release_upstream: Option<&std::sync::mpsc::Sender<PrimePpSignal>>,
5031 outgoing: &std::sync::mpsc::Sender<PrimePpSignal>,
5032 released_downstream: &std::sync::mpsc::Receiver<PrimePpSignal>,
5033 stage: usize,
5034 seq_end: usize,
5035 fence: &[usize],
5036 initial_base: usize,
5037 ) -> Result<(), Box<dyn std::error::Error>> {
5038 debug_assert_eq!(incoming.is_some(), stage > 0);
5039 debug_assert_eq!(release_upstream.is_some(), stage > 0);
5040 let mut cache = cache
5041 .lock()
5042 .map_err(|_| "prime PP cache stage lock poisoned")?;
5043 let mut credits = PrimePpWaveCredits::default();
5044 for (wave_index, wave) in waves.iter().enumerate() {
5045 let incoming = match incoming {
5046 Some(receiver) => Some(recv_prime_pp_signal(
5047 receiver,
5048 PrimePpWaveSlot {
5049 wave: wave_index,
5050 slot: 0,
5051 },
5052 false,
5053 "prime PP stage input",
5054 )?),
5055 None => None,
5056 };
5057 let sent = self.prime_ppn_wave_stage(
5058 e,
5059 rt,
5060 wave,
5061 &mut cache,
5062 stage,
5063 incoming,
5064 release_upstream,
5065 &mut credits,
5066 released_downstream,
5067 seq_end,
5068 fence,
5069 initial_base,
5070 )?;
5071 send_prime_pp_signal(outgoing, PrimePpSignal::Slot(sent), "prime PP stage output")?;
5072 }
5073 while let Some(expected) = credits.pending.front().copied() {
5074 let released = recv_prime_pp_signal(
5075 released_downstream,
5076 expected,
5077 true,
5078 "prime PP final slot release",
5079 )?;
5080 credits.record_release(released)?;
5081 }
5082 Ok(())
5083 }
5084
5085 #[allow(clippy::too_many_arguments)]
5086 fn prime_ppn_wave_stage(
5087 &self,
5088 e: &Engine,
5089 rt: &crate::pp::PpNRt,
5090 wave: &PrimePpWave<'_>,
5091 cache: &mut Cache,
5092 stage: usize,
5093 incoming: Option<PrimePpWaveSlot>,
5094 release_upstream: Option<&std::sync::mpsc::Sender<PrimePpSignal>>,
5095 credits: &mut PrimePpWaveCredits,
5096 released_downstream: &std::sync::mpsc::Receiver<PrimePpSignal>,
5097 seq_end: usize,
5098 fence: &[usize],
5099 initial_base: usize,
5100 ) -> Result<PrimePpWaveSlot, Box<dyn std::error::Error>> {
5101 debug_assert!(stage + 1 < fence.len() - 1);
5102 let t = wave.end - wave.start;
5103 let base = initial_base + wave.start;
5104 debug_assert_eq!(cache.pos, base, "prime PP stage advanced out of order");
5105 let n_embd = self.cfg.n_embd as usize;
5106 let payload = t * n_embd;
5107 let positions: Vec<i32> = (base as i32..(base + t) as i32).collect();
5108 rt.bind_stage(stage)?;
5109 let _stage = rt.enter(stage);
5110 let engine = rt.engine(stage, e);
5111 let positions_d = engine.htod_i32(&positions)?;
5112 let x = if stage == 0 {
5113 debug_assert!(incoming.is_none());
5114 self.embed(engine, wave.tokens)?
5115 } else {
5116 let incoming = incoming.ok_or("prime PP stage has no incoming boundary slot")?;
5117 let x = rt.rx(stage - 1, incoming.slot, payload)?;
5118 send_prime_pp_signal(
5119 release_upstream.ok_or("prime PP stage has no upstream release channel")?,
5120 PrimePpSignal::Slot(incoming),
5121 "prime PP upstream slot release",
5122 )?;
5123 x
5124 };
5125 let x = {
5126 let _wave_cell = crate::pp::enter_pp_wave_cell();
5127 let _overlap = crate::pp::enter_prime_pipe_stage();
5128 self.prime_layers(
5129 engine,
5130 x,
5131 fence[stage],
5132 fence[stage + 1],
5133 &positions_d,
5134 t,
5135 base,
5136 cache,
5137 seq_end,
5138 )?
5139 };
5140 if let Some(expected) = credits.release_required() {
5141 let released =
5142 recv_prime_pp_signal(released_downstream, expected, true, "prime PP slot credit")?;
5143 credits.record_release(released)?;
5144 }
5145 let sent = PrimePpWaveSlot {
5146 wave: credits.next_wave,
5147 slot: rt.tx_pipelined(stage, &x, payload)?,
5148 };
5149 credits.record_send(sent)?;
5150 cache.pos = base + t;
5151 Ok(sent)
5152 }
5153
5154 #[allow(clippy::too_many_arguments)]
5155 #[allow(clippy::type_complexity)] fn prime_ppn_wave_final(
5157 &self,
5158 e: &Engine,
5159 rt: &crate::pp::PpNRt,
5160 wave: &PrimePpWave<'_>,
5161 cache: &mut Cache,
5162 incoming: PrimePpWaveSlot,
5163 release_upstream: &std::sync::mpsc::Sender<PrimePpSignal>,
5164 seq_end: usize,
5165 fence: &[usize],
5166 initial_base: usize,
5167 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5168 let stage = fence.len() - 2;
5169 let t = wave.end - wave.start;
5170 let base = initial_base + wave.start;
5171 debug_assert_eq!(cache.pos, base, "prime PP head stage advanced out of order");
5172 let n_embd = self.cfg.n_embd as usize;
5173 let payload = t * n_embd;
5174 let positions: Vec<i32> = (base as i32..(base + t) as i32).collect();
5175 rt.bind_stage(stage)?;
5176 let _stage = rt.enter(stage);
5177 let engine = rt.engine(stage, e);
5178 let positions_d = engine.htod_i32(&positions)?;
5179 let x = rt.rx(stage - 1, incoming.slot, payload)?;
5180 send_prime_pp_signal(
5181 release_upstream,
5182 PrimePpSignal::Slot(incoming),
5183 "prime PP head slot release",
5184 )?;
5185 let _wave_cell = crate::pp::enter_pp_wave_cell();
5186 let _overlap = crate::pp::enter_prime_pipe_stage();
5187 let x = self.prime_layers(
5188 engine,
5189 x,
5190 fence[stage],
5191 fence[stage + 1],
5192 &positions_d,
5193 t,
5194 base,
5195 cache,
5196 seq_end,
5197 )?;
5198 self.prime_chunk_epilogue(engine, x, t, cache)
5199 }
5200
5201 fn gdn_hk(e: &Engine, t: usize, num_v: usize, num_k: usize) -> usize {
5208 if Engine::gdn_db_on()
5209 && Engine::gdn_chunked_enabled()
5210 && t >= 16
5211 && e.gdn_mma_enabled(Engine::gdn_chunk_size())
5212 && num_k * 2 == num_v
5213 {
5214 num_k
5215 } else {
5216 num_v
5217 }
5218 }
5219
5220 fn f16out_on(e: &Engine, t: usize) -> bool {
5225 crate::f16_ffi::pp_f16_enabled()
5226 && t >= 16
5227 && !e.verify_exact_on()
5228 && std::env::var("MEMRA_F16OUT").as_deref() != Ok("0")
5229 }
5230
5231 pub fn prime_slabs_get(
5239 &self,
5240 e: &Engine,
5241 t: usize,
5242 n_embd: usize,
5243 n_ff_max: usize,
5244 ) -> Result<std::sync::Arc<std::sync::Mutex<PrimeSlabs>>, Box<dyn std::error::Error>> {
5245 let mut slabs = self.prime_slabs.lock().unwrap();
5246 let dev = e.ctx().ordinal();
5247 let need_new = match slabs.get(&dev) {
5248 None => true,
5249 Some(sl) => sl.lock().unwrap().t_cap < t,
5250 };
5251 if need_new {
5252 slabs.insert(
5253 dev,
5254 std::sync::Arc::new(std::sync::Mutex::new(PrimeSlabs {
5255 t_cap: t,
5256 h: e.uninit(t * n_embd)?,
5257 x1: e.uninit(t * n_embd)?,
5258 z: e.uninit(t * n_embd)?,
5259 act: e.uninit(t * n_ff_max)?,
5260 xa: e.uninit(t * n_embd)?,
5261 xb: e.uninit(t * n_embd)?,
5262 h16: e.alloc_u8_uninit(t * n_embd * 2)?,
5263 z16: e.alloc_u8_uninit(t * n_embd * 2)?,
5264 gate: e.uninit(t * n_ff_max)?,
5265 up: e.uninit(t * n_ff_max)?,
5266 ffn_out: e.uninit(t * n_embd)?,
5267 seg_glue: Vec::new(),
5268 mixed: e.uninit(t * n_embd)?,
5269 seg_mid: Vec::new(),
5270 seg_t: 0,
5271 })),
5272 );
5273 }
5274 Ok(slabs.get(&dev).expect("prime slab inserted").clone())
5275 }
5276
5277 #[allow(clippy::type_complexity)] fn prime_chunk(
5282 &self,
5283 e: &Engine,
5284 tokens: &[u32],
5285 cache: &mut Cache,
5286 seq_end: usize,
5287 chunk_off: usize,
5288 overlay: Option<&crate::vision::EmbedOverlay>,
5289 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5290 if crate::pp::pp_host_bounce_active()
5291 && (self.uses_gemma_program() || !crate::pp::prime_pp_on())
5292 {
5293 return Err(
5294 "prime_chunk: refused with MEMRA_PP_HOST_BOUNCE=1 because this configuration \
5295 has no active prime stage split and would peer-read remote weights; keep \
5296 MEMRA_PRIME_PP enabled and use a PP-prime-supported model"
5297 .into(),
5298 );
5299 }
5300 if !self.uses_gemma_program()
5309 && !crate::pp::pp2_streams_off()
5310 && crate::pp::prime_pp_on()
5311 && let Some(fence) = crate::pp::pp_cuts(self.layers.len())
5312 {
5313 if overlay.is_some() {
5314 return Err("vision embedding overlay + PP prime unsupported (v1); \
5315 run single-device or MEMRA_PRIME_PP=0"
5316 .into());
5317 }
5318 return self.prime_chunk_ppn(e, tokens, cache, seq_end, &fence);
5319 }
5320 if crate::pp::pp_host_bounce_active() {
5321 return Err(
5322 "prime_chunk: MEMRA_PP_HOST_BOUNCE=1 found no valid prime stage split; \
5323 refusing an unsplit remote-weight walk"
5324 .into(),
5325 );
5326 }
5327 let t = tokens.len();
5328 let base = cache.pos;
5329 debug_assert!(
5330 seq_end >= base + t,
5331 "prime_chunk: seq_end must cover this chunk"
5332 );
5333 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
5334 let pos_d = e.htod_i32(&pos)?;
5335
5336 let mut x_embed = self.embed(e, tokens)?; if let Some(ov) = overlay {
5338 ov.splice_into(e, &mut x_embed, chunk_off, t, self.cfg.n_embd as usize)?;
5347 }
5348 let x = self.prime_layers(
5349 e,
5350 x_embed,
5351 0,
5352 self.layers.len(),
5353 &pos_d,
5354 t,
5355 base,
5356 cache,
5357 seq_end,
5358 )?;
5359 self.prime_chunk_epilogue(e, x, t, cache)
5360 }
5361
5362 #[allow(clippy::too_many_arguments)]
5378 fn prime_layers(
5379 &self,
5380 e: &Engine,
5381 x_in: CudaSlice<f32>,
5382 lo: usize,
5383 hi: usize,
5384 pos_d: &CudaSlice<i32>,
5385 t: usize,
5386 base: usize,
5387 cache: &mut Cache,
5388 seq_end: usize,
5389 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5390 let cfg = &self.cfg;
5391 let n_embd = cfg.n_embd as usize;
5392 let eps = cfg.rms_eps;
5393 let f16fuse = crate::f16_ffi::pp_f16_enabled() && t >= 16;
5397 let n_ff_max = self
5404 .layers
5405 .iter()
5406 .map(|l| match &l.ffn {
5407 crate::hybrid::Ffn::Dense { ffn_gate, .. } => ffn_gate.out_features(),
5408 _ => n_embd,
5409 })
5410 .max()
5411 .unwrap_or(n_embd)
5412 .max(n_embd);
5413 let use_slabs = std::env::var("MEMRA_PRIME_SLABS").as_deref() != Ok("0");
5414 let slab = if use_slabs {
5415 Some(self.prime_slabs_get(e, t, n_embd, n_ff_max)?)
5416 } else {
5417 None
5418 };
5419 let mut slab_guard = slab.as_ref().map(|sl| sl.lock().unwrap());
5420 let mut x_own; type SlabRefs<'a> = (
5422 &'a mut CudaSlice<f32>,
5423 &'a mut CudaSlice<f32>,
5424 &'a mut CudaSlice<f32>,
5425 &'a mut CudaSlice<f32>,
5426 &'a mut CudaSlice<u8>,
5427 &'a mut CudaSlice<u8>,
5428 &'a mut CudaSlice<f32>,
5429 &'a mut CudaSlice<f32>,
5430 &'a mut CudaSlice<f32>,
5431 );
5432 let (mut x_cur, mut x_nxt, sl): (
5433 &mut CudaSlice<f32>,
5434 &mut CudaSlice<f32>,
5435 Option<SlabRefs>,
5436 );
5437 #[allow(clippy::type_complexity)]
5438 let mut seg: Option<(
5440 &mut Vec<Option<cudarc::driver::CudaGraph>>,
5441 &mut Vec<Option<cudarc::driver::CudaGraph>>,
5442 &mut CudaSlice<f32>,
5443 &mut usize,
5444 )> = None;
5445 let mut x_own2;
5446 match slab_guard.as_mut() {
5447 Some(g) => {
5448 let slabs = &mut **g;
5449 e.copy_into(&mut slabs.xa, 0, &x_in, t * n_embd)?;
5450 let PrimeSlabs {
5451 xa,
5452 xb,
5453 h,
5454 x1,
5455 z,
5456 act,
5457 h16,
5458 z16,
5459 gate,
5460 up,
5461 ffn_out,
5462 seg_glue,
5463 mixed,
5464 seg_mid,
5465 seg_t,
5466 ..
5467 } = slabs;
5468 x_cur = xa;
5469 x_nxt = xb;
5470 seg = Some((seg_glue, seg_mid, mixed, seg_t));
5471 sl = Some((h, x1, z, act, h16, z16, gate, up, ffn_out));
5472 }
5473 None => {
5474 x_own = x_in;
5475 x_own2 = e.uninit(t * n_embd)?;
5476 x_cur = &mut x_own;
5477 x_nxt = &mut x_own2;
5478 sl = None;
5479 }
5480 }
5481 let mut alloc_h;
5482 let mut alloc_x1;
5483 let mut alloc_z;
5484 let mut alloc_act;
5485 let mut alloc_h16;
5486 let mut alloc_z16;
5487 let mut alloc_gate;
5488 let mut alloc_up;
5489 let mut alloc_fo;
5490 let (h, x1, z, act): (
5491 &mut CudaSlice<f32>,
5492 &mut CudaSlice<f32>,
5493 &mut CudaSlice<f32>,
5494 &mut CudaSlice<f32>,
5495 );
5496 let (h16, z16): (&mut CudaSlice<u8>, &mut CudaSlice<u8>);
5497 let (sl_gate, sl_up, sl_fo): (
5498 &mut CudaSlice<f32>,
5499 &mut CudaSlice<f32>,
5500 &mut CudaSlice<f32>,
5501 );
5502 match sl {
5503 Some((a, b, c, d, e16, f16b, g, u, fo)) => {
5504 h = a;
5505 x1 = b;
5506 z = c;
5507 act = d;
5508 h16 = e16;
5509 z16 = f16b;
5510 sl_gate = g;
5511 sl_up = u;
5512 sl_fo = fo;
5513 }
5514 None => {
5515 alloc_h = e.uninit(t * n_embd)?;
5516 alloc_x1 = e.uninit(t * n_embd)?;
5517 alloc_z = e.uninit(t * n_embd)?;
5518 alloc_act = e.uninit(t * n_ff_max)?;
5519 alloc_h16 = e.alloc_u8_uninit(t * n_embd * 2)?;
5520 alloc_z16 = e.alloc_u8_uninit(t * n_embd * 2)?;
5521 alloc_gate = e.uninit(t * n_ff_max)?;
5522 alloc_up = e.uninit(t * n_ff_max)?;
5523 alloc_fo = e.uninit(t * n_embd)?;
5524 h = &mut alloc_h;
5525 x1 = &mut alloc_x1;
5526 z = &mut alloc_z;
5527 act = &mut alloc_act;
5528 h16 = &mut alloc_h16;
5529 z16 = &mut alloc_z16;
5530 sl_gate = &mut alloc_gate;
5531 sl_up = &mut alloc_up;
5532 sl_fo = &mut alloc_fo;
5533 }
5534 }
5535 let n_layers = self.layers.len();
5540 let use_seg = f16fuse
5550 && seg.is_some()
5551 && !self.uses_sliding_gated_moe_program()
5552 && lo == 0
5553 && hi == n_layers
5554 && std::env::var("MEMRA_PRIME_SEG").as_deref() == Ok("1")
5555 && {
5561 let ok = crate::spec::graph_launch_headroom_ok(e);
5562 if !ok {
5563 static NOTED: std::sync::Once = std::sync::Once::new();
5564 NOTED.call_once(|| crate::spec::graph_replay_suspended_note("prime-seg"));
5565 }
5566 ok
5567 };
5568 if let Some((sg, sm, _, st)) = seg.as_mut()
5569 && **st != t
5570 {
5571 sg.clear();
5572 sg.extend((0..n_layers).map(|_| None));
5573 sm.clear();
5574 sm.extend((0..n_layers).map(|_| None));
5575 **st = t;
5576 }
5577 {
5578 let layer_lo = &self.layers[lo];
5579 if f16fuse {
5580 e.rms_norm_f16out(
5581 x_cur,
5582 layer_lo.attn_norm.float_data(),
5583 h,
5584 h16,
5585 n_embd,
5586 t,
5587 eps,
5588 )?;
5589 } else {
5590 e.rms_norm(x_cur, layer_lo.attn_norm.float_data(), h, n_embd, t, eps)?;
5591 }
5592 }
5593 let anat = Self::prime_anatomy_on();
5594 let mut anat_last = if anat {
5595 e.stream().synchronize()?;
5596 Some(std::time::Instant::now())
5597 } else {
5598 None
5599 };
5600 macro_rules! anat_mark {
5602 ($slot:expr) => {
5603 if let Some(ts) = anat_last.as_mut() {
5604 e.stream().synchronize()?;
5605 Self::prime_anatomy_slots()[$slot].fetch_add(
5606 ts.elapsed().as_nanos() as u64,
5607 std::sync::atomic::Ordering::Relaxed,
5608 );
5609 *ts = std::time::Instant::now();
5610 }
5611 };
5612 }
5613 for il in lo..hi {
5614 let layer = &self.layers[il];
5615 let hx16 = if f16fuse { Some(&*h16) } else { None };
5616 if use_seg {
5617 let (pre, pre16, w_out) = match &layer.mixer {
5620 Mixer::Full(fa) => {
5621 let g3 = match hx16 {
5622 Some(xh) => e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], h, xh, t)?,
5623 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
5624 };
5625 let (pre, pre16) =
5626 self.full_attn_prime_core_inner(e, fa, g3, pos_d, t, cache, il)?;
5627 (pre, pre16, &fa.wo)
5628 }
5629 Mixer::Mla(_) => crate::hybrid::mla_path_unimplemented("core-split prime"),
5630 Mixer::Kda(_) => {
5631 crate::hybrid::kda_path_unimplemented("core-split captured prime")
5632 }
5633 Mixer::Linear(la) => {
5634 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
5635 let g4 = match hx16 {
5636 Some(xh) => e.matmul_group_xh(&ws, h, xh, t)?,
5637 None => e.matmul_group(&ws, h, t)?,
5638 };
5639 let (pre, pre16) =
5640 self.linear_attn_prime_core_pad_inner(e, la, g4, t, cache, il, None)?;
5641 (pre, pre16, &la.ssm_out)
5642 }
5643 };
5644 {
5645 let (_, sm, mslab, _) = seg.as_mut().unwrap();
5646 let pre_n = pre.len() / t;
5647 let xh_pre = match pre16 {
5648 Some(x) => x,
5649 None => e.f16_act(&pre, t * pre_n, pre_n)?,
5650 };
5651 if !e.try_f16_gemm_pre_into(w_out, &xh_pre, t, mslab)? {
5652 let y = e.matmul(w_out, &pre, t)?;
5653 e.copy_into(mslab, 0, &y, t * n_embd)?;
5654 }
5655 if sm[il].is_none() {
5656 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
5657 let w_post = layer.post_attn_norm.float_data();
5658 e.stream().synchronize()?;
5659 e.stream()
5660 .begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
5661 let r = (|| -> Result<(), Box<dyn std::error::Error>> {
5662 e.add(x_cur, mslab, x1, t * n_embd)?;
5663 e.rms_norm_f16out(x1, w_post, z, z16, n_embd, t, eps)?;
5664 Ok(())
5665 })();
5666 let g = e.stream().end_capture(
5667 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH);
5668 r?;
5669 sm[il] = Some(g?.ok_or("S-mid capture produced no graph")?);
5670 }
5671 sm[il].as_ref().unwrap().launch()?;
5672 }
5673 } else {
5674 let mixed = match &layer.mixer {
5675 Mixer::Full(fa) => {
5676 let y =
5677 self.full_attn_prime(e, fa, h, hx16, pos_d, t, cache, il, seq_end)?;
5678 anat_mark!(0);
5679 y
5680 }
5681 Mixer::Linear(la) => {
5682 let y = self.linear_attn_prime(e, la, h, hx16, t, cache, il)?;
5683 anat_mark!(1);
5684 y
5685 }
5686 Mixer::Mla(mla) => self.mla_attn_cached(e, mla, h, pos_d, t, il, cache)?,
5687 Mixer::Kda(la) => {
5688 let y = crate::kda::kda_prime_cached(e, la, h, t, eps, cache, il)?;
5689 anat_mark!(1);
5690 y
5691 }
5692 };
5693 if f16fuse {
5694 e.add_rms_norm_f16out(
5697 x_cur,
5698 &mixed,
5699 layer.post_attn_norm.float_data(),
5700 x1,
5701 z,
5702 z16,
5703 n_embd,
5704 t,
5705 eps,
5706 )?;
5707 } else {
5708 e.add(x_cur, &mixed, x1, t * n_embd)?;
5709 e.rms_norm(x1, layer.post_attn_norm.float_data(), z, n_embd, t, eps)?;
5710 }
5711 anat_mark!(4);
5712 }
5713 let zx16 = if f16fuse { Some(&*z16) } else { None };
5714 match &layer.ffn {
5715 crate::hybrid::Ffn::Dense {
5716 ffn_gate,
5717 ffn_up,
5718 ffn_down,
5719 } => {
5720 let n_ff = ffn_gate.out_features();
5721 let mut into_ok = false;
5724 if let Some(xh) = zx16 {
5725 into_ok = e.try_f16_gemm_pre_into(ffn_gate, xh, t, sl_gate)?
5726 && e.try_f16_gemm_pre_into(ffn_up, xh, t, sl_up)?;
5727 }
5728 if !into_ok {
5729 let mut g2 = match zx16 {
5730 Some(xh) => e.matmul_group_xh(&[ffn_gate, ffn_up], z, xh, t)?,
5731 None => e.matmul_group(&[ffn_gate, ffn_up], z, t)?,
5732 };
5733 let up_y = g2.pop().unwrap();
5734 let gate_y = g2.pop().unwrap();
5735 e.copy_into(sl_gate, 0, &gate_y, t * n_ff)?;
5736 e.copy_into(sl_up, 0, &up_y, t * n_ff)?;
5737 }
5738 let d_lim = self.cfg.clamp_shexp_at(il as u32);
5743 let act16 = if Self::f16out_on(e, t) && self.cfg.m3.is_none() && d_lim.is_none()
5744 {
5745 let mut a16 = e.alloc_u8_uninit(t * n_ff * 2)?;
5746 e.silu_mul_f16out(sl_gate, sl_up, act, &mut a16, t * n_ff)?;
5747 Some(a16)
5748 } else {
5749 Self::ffn_act_lim(
5750 e,
5751 &self.cfg,
5752 sl_gate,
5753 sl_up,
5754 1.0,
5755 1.0,
5756 d_lim,
5757 act,
5758 t * n_ff,
5759 )?;
5760 None
5761 };
5762 let xh_act = match act16 {
5764 Some(x) => x,
5765 None => e.f16_act(act, t * n_ff, n_ff)?,
5766 };
5767 if !e.try_f16_gemm_pre_into(ffn_down, &xh_act, t, sl_fo)? {
5768 let y = e.matmul(ffn_down, &*act, t)?;
5769 e.copy_into(sl_fo, 0, &y, t * n_embd)?;
5770 }
5771 }
5772 crate::hybrid::Ffn::Moe(m) => {
5773 let y = self.moe_ffn_il_prefill(e, m, z, t, il as u16)?;
5774 e.copy_into(sl_fo, 0, &y, t * n_embd)?;
5775 anat_mark!(2);
5776 }
5777 }
5778 if let (crate::hybrid::Ffn::Dense { .. }, true) = (&layer.ffn, anat) {
5779 anat_mark!(3);
5780 }
5781 if use_seg && il + 1 < hi {
5782 let w_next = self.layers[il + 1].attn_norm.float_data();
5784 let (sg, _, _, _) = seg.as_mut().unwrap();
5785 if sg[il].is_none() {
5786 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
5787 e.stream().synchronize()?;
5788 e.stream()
5789 .begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
5790 let r = (|| -> Result<(), Box<dyn std::error::Error>> {
5791 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
5792 e.rms_norm_f16out(x_nxt, w_next, h, h16, n_embd, t, eps)?;
5793 Ok(())
5794 })();
5795 let g = e.stream().end_capture(
5796 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH,
5797 );
5798 r?;
5799 sg[il] = Some(g?.ok_or("S-glue capture produced no graph")?);
5800 }
5801 sg[il].as_ref().unwrap().launch()?;
5802 } else {
5803 if il + 1 < hi {
5804 let w_next = self.layers[il + 1].attn_norm.float_data();
5805 if f16fuse {
5806 e.add_rms_norm_f16out(x1, sl_fo, w_next, x_nxt, h, h16, n_embd, t, eps)?;
5807 } else {
5808 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
5809 e.rms_norm(x_nxt, w_next, h, n_embd, t, eps)?;
5810 }
5811 } else {
5812 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
5813 }
5814 }
5815 anat_mark!(4);
5816 if let Some(path) = Self::prime_trace_path() {
5822 let row = base + t - 1;
5823 let host = e.dtoh(x_nxt)?;
5824 let last = &host[(t - 1) * n_embd..t * n_embd];
5825 use std::io::Write as _;
5826 let mut f = std::fs::OpenOptions::new()
5827 .create(true)
5828 .append(true)
5829 .open(path)?;
5830 let mut h64: u64 = 0xcbf29ce484222325;
5831 for v in last {
5832 h64 ^= v.to_bits() as u64;
5833 h64 = h64.wrapping_mul(0x100000001b3);
5834 }
5835 writeln!(
5836 f,
5837 "{{\"pos\":{row},\"layer\":{il},\"t\":{t},\"base\":{base},\
5838 \"hash\":\"{h64:016x}\",\"v0\":{:.9e},\"v1\":{:.9e},\"v2\":{:.9e}}}",
5839 last[0], last[1], last[2]
5840 )?;
5841 }
5842 self.dflash_tap(e, cache, il, x_nxt, t)?;
5845 std::mem::swap(&mut x_cur, &mut x_nxt);
5846 }
5847 if anat {
5848 let s = Self::prime_anatomy_slots();
5849 let ms = |i: usize| s[i].load(std::sync::atomic::Ordering::Relaxed) as f64 / 1.0e6;
5850 eprintln!(
5851 "[prime-anatomy] cumulative ms: attn_full={:.1} gdn_linear={:.1} moe={:.1} \
5852 dense={:.1} norms_adds={:.1} (t={t}, layers {lo}..{hi})",
5853 ms(0),
5854 ms(1),
5855 ms(2),
5856 ms(3),
5857 ms(4)
5858 );
5859 }
5860 let mut x = e.uninit(t * n_embd)?;
5862 e.copy_into(&mut x, 0, x_cur, t * n_embd)?;
5863 drop(slab_guard);
5864 Ok(x)
5865 }
5866
5867 #[allow(clippy::type_complexity)] fn prime_chunk_epilogue(
5873 &self,
5874 e: &Engine,
5875 x: CudaSlice<f32>,
5876 t: usize,
5877 cache: &mut Cache,
5878 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5879 let n_embd = self.cfg.n_embd as usize;
5880 let eps = self.cfg.rms_eps;
5881 let mut h_seed = e.uninit(n_embd)?;
5885 if !crate::spec::spec_hpost() {
5886 e.copy_view_into(
5887 &mut h_seed,
5888 0,
5889 &x.slice((t - 1) * n_embd..t * n_embd),
5890 n_embd,
5891 )?;
5892 }
5893 let mut hn = e.uninit(t * n_embd)?;
5895 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
5896 if crate::spec::spec_hpost() {
5897 e.copy_view_into(
5898 &mut h_seed,
5899 0,
5900 &hn.slice((t - 1) * n_embd..t * n_embd),
5901 n_embd,
5902 )?;
5903 }
5904 let last = e.view(&hn, t * n_embd);
5905 let last_row = last.slice((t - 1) * n_embd..t * n_embd);
5906 let mut hlast = e.uninit(n_embd)?;
5907 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
5908 let logits = e.matmul(&self.output, &hlast, 1)?;
5909 cache.pos += t;
5910 Ok((
5913 e.dtoh(&logits)?,
5914 h_seed,
5915 if crate::spec::spec_hpost() { hn } else { x },
5916 ))
5917 }
5918
5919 pub fn hidden_postnorm_row(
5925 &self,
5926 e: &Engine,
5927 hiddens: &CudaSlice<f32>,
5928 row: usize,
5929 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
5930 let n_embd = self.cfg.n_embd as usize;
5931 let mut x1 = e.uninit(n_embd)?;
5932 e.copy_view_into(
5933 &mut x1,
5934 0,
5935 &hiddens.slice(row * n_embd..(row + 1) * n_embd),
5936 n_embd,
5937 )?;
5938 if crate::spec::spec_hpost() {
5939 return e.dtoh(&x1);
5940 }
5941 let mut hn = e.uninit(n_embd)?;
5942 e.rms_norm(
5943 &x1,
5944 self.output_norm.float_data(),
5945 &mut hn,
5946 n_embd,
5947 1,
5948 self.cfg.rms_eps,
5949 )?;
5950 e.dtoh(&hn)
5951 }
5952
5953 #[allow(clippy::type_complexity)] fn prime_chunk_ppn(
5978 &self,
5979 e: &Engine,
5980 tokens: &[u32],
5981 cache: &mut Cache,
5982 seq_end: usize,
5983 fence: &[usize],
5984 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5985 let rt = crate::pp::PpNRt::get(e)?;
5986 let n_st = fence.len() - 1;
5987 assert_eq!(
5988 rt.n_stages(),
5989 n_st,
5990 "PpNRt stage count {} != fence stages {n_st}",
5991 rt.n_stages()
5992 );
5993 let n_embd = self.cfg.n_embd as usize;
5994 let t = tokens.len();
5995 let base = cache.pos;
5996 debug_assert!(
5997 seq_end >= base + t,
5998 "prime_chunk_ppn: seq_end must cover this chunk"
5999 );
6000 let payload = t * n_embd;
6001 let caller_stream = e.stream();
6005 rt.fence_stages_behind(&caller_stream)?;
6006
6007 if n_st == 2 {
6008 let slot =
6009 self.prime_pp2_stage0_enqueue(e, rt, tokens, cache, seq_end, fence, base, false)?;
6010 let x =
6011 self.prime_pp2_stage1_enqueue(e, rt, slot, t, cache, seq_end, fence, base, false)?;
6012 let out = {
6013 rt.bind_stage(1)?;
6014 let _st1 = rt.enter(1);
6015 let e1 = rt.engine(1, e);
6016 self.prime_chunk_epilogue(e1, x, t, cache)?
6017 };
6018 rt.publish_to(1, &caller_stream)?;
6019 crate::pp::PRIME_SPLIT_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
6020 return Ok(out);
6026 }
6027
6028 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
6029
6030 let mut slot = {
6032 let _st0 = rt.enter(0);
6033 let e0 = rt.engine(0, e);
6034 let pos_d = e0.htod_i32(&pos)?;
6035 let x = self.embed(e0, tokens)?;
6036 let x =
6037 self.prime_layers(e0, x, fence[0], fence[1], &pos_d, t, base, cache, seq_end)?;
6038 rt.tx(0, &x, payload)?
6039 };
6041
6042 for s in 1..n_st - 1 {
6044 let _st = rt.enter(s);
6045 let es = rt.engine(s, e);
6046 let pos_d = es.htod_i32(&pos)?;
6047 let x = rt.rx(s - 1, slot, payload)?;
6048 let x = self.prime_layers(
6049 es,
6050 x,
6051 fence[s],
6052 fence[s + 1],
6053 &pos_d,
6054 t,
6055 base,
6056 cache,
6057 seq_end,
6058 )?;
6059 slot = rt.tx(s, &x, payload)?;
6060 }
6061
6062 let _stl = rt.enter(n_st - 1);
6064 let el = rt.engine(n_st - 1, e);
6065 let pos_d = el.htod_i32(&pos)?;
6066 let x = rt.rx(n_st - 2, slot, payload)?;
6067 let x = self.prime_layers(
6068 el,
6069 x,
6070 fence[n_st - 1],
6071 fence[n_st],
6072 &pos_d,
6073 t,
6074 base,
6075 cache,
6076 seq_end,
6077 )?;
6078 let out = self.prime_chunk_epilogue(el, x, t, cache)?;
6079 rt.publish_to(n_st - 1, &caller_stream)?;
6085 crate::pp::PRIME_SPLIT_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
6086 Ok(out)
6087 }
6088
6089 #[allow(clippy::too_many_arguments)] fn prime_pp2_stage0_enqueue(
6091 &self,
6092 e: &Engine,
6093 rt: &crate::pp::PpNRt,
6094 tokens: &[u32],
6095 cache: &mut Cache,
6096 seq_end: usize,
6097 fence: &[usize],
6098 base: usize,
6099 pipelined: bool,
6100 ) -> Result<usize, Box<dyn std::error::Error>> {
6101 let t = tokens.len();
6102 let n_embd = self.cfg.n_embd as usize;
6103 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
6104 rt.bind_stage(0)?;
6105 let _st0 = rt.enter(0);
6106 let e0 = rt.engine(0, e);
6107 let pos_d = e0.htod_i32(&pos)?;
6108 let x = self.embed(e0, tokens)?;
6109 let _overlap = pipelined.then(crate::pp::enter_prime_pipe_stage);
6110 let x = self.prime_layers(e0, x, fence[0], fence[1], &pos_d, t, base, cache, seq_end)?;
6111 if pipelined {
6112 rt.tx_pipelined(0, &x, t * n_embd)
6113 } else {
6114 rt.tx(0, &x, t * n_embd)
6115 }
6116 }
6117
6118 #[allow(clippy::too_many_arguments)] fn prime_pp2_stage1_enqueue(
6120 &self,
6121 e: &Engine,
6122 rt: &crate::pp::PpNRt,
6123 slot: usize,
6124 t: usize,
6125 cache: &mut Cache,
6126 seq_end: usize,
6127 fence: &[usize],
6128 base: usize,
6129 pipelined: bool,
6130 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6131 let n_embd = self.cfg.n_embd as usize;
6132 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
6133 rt.bind_stage(1)?;
6134 let _st1 = rt.enter(1);
6135 let e1 = rt.engine(1, e);
6136 let pos_d = e1.htod_i32(&pos)?;
6137 let x = rt.rx(0, slot, t * n_embd)?;
6138 let _overlap = pipelined.then(crate::pp::enter_prime_pipe_stage);
6139 self.prime_layers(e1, x, fence[1], fence[2], &pos_d, t, base, cache, seq_end)
6140 }
6141
6142 #[allow(clippy::too_many_arguments)] pub fn prime_chunk_captured(
6159 &self,
6160 e: &Engine,
6161 x_in: &CudaSlice<f32>,
6162 pos_d: &CudaSlice<i32>,
6163 t: usize,
6164 cache: &mut Cache,
6165 len_d: &CudaSlice<i32>,
6166 logits_out: &mut CudaSlice<f32>,
6167 h_seed_out: &mut CudaSlice<f32>,
6168 ) -> Result<(), Box<dyn std::error::Error>> {
6169 self.refuse_hyper("prime_chunk_captured")?;
6170 cache.ensure_usable("prime_chunk_captured")?;
6171 let cfg = &self.cfg;
6172 let n_embd = cfg.n_embd as usize;
6173 let eps = cfg.rms_eps;
6174 let f16fuse = crate::f16_ffi::pp_f16_enabled() && t >= 16;
6175 let mut x = e.uninit(t * n_embd)?;
6176 e.copy_into(&mut x, 0, x_in, t * n_embd)?;
6177 for (il, layer) in self.layers.iter().enumerate() {
6178 let mut h = e.uninit(t * n_embd)?;
6179 let mut hx16: Option<CudaSlice<u8>> = None;
6180 if f16fuse {
6181 let mut b16 = e.alloc_u8_uninit(t * n_embd * 2)?;
6182 e.rms_norm_f16out(
6183 &x,
6184 layer.attn_norm.float_data(),
6185 &mut h,
6186 &mut b16,
6187 n_embd,
6188 t,
6189 eps,
6190 )?;
6191 hx16 = Some(b16);
6192 } else {
6193 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
6194 }
6195 let mixed = match &layer.mixer {
6196 Mixer::Full(fa) => {
6200 self.full_attn_prime(e, fa, &h, hx16.as_ref(), pos_d, t, cache, il, t)?
6201 }
6202 Mixer::Mla(_) => crate::hybrid::mla_path_unimplemented("captured-graph prime"),
6203 Mixer::Kda(_) => crate::hybrid::kda_path_unimplemented("captured prime chunk"),
6204 Mixer::Linear(la) => {
6205 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
6206 let g4 = match hx16.as_ref() {
6207 Some(xh) => e.matmul_group_xh(&ws, &h, xh, t)?,
6208 None => e.matmul_group(&ws, &h, t)?,
6209 };
6210 self.linear_attn_prime_core_pad(e, la, g4, t, cache, il, Some(len_d))?
6211 }
6212 };
6213 let mut x1 = e.uninit(t * n_embd)?;
6214 e.add(&x, &mixed, &mut x1, t * n_embd)?;
6215 let mut z = e.uninit(t * n_embd)?;
6216 let mut zx16: Option<CudaSlice<u8>> = None;
6217 if f16fuse {
6218 let mut b16 = e.alloc_u8_uninit(t * n_embd * 2)?;
6219 e.rms_norm_f16out(
6220 &x1,
6221 layer.post_attn_norm.float_data(),
6222 &mut z,
6223 &mut b16,
6224 n_embd,
6225 t,
6226 eps,
6227 )?;
6228 zx16 = Some(b16);
6229 } else {
6230 e.rms_norm(
6231 &x1,
6232 layer.post_attn_norm.float_data(),
6233 &mut z,
6234 n_embd,
6235 t,
6236 eps,
6237 )?;
6238 }
6239 let ffn_out = match &layer.ffn {
6240 crate::hybrid::Ffn::Dense {
6241 ffn_gate,
6242 ffn_up,
6243 ffn_down,
6244 } => {
6245 let n_ff = ffn_gate.out_features();
6246 let mut g2 = match &zx16 {
6247 Some(xh) => e.matmul_group_xh(&[ffn_gate, ffn_up], &z, xh, t)?,
6248 None => e.matmul_group(&[ffn_gate, ffn_up], &z, t)?,
6249 };
6250 let up = g2.pop().unwrap();
6251 let gate = g2.pop().unwrap();
6252 let mut act = e.uninit(t * n_ff)?;
6253 Self::ffn_act_lim(
6255 e,
6256 &self.cfg,
6257 &gate,
6258 &up,
6259 1.0,
6260 1.0,
6261 self.cfg.clamp_shexp_at(il as u32),
6262 &mut act,
6263 t * n_ff,
6264 )?;
6265 e.matmul(ffn_down, &act, t)?
6266 }
6267 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il_prefill(e, m, &z, t, il as u16)?,
6268 };
6269 let mut x2 = e.uninit(t * n_embd)?;
6270 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
6271 x = x2;
6272 }
6273 if !crate::spec::spec_hpost() {
6275 e.row_gather_dev(&x, h_seed_out, len_d, n_embd)?;
6276 }
6277 let mut hn = e.uninit(t * n_embd)?;
6278 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
6279 if crate::spec::spec_hpost() {
6280 e.row_gather_dev(&hn, h_seed_out, len_d, n_embd)?;
6281 }
6282 let mut hlast = e.uninit(n_embd)?;
6283 e.row_gather_dev(&hn, &mut hlast, len_d, n_embd)?;
6284 let logits = e.matmul(&self.output, &hlast, 1)?;
6285 let nv = logits.len();
6286 e.copy_into(logits_out, 0, &logits, nv)?;
6287 Ok(())
6288 }
6289
6290 fn step35_prime_batch_on() -> bool {
6291 std::env::var("MEMRA_STEP35_PRIME_BATCH").as_deref() != Ok("0")
6292 }
6293
6294 #[allow(clippy::too_many_arguments)]
6297 #[allow(clippy::too_many_arguments)]
6302 fn step35_prime_batch_layers(
6303 &self,
6304 e: &Engine,
6305 mut x: CudaSlice<f32>,
6306 lo: usize,
6307 hi: usize,
6308 ts: &[usize],
6309 offs: &[usize],
6310 seq_ends: &[usize],
6311 pos_ds: &[CudaSlice<i32>],
6312 caches: &mut [&mut Cache],
6313 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6314 let cfg = &self.cfg;
6315 let n_embd = cfg.n_embd as usize;
6316 let eps = cfg.rms_eps;
6317 let b = ts.len();
6318 let total: usize = ts.iter().sum();
6319 let f16fuse = crate::f16_ffi::pp_f16_enabled() && total >= 16;
6320
6321 let split = |e: &Engine,
6322 y: &CudaSlice<f32>,
6323 dim: usize|
6324 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
6325 let mut out = Vec::with_capacity(b);
6326 for s in 0..b {
6327 let mut ys = e.uninit(ts[s] * dim)?;
6328 e.copy_view_into(
6329 &mut ys,
6330 0,
6331 &y.slice(offs[s] * dim..(offs[s] + ts[s]) * dim),
6332 ts[s] * dim,
6333 )?;
6334 out.push(ys);
6335 }
6336 Ok(out)
6337 };
6338
6339 let prof = std::env::var("MEMRA_PRIME_PROF").as_deref() == Ok("1");
6344 let mut ph = [0f64; 4]; let mark = |e: &Engine, acc: usize, t0: &mut std::time::Instant, ph: &mut [f64; 4]| {
6346 if prof {
6347 let _ = e.stream().synchronize();
6348 ph[acc] += t0.elapsed().as_secs_f64() * 1e3;
6349 *t0 = std::time::Instant::now();
6350 }
6351 };
6352 let mut pt = std::time::Instant::now();
6353 for il in lo..hi {
6354 let layer = &self.layers[il];
6355 let Mixer::Full(fa) = &layer.mixer else {
6356 return Err(format!("step35 layer {il} is not full-attn — corrupt config").into());
6357 };
6358
6359 let mut h = e.uninit(total * n_embd)?;
6360 let mut hx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
6361 if f16fuse {
6362 e.rms_norm_f16out(
6363 &x,
6364 layer.attn_norm.float_data(),
6365 &mut h,
6366 &mut hx16,
6367 n_embd,
6368 total,
6369 eps,
6370 )?;
6371 } else {
6372 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, total, eps)?;
6373 }
6374
6375 let gate_w = fa
6379 .attn_gate
6380 .as_ref()
6381 .ok_or("step35 layer is missing attn_gate.weight (head-wise attention gate)")?;
6382 let mut g4 = if f16fuse {
6383 e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv, gate_w], &h, &hx16, total)?
6384 } else {
6385 e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv, gate_w], &h, total)?
6386 };
6387 let gate = g4.pop().unwrap();
6388 let mut parts: Vec<Vec<CudaSlice<f32>>> =
6389 (0..b).map(|_| Vec::with_capacity(3)).collect();
6390 for (w, y) in [&fa.wq, &fa.wk, &fa.wv].iter().zip(g4) {
6391 for (s, ys) in split(e, &y, w.out_features())?.into_iter().enumerate() {
6392 parts[s].push(ys);
6393 }
6394 }
6395 let gates = split(e, &gate, gate_w.out_features())?;
6396 let geometry = self.step35_geom(il);
6397 let hd = geometry.head_dim_k as usize;
6398 let nh = geometry.n_head as usize;
6399 let mut ag_cat = e.uninit(total * nh * hd)?;
6400 for (s, (g3s, gate)) in parts.into_iter().zip(gates).enumerate() {
6401 mark(e, 0, &mut pt, &mut ph);
6402 let ag = self.step35_attn_pre_wo(
6403 e,
6404 fa,
6405 g3s,
6406 None,
6407 Some(&gate),
6408 &pos_ds[s],
6409 ts[s],
6410 Some(&mut *caches[s]),
6411 il,
6412 seq_ends[s],
6413 )?;
6414 e.copy_into(&mut ag_cat, offs[s] * nh * hd, &ag, ts[s] * nh * hd)?;
6415 }
6416 mark(e, 1, &mut pt, &mut ph);
6417 let mixed = e.matmul(&fa.wo, &ag_cat, total)?;
6418
6419 let mut x1 = e.uninit(total * n_embd)?;
6420 let mut z = e.uninit(total * n_embd)?;
6421 let mut zx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
6422 if f16fuse {
6423 e.add_rms_norm_f16out(
6424 &x,
6425 &mixed,
6426 layer.post_attn_norm.float_data(),
6427 &mut x1,
6428 &mut z,
6429 &mut zx16,
6430 n_embd,
6431 total,
6432 eps,
6433 )?;
6434 } else {
6435 e.add(&x, &mixed, &mut x1, total * n_embd)?;
6436 e.rms_norm(
6437 &x1,
6438 layer.post_attn_norm.float_data(),
6439 &mut z,
6440 n_embd,
6441 total,
6442 eps,
6443 )?;
6444 }
6445
6446 mark(e, 2, &mut pt, &mut ph);
6447 let ffn_out = match &layer.ffn {
6448 crate::hybrid::Ffn::Dense {
6449 ffn_gate,
6450 ffn_up,
6451 ffn_down,
6452 } => {
6453 let n_ff = ffn_gate.out_features();
6454 let mut g2 = if f16fuse {
6455 e.matmul_group_xh(&[ffn_gate, ffn_up], &z, &zx16, total)?
6456 } else {
6457 e.matmul_group(&[ffn_gate, ffn_up], &z, total)?
6458 };
6459 let up = g2.pop().unwrap();
6460 let gate = g2.pop().unwrap();
6461 let mut act = e.uninit(total * n_ff)?;
6462 let d_lim = cfg.clamp_shexp_at(il as u32);
6463 if Self::f16out_on(e, total) && cfg.m3.is_none() && d_lim.is_none() {
6464 let mut a16 = e.alloc_u8_uninit(total * n_ff * 2)?;
6465 e.silu_mul_f16out(&gate, &up, &mut act, &mut a16, total * n_ff)?;
6466 match e.try_f16_gemm_pre(ffn_down, &a16, total)? {
6467 Some(y) => y,
6468 None => e.matmul(ffn_down, &act, total)?,
6469 }
6470 } else {
6471 Self::ffn_act_lim(
6472 e,
6473 cfg,
6474 &gate,
6475 &up,
6476 1.0,
6477 1.0,
6478 d_lim,
6479 &mut act,
6480 total * n_ff,
6481 )?;
6482 e.matmul(ffn_down, &act, total)?
6483 }
6484 }
6485 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il(e, m, &z, total, il as u16)?,
6486 };
6487 let mut x2 = e.uninit(total * n_embd)?;
6488 e.add(&x1, &ffn_out, &mut x2, total * n_embd)?;
6489 x = x2;
6490 mark(e, 3, &mut pt, &mut ph);
6491 }
6492 if prof {
6493 eprintln!(
6494 "[prime-prof] t={total} layers={} norm+qkv={:.0}ms attn={:.0}ms o_proj={:.0}ms moe={:.0}ms",
6495 hi - lo,
6496 ph[0],
6497 ph[1],
6498 ph[2],
6499 ph[3]
6500 );
6501 }
6502 Ok(x)
6503 }
6504
6505 #[allow(clippy::type_complexity)] fn step35_prime_batch_epilogue(
6507 &self,
6508 e: &Engine,
6509 x: CudaSlice<f32>,
6510 ts: &[usize],
6511 offs: &[usize],
6512 caches: &mut [&mut Cache],
6513 ) -> Result<Vec<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6514 let n_embd = self.cfg.n_embd as usize;
6515 let total: usize = ts.iter().sum();
6516 let mut hn = e.uninit(total * n_embd)?;
6517 e.rms_norm(
6518 &x,
6519 self.output_norm.float_data(),
6520 &mut hn,
6521 n_embd,
6522 total,
6523 self.cfg.rms_eps,
6524 )?;
6525
6526 let hidden_src = if crate::spec::spec_hpost() { &hn } else { &x };
6527 let mut out = Vec::with_capacity(ts.len());
6528 for s in 0..ts.len() {
6529 let mut hidden = e.uninit(ts[s] * n_embd)?;
6530 e.copy_view_into(
6531 &mut hidden,
6532 0,
6533 &hidden_src.slice(offs[s] * n_embd..(offs[s] + ts[s]) * n_embd),
6534 ts[s] * n_embd,
6535 )?;
6536 let last0 = (offs[s] + ts[s] - 1) * n_embd;
6537 let mut h_seed = e.uninit(n_embd)?;
6538 e.copy_view_into(
6539 &mut h_seed,
6540 0,
6541 &hidden_src.slice(last0..last0 + n_embd),
6542 n_embd,
6543 )?;
6544 let mut hlast = e.uninit(n_embd)?;
6546 e.copy_view_into(&mut hlast, 0, &hn.slice(last0..last0 + n_embd), n_embd)?;
6547 let logits = e.dtoh(&e.matmul(&self.output, &hlast, 1)?)?;
6548 caches[s].pos += ts[s];
6549 out.push((logits, h_seed, hidden));
6550 }
6551 Ok(out)
6552 }
6553
6554 #[allow(clippy::type_complexity)] fn step35_prime_cache_batch(
6560 &self,
6561 e: &Engine,
6562 prompts: &[&[u32]],
6563 caches: &mut [&mut Cache],
6564 seq_ends: &[usize],
6565 ) -> Result<Vec<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6566 assert_eq!(
6567 seq_ends.len(),
6568 prompts.len(),
6569 "step35 batched prime: one seq_end per sequence"
6570 );
6571 validate_step_prime_batch_modes(
6572 step_tp_prefill_enabled()?,
6573 step_ep_grouped_prefill_enabled()?,
6574 )?;
6575 if crate::pp::pp_host_bounce_active()
6576 && (!crate::pp::prime_pp_on() || crate::pp::pp_cuts(self.layers.len()).is_none())
6577 {
6578 return Err(
6579 "step35_prime_cache_batch: MEMRA_PP_HOST_BOUNCE=1 requires a valid prime \
6580 stage split; refusing an unsplit remote-weight walk"
6581 .into(),
6582 );
6583 }
6584 if !Self::step35_prime_batch_on() {
6585 return Err("step35 batched prime is disabled (MEMRA_STEP35_PRIME_BATCH=0)".into());
6586 }
6587 if prompts.len() > 1 && caches.iter().any(|c| c.pos != 0) {
6592 return Err(
6593 "step35 batched prime supports continuation only at B=1; a cross-request batch \
6594 at mixed positions requires per-request queued_after"
6595 .into(),
6596 );
6597 }
6598
6599 let ts: Vec<usize> = prompts.iter().map(|p| p.len()).collect();
6600 for &t in &ts {
6601 assert!(
6602 t >= PRIME_MIN_T,
6603 "step35 batched prime needs T >= {PRIME_MIN_T}"
6604 );
6605 }
6606 for (s, c) in caches.iter().enumerate() {
6607 assert!(
6610 c.pos + ts[s] <= c.max_ctx,
6611 "step35 batched prime exceeds cache max_ctx"
6612 );
6613 assert!(
6614 seq_ends[s] >= c.pos + ts[s],
6615 "step35 batched prime: seq_end must cover this chunk"
6616 );
6617 }
6618 let mut transaction = CacheTaintGuard::arm(caches);
6619 let legacy_tsend = std::env::var("MEMRA_STEP35_PRIME_BATCH_TSEND").as_deref() == Ok("1");
6626 let seq_ends_eff: Vec<usize> = if legacy_tsend {
6627 ts.clone()
6628 } else {
6629 seq_ends.to_vec()
6630 };
6631 let offs: Vec<usize> = ts
6632 .iter()
6633 .scan(0usize, |a, &t| {
6634 let o = *a;
6635 *a += t;
6636 Some(o)
6637 })
6638 .collect();
6639 let total: usize = ts.iter().sum();
6640 let payload = total * self.cfg.n_embd as usize;
6641 let cat_tokens: Vec<u32> = prompts.iter().flat_map(|p| p.iter().copied()).collect();
6642 let positions: Vec<Vec<i32>> = ts
6648 .iter()
6649 .zip(caches.iter())
6650 .map(|(&t, c)| {
6651 let base = c.pos as i32;
6652 (0..t as i32).map(|i| base + i).collect()
6653 })
6654 .collect();
6655 let upload_positions =
6656 |e: &Engine| -> Result<Vec<CudaSlice<i32>>, Box<dyn std::error::Error>> {
6657 positions
6658 .iter()
6659 .map(|p| e.htod_i32(p))
6660 .collect::<Result<_, _>>()
6661 };
6662
6663 static ONCE: std::sync::Once = std::sync::Once::new();
6664 ONCE.call_once(|| {
6665 eprintln!(
6666 "[step35-prime-batch] first concat prime: B={} tokens={total}",
6667 prompts.len()
6668 );
6669 });
6670
6671 let out = if !crate::pp::pp2_streams_off() && crate::pp::prime_pp_on() {
6672 if let Some(fence) = crate::pp::pp_cuts(self.layers.len()) {
6673 let rt = crate::pp::PpNRt::get(e)?;
6674 let n_st = fence.len() - 1;
6675 assert_eq!(
6676 rt.n_stages(),
6677 n_st,
6678 "step35 prime batch stage count mismatch"
6679 );
6680 let caller_stream = e.stream();
6681 rt.fence_stages_behind(&caller_stream)?;
6682
6683 let mut slot = {
6684 let _st0 = rt.enter(0);
6685 let e0 = rt.engine(0, e);
6686 let pos_ds = upload_positions(e0)?;
6687 let x = self.embed(e0, &cat_tokens)?;
6688 let x = self.step35_prime_batch_layers(
6689 e0,
6690 x,
6691 fence[0],
6692 fence[1],
6693 &ts,
6694 &offs,
6695 &seq_ends_eff,
6696 &pos_ds,
6697 caches,
6698 )?;
6699 rt.tx(0, &x, payload)?
6700 };
6701 for s in 1..n_st - 1 {
6702 let _st = rt.enter(s);
6703 let es = rt.engine(s, e);
6704 let pos_ds = upload_positions(es)?;
6705 let x = rt.rx(s - 1, slot, payload)?;
6706 let x = self.step35_prime_batch_layers(
6707 es,
6708 x,
6709 fence[s],
6710 fence[s + 1],
6711 &ts,
6712 &offs,
6713 &seq_ends_eff,
6714 &pos_ds,
6715 caches,
6716 )?;
6717 slot = rt.tx(s, &x, payload)?;
6718 }
6719
6720 let _stl = rt.enter(n_st - 1);
6721 let el = rt.engine(n_st - 1, e);
6722 let pos_ds = upload_positions(el)?;
6723 let x = rt.rx(n_st - 2, slot, payload)?;
6724 let x = self.step35_prime_batch_layers(
6725 el,
6726 x,
6727 fence[n_st - 1],
6728 fence[n_st],
6729 &ts,
6730 &offs,
6731 &seq_ends_eff,
6732 &pos_ds,
6733 caches,
6734 )?;
6735 let out = self.step35_prime_batch_epilogue(el, x, &ts, &offs, caches)?;
6736 rt.publish_to(n_st - 1, &caller_stream)?;
6737 crate::pp::STEP35_PRIME_BATCH_SPLITS
6738 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
6739 out
6740 } else {
6741 let pos_ds = upload_positions(e)?;
6742 let x = self.embed(e, &cat_tokens)?;
6743 let x = self.step35_prime_batch_layers(
6744 e,
6745 x,
6746 0,
6747 self.layers.len(),
6748 &ts,
6749 &offs,
6750 &seq_ends_eff,
6751 &pos_ds,
6752 caches,
6753 )?;
6754 self.step35_prime_batch_epilogue(e, x, &ts, &offs, caches)?
6755 }
6756 } else {
6757 let pos_ds = upload_positions(e)?;
6758 let x = self.embed(e, &cat_tokens)?;
6759 let x = self.step35_prime_batch_layers(
6760 e,
6761 x,
6762 0,
6763 self.layers.len(),
6764 &ts,
6765 &offs,
6766 &seq_ends_eff,
6767 &pos_ds,
6768 caches,
6769 )?;
6770 self.step35_prime_batch_epilogue(e, x, &ts, &offs, caches)?
6771 };
6772 crate::pp::STEP35_PRIME_BATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
6773 transaction.commit();
6774 Ok(out)
6775 }
6776
6777 #[allow(clippy::type_complexity)] pub fn prime_cache_batch(
6795 &self,
6796 e: &Engine,
6797 prompts: &[&[u32]],
6798 caches: &mut [&mut Cache],
6799 ) -> Result<Vec<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6800 let events_before = crate::progress::events();
6807 let out = self.prime_cache_batch_inner(e, prompts, caches);
6808 if out.is_ok() && crate::progress::events() == events_before {
6809 crate::progress::note_prime_rows(prompts.iter().map(|p| p.len()).sum());
6810 }
6811 out
6812 }
6813
6814 #[allow(clippy::type_complexity)] fn prime_cache_batch_inner(
6816 &self,
6817 e: &Engine,
6818 prompts: &[&[u32]],
6819 caches: &mut [&mut Cache],
6820 ) -> Result<Vec<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
6821 self.refuse_hyper("prime_cache_batch")?;
6822 for cache in caches.iter() {
6823 cache.ensure_usable("prime_cache_batch")?;
6824 }
6825 if crate::pp::pp_cuts(self.layers.len()).is_some()
6826 && !self.rewrite_allowed(memra_gguf::execution_manifest::RewriteSurface::Pipeline)
6827 {
6828 return Err("pipeline rewrite is not qualified for batched prime".into());
6829 }
6830 if !self.rewrite_allowed(memra_gguf::execution_manifest::RewriteSurface::CarriedPrime) {
6831 if !self.rewrite_allowed(memra_gguf::execution_manifest::RewriteSurface::DecodeEager) {
6832 return Err("neither batched-prime nor eager rewrite is qualified".into());
6833 }
6834 if prompts.len() != caches.len() {
6835 return Err("prime fallback prompt/cache shape mismatch".into());
6836 }
6837 static ONCE: std::sync::Once = std::sync::Once::new();
6838 ONCE.call_once(|| {
6839 eprintln!(
6840 "[rewrite] carried-prime.v1 unqualified; using individual native eager primes"
6841 );
6842 });
6843 let mut transaction = CacheTaintGuard::arm(caches);
6844 let result: Result<Vec<_>, Box<dyn std::error::Error>> = prompts
6845 .iter()
6846 .copied()
6847 .zip(caches.iter_mut())
6848 .map(|(prompt, cache)| self.prime_cache(e, prompt, cache, 0))
6849 .collect();
6850 if result.is_ok() {
6851 transaction.commit();
6852 }
6853 return result;
6854 }
6855 let _pp_walk =
6856 if crate::pp::pp_cuts(self.layers.len()).is_some() && !crate::pp::pp2_streams_off() {
6857 let rt = crate::pp::PpNRt::get(e)?;
6858 Some(rt.acquire_walk("prime_cache_batch")?)
6859 } else {
6860 None
6861 };
6862 let cfg = &self.cfg;
6863 let n_embd = cfg.n_embd as usize;
6864 let eps = cfg.rms_eps;
6865 let b = prompts.len();
6866 assert!(b >= 1 && b == caches.len());
6867 let pos0s: Vec<usize> = caches.iter().map(|c| c.pos).collect();
6868 let carried = pos0s.iter().any(|&p| p > 0);
6869 if self.uses_gemma_program() {
6875 return Err(
6876 "prime_cache_batch: gemma4 has no batched prime core (per-layer \
6877 swa/global geometry, softcapped head) — use gemma4_prime per sequence"
6878 .into(),
6879 );
6880 }
6881 if self.uses_sliding_gated_moe_program() {
6884 let seq_ends: Vec<usize> = caches
6889 .iter()
6890 .zip(prompts.iter())
6891 .map(|(c, p)| c.pos + p.len())
6892 .collect();
6893 return self.step35_prime_cache_batch(e, prompts, caches, &seq_ends);
6894 }
6895 if crate::pp::pp_cuts(self.layers.len()).is_some() && !crate::pp::pp2_streams_off() {
6896 let rt = crate::pp::PpNRt::get(e)?;
6897 if rt.cross_device() {
6898 return Err(
6899 "prime_cache_batch: generic dense concat prime has no cross-device PP split; use individual prime_cache calls"
6900 .into(),
6901 );
6902 }
6903 }
6904 let ts: Vec<usize> = prompts.iter().map(|p| p.len()).collect();
6905 for &t in &ts {
6906 assert!(
6907 t >= PRIME_MIN_T,
6908 "prime_cache_batch needs T >= {PRIME_MIN_T}"
6909 );
6910 }
6911 for (s, c) in caches.iter().enumerate() {
6912 assert!(
6913 c.pos + ts[s] <= c.max_ctx,
6914 "prime_cache_batch: prompt exceeds cache max_ctx"
6915 );
6916 }
6917 let mut transaction = CacheTaintGuard::arm(caches);
6918 let total: usize = ts.iter().sum();
6919 let offs: Vec<usize> = ts
6920 .iter()
6921 .scan(0usize, |a, &t| {
6922 let o = *a;
6923 *a += t;
6924 Some(o)
6925 })
6926 .collect();
6927 let pos_ds: Vec<CudaSlice<i32>> = ts
6929 .iter()
6930 .zip(&pos0s)
6931 .map(|(&t, &p0)| e.htod_i32(&(p0 as i32..(p0 + t) as i32).collect::<Vec<_>>()))
6932 .collect::<Result<_, _>>()?;
6933 let split = |e: &Engine,
6935 y: &CudaSlice<f32>,
6936 dim: usize|
6937 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
6938 let mut out = Vec::with_capacity(b);
6939 for s in 0..b {
6940 let mut ys = e.uninit(ts[s] * dim)?;
6941 e.copy_view_into(
6942 &mut ys,
6943 0,
6944 &y.slice(offs[s] * dim..(offs[s] + ts[s]) * dim),
6945 ts[s] * dim,
6946 )?;
6947 out.push(ys);
6948 }
6949 Ok(out)
6950 };
6951
6952 let cat_tokens: Vec<u32> = prompts.iter().flat_map(|p| p.iter().copied()).collect();
6953 let mut x = self.embed(e, &cat_tokens)?; for (il, layer) in self.layers.iter().enumerate() {
6955 let mut h = e.uninit(total * n_embd)?;
6956 let mut hx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
6957 e.rms_norm_f16out(
6958 &x,
6959 layer.attn_norm.float_data(),
6960 &mut h,
6961 &mut hx16,
6962 n_embd,
6963 total,
6964 eps,
6965 )?;
6966 let mut mixed = e.uninit(total * n_embd)?;
6968 match &layer.mixer {
6969 Mixer::Full(fa) => {
6970 let g3 = e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], &h, &hx16, total)?;
6971 let geometry = self.cfg.full_attention_geometry_at(il as u32);
6977 let (n_head, n_head_kv, head_dim) = (
6978 geometry.n_head as usize,
6979 geometry.n_head_kv as usize,
6980 geometry.head_dim_k as usize,
6981 );
6982 let fa_scale = geometry.attention_scale();
6983 let use_favl = !carried
6984 && (2..=8).contains(&b)
6985 && (head_dim == 256 || head_dim == 128)
6986 && geometry.attention_gate == memra_gguf::config::AttentionGateKind::FusedQ
6987 && std::env::var("MEMRA_NOFA").is_err()
6988 && std::env::var("MEMRA_FA_FLOOR").is_err()
6989 && std::env::var("MEMRA_FA_PP_W2").as_deref() != Ok("1")
6990 && std::env::var("MEMRA_FA_BF16KV").as_deref() != Ok("0")
6991 && std::env::var("MEMRA_FA_VL").as_deref() != Ok("0");
6992 if use_favl {
6993 let (qf_w, kf_w, vf_w) = (
6994 fa.wq.out_features(),
6995 fa.wk.out_features(),
6996 fa.wv.out_features(),
6997 );
6998 memra_gguf::config::check_fused_q_gate_extent(qf_w, head_dim, n_head, 1)?;
7003 struct APre {
7004 q: CudaSlice<f32>,
7005 gate: Option<CudaSlice<f32>>,
7006 qn: CudaSlice<f32>,
7007 kn: CudaSlice<f32>,
7008 }
7009 let mut aps = Vec::with_capacity(b);
7010 for &t in ts.iter().take(b) {
7011 aps.push(APre {
7012 q: e.uninit(t * n_head * head_dim)?,
7013 gate: Some(e.uninit(t * n_head * head_dim)?),
7014 qn: e.uninit(t * n_head * head_dim)?,
7015 kn: e.uninit(t * n_head_kv * head_dim)?,
7016 });
7017 }
7018 let (kv_dim_k, kv_dim_v, ktb, vtb) = {
7019 let kvl = caches[0].kv[il].as_ref().unwrap();
7020 (kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes)
7021 };
7022 let pargs: Vec<crate::AttnPreVl> = (0..b)
7023 .map(|s| {
7024 let (o, t) = (offs[s], ts[s]);
7025 let kvl = caches[s].kv[il].as_ref().unwrap();
7026 assert!(
7027 kvl.len == 0 && kvl.len + t <= caches[s].max_ctx,
7028 "prime_cache_batch attn vl: fresh + capacity"
7029 );
7030 crate::AttnPreVl {
7031 qf: e.addr_f32v(&g3[0].slice(o * qf_w..(o + t) * qf_w)),
7032 kf: e.addr_f32v(&g3[1].slice(o * kf_w..(o + t) * kf_w)),
7033 vf: e.addr_f32v(&g3[2].slice(o * vf_w..(o + t) * vf_w)),
7034 q: e.addr_f32(&aps[s].q),
7035 gate: e.addr_f32(aps[s].gate.as_ref().unwrap()),
7036 qn: e.addr_f32(&aps[s].qn),
7037 kn: e.addr_f32(&aps[s].kn),
7038 kc: e.addr_u8(&kvl.k),
7039 vc: e.addr_u8(&kvl.v),
7040 t: t as i32,
7041 pad: 0,
7042 }
7043 })
7044 .collect();
7045 e.attn_pre_vl8(
7046 &pargs,
7047 fa.q_norm.float_data(),
7048 fa.k_norm.float_data(),
7049 head_dim,
7050 geometry.n_rot as usize,
7051 n_head,
7052 n_head_kv,
7053 self.cfg.rms_eps,
7054 geometry.rope_base,
7055 1.0,
7056 kv_dim_k,
7057 kv_dim_v,
7058 ktb,
7059 vtb,
7060 )?;
7061 for s in 0..b {
7062 let kvl = caches[s].kv[il].as_mut().unwrap();
7063 kvl.len += ts[s];
7064 let new_len = kvl.len as i32;
7065 e.set_i32_one(&mut kvl.len_d, new_len)?;
7066 }
7067 let mut attns = Vec::with_capacity(b);
7068 let mut mirrors = Vec::with_capacity(b);
7069 for &t in ts.iter().take(b) {
7070 attns.push(e.uninit(t * n_head * head_dim)?);
7071 let n = t * n_head_kv * head_dim;
7072 mirrors.push((e.alloc_u8_uninit(n * 2)?, e.alloc_u8_uninit(n * 2)?));
7073 }
7074 let fa3_on = match std::env::var("MEMRA_FA3").as_deref() {
7077 Ok("0") => false,
7078 Ok("1") => {
7082 crate::refuse_portable_force(
7083 "MEMRA_FA3=1",
7084 "the sm_90a fa3/bf16 kernels",
7085 );
7086 true
7087 }
7088 _ => cfg!(memra_hopper_mma),
7089 };
7090 if fa3_on {
7091 let mut q16s = Vec::with_capacity(b);
7092 let mut v16s = Vec::with_capacity(b);
7093 for s in 0..b {
7094 let t = ts[s];
7095 let mut q16 = e.alloc_u8_uninit(t * n_head * head_dim * 2)?;
7096 e.f32_to_bf16_into(&aps[s].qn, &mut q16, t * n_head * head_dim)?;
7097 let mut k16 = e.alloc_u8_uninit(t * n_head_kv * head_dim * 2)?;
7098 e.f32_to_bf16_into(&aps[s].kn, &mut k16, t * n_head_kv * head_dim)?;
7099 let mut v16 = e.alloc_u8_uninit(t * n_head_kv * head_dim * 2)?;
7100 e.f32_to_bf16_v(
7101 &g3[2].slice(offs[s] * vf_w..(offs[s] + t) * vf_w),
7102 &mut v16,
7103 t * n_head_kv * head_dim,
7104 )?;
7105 q16s.push(q16);
7106 v16s.push((k16, v16));
7107 }
7108 let mut qp = [core::ptr::null::<core::ffi::c_void>(); 8];
7109 let mut kp = qp;
7110 let mut vp = qp;
7111 let mut op = [core::ptr::null_mut::<f32>(); 8];
7112 let mut tsv = [0i32; 8];
7113 for s in 0..b {
7114 qp[s] = e.addr_u8(&q16s[s]) as *const core::ffi::c_void;
7115 kp[s] = e.addr_u8(&v16s[s].0) as *const core::ffi::c_void;
7116 vp[s] = e.addr_u8(&v16s[s].1) as *const core::ffi::c_void;
7117 op[s] = e.addr_f32(&attns[s]) as *mut f32;
7118 tsv[s] = ts[s] as i32;
7119 }
7120 let rc = unsafe {
7121 crate::fa3_vl_raw(
7122 qp.as_ptr(),
7123 kp.as_ptr(),
7124 vp.as_ptr(),
7125 op.as_ptr(),
7126 tsv.as_ptr(),
7127 b as i32,
7128 n_head as i32,
7129 n_head_kv as i32,
7130 head_dim as i32,
7131 fa_scale,
7132 e.stream().cu_stream() as *mut core::ffi::c_void,
7133 )
7134 };
7135 if rc != 0 {
7136 return Err(format!("memra_fa3_vl rc={rc}").into());
7137 }
7138 } else {
7139 let fargs: Vec<crate::FaSeqVl> = (0..b)
7140 .map(|s| crate::FaSeqVl {
7141 q: e.addr_f32(&aps[s].qn),
7142 k16: e.addr_u8(&mirrors[s].0),
7143 v16: e.addr_u8(&mirrors[s].1),
7144 o: e.addr_f32(&attns[s]),
7145 kf: e.addr_f32(&aps[s].kn),
7146 vf: e.addr_f32v(
7147 &g3[2].slice(offs[s] * vf_w..(offs[s] + ts[s]) * vf_w),
7148 ),
7149 t: ts[s] as i32,
7150 pad: 0,
7151 })
7152 .collect();
7153 e.fa_prefill_vl8(&fargs, head_dim, n_head, n_head_kv, fa_scale)?;
7154 }
7155 for (s, attn) in attns.into_iter().enumerate() {
7156 let (attn_g, ag16) = self.full_attn_prime_post_fa(
7157 e,
7158 attn,
7159 &aps[s].gate,
7160 ts[s],
7161 n_head,
7162 head_dim,
7163 )?;
7164 let mut done = false;
7165 if let Some(xh) = &ag16 {
7166 done = e.try_f16_gemm_pre_into_off(
7167 &fa.wo,
7168 xh,
7169 ts[s],
7170 &mut mixed,
7171 offs[s] * n_embd,
7172 )?;
7173 }
7174 if !done {
7175 let m = e.matmul(&fa.wo, &attn_g, ts[s])?;
7176 e.copy_into(&mut mixed, offs[s] * n_embd, &m, ts[s] * n_embd)?;
7177 }
7178 }
7179 } else {
7180 let mut parts: Vec<Vec<CudaSlice<f32>>> =
7181 (0..b).map(|_| Vec::new()).collect();
7182 for (w, y) in [&fa.wq, &fa.wk, &fa.wv].iter().zip(g3) {
7183 for (s, ys) in split(e, &y, w.out_features())?.into_iter().enumerate() {
7184 parts[s].push(ys);
7185 }
7186 }
7187 for (s, g3s) in parts.into_iter().enumerate() {
7188 let (attn_g, ag16) = self.full_attn_prime_core_inner(
7190 e, fa, g3s, &pos_ds[s], ts[s], caches[s], il,
7191 )?;
7192 let mut done = false;
7193 if let Some(xh) = &ag16 {
7194 done = e.try_f16_gemm_pre_into_off(
7195 &fa.wo,
7196 xh,
7197 ts[s],
7198 &mut mixed,
7199 offs[s] * n_embd,
7200 )?;
7201 }
7202 if !done {
7203 let m = e.matmul(&fa.wo, &attn_g, ts[s])?;
7204 e.copy_into(&mut mixed, offs[s] * n_embd, &m, ts[s] * n_embd)?;
7205 }
7206 }
7207 }
7208 }
7209 Mixer::Mla(_) => crate::hybrid::mla_path_unimplemented("batched cache prime"),
7210 Mixer::Kda(_) => crate::hybrid::kda_path_unimplemented("batched prime"),
7211 Mixer::Linear(la) => {
7212 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
7217 let g4 = e.matmul_group_xh(&ws, &h, &hx16, total)?;
7218 let outs =
7219 self.linear_attn_prime_core_batch(e, la, &g4, &offs, &ts, caches, il)?;
7220 for (s, (gn, gn16)) in outs.into_iter().enumerate() {
7221 let (o, t) = (offs[s], ts[s]);
7222 let mut done = false;
7223 if let Some(xh) = &gn16 {
7224 done = e.try_f16_gemm_pre_into_off(
7225 &la.ssm_out,
7226 xh,
7227 t,
7228 &mut mixed,
7229 o * n_embd,
7230 )?;
7231 }
7232 if !done {
7233 let m = e.matmul(&la.ssm_out, &gn, t)?;
7234 e.copy_into(&mut mixed, o * n_embd, &m, t * n_embd)?;
7235 }
7236 }
7237 }
7238 }
7239 let mut x1 = e.uninit(total * n_embd)?;
7240 let mut z = e.uninit(total * n_embd)?;
7241 let mut zx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
7242 e.add_rms_norm_f16out(
7243 &x,
7244 &mixed,
7245 layer.post_attn_norm.float_data(),
7246 &mut x1,
7247 &mut z,
7248 &mut zx16,
7249 n_embd,
7250 total,
7251 eps,
7252 )?;
7253 let ffn_out = match &layer.ffn {
7254 crate::hybrid::Ffn::Dense {
7255 ffn_gate,
7256 ffn_up,
7257 ffn_down,
7258 } => {
7259 let n_ff = ffn_gate.out_features();
7260 let mut g2 = e.matmul_group_xh(&[ffn_gate, ffn_up], &z, &zx16, total)?;
7261 let up = g2.pop().unwrap();
7262 let gate = g2.pop().unwrap();
7263 let mut act = e.uninit(total * n_ff)?;
7264 let d_lim = self.cfg.clamp_shexp_at(il as u32);
7268 if Self::f16out_on(e, total) && self.cfg.m3.is_none() && d_lim.is_none() {
7269 let mut a16 = e.alloc_u8_uninit(total * n_ff * 2)?;
7270 e.silu_mul_f16out(&gate, &up, &mut act, &mut a16, total * n_ff)?;
7271 match e.try_f16_gemm_pre(ffn_down, &a16, total)? {
7272 Some(y) => y,
7273 None => e.matmul(ffn_down, &act, total)?,
7274 }
7275 } else {
7276 Self::ffn_act_lim(
7277 e,
7278 &self.cfg,
7279 &gate,
7280 &up,
7281 1.0,
7282 1.0,
7283 d_lim,
7284 &mut act,
7285 total * n_ff,
7286 )?;
7287 e.matmul(ffn_down, &act, total)?
7288 }
7289 }
7290 crate::hybrid::Ffn::Moe(m) => {
7291 self.moe_ffn_il_prefill(e, m, &z, total, il as u16)?
7292 }
7293 };
7294 let mut x2 = e.uninit(total * n_embd)?;
7295 e.add(&x1, &ffn_out, &mut x2, total * n_embd)?;
7296 x = x2;
7297 }
7298 let mut hn = e.uninit(total * n_embd)?;
7300 e.rms_norm(
7301 &x,
7302 self.output_norm.float_data(),
7303 &mut hn,
7304 n_embd,
7305 total,
7306 eps,
7307 )?;
7308 let mut hcat = e.uninit(b * n_embd)?;
7314 for s in 0..b {
7315 let last0 = (offs[s] + ts[s] - 1) * n_embd;
7316 e.copy_view_into(
7317 &mut hcat,
7318 s * n_embd,
7319 &hn.slice(last0..last0 + n_embd),
7320 n_embd,
7321 )?;
7322 }
7323 let logits_cat = if b >= 2 {
7324 e.try_f16_gemm(&self.output, &hcat, b)?
7325 } else {
7326 None
7327 };
7328 let logits_host: Option<Vec<f32>> = match &logits_cat {
7329 Some(lc) => Some(e.dtoh(lc)?),
7330 None => None,
7331 };
7332 let n_vocab = self.output.out_features();
7333 let mut hidden_all = if crate::spec::spec_hpost() {
7334 split(e, &hn, n_embd)?
7335 } else {
7336 split(e, &x, n_embd)?
7337 };
7338 let mut out = Vec::with_capacity(b);
7339 for s in 0..b {
7340 let last0 = (offs[s] + ts[s] - 1) * n_embd;
7341 let mut h_seed = e.uninit(n_embd)?;
7342 if !crate::spec::spec_hpost() {
7343 e.copy_view_into(&mut h_seed, 0, &x.slice(last0..last0 + n_embd), n_embd)?;
7344 } else {
7345 e.copy_view_into(&mut h_seed, 0, &hn.slice(last0..last0 + n_embd), n_embd)?;
7346 }
7347 let logits = match &logits_host {
7348 Some(lh) => lh[s * n_vocab..(s + 1) * n_vocab].to_vec(),
7349 None => {
7350 let mut hlast = e.uninit(n_embd)?;
7351 e.copy_view_into(&mut hlast, 0, &hn.slice(last0..last0 + n_embd), n_embd)?;
7352 e.dtoh(&e.matmul(&self.output, &hlast, 1)?)?
7353 }
7354 };
7355 caches[s].pos += ts[s];
7356 out.push((logits, h_seed, hidden_all.remove(0)));
7357 }
7358 transaction.commit();
7359 Ok(out)
7360 }
7361
7362 #[allow(clippy::too_many_arguments)]
7373 fn full_attn_prime(
7374 &self,
7375 e: &Engine,
7376 fa: &FullAttnLayer,
7377 h: &CudaSlice<f32>,
7378 hx: Option<&CudaSlice<u8>>,
7379 pos_d: &CudaSlice<i32>,
7380 t: usize,
7381 cache: &mut Cache,
7382 il: usize,
7383 seq_end: usize,
7384 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7385 if self.uses_sliding_gated_moe_program() {
7386 return self.step35_attn_prime(e, fa, h, hx, pos_d, t, cache, il, seq_end);
7387 }
7388 let g3 = match hx {
7393 Some(xh) => e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], h, xh, t)?,
7394 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
7395 };
7396 self.full_attn_prime_core(e, fa, g3, pos_d, t, cache, il)
7397 }
7398
7399 #[allow(clippy::too_many_arguments)] fn full_attn_prime_core(
7404 &self,
7405 e: &Engine,
7406 fa: &FullAttnLayer,
7407 g3: Vec<CudaSlice<f32>>,
7408 pos_d: &CudaSlice<i32>,
7409 t: usize,
7410 cache: &mut Cache,
7411 il: usize,
7412 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7413 let (attn_g, ag16) = self.full_attn_prime_core_inner(e, fa, g3, pos_d, t, cache, il)?;
7414 if let Some(xh) = &ag16
7415 && let Some(y) = e.try_f16_gemm_pre(&fa.wo, xh, t)?
7416 {
7417 return Ok(y);
7418 }
7419 e.matmul(&fa.wo, &attn_g, t)
7420 }
7421
7422 #[allow(clippy::type_complexity)] #[allow(clippy::too_many_arguments)] fn full_attn_prime_core_inner(
7425 &self,
7426 e: &Engine,
7427 fa: &FullAttnLayer,
7428 g3: Vec<CudaSlice<f32>>,
7429 pos_d: &CudaSlice<i32>,
7430 t: usize,
7431 cache: &mut Cache,
7432 il: usize,
7433 ) -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
7434 let cfg = &self.cfg;
7435 let geometry = cfg.full_attention_geometry_at(il as u32);
7436 let n_head = geometry.n_head as usize;
7437 let n_head_kv = geometry.n_head_kv as usize;
7438 let head_dim = geometry.head_dim_k as usize;
7439 let scale = geometry.attention_scale();
7440 let (pre, base_len) = self.full_attn_prime_pre_fa(e, fa, g3, pos_d, t, cache, il)?;
7441 let AttnPre { q, k, v, gate } = pre;
7442 let mut attn = e.uninit(t * n_head * head_dim)?;
7443 self.full_attn_prime_fa_dispatch(
7444 e, &q, &k, &v, &mut attn, base_len, t, cache, il, head_dim, n_head, n_head_kv, scale,
7445 )?;
7446 self.full_attn_prime_post_fa(e, attn, &gate, t, n_head, head_dim)
7447 }
7448
7449 #[allow(clippy::type_complexity)]
7453 #[allow(clippy::too_many_arguments)] fn full_attn_prime_pre_fa(
7455 &self,
7456 e: &Engine,
7457 fa: &FullAttnLayer,
7458 mut g3: Vec<CudaSlice<f32>>,
7459 pos_d: &CudaSlice<i32>,
7460 t: usize,
7461 cache: &mut Cache,
7462 il: usize,
7463 ) -> Result<(AttnPre, usize), Box<dyn std::error::Error>> {
7464 let cfg = &self.cfg;
7465 let geometry = cfg.full_attention_geometry_at(il as u32);
7466 let n_head = geometry.n_head as usize;
7467 let n_head_kv = geometry.n_head_kv as usize;
7468 let head_dim = geometry.head_dim_k as usize;
7469 let eps = cfg.rms_eps;
7470
7471 let gated = geometry.attention_gate == memra_gguf::config::AttentionGateKind::FusedQ;
7475 let v = g3.pop().unwrap();
7476 let mut k = g3.pop().unwrap();
7477 let qf = g3.pop().unwrap();
7478 let (mut q, gate) = if gated {
7479 let mut q = e.uninit(t * n_head * head_dim)?;
7480 let mut gate = e.uninit(t * n_head * head_dim)?;
7481 e.q_gate_split(&qf, &mut q, &mut gate, head_dim, n_head, t)?;
7482 (q, Some(gate))
7483 } else {
7484 (qf, None)
7485 };
7486
7487 let mut qn = e.uninit(t * n_head * head_dim)?;
7488 e.rms_norm(
7489 &q,
7490 fa.q_norm.float_data(),
7491 &mut qn,
7492 head_dim,
7493 n_head * t,
7494 eps,
7495 )?;
7496 q = qn;
7497 let mut kn = e.uninit(t * n_head_kv * head_dim)?;
7498 e.rms_norm(
7499 &k,
7500 fa.k_norm.float_data(),
7501 &mut kn,
7502 head_dim,
7503 n_head_kv * t,
7504 eps,
7505 )?;
7506 k = kn;
7507 let rope_dims = geometry.n_rot as usize;
7508 e.rope_neox(
7509 &mut q,
7510 pos_d,
7511 head_dim,
7512 rope_dims,
7513 n_head,
7514 t,
7515 geometry.rope_base,
7516 1.0,
7517 )?;
7518 e.rope_neox(
7519 &mut k,
7520 pos_d,
7521 head_dim,
7522 rope_dims,
7523 n_head_kv,
7524 t,
7525 geometry.rope_base,
7526 1.0,
7527 )?;
7528
7529 {
7532 let kvl = cache.kv[il].as_mut().unwrap();
7533 assert!(kvl.len + t <= cache.max_ctx, "prime_cache: KV overflow");
7534 e.append_kv_quantized_rows(
7535 &k,
7536 &v,
7537 &mut kvl.k,
7538 &mut kvl.v,
7539 kvl.len,
7540 t,
7541 kvl.kv_dim_k,
7542 kvl.kv_dim_v,
7543 kvl.k_tok_bytes,
7544 kvl.v_tok_bytes,
7545 crate::Engine::kv_fp8_on(),
7546 )?;
7547 kvl.len += t;
7548 let new_len = kvl.len as i32;
7549 e.set_i32_one(&mut kvl.len_d, new_len)?;
7550 }
7551
7552 let base_len = {
7553 let kvl = cache.kv[il].as_ref().unwrap();
7554 kvl.len - t };
7556 Ok((AttnPre { q, k, v, gate }, base_len))
7557 }
7558
7559 #[allow(clippy::too_many_arguments)]
7566 fn full_attn_prime_fa_dispatch(
7567 &self,
7568 e: &Engine,
7569 q: &CudaSlice<f32>,
7570 k: &CudaSlice<f32>,
7571 v: &CudaSlice<f32>,
7572 attn: &mut CudaSlice<f32>,
7573 base_len: usize,
7574 t: usize,
7575 cache: &mut Cache,
7576 il: usize,
7577 head_dim: usize,
7578 n_head: usize,
7579 n_head_kv: usize,
7580 scale: f32,
7581 ) -> Result<(), Box<dyn std::error::Error>> {
7582 if base_len == 0 && std::env::var("MEMRA_PRIME_F32CHUNK0").as_deref() == Ok("1") {
7595 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
7596 e.sdpa_naive(
7597 q, k, v, attn, head_dim, n_head, n_head_kv, t, t, scale, true,
7598 )?;
7599 } else {
7600 e.fa_prefill(
7601 q, k, v, attn, head_dim, n_head, n_head_kv, t, t, scale, true,
7602 )?;
7603 }
7604 return Ok(());
7605 }
7606 let kvl = cache.kv[il].as_ref().unwrap();
7607 let t_kv = base_len + t;
7608 let k_view = e.view_u8(&kvl.k, t_kv * kvl.k_tok_bytes);
7609 let v_view = e.view_u8(&kvl.v, t_kv * kvl.v_tok_bytes);
7610 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
7614 e.sdpa_naive_quantized_view(
7615 q,
7616 &k_view,
7617 &v_view,
7618 attn,
7619 head_dim,
7620 n_head,
7621 n_head_kv,
7622 t,
7623 t_kv,
7624 scale,
7625 true,
7626 kvl.k_tok_bytes,
7627 kvl.v_tok_bytes,
7628 )?;
7629 return Ok(());
7630 }
7631 let deqw = std::env::var("MEMRA_PRIME_DEQW")
7639 .map(|v| v != "0")
7640 .unwrap_or(true);
7641 if deqw {
7642 e.fa_prefill_view_ws(
7643 q,
7644 &k_view,
7645 &v_view,
7646 attn,
7647 head_dim,
7648 n_head,
7649 n_head_kv,
7650 t,
7651 t_kv,
7652 scale,
7653 true,
7654 kvl.k_tok_bytes,
7655 kvl.v_tok_bytes,
7656 crate::Engine::kv_fp8_on(),
7657 )?;
7658 } else {
7659 e.fa_prefill_view(
7660 q,
7661 &k_view,
7662 &v_view,
7663 attn,
7664 head_dim,
7665 n_head,
7666 n_head_kv,
7667 t,
7668 t_kv,
7669 scale,
7670 true,
7671 kvl.k_tok_bytes,
7672 kvl.v_tok_bytes,
7673 crate::Engine::kv_fp8_on(),
7674 )?;
7675 }
7676 Ok(())
7677 }
7678
7679 #[allow(clippy::type_complexity)] fn full_attn_prime_post_fa(
7683 &self,
7684 e: &Engine,
7685 attn: CudaSlice<f32>,
7686 gate: &Option<CudaSlice<f32>>,
7687 t: usize,
7688 n_head: usize,
7689 head_dim: usize,
7690 ) -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
7691 let (attn_g, ag16) = match gate {
7692 Some(gate) => {
7693 let n = t * n_head * head_dim;
7694 let mut ag = e.uninit(n)?;
7695 if Self::f16out_on(e, t) {
7696 let mut a16 = e.alloc_u8_uninit(n * 2)?;
7697 e.sig_mul_f16out(&attn, gate, &mut ag, &mut a16, n)?;
7698 (ag, Some(a16))
7699 } else {
7700 let mut gsig = e.uninit(n)?;
7701 e.sigmoid(gate, &mut gsig, n)?;
7702 e.mul(&attn, &gsig, &mut ag, n)?;
7703 (ag, None)
7704 }
7705 }
7706 None => (attn, None),
7707 };
7708 Ok((attn_g, ag16))
7709 }
7710
7711 #[allow(clippy::too_many_arguments)] fn linear_attn_prime(
7719 &self,
7720 e: &Engine,
7721 la: &LinearAttnLayer,
7722 h: &CudaSlice<f32>,
7723 hx: Option<&CudaSlice<u8>>,
7724 t: usize,
7725 cache: &mut Cache,
7726 il: usize,
7727 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7728 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
7730 let g4 = match hx {
7731 Some(xh) => e.matmul_group_xh(&ws, h, xh, t)?,
7732 None => e.matmul_group(&ws, h, t)?,
7733 };
7734 self.linear_attn_prime_core(e, la, g4, t, cache, il)
7735 }
7736
7737 fn linear_attn_prime_core(
7739 &self,
7740 e: &Engine,
7741 la: &LinearAttnLayer,
7742 mut g4: Vec<CudaSlice<f32>>,
7743 t: usize,
7744 cache: &mut Cache,
7745 il: usize,
7746 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7747 self.linear_attn_prime_core_pad(e, la, std::mem::take(&mut g4), t, cache, il, None)
7748 }
7749
7750 #[allow(clippy::too_many_arguments)]
7754 #[allow(clippy::type_complexity)] fn linear_attn_prime_core_pad_inner(
7756 &self,
7757 e: &Engine,
7758 la: &LinearAttnLayer,
7759 mut g4: Vec<CudaSlice<f32>>,
7760 t: usize,
7761 cache: &mut Cache,
7762 il: usize,
7763 pad_len: Option<&CudaSlice<i32>>,
7764 ) -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
7765 let geometry = la.geometry;
7767 let d_state = geometry.key_head_dim as usize;
7768 let num_k = geometry.key_heads as usize;
7769 let num_v = geometry.value_heads as usize;
7770 let key_dim = d_state * num_k;
7771 let value_dim = geometry.value_head_dim as usize * num_v;
7772 let conv_dim = key_dim * 2 + value_dim;
7773 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(
7778 e,
7779 la,
7780 &qkv_mixed.slice(0..t * conv_dim),
7781 &z.slice(0..t * value_dim),
7782 &beta_raw.slice(0..t * num_v),
7783 &alpha.slice(0..t * num_v),
7784 t,
7785 cache,
7786 il,
7787 pad_len,
7788 )
7789 }
7790
7791 #[allow(clippy::too_many_arguments)]
7794 fn linear_attn_gdn_prep(
7795 &self,
7796 e: &Engine,
7797 la: &LinearAttnLayer,
7798 qkv_mixed: &cudarc::driver::CudaView<f32>,
7799 beta_raw: &cudarc::driver::CudaView<f32>,
7800 alpha: &cudarc::driver::CudaView<f32>,
7801 t: usize,
7802 cache: &mut Cache,
7803 il: usize,
7804 pad_len: Option<&CudaSlice<i32>>,
7805 ) -> Result<GdnPrep, Box<dyn std::error::Error>> {
7806 let cfg = &self.cfg;
7807 let geometry = la.geometry;
7808 let d_state = geometry.key_head_dim as usize;
7809 let num_k = geometry.key_heads as usize;
7810 let num_v = geometry.value_heads as usize;
7811 let d_conv = geometry.conv_kernel as usize;
7812 let key_dim = d_state * num_k; let value_dim = geometry.value_head_dim as usize * num_v;
7814 let conv_dim = key_dim * 2 + value_dim; let eps = cfg.rms_eps;
7816 debug_assert!(
7817 t >= d_conv - 1,
7818 "stateful conv needs T >= pad (PRIME_MIN_T gates)"
7819 );
7820
7821 let rl = cache.recur[il].as_mut().unwrap();
7826 let hk = Self::gdn_hk(e, t, num_v, num_k);
7827 let conv_fuse = std::env::var("MEMRA_CONV_FUSE").as_deref() != Ok("0");
7828 let hk = if conv_fuse { hk } else { num_v }; let mut q_g = e.uninit(d_state * hk * t)?;
7830 let mut k_g = e.uninit(d_state * hk * t)?;
7831 let mut v_g = e.uninit(d_state * num_v * t)?;
7832 if conv_fuse {
7833 e.ssm_conv1d_gdn_state_pad(
7834 qkv_mixed,
7835 &mut rl.conv_state,
7836 la.ssm_conv1d.float_data(),
7837 &mut q_g,
7838 &mut k_g,
7839 &mut v_g,
7840 conv_dim,
7841 t,
7842 d_conv,
7843 d_state,
7844 num_v,
7845 num_k,
7846 key_dim,
7847 hk,
7848 pad_len,
7849 )?;
7850 } else {
7851 let mut conv_out = e.uninit(conv_dim * t)?; e.ssm_conv1d_tm_state_pad_v(
7853 qkv_mixed,
7854 &mut rl.conv_state,
7855 la.ssm_conv1d.float_data(),
7856 &mut conv_out,
7857 conv_dim,
7858 t,
7859 d_conv,
7860 pad_len,
7861 )?;
7862 e.qkv_to_gdn_repack(
7863 &conv_out, &mut q_g, &mut k_g, &mut v_g, d_state, num_v, num_k, key_dim, t,
7864 )?;
7865 }
7866 let mut q_l2 = e.uninit(d_state * hk * t)?;
7867 let qb16 = if Engine::l2_v2_on(d_state) && e.gdn_wgmma_on(32) {
7871 let mut qb = e.alloc_u8_uninit(d_state * hk * t * 2)?;
7872 e.l2_norm_pp(&q_g, &mut q_l2, Some(&mut qb), d_state, hk * t, eps)?;
7873 Some(qb)
7874 } else {
7875 e.l2_norm_pp(&q_g, &mut q_l2, None, d_state, hk * t, eps)?;
7876 None
7877 };
7878 let mut k_l2 = e.uninit(d_state * hk * t)?;
7879 let kb16 = if Engine::l2_v2_on(d_state) {
7881 let mut kb = e.alloc_u8_uninit(d_state * hk * t * 2)?;
7882 e.l2_norm_pp(&k_g, &mut k_l2, Some(&mut kb), d_state, hk * t, eps)?;
7883 Some(kb)
7884 } else {
7885 e.l2_norm_pp(&k_g, &mut k_l2, None, d_state, hk * t, eps)?;
7886 None
7887 };
7888 let mut beta = e.uninit(t * num_v)?;
7889 e.sigmoid_v(beta_raw, &mut beta, t * num_v)?;
7890 let mut g_log = e.uninit(t * num_v)?;
7891 e.gdn_glog_v(
7892 alpha,
7893 la.ssm_dt.float_data(),
7894 la.ssm_a.float_data(),
7895 &mut g_log,
7896 num_v,
7897 t,
7898 )?;
7899 if let Some(len_d) = pad_len {
7900 e.gdn_pad_mask(&mut beta, &mut g_log, len_d, num_v, t)?;
7901 }
7902 Ok(GdnPrep {
7903 hk,
7904 q_l2,
7905 k_l2,
7906 v_g,
7907 beta,
7908 g_log,
7909 kb16,
7910 qb16,
7911 })
7912 }
7913
7914 #[allow(clippy::too_many_arguments)]
7919 #[allow(clippy::type_complexity)] fn linear_attn_prime_core_batch(
7921 &self,
7922 e: &Engine,
7923 la: &LinearAttnLayer,
7924 g4: &[CudaSlice<f32>],
7925 offs: &[usize],
7926 ts: &[usize],
7927 caches: &mut [&mut Cache],
7928 il: usize,
7929 ) -> Result<Vec<(CudaSlice<f32>, Option<CudaSlice<u8>>)>, Box<dyn std::error::Error>> {
7930 let geometry = la.geometry;
7931 let d_state = geometry.key_head_dim as usize;
7932 let num_k = geometry.key_heads as usize;
7933 let num_v = geometry.value_heads as usize;
7934 let d_conv = geometry.conv_kernel as usize;
7935 let key_dim = d_state * num_k;
7936 let value_dim = geometry.value_head_dim as usize * num_v;
7937 let conv_dim = key_dim * 2 + value_dim;
7938 let eps = self.cfg.rms_eps;
7939 let scale = 1.0 / (d_state as f32).sqrt();
7940 let b = ts.len();
7941 let c = Engine::gdn_chunk_size();
7942 let carried = caches.iter().any(|c| c.pos > 0);
7945 let use_vl = !carried
7946 && (2..=8).contains(&b)
7947 && Engine::gdn_chunked_enabled()
7948 && ts.iter().all(|&t| t >= 16)
7949 && e.gdn_mma_enabled(c)
7950 && std::env::var("MEMRA_GDN_VL").as_deref() != Ok("0");
7951 if !use_vl {
7952 return (0..b)
7953 .map(|s| {
7954 let (o, t) = (offs[s], ts[s]);
7955 self.linear_attn_prime_core_pad_view(
7956 e,
7957 la,
7958 &g4[0].slice(o * conv_dim..(o + t) * conv_dim),
7959 &g4[1].slice(o * value_dim..(o + t) * value_dim),
7960 &g4[2].slice(o * num_v..(o + t) * num_v),
7961 &g4[3].slice(o * num_v..(o + t) * num_v),
7962 t,
7963 caches[s],
7964 il,
7965 None,
7966 )
7967 })
7968 .collect();
7969 }
7970 struct SeqBufs {
7974 conv_out: CudaSlice<f32>,
7975 q_g: CudaSlice<f32>,
7976 k_g: CudaSlice<f32>,
7977 v_g: CudaSlice<f32>,
7978 q_l2: CudaSlice<f32>,
7979 k_l2: CudaSlice<f32>,
7980 beta: CudaSlice<f32>,
7981 g_log: CudaSlice<f32>,
7982 gn: CudaSlice<f32>,
7983 gn16: CudaSlice<u8>,
7984 }
7985 let f16o = Self::f16out_on(e, 16);
7986 let hk = Self::gdn_hk(e, 16, num_v, num_k); let mut sb = Vec::with_capacity(b);
7988 let mut pres = Vec::with_capacity(b);
7989 for &t in ts.iter().take(b) {
7990 sb.push(SeqBufs {
7991 conv_out: e.uninit(conv_dim * t)?,
7992 q_g: e.uninit(d_state * hk * t)?,
7993 k_g: e.uninit(d_state * hk * t)?,
7994 v_g: e.uninit(d_state * num_v * t)?,
7995 q_l2: e.uninit(d_state * hk * t)?,
7996 k_l2: e.uninit(d_state * hk * t)?,
7997 beta: e.uninit(t * num_v)?,
7998 g_log: e.uninit(t * num_v)?,
7999 gn: e.uninit(d_state * num_v * t)?,
8000 gn16: e.alloc_u8_uninit(d_state * num_v * t * 2)?,
8001 });
8002 pres.push(e.gdn_chunk_alloc(num_v, t, c, hk)?);
8003 }
8004 let prep_args: Vec<crate::GdnPrepVl> = (0..b)
8005 .map(|s| {
8006 let (o, t) = (offs[s], ts[s]);
8007 let rl = caches[s].recur[il].as_ref().unwrap();
8008 crate::GdnPrepVl {
8009 qkv: e.addr_f32v(&g4[0].slice(o * conv_dim..(o + t) * conv_dim)),
8010 conv_state: e.addr_f32(&rl.conv_state),
8011 conv_out: e.addr_f32(&sb[s].conv_out),
8012 q_g: e.addr_f32(&sb[s].q_g),
8013 k_g: e.addr_f32(&sb[s].k_g),
8014 v_g: e.addr_f32(&sb[s].v_g),
8015 q_l2: e.addr_f32(&sb[s].q_l2),
8016 k_l2: e.addr_f32(&sb[s].k_l2),
8017 beta_raw: e.addr_f32v(&g4[2].slice(o * num_v..(o + t) * num_v)),
8018 alpha: e.addr_f32v(&g4[3].slice(o * num_v..(o + t) * num_v)),
8019 beta: e.addr_f32(&sb[s].beta),
8020 g_log: e.addr_f32(&sb[s].g_log),
8021 o: e.addr_f32(&pres[s].o),
8022 z: e.addr_f32v(&g4[1].slice(o * value_dim..(o + t) * value_dim)),
8023 gn: e.addr_f32(&sb[s].gn),
8024 gn16: e.addr_u8(&sb[s].gn16),
8025 kb16: if Engine::l2_v2_on(d_state) {
8026 e.addr_u8(&pres[s].kb16)
8027 } else {
8028 0
8029 },
8030 qb16: if Engine::l2_v2_on(d_state) && e.gdn_wgmma_on(c) {
8031 e.addr_u8(&pres[s].qb16)
8032 } else {
8033 0
8034 },
8035 t: t as i32,
8036 pad: 0,
8037 }
8038 })
8039 .collect();
8040 let args: Vec<crate::GdnSeqVl> = (0..b)
8041 .map(|s| {
8042 let rl = caches[s].recur[il].as_ref().unwrap();
8043 crate::GdnSeqVl {
8044 kb16: e.addr_u8(&pres[s].kb16),
8045 gcum: e.addr_f32(&pres[s].gcum),
8046 beta: e.addr_f32(&sb[s].beta),
8047 u: e.addr_f32(&pres[s].u),
8048 wb16: e.addr_u8(&pres[s].wb16),
8049 y: e.addr_u8(&pres[s].y16),
8050 ssnap: e.addr_u8(&pres[s].ssnap16),
8051 state_in: e.addr_f32(&rl.ssm_state),
8052 state_out: e.addr_f32(&rl.ssm_state_alt),
8053 q: e.addr_f32(&sb[s].q_l2),
8054 p: e.addr_f32(&pres[s].p),
8055 o: e.addr_f32(&pres[s].o),
8056 k: e.addr_f32(&sb[s].k_l2),
8057 v: e.addr_f32(&sb[s].v_g),
8058 g: e.addr_f32(&sb[s].g_log),
8059 a: e.addr_f32(&pres[s].a),
8060 w: e.addr_f32(&pres[s].w),
8061 t: ts[s] as i32,
8062 nc: pres[s].nc as i32,
8063 }
8064 })
8065 .collect();
8066 e.gdn_prep_vl8(
8067 &prep_args,
8068 la.ssm_conv1d.float_data(),
8069 la.ssm_dt.float_data(),
8070 la.ssm_a.float_data(),
8071 conv_dim,
8072 d_conv,
8073 d_state,
8074 num_v,
8075 num_k,
8076 key_dim,
8077 hk,
8078 eps,
8079 )?;
8080 if !Engine::l2_v2_on(d_state) {
8083 e.gdn_mirror_vl8(&args, num_v, 0, hk)?;
8084 }
8085 let wq8: Option<crate::GdnWVl8> = if e.gdn_wgmma_on(c) {
8087 if !Engine::l2_v2_on(d_state) {
8089 for s in 0..b {
8090 e.f32_to_bf16_into(&sb[s].q_l2, &mut pres[s].qb16, d_state * hk * ts[s])?;
8091 }
8092 }
8093 let mut wa = [crate::GdnWVl::default(); 8];
8094 for s in 0..b {
8095 wa[s] = crate::GdnWVl {
8096 qb16: e.addr_u8(&pres[s].qb16),
8097 pb16: e.addr_u8(&pres[s].pb16),
8098 };
8099 }
8100 Some(crate::GdnWVl8(wa))
8101 } else {
8102 None
8103 };
8104 e.gdn_chunk_k123_vl8(&args, num_v, hk, wq8.as_ref())?;
8105 e.gdn_chunk_vl8(&args, num_v, scale, hk, wq8.as_ref())?;
8106 if f16o {
8107 e.gdn_tail_vl8(&prep_args, la.ssm_norm.float_data(), d_state, num_v, eps)?;
8108 }
8109 let mut out = Vec::with_capacity(b);
8111 for (s, bufs) in sb.into_iter().enumerate() {
8112 let rl = caches[s].recur[il].as_mut().unwrap();
8113 std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
8114 let (o, t) = (offs[s], ts[s]);
8115 let SeqBufs { mut gn, gn16, .. } = bufs;
8116 if f16o {
8117 out.push((gn, Some(gn16)));
8118 } else {
8119 let z_v = g4[1].slice(o * value_dim..(o + t) * value_dim);
8120 e.gated_rmsnorm_zv(
8121 &pres[s].o,
8122 la.ssm_norm.float_data(),
8123 &z_v,
8124 &mut gn,
8125 d_state,
8126 num_v * t,
8127 eps,
8128 )?;
8129 out.push((gn, None));
8130 }
8131 }
8132 Ok(out)
8133 }
8134
8135 #[allow(clippy::too_many_arguments)]
8139 #[allow(clippy::type_complexity)] fn linear_attn_prime_core_pad_view(
8141 &self,
8142 e: &Engine,
8143 la: &LinearAttnLayer,
8144 qkv_mixed: &cudarc::driver::CudaView<f32>,
8145 z: &cudarc::driver::CudaView<f32>,
8146 beta_raw: &cudarc::driver::CudaView<f32>,
8147 alpha: &cudarc::driver::CudaView<f32>,
8148 t: usize,
8149 cache: &mut Cache,
8150 il: usize,
8151 pad_len: Option<&CudaSlice<i32>>,
8152 ) -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
8153 let cfg = &self.cfg;
8154 let geometry = la.geometry;
8155 let d_state = geometry.key_head_dim as usize;
8156 let num_v = geometry.value_heads as usize;
8157 let eps = cfg.rms_eps;
8158 let scale = 1.0 / (d_state as f32).sqrt();
8159
8160 let prep =
8161 self.linear_attn_gdn_prep(e, la, qkv_mixed, beta_raw, alpha, t, cache, il, pad_len)?;
8162
8163 let mut o = e.uninit(d_state * num_v * t)?;
8169 let rl = cache.recur[il].as_mut().unwrap();
8170 {
8171 let crate::cache::RecurLayer {
8172 ssm_state,
8173 ssm_state_alt,
8174 ..
8175 } = rl;
8176 e.gdn_scan_prefill(
8177 &prep.q_l2,
8178 &prep.k_l2,
8179 &prep.v_g,
8180 &prep.g_log,
8181 &prep.beta,
8182 prep.kb16.as_ref(),
8183 prep.qb16.as_ref(),
8184 ssm_state,
8185 ssm_state_alt,
8186 &mut o,
8187 num_v,
8188 t,
8189 scale,
8190 prep.hk,
8191 )?;
8192 }
8193 std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
8194
8195 let mut gn = e.uninit(d_state * num_v * t)?;
8198 let gn16 = if Self::f16out_on(e, t) {
8199 let mut g16 = e.alloc_u8_uninit(d_state * num_v * t * 2)?;
8200 e.gated_rmsnorm_f16out_zv(
8201 &o,
8202 la.ssm_norm.float_data(),
8203 z,
8204 &mut gn,
8205 &mut g16,
8206 d_state,
8207 num_v * t,
8208 eps,
8209 )?;
8210 Some(g16)
8211 } else {
8212 e.gated_rmsnorm_zv(
8213 &o,
8214 la.ssm_norm.float_data(),
8215 z,
8216 &mut gn,
8217 d_state,
8218 num_v * t,
8219 eps,
8220 )?;
8221 None
8222 };
8223 Ok((gn, gn16))
8224 }
8225
8226 #[allow(clippy::too_many_arguments)]
8228 fn linear_attn_prime_core_pad(
8229 &self,
8230 e: &Engine,
8231 la: &LinearAttnLayer,
8232 g4: Vec<CudaSlice<f32>>,
8233 t: usize,
8234 cache: &mut Cache,
8235 il: usize,
8236 pad_len: Option<&CudaSlice<i32>>,
8237 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8238 let (gn, gn16) = self.linear_attn_prime_core_pad_inner(e, la, g4, t, cache, il, pad_len)?;
8239 if let Some(xh) = &gn16
8240 && let Some(y) = e.try_f16_gemm_pre(&la.ssm_out, xh, t)?
8241 {
8242 return Ok(y);
8243 }
8244 e.matmul(&la.ssm_out, &gn, t)
8245 }
8246
8247 pub fn full_attn(
8252 &self,
8253 e: &Engine,
8254 fa: &FullAttnLayer,
8255 h: &CudaSlice<f32>,
8256 pos_d: &CudaSlice<i32>,
8257 t: usize,
8258 il: usize,
8259 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8260 if self.uses_sliding_gated_moe_program() {
8261 return self.step35_attn(e, fa, h, pos_d, t, il);
8262 }
8263 let cfg = &self.cfg;
8264 let _n_embd = cfg.n_embd as usize;
8265 let geometry = cfg.full_attention_geometry_at(il as u32);
8266 let n_head = geometry.n_head as usize;
8267 let n_head_kv = geometry.n_head_kv as usize;
8268 let head_dim = geometry.head_dim_k as usize;
8269 let eps = cfg.rms_eps;
8270 let scale = geometry.attention_scale();
8271
8272 let gated = geometry.attention_gate == memra_gguf::config::AttentionGateKind::FusedQ;
8275 let mut g3 = match self.full_attn_tp_qkv(e, fa, h, t)? {
8279 Some(g3) => g3,
8280 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
8281 };
8282 let v = g3.pop().unwrap();
8283 let mut k = g3.pop().unwrap();
8284 let qf = g3.pop().unwrap();
8285 let (mut q, gate) = if gated {
8286 let mut q = e.uninit(t * n_head * head_dim)?;
8287 let mut gate = e.uninit(t * n_head * head_dim)?;
8288 e.q_gate_split(&qf, &mut q, &mut gate, head_dim, n_head, t)?;
8289 (q, Some(gate))
8290 } else {
8291 (qf, None)
8292 };
8293
8294 let mut qn = e.uninit(t * n_head * head_dim)?;
8296 e.rms_norm(
8297 &q,
8298 fa.q_norm.float_data(),
8299 &mut qn,
8300 head_dim,
8301 n_head * t,
8302 eps,
8303 )?;
8304 q = qn;
8305 let mut kn = e.uninit(t * n_head_kv * head_dim)?;
8306 e.rms_norm(
8307 &k,
8308 fa.k_norm.float_data(),
8309 &mut kn,
8310 head_dim,
8311 n_head_kv * t,
8312 eps,
8313 )?;
8314 k = kn;
8315 let rope_dims = geometry.n_rot as usize;
8316 e.rope_neox(
8317 &mut q,
8318 pos_d,
8319 head_dim,
8320 rope_dims,
8321 n_head,
8322 t,
8323 geometry.rope_base,
8324 1.0,
8325 )?;
8326 e.rope_neox(
8327 &mut k,
8328 pos_d,
8329 head_dim,
8330 rope_dims,
8331 n_head_kv,
8332 t,
8333 geometry.rope_base,
8334 1.0,
8335 )?;
8336
8337 let mut attn = e.uninit(t * n_head * head_dim)?;
8339 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
8342 e.sdpa_naive(
8344 &q, &k, &v, &mut attn, head_dim, n_head, n_head_kv, t, t, scale, true,
8345 )?;
8346 } else {
8347 e.fa_prefill(
8348 &q, &k, &v, &mut attn, head_dim, n_head, n_head_kv, t, t, scale, true,
8349 )?;
8350 }
8351
8352 let attn_g = match &gate {
8354 Some(gate) => {
8355 let mut gsig = e.uninit(t * n_head * head_dim)?;
8356 e.sigmoid(gate, &mut gsig, t * n_head * head_dim)?;
8357 let mut ag = e.uninit(t * n_head * head_dim)?;
8358 e.mul(&attn, &gsig, &mut ag, t * n_head * head_dim)?;
8359 ag
8360 }
8361 None => attn,
8362 };
8363
8364 self.full_attn_o(e, fa, &attn_g, t)
8366 }
8367
8368 fn mla_split_operand<'w>(
8376 w: &'w crate::model::GpuTensor,
8377 name: &str,
8378 il: usize,
8379 ) -> &'w CudaSlice<f32> {
8380 match w {
8381 crate::model::GpuTensor::Float { data, .. } => data,
8382 _ => panic!(
8383 "layer {il}: MLA conversion-split operand {name} is not f32-resident. The 3D \
8384 (d_nope|kv_rank, kv_rank|d_v, n_head) splits have no quantized resident layout: \
8385 a quantized 3D tensor mis-derives row_bytes in the generic 2D Quant arm, so the \
8386 source must dequantize the fused kv_b_proj (TensorTransform::SplitMlaKv). \
8387 Reaching this means both the loader rank guard and MlaAttnLayer::load's \
8388 residency audit were bypassed"
8389 ),
8390 }
8391 }
8392
8393 #[allow(clippy::too_many_arguments)]
8401 fn mla_attn_core(
8405 &self,
8406 e: &Engine,
8407 mla: &crate::hybrid::MlaAttnLayer,
8408 h: &CudaSlice<f32>,
8409 pos_d: &CudaSlice<i32>,
8410 t: usize,
8411 il: usize,
8412 latent: &mut CudaSlice<f32>,
8413 index_plane: Option<IndexerPlanes<'_>>,
8414 slot: usize,
8415 rows_exact: bool,
8416 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8417 let attn = self.mla_attn_core_pre_wo(
8418 e,
8419 mla,
8420 h,
8421 pos_d,
8422 t,
8423 il,
8424 latent,
8425 index_plane,
8426 slot,
8427 rows_exact,
8428 )?;
8429 if rows_exact {
8433 e.matmul_rows_exact(&mla.wo, &attn, t)
8434 } else {
8435 e.matmul(&mla.wo, &attn, t)
8436 }
8437 }
8438
8439 #[allow(clippy::too_many_arguments)]
8445 #[allow(clippy::too_many_arguments)]
8452 fn mla_attn_core_pre_wo(
8453 &self,
8454 e: &Engine,
8455 mla: &crate::hybrid::MlaAttnLayer,
8456 h: &CudaSlice<f32>,
8457 pos_d: &CudaSlice<i32>,
8458 t: usize,
8459 il: usize,
8460 latent: &mut CudaSlice<f32>,
8461 index_plane: Option<IndexerPlanes<'_>>,
8462 slot: usize,
8463 rows_exact: bool,
8464 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8465 if t == 1 && Engine::mla_seg_ws_on() {
8469 let g = mla.geom;
8470 let q_lora = mla.wq_b.in_features();
8471 let mut ws = e.mla_seg_ws_take(g.n_head, g.d_nope, g.d_rope, g.kv_rank, q_lora)?;
8472 let out = (|| -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8473 if MLA_SEG_WS_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed) == 0 {
8474 eprintln!(
8475 "[mla-seg-ws] engaged: the T=1 MLA PRE segment writes the session's \
8476 stable handoff buffers (MEMRA_MLA_SEG_WS=1)"
8477 );
8478 }
8479 self.mla_seg_pre(e, mla, h, pos_d, t, il, rows_exact, Some(&mut ws))?;
8480 let gathered = self.mla_seg_mid(
8481 e,
8482 mla,
8483 h,
8484 MlaMidIn {
8485 q_an: &ws.q_an,
8486 c_kv_n: &ws.c_kv_n,
8487 k_pe: &ws.k_pe,
8488 },
8489 latent,
8490 index_plane,
8491 t,
8492 il,
8493 slot,
8494 rows_exact,
8495 )?;
8496 self.mla_seg_post(
8497 e, mla, &ws.q_nope, &ws.q_pe, gathered, latent, t, il, slot, rows_exact,
8498 )
8499 })();
8500 e.mla_seg_ws_put(ws);
8501 return out;
8502 }
8503 let pre = self
8504 .mla_seg_pre(e, mla, h, pos_d, t, il, rows_exact, None)?
8505 .expect("the ws-free PRE segment always returns its own buffers");
8506 let gathered = self.mla_seg_mid(
8507 e,
8508 mla,
8509 h,
8510 MlaMidIn {
8511 q_an: &pre.q_an,
8512 c_kv_n: &pre.c_kv_n,
8513 k_pe: &pre.k_pe,
8514 },
8515 latent,
8516 index_plane,
8517 t,
8518 il,
8519 slot,
8520 rows_exact,
8521 )?;
8522 self.mla_seg_post(
8523 e,
8524 mla,
8525 &pre.q_nope,
8526 &pre.q_pe,
8527 gathered,
8528 latent,
8529 t,
8530 il,
8531 slot,
8532 rows_exact,
8533 )
8534 }
8535
8536 #[allow(clippy::too_many_arguments)]
8540 fn mla_seg_pre(
8541 &self,
8542 e: &Engine,
8543 mla: &crate::hybrid::MlaAttnLayer,
8544 h: &CudaSlice<f32>,
8545 pos_d: &CudaSlice<i32>,
8546 t: usize,
8547 il: usize,
8548 rows_exact: bool,
8549 ws: Option<&mut MlaSegWs>,
8550 ) -> Result<Option<MlaPreOut>, Box<dyn std::error::Error>> {
8551 let g = mla.geom;
8552 let cfg = &self.cfg;
8553 let eps = cfg.rms_eps;
8554 let base = cfg.rope_freq_base;
8555 let (nh, dn, dr, r) = (g.n_head, g.d_nope, g.d_rope, g.kv_rank);
8556 assert_eq!(
8557 g.latent_dim,
8558 r + dr,
8559 "layer {il}: MlaGeom latent_dim disagrees with kv_rank + d_rope"
8560 );
8561 let mm = |w: &crate::model::GpuTensor,
8566 x: &CudaSlice<f32>|
8567 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8568 if rows_exact {
8569 e.matmul_rows_exact(w, x, t)
8570 } else {
8571 e.matmul(w, x, t)
8572 }
8573 };
8574
8575 let q_lora = mla.wq_b.in_features();
8577 if let Some(ws) = ws {
8582 if t != 1 || ws.sig != (nh, dn, dr, r, q_lora) {
8583 return Err(format!(
8584 "layer {il}: the MLA segment workspace is sized {:?} for t=1 and this call is \
8585 {:?} at t={t}",
8586 ws.sig,
8587 (nh, dn, dr, r, q_lora)
8588 )
8589 .into());
8590 }
8591 let q_a = mm(&mla.wq_a, h)?;
8592 e.rms_norm(
8593 &q_a,
8594 mla.q_a_norm.float_data(),
8595 &mut ws.q_an,
8596 q_lora,
8597 t,
8598 eps,
8599 )?;
8600 let q = mm(&mla.wq_b, &ws.q_an)?;
8601 e.mla_split_latent(&q, &mut ws.q_nope, &mut ws.q_pe, t * nh, dn, dr)?;
8602 e.mla_rope_interleaved(&mut ws.q_pe, pos_d, t, nh, dr, base)?;
8603 let kv = mm(&mla.wkv_a, h)?;
8604 let mut c_kv = e.uninit(t * r)?;
8605 e.mla_split_latent(&kv, &mut c_kv, &mut ws.k_pe, t, r, dr)?;
8606 e.rms_norm(&c_kv, mla.kv_a_norm.float_data(), &mut ws.c_kv_n, r, t, eps)?;
8607 e.mla_rope_interleaved(&mut ws.k_pe, pos_d, t, 1, dr, base)?;
8608 return Ok(None);
8609 }
8610 let q_a = mm(&mla.wq_a, h)?;
8611 let mut q_an = e.uninit(t * q_lora)?;
8612 e.rms_norm(&q_a, mla.q_a_norm.float_data(), &mut q_an, q_lora, t, eps)?;
8613 let q = mm(&mla.wq_b, &q_an)?;
8614 let mut q_nope = e.uninit(t * nh * dn)?;
8617 let mut q_pe = e.uninit((t * nh * dr).max(1))?;
8618 e.mla_split_latent(&q, &mut q_nope, &mut q_pe, t * nh, dn, dr)?;
8619 e.mla_rope_interleaved(&mut q_pe, pos_d, t, nh, dr, base)?;
8622
8623 let kv = mm(&mla.wkv_a, h)?;
8625 let mut c_kv = e.uninit(t * r)?;
8626 let mut k_pe = e.uninit((t * dr).max(1))?;
8627 e.mla_split_latent(&kv, &mut c_kv, &mut k_pe, t, r, dr)?;
8628 let mut c_kv_n = e.uninit(t * r)?;
8629 e.rms_norm(&c_kv, mla.kv_a_norm.float_data(), &mut c_kv_n, r, t, eps)?;
8630 e.mla_rope_interleaved(&mut k_pe, pos_d, t, 1, dr, base)?;
8631
8632 Ok(Some(MlaPreOut {
8636 q_nope,
8637 q_pe,
8638 q_an,
8639 c_kv_n,
8640 k_pe,
8641 }))
8642 }
8643
8644 #[allow(clippy::too_many_arguments, clippy::type_complexity)]
8648 fn mla_seg_mid(
8649 &self,
8650 e: &Engine,
8651 mla: &crate::hybrid::MlaAttnLayer,
8652 h: &CudaSlice<f32>,
8653 planes: MlaMidIn<'_>,
8654 latent: &mut CudaSlice<f32>,
8655 index_plane: Option<IndexerPlanes<'_>>,
8656 t: usize,
8657 il: usize,
8658 slot: usize,
8659 rows_exact: bool,
8660 ) -> Result<Option<(CudaSlice<i32>, usize)>, Box<dyn std::error::Error>> {
8661 let g = mla.geom;
8662 let (dr, r) = (g.d_rope, g.kv_rank);
8663 e.mla_append_latent(latent, planes.c_kv_n, planes.k_pe, slot, t, r, dr)?;
8664 let q_an = planes.q_an;
8665 let gathered = match (&mla.index, index_plane) {
8666 (Some(indexer), Some(plane)) => {
8667 Some(self.mla_kpool_select(e, indexer, h, q_an, plane, t, slot, il, rows_exact)?)
8668 }
8669 (Some(_), None) => {
8670 return Err(format!(
8671 "layer {il} declares a DSA k-pool indexer but no indexer state plane was \
8672 supplied — the ModelPlan must declare StatePlan::LatentKvCache with a \
8673 non-zero index_width for it"
8674 )
8675 .into());
8676 }
8677 (None, _) => None,
8678 };
8679
8680 Ok(gathered)
8682 }
8683
8684 #[allow(clippy::too_many_arguments)]
8688 fn mla_seg_post(
8689 &self,
8690 e: &Engine,
8691 mla: &crate::hybrid::MlaAttnLayer,
8692 q_nope: &CudaSlice<f32>,
8693 q_pe: &CudaSlice<f32>,
8694 gathered: Option<(CudaSlice<i32>, usize)>,
8695 latent: &CudaSlice<f32>,
8696 t: usize,
8697 il: usize,
8698 slot: usize,
8699 rows_exact: bool,
8700 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8701 let g = mla.geom;
8702 let (nh, dn, dr, dv, r) = (g.n_head, g.d_nope, g.d_rope, g.d_v, g.kv_rank);
8703 let t_kv = slot + t;
8704 let wk_b = Self::mla_split_operand(&mla.wk_b, "attn_k_b", il);
8705 let wv_b = Self::mla_split_operand(&mla.wv_b, "attn_v_b", il);
8706
8707 if mla.tp_shard
8730 && gathered.is_some()
8731 && dr == 0
8732 && r == 512
8733 && t >= 16
8734 && !crate::portable_mma_gated()
8735 && mla_tc_prefill_enabled()
8736 {
8737 static TP_TC_DECLINE: std::sync::Once = std::sync::Once::new();
8738 TP_TC_DECLINE.call_once(|| {
8739 eprintln!(
8740 "[mla-tc-prefill] DECLINED on glm5-TP head shards: the door's gate ran \
8741 on full-head geometry; shards ride the f32 prefill kernels until the \
8742 TP composition gate lands (pin MEMRA_MLA_TC_PREFILL=0 to silence)"
8743 );
8744 });
8745 }
8746 if let Some((idx, slots)) = &gathered
8747 && dr == 0
8748 && r == 512
8749 && t >= 16
8750 && !rows_exact && !crate::portable_mma_gated()
8752 && !mla.tp_shard
8753 && mla_tc_prefill_enabled()
8754 && let Some(attn) = self.mla_tc_prefill_chain(
8755 e, wk_b, wv_b, q_nope, latent, idx, *slots, t, t_kv, nh, dn, dv, r, g.scale,
8756 )?
8757 {
8758 return Ok(attn);
8759 }
8760
8761 let mut q_lat = e.uninit(t * nh * r)?;
8762 e.mla_absorb_q(q_nope, wk_b, &mut q_lat, t, nh, dn, r)?;
8763 let mut o_lat = e.uninit(t * nh * r)?;
8764 match &gathered {
8765 Some((idx, slots)) => e.mla_attn_gathered(
8766 &q_lat, q_pe, latent, idx, &mut o_lat, nh, r, dr, t, *slots, g.scale,
8767 )?,
8768 None => e.mla_attn_absorbed(
8769 &q_lat, q_pe, latent, &mut o_lat, nh, r, dr, t, t_kv, g.scale,
8770 )?,
8771 }
8772 let mut attn = e.uninit(t * nh * dv)?;
8773 e.mla_decompress_v(&o_lat, wv_b, &mut attn, t, nh, dv, r)?;
8774
8775 Ok(attn)
8776 }
8777
8778 #[allow(clippy::too_many_arguments)]
8802 fn mla_tc_prefill_chain(
8803 &self,
8804 e: &Engine,
8805 wk_b: &CudaSlice<f32>,
8806 wv_b: &CudaSlice<f32>,
8807 q_nope: &CudaSlice<f32>,
8808 latent: &CudaSlice<f32>,
8809 idx: &CudaSlice<i32>,
8810 width: usize,
8811 t: usize,
8812 t_kv: usize,
8813 nh: usize,
8814 dn: usize,
8815 dv: usize,
8816 r: usize,
8817 scale: f32,
8818 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
8819 fn declined(stage: &str, m: usize, n: usize, k: usize, batch: usize) {
8822 type ShapeSet = std::collections::HashSet<(usize, usize, usize, usize)>;
8823 static SAID: std::sync::Mutex<Option<ShapeSet>> = std::sync::Mutex::new(None);
8824 let mut g = SAID.lock().unwrap();
8825 if g.get_or_insert_with(std::collections::HashSet::new)
8826 .insert((m, n, k, batch))
8827 {
8828 eprintln!(
8829 "[mla-tc-prefill] DECLINED at {stage} m={m} n={n} k={k} batch={batch} \
8830 (no cuBLASLt heuristic) — this call falls back to the f32 MLA kernels"
8831 );
8832 }
8833 }
8834 for (name, n) in [
8838 ("wk_b", nh * r * dn),
8839 ("wv_b", nh * dv * r),
8840 ("q_nope", t * nh * dn),
8841 ("latent", t_kv * r),
8842 ] {
8843 debug_assert!(
8844 n.is_multiple_of(4),
8845 "mla-tc-prefill: {name} elems {n} % 4 != 0"
8846 );
8847 let _ = (name, n);
8848 }
8849 let wk_bf = e.f32_to_bf16(wk_b, nh * r * dn)?;
8850 let wv_bf = e.f32_to_bf16(wv_b, nh * dv * r)?;
8851 let qn_bf = e.f32_to_bf16(q_nope, t * nh * dn)?;
8852 let mut q_lat_bf = e.alloc_u8_uninit(t * nh * r * 2)?;
8856 if !e.mla_bf16_gemm_sb_bf16out(
8857 &wk_bf,
8858 &qn_bf,
8859 &mut q_lat_bf,
8860 t,
8861 r,
8862 dn,
8863 nh * dn,
8864 dn,
8865 nh * r,
8866 r,
8867 nh,
8868 )? {
8869 declined("absorb", t, r, dn, nh);
8870 return Ok(None);
8871 }
8872 let cache_bf = e.f32_to_bf16(latent, t_kv * r)?;
8874 let mut o_lat = e.uninit(t * nh * r)?;
8875 e.mla_attn_gathered_tc(
8876 &q_lat_bf, &cache_bf, idx, &mut o_lat, nh, r, t, width, scale,
8877 )?;
8878 let o_bf = e.f32_to_bf16(&o_lat, t * nh * r)?;
8881 let mut attn = e.uninit(t * nh * dv)?;
8882 if !e.mla_bf16_gemm_sb_f32out(
8883 &wv_bf,
8884 &o_bf,
8885 &mut attn,
8886 t,
8887 dv,
8888 r,
8889 nh * r,
8890 r,
8891 nh * dv,
8892 dv,
8893 nh,
8894 )? {
8895 declined("decompress", t, dv, r, nh);
8896 return Ok(None);
8897 }
8898 crate::MLA_TC_PREFILL_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
8899 {
8900 static ANNOUNCED: std::sync::Once = std::sync::Once::new();
8901 ANNOUNCED.call_once(|| {
8902 eprintln!(
8903 "[mla-tc-prefill] engaged: absorb/decompress = strided-batched bf16 TC \
8904 GEMMs, attention = fa_mla_gathered_bf16 (t={t}, t_kv={t_kv}, nh={nh}, \
8905 width={width}); dispatches counted in MLA_TC_PREFILL_DISPATCHES"
8906 );
8907 });
8908 }
8909 Ok(Some(attn))
8910 }
8911
8912 #[allow(clippy::too_many_arguments)]
8914 fn mla_kpool_select(
8915 &self,
8916 e: &Engine,
8917 indexer: &crate::hybrid::MlaIndexer,
8918 h: &CudaSlice<f32>,
8919 q_resid: &CudaSlice<f32>,
8920 plane: IndexerPlanes<'_>,
8921 t: usize,
8922 slot: usize,
8923 il: usize,
8924 rows_exact: bool,
8925 ) -> Result<(CudaSlice<i32>, usize), Box<dyn std::error::Error>> {
8926 Self::mla_kpool_indices_ex(e, indexer, h, q_resid, plane, t, slot, rows_exact).map_err(
8927 |source| -> Box<dyn std::error::Error> {
8928 format!("layer {il}: DSA k-pool selection failed: {source}").into()
8929 },
8930 )
8931 }
8932
8933 #[allow(clippy::too_many_arguments)]
8950 pub fn mla_kpool_indices(
8951 e: &Engine,
8952 indexer: &crate::hybrid::MlaIndexer,
8953 h: &CudaSlice<f32>,
8954 q_resid: &CudaSlice<f32>,
8955 plane: IndexerPlanes<'_>,
8956 t: usize,
8957 slot: usize,
8958 ) -> Result<(CudaSlice<i32>, usize), Box<dyn std::error::Error>> {
8959 Self::mla_kpool_indices_ex(e, indexer, h, q_resid, plane, t, slot, false)
8960 }
8961
8962 #[allow(clippy::too_many_arguments)]
8967 pub fn mla_kpool_indices_ex(
8968 e: &Engine,
8969 indexer: &crate::hybrid::MlaIndexer,
8970 h: &CudaSlice<f32>,
8971 q_resid: &CudaSlice<f32>,
8972 plane: IndexerPlanes<'_>,
8973 t: usize,
8974 slot: usize,
8975 rows_exact: bool,
8976 ) -> Result<(CudaSlice<i32>, usize), Box<dyn std::error::Error>> {
8977 let mm = |w: &crate::model::GpuTensor,
8978 x: &CudaSlice<f32>|
8979 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8980 if rows_exact {
8981 e.matmul_rows_exact(w, x, t)
8982 } else {
8983 e.matmul(w, x, t)
8984 }
8985 };
8986 const INDEX_NORM_EPS: f32 = 1e-5;
8990
8991 let ig = indexer.geom;
8992 let d = ig.head_dim;
8993 let t_kv = slot + t;
8994 let IndexerPlanes {
8995 state: plane,
8996 pool_keys: pool_key_plane,
8997 ready: pools_ready,
8998 state_ring_rows,
8999 capacity_tokens,
9000 } = plane;
9001
9002 let ring = if state_ring_rows == 0 {
9010 0
9011 } else {
9012 state_ring_rows / ig.pool * ig.pool
9013 };
9014 if state_ring_rows > 0 && ring == 0 {
9015 return Err(format!(
9016 "indexer tail ring of {state_ring_rows} rows cannot hold one pool of {}; \
9017 raise MEMRA_DSA_INDEX_RING or set it to 0 for the flat plane",
9018 ig.pool
9019 )
9020 .into());
9021 }
9022 if *pools_ready > slot / ig.pool {
9027 return Err(format!(
9028 "resident k-pool key plane claims {} finished pools but the cache holds only {} \
9029 complete pools before this call ({slot} rows / pool {}) — a rewind reduced the \
9030 latent length without clamping index_pools_ready",
9031 *pools_ready,
9032 slot / ig.pool,
9033 ig.pool
9034 )
9035 .into());
9036 }
9037
9038 let k_raw = mm(&indexer.wk, h)?;
9041 let mut k_norm = e.uninit(t * d)?;
9042 e.layer_norm_bias(
9043 &k_raw,
9044 indexer.k_norm_w.float_data(),
9045 indexer.k_norm_b.float_data(),
9046 &mut k_norm,
9047 d,
9048 t,
9049 INDEX_NORM_EPS,
9050 )?;
9051 let gate = mm(&indexer.kpool_gate, h)?;
9052
9053 let n_pools = t_kv / ig.pool;
9059 let select_k = ig.select_k(n_pools);
9060 let width = ig.index_width(n_pools);
9061 let capacity_pools = capacity_tokens / ig.pool;
9069 let need = (capacity_pools * d).max(n_pools * d).max(1);
9070 if pool_key_plane.as_ref().is_none_or(|k| k.len() < need) {
9071 *pool_key_plane = Some(e.uninit(need)?);
9072 *pools_ready = 0;
9073 }
9074 let pool_keys = pool_key_plane
9075 .as_mut()
9076 .expect("resident pool-key plane just allocated");
9077
9078 let ape = indexer.kpool_ape.float_data();
9091 let mut cur = slot;
9092 let mut appended = 0usize;
9093 while appended < t {
9094 let take =
9095 crate::cache::index_ring_take(ring, ig.pool, *pools_ready, cur, t - appended)
9096 .ok_or_else(|| -> Box<dyn std::error::Error> {
9097 format!(
9098 "indexer tail ring lapped: {ring} rows cannot hold the {} rows still \
9099 owed to unbuilt pools at row {cur} (pools_ready {}, pool {}, slot \
9100 {slot}, t {t}). The pool-key plane was reset or the cache rewound \
9101 without clamping index_pools_ready, so rows this call must read were \
9102 already overwritten. Raise MEMRA_DSA_INDEX_RING, or set \
9103 MEMRA_DSA_INDEX_RING=0 for the flat plane",
9104 cur.saturating_sub((*pools_ready).saturating_mul(ig.pool)),
9105 *pools_ready,
9106 ig.pool
9107 )
9108 .into()
9109 })?;
9110 debug_assert!(take > 0 && appended + take <= t);
9111 e.mla_index_append(plane, &k_norm, &gate, appended, cur, take, d, d, ring)?;
9112 cur += take;
9113 appended += take;
9114 let ready_now = cur / ig.pool;
9115 e.mla_kpool_pool_keys(
9116 plane,
9117 ape,
9118 pool_keys,
9119 (*pools_ready).min(ready_now),
9120 ready_now,
9121 ig.pool,
9122 d,
9123 ring,
9124 )?;
9125 *pools_ready = ready_now;
9126 }
9127 debug_assert!(t == 0 || *pools_ready == n_pools);
9128 let pool_keys = &*pool_keys;
9129
9130 let q_index = mm(&indexer.wq_b, q_resid)?;
9132 let head_weights = mm(&indexer.weights_proj, h)?;
9133 let mut score = e.uninit((t * n_pools).max(1))?;
9134 e.mla_kpool_score(
9135 &q_index,
9136 pool_keys,
9137 &head_weights,
9138 &mut score,
9139 t,
9140 ig.heads,
9141 d,
9142 n_pools,
9143 ig.pool,
9144 slot,
9145 (d as f32).powf(-0.5),
9146 (ig.heads as f32).powf(-0.5),
9147 )?;
9148 let mut idx = e.uninit_i32(t * width)?;
9149 e.mla_kpool_select(
9150 &score,
9151 &mut idx,
9152 t,
9153 n_pools,
9154 ig.pool,
9155 select_k,
9156 width,
9157 slot,
9158 ig.always_select_tail,
9159 )?;
9160 Ok((idx, width))
9161 }
9162
9163 pub fn mla_attn(
9166 &self,
9167 e: &Engine,
9168 mla: &crate::hybrid::MlaAttnLayer,
9169 h: &CudaSlice<f32>,
9170 pos_d: &CudaSlice<i32>,
9171 t: usize,
9172 il: usize,
9173 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9174 if mla.tp.is_some() {
9175 return Err(format!(
9176 "layer {il}: MLA layer is glm5-TP-sharded (MEMRA_GLM5_TP): the stateless \
9177 mixer path is unwired for a head shard"
9178 )
9179 .into());
9180 }
9181 let mut latent = e.uninit(t * mla.geom.latent_dim)?;
9182 let mut index_plane = match mla.index.as_ref() {
9183 Some(indexer) => Some(e.uninit(t * indexer.geom.state_width())?),
9184 None => None,
9185 };
9186 let mut pool_keys = None;
9189 let mut pools_ready = 0usize;
9190 let planes = index_plane.as_mut().map(|state| IndexerPlanes {
9191 state,
9192 pool_keys: &mut pool_keys,
9193 ready: &mut pools_ready,
9194 state_ring_rows: 0,
9196 capacity_tokens: t,
9197 });
9198 self.mla_attn_core(e, mla, h, pos_d, t, il, &mut latent, planes, 0, false)
9199 }
9200
9201 #[allow(clippy::too_many_arguments)] pub fn mla_attn_cached(
9206 &self,
9207 e: &Engine,
9208 mla: &crate::hybrid::MlaAttnLayer,
9209 h: &CudaSlice<f32>,
9210 pos_d: &CudaSlice<i32>,
9211 t: usize,
9212 il: usize,
9213 cache: &mut Cache,
9214 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9215 self.mla_attn_cached_inner(e, mla, h, pos_d, t, il, cache, false)
9216 }
9217
9218 #[allow(clippy::too_many_arguments)] pub fn mla_attn_cached_rows_exact(
9229 &self,
9230 e: &Engine,
9231 mla: &crate::hybrid::MlaAttnLayer,
9232 h: &CudaSlice<f32>,
9233 pos_d: &CudaSlice<i32>,
9234 t: usize,
9235 il: usize,
9236 cache: &mut Cache,
9237 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9238 self.mla_attn_cached_inner(e, mla, h, pos_d, t, il, cache, true)
9239 }
9240
9241 #[allow(clippy::too_many_arguments)] fn mla_attn_cached_inner(
9243 &self,
9244 e: &Engine,
9245 mla: &crate::hybrid::MlaAttnLayer,
9246 h: &CudaSlice<f32>,
9247 pos_d: &CudaSlice<i32>,
9248 t: usize,
9249 il: usize,
9250 cache: &mut Cache,
9251 rows_exact: bool,
9252 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9253 if mla.tp.is_some() {
9258 return Err(format!(
9259 "layer {il}: MLA layer is glm5-TP-sharded (MEMRA_GLM5_TP): the plain mixer \
9260 path is unwired for a head shard — only the TP decode/prime walk may \
9261 execute it (rows_exact={rows_exact})"
9262 )
9263 .into());
9264 }
9265 let max_ctx = cache.max_ctx;
9268 let layer = cache.latent[il].as_mut().ok_or_else(|| {
9269 format!(
9270 "layer {il} is Mixer::Mla but the cache has no latent plane — the ModelPlan \
9271 must declare StatePlan::LatentKvCache for it"
9272 )
9273 })?;
9274 let attn =
9275 self.mla_attn_cached_pre_wo(e, mla, h, pos_d, t, il, layer, max_ctx, rows_exact)?;
9276 if rows_exact {
9279 e.matmul_rows_exact(&mla.wo, &attn, t)
9280 } else {
9281 e.matmul(&mla.wo, &attn, t)
9282 }
9283 }
9284
9285 #[allow(clippy::too_many_arguments)]
9290 pub(crate) fn mla_attn_cached_pre_wo(
9291 &self,
9292 e: &Engine,
9293 mla: &crate::hybrid::MlaAttnLayer,
9294 h: &CudaSlice<f32>,
9295 pos_d: &CudaSlice<i32>,
9296 t: usize,
9297 il: usize,
9298 layer: &mut memra_kv::LatentKvLayer,
9299 max_ctx: usize,
9300 rows_exact: bool,
9301 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9302 let slot = layer.len;
9303 let width = layer.width;
9304 assert_eq!(
9305 width, mla.geom.latent_dim,
9306 "layer {il}: cache latent width {width} != MlaGeom latent_dim {}",
9307 mla.geom.latent_dim
9308 );
9309 let capacity = layer.rows.len() / width;
9310 if slot + t > capacity {
9311 return Err(format!(
9312 "layer {il}: latent cache overflow — {slot} + {t} rows exceeds capacity {capacity}"
9313 )
9314 .into());
9315 }
9316 if mla.index.is_some() && layer.index_rows.is_none() {
9317 return Err(format!(
9318 "layer {il} loaded a DSA k-pool indexer but its latent cache carries no indexer \
9319 state plane — StatePlan::LatentKvCache declared index_width 0 for a layer whose \
9320 SparseIndexPlan is Own {{ kpool: Some(..) }}"
9321 )
9322 .into());
9323 }
9324 let mut rows = std::mem::replace(&mut layer.rows, e.uninit(0)?);
9329 let mut index_rows = layer.index_rows.take();
9330 let mut pool_keys = layer.index_pool_keys.take();
9331 let mut pools_ready = layer.index_pools_ready;
9332 let index_ring_rows = layer.index_ring_rows.unwrap_or(0);
9333 let planes = index_rows.as_mut().map(|state| IndexerPlanes {
9334 state,
9335 pool_keys: &mut pool_keys,
9336 ready: &mut pools_ready,
9337 state_ring_rows: index_ring_rows,
9338 capacity_tokens: max_ctx,
9339 });
9340 let out =
9341 self.mla_attn_core_pre_wo(e, mla, h, pos_d, t, il, &mut rows, planes, slot, rows_exact);
9342 layer.rows = rows;
9343 layer.index_rows = index_rows;
9344 layer.index_pool_keys = pool_keys;
9345 layer.index_pools_ready = if out.is_ok() {
9350 pools_ready
9351 } else if let Some(indexer) = mla.index.as_ref() {
9352 pools_ready.min(layer.len / indexer.geom.pool)
9353 } else {
9354 pools_ready
9355 };
9356 let out = out?;
9357 if let Some(indexer) = mla.index.as_ref() {
9362 let pool = indexer.geom.pool;
9363 if layer.index_pool != 0 && layer.index_pool != pool {
9364 return Err(format!(
9365 "layer {il}: resident indexer pool {} != loaded geometry pool {pool}",
9366 layer.index_pool,
9367 )
9368 .into());
9369 }
9370 layer.index_pool = pool;
9371 }
9372 layer.len = slot + t;
9373 let len_i32 = i32::try_from(layer.len).map_err(|_| "latent length exceeds i32 mirror")?;
9374 e.i32_mirror_store(&mut layer.len_d, len_i32)?;
9377 Ok(out)
9378 }
9379
9380 #[allow(clippy::too_many_arguments)]
9388 pub(crate) fn mla_tp_attn_cached(
9389 &self,
9390 e: &Engine,
9391 mla: &crate::hybrid::MlaAttnLayer,
9392 h: &CudaSlice<f32>,
9393 pos_d: &CudaSlice<i32>,
9394 t: usize,
9395 il: usize,
9396 cache: &mut Cache,
9397 rows_exact: bool,
9398 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9399 let tp = mla
9400 .tp
9401 .as_ref()
9402 .ok_or("mla_tp_attn_cached called on an unsharded layer")?;
9403 let rt = &tp.rt;
9404 let ranks = tp.ranks();
9405 let g = mla.geom; let hl = g.n_head;
9407 let dv = g.d_v;
9408 let full_heads = tp.full_heads;
9409 let n_embd = tp.n_embd;
9410 let hh = n_embd / ranks;
9411 let max_ctx = cache.max_ctx;
9412
9413 let hop = rt.hop(e);
9417 let h_peers = crate::tp_transport::fanout_f32(&hop, h, h.len())?;
9418 let pos_peers = crate::tp_transport::fanout_i32(&hop, pos_d, pos_d.len())?;
9419
9420 {
9422 let canonical = cache.latent[il].as_ref().ok_or_else(|| {
9423 format!("layer {il}: glm5 TP MLA walk found no canonical latent plane")
9424 })?;
9425 crate::glm5_tp::ensure_mla_peer_latent(
9426 rt,
9427 canonical,
9428 &mut cache.glm5_tp_latent_peer[il],
9429 )?;
9430 }
9431
9432 let mut attn: Vec<Option<CudaSlice<f32>>> = (0..ranks).map(|_| None).collect();
9439 for r in 1..ranks {
9440 let layer = &mut cache.glm5_tp_latent_peer[il].as_mut().unwrap()[r - 1];
9441 attn[r] = Some(self.mla_attn_cached_pre_wo(
9442 &rt.peers[r - 1],
9443 &tp.peers[r - 1],
9444 &h_peers[r - 1],
9445 &pos_peers[r - 1],
9446 t,
9447 il,
9448 layer,
9449 max_ctx,
9450 rows_exact,
9451 )?);
9452 }
9453 attn[0] = {
9454 let layer = cache.latent[il].as_mut().unwrap();
9455 Some(self.mla_attn_cached_pre_wo(e, mla, h, pos_d, t, il, layer, max_ctx, rows_exact)?)
9456 };
9457
9458 let part = hl * dv;
9461 debug_assert_eq!(full_heads * dv, ranks * part);
9462 let attn_refs: Vec<&CudaSlice<f32>> = attn
9463 .iter()
9464 .map(|a| a.as_ref().expect("filled above"))
9465 .collect();
9466 let fulls = crate::tp_transport::gather_parts(&hop, &attn_refs, t, part)?;
9467
9468 let mut ys = Vec::with_capacity(ranks);
9471 if rows_exact {
9472 ys.push(e.matmul_rows_exact(&mla.wo, &fulls[0], t)?);
9473 for r in 1..ranks {
9474 ys.push(rt.peers[r - 1].matmul_rows_exact(&tp.peers[r - 1].wo, &fulls[r], t)?);
9475 }
9476 } else {
9477 ys.push(e.matmul(&mla.wo, &fulls[0], t)?);
9478 for r in 1..ranks {
9479 ys.push(rt.peers[r - 1].matmul(&tp.peers[r - 1].wo, &fulls[r], t)?);
9480 }
9481 }
9482 debug_assert_eq!(n_embd, ranks * hh);
9484 let y_refs: Vec<&CudaSlice<f32>> = ys.iter().collect();
9485 crate::tp_transport::concat_parts_on_root(&hop, &y_refs, t, hh)
9486 }
9487
9488 pub fn linear_attn(
9490 &self,
9491 e: &Engine,
9492 la: &LinearAttnLayer,
9493 h: &CudaSlice<f32>,
9494 t: usize,
9495 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9496 let cfg = &self.cfg;
9497 let _n_embd = cfg.n_embd as usize;
9498 let geometry = la.geometry;
9499 let d_state = geometry.key_head_dim as usize;
9500 let num_k = geometry.key_heads as usize;
9501 let num_v = geometry.value_heads as usize;
9502 let d_conv = geometry.conv_kernel as usize;
9503 let head_k = d_state;
9504 let head_v = geometry.value_head_dim as usize;
9505 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;
9509 let scale = 1.0 / (d_state as f32).sqrt();
9510
9511 let mut g4 = e.matmul_group(
9514 &[&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha],
9515 h,
9516 t,
9517 )?;
9518 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);
9530 let mut q_g = e.uninit(d_state * num_v * t)?;
9531 let mut k_g = e.uninit(d_state * num_v * t)?;
9532 let mut v_g = e.uninit(d_state * num_v * t)?;
9533 e.ssm_conv1d_gdn(
9534 &qkv_mixed,
9535 la.ssm_conv1d.float_data(),
9536 &mut q_g,
9537 &mut k_g,
9538 &mut v_g,
9539 conv_dim,
9540 t,
9541 d_conv,
9542 d_state,
9543 num_v,
9544 num_k,
9545 key_dim,
9546 )?;
9547 let mut q_l2 = e.uninit(d_state * num_v * t)?;
9549 e.l2_norm(&q_g, &mut q_l2, d_state, num_v * t, eps)?;
9550 let mut k_l2 = e.uninit(d_state * num_v * t)?;
9551 e.l2_norm(&k_g, &mut k_l2, d_state, num_v * t, eps)?;
9552 let v_gd = v_g;
9553
9554 let mut beta = e.uninit(t * num_v)?;
9557 e.sigmoid(&beta_raw, &mut beta, t * num_v)?;
9558 let mut g_log = e.uninit(t * num_v)?;
9560 e.gdn_glog(
9561 &alpha,
9562 la.ssm_dt.float_data(),
9563 la.ssm_a.float_data(),
9564 &mut g_log,
9565 num_v,
9566 t,
9567 )?;
9568
9569 let state_in = e.zeros(d_state * d_state * num_v)?; let mut state_out = e.zeros(d_state * d_state * num_v)?;
9572 let mut o = e.uninit(d_state * num_v * t)?;
9573 e.gdn_scan_prefill(
9574 &q_l2,
9575 &k_l2,
9576 &v_gd,
9577 &g_log,
9578 &beta,
9579 None,
9580 None,
9581 &state_in,
9582 &mut state_out,
9583 &mut o,
9584 num_v,
9585 t,
9586 scale,
9587 num_v,
9588 )?;
9589
9590 let mut gn = e.uninit(d_state * num_v * t)?;
9595 e.gated_rmsnorm(
9596 &o,
9597 la.ssm_norm.float_data(),
9598 &z,
9599 &mut gn,
9600 d_state,
9601 num_v * t,
9602 eps,
9603 )?;
9604
9605 let out = e.matmul(&la.ssm_out, &gn, t)?;
9609 Ok(out)
9610 }
9611}
9612
9613impl HybridModel {
9614 pub fn moe_ffn_il(
9625 &self,
9626 e: &Engine,
9627 m: &MoeWeights,
9628 z: &CudaSlice<f32>,
9629 t: usize,
9630 il: u16,
9631 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9632 Self::moe_ffn_inner(
9633 e,
9634 m,
9635 z,
9636 None,
9637 t,
9638 &self.cfg,
9639 il,
9640 self.max_moe_block(),
9641 false,
9642 None,
9643 self.uses_sliding_gated_moe_program(),
9644 false,
9645 )
9646 }
9647
9648 pub fn moe_ffn_il_prefill(
9651 &self,
9652 e: &Engine,
9653 m: &MoeWeights,
9654 z: &CudaSlice<f32>,
9655 t: usize,
9656 il: u16,
9657 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9658 Self::moe_ffn_inner(
9659 e,
9660 m,
9661 z,
9662 None,
9663 t,
9664 &self.cfg,
9665 il,
9666 self.max_moe_block(),
9667 true,
9668 Some(&self.step_grouped_prefill),
9669 self.uses_sliding_gated_moe_program(),
9670 false,
9671 )
9672 }
9673
9674 pub fn moe_ffn_il_zq8(
9678 &self,
9679 e: &Engine,
9680 m: &MoeWeights,
9681 z: &CudaSlice<f32>,
9682 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
9683 t: usize,
9684 il: u16,
9685 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9686 Self::moe_ffn_inner(
9687 e,
9688 m,
9689 z,
9690 zq8,
9691 t,
9692 &self.cfg,
9693 il,
9694 self.max_moe_block(),
9695 false,
9696 None,
9697 self.uses_sliding_gated_moe_program(),
9698 false,
9699 )
9700 }
9701
9702 pub(crate) fn moe_ffn_il_zq8_vrows(
9707 &self,
9708 e: &Engine,
9709 m: &MoeWeights,
9710 z: &CudaSlice<f32>,
9711 t: usize,
9712 il: u16,
9713 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9714 Self::moe_ffn_inner(
9715 e,
9716 m,
9717 z,
9718 None,
9719 t,
9720 &self.cfg,
9721 il,
9722 self.max_moe_block(),
9723 false,
9724 None,
9725 self.uses_sliding_gated_moe_program(),
9726 true,
9727 )
9728 }
9729
9730 pub(crate) fn moe_ffn(
9738 e: &Engine,
9739 m: &MoeWeights,
9740 z: &CudaSlice<f32>,
9741 t: usize,
9742 cfg: &ModelConfig,
9743 il: u16,
9744 max_block: usize,
9745 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9746 Self::moe_ffn_inner(
9747 e, m, z, None, t, cfg, il, max_block, false, None, false, false,
9748 )
9749 }
9750
9751 #[allow(clippy::too_many_arguments)]
9752 #[allow(clippy::map_entry)] pub(crate) fn moe_ffn_inner(
9754 e: &Engine,
9755 m: &MoeWeights,
9756 z: &CudaSlice<f32>,
9757 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
9758 t: usize,
9759 cfg: &ModelConfig,
9760 il: u16,
9761 max_block: usize,
9762 prefill: bool,
9763 grouped_prefill: Option<&std::sync::Mutex<crate::hybrid::StepEpGroupedPrefill>>,
9764 sliding_gated_moe: bool,
9765 vrows: bool,
9766 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9767 let worker_io = crate::spill_pread::worker_enabled();
9768 let epoch_lfu = std::env::var_os("MEMRA_MOE_LFU_DECAY").is_some();
9769 if Engine::moe_cache_enabled() && (worker_io || epoch_lfu) {
9770 e.with_moe_cache(max_block, |cache, _| {
9771 cache.begin_forward_epoch(il, t);
9772 if worker_io {
9773 cache.begin_worker_scope();
9774 }
9775 Ok(())
9776 })?;
9777 }
9778 if let Some(ep) = &m.glm5_ep {
9779 return Self::moe_ffn_glm5_ep(e, m, ep, z, zq8, t, cfg, il, prefill);
9784 }
9785 if m.step_ep.is_some() || m.step_tp.is_some() {
9786 let moe = cfg
9787 .moe
9788 .as_ref()
9789 .ok_or("Step distributed execution requires MoE model metadata")?;
9790 let n_embd = cfg.n_embd as usize;
9791 let n_expert = moe.expert_count as usize;
9792 let n_used = moe.expert_used_count as usize;
9793 let sigmoid = cfg
9794 .sigmoid_router()
9795 .ok_or("Step distributed execution requires the Step sigmoid router")?;
9796 let logits = Self::moe_router_logits(e, m, z, t, cfg)?;
9797 Self::trace_sigmoid_router_logits(e, il, t, n_expert, n_used, &logits, m, sigmoid)?;
9798 let grouped_prefill_requested = prefill && step_ep_grouped_prefill_enabled()?;
9799 if grouped_prefill_requested && !step_tp_prefill_enabled()? {
9800 return Err(
9801 "MEMRA_STEP_EP_GROUPED_PREFILL=1 requires MEMRA_STEP_TP_PREFILL=1".into(),
9802 );
9803 }
9804 if grouped_prefill_requested && !step_grouped_prefill_shape(true, prefill, t) {
9805 return Err(format!(
9806 "Step grouped prefill tokens {t} are outside the qualified {}..={} range",
9807 PRIME_MIN_T,
9808 crate::cache::PRIME_CHUNK_MAX_TOKENS,
9809 )
9810 .into());
9811 }
9812 let grouped_decode_shape = step_grouped_decode_shape(prefill, t);
9813 let grouped_prefill_shape =
9814 step_grouped_prefill_shape(grouped_prefill_requested, prefill, t);
9815 if let Some(ep) = m.step_ep.as_ref().filter(|ep| {
9816 ep.grouped_decode.is_some() && (grouped_decode_shape || grouped_prefill_shape)
9817 }) {
9818 let (selected, route_weights) =
9819 Self::moe_route_sigmoid_cfg(e, &logits, t, n_expert, n_used, m, sigmoid)?;
9820 crate::moesd::record_host_routes(il, n_expert, n_used, &selected)?;
9821 Self::trace_moe_routes(il, t, &selected, &route_weights)?;
9822 Self::trace_moe_input(e, il, t, n_embd, z)?;
9823 let selected = selected
9824 .iter()
9825 .map(|&expert| expert as usize)
9826 .collect::<Vec<_>>();
9827
9828 e.stream().synchronize()?;
9831 let execute = |state: &mut crate::hybrid::StepEpGroupedDecode| {
9832 state.projection.set_activation_limit(ep.activation_limit)?;
9833 ep.runtime
9834 .refresh_step_grouped_expert_parallel_gate_from_root_device(
9835 ep.experts.e4m3()?,
9836 &mut state.projection,
9837 z,
9838 t,
9839 &selected,
9840 )?;
9841 ep.runtime.refresh_step_grouped_expert_parallel_combine(
9842 &state.projection,
9843 &mut state.combine,
9844 &route_weights,
9845 )?;
9846 ep.runtime.execute_step_grouped_expert_parallel_gate(
9847 ep.experts.e4m3()?,
9848 &mut state.projection,
9849 )?;
9850 ep.runtime.execute_step_grouped_expert_parallel_combine(
9851 &state.projection,
9852 &mut state.combine,
9853 )?;
9854 let mut output = ep.runtime.copy_step_grouped_expert_parallel_combine_root(
9855 &state.projection,
9856 &state.combine,
9857 e,
9858 )?;
9859 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut output)?;
9860 if prefill {
9861 e.stream().synchronize()?;
9864 }
9865 eprintln!(
9866 "[step-tp-ep-grouped] execute layer={il} tokens={t} devices={:?} \
9867 attention_layout=tensor-parallel expert_layout=expert-parallel \
9868 expert_transport={} native_p2p=true route_control=host-narrow \
9869 input=root-device projection_workspaces=persistent \
9870 combine=root-device output=owning-stage-device \
9871 prefill={prefill} batched_decode=false capacity={} \
9872 performance_claim=false",
9873 ep.devices,
9874 ep.runtime.transport_label(),
9875 state.projection.max_tokens(),
9876 );
9877 Ok::<_, Box<dyn std::error::Error>>(output)
9878 };
9879
9880 if grouped_prefill_shape {
9881 let grouped_prefill = grouped_prefill
9882 .ok_or("Step grouped prefill has no model-scoped executor")?;
9883 let mut shared = grouped_prefill
9884 .lock()
9885 .map_err(|_| "Step grouped prefill state lock is poisoned")?;
9886 let needs_prepare = shared.state.as_ref().is_none_or(|state| {
9887 state.devices != ep.devices
9888 || state.grouped.projection.max_tokens() < t
9889 || state.grouped.projection.input_width() != n_embd
9890 || state.grouped.projection.expert_width()
9891 != moe.expert_ff_length as usize
9892 });
9893 if needs_prepare {
9894 let seed_input = vec![0.0f32; n_embd];
9895 let seed_selected = &selected[..n_used];
9896 let seed_weights = &route_weights[..n_used];
9897 let projection = ep
9898 .runtime
9899 .prepare_step_grouped_expert_parallel_gate_with_capacity(
9900 ep.experts.e4m3()?,
9901 &seed_input,
9902 1,
9903 seed_selected,
9904 ep.activation_limit,
9905 t,
9906 )?;
9907 let combine = ep.runtime.prepare_step_grouped_expert_parallel_combine(
9908 &projection,
9909 seed_weights,
9910 )?;
9911 shared.state = Some(crate::hybrid::StepEpGroupedPrefillState {
9912 devices: ep.devices.clone(),
9913 grouped: crate::hybrid::StepEpGroupedDecode {
9914 projection,
9915 combine,
9916 },
9917 });
9918 eprintln!(
9919 "[step-tp-ep-grouped-prefill] prepare capacity={t} devices={:?} \
9920 shared_across_layers=true performance_claim=false",
9921 ep.devices,
9922 );
9923 }
9924 return execute(
9925 &mut shared
9926 .state
9927 .as_mut()
9928 .expect("Step grouped prefill state prepared above")
9929 .grouped,
9930 );
9931 }
9932
9933 let mut grouped = ep
9934 .grouped_decode
9935 .as_ref()
9936 .expect("grouped decode presence checked above")
9937 .lock()
9938 .map_err(|_| "Step grouped decode state lock is poisoned")?;
9939 return execute(&mut grouped);
9940 }
9941 if grouped_prefill_shape {
9942 return Err(
9943 "Step grouped prefill requires native-P2P expert-owner device arithmetic"
9944 .into(),
9945 );
9946 }
9947 if t >= 16
9962 && crate::step_gemm_prime_on()
9963 && let Some(tp) = &m.step_tp
9964 && let crate::hybrid::StepTpExpertBank::Nvfp4(bank) = &tp.experts
9965 {
9966 let mprof = std::env::var("MEMRA_PRIME_PROF").as_deref() == Ok("1") && t >= 16;
9974 let mut mt = std::time::Instant::now();
9975 let (selected, route_weights) =
9976 Self::moe_route_sigmoid_cfg(e, &logits, t, n_expert, n_used, m, sigmoid)?;
9977 let sel_i32: Vec<i32> = selected.iter().map(|&x| x as i32).collect();
9978 let d_router = if mprof {
9979 let _ = e.stream().synchronize();
9980 let v = mt.elapsed().as_secs_f64() * 1e3;
9981 mt = std::time::Instant::now();
9982 v
9983 } else {
9984 0.0
9985 };
9986 let mdet =
10000 std::env::var("MEMRA_MOE_DETERM").as_deref() == Ok("1") && t >= 16 && il < 4;
10001 if mdet {
10002 let a = tp.runtime.run_tensor_parallel_routes_nvfp4_prime_grouped(
10003 bank,
10004 e,
10005 z,
10006 t,
10007 &sel_i32,
10008 &route_weights,
10009 n_used,
10010 tp.activation_limit,
10011 )?;
10012 let b = tp.runtime.run_tensor_parallel_routes_nvfp4_prime_grouped(
10013 bank,
10014 e,
10015 z,
10016 t,
10017 &sel_i32,
10018 &route_weights,
10019 n_used,
10020 tp.activation_limit,
10021 )?;
10022 let (ha, hb) = (e.dtoh(&a)?, e.dtoh(&b)?);
10023 let mut md = 0.0f32;
10024 let mut ndiff = 0usize;
10025 for (x, y) in ha.iter().zip(hb.iter()) {
10026 let d = (x - y).abs();
10027 if d > 0.0 {
10028 ndiff += 1;
10029 }
10030 if d > md {
10031 md = d;
10032 }
10033 }
10034 eprintln!(
10035 "[moe-determ] il={il} t={t} maxdiff={md:.3e} \
10036 differing={ndiff}/{} -> {}",
10037 ha.len(),
10038 if ndiff == 0 {
10039 "IDENTICAL"
10040 } else {
10041 "NONDETERMINISTIC"
10042 }
10043 );
10044 }
10045 let mut output = tp.runtime.run_tensor_parallel_routes_nvfp4_prime_grouped(
10046 bank,
10047 e,
10048 z,
10049 t,
10050 &sel_i32,
10051 &route_weights,
10052 n_used,
10053 tp.activation_limit,
10054 )?;
10055 let d_gemm = if mprof {
10056 let _ = e.stream().synchronize();
10057 let v = mt.elapsed().as_secs_f64() * 1e3;
10058 mt = std::time::Instant::now();
10059 v
10060 } else {
10061 0.0
10062 };
10063 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut output)?;
10064 if mprof {
10065 let _ = e.stream().synchronize();
10066 let d_shared = mt.elapsed().as_secs_f64() * 1e3;
10067 eprintln!(
10071 "[moe-prof] il={il} t={t} router={d_router:.1}ms \
10072 gemm={d_gemm:.1}ms shared={d_shared:.1}ms"
10073 );
10074 }
10075 return Ok(output);
10076 }
10077 if t == 1
10078 && crate::tp::step_nvfp4_dev_routes_enabled()?
10079 && crate::tp::step_tp_dev_router_enabled()?
10080 && let Some(tp) = &m.step_tp
10081 && let crate::hybrid::StepTpExpertBank::Nvfp4(bank) = &tp.experts
10082 {
10083 let (sf, route_norm) = sigmoid;
10084 static D1_ROUTER: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10091 let d1_router = *D1_ROUTER
10092 .get_or_init(|| std::env::var("MEMRA_DEV1_ROUTER").as_deref() == Ok("1"));
10093 if d1_router {
10094 let (sf_h, rn_h) = sigmoid;
10095 let n_ex = m.gate_exps.n_expert;
10096 let act_ct = m.active_count();
10097 let _ = tp.runtime.nvfp4_routes_prestage_with(
10098 bank,
10099 e,
10100 z,
10101 |rank1, in1, sel1, w1| {
10102 let mut guard = DEV1_ROUTER_REPS
10103 .lock()
10104 .map_err(|_| "dev1 router replica lock")?;
10105 let (reps, scratch) =
10106 guard.get_or_insert_with(|| (Default::default(), None));
10107 if !reps.contains_key(&il) {
10108 use cudarc::driver::DevicePtr;
10109 let (g1, p1, a1) = (
10110 rank1.htod(&vec![0.0f32; n_ex * n_embd])?,
10111 rank1.htod(&vec![0.0f32; n_ex])?,
10112 rank1.alloc_u8_uninit(n_ex)?,
10113 );
10114 for (src, dst_len, dst) in [
10115 (
10116 {
10117 let s = e.stream();
10118 let (p, _g) = m.gate_inp.float_data().device_ptr(&s);
10119 p
10120 },
10121 n_ex * n_embd * 4,
10122 {
10123 let s = rank1.stream();
10124 let (p, _g) = g1.device_ptr(&s);
10125 p
10126 },
10127 ),
10128 (
10129 {
10130 let s = e.stream();
10131 let (p, _g) = m.exp_probs_b_dev.device_ptr(&s);
10132 p
10133 },
10134 n_ex * 4,
10135 {
10136 let s = rank1.stream();
10137 let (p, _g) = p1.device_ptr(&s);
10138 p
10139 },
10140 ),
10141 (
10142 {
10143 let s = e.stream();
10144 let (p, _g) = m.active_experts_dev.device_ptr(&s);
10145 p
10146 },
10147 n_ex,
10148 {
10149 let s = rank1.stream();
10150 let (p, _g) = a1.device_ptr(&s);
10151 p
10152 },
10153 ),
10154 ] {
10155 crate::tp::raw_copy_bytes(dst, src, dst_len, rank1)?;
10156 }
10157 rank1.stream().synchronize()?;
10158 reps.insert(il, (g1, p1, a1));
10159 }
10160 if scratch.is_none() {
10161 *scratch = Some(rank1.htod(&vec![0.0f32; n_ex])?);
10162 }
10163 let (g1, p1, a1) = reps.get(&il).expect("armed above");
10164 let logits1 = scratch.as_mut().expect("armed above");
10165 rank1.router_gemv_into(g1, in1, logits1, n_embd, n_ex, 1)?;
10166 rank1.moe_router_sigmoid_topk_into(
10167 logits1, 1, n_ex, n_used, act_ct, p1, a1, sf_h, rn_h, sel1, w1,
10168 )?;
10169 Ok(true)
10170 },
10171 )?;
10172 } else {
10173 let _ = tp.runtime.nvfp4_routes_prestage(bank, e, z)?;
10174 }
10175 #[allow(clippy::type_complexity)] static SELW: std::sync::Mutex<Option<(usize, CudaSlice<i32>, CudaSlice<f32>)>> =
10180 std::sync::Mutex::new(None);
10181 let mut selw = SELW.lock().map_err(|_| "selw lock poisoned")?;
10182 if selw.as_ref().is_none_or(|(d, ..)| *d != e.ctx().ordinal()) {
10183 *selw = Some((
10184 e.ctx().ordinal(),
10185 e.htod_i32(&vec![0i32; n_used])?,
10186 e.htod(&vec![0.0f32; n_used])?,
10187 ));
10188 }
10189 let (_, sel_d, w_d) = selw.as_mut().expect("armed above");
10190 e.moe_router_sigmoid_topk_into(
10191 &logits,
10192 t,
10193 n_expert,
10194 n_used,
10195 m.active_count(),
10196 &m.exp_probs_b_dev,
10197 &m.active_experts_dev,
10198 sf,
10199 route_norm,
10200 sel_d,
10201 w_d,
10202 )?;
10203 crate::moesd::record_device_routes(e, il, n_expert, n_used, sel_d)?;
10204 if std::env::var("MEMRA_MOE_TRACE").is_ok()
10213 || std::env::var("MEMRA_MOE_WEIGHT_TRACE").is_ok()
10214 {
10215 return Err("MEMRA_MOE_TRACE/MEMRA_MOE_WEIGHT_TRACE cannot trace the \
10216 device-routed step TP walk (selection never returns to host; \
10217 tracing would add a new sync). Route through the host-router \
10218 arm — refused rather than silently dropping rows"
10219 .into());
10220 }
10221 crate::moe_sel_dump::refuse_device_only("the device-routed step TP walk")?;
10222 static SHEXP_OV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10226 let shexp_ov = *SHEXP_OV
10227 .get_or_init(|| std::env::var("MEMRA_SHEXP_OVERLAP").as_deref() == Ok("1"));
10228 static SHEXP_D1: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10232 let shexp_d1 = *SHEXP_D1
10233 .get_or_init(|| std::env::var("MEMRA_SHEXP_DEV1").as_deref() == Ok("1"))
10234 && tp.runtime.rank_engine(1).is_some();
10235 static TAIL3: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10239 let tail3 =
10240 *TAIL3.get_or_init(|| std::env::var("MEMRA_TAIL_ADD3").as_deref() != Ok("0"));
10241 let mut ov_issued = false;
10242 let mut d1_issued = false;
10243 let mut tail_folded = false;
10244 let mut output = if shexp_d1 {
10245 let rank1 = tp.runtime.rank_engine(1).expect("checked above");
10246 tp.runtime
10247 .run_tensor_parallel_routes_nvfp4_device_routed_prejoin(
10248 bank,
10249 e,
10250 z,
10251 sel_d,
10252 w_d,
10253 n_used,
10254 tp.activation_limit,
10255 || {
10256 d1_issued =
10257 Self::shexp_dev1_issue(e, rank1, m, z, cfg, il, n_embd)?;
10258 Ok(())
10259 },
10260 )?
10261 } else if shexp_ov {
10262 let post_add = if tail3 {
10267 Self::shexp_overlap_tail_ptrs(e, m, cfg, n_embd)?
10268 } else {
10269 None
10270 };
10271 let used_post = post_add.is_some();
10272 let out = tp
10273 .runtime
10274 .run_tensor_parallel_routes_nvfp4_device_routed_prejoin_add3(
10275 bank,
10276 e,
10277 z,
10278 sel_d,
10279 w_d,
10280 n_used,
10281 tp.activation_limit,
10282 || {
10283 ov_issued = Self::shexp_overlap_issue(e, m, z, cfg, il, n_embd)?;
10284 Ok(())
10285 },
10286 post_add,
10287 )?;
10288 if used_post && ov_issued {
10293 tail_folded = true; }
10295 out
10296 } else {
10297 tp.runtime.run_tensor_parallel_routes_nvfp4_device_routed(
10298 bank,
10299 e,
10300 z,
10301 sel_d,
10302 w_d,
10303 n_used,
10304 tp.activation_limit,
10305 )?
10306 };
10307 if output.len() != t * n_embd {
10308 return Err(format!(
10309 "Step tp routed output has {} values, expected {t}x{n_embd}",
10310 output.len()
10311 )
10312 .into());
10313 }
10314 if tail_folded {
10315 } else if d1_issued {
10317 Self::shexp_dev1_apply(e, &mut output, n_embd)?;
10318 } else if ov_issued {
10319 Self::shexp_overlap_apply(e, &mut output, n_embd)?;
10320 } else {
10321 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut output)?;
10322 }
10323 static DR_LOGGED: std::sync::atomic::AtomicU64 =
10324 std::sync::atomic::AtomicU64::new(0);
10325 let layer_bit = 1u64 << (il as u64 % 64);
10326 if DR_LOGGED.fetch_or(layer_bit, std::sync::atomic::Ordering::Relaxed) & layer_bit
10327 == 0
10328 {
10329 eprintln!(
10330 "[step-tp] execute layer={il} tokens={t} devices={:?} \
10331 expert_transport={} native_p2p={} router=device \
10332 activation=host-canonical accumulation=host-canonical \
10333 output=e-device io=device performance_claim=false \
10334 (logged once per layer)",
10335 tp.devices,
10336 tp.runtime.transport_label(),
10337 tp.runtime.native_p2p(),
10338 );
10339 }
10340 return Ok(output);
10341 }
10342 let automatic_ep_device_router = crate::tp::parallel_ep_device_router_enabled()?;
10343 let automatic_ep_q8_act = crate::tp::parallel_ep_q8_act_enabled()?;
10344 let automatic_ep_q8_scope = crate::tp::parallel_ep_q8_scope()?;
10345 crate::tp::parallel_ep_q8_gu_paired_enabled(
10346 automatic_ep_q8_act,
10347 automatic_ep_q8_scope,
10348 )?;
10349 let automatic_ep_q8_active =
10350 automatic_ep_q8_act && t <= crate::tp::NVFP4_EP_Q8_BATCH_CAP;
10351 if automatic_ep_q8_scope.is_some() && !automatic_ep_q8_act {
10352 return Err(
10353 "MEMRA_PARALLEL_EP_Q8_SCOPE requires MEMRA_PARALLEL_EP_Q8_ACT=1".into(),
10354 );
10355 }
10356 if automatic_ep_q8_act && !automatic_ep_device_router {
10357 return Err(
10358 "MEMRA_PARALLEL_EP_Q8_ACT=1 requires MEMRA_PARALLEL_EP_DEVICE_ROUTER=1".into(),
10359 );
10360 }
10361 if automatic_ep_q8_act && m.step_ep.as_ref().is_none_or(|ep| !ep.nvfp4_device_routes) {
10362 return Err(
10363 "MEMRA_PARALLEL_EP_Q8_ACT=1 requires automatic W4A16 whole-expert EP".into(),
10364 );
10365 }
10366 if t <= crate::tp::NVFP4_EP_DEVICE_ROUTER_BATCH_CAP
10367 && automatic_ep_device_router
10368 && let Some(ep) = &m.step_ep
10369 && ep.nvfp4_device_routes
10370 {
10371 let bank = match &ep.experts {
10372 crate::hybrid::StepEpExpertBank::Nvfp4(bank) => bank,
10373 crate::hybrid::StepEpExpertBank::E4m3(_) => {
10374 return Err("W4A16 device-routed EP reached an E4M3 expert bank".into());
10375 }
10376 };
10377 let pairs = t
10378 .checked_mul(n_used)
10379 .ok_or("W4A16 device-routed EP pair count overflow")?;
10380 let capacity = crate::tp::NVFP4_EP_DEVICE_BATCH_CAP * n_used;
10381 type EpSelwByDevice =
10385 std::collections::HashMap<usize, (usize, CudaSlice<i32>, CudaSlice<f32>)>;
10386 static EP_SELW: std::sync::Mutex<Option<EpSelwByDevice>> =
10387 std::sync::Mutex::new(None);
10388 let mut selw = EP_SELW
10389 .lock()
10390 .map_err(|_| "automatic EP device-router workspace lock poisoned")?;
10391 let device = e.ctx().ordinal();
10392 let workspaces = selw.get_or_insert_with(Default::default);
10393 if workspaces
10394 .get(&device)
10395 .is_none_or(|(cap, ..)| *cap < capacity)
10396 {
10397 workspaces.insert(
10398 device,
10399 (
10400 capacity,
10401 e.htod_i32(&vec![0i32; capacity])?,
10402 e.htod(&vec![0.0f32; capacity])?,
10403 ),
10404 );
10405 }
10406 let (_, sel_d, w_d) = workspaces.get_mut(&device).expect("armed above");
10407 let (sf, route_norm) = sigmoid;
10408 e.moe_router_sigmoid_topk_into(
10409 &logits,
10410 t,
10411 n_expert,
10412 n_used,
10413 m.active_count(),
10414 &m.exp_probs_b_dev,
10415 &m.active_experts_dev,
10416 sf,
10417 route_norm,
10418 sel_d,
10419 w_d,
10420 )?;
10421 crate::moesd::record_device_routes(e, il, n_expert, n_used, sel_d)?;
10422 crate::moe_sel_dump::refuse_device_only(
10423 "the automatic W4A16 device-routed EP walk",
10424 )?;
10425 static SHEXP_OV_AUTO: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10426 let shexp_ov = t == 1
10427 && *SHEXP_OV_AUTO
10428 .get_or_init(|| std::env::var("MEMRA_SHEXP_OVERLAP").as_deref() == Ok("1"));
10429 let mut ov_issued = false;
10430 let mut output = if shexp_ov {
10431 ep.runtime
10432 .run_routed_experts_nvfp4_w4a16_device_routed_prejoin(
10433 bank,
10434 e,
10435 z,
10436 sel_d,
10437 w_d,
10438 t,
10439 n_used,
10440 ep.activation_limit,
10441 || {
10442 ov_issued = Self::shexp_overlap_issue(e, m, z, cfg, il, n_embd)?;
10443 Ok(())
10444 },
10445 )?
10446 } else {
10447 ep.runtime.run_routed_experts_nvfp4_w4a16_device_routed(
10448 bank,
10449 e,
10450 z,
10451 sel_d,
10452 w_d,
10453 t,
10454 n_used,
10455 ep.activation_limit,
10456 )?
10457 };
10458 if output.len() != t * n_embd {
10459 return Err(format!(
10460 "W4A16 device-routed EP output has {} values, expected \
10461 {t}x{n_embd}={}",
10462 output.len(),
10463 t * n_embd,
10464 )
10465 .into());
10466 }
10467 if ov_issued {
10468 Self::shexp_overlap_apply(e, &mut output, n_embd)?;
10469 } else {
10470 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut output)?;
10471 }
10472 static DEVICE_ROUTER_LOGGED: std::sync::atomic::AtomicU64 =
10473 std::sync::atomic::AtomicU64::new(0);
10474 let layer_bit = 1u64 << (il as u64 % 64);
10475 if DEVICE_ROUTER_LOGGED.fetch_or(layer_bit, std::sync::atomic::Ordering::Relaxed)
10476 & layer_bit
10477 == 0
10478 {
10479 eprintln!(
10480 "[parallel-ep] execute layer={il} tokens={t} devices={:?} \
10481 router=device expert_transport={} native_p2p={} \
10482 activation=bf16-rounded accumulation={} output=e-device \
10483 performance_claim=false (logged once per layer)",
10484 ep.devices,
10485 ep.runtime.transport_label(),
10486 ep.runtime.native_p2p(),
10487 if automatic_ep_q8_active {
10488 "token-slot-order-q8"
10489 } else {
10490 "token-slot-order"
10491 },
10492 );
10493 }
10494 debug_assert!(pairs <= capacity);
10495 return Ok(output);
10496 }
10497 static ROUTE_NS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
10501 static ROUTE_CALLS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
10502 let route_timing = std::env::var("MEMRA_STEP_TP_TIMING").as_deref() == Ok("1");
10503 let route_started = route_timing.then(std::time::Instant::now);
10504 let (selected, route_weights, input) = Self::moe_route_sigmoid_with_input(
10505 e,
10506 &logits,
10507 z,
10508 t,
10509 n_embd,
10510 n_expert,
10511 n_used,
10512 m.exp_probs_b.as_deref(),
10513 sigmoid,
10514 m.active_experts.as_deref(),
10515 )?;
10516 if let Some(started) = route_started {
10517 use std::sync::atomic::Ordering;
10518 let ns = ROUTE_NS.fetch_add(started.elapsed().as_nanos() as u64, Ordering::Relaxed)
10519 + started.elapsed().as_nanos() as u64;
10520 let calls = ROUTE_CALLS.fetch_add(1, Ordering::Relaxed) + 1;
10521 if calls.is_multiple_of(430) {
10522 eprintln!(
10523 "[moe-route-timing] calls={calls} total_ms={:.1} avg_us={:.1}",
10524 ns as f64 / 1.0e6,
10525 ns as f64 / calls as f64 / 1.0e3,
10526 );
10527 }
10528 }
10529 crate::moesd::record_host_routes(il, n_expert, n_used, &selected)?;
10530 Self::trace_moe_routes(il, t, &selected, &route_weights)?;
10531 Self::trace_moe_input(e, il, t, n_embd, z)?;
10532 let selected = selected
10533 .iter()
10534 .map(|&expert| expert as usize)
10535 .collect::<Vec<_>>();
10536 if t <= crate::tp::NVFP4_EP_DEVICE_BATCH_CAP
10537 && let Some(ep) = &m.step_ep
10538 && ep.nvfp4_device_routes
10539 {
10540 let bank = match &ep.experts {
10541 crate::hybrid::StepEpExpertBank::Nvfp4(bank) => bank,
10542 crate::hybrid::StepEpExpertBank::E4m3(_) => {
10543 return Err("W4A16 NVFP4 device EP reached an E4M3 expert bank".into());
10544 }
10545 };
10546 let mut output = ep.runtime.run_routed_experts_nvfp4_w4a16_device_io(
10547 bank,
10548 e,
10549 z,
10550 t,
10551 &selected,
10552 &route_weights,
10553 n_used,
10554 ep.activation_limit,
10555 )?;
10556 if output.len() != t * n_embd {
10557 return Err(format!(
10558 "W4A16 NVFP4 EP routed output has {} values, expected \
10559 {t}x{n_embd}={}",
10560 output.len(),
10561 t * n_embd,
10562 )
10563 .into());
10564 }
10565 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut output)?;
10566 static W4A16_EP_LOGGED: std::sync::atomic::AtomicU64 =
10567 std::sync::atomic::AtomicU64::new(0);
10568 let layer_bit = 1u64 << (il as u64 % 64);
10569 if W4A16_EP_LOGGED.fetch_or(layer_bit, std::sync::atomic::Ordering::Relaxed)
10570 & layer_bit
10571 == 0
10572 {
10573 eprintln!(
10574 "[step-ep] execute layer={il} tokens={t} devices={:?} \
10575 expert_transport={} native_p2p={} activation=bf16-rounded \
10576 accumulation={} output=e-device \
10577 performance_claim=false (logged once per layer)",
10578 ep.devices,
10579 ep.runtime.transport_label(),
10580 ep.runtime.native_p2p(),
10581 if t == 1 {
10582 "owner-grouped-rank-order"
10583 } else {
10584 "token-slot-order"
10585 },
10586 );
10587 }
10588 return Ok(output);
10589 }
10590 if t == 1
10595 && crate::tp::step_nvfp4_dev_routes_enabled()?
10596 && let Some(tp) = &m.step_tp
10597 && let crate::hybrid::StepTpExpertBank::Nvfp4(bank) = &tp.experts
10598 {
10599 let mut output = tp.runtime.run_tensor_parallel_routes_nvfp4_device_io(
10600 bank,
10601 e,
10602 z,
10603 &selected,
10604 &route_weights,
10605 n_used,
10606 tp.activation_limit,
10607 )?;
10608 if output.len() != t * n_embd {
10609 return Err(format!(
10610 "Step tp routed output has {} values, expected {t}x{n_embd}",
10611 output.len()
10612 )
10613 .into());
10614 }
10615 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut output)?;
10616 static IO_LOGGED: std::sync::atomic::AtomicU64 =
10617 std::sync::atomic::AtomicU64::new(0);
10618 let layer_bit = 1u64 << (il as u64 % 64);
10619 if IO_LOGGED.fetch_or(layer_bit, std::sync::atomic::Ordering::Relaxed) & layer_bit
10620 == 0
10621 {
10622 eprintln!(
10623 "[step-tp] execute layer={il} tokens={t} devices={:?} \
10624 expert_transport={} native_p2p={} activation=host-canonical \
10625 accumulation=host-canonical output=e-device io=device \
10626 performance_claim=false (logged once per layer)",
10627 tp.devices,
10628 tp.runtime.transport_label(),
10629 tp.runtime.native_p2p(),
10630 );
10631 }
10632 return Ok(output);
10633 }
10634 let (routed, mode, devices, transport, native_p2p) = if let Some(tp) = &m.step_tp {
10635 (
10636 match &tp.experts {
10637 crate::hybrid::StepTpExpertBank::E4m3(bank) => {
10638 tp.runtime.run_tensor_parallel_routes(
10639 bank,
10640 &input,
10641 t,
10642 &selected,
10643 &route_weights,
10644 n_used,
10645 )?
10646 }
10647 crate::hybrid::StepTpExpertBank::Nvfp4(bank) => {
10648 if t == 1 && crate::tp::step_nvfp4_dev_routes_enabled()? {
10649 tp.runtime.run_tensor_parallel_routes_nvfp4_device(
10650 bank,
10651 &input,
10652 &selected,
10653 &route_weights,
10654 n_used,
10655 tp.activation_limit,
10656 )?
10657 } else {
10658 tp.runtime.run_tensor_parallel_routes_nvfp4(
10659 bank,
10660 &input,
10661 t,
10662 &selected,
10663 &route_weights,
10664 n_used,
10665 tp.activation_limit,
10666 )?
10667 }
10668 }
10669 },
10670 "tp",
10671 &tp.devices,
10672 tp.runtime.transport_label(),
10673 tp.runtime.native_p2p(),
10674 )
10675 } else {
10676 let ep = m
10677 .step_ep
10678 .as_ref()
10679 .ok_or("Step distributed runtime has no EP or TP state")?;
10680 (
10681 match &ep.experts {
10682 crate::hybrid::StepEpExpertBank::E4m3(bank) => {
10683 ep.runtime.run_routed_experts(
10684 bank,
10685 &input,
10686 t,
10687 &selected,
10688 &route_weights,
10689 n_used,
10690 ep.activation_limit,
10691 )?
10692 }
10693 crate::hybrid::StepEpExpertBank::Nvfp4(bank) => {
10694 ep.runtime.run_routed_experts_nvfp4(
10695 bank,
10696 &input,
10697 t,
10698 &selected,
10699 &route_weights,
10700 n_used,
10701 ep.activation_limit,
10702 )?
10703 }
10704 },
10705 if ep.configured_by_tp { "tp-ep" } else { "ep" },
10706 &ep.devices,
10707 ep.runtime.transport_label(),
10708 ep.runtime.native_p2p(),
10709 )
10710 };
10711 if routed.len() != t * n_embd {
10712 return Err(format!(
10713 "Step {mode} routed output has {} values, expected {t}x{n_embd}",
10714 routed.len()
10715 )
10716 .into());
10717 }
10718 let mut output = e.htod(&routed)?;
10719 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut output)?;
10720 static STEP_LOGGED: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
10723 let layer_bit = 1u64 << (il as u64 % 64);
10724 if STEP_LOGGED.fetch_or(layer_bit, std::sync::atomic::Ordering::Relaxed) & layer_bit
10725 == 0
10726 {
10727 eprintln!(
10728 "[step-{mode}] execute layer={il} tokens={t} devices={devices:?} \
10729 expert_transport={transport} native_p2p={native_p2p} \
10730 activation={} accumulation={} output={} \
10731 performance_claim=false (logged once per layer)",
10732 if let Some(ep) = &m.step_ep {
10733 ep.runtime.expert_activation_label()
10734 } else {
10735 "host-canonical"
10736 },
10737 if let Some(ep) = &m.step_ep {
10738 ep.runtime.expert_accumulation_label()
10739 } else {
10740 "host-canonical"
10741 },
10742 if let Some(ep) = &m.step_ep {
10743 ep.runtime.expert_output_label()
10744 } else {
10745 "host-accumulated"
10746 },
10747 );
10748 if let Some(ep) = &m.step_ep
10749 && let Some(limit) = ep.activation_limit
10750 {
10751 eprintln!(
10752 "[step-ep-clamp] execute layer={il} tokens={t} routed_clamp={limit} \
10753 formula=min-silu-times-clamped-up performance_claim=false"
10754 );
10755 }
10756 }
10757 return Ok(output);
10758 }
10759 if Self::sigmoid_resident_dev_eligible(e, m, cfg, sliding_gated_moe) {
10760 let moe = cfg.moe.as_ref().unwrap();
10761 let n_expert = moe.expert_count as usize;
10762 let n_used = moe.expert_used_count as usize;
10763 let sigmoid = cfg.sigmoid_router().unwrap();
10764 let logits = Self::moe_router_logits(e, m, z, t, cfg)?;
10765 Self::trace_sigmoid_router_logits(e, il, t, n_expert, n_used, &logits, m, sigmoid)?;
10766 return Self::moe_ffn_sigmoid_dev(e, m, z, zq8, &logits, t, cfg, il, sigmoid);
10767 }
10768 if prefill && t > MOE_DEV_MAX_T && cfg.sigmoid_router().is_some() && cfg.glm5.is_some() {
10790 static GPF_ANNOUNCED: std::sync::atomic::AtomicU8 = std::sync::atomic::AtomicU8::new(0);
10793 let enabled = moe_grouped_prefill_enabled();
10794 let bit = 1u8 << u8::from(enabled);
10795 if GPF_ANNOUNCED.fetch_or(bit, std::sync::atomic::Ordering::Relaxed) & bit == 0 {
10796 eprintln!(
10797 "[moe-grouped-prefill] flag={} t={t} il={il} (announce printed in both \
10798 arms; engagement is the per-layer execute line + the dispatch counter)",
10799 if enabled { "on" } else { "off" },
10800 );
10801 }
10802 if enabled
10803 && let Some(out) = Self::moe_ffn_grouped_prefill_sigmoid(e, m, z, t, cfg, il)?
10804 {
10805 return Ok(out);
10806 }
10807 }
10808 if t > 1 && moe_grouped_enabled(cfg, prefill) {
10811 let grouped_out = Self::moe_ffn_grouped(e, m, z, t, cfg, il, max_block)?;
10812 if std::env::var("MEMRA_MOE_GATE").is_ok() {
10817 let seq_out = Self::moe_ffn_sequential(e, m, z, t, cfg, il, max_block)?;
10818 let g_host = e.dtoh(&grouped_out)?;
10819 let s_host = e.dtoh(&seq_out)?;
10820 let g_bytes: &[u8] = unsafe {
10821 std::slice::from_raw_parts(g_host.as_ptr() as *const u8, g_host.len() * 4)
10822 };
10823 let s_bytes: &[u8] = unsafe {
10824 std::slice::from_raw_parts(s_host.as_ptr() as *const u8, s_host.len() * 4)
10825 };
10826 if g_bytes == s_bytes {
10827 println!("moe-gate il={il} t={t} BYTE-IDENTICAL");
10828 } else {
10829 let diffs = g_host
10830 .iter()
10831 .zip(s_host.iter())
10832 .enumerate()
10833 .filter(|(_, (a, b))| a != b)
10834 .count();
10835 let maxdiff = g_host
10836 .iter()
10837 .zip(s_host.iter())
10838 .map(|(a, b)| (a - b).abs())
10839 .fold(0.0f32, f32::max);
10840 panic!(
10841 "moe-gate il={il} t={t} MISMATCH: {diffs}/{} elems differ, maxdiff={maxdiff:.6e}",
10842 g_host.len()
10843 );
10844 }
10845 }
10846 return Ok(grouped_out);
10847 }
10848 Self::moe_ffn_sequential_zq8(e, m, z, zq8, t, cfg, il, max_block, vrows)
10849 }
10850
10851 fn sum_f32(v: &[f32]) -> String {
10856 let mut h: u64 = 0xcbf2_9ce4_8422_2325;
10857 let (mut nz, mut absmax) = (0usize, 0f32);
10858 for x in v {
10859 h ^= x.to_bits() as u64;
10860 h = h.wrapping_mul(0x100_0000_01b3);
10861 if *x != 0.0 {
10862 nz += 1;
10863 }
10864 if x.abs() > absmax {
10865 absmax = x.abs();
10866 }
10867 }
10868 format!("0x{h:016x}/nz{nz}of{}/max{absmax:.4e}", v.len())
10869 }
10870
10871 fn sum_i8(v: &[i8]) -> String {
10872 let mut h: u64 = 0xcbf2_9ce4_8422_2325;
10873 let mut nz = 0usize;
10874 for x in v {
10875 h ^= *x as u8 as u64;
10876 h = h.wrapping_mul(0x100_0000_01b3);
10877 if *x != 0 {
10878 nz += 1;
10879 }
10880 }
10881 format!("0x{h:016x}/nz{nz}of{}", v.len())
10882 }
10883
10884 #[allow(clippy::too_many_arguments)] fn trace_moe_act(
10892 e: &Engine,
10893 arm: &str,
10894 il: u16,
10895 t: usize,
10896 z: &CudaSlice<f32>,
10897 zq: &CudaSlice<i8>,
10898 zd: &CudaSlice<f32>,
10899 ) {
10900 if !crate::glm5_graph_trace_on()
10901 || crate::glm5_graph_capture_open()
10902 || !crate::glm5_trace_take_slot("act", arm, il)
10903 {
10904 return;
10905 }
10906 let (zs, qs, ds) = (e.dtoh(z), e.dtoh_i8(zq), e.dtoh(zd));
10907 match (zs, qs, ds) {
10908 (Ok(zv), Ok(qv), Ok(dv)) => eprintln!(
10909 "[glm5-vrows-act] arm={arm} il={il} t={t} z={} zq={} zd={}",
10910 Self::sum_f32(&zv),
10911 Self::sum_i8(&qv),
10912 Self::sum_f32(&dv),
10913 ),
10914 _ => eprintln!("[glm5-vrows-act] arm={arm} il={il} readback failed"),
10915 }
10916 }
10917
10918 fn trace_moe_out(e: &Engine, arm: &str, il: u16, out: &CudaSlice<f32>) {
10920 if !crate::glm5_graph_trace_on()
10921 || crate::glm5_graph_capture_open()
10922 || !crate::glm5_trace_take_slot("out", arm, il)
10923 {
10924 return;
10925 }
10926 match e.dtoh(out) {
10927 Ok(v) => eprintln!(
10928 "[glm5-vrows-out] arm={arm} il={il} out={}",
10929 Self::sum_f32(&v)
10930 ),
10931 Err(err) => eprintln!("[glm5-vrows-out] arm={arm} il={il} readback failed ({err})"),
10932 }
10933 }
10934
10935 #[allow(clippy::too_many_arguments)] fn dump_moe_t1_inputs(
10946 e: &Engine,
10947 arm: &str,
10948 m: &MoeWeights,
10949 cfg: &ModelConfig,
10950 il: u16,
10951 t: usize,
10952 n_used: usize,
10953 n_expert: usize,
10954 sel: &[u32],
10955 w: &[f32],
10956 ) {
10957 if !crate::glm5_trace_take_slot("shape", arm, il) {
10958 return;
10959 }
10960 let mac = |x: &crate::model::HostExps| -> Vec<f32> {
10961 sel.iter().map(|&ex| x.macro_scale(ex as usize)).collect()
10962 };
10963 eprintln!(
10964 "[glm5-vrows-t1] arm={arm} dev={} il={il} t={t} n_used={n_used} n_pairs={} \
10965 n_expert={n_expert} limit={:?} gu_il={:?} rp={:?} qtypes=({},{},{}) row_bytes=({},{},{}) \
10966 strides=({},{},{}) macros={} sel={sel:?} w={w:?} mac_g={:?} mac_u={:?} mac_d={:?}",
10967 e.ctx().ordinal(),
10968 t * n_used,
10969 cfg.clamp_exp_at(il as u32),
10970 m.dev_exps.as_ref().map(|d| d.gu_il),
10971 m.dev_exps.as_ref().map(|d| d.rp),
10972 m.gate_exps.qtype,
10973 m.up_exps.qtype,
10974 m.down_exps.qtype,
10975 m.gate_exps.row_bytes,
10976 m.up_exps.row_bytes,
10977 m.down_exps.row_bytes,
10978 m.gate_exps.expert_stride,
10979 m.up_exps.expert_stride,
10980 m.down_exps.expert_stride,
10981 m.gate_exps.macros.is_some(),
10982 mac(&m.gate_exps),
10983 mac(&m.up_exps),
10984 mac(&m.down_exps),
10985 );
10986 }
10987
10988 pub(crate) fn glm5_t1_dev_moe_ready(
10995 e: &Engine,
10996 layer: &crate::hybrid::HybridLayer,
10997 cfg: &ModelConfig,
10998 il: usize,
10999 ) -> bool {
11000 let crate::hybrid::Ffn::Moe(m) = &layer.ffn else {
11001 return true;
11004 };
11005 let Some(moe) = cfg.moe.as_ref() else {
11006 return false;
11007 };
11008 let slab_local = m
11009 .dev_exps
11010 .as_ref()
11011 .is_some_and(|d| !d.gu_il && moe_slab_enabled() && d.dev == e.ctx().ordinal());
11012 slab_local
11013 && m.has_uniform_expert_layout()
11014 && moe_q8_enabled_for_model(cfg, m)
11015 && moe.expert_used_count as usize <= 8
11016 && matches!(cfg.clamp_exp_at(il as u32), Some(SwigluClamp::Pre(l)) if l > 1e-6)
11017 && !crate::cpu_experts::configured()
11018 }
11019
11020 fn sigmoid_resident_dev_eligible(
11021 e: &Engine,
11022 m: &MoeWeights,
11023 cfg: &ModelConfig,
11024 sliding_gated_moe: bool,
11025 ) -> bool {
11026 let Some(moe) = cfg.moe.as_ref() else {
11027 return false;
11028 };
11029 static OBSERVATION_MODE: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
11032 let observation_mode = *OBSERVATION_MODE.get_or_init(|| {
11033 std::env::var("MEMRA_MOE_STATS").is_ok()
11034 || std::env::var("MEMRA_MOE_TRACE").is_ok()
11035 || std::env::var("MEMRA_MOE_WEIGHT_TRACE").is_ok()
11036 || std::env::var("MEMRA_MOE_INPUT_TRACE_DIR").is_ok()
11037 || std::env::var("MEMRA_MOE_GATE").is_ok()
11038 });
11039 let resident_layout_supported = m.dev_exps.as_ref().is_some_and(|dev| {
11040 if dev.dev != e.ctx().ordinal() {
11041 return false;
11042 }
11043 let q8 = moe_q8_enabled_for_model(cfg, m);
11044 let fp8 = dev.fp8_blk.is_some()
11045 && m.gate_exps.qtype == crate::QT_F8_E4M3_BLK
11046 && m.up_exps.qtype == crate::QT_F8_E4M3_BLK
11047 && m.down_exps.qtype == crate::QT_F8_E4M3_BLK;
11048 q8 || fp8
11049 });
11050 sliding_gated_moe
11051 && sigmoid_router_enabled()
11052 && moe_dev_enabled()
11053 && moe_slab_enabled()
11054 && !observation_mode
11055 && moe.expert_used_count <= 8
11056 && m.has_uniform_expert_layout()
11057 && m.gate_exps.macros.is_none()
11058 && m.up_exps.macros.is_none()
11059 && m.down_exps.macros.is_none()
11060 && !m.has_macros
11061 && resident_layout_supported
11062 }
11063
11064 pub(crate) fn moe_ffn_sequential(
11066 e: &Engine,
11067 m: &MoeWeights,
11068 z: &CudaSlice<f32>,
11069 t: usize,
11070 cfg: &ModelConfig,
11071 il: u16,
11072 max_block: usize,
11073 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
11074 Self::moe_ffn_sequential_zq8(e, m, z, None, t, cfg, il, max_block, false)
11075 }
11076
11077 fn moe_router_logits(
11081 e: &Engine,
11082 m: &MoeWeights,
11083 z: &CudaSlice<f32>,
11084 t: usize,
11085 cfg: &ModelConfig,
11086 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
11087 if t < PRIME_MIN_T {
11088 if crate::router_kernel_on() {
11090 e.router_gemv(
11091 m.gate_inp.float_data(),
11092 z,
11093 cfg.n_embd as usize,
11094 m.gate_exps.n_expert,
11095 t,
11096 )
11097 } else {
11098 e.matmul_decode_exact(&m.gate_inp, z, t)
11099 }
11100 } else if crate::router_prefill_exact_on() && crate::router_kernel_on() {
11101 e.router_gemv(
11102 m.gate_inp.float_data(),
11103 z,
11104 cfg.n_embd as usize,
11105 m.gate_exps.n_expert,
11106 t,
11107 )
11108 } else {
11109 e.matmul(&m.gate_inp, z, t)
11110 }
11111 }
11112
11113 fn trace_moe_routes(
11121 il: u16,
11122 t: usize,
11123 sel_all: &[u32],
11124 weights: &[f32],
11125 ) -> Result<(), Box<dyn std::error::Error>> {
11126 use std::io::Write as _;
11127 crate::moe_sel_dump::record_host(il, t, sel_all, weights)?;
11133 if let Ok(path) = std::env::var("MEMRA_MOE_TRACE") {
11134 let mut f = std::fs::OpenOptions::new()
11135 .create(true)
11136 .append(true)
11137 .open(path)?;
11138 let ids: Vec<String> = sel_all.iter().map(|s| s.to_string()).collect();
11139 writeln!(f, "{} {} {}", il, t, ids.join(","))?;
11140 }
11141 if let Ok(path) = std::env::var("MEMRA_MOE_WEIGHT_TRACE") {
11142 let mut f = std::fs::OpenOptions::new()
11143 .create(true)
11144 .append(true)
11145 .open(path)?;
11146 let pairs: Vec<String> = sel_all
11147 .iter()
11148 .zip(weights)
11149 .map(|(expert, weight)| format!("{expert}:{weight:.9}"))
11150 .collect();
11151 writeln!(f, "{} {} {}", il, t, pairs.join(","))?;
11152 }
11153 Ok(())
11154 }
11155
11156 #[allow(clippy::too_many_arguments)]
11157 fn trace_sigmoid_router_logits(
11158 e: &Engine,
11159 il: u16,
11160 t: usize,
11161 n_expert: usize,
11162 n_used: usize,
11163 logits: &CudaSlice<f32>,
11164 m: &MoeWeights,
11165 (scaling_factor, route_norm): (f32, bool),
11166 ) -> Result<(), Box<dyn std::error::Error>> {
11167 if !crate::sigrouter_contract::served_logit_trace_enabled() || t != 1 {
11168 return Ok(());
11169 }
11170 let logits = e.dtoh(logits)?;
11171 let active: Vec<u8> = m
11172 .active_experts
11173 .as_ref()
11174 .map(|mask| mask.iter().map(|&enabled| u8::from(enabled)).collect())
11175 .unwrap_or_else(|| vec![1; n_expert]);
11176 let bias = m.exp_probs_b.clone().unwrap_or_else(|| vec![0.0; n_expert]);
11177 crate::sigrouter_contract::capture_served_logits(
11178 il as u32,
11179 t,
11180 n_expert,
11181 n_used,
11182 scaling_factor,
11183 route_norm,
11184 &active,
11185 &bias,
11186 &logits,
11187 )?;
11188 Ok(())
11189 }
11190
11191 fn trace_moe_input(
11196 e: &Engine,
11197 il: u16,
11198 t: usize,
11199 n_embd: usize,
11200 z: &CudaSlice<f32>,
11201 ) -> Result<(), Box<dyn std::error::Error>> {
11202 use std::io::Write as _;
11203 let Ok(dir) = std::env::var("MEMRA_MOE_INPUT_TRACE_DIR") else {
11204 return Ok(());
11205 };
11206 let values = active_matrix_values(z.len(), t, n_embd, "MoE input trace activation")?;
11207 let host = e.dtoh_view(&z.slice(0..values))?;
11208 let bytes = unsafe {
11209 std::slice::from_raw_parts(
11210 host.as_ptr().cast::<u8>(),
11211 host.len() * std::mem::size_of::<f32>(),
11212 )
11213 };
11214 let state = MOE_INPUT_TRACE_WRITER.get_or_init(|| std::sync::Mutex::new(None));
11215 let mut state = state
11216 .lock()
11217 .map_err(|_| "MoE input trace writer lock is poisoned")?;
11218 if state.is_none() {
11219 let dir = std::path::PathBuf::from(&dir);
11220 std::fs::create_dir_all(&dir)?;
11221 let index = std::fs::OpenOptions::new()
11222 .create(true)
11223 .append(true)
11224 .open(dir.join("index.jsonl"))?;
11225 *state = Some(MoeInputTraceWriter {
11226 dir,
11227 index,
11228 payloads: std::collections::HashMap::new(),
11229 });
11230 }
11231 let writer = state.as_mut().unwrap();
11232 if writer.dir != std::path::Path::new(&dir) {
11233 return Err("MEMRA_MOE_INPUT_TRACE_DIR changed after capture started".into());
11234 }
11235 let file_name = format!("layer-{il:03}.f32");
11236 if !writer.payloads.contains_key(&il) {
11237 let payload = std::fs::OpenOptions::new()
11238 .create(true)
11239 .append(true)
11240 .open(writer.dir.join(&file_name))?;
11241 let offset = payload.metadata()?.len();
11242 writer.payloads.insert(il, (payload, offset));
11243 }
11244 let (payload, offset) = writer.payloads.get_mut(&il).unwrap();
11245 let row_offset = *offset;
11246 payload.write_all(bytes)?;
11247 *offset += bytes.len() as u64;
11248 writeln!(
11249 writer.index,
11250 "{{\"format\":\"memra-moe-input-trace-v1\",\"layer\":{il},\"tokens\":{t},\
11251 \"hidden_size\":{n_embd},\"file\":\"{file_name}\",\"offset\":{row_offset},\
11252 \"payload_bytes\":{}}}",
11253 bytes.len()
11254 )?;
11255 Ok(())
11256 }
11257
11258 #[allow(clippy::too_many_arguments)]
11259 #[allow(clippy::too_many_arguments)]
11260 pub(crate) fn moe_ffn_sequential_zq8(
11262 e: &Engine,
11263 m: &MoeWeights,
11264 z: &CudaSlice<f32>,
11265 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
11266 t: usize,
11267 cfg: &ModelConfig,
11268 il: u16,
11269 max_block: usize,
11270 vrows: bool,
11271 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
11272 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
11273 let moe = cfg.moe.as_ref().unwrap();
11274 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);
11281 debug_assert_eq!(m.gate_exps.out_f, n_ff_exp);
11282 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);
11285
11286 let lim_exp = cfg.clamp_exp_at(il as u32);
11289 let lim_shexp = cfg.clamp_shexp_at(il as u32);
11290 let use_cache = Engine::moe_cache_enabled();
11291 let uniform_experts = m.has_uniform_expert_layout();
11292 let moe_q8 = uniform_experts && moe_q8_enabled_for_model(cfg, m);
11293 let cpu_expert_requested = crate::cpu_experts::configured();
11300 if cpu_expert_requested && (cfg.hy3.is_none() || cfg.m3.is_some()) {
11301 return Err(std::io::Error::other(
11302 "MEMRA_CPU_EXPERT_LIB is experimental and currently gated to Hy3",
11303 )
11304 .into());
11305 }
11306 let cpu_hybrid = cpu_expert_requested && t < PRIME_MIN_T && m.dev_exps.is_none();
11307 let freeze_cpu_residency = cpu_expert_requested
11313 && std::env::var("MEMRA_CPU_EXPERT_FREEZE_CACHE").as_deref() == Ok("1");
11314 let caller_warms_before_freeze = std::env::var("MEMRA_CPU_EXPERT_FREEZE_WARMUP_TOKENS")
11315 .ok()
11316 .and_then(|value| value.parse::<usize>().ok())
11317 .is_some_and(|tokens| tokens > 0);
11318 if cpu_hybrid && freeze_cpu_residency && !caller_warms_before_freeze {
11319 e.freeze_moe_cache();
11320 }
11321 let cache_frozen = use_cache && e.moe_cache_frozen();
11322 let cache_dispatch = use_cache && (!cache_frozen || cpu_hybrid);
11323
11324 let logits = Self::moe_router_logits(e, m, z, t, cfg)?;
11327 if let Some(sig) = cfg.sigmoid_router() {
11328 Self::trace_sigmoid_router_logits(e, il, t, n_expert, n_used, &logits, m, sig)?;
11329 }
11330
11331 let no_exp_macros = m.gate_exps.macros.is_none()
11370 && m.up_exps.macros.is_none()
11371 && m.down_exps.macros.is_none();
11372 if cfg.sigmoid_router().is_none()
11376 && cfg.m3.is_none()
11377 && cfg.hy3.is_none()
11378 && !cfg.swiglu_clamped_at(il as u32)
11379 && no_exp_macros
11380 && t > MOE_DEV_MAX_T
11384 && m.dev_exps.is_some()
11385 && moe_q8_enabled_for_model(cfg, m)
11386 && std::env::var("MEMRA_MOE_PAIRS")
11387 .map(|v| v != "0")
11388 .unwrap_or(true)
11389 && std::env::var("MEMRA_MOE_STATS").is_err()
11390 {
11391 return Self::moe_ffn_pairs(e, m, z, &logits, t, cfg);
11392 }
11393
11394 let dev_ok = uniform_experts
11412 && cfg.sigmoid_router().is_none()
11413 && cfg.m3.is_none()
11414 && cfg.hy3.is_none()
11415 && !cfg.swiglu_clamped_at(il as u32);
11416 let observe_routes = std::env::var("MEMRA_MOE_STATS").is_ok()
11420 || std::env::var("MEMRA_MOE_TRACE").is_ok()
11421 || std::env::var("MEMRA_MOE_WEIGHT_TRACE").is_ok()
11422 || std::env::var("MEMRA_MOE_INPUT_TRACE_DIR").is_ok();
11423 if dev_ok
11424 && t <= MOE_DEV_MAX_T
11425 && m.dev_exps.is_some()
11426 && n_used <= 8
11427 && moe_dev_enabled()
11428 && !observe_routes
11429 {
11430 return Self::moe_ffn_dev(e, m, z, zq8, &logits, t, cfg, il, max_block);
11431 }
11432 if dev_ok && use_cache && n_used <= 8 && moe_dev_enabled() && !observe_routes {
11433 let row_ok = e.with_moe_cache(max_block, |c, eng| {
11434 if moe_prewarm_enabled() {
11435 c.prewarm_layer(il, m, eng)?;
11436 }
11437 Ok(c.layer_dev_row(il, n_expert, eng)?.is_some())
11438 })?;
11439 if row_ok {
11440 return Self::moe_ffn_dev(e, m, z, zq8, &logits, t, cfg, il, max_block);
11441 }
11442 }
11443
11444 let slab_local = m
11450 .dev_exps
11451 .as_ref()
11452 .filter(|d| !d.gu_il && moe_slab_enabled() && d.dev == e.ctx().ordinal());
11453 let slab_bases = slab_local.map(|d| {
11454 use cudarc::driver::DevicePtr;
11455 let s = e.stream();
11456 let (pg, _g0) = d.gate.device_ptr(&s);
11457 let (pu, _g1) = d.up.device_ptr(&s);
11458 let (pd, _g2) = d.down.device_ptr(&s);
11459 (pg, pu, pd)
11460 });
11461 let slab_rp = slab_local.is_some_and(|d| d.rp);
11463 let worker_disk_prefetch =
11469 cache_dispatch && crate::spill_pread::worker_enabled() && !cpu_hybrid;
11470 let promote_worker_h2d =
11471 t == 1 && worker_disk_prefetch && crate::spill_pread::copy_h2d_enabled();
11472 let vrows_t1_dev = t == 1
11513 && (crate::glm5_graph_capture_open() || crate::glm5_vrows_t1_dev_forced())
11519 && !crate::glm5_graph_host_moe()
11524 && !promote_worker_h2d
11525 && sigmoid_router_enabled()
11526 && cfg.sigmoid_router().is_some()
11527 && !observe_routes
11528 && !memra_reference::hidden_trace::enabled()
11529 && !crate::moesd::capture_active();
11530 let vrows_dev = ((vrows && t >= 2 && crate::moe_vrows_dev_tables_on()) || vrows_t1_dev)
11548 && slab_bases.is_some()
11549 && moe_q8
11550 && uniform_experts
11551 && n_used <= 8
11552 && cfg.sigmoid_router().is_some()
11553 && matches!(lim_exp, Some(SwigluClamp::Pre(l)) if l > 1e-6)
11554 && !cpu_hybrid
11555 && sigmoid_router_enabled()
11556 && !observe_routes
11557 && !memra_reference::hidden_trace::enabled()
11558 && !crate::moesd::capture_active();
11559 let mut sel_dev: Option<(CudaSlice<i32>, CudaSlice<f32>)> = None;
11562 let (sel_all, w_all, routed_cpu_input) = if vrows_dev {
11563 let (sf, route_norm) = cfg
11564 .sigmoid_router()
11565 .expect("vrows_dev carries cfg.sigmoid_router().is_some()");
11566 sel_dev = Some(e.moe_router_sigmoid_topk(
11567 &logits,
11568 t,
11569 n_expert,
11570 n_used,
11571 m.active_count(),
11572 &m.exp_probs_b_dev,
11573 &m.active_experts_dev,
11574 sf,
11575 route_norm,
11576 )?);
11577 crate::MOE_VROWS_ROUTER_SYNCS_AVOIDED
11578 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11579 if let Some((si, sw)) = sel_dev.as_ref() {
11584 crate::moe_sel_dump::record_device(e, il, t, n_used, si, sw)?;
11595 if !crate::glm5_graph_capture_open() {
11604 crate::glm5_sel_ledger::prearm(e, il, n_used)?;
11605 }
11606 crate::glm5_sel_ledger::record_device(e, il, si, sw)?;
11607 }
11608 if crate::glm5_graph_trace_on() && !crate::glm5_graph_capture_open() {
11612 let (sel_h, w_h) = match sel_dev.as_ref() {
11613 Some((si, sw)) => (
11614 e.dtoh_i32(si)?
11615 .iter()
11616 .map(|&x| x as u32)
11617 .collect::<Vec<_>>(),
11618 e.dtoh(sw)?,
11619 ),
11620 None => (Vec::new(), Vec::new()),
11621 };
11622 Self::dump_moe_t1_inputs(
11623 e, "device", m, cfg, il, t, n_used, n_expert, &sel_h, &w_h,
11624 );
11625 }
11626 (Vec::new(), Vec::new(), None)
11627 } else if let Some(sig) = cfg.sigmoid_router() {
11628 if cpu_hybrid {
11629 let (sel, w, input) = Self::moe_route_sigmoid_with_input(
11630 e,
11631 &logits,
11632 z,
11633 t,
11634 n_embd,
11635 n_expert,
11636 n_used,
11637 m.exp_probs_b.as_deref(),
11638 sig,
11639 m.active_experts.as_deref(),
11640 )?;
11641 (sel, w, Some(input))
11642 } else {
11643 let (sel, w) =
11644 Self::moe_route_sigmoid_cfg(e, &logits, t, n_expert, n_used, m, sig)?;
11645 (sel, w, None)
11646 }
11647 } else {
11648 let (sel, w) =
11649 Self::moe_route_cfg(e, &logits, t, n_expert, n_used, m.active_experts.as_deref())?;
11650 (sel, w, None)
11651 };
11652 if crate::glm5_graph_trace_on() && !crate::glm5_graph_capture_open() && !sel_all.is_empty()
11656 {
11657 let last = (t - 1) * n_used;
11658 Self::dump_moe_t1_inputs(
11659 e,
11660 "host",
11661 m,
11662 cfg,
11663 il,
11664 t,
11665 n_used,
11666 n_expert,
11667 &sel_all[last..last + n_used],
11668 &w_all[last..last + n_used],
11669 );
11670 }
11671 crate::moesd::record_host_routes(il, n_expert, n_used, &sel_all)?;
11672 if !sel_all.is_empty() && crate::glm5_sel_ledger::armed() {
11674 let last = (t - 1) * n_used;
11675 crate::glm5_sel_ledger::record_host(
11676 e.ctx().ordinal(),
11677 il,
11678 &sel_all[last..last + n_used],
11679 &w_all[last..last + n_used],
11680 );
11681 }
11682 if memra_reference::hidden_trace::enabled() {
11683 memra_reference::hidden_trace::emit_last_row(
11684 "router",
11685 il as i64,
11686 t,
11687 n_expert,
11688 &e.dtoh(&logits)?,
11689 );
11690 let last = (t - 1) * n_used;
11691 let mut route = Vec::with_capacity(n_used * 2);
11692 for slot in 0..n_used {
11693 route.push(sel_all[last + slot] as f32);
11694 route.push(w_all[last + slot]);
11695 }
11696 memra_reference::hidden_trace::emit_last_row("route", il as i64, 1, n_used * 2, &route);
11697 }
11698
11699 Self::trace_moe_routes(il, t, &sel_all, &w_all)?;
11703 Self::trace_moe_input(e, il, t, n_embd, z)?;
11704
11705 if promote_worker_h2d {
11717 let mut selected_blocks = Vec::with_capacity(n_used * 3);
11718 for &ex in sel_all.iter().take(n_used) {
11719 let ex = ex as u16;
11720 selected_blocks.extend([
11721 BlockId::new(il, PROJ_GATE, ex),
11722 BlockId::new(il, PROJ_UP, ex),
11723 BlockId::new(il, PROJ_DOWN, ex),
11724 ]);
11725 }
11726 for &ex in sel_all.iter().take(n_used) {
11727 Self::moe_prefetch_disk_expert(e, il, ex as usize, m, max_block, &selected_blocks)?;
11728 }
11729 e.with_moe_cache(max_block, |cache, eng| {
11730 cache.promote_worker_reads_at_safe_boundary(
11731 &selected_blocks,
11732 &selected_blocks,
11733 eng,
11734 )?;
11735 Ok(())
11736 })?;
11737 }
11738
11739 if t > 1 && std::env::var("MEMRA_MOE_STATS").is_ok() {
11742 let mut cnt = vec![0u32; n_expert];
11743 for &s in sel_all.iter() {
11744 cnt[s as usize] += 1;
11745 }
11746 let total = sel_all.len() as f64;
11747 let mut h = 0.0f64;
11748 let mut active = 0usize;
11749 for &c in &cnt {
11750 if c > 0 {
11751 active += 1;
11752 let p = c as f64 / total;
11753 h -= p * p.log2();
11754 }
11755 }
11756 let maxc = cnt.iter().copied().max().unwrap_or(0);
11757 println!(
11758 "moe-stats il={} t={} assignments={} active={}/{} entropy={:.3}b (max {:.3}b) mean_tok_per_active={:.2} max_tok_per_expert={}",
11759 il,
11760 t,
11761 sel_all.len(),
11762 active,
11763 n_expert,
11764 h,
11765 (n_expert as f64).log2(),
11766 total / active.max(1) as f64,
11767 maxc
11768 );
11769 }
11770
11771 let gdec_may_fire = uniform_experts
11784 && use_cache
11785 && n_used <= 8
11786 && gdec_enabled()
11787 && !cfg.swiglu_clamped_at(il as u32);
11788 let slab_fused_may_fire = slab_bases.is_some()
11816 && n_used <= 8
11817 && gdec_enabled()
11818 && !cfg.swiglu_clamped_at(il as u32)
11819 && cfg.m3.is_none()
11820 && no_exp_macros
11821 && moe_q8;
11822 let fused_epi_common = n_used <= 8
11857 && moe_q8
11858 && cfg.m3.is_none()
11859 && cfg.sigmoid_router().is_some()
11860 && matches!(lim_exp, Some(SwigluClamp::Pre(l)) if l > 1e-6)
11861 && moe_fused_epi_enabled();
11862 let fused_epi_may_fire = fused_epi_common
11863 && uniform_experts
11864 && use_cache
11865 && cache_dispatch
11866 && slab_local.is_none();
11867 let fused_epi_slab_may_fire = fused_epi_common && slab_bases.is_some();
11868 let vrows_fires = ((vrows && t >= 2) || vrows_t1_dev)
11884 && slab_bases.is_some()
11885 && moe_q8
11886 && uniform_experts
11887 && n_used <= 8
11888 && cfg.sigmoid_router().is_some()
11889 && matches!(lim_exp, Some(SwigluClamp::Pre(l)) if l > 1e-6)
11890 && !cpu_hybrid;
11891 if t == 1
11897 && !vrows_fires
11898 && (crate::glm5_decode_graph_on() || crate::glm5_graph_trace_on())
11899 && crate::glm5_trace_take_slot("deny", "t1", il)
11900 {
11901 eprintln!(
11902 "[glm5-vrows-t1-deny] dev={} il={il} capture_open={} t1_dev={vrows_t1_dev} \
11903 forced={} slab_bases={} moe_q8={moe_q8} uniform={uniform_experts} \
11904 n_used={n_used} sigmoid_cfg={} pre_clamp={} cpu_hybrid={cpu_hybrid} \
11905 promote_worker_h2d={promote_worker_h2d} observe_routes={observe_routes} \
11906 sig_router_env={} — the T=1 device-table MoE arm stood down and this layer \
11907 takes the host-readback path. Benign with capture_open=false (an eager gap \
11908 layer); with capture_open=true it is the wrong-answer class the router guard \
11909 refuses, and the first denied conjunct on the line is the reason.",
11910 e.ctx().ordinal(),
11911 crate::glm5_graph_capture_open(),
11912 crate::glm5_vrows_t1_dev_forced(),
11913 slab_bases.is_some(),
11914 cfg.sigmoid_router().is_some(),
11915 matches!(lim_exp, Some(SwigluClamp::Pre(l)) if l > 1e-6),
11916 sigmoid_router_enabled(),
11917 );
11918 }
11919 let mut moe_out = if gdec_may_fire
11924 || slab_fused_may_fire
11925 || fused_epi_may_fire
11926 || fused_epi_slab_may_fire
11927 || vrows_fires
11928 {
11929 e.uninit(t * n_embd)?
11930 } else {
11931 e.zeros(t * n_embd)?
11932 };
11933 let cpu_input = if cpu_hybrid {
11936 Some(routed_cpu_input.ok_or("CPU expert routing did not return the MoE input")?)
11937 } else {
11938 None
11939 };
11940
11941 if vrows_dev && !vrows_fires {
11946 return Err(
11947 "MEMRA_MOE_VROWS_DEV_TABLES routed device-only but the verify-rows arm did not \
11948 fire: the door-D and vrows_fires predicates disagree"
11949 .into(),
11950 );
11951 }
11952 if vrows_fires {
11953 let Some(SwigluClamp::Pre(limit)) = lim_exp else {
11954 return Err(
11955 "verify-rows MoE arm fired without a live PRE clamp: the predicate and \
11956 the dispatch disagree"
11957 .into(),
11958 );
11959 };
11960 let bases = slab_bases.expect("vrows_fires carries slab_bases.is_some()");
11961 let sel = match sel_dev.as_ref() {
11962 Some((si, sw)) => VrowsSel::Dev(si, sw),
11963 None => VrowsSel::Host(&sel_all, &w_all),
11964 };
11965 Self::moe_vrows_pairs_q8(
11966 e,
11967 m,
11968 z,
11969 sel,
11970 il,
11971 bases,
11972 slab_rp,
11973 t,
11974 n_embd,
11975 n_ff_exp,
11976 n_used,
11977 limit,
11978 &mut moe_out,
11979 )?;
11980 if memra_reference::hidden_trace::enabled() {
11981 memra_reference::hidden_trace::emit_last_row(
11982 "routed",
11983 il as i64,
11984 t,
11985 n_embd,
11986 &e.dtoh(&moe_out)?,
11987 );
11988 }
11989 Self::moe_shexp_add(e, m, z, zq8, t, cfg, lim_shexp, &mut moe_out)?;
11990 return Ok(moe_out);
11991 }
11992
11993 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;
12001 let mut scratch_u: Option<CudaSlice<u8>> = None;
12002 let mut scratch_d: Option<CudaSlice<u8>> = None;
12003 let page_window = moe_page_prefetch_window();
12011
12012 for tok in 0..t {
12015 let sel = &sel_all[tok * n_used..(tok + 1) * n_used];
12016 let w = &w_all[tok * n_used..(tok + 1) * n_used];
12017 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd); let mut tok_q8: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
12019
12020 let no_macros = m.gate_exps.macros.is_none()
12034 && m.up_exps.macros.is_none()
12035 && m.down_exps.macros.is_none();
12036 if slab_fused_may_fire {
12046 let (pg, pu, pd) = slab_bases.unwrap();
12047 let mut gp = [0u64; 8];
12048 let mut up = [0u64; 8];
12049 let mut dp = [0u64; 8];
12050 for (j, &ex) in sel.iter().enumerate() {
12051 let ex = ex as usize;
12052 gp[j] = pg + (ex * m.gate_exps.expert_stride) as u64;
12053 up[j] = pu + (ex * m.up_exps.expert_stride) as u64;
12054 dp[j] = pd + (ex * m.down_exps.expert_stride) as u64;
12055 }
12056 let mut wv = [0f32; 8];
12057 wv[..n_used].copy_from_slice(w);
12058 if tok_q8.is_none() {
12059 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
12060 }
12061 let (zq, zd) = tok_q8.as_ref().unwrap();
12062 Self::trace_moe_act(e, "host", il, t, z, zq, zd);
12063 let act = e.moe_gate_up_silu8_q8(
12064 crate::WPtr8(gp),
12065 crate::WPtr8(up),
12066 zq,
12067 zd,
12068 n_embd,
12069 n_ff_exp,
12070 n_used,
12071 m.gate_exps.qtype,
12072 m.up_exps.qtype,
12073 m.gate_exps.row_bytes,
12074 m.up_exps.row_bytes,
12075 )?;
12076 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
12077 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
12078 e.moe_down8_fma_q8(
12079 crate::WPtr8(dp),
12080 crate::F32x8(wv),
12081 &aq2,
12082 &ad2,
12083 &mut dst,
12084 n_ff_exp,
12085 n_embd,
12086 n_used,
12087 m.down_exps.qtype,
12088 m.down_exps.row_bytes,
12089 )?;
12090 continue;
12091 }
12092 if fused_epi_slab_may_fire {
12097 let Some(SwigluClamp::Pre(limit)) = lim_exp else {
12098 return Err(
12099 "fused MoE epilogue (slab) fired without a live PRE clamp: the \
12100 predicate and the dispatch disagree"
12101 .into(),
12102 );
12103 };
12104 let (pg, pu, pd) = slab_bases.unwrap();
12105 let mut g = [0u64; 8];
12106 let mut u = [0u64; 8];
12107 let mut d = [0u64; 8];
12108 for (j, &ex) in sel.iter().enumerate() {
12109 let ex = ex as usize;
12110 g[j] = pg + (ex * m.gate_exps.expert_stride) as u64;
12111 u[j] = pu + (ex * m.up_exps.expert_stride) as u64;
12112 d[j] = pd + (ex * m.down_exps.expert_stride) as u64;
12113 }
12114 if tok_q8.is_none() {
12115 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
12116 }
12117 let (zq, zd) = tok_q8.as_ref().unwrap();
12118 Self::trace_moe_act(e, "host", il, t, z, zq, zd);
12119 Self::moe_fused_epi_launch(
12120 e,
12121 m,
12122 zq,
12123 zd,
12124 sel,
12125 w,
12126 g,
12127 u,
12128 d,
12129 &mut moe_out,
12130 tok,
12131 n_embd,
12132 n_ff_exp,
12133 n_used,
12134 limit,
12135 slab_rp,
12136 )?;
12137 continue;
12138 }
12139 if fused_epi_may_fire {
12144 let Some(SwigluClamp::Pre(limit)) = lim_exp else {
12145 return Err(
12146 "fused MoE epilogue fired without a live PRE clamp: the predicate and \
12147 the dispatch disagree"
12148 .into(),
12149 );
12150 };
12151 if tok_q8.is_none() {
12152 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
12153 }
12154 let (zq, zd) = tok_q8.as_ref().unwrap();
12155 Self::trace_moe_act(e, "host", il, t, z, zq, zd);
12156 if Self::moe_fused_epi_token_q8(
12157 e,
12158 m,
12159 il,
12160 max_block,
12161 zq,
12162 zd,
12163 sel,
12164 w,
12165 &mut moe_out,
12166 tok,
12167 n_embd,
12168 n_ff_exp,
12169 n_used,
12170 limit,
12171 )? {
12172 continue;
12173 }
12174 }
12175 if gdec_may_fire && moe_q8 && cfg.m3.is_none() && no_macros {
12176 if tok_q8.is_none() {
12177 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
12178 }
12179 let (zq, zd) = tok_q8.as_ref().unwrap();
12180 Self::trace_moe_act(e, "host", il, t, z, zq, zd);
12181 if Self::moe_gdec_token_q8(
12182 e,
12183 m,
12184 il,
12185 max_block,
12186 zq,
12187 zd,
12188 sel,
12189 w,
12190 &mut moe_out,
12191 tok,
12192 n_embd,
12193 n_ff_exp,
12194 n_used,
12195 )? {
12196 continue;
12197 }
12198 } else if gdec_may_fire
12199 && cfg.m3.is_none()
12200 && no_macros
12201 && Self::moe_gdec_token(
12202 e,
12203 m,
12204 il,
12205 max_block,
12206 &zt,
12207 sel,
12208 w,
12209 &mut moe_out,
12210 tok,
12211 n_embd,
12212 n_ff_exp,
12213 n_used,
12214 )?
12215 {
12216 continue;
12217 }
12218
12219 if gdec_may_fire || slab_fused_may_fire || fused_epi_may_fire || fused_epi_slab_may_fire
12225 {
12226 let mut row = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
12227 e.memset_zeros_view(&mut row)?;
12228 }
12229
12230 let mut cpu_mask = vec![false; sel.len()];
12236 let cpu_worker = if let Some(host_input) = cpu_input.as_ref() {
12237 let gpu_resident = if use_cache {
12238 e.with_moe_cache(max_block, |cache, _| {
12239 Ok(sel
12240 .iter()
12241 .map(|&expert| {
12242 let expert = expert as u16;
12243 [PROJ_GATE, PROJ_UP, PROJ_DOWN]
12244 .into_iter()
12245 .filter(|&projection| {
12246 cache
12247 .resident(BlockId::new(il, projection, expert))
12248 .is_some()
12249 })
12250 .count()
12251 })
12252 .collect::<Vec<_>>())
12253 })?
12254 } else {
12255 vec![0; sel.len()]
12256 };
12257 let mut cpu_selected = Vec::new();
12258 for (index, (&expert, &route_weight)) in sel.iter().zip(w).enumerate() {
12259 if gpu_resident[index] != 3 {
12260 cpu_mask[index] = true;
12261 crate::cpu_experts::record_incomplete_gpu_residency(gpu_resident[index]);
12262 let expert = expert as usize;
12263 cpu_selected.push((expert, route_weight));
12264 }
12265 }
12266 if crate::cpu_experts::predictor_enabled() {
12267 let row = &host_input[tok * n_embd..(tok + 1) * n_embd];
12271 crate::cpu_experts::predictor_submit(il, row);
12272 }
12273 if cpu_selected.is_empty() {
12274 None
12275 } else {
12276 let row = &host_input[tok * n_embd..(tok + 1) * n_embd];
12277 let job = crate::cpu_experts::prepare_job(m, il, &cpu_selected, row)
12278 .map_err(std::io::Error::other)?;
12279 Some(crate::cpu_experts::submit(job).map_err(std::io::Error::other)?)
12280 }
12281 } else {
12282 None
12283 };
12284
12285 let worker_window = worker_disk_prefetch
12286 .then(worker_prefetch_window)
12287 .unwrap_or(0);
12288 for (j, &ex) in sel.iter().enumerate() {
12289 if cpu_mask[j] {
12290 continue;
12291 }
12292 let ex = ex as usize;
12293 if let Some(d) = slab_local {
12300 let gl = m.gate_exps.expert_layout(ex);
12301 let ul = m.up_exps.expert_layout(ex);
12302 let dl = m.down_exps.expert_layout(ex);
12303 let (g0, u0, d0) = (
12304 ex * m.gate_exps.expert_stride,
12305 ex * m.up_exps.expert_stride,
12306 ex * m.down_exps.expert_stride,
12307 );
12308 let (gate, up) = if moe_q8 {
12309 if tok_q8.is_none() {
12310 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
12311 }
12312 let (zq, zd) = tok_q8.as_ref().unwrap();
12313 Self::trace_moe_act(e, "host", il, t, z, zq, zd);
12314 (
12315 e.qmatvec_expert_q8(
12316 &d.gate,
12317 g0..g0 + gl.len,
12318 zq,
12319 zd,
12320 1,
12321 m.gate_exps.in_f,
12322 m.gate_exps.out_f,
12323 gl.qtype,
12324 gl.row_bytes,
12325 )?,
12326 e.qmatvec_expert_q8(
12327 &d.up,
12328 u0..u0 + ul.len,
12329 zq,
12330 zd,
12331 1,
12332 m.up_exps.in_f,
12333 m.up_exps.out_f,
12334 ul.qtype,
12335 ul.row_bytes,
12336 )?,
12337 )
12338 } else {
12339 (
12340 m.qmatvec_view(
12341 e,
12342 &d.gate,
12343 g0..g0 + gl.len,
12344 &zt,
12345 1,
12346 m.gate_exps.in_f,
12347 m.gate_exps.out_f,
12348 gl.qtype,
12349 gl.row_bytes,
12350 )?,
12351 m.qmatvec_view(
12352 e,
12353 &d.up,
12354 u0..u0 + ul.len,
12355 &zt,
12356 1,
12357 m.up_exps.in_f,
12358 m.up_exps.out_f,
12359 ul.qtype,
12360 ul.row_bytes,
12361 )?,
12362 )
12363 };
12364 let mut act = e.uninit(n_ff_exp)?;
12365 Self::ffn_act_lim(
12366 e,
12367 cfg,
12368 &gate,
12369 &up,
12370 m.gate_exps.macro_scale(ex),
12371 m.up_exps.macro_scale(ex),
12372 lim_exp,
12373 &mut act,
12374 n_ff_exp,
12375 )?;
12376 let y = if moe_q8 {
12377 let (aq2, ad2) = e.quantize_q8_1(&act, 1, n_ff_exp)?;
12378 e.qmatvec_expert_q8(
12379 &d.down,
12380 d0..d0 + dl.len,
12381 &aq2,
12382 &ad2,
12383 1,
12384 m.down_exps.in_f,
12385 m.down_exps.out_f,
12386 dl.qtype,
12387 dl.row_bytes,
12388 )?
12389 } else {
12390 let actv = act.slice(0..n_ff_exp);
12391 m.qmatvec_view(
12392 e,
12393 &d.down,
12394 d0..d0 + dl.len,
12395 &actv,
12396 1,
12397 m.down_exps.in_f,
12398 m.down_exps.out_f,
12399 dl.qtype,
12400 dl.row_bytes,
12401 )?
12402 };
12403 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
12404 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
12405 continue;
12406 }
12407 for next in page_prefetch_positions(j, sel.len(), page_window) {
12408 Self::moe_prefetch_host_expert(sel[next] as usize, m);
12409 }
12410 let keep = [
12411 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_GATE, ex as u16),
12412 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_UP, ex as u16),
12413 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_DOWN, ex as u16),
12414 ];
12415 if worker_disk_prefetch && worker_window > 0 {
12416 for next in worker_prefetch_positions(j, sel.len(), worker_window) {
12417 Self::moe_prefetch_disk_expert(
12418 e,
12419 il,
12420 sel[next] as usize,
12421 m,
12422 max_block,
12423 &keep,
12424 )?;
12425 }
12426 } else if cache_dispatch
12427 && !cpu_hybrid
12428 && moe_prefetch_enabled()
12429 && j + 1 < sel.len()
12430 {
12431 let next = sel[j + 1] as usize;
12432 Self::moe_prefetch_expert(e, il, next, m, max_block, &keep)?;
12433 }
12434 let [gate_q8, up_q8, down_q8] = [moe_q8; 3];
12435 if cache_dispatch && (gate_q8 || up_q8 || down_q8) {
12436 if (gate_q8 || up_q8) && tok_q8.is_none() {
12439 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
12440 }
12441 let gate = if gate_q8 {
12442 let (zq, zd) = tok_q8.as_ref().unwrap();
12443 Self::trace_moe_act(e, "host", il, t, z, zq, zd);
12444 Self::moe_cached_gemm_q8(e, il, PROJ_GATE, ex, m, max_block, zq, zd)?
12445 } else {
12446 Self::moe_cached_gemm(e, il, PROJ_GATE, ex, m, max_block, &zt)?
12447 };
12448 let up = if up_q8 {
12449 let (zq, zd) = tok_q8.as_ref().unwrap();
12450 Self::trace_moe_act(e, "host", il, t, z, zq, zd);
12451 Self::moe_cached_gemm_q8(e, il, PROJ_UP, ex, m, max_block, zq, zd)?
12452 } else {
12453 Self::moe_cached_gemm(e, il, PROJ_UP, ex, m, max_block, &zt)?
12454 };
12455 let mut act = e.uninit(n_ff_exp)?;
12456 Self::ffn_act_lim(
12457 e,
12458 cfg,
12459 &gate,
12460 &up,
12461 m.gate_exps.macro_scale(ex),
12462 m.up_exps.macro_scale(ex),
12463 lim_exp,
12464 &mut act,
12465 n_ff_exp,
12466 )?;
12467 let y = if down_q8 {
12468 let (aq2, ad2) = e.quantize_q8_1(&act, 1, n_ff_exp)?;
12469 Self::moe_cached_gemm_q8(e, il, PROJ_DOWN, ex, m, max_block, &aq2, &ad2)?
12470 } else {
12471 let actv = act.slice(0..n_ff_exp);
12472 Self::moe_cached_gemm(e, il, PROJ_DOWN, ex, m, max_block, &actv)?
12473 };
12474 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
12475 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
12477 } else if cache_dispatch {
12478 let gate = Self::moe_cached_gemm(e, il, PROJ_GATE, ex, m, max_block, &zt)?;
12483 let up = Self::moe_cached_gemm(e, il, PROJ_UP, ex, m, max_block, &zt)?;
12484 let mut act = e.uninit(n_ff_exp)?; Self::ffn_act_lim(
12486 e,
12487 cfg,
12488 &gate,
12489 &up,
12490 m.gate_exps.macro_scale(ex),
12491 m.up_exps.macro_scale(ex),
12492 lim_exp,
12493 &mut act,
12494 n_ff_exp,
12495 )?;
12496 let actv = act.slice(0..n_ff_exp);
12497 let y = Self::moe_cached_gemm(e, il, PROJ_DOWN, ex, m, max_block, &actv)?;
12498 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
12499 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
12501 } else if cache_frozen {
12502 let gate = Self::moe_frozen_gemm(
12507 e,
12508 il,
12509 PROJ_GATE,
12510 ex,
12511 m,
12512 max_block,
12513 &zt,
12514 &mut scratch_g,
12515 g_len,
12516 )?;
12517 let up = Self::moe_frozen_gemm(
12518 e,
12519 il,
12520 PROJ_UP,
12521 ex,
12522 m,
12523 max_block,
12524 &zt,
12525 &mut scratch_u,
12526 u_len,
12527 )?;
12528 let mut act = e.uninit(n_ff_exp)?;
12529 Self::ffn_act_lim(
12530 e,
12531 cfg,
12532 &gate,
12533 &up,
12534 m.gate_exps.macro_scale(ex),
12535 m.up_exps.macro_scale(ex),
12536 lim_exp,
12537 &mut act,
12538 n_ff_exp,
12539 )?;
12540 let actv = act.slice(0..n_ff_exp);
12541 let y = Self::moe_frozen_gemm(
12542 e,
12543 il,
12544 PROJ_DOWN,
12545 ex,
12546 m,
12547 max_block,
12548 &actv,
12549 &mut scratch_d,
12550 d_len,
12551 )?;
12552 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
12553 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
12554 } else {
12555 if scratch_g.is_none() {
12559 scratch_g = Some(e.alloc_u8_uninit(g_len)?);
12560 scratch_u = Some(e.alloc_u8_uninit(u_len)?);
12561 scratch_d = Some(e.alloc_u8_uninit(d_len)?);
12562 }
12563 let (sg, su, sd) = (
12564 scratch_g.as_mut().unwrap(),
12565 scratch_u.as_mut().unwrap(),
12566 scratch_d.as_mut().unwrap(),
12567 );
12568 let gl = m.gate_exps.expert_layout(ex);
12569 let ul = m.up_exps.expert_layout(ex);
12570 let dl = m.down_exps.expert_layout(ex);
12571 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
12572 let gate = m.qmatvec_view(
12573 e,
12574 sg,
12575 0..gl.len,
12576 &zt,
12577 1,
12578 m.gate_exps.in_f,
12579 m.gate_exps.out_f,
12580 gl.qtype,
12581 gl.row_bytes,
12582 )?;
12583
12584 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
12585 let up = m.qmatvec_view(
12586 e,
12587 su,
12588 0..ul.len,
12589 &zt,
12590 1,
12591 m.up_exps.in_f,
12592 m.up_exps.out_f,
12593 ul.qtype,
12594 ul.row_bytes,
12595 )?;
12596
12597 let mut act = e.uninit(n_ff_exp)?; Self::ffn_act_lim(
12599 e,
12600 cfg,
12601 &gate,
12602 &up,
12603 m.gate_exps.macro_scale(ex),
12604 m.up_exps.macro_scale(ex),
12605 lim_exp,
12606 &mut act,
12607 n_ff_exp,
12608 )?;
12609
12610 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
12611 let actv = act.slice(0..n_ff_exp);
12612 let y = m.qmatvec_view(
12613 e,
12614 sd,
12615 0..dl.len,
12616 &actv,
12617 1,
12618 m.down_exps.in_f,
12619 m.down_exps.out_f,
12620 dl.qtype,
12621 dl.row_bytes,
12622 )?;
12623
12624 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
12625 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
12626 }
12627 }
12628 if let Some(worker) = cpu_worker {
12629 let cpu_output = worker.wait().map_err(std::io::Error::other)?;
12630 let cpu_output = e.htod(&cpu_output)?;
12631 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
12632 e.axpy_into(&cpu_output, 1.0, &mut dst, n_embd)?;
12633 }
12634 if cpu_hybrid && !cache_frozen && cpu_expert_profile_admit_enabled() {
12635 for (j, &ex) in sel.iter().enumerate() {
12636 if cpu_mask[j] {
12637 Self::moe_profile_admit_expert(e, il, ex as usize, m, max_block)?;
12638 }
12639 }
12640 }
12641 }
12642
12643 if memra_reference::hidden_trace::enabled() {
12644 memra_reference::hidden_trace::emit_last_row(
12645 "routed",
12646 il as i64,
12647 t,
12648 n_embd,
12649 &e.dtoh(&moe_out)?,
12650 );
12651 }
12652
12653 Self::trace_moe_out(e, "host", il, &moe_out);
12654 Self::moe_shexp_add(e, m, z, zq8, t, cfg, lim_shexp, &mut moe_out)?;
12655
12656 Ok(moe_out)
12657 }
12658
12659 #[allow(clippy::too_many_arguments)]
12679 fn moe_ffn_glm5_ep(
12680 e: &Engine,
12681 m: &MoeWeights,
12682 ep: &crate::glm5_tp::Glm5EpExps,
12683 z: &CudaSlice<f32>,
12684 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
12685 t: usize,
12686 cfg: &ModelConfig,
12687 il: u16,
12688 prefill: bool,
12689 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12690 let moe = cfg
12691 .moe
12692 .as_ref()
12693 .ok_or("glm5 EP execution requires MoE model metadata")?;
12694 let n_embd = cfg.n_embd as usize;
12695 let n_expert = moe.expert_count as usize;
12696 let n_used = moe.expert_used_count as usize;
12697 let n_ff_exp = moe.expert_ff_length as usize;
12698 let sig = cfg
12699 .sigmoid_router()
12700 .ok_or("glm5 EP execution requires the sigmoid router")?;
12701 let lim_exp = cfg.clamp_exp_at(il as u32);
12702 let lim_shexp = cfg.clamp_shexp_at(il as u32);
12703 if ep.slabs.iter().map(|s| s.n_experts).sum::<usize>() != n_expert {
12704 return Err(format!(
12705 "glm5 EP slabs cover {:?} experts, model declares {n_expert}",
12706 ep.slabs.iter().map(|s| s.n_experts).collect::<Vec<_>>()
12707 )
12708 .into());
12709 }
12710 let rt = &ep.rt;
12711 let ranks = ep.ranks();
12712
12713 let logits = Self::moe_router_logits(e, m, z, t, cfg)?;
12715 let (sel_all, w_all) =
12716 Self::moe_route_sigmoid_cfg(e, &logits, t, n_expert, n_used, m, sig)?;
12717 crate::moesd::record_host_routes(il, n_expert, n_used, &sel_all)?;
12718 Self::trace_moe_routes(il, t, &sel_all, &w_all)?;
12721
12722 if prefill && t > MOE_DEV_MAX_T {
12728 static EPGP_ANNOUNCED: std::sync::atomic::AtomicU8 =
12729 std::sync::atomic::AtomicU8::new(0);
12730 let enabled = crate::ep_grouped_prime_on() && moe_grouped_prefill_enabled();
12731 let bit = 1u8 << u8::from(enabled);
12732 if EPGP_ANNOUNCED.fetch_or(bit, std::sync::atomic::Ordering::Relaxed) & bit == 0 {
12733 eprintln!(
12734 "[glm5-ep-grouped-prime] flag={} t={t} il={il} (announce printed in both \
12735 arms; engagement is the dispatch counter + per-layer execute line)",
12736 if enabled { "on" } else { "off" },
12737 );
12738 }
12739 if enabled
12740 && let Some(mut out) =
12741 Self::moe_ffn_glm5_ep_grouped_prime(e, m, ep, z, &sel_all, &w_all, t, cfg, il)?
12742 {
12743 Self::moe_shexp_add(e, m, z, zq8, t, cfg, lim_shexp, &mut out)?;
12744 return Ok(out);
12745 }
12746 }
12747
12748 if t > 1 && !prefill {
12756 static EP_VERIFY_MARKED: std::sync::atomic::AtomicBool =
12757 std::sync::atomic::AtomicBool::new(false);
12758 if !EP_VERIFY_MARKED.swap(true, std::sync::atomic::Ordering::Relaxed) {
12759 eprintln!(
12760 "[glm5-tp-ep] verify rows ride the SEQUENTIAL EP walk (t={t}): the \
12761 batched vrows MoE pair is preempted by EP; the EP-aware vrows arm is \
12762 the named lever performance_claim=false"
12763 );
12764 }
12765 }
12766 if crate::ep_diet_on() {
12769 let mut out = Self::moe_ffn_glm5_ep_diet(e, m, ep, z, &sel_all, &w_all, t, cfg, il)?;
12770 Self::moe_shexp_add(e, m, z, zq8, t, cfg, lim_shexp, &mut out)?;
12771 return Ok(out);
12772 }
12773
12774 let mut moe_out = e.zeros(t * n_embd)?;
12775 use crate::tp_transport::TpTransport as TpXport;
12776 let hop = ep.rt.hop(e);
12777 let z_host = match hop.transport {
12786 TpXport::HostCanonical => Some(crate::tp_transport::host_stage_block(
12787 &hop,
12788 0,
12789 z,
12790 t * n_embd,
12791 )?),
12792 TpXport::PeerPull => None,
12793 };
12794 let z_peer_bulks = match hop.transport {
12795 TpXport::PeerPull => Some(crate::tp_transport::fanout_f32(&hop, z, t * n_embd)?),
12796 TpXport::HostCanonical => None,
12797 };
12798 for tok in 0..t {
12799 let sel = &sel_all[tok * n_used..(tok + 1) * n_used];
12800 let w = &w_all[tok * n_used..(tok + 1) * n_used];
12801 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
12802 let z_peer_row_holders: Option<Vec<CudaSlice<f32>>> = match &z_host {
12806 Some(h) => {
12807 let mut rows = Vec::with_capacity(ranks - 1);
12808 for r in 1..ranks {
12809 rows.push(crate::tp_transport::host_row_to(
12810 &hop,
12811 r,
12812 &h[tok * n_embd..(tok + 1) * n_embd],
12813 )?);
12814 }
12815 Some(rows)
12816 }
12817 None => None,
12818 };
12819 for (j, &ex) in sel.iter().enumerate() {
12822 let ex = ex as usize;
12823 let owner = ep.owner(ex);
12824 if owner != 0 {
12825 crate::glm5_tp::GLM5_EP_PEER_SLOT_DISPATCHES
12827 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
12828 if matches!(
12831 crate::glm5_tp::gate_red(),
12832 Ok(Some(crate::glm5_tp::GateRed::SkipPeerCombine))
12833 ) {
12834 continue;
12835 }
12836 }
12837 let zin_holder;
12838 let (dev, slab, zin) = if owner == 0 {
12839 (e, &ep.slabs[0], &zt)
12840 } else {
12841 zin_holder = match (&z_peer_row_holders, &z_peer_bulks) {
12842 (Some(rows), _) => rows[owner - 1].slice(0..n_embd),
12843 (None, Some(bulks)) => {
12844 bulks[owner - 1].slice(tok * n_embd..(tok + 1) * n_embd)
12845 }
12846 (None, None) => {
12847 return Err(
12848 "glm5 EP: neither transport arm staged the peer activation".into(),
12849 );
12850 }
12851 };
12852 (
12853 crate::glm5_tp::rank_engine(e, rt, owner),
12854 &ep.slabs[owner],
12855 &zin_holder,
12856 )
12857 };
12858 let local = ep.local_of[ex] as usize;
12862 let gl = m.gate_exps.expert_stride;
12863 let ul = m.up_exps.expert_stride;
12864 let dl = m.down_exps.expert_stride;
12865 let gate = dev.qmatvec_view(
12866 &slab.gate,
12867 local * gl..(local + 1) * gl,
12868 zin,
12869 1,
12870 m.gate_exps.in_f,
12871 m.gate_exps.out_f,
12872 m.gate_exps.qtype,
12873 m.gate_exps.row_bytes,
12874 )?;
12875 let up = dev.qmatvec_view(
12876 &slab.up,
12877 local * ul..(local + 1) * ul,
12878 zin,
12879 1,
12880 m.up_exps.in_f,
12881 m.up_exps.out_f,
12882 m.up_exps.qtype,
12883 m.up_exps.row_bytes,
12884 )?;
12885 let mut act = dev.uninit(n_ff_exp)?; Self::ffn_act_lim(
12887 dev,
12888 cfg,
12889 &gate,
12890 &up,
12891 m.gate_exps.macro_scale(ex),
12892 m.up_exps.macro_scale(ex),
12893 lim_exp,
12894 &mut act,
12895 n_ff_exp,
12896 )?;
12897 let actv = act.slice(0..n_ff_exp);
12898 let y = dev.qmatvec_view(
12899 &slab.down,
12900 local * dl..(local + 1) * dl,
12901 &actv,
12902 1,
12903 m.down_exps.in_f,
12904 m.down_exps.out_f,
12905 m.down_exps.qtype,
12906 m.down_exps.row_bytes,
12907 )?;
12908 let y_root = if owner == 0 {
12913 y
12914 } else {
12915 crate::tp_transport::return_row_to_root(&hop, owner, &y, n_embd)?
12916 };
12917 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
12918 e.axpy_into(
12919 &y_root,
12920 w[j] * m.down_exps.macro_scale(ex),
12921 &mut dst,
12922 n_embd,
12923 )?;
12924 }
12925 }
12926
12927 Self::moe_shexp_add(e, m, z, zq8, t, cfg, lim_shexp, &mut moe_out)?;
12928 Ok(moe_out)
12929 }
12930
12931 #[allow(clippy::too_many_arguments)]
12959 fn moe_ffn_glm5_ep_diet(
12960 e: &Engine,
12961 m: &MoeWeights,
12962 ep: &crate::glm5_tp::Glm5EpExps,
12963 z: &CudaSlice<f32>,
12964 sel_all: &[u32],
12965 w_all: &[f32],
12966 t: usize,
12967 cfg: &ModelConfig,
12968 il: u16,
12969 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
12970 use std::sync::atomic::Ordering;
12971 let moe = cfg
12972 .moe
12973 .as_ref()
12974 .ok_or("glm5 EP execution requires MoE model metadata")?;
12975 let n_embd = cfg.n_embd as usize;
12976 let n_used = moe.expert_used_count as usize;
12977 let n_ff_exp = moe.expert_ff_length as usize;
12978 let lim_exp = cfg.clamp_exp_at(il as u32);
12979 let rt = &ep.rt;
12980 let n_pairs = t * n_used;
12981 if sel_all.len() < n_pairs || w_all.len() < n_pairs || z.len() < t * n_embd {
12982 return Err("glm5 EP diet geometry".into());
12983 }
12984 let red_skip_peer = matches!(
12985 crate::glm5_tp::gate_red(),
12986 Ok(Some(crate::glm5_tp::GateRed::SkipPeerCombine))
12987 );
12988
12989 let ranks = ep.ranks();
12995 let mut per_rank = vec![0usize; ranks];
12996 for &s in sel_all.iter().take(n_pairs) {
12997 let ex = s as usize;
12998 if ex >= ep.owner_of.len() {
12999 return Err(format!("glm5 EP diet: selection {ex} outside the bank").into());
13000 }
13001 per_rank[ep.owner(ex)] += 1;
13002 }
13003 let mut base = vec![0usize; ranks];
13004 for r in 1..ranks {
13005 base[r] = base[r - 1] + per_rank[r - 1];
13006 }
13007 let mut ids = vec![0i32; n_pairs];
13008 {
13009 let mut k = vec![0usize; ranks];
13010 for (p, id) in ids.iter_mut().enumerate() {
13011 let r = ep.owner(sel_all[p] as usize);
13012 *id = (base[r] + k[r]) as i32;
13013 k[r] += 1;
13014 }
13015 }
13016
13017 crate::glm5_tp::GLM5_EP_DIET_DISPATCHES.fetch_add(1, Ordering::Relaxed);
13018 for r in 1..ranks {
13019 crate::glm5_tp::GLM5_EP_DIET_FANOUT_UPLOADS_AVOIDED.fetch_add(
13020 if per_rank[r] > 0 {
13021 (t - 1) as u64
13022 } else {
13023 t as u64
13024 },
13025 Ordering::Relaxed,
13026 );
13027 }
13028 static EP_DIET_MARKED: std::sync::atomic::AtomicBool =
13029 std::sync::atomic::AtomicBool::new(false);
13030 if !EP_DIET_MARKED.swap(true, Ordering::Relaxed) {
13031 eprintln!(
13032 "[glm5-ep-diet] engaged: bulk fan-out + compact peer staging + single \
13033 slot-ordered scatter combine; per-slot host round-trips removed \
13034 transport={} performance_claim=false",
13035 ep.rt.transport.name(),
13036 );
13037 }
13038
13039 let expert_row = |dev: &Engine,
13042 slab: &crate::glm5_tp::EpRankSlab,
13043 zin: &cudarc::driver::CudaView<f32>,
13044 ex: usize|
13045 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
13046 let local = ep.local_of[ex] as usize;
13047 let gl = m.gate_exps.expert_stride;
13048 let ul = m.up_exps.expert_stride;
13049 let dl = m.down_exps.expert_stride;
13050 let gate = dev.qmatvec_view(
13051 &slab.gate,
13052 local * gl..(local + 1) * gl,
13053 zin,
13054 1,
13055 m.gate_exps.in_f,
13056 m.gate_exps.out_f,
13057 m.gate_exps.qtype,
13058 m.gate_exps.row_bytes,
13059 )?;
13060 let up = dev.qmatvec_view(
13061 &slab.up,
13062 local * ul..(local + 1) * ul,
13063 zin,
13064 1,
13065 m.up_exps.in_f,
13066 m.up_exps.out_f,
13067 m.up_exps.qtype,
13068 m.up_exps.row_bytes,
13069 )?;
13070 let mut act = dev.uninit(n_ff_exp)?; Self::ffn_act_lim(
13072 dev,
13073 cfg,
13074 &gate,
13075 &up,
13076 m.gate_exps.macro_scale(ex),
13077 m.up_exps.macro_scale(ex),
13078 lim_exp,
13079 &mut act,
13080 n_ff_exp,
13081 )?;
13082 let actv = act.slice(0..n_ff_exp);
13083 dev.qmatvec_view(
13084 &slab.down,
13085 local * dl..(local + 1) * dl,
13086 &actv,
13087 1,
13088 m.down_exps.in_f,
13089 m.down_exps.out_f,
13090 m.down_exps.qtype,
13091 m.down_exps.row_bytes,
13092 )
13093 };
13094
13095 let hop = ep.rt.hop(e);
13099 let mut y_peer_blks: Vec<Option<CudaSlice<f32>>> = (0..ranks).map(|_| None).collect();
13100 for r in 1..ranks {
13101 if per_rank[r] == 0 {
13102 continue;
13103 }
13104 let dev = crate::glm5_tp::rank_engine(e, rt, r);
13105 let z_r = crate::tp_transport::fanout_f32_to(&hop, r, z, t * n_embd)?;
13108 let mut blk = if red_skip_peer {
13111 dev.zeros(per_rank[r] * n_embd)?
13112 } else {
13113 dev.uninit(per_rank[r] * n_embd)?
13114 };
13115 let mut k = 0usize;
13116 for tok in 0..t {
13117 let sel = &sel_all[tok * n_used..(tok + 1) * n_used];
13118 for &ex in sel.iter() {
13119 let ex = ex as usize;
13120 if ep.owner(ex) != r {
13121 continue;
13122 }
13123 crate::glm5_tp::GLM5_EP_PEER_SLOT_DISPATCHES.fetch_add(1, Ordering::Relaxed);
13125 crate::glm5_tp::GLM5_EP_DIET_PEER_ROUNDTRIPS_AVOIDED
13126 .fetch_add(1, Ordering::Relaxed);
13127 if red_skip_peer {
13128 k += 1;
13129 continue;
13130 }
13131 let zt_r = z_r.slice(tok * n_embd..(tok + 1) * n_embd);
13132 let y = expert_row(dev, &ep.slabs[r], &zt_r, ex)?;
13133 dev.copy_into(&mut blk, k * n_embd, &y, n_embd)?;
13134 k += 1;
13135 }
13136 }
13137 y_peer_blks[r] = Some(blk);
13138 }
13139
13140 let mut y_all = e.uninit(n_pairs * n_embd)?;
13143 {
13144 let mut k = 0usize;
13145 for tok in 0..t {
13146 let sel = &sel_all[tok * n_used..(tok + 1) * n_used];
13147 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
13148 for &ex in sel.iter() {
13149 let ex = ex as usize;
13150 if ep.owner(ex) != 0 {
13151 continue;
13152 }
13153 let y = expert_row(e, &ep.slabs[0], &zt, ex)?;
13154 e.copy_into(&mut y_all, k * n_embd, &y, n_embd)?;
13155 k += 1;
13156 }
13157 }
13158 }
13159
13160 for r in 1..ranks {
13165 if let Some(blk) = &y_peer_blks[r] {
13166 crate::tp_transport::return_block_to_root(
13167 &hop,
13168 r,
13169 blk,
13170 &mut y_all,
13171 base[r] * n_embd,
13172 per_rank[r] * n_embd,
13173 )?;
13174 crate::glm5_tp::GLM5_EP_DIET_BULK_RETURNS.fetch_add(1, Ordering::Relaxed);
13175 }
13176 }
13177
13178 let mut wd = vec![0f32; n_pairs];
13182 for ((&id, &w), &s) in ids.iter().zip(w_all.iter()).zip(sel_all.iter()) {
13183 wd[id as usize] = w * m.down_exps.macro_scale(s as usize);
13184 }
13185 let pw = e.htod(&wd)?;
13186 let toff: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
13187 let toff_d = e.htod_i32(&toff)?;
13188 let ids_d = e.htod_i32(&ids)?;
13189 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)?;
13191 Ok(moe_out)
13192 }
13193
13194 #[allow(clippy::too_many_arguments)]
13213 fn moe_ffn_glm5_ep_grouped_prime(
13214 e: &Engine,
13215 m: &MoeWeights,
13216 ep: &crate::glm5_tp::Glm5EpExps,
13217 z: &CudaSlice<f32>,
13218 sel_all: &[u32],
13219 w_all: &[f32],
13220 t: usize,
13221 cfg: &ModelConfig,
13222 il: u16,
13223 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
13224 use std::sync::atomic::Ordering;
13225 let moe = cfg
13226 .moe
13227 .as_ref()
13228 .ok_or("glm5 EP grouped prime requires MoE model metadata")?;
13229 let n_embd = cfg.n_embd as usize;
13230 let n_expert = moe.expert_count as usize;
13231 let n_used = moe.expert_used_count as usize;
13232 let n_ff_exp = moe.expert_ff_length as usize;
13233 if crate::moe_f16g_mode() == 0 || std::env::var("MEMRA_MOE_GATE").is_ok() {
13237 return Ok(None);
13238 }
13239 if !(f16g_proj_ok(m.gate_exps.qtype, n_embd)
13240 && f16g_proj_ok(m.up_exps.qtype, n_embd)
13241 && f16g_proj_ok(m.down_exps.qtype, n_ff_exp))
13242 {
13243 return Ok(None);
13244 }
13245 if n_expert > 512 || n_used == 0 || n_used > 8 {
13246 return Ok(None);
13247 }
13248 let lim_exp = cfg.clamp_exp_at(il as u32);
13249 if matches!(lim_exp, Some(SwigluClamp::Post(_))) {
13250 return Err(
13251 "EP grouped prime is qualified for the PRE-clamped SwiGLU form only; \
13252 a POST-clamp layer must ride the sequential arm"
13253 .into(),
13254 );
13255 }
13256 let n_pairs = t * n_used;
13257 if sel_all.len() < n_pairs || w_all.len() < n_pairs || z.len() < t * n_embd {
13258 return Err("EP grouped prime geometry".into());
13259 }
13260 let rt = &ep.rt;
13261 let red_skip_peer = matches!(
13262 crate::glm5_tp::gate_red(),
13263 Ok(Some(crate::glm5_tp::GateRed::SkipPeerCombine))
13264 );
13265
13266 let rank_pass = |dev: &Engine,
13271 rank: u8,
13272 ptr_row: &CudaSlice<u64>,
13273 z_dev: &CudaSlice<f32>|
13274 -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
13275 let mut buckets_l: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
13278 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];
13282 for p in 0..n_pairs {
13283 let ex = sel_all[p] as usize;
13284 if ex >= n_expert {
13285 return Err(format!("EP grouped prime selection {ex} >= {n_expert}").into());
13286 }
13287 if ep.owner(ex) != rank as usize {
13288 continue;
13289 }
13290 let l = local_tok.len() as i32;
13291 buckets_l[ex].push(l);
13292 let tok = p / n_used;
13293 local_tok.push(tok as i32);
13294 local_ex.push(ex);
13295 local_wd.push(w_all[p] * m.down_exps.macro_scale(ex));
13296 local_count_per_tok[tok] += 1;
13297 }
13298 let n_owned = local_tok.len();
13299 if n_owned == 0 {
13300 return Ok(None);
13301 }
13302 let mut ex_ids: Vec<i32> = Vec::new();
13303 let mut ex_off: Vec<i32> = vec![0];
13304 let mut ex_pairs: Vec<i32> = Vec::with_capacity(n_owned); let mut csr_tok: Vec<i32> = Vec::with_capacity(n_owned);
13306 for (e_id, b) in buckets_l.iter().enumerate() {
13307 if !b.is_empty() {
13308 ex_ids.push(e_id as i32);
13309 for &l in b {
13310 ex_pairs.push(l);
13311 csr_tok.push(local_tok[l as usize]);
13312 }
13313 ex_off.push(ex_pairs.len() as i32);
13314 }
13315 }
13316 let n_active = ex_ids.len();
13317 if n_active == 0 || n_active > 512 {
13318 return Err(format!("EP grouped prime n_active {n_active} outside 1..=512").into());
13319 }
13320
13321 let exi = dev.htod_i32(&ex_ids)?;
13322 let exo = dev.htod_i32(&ex_off)?;
13323 let exp_d = dev.htod_i32(&ex_pairs)?;
13324 let csr_tok_d = dev.htod_i32(&csr_tok)?;
13325
13326 let (z16, zs) = dev.moe_f16g_act(z_dev, Some(&csr_tok_d), n_embd, n_owned)?;
13328 let mut g = dev.moe_f16_grouped(
13329 ptr_row,
13330 0,
13331 n_expert,
13332 &exi,
13333 &ex_off,
13334 &exo,
13335 &z16,
13336 &zs,
13337 n_embd,
13338 n_ff_exp,
13339 n_active,
13340 n_owned,
13341 m.gate_exps.qtype,
13342 m.gate_exps.row_bytes,
13343 )?;
13344 if m.gate_exps.macros.is_some() {
13345 let mg: Vec<f32> = ex_pairs
13346 .iter()
13347 .map(|&l| m.gate_exps.macro_scale(local_ex[l as usize]))
13348 .collect();
13349 let mg_d = dev.htod(&mg)?;
13350 dev.scale_rows(&mut g, &mg_d, n_ff_exp, n_owned)?;
13351 }
13352 let mut u = dev.moe_f16_grouped(
13353 ptr_row,
13354 1,
13355 n_expert,
13356 &exi,
13357 &ex_off,
13358 &exo,
13359 &z16,
13360 &zs,
13361 n_embd,
13362 n_ff_exp,
13363 n_active,
13364 n_owned,
13365 m.up_exps.qtype,
13366 m.up_exps.row_bytes,
13367 )?;
13368 if m.up_exps.macros.is_some() {
13369 let mu: Vec<f32> = ex_pairs
13370 .iter()
13371 .map(|&l| m.up_exps.macro_scale(local_ex[l as usize]))
13372 .collect();
13373 let mu_d = dev.htod(&mu)?;
13374 dev.scale_rows(&mut u, &mu_d, n_ff_exp, n_owned)?;
13375 }
13376
13377 let act = match lim_exp {
13379 Some(SwigluClamp::Pre(limit)) => {
13380 let mut a = dev.uninit(n_owned * n_ff_exp)?;
13381 dev.swiglu_preclamped_mul_scaled(
13382 &g,
13383 &u,
13384 1.0,
13385 1.0,
13386 limit,
13387 &mut a,
13388 n_owned * n_ff_exp,
13389 )?;
13390 a
13391 }
13392 None => dev.moe_pairs_silu_mul(&g, &u, n_owned * n_ff_exp)?,
13393 Some(SwigluClamp::Post(_)) => unreachable!("refused before any launch"),
13394 };
13395
13396 let (a16, a_s) = dev.moe_f16g_act(&act, None, n_ff_exp, n_owned)?;
13398 let d_csr = dev.moe_f16_grouped(
13399 ptr_row,
13400 2,
13401 n_expert,
13402 &exi,
13403 &ex_off,
13404 &exo,
13405 &a16,
13406 &a_s,
13407 n_ff_exp,
13408 n_embd,
13409 n_active,
13410 n_owned,
13411 m.down_exps.qtype,
13412 m.down_exps.row_bytes,
13413 )?;
13414 let y_local = dev.rows_permute(&d_csr, &exp_d, n_owned, n_embd)?;
13415 let mut toff: Vec<i32> = Vec::with_capacity(t + 1);
13416 let mut acc = 0i32;
13417 toff.push(0);
13418 for &c in &local_count_per_tok {
13419 acc += c;
13420 toff.push(acc);
13421 }
13422 let tids: Vec<i32> = (0..n_owned as i32).collect();
13423 let pw = dev.htod(&local_wd)?;
13424 let toff_d = dev.htod_i32(&toff)?;
13425 let tids_d = dev.htod_i32(&tids)?;
13426 let mut partial = dev.uninit(t * n_embd)?; dev.moe_pairs_scatter(&y_local, &pw, &toff_d, &tids_d, &mut partial, t, n_embd)?;
13428 Ok(Some(partial))
13429 };
13430
13431 let ranks = ep.ranks();
13436 let n_peer_pairs = sel_all
13437 .iter()
13438 .take(n_pairs)
13439 .filter(|&&ex| ep.owner(ex as usize) != 0)
13440 .count() as u64;
13441 crate::glm5_tp::GLM5_EP_PEER_SLOT_DISPATCHES.fetch_add(n_peer_pairs, Ordering::Relaxed);
13442 let hop = ep.rt.hop(e);
13443 let mut peer_partials: Vec<Option<CudaSlice<f32>>> = (0..ranks).map(|_| None).collect();
13444 for r in 1..ranks {
13445 let rank_owns_pairs = sel_all
13446 .iter()
13447 .take(n_pairs)
13448 .any(|&ex| ep.owner(ex as usize) == r);
13449 if !rank_owns_pairs {
13450 continue;
13451 }
13452 let dev = crate::glm5_tp::rank_engine(e, rt, r);
13453 let z_r = crate::tp_transport::fanout_f32_to(&hop, r, z, t * n_embd)?;
13455 dev.bind_runtime_device(dev.ctx().ordinal() as i32)?;
13456 let res = rank_pass(dev, r as u8, &ep.ptr_rows[r], &z_r);
13457 e.bind_runtime_device(e.ctx().ordinal() as i32)?;
13458 peer_partials[r] = res?;
13459 }
13460 let root_partial = rank_pass(e, 0, &ep.ptr_rows[0], z)?;
13461
13462 let mut out = match root_partial {
13468 Some(p) => p,
13469 None => e.zeros(t * n_embd)?,
13470 };
13471 for r in 1..ranks {
13472 if let Some(pp) = &peer_partials[r]
13473 && !red_skip_peer
13474 {
13475 let pp_root = crate::tp_transport::return_row_to_root(&hop, r, pp, t * n_embd)?;
13476 let mut dst = out.slice_mut(0..t * n_embd);
13477 e.axpy_into(&pp_root, 1.0, &mut dst, t * n_embd)?;
13478 crate::glm5_tp::GLM5_EP_DIET_BULK_RETURNS.fetch_add(1, Ordering::Relaxed);
13479 }
13480 }
13481 crate::glm5_tp::GLM5_EP_GROUPED_PRIME_DISPATCHES.fetch_add(1, Ordering::Relaxed);
13482 static EPGP_LOGGED: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
13483 let layer_bit = 1u64 << (il as u64 % 64);
13484 if EPGP_LOGGED.fetch_or(layer_bit, Ordering::Relaxed) & layer_bit == 0 {
13485 eprintln!(
13486 "[glm5-ep-grouped-prime] execute layer={il} tokens={t} \
13487 provenance=ep-rank-slabs router=sigmoid-host-oracle epilogue=pre-clamped \
13488 combine=rank-partial-add transport={} performance_claim=false \
13489 (logged once per layer)",
13490 hop.transport.name(),
13491 );
13492 }
13493 Ok(Some(out))
13494 }
13495
13496 #[allow(clippy::too_many_arguments)]
13504 fn moe_shexp_add(
13505 e: &Engine,
13506 m: &MoeWeights,
13507 z: &CudaSlice<f32>,
13508 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
13509 t: usize,
13510 cfg: &ModelConfig,
13511 lim_shexp: Option<memra_gguf::config::SwigluClamp>,
13512 moe_out: &mut CudaSlice<f32>,
13513 ) -> Result<(), Box<dyn std::error::Error>> {
13514 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
13515 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
13516 {
13517 let n_embd = cfg.n_embd as usize;
13518 let n_ff_sh = gate_shexp.out_features(); let verify_t = t > 1 && t < PRIME_MIN_T;
13527 let (sg_gate, sg_up) = if t == 1 {
13528 shexp_gate_up_t1(e, gate_shexp, up_shexp, z, zq8)?
13529 } else if verify_t {
13530 (
13531 e.matmul_decode_exact(gate_shexp, z, t)?,
13532 e.matmul_decode_exact(up_shexp, z, t)?,
13533 )
13534 } else {
13535 (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?) };
13537 let mut sa = e.uninit(t * n_ff_sh)?; Self::ffn_act_lim(
13539 e,
13540 cfg,
13541 &sg_gate,
13542 &sg_up,
13543 1.0,
13544 1.0,
13545 lim_shexp,
13546 &mut sa,
13547 t * n_ff_sh,
13548 )?;
13549 let sh = if verify_t {
13550 e.matmul_decode_exact(down_shexp, &sa, t)?
13551 } else {
13552 e.matmul(down_shexp, &sa, t)?
13553 }; if m.gate_inp_shexp.is_none() && crate::htod_diet_on() {
13572 crate::HTOD_DIET_AVOIDED.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
13573 e.add_scaled_rows_ones(&sh, moe_out, n_embd, t)?;
13574 return Ok(());
13575 }
13576 let g = match &m.gate_inp_shexp {
13577 Some(gate_inp_shexp) => {
13578 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
13579 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
13580 } else {
13581 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
13582 let mut g = e.uninit(t)?; e.sigmoid(&gs, &mut g, t)?;
13584 g
13585 }
13586 }
13587 None => e.htod(&vec![1.0f32; t])?,
13588 };
13589 e.add_scaled_rows(&sh, &g, moe_out, n_embd, t)?;
13591 }
13592 Ok(())
13593 }
13594
13595 pub fn stage1_h2d_per_token(&self) -> u64 {
13598 use crate::hybrid::Ffn;
13599 let n_used = self
13600 .cfg
13601 .moe
13602 .as_ref()
13603 .map(|m| m.expert_used_count as u64)
13604 .unwrap_or(0);
13605 let mut bytes = 0u64;
13606 for l in self.layers.iter() {
13607 if let Ffn::Moe(m) = &l.ffn {
13608 bytes += n_used
13609 * (m.gate_exps.max_expert_bytes()
13610 + m.up_exps.max_expert_bytes()
13611 + m.down_exps.max_expert_bytes()) as u64;
13612 }
13613 }
13614 bytes
13615 }
13616
13617 pub(crate) fn max_moe_block(&self) -> usize {
13621 use crate::hybrid::Ffn;
13622 let mut mx = 0usize;
13623 let mut scan = |ffn: &Ffn| {
13624 if let Ffn::Moe(m) = ffn {
13625 mx = mx
13626 .max(m.gate_exps.max_expert_bytes())
13627 .max(m.up_exps.max_expert_bytes())
13628 .max(m.down_exps.max_expert_bytes());
13629 }
13630 };
13631 for l in self.layers.iter() {
13632 scan(&l.ffn);
13633 }
13634 if let Some(mtp) = self.mtp.as_ref() {
13635 scan(&mtp.ffn);
13636 }
13637 mx
13638 }
13639
13640 pub(crate) fn moe_cache_block_sizes(&self) -> Vec<usize> {
13643 use crate::hybrid::Ffn;
13644 let mut sizes = Vec::new();
13645 let mut scan = |ffn: &Ffn| {
13646 let Ffn::Moe(m) = ffn else { return };
13647 for ex in 0..m.gate_exps.n_expert {
13648 if m.active_experts.as_ref().is_some_and(|active| !active[ex]) {
13649 continue;
13650 }
13651 for exps in [&m.gate_exps, &m.up_exps, &m.down_exps] {
13652 let len = exps.expert_layout(ex).len;
13653 if len > 0 {
13654 sizes.push(len);
13655 }
13656 }
13657 }
13658 };
13659 for layer in &self.layers {
13660 scan(&layer.ffn);
13661 }
13662 if let Some(mtp) = &self.mtp {
13663 scan(&mtp.ffn);
13664 }
13665 sizes
13666 }
13667
13668 pub fn save_cpu_expert_residency_profile(
13674 &self,
13675 e: &Engine,
13676 path: &std::path::Path,
13677 ) -> Result<(), Box<dyn std::error::Error>> {
13678 let Some(ids) = e.export_moe_residency() else {
13679 return Err("no MoE residency cache to persist".into());
13680 };
13681 let mut body = format!(
13682 "memra-freeze-profile v1 max_block={} blocks={}\n",
13683 self.max_moe_block(),
13684 ids.len()
13685 );
13686 for (layer, proj, ex) in &ids {
13687 body.push_str(&format!("{layer} {proj} {ex}\n"));
13688 }
13689 let tmp = path.with_extension("tmp");
13690 std::fs::write(&tmp, body)?;
13691 std::fs::rename(&tmp, path)?;
13692 println!(
13693 "[moe-cache] freeze profile saved: {} blocks -> {}",
13694 ids.len(),
13695 path.display()
13696 );
13697 Ok(())
13698 }
13699
13700 pub fn restore_cpu_expert_residency_profile(
13704 &self,
13705 e: &Engine,
13706 path: &std::path::Path,
13707 ) -> Result<bool, Box<dyn std::error::Error>> {
13708 use crate::hybrid::Ffn;
13709 use crate::moe_cache::BlockId;
13710 let Ok(content) = std::fs::read_to_string(path) else {
13711 return Ok(false);
13712 };
13713 let mut lines = content.lines();
13714 let Some(header) = lines.next() else {
13715 return Ok(false);
13716 };
13717 let expected = format!("memra-freeze-profile v1 max_block={}", self.max_moe_block());
13718 if !header.starts_with(&expected) {
13719 println!(
13720 "[moe-cache] freeze profile ignored (geometry mismatch): {}",
13721 path.display()
13722 );
13723 return Ok(false);
13724 }
13725 let mut by_layer: std::collections::HashMap<u16, Vec<BlockId>> =
13726 std::collections::HashMap::new();
13727 for line in lines {
13728 let mut fields = line.split_whitespace();
13729 let (Some(layer), Some(proj), Some(ex)) = (fields.next(), fields.next(), fields.next())
13730 else {
13731 continue;
13732 };
13733 let (Ok(layer), Ok(proj), Ok(ex)) =
13734 (layer.parse::<u16>(), proj.parse::<u8>(), ex.parse::<u16>())
13735 else {
13736 continue;
13737 };
13738 by_layer
13739 .entry(layer)
13740 .or_default()
13741 .push(BlockId::new(layer, proj, ex));
13742 }
13743 let requested: usize = by_layer.values().map(Vec::len).sum();
13744 if requested == 0 {
13745 return Ok(false);
13746 }
13747 let max_block = self.max_moe_block();
13748 let mut restaged = 0usize;
13749 let mut stage_layer =
13750 |layer_index: u16, ffn: &Ffn| -> Result<(), Box<dyn std::error::Error>> {
13751 let Ffn::Moe(m) = ffn else { return Ok(()) };
13752 let Some(ids) = by_layer.get(&layer_index) else {
13753 return Ok(());
13754 };
13755 e.with_moe_cache(max_block, |cache, eng| {
13756 for id in ids {
13757 if cache.restage_block(*id, m, eng)? {
13758 restaged += 1;
13759 }
13760 }
13761 Ok(())
13762 })
13763 };
13764 for (index, layer) in self.layers.iter().enumerate() {
13765 stage_layer(index as u16, &layer.ffn)?;
13766 }
13767 if let Some(mtp) = self.mtp.as_ref() {
13768 stage_layer(u16::MAX, &mtp.ffn)?;
13769 }
13770 e.freeze_moe_cache();
13771 println!(
13772 "[moe-cache] freeze profile restored: {restaged}/{requested} blocks restaged from {}",
13773 path.display()
13774 );
13775 Ok(true)
13776 }
13777
13778 pub fn freeze_cpu_expert_residency(
13780 &self,
13781 e: &Engine,
13782 ) -> Result<(), Box<dyn std::error::Error>> {
13783 e.freeze_moe_cache();
13784 Ok(())
13785 }
13786
13787 pub fn ffn_act(
13795 e: &Engine,
13796 cfg: &ModelConfig,
13797 gate: &CudaSlice<f32>,
13798 up: &CudaSlice<f32>,
13799 act: &mut CudaSlice<f32>,
13800 n: usize,
13801 ) -> Result<(), Box<dyn std::error::Error>> {
13802 Self::ffn_act_scaled(e, cfg, gate, up, 1.0, 1.0, act, n)
13803 }
13804
13805 #[allow(clippy::too_many_arguments)]
13809 pub(crate) fn ffn_act_scaled(
13810 e: &Engine,
13811 cfg: &ModelConfig,
13812 gate: &CudaSlice<f32>,
13813 up: &CudaSlice<f32>,
13814 gs: f32,
13815 us: f32,
13816 act: &mut CudaSlice<f32>,
13817 n: usize,
13818 ) -> Result<(), Box<dyn std::error::Error>> {
13819 Self::ffn_act_lim(e, cfg, gate, up, gs, us, None, act, n)
13820 }
13821
13822 #[allow(clippy::too_many_arguments)]
13833 pub(crate) fn ffn_act_lim(
13834 e: &Engine,
13835 cfg: &ModelConfig,
13836 gate: &CudaSlice<f32>,
13837 up: &CudaSlice<f32>,
13838 gs: f32,
13839 us: f32,
13840 limit: Option<SwigluClamp>,
13841 act: &mut CudaSlice<f32>,
13842 n: usize,
13843 ) -> Result<(), Box<dyn std::error::Error>> {
13844 if let Some(m3) = cfg.m3.as_ref() {
13845 debug_assert!(
13846 limit.is_none(),
13847 "m3 swigluoai and the step35/glm5_next clamps are different archs"
13848 );
13849 return e.swigluoai_mul_scaled(
13850 gate,
13851 up,
13852 gs,
13853 us,
13854 m3.swiglu_alpha,
13855 m3.swiglu_limit,
13856 act,
13857 n,
13858 );
13859 }
13860 match limit {
13861 Some(SwigluClamp::Post(l)) => {
13862 return e.swiglu_clamped_mul_scaled(gate, up, gs, us, l, act, n);
13863 }
13864 Some(SwigluClamp::Pre(l)) => {
13865 return e.swiglu_preclamped_mul_scaled(gate, up, gs, us, l, act, n);
13866 }
13867 None => {}
13868 }
13869 if gs == 1.0 && us == 1.0 {
13870 return e.silu_mul(gate, up, act, n);
13871 }
13872 e.silu_mul_scaled(gate, up, gs, us, act, n)
13873 }
13874
13875 fn fused_post_limit(lim: Option<SwigluClamp>) -> Result<Option<f32>, ()> {
13881 match lim {
13882 None => Ok(None),
13883 Some(SwigluClamp::Post(l)) => Ok(Some(l)),
13884 Some(SwigluClamp::Pre(_)) => Err(()),
13885 }
13886 }
13887
13888 fn moe_route(
13894 e: &Engine,
13895 logits: &CudaSlice<f32>,
13896 t: usize,
13897 n_expert: usize,
13898 n_used: usize,
13899 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
13900 Self::moe_route_cfg(e, logits, t, n_expert, n_used, None)
13901 }
13902
13903 #[allow(clippy::too_many_arguments)]
13911 fn moe_route_sigmoid_cfg(
13912 e: &Engine,
13913 logits: &CudaSlice<f32>,
13914 t: usize,
13915 n_expert: usize,
13916 n_used: usize,
13917 m: &MoeWeights,
13918 (sf, route_norm): (f32, bool),
13919 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
13920 if sigmoid_router_enabled() {
13921 return e.moe_router_sigmoid_topk_host(
13922 logits,
13923 t,
13924 n_expert,
13925 n_used,
13926 m.active_count(),
13927 &m.exp_probs_b_dev,
13928 &m.active_experts_dev,
13929 sf,
13930 route_norm,
13931 );
13932 }
13933 let lg = e.dtoh(logits)?;
13934 Self::moe_route_sigmoid_host(
13935 &lg,
13936 t,
13937 n_expert,
13938 n_used,
13939 m.exp_probs_b.as_deref(),
13940 sf,
13941 route_norm,
13942 m.active_experts.as_deref(),
13943 )
13944 }
13945
13946 #[allow(clippy::excessive_precision)] fn moe_route_cfg(
13950 e: &Engine,
13951 logits: &CudaSlice<f32>,
13952 t: usize,
13953 n_expert: usize,
13954 n_used: usize,
13955 active: Option<&[bool]>,
13956 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
13957 if active.is_none() && !matches!(std::env::var("MEMRA_FUSED_ROUTER").as_deref(), Ok("0")) {
13960 return e.moe_router_topk_host(logits, t, n_expert, n_used);
13961 }
13962 let lg = e.dtoh(logits)?; let mut sel = vec![0u32; t * n_used];
13965 let mut w_out = vec![0f32; t * n_used];
13966 for tok in 0..t {
13967 let row = &lg[tok * n_expert..(tok + 1) * n_expert];
13968 let maxl = row
13970 .iter()
13971 .enumerate()
13972 .filter(|(i, _)| active.is_none_or(|mask| mask[*i]))
13973 .map(|(_, &x)| x)
13974 .fold(f32::NEG_INFINITY, f32::max);
13975 let mut probs = vec![0f32; n_expert];
13976 let mut den = 0f32;
13977 for i in 0..n_expert {
13978 if active.is_some_and(|mask| !mask[i]) {
13979 continue;
13980 }
13981 let x = (row[i] - maxl).exp();
13982 probs[i] = x;
13983 den += x;
13984 }
13985 for p in probs.iter_mut() {
13986 *p /= den;
13987 }
13988 let mut idx: Vec<usize> = (0..n_expert)
13990 .filter(|&i| active.is_none_or(|mask| mask[i]))
13991 .collect();
13992 idx.sort_by(|&a, &b| probs[b].total_cmp(&probs[a]).then(a.cmp(&b)));
13993 let sl = &idx[..n_used];
13994 let mut wv: Vec<f32> = sl.iter().map(|&i| probs[i]).collect();
13995 let mut ws: f32 = wv.iter().sum();
13996 ws = ws.max(6.103515625e-5_f32); for x in wv.iter_mut() {
13998 *x /= ws;
13999 }
14000 for j in 0..n_used {
14001 sel[tok * n_used + j] = sl[j] as u32;
14002 w_out[tok * n_used + j] = wv[j];
14003 }
14004 }
14005 Ok((sel, w_out))
14006 }
14007
14008 #[allow(clippy::too_many_arguments)]
14009 #[allow(clippy::type_complexity)] fn moe_route_sigmoid_with_input(
14011 e: &Engine,
14012 logits: &CudaSlice<f32>,
14013 input: &CudaSlice<f32>,
14014 t: usize,
14015 in_features: usize,
14016 n_expert: usize,
14017 n_used: usize,
14018 bias: Option<&[f32]>,
14019 (sf, route_norm): (f32, bool),
14020 active: Option<&[bool]>,
14021 ) -> Result<(Vec<u32>, Vec<f32>, Vec<f32>), Box<dyn std::error::Error>> {
14022 let logit_values =
14023 active_matrix_values(logits.len(), t, n_expert, "sigmoid router logits")?;
14024 let input_values =
14025 active_matrix_values(input.len(), t, in_features, "sigmoid router input")?;
14026 let (lg, input) = e.dtoh_pair_views(
14027 &logits.slice(0..logit_values),
14028 &input.slice(0..input_values),
14029 )?;
14030 let (sel, w) =
14031 Self::moe_route_sigmoid_host(&lg, t, n_expert, n_used, bias, sf, route_norm, active)?;
14032 Ok((sel, w, input))
14033 }
14034
14035 pub fn start_moe_prefetch_predictor(
14040 &self,
14041 e: &Engine,
14042 cfg: &ModelConfig,
14043 ) -> Result<(), Box<dyn std::error::Error>> {
14044 use crate::hybrid::Ffn;
14045 let Some(sig) = cfg.sigmoid_router() else {
14046 return Err("prefetch predictor requires a sigmoid-router arch".into());
14047 };
14048 let resident: std::collections::HashSet<(u16, u8, u16)> = e
14049 .export_moe_residency()
14050 .ok_or("prefetch predictor needs the frozen MoE residency cache")?
14051 .into_iter()
14052 .collect();
14053 let mut layers = Vec::new();
14054 for (index, layer) in self.layers.iter().enumerate() {
14055 let Ffn::Moe(m) = &layer.ffn else { continue };
14056 let crate::model::GpuTensor::Float { data, .. } = &m.gate_inp else {
14057 continue;
14058 };
14059 let router = e.dtoh(data)?;
14060 let n_expert = m.gate_exps.n_expert;
14061 let n_embd = m.gate_exps.in_f;
14062 if router.len() != n_embd * n_expert {
14063 continue;
14064 }
14065 let build = |exps: &crate::model::HostExps| {
14066 (0..n_expert)
14067 .map(|expert| crate::cpu_experts::predictor_projection(exps, expert))
14068 .collect::<Vec<_>>()
14069 };
14070 layers.push((
14071 index as u16,
14072 crate::cpu_experts::PredictLayerInit {
14073 router,
14074 bias: m.exp_probs_b.clone(),
14075 active: m.active_experts.clone(),
14076 n_embd,
14077 n_used: cfg
14078 .moe
14079 .as_ref()
14080 .map(|moe| moe.expert_used_count as usize)
14081 .ok_or("prefetch predictor requires MoE config")?,
14082 sig,
14083 weights_n_expert: n_expert,
14084 gate: build(&m.gate_exps),
14085 up: build(&m.up_exps),
14086 down: build(&m.down_exps),
14087 },
14088 ));
14089 }
14090 crate::cpu_experts::start_prefetch_predictor(layers, resident).map_err(|error| error.into())
14091 }
14092
14093 #[allow(clippy::too_many_arguments)]
14096 pub fn moe_route_sigmoid_host_public(
14097 logits: &[f32],
14098 t: usize,
14099 n_expert: usize,
14100 n_used: usize,
14101 bias: Option<&[f32]>,
14102 sf: f32,
14103 route_norm: bool,
14104 active: Option<&[bool]>,
14105 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
14106 Self::moe_route_sigmoid_host(logits, t, n_expert, n_used, bias, sf, route_norm, active)
14107 }
14108
14109 #[allow(clippy::too_many_arguments)]
14110 fn moe_route_sigmoid_host(
14111 lg: &[f32],
14112 t: usize,
14113 n_expert: usize,
14114 n_used: usize,
14115 bias: Option<&[f32]>,
14116 sf: f32,
14117 route_norm: bool,
14118 active: Option<&[bool]>,
14119 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
14120 let active_count = active
14121 .map(|mask| mask.iter().filter(|&&enabled| enabled).count())
14122 .unwrap_or(n_expert);
14123 crate::sigrouter_contract::validate_active_count(n_used, active_count)?;
14124 if lg.len() != t * n_expert {
14125 return Err(format!(
14126 "sigmoid router logits length mismatch: got {}, expected {}",
14127 lg.len(),
14128 t * n_expert,
14129 )
14130 .into());
14131 }
14132 let mut sel = vec![0u32; t * n_used];
14133 let mut w_out = vec![0f32; t * n_used];
14134 for tok in 0..t {
14135 let row = &lg[tok * n_expert..(tok + 1) * n_expert];
14136 let scores: Vec<f32> = row.iter().map(|&x| 1.0 / (1.0 + (-x).exp())).collect();
14137 let selsc: Vec<f32> = match bias {
14139 Some(b) => scores.iter().zip(b).map(|(s, bb)| s + bb).collect(),
14140 None => scores.clone(),
14141 };
14142 let mut idx: Vec<usize> = (0..n_expert)
14143 .filter(|&i| active.is_none_or(|mask| mask[i]))
14144 .collect();
14145 idx.sort_by(|&a, &b| selsc[b].total_cmp(&selsc[a]).then(a.cmp(&b)));
14146 let sl = &idx[..n_used];
14147 let mut wv: Vec<f32> = sl.iter().map(|&i| scores[i]).collect();
14148 if route_norm {
14149 let ws: f32 = wv.iter().sum::<f32>().max(1e-20);
14150 for x in wv.iter_mut() {
14151 *x = *x / ws * sf;
14152 }
14153 } else {
14154 for x in wv.iter_mut() {
14155 *x *= sf;
14156 }
14157 }
14158 for j in 0..n_used {
14159 sel[tok * n_used + j] = sl[j] as u32;
14160 w_out[tok * n_used + j] = wv[j];
14161 }
14162 }
14163 Ok((sel, w_out))
14164 }
14165
14166 #[allow(clippy::too_many_arguments)]
14170 fn moe_ffn_sigmoid_dev(
14171 e: &Engine,
14172 m: &MoeWeights,
14173 z: &CudaSlice<f32>,
14174 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
14175 logits: &CudaSlice<f32>,
14176 t: usize,
14177 cfg: &ModelConfig,
14178 il: u16,
14179 (scaling_factor, route_norm): (f32, bool),
14180 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
14181 crate::moe_rp_refuse(
14182 m.dev_exps.as_ref().is_some_and(|d| d.rp),
14183 "moe_ffn_sigmoid_dev",
14184 )?; let moe = cfg.moe.as_ref().unwrap();
14186 let n_embd = cfg.n_embd as usize;
14187 let n_expert = moe.expert_count as usize;
14188 let n_used = moe.expert_used_count as usize;
14189 let n_ff_exp = moe.expert_ff_length as usize;
14190 let dev = m.dev_exps.as_ref().unwrap();
14191 debug_assert_eq!(dev.dev, e.ctx().ordinal());
14192 debug_assert!(m.has_uniform_expert_layout());
14193 debug_assert!(!m.has_macros);
14194
14195 let (sel_d, w_d) = e.moe_router_sigmoid_topk(
14196 logits,
14197 t,
14198 n_expert,
14199 n_used,
14200 m.active_count(),
14201 &m.exp_probs_b_dev,
14202 &m.active_experts_dev,
14203 scaling_factor,
14204 route_norm,
14205 )?;
14206 crate::moesd::record_device_routes(e, il, n_expert, n_used, &sel_d)?;
14207 crate::moe_sel_dump::record_device(e, il, t, n_used, &sel_d, &w_d)?;
14208 if let Some(fp8) = dev.fp8_blk.as_ref() {
14209 debug_assert_eq!(m.gate_exps.qtype, crate::QT_F8_E4M3_BLK);
14210 debug_assert_eq!(m.up_exps.qtype, crate::QT_F8_E4M3_BLK);
14211 debug_assert_eq!(m.down_exps.qtype, crate::QT_F8_E4M3_BLK);
14212 debug_assert_eq!(fp8.gate.rows, m.gate_exps.out_f.div_ceil(128));
14213 debug_assert_eq!(fp8.up.rows, m.up_exps.out_f.div_ceil(128));
14214 debug_assert_eq!(fp8.down.rows, m.down_exps.out_f.div_ceil(128));
14215
14216 let selected = e.dtoh_i32(&sel_d)?;
14223 let route_weights = e.dtoh(&w_d)?;
14224 let mut moe_out = e.zeros(t * n_embd)?;
14225 for tok in 0..t {
14226 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
14227 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
14228 for j in 0..n_used {
14229 let pair = tok * n_used + j;
14230 let expert = selected[pair] as usize;
14231 let gate = Self::moe_resident_fp8_e4m3(
14232 e,
14233 &m.gate_exps,
14234 &dev.gate,
14235 &fp8.gate,
14236 expert,
14237 &zt,
14238 1,
14239 )?;
14240 let up = Self::moe_resident_fp8_e4m3(
14241 e, &m.up_exps, &dev.up, &fp8.up, expert, &zt, 1,
14242 )?;
14243 let mut act = e.uninit(n_ff_exp)?;
14244 Self::ffn_act_lim(
14245 e,
14246 cfg,
14247 &gate,
14248 &up,
14249 1.0,
14250 1.0,
14251 cfg.clamp_exp_at(il as u32),
14252 &mut act,
14253 n_ff_exp,
14254 )?;
14255 let act = act.slice(0..n_ff_exp);
14256 let down = Self::moe_resident_fp8_e4m3(
14257 e,
14258 &m.down_exps,
14259 &dev.down,
14260 &fp8.down,
14261 expert,
14262 &act,
14263 1,
14264 )?;
14265 e.axpy_into(&down, route_weights[pair], &mut dst, n_embd)?;
14266 }
14267 }
14268 if std::env::var("MEMRA_SIG_ROUTER_DISPATCH_TRACE").as_deref() == Ok("1") {
14269 eprintln!(
14270 "[sigrouter-dev] layer={il} tokens={t} experts={n_expert} used={n_used} \
14271 native=fp8blk-w8a8-e4m3-reference clamp={}",
14272 cfg.clamp_exp_at(il as u32).is_some(),
14273 );
14274 }
14275 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut moe_out)?;
14276 return Ok(moe_out);
14277 }
14278 let (gate_row_bytes, up_row_bytes) = if dev.gu_il {
14279 let combined = m.gate_exps.row_bytes + m.up_exps.row_bytes;
14280 (combined, combined)
14281 } else {
14282 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
14283 };
14284 let (zq, zd) = match (t, zq8) {
14285 (1, Some((q, d))) => (q.clone(), d.clone()),
14286 _ => e.quantize_q8_1(z, t, n_embd)?,
14287 };
14288 let n_pairs = t * n_used;
14289 let mut moe_out = if cfg.clamp_exp_at(il as u32).is_some() {
14290 let pair_tok: Vec<i32> = (0..n_pairs).map(|pair| (pair / n_used) as i32).collect();
14294 let pair_tok_d = e.htod_i32(&pair_tok)?;
14295 let gate = e.moe_pairs_matvec_q8(
14296 &dev.ptr_row,
14297 0,
14298 &pair_tok_d,
14299 &sel_d,
14300 &zq,
14301 &zd,
14302 n_embd,
14303 n_ff_exp,
14304 n_expert,
14305 n_pairs,
14306 m.gate_exps.qtype,
14307 gate_row_bytes,
14308 )?;
14309 let up = e.moe_pairs_matvec_q8(
14310 &dev.ptr_row,
14311 1,
14312 &pair_tok_d,
14313 &sel_d,
14314 &zq,
14315 &zd,
14316 n_embd,
14317 n_ff_exp,
14318 n_expert,
14319 n_pairs,
14320 m.up_exps.qtype,
14321 up_row_bytes,
14322 )?;
14323 let mut act = e.uninit(n_pairs * n_ff_exp)?;
14324 Self::ffn_act_lim(
14325 e,
14326 cfg,
14327 &gate,
14328 &up,
14329 1.0,
14330 1.0,
14331 cfg.clamp_exp_at(il as u32),
14332 &mut act,
14333 n_pairs * n_ff_exp,
14334 )?;
14335 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
14336 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
14337 let pair_self_d = e.htod_i32(&pair_self)?;
14338 let down = e.moe_pairs_matvec_q8(
14339 &dev.ptr_row,
14340 2,
14341 &pair_self_d,
14342 &sel_d,
14343 &aq2,
14344 &ad2,
14345 n_ff_exp,
14346 n_embd,
14347 n_expert,
14348 n_pairs,
14349 m.down_exps.qtype,
14350 m.down_exps.row_bytes,
14351 )?;
14352 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
14353 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
14354 let tok_off_d = e.htod_i32(&tok_off)?;
14355 let tok_ids_d = e.htod_i32(&tok_ids)?;
14356 let mut output = e.uninit(t * n_embd)?;
14357 e.moe_pairs_scatter(&down, &w_d, &tok_off_d, &tok_ids_d, &mut output, t, n_embd)?;
14358 output
14359 } else {
14360 let act = e.moe_gate_up_silu8_dev_q8_rows(
14361 &dev.ptr_row,
14362 &sel_d,
14363 &zq,
14364 &zd,
14365 t,
14366 n_embd,
14367 n_ff_exp,
14368 n_used,
14369 n_expert,
14370 m.gate_exps.qtype,
14371 m.up_exps.qtype,
14372 gate_row_bytes,
14373 up_row_bytes,
14374 &m.dev_macros,
14375 )?;
14376 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
14377 let mut output = e.uninit(t * n_embd)?;
14378 e.moe_down8_fma_dev_q8_rows_g(
14379 &dev.ptr_row,
14380 &sel_d,
14381 &w_d,
14382 &aq2,
14383 &ad2,
14384 &mut output,
14385 t,
14386 n_ff_exp,
14387 n_embd,
14388 n_used,
14389 n_expert,
14390 m.down_exps.qtype,
14391 m.down_exps.row_bytes,
14392 )?;
14393 output
14394 };
14395
14396 if std::env::var("MEMRA_SIG_ROUTER_DISPATCH_TRACE").as_deref() == Ok("1") {
14397 eprintln!(
14398 "[sigrouter-dev] layer={il} tokens={t} experts={n_expert} used={n_used} clamp={} gu_il={}",
14399 cfg.clamp_exp_at(il as u32).is_some(),
14400 dev.gu_il,
14401 );
14402 }
14403 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut moe_out)?;
14404 Ok(moe_out)
14405 }
14406
14407 #[allow(clippy::too_many_arguments)]
14408 fn moe_resident_fp8_e4m3(
14409 e: &Engine,
14410 exps: &crate::model::HostExps,
14411 bytes: &CudaSlice<u8>,
14412 scales: &crate::hybrid::DevExpertFp8ProjectionScales,
14413 expert: usize,
14414 x: &cudarc::driver::CudaView<f32>,
14415 m: usize,
14416 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
14417 let layout = exps.expert_layout(expert);
14418 debug_assert_eq!(layout.qtype, crate::QT_F8_E4M3_BLK);
14419 debug_assert_eq!(scales.rows * scales.cols, scales.expert_stride);
14420 let byte_start = expert * exps.expert_stride;
14421 let scale_start = expert * scales.expert_stride;
14422 let weight = bytes.slice(byte_start..byte_start + layout.len);
14423 let scale = scales
14424 .scales
14425 .slice(scale_start..scale_start + scales.expert_stride);
14426 e.qmatvec_mmq_fp8_blk_view(&weight, &scale, x, m, exps.in_f, exps.out_f)
14427 }
14428
14429 fn moe_ffn_pairs(
14438 e: &Engine,
14439 m: &MoeWeights,
14440 z: &CudaSlice<f32>,
14441 logits: &CudaSlice<f32>,
14442 t: usize,
14443 cfg: &ModelConfig,
14444 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
14445 crate::moe_rp_refuse(m.dev_exps.as_ref().is_some_and(|d| d.rp), "moe_ffn_pairs")?; let moe = cfg.moe.as_ref().unwrap();
14447 let n_embd = cfg.n_embd as usize;
14448 let n_expert = moe.expert_count as usize;
14449 let n_used = moe.expert_used_count as usize;
14450 let n_ff_exp = moe.expert_ff_length as usize;
14451 debug_assert!(
14456 !cfg.swiglu_clamped_anywhere(),
14457 "moe_ffn_pairs has no per-layer clamp: fused epilogues are plain SiLU"
14458 );
14459 let dev = m.dev_exps.as_ref().unwrap();
14460 let (rbg_d, rbu_d) = if dev.gu_il {
14462 let sxx = m.gate_exps.row_bytes + m.up_exps.row_bytes;
14463 (sxx, sxx)
14464 } else {
14465 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
14466 };
14467
14468 let (sel_all, w_all) = Self::moe_route(e, logits, t, n_expert, n_used)?;
14469 let n_pairs = t * n_used;
14470 let pair_tok: Vec<i32> = (0..n_pairs).map(|p| (p / n_used) as i32).collect();
14473 let pair_ex: Vec<i32> = sel_all.iter().map(|&x| x as i32).collect();
14474 let pair_w: Vec<f32> = w_all.clone();
14475 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
14476 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
14477 let pt = e.htod_i32(&pair_tok)?;
14478 let px = e.htod_i32(&pair_ex)?;
14479 let pw = e.htod(&pair_w)?;
14480 let toff = e.htod_i32(&tok_off)?;
14481 let tids = e.htod_i32(&tok_ids)?;
14482
14483 let mut by_ex: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
14487 for p in 0..n_pairs {
14488 by_ex[pair_ex[p] as usize].push(p as i32);
14489 }
14490 let mut ex_ids: Vec<i32> = Vec::new();
14491 let mut ex_off: Vec<i32> = vec![0];
14492 let mut ex_pairs: Vec<i32> = Vec::with_capacity(n_pairs);
14493 for (ex, list) in by_ex.iter().enumerate() {
14494 if list.is_empty() {
14495 continue;
14496 }
14497 ex_ids.push(ex as i32);
14498 ex_pairs.extend_from_slice(list);
14499 ex_off.push(ex_pairs.len() as i32);
14500 }
14501 let n_active = ex_ids.len();
14502 let exi = e.htod_i32(&ex_ids)?;
14503 let exo = e.htod_i32(&ex_off)?;
14504 let exp_d = e.htod_i32(&ex_pairs)?;
14505 let _ = &px; static MMA_T: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
14526 let mma_t = *MMA_T.get_or_init(|| {
14527 std::env::var("MEMRA_MOE_MMA_T")
14528 .ok()
14529 .and_then(|v| v.parse().ok())
14530 .unwrap_or(16)
14531 });
14532 let use_mma = std::env::var("MEMRA_MOE_MMA")
14533 .map(|v| v != "0")
14534 .unwrap_or(true)
14535 && t >= mma_t
14536 && q8_expert_dec_supported(m.gate_exps.qtype)
14537 && q8_expert_dec_supported(m.up_exps.qtype)
14538 && q8_expert_dec_supported(m.down_exps.qtype)
14539 && n_embd.is_multiple_of(256)
14540 && n_ff_exp.is_multiple_of(256);
14541 let mma_capable = q8_expert_dec_supported(m.gate_exps.qtype)
14557 && q8_expert_dec_supported(m.up_exps.qtype)
14558 && q8_expert_dec_supported(m.down_exps.qtype)
14559 && n_embd.is_multiple_of(256)
14560 && n_ff_exp.is_multiple_of(256);
14561 let f16g_mode = crate::moe_f16g_mode();
14562 let f16g = f16g_mode != 0
14563 && t >= mma_t
14564 && (f16g_mode != 3 || !mma_capable)
14565 && f16g_proj_ok(m.gate_exps.qtype, n_embd)
14566 && f16g_proj_ok(m.up_exps.qtype, n_embd)
14567 && f16g_proj_ok(m.down_exps.qtype, n_ff_exp);
14568 if use_mma || f16g {
14569 let y_down = if f16g {
14577 let csr_tok: Vec<i32> = ex_pairs.iter().map(|&p| p / n_used as i32).collect();
14581 let csr_tok_d = e.htod_i32(&csr_tok)?;
14582 let (z_f16, z_s) = e.moe_f16g_act(z, Some(&csr_tok_d), n_embd, n_pairs)?;
14583 let g_csr = e.moe_f16_grouped(
14584 &dev.ptr_row,
14585 0,
14586 n_expert,
14587 &exi,
14588 &ex_off,
14589 &exo,
14590 &z_f16,
14591 &z_s,
14592 n_embd,
14593 n_ff_exp,
14594 n_active,
14595 n_pairs,
14596 m.gate_exps.qtype,
14597 rbg_d,
14598 )?;
14599 let u_csr = e.moe_f16_grouped(
14600 &dev.ptr_row,
14601 1,
14602 n_expert,
14603 &exi,
14604 &ex_off,
14605 &exo,
14606 &z_f16,
14607 &z_s,
14608 n_embd,
14609 n_ff_exp,
14610 n_active,
14611 n_pairs,
14612 m.up_exps.qtype,
14613 rbu_d,
14614 )?;
14615 let act_csr = e.moe_pairs_silu_mul(&g_csr, &u_csr, n_pairs * n_ff_exp)?;
14616 let (a_f16, a_s) = e.moe_f16g_act(&act_csr, None, n_ff_exp, n_pairs)?;
14617 let d_csr = e.moe_f16_grouped(
14618 &dev.ptr_row,
14619 2,
14620 n_expert,
14621 &exi,
14622 &ex_off,
14623 &exo,
14624 &a_f16,
14625 &a_s,
14626 n_ff_exp,
14627 n_embd,
14628 n_active,
14629 n_pairs,
14630 m.down_exps.qtype,
14631 m.down_exps.row_bytes,
14632 )?;
14633 e.rows_permute(&d_csr, &exp_d, n_pairs, n_embd)?
14634 } else {
14635 let z_scr = e.mmq_iq_quantize_act(z, n_embd, t)?;
14637 let gate = e.mmq_iq_experts(
14638 &dev.ptr_row,
14639 0,
14640 n_expert,
14641 &exi,
14642 &exo,
14643 &exp_d,
14644 &pt,
14645 &z_scr,
14646 n_embd,
14647 n_ff_exp,
14648 n_active,
14649 n_pairs,
14650 t,
14651 m.gate_exps.qtype,
14652 rbg_d,
14653 )?;
14654 let up = e.mmq_iq_experts(
14655 &dev.ptr_row,
14656 1,
14657 n_expert,
14658 &exi,
14659 &exo,
14660 &exp_d,
14661 &pt,
14662 &z_scr,
14663 n_embd,
14664 n_ff_exp,
14665 n_active,
14666 n_pairs,
14667 t,
14668 m.up_exps.qtype,
14669 rbu_d,
14670 )?;
14671 let a_scr = if crate::moe_fuse_actq_on() {
14677 e.mmq_iq_fused_act_quant(&gate, &up, n_ff_exp, n_pairs, 0)?
14678 } else {
14679 let act = e.moe_pairs_silu_mul(&gate, &up, n_pairs * n_ff_exp)?;
14680 e.mmq_iq_quantize_act(&act, n_ff_exp, n_pairs)?
14681 };
14682 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
14683 let pself = e.htod_i32(&pair_self)?;
14684 e.mmq_iq_experts(
14685 &dev.ptr_row,
14686 2,
14687 n_expert,
14688 &exi,
14689 &exo,
14690 &exp_d,
14691 &pself,
14692 &a_scr,
14693 n_ff_exp,
14694 n_embd,
14695 n_active,
14696 n_pairs,
14697 n_pairs,
14698 m.down_exps.qtype,
14699 m.down_exps.row_bytes,
14700 )?
14701 };
14702 let mut moe_out = e.uninit(t * n_embd)?;
14703 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
14704 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
14705 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
14706 {
14707 let n_ff_sh = gate_shexp.out_features();
14708 let sg_gate = e.matmul(gate_shexp, z, t)?;
14709 let sg_up = e.matmul(up_shexp, z, t)?;
14710 let mut sa = e.uninit(t * n_ff_sh)?;
14711 Self::ffn_act(e, cfg, &sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
14712 let sh = e.matmul(down_shexp, &sa, t)?;
14713 let g = match &m.gate_inp_shexp {
14719 Some(gate_inp_shexp) if crate::router_prefill_exact_on() => {
14720 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
14721 }
14722 Some(gate_inp_shexp) => {
14723 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
14724 let mut g = e.uninit(t)?;
14725 e.sigmoid(&gs, &mut g, t)?;
14726 g
14727 }
14728 None => e.htod(&vec![1.0f32; t])?,
14729 };
14730 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
14731 }
14732 return Ok(moe_out);
14733 }
14734
14735 let dec = std::env::var("MEMRA_MOE_DEC")
14738 .map(|v| v != "0")
14739 .unwrap_or(true);
14740 let matvec = |proj,
14741 exi: &_,
14742 exo: &_,
14743 exp_d: &_,
14744 pt: &_,
14745 aq: &_,
14746 ad: &_,
14747 inf,
14748 outf,
14749 qtype,
14750 rb|
14751 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
14752 let dec = dec && q8_expert_dec_supported(qtype);
14754 if dec {
14755 e.moe_pairs_matvec_q8_dec(
14756 &dev.ptr_row,
14757 proj,
14758 exi,
14759 exo,
14760 exp_d,
14761 pt,
14762 aq,
14763 ad,
14764 inf,
14765 outf,
14766 n_expert,
14767 n_active,
14768 n_pairs,
14769 qtype,
14770 rb,
14771 )
14772 } else {
14773 e.moe_pairs_matvec_q8_em(
14774 &dev.ptr_row,
14775 proj,
14776 exi,
14777 exo,
14778 exp_d,
14779 pt,
14780 aq,
14781 ad,
14782 inf,
14783 outf,
14784 n_expert,
14785 n_active,
14786 n_pairs,
14787 qtype,
14788 rb,
14789 )
14790 }
14791 };
14792 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
14793 let gate = matvec(
14794 0,
14795 &exi,
14796 &exo,
14797 &exp_d,
14798 &pt,
14799 &zq,
14800 &zd,
14801 n_embd,
14802 n_ff_exp,
14803 m.gate_exps.qtype,
14804 rbg_d,
14805 )?;
14806 let up = matvec(
14807 1,
14808 &exi,
14809 &exo,
14810 &exp_d,
14811 &pt,
14812 &zq,
14813 &zd,
14814 n_embd,
14815 n_ff_exp,
14816 m.up_exps.qtype,
14817 rbu_d,
14818 )?;
14819 let act = e.moe_pairs_silu_mul(&gate, &up, n_pairs * n_ff_exp)?;
14820 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
14821 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
14823 let pself = e.htod_i32(&pair_self)?;
14824 let y_down = matvec(
14825 2,
14826 &exi,
14827 &exo,
14828 &exp_d,
14829 &pself,
14830 &aq2,
14831 &ad2,
14832 n_ff_exp,
14833 n_embd,
14834 m.down_exps.qtype,
14835 m.down_exps.row_bytes,
14836 )?;
14837 let mut moe_out = e.uninit(t * n_embd)?; e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
14839
14840 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
14844 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
14845 {
14846 let n_ff_sh = gate_shexp.out_features();
14847 let step_exact = true;
14851 let verify_t = step_exact && t > 1 && t < PRIME_MIN_T;
14852 let (sg_gate, sg_up) = if step_exact && t == 1 {
14853 shexp_gate_up_t1(e, gate_shexp, up_shexp, z, None)?
14854 } else if verify_t {
14855 let mut fused = None;
14856 if crate::spec::spec_fused_t()
14857 && (2..=4).contains(&t)
14858 && e.uses_q8_1_fast(gate_shexp)
14859 && e.uses_q8_1_fast(up_shexp)
14860 {
14861 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
14862 fused = e.matmul_q8_fused2_t(gate_shexp, up_shexp, &zq, &zd, t)?;
14863 }
14864 match fused {
14865 Some(pair) => pair,
14866 None => (
14867 e.matmul_decode_exact(gate_shexp, z, t)?,
14868 e.matmul_decode_exact(up_shexp, z, t)?,
14869 ),
14870 }
14871 } else {
14872 (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?)
14873 };
14874 let mut sa = e.uninit(t * n_ff_sh)?;
14875 e.silu_mul(&sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
14876 let sh = if verify_t {
14877 e.matmul_decode_exact(down_shexp, &sa, t)?
14878 } else {
14879 e.matmul(down_shexp, &sa, t)?
14880 };
14881 let g = match &m.gate_inp_shexp {
14886 Some(gate_inp_shexp) if crate::router_prefill_exact_on() => {
14887 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
14888 }
14889 Some(gate_inp_shexp) => {
14890 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
14891 let mut g = e.uninit(t)?;
14892 e.sigmoid(&gs, &mut g, t)?;
14893 g
14894 }
14895 None => e.htod(&vec![1.0f32; t])?,
14896 };
14897 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
14898 }
14899 Ok(moe_out)
14900 }
14901
14902 #[allow(clippy::too_many_arguments)]
14904 #[allow(clippy::too_many_arguments)]
14905 fn moe_ffn_dev(
14906 e: &Engine,
14907 m: &MoeWeights,
14908 z: &CudaSlice<f32>,
14909 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
14910 logits: &CudaSlice<f32>,
14911 t: usize,
14912 cfg: &ModelConfig,
14913 il: u16,
14914 max_block: usize,
14915 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
14916 crate::moe_rp_refuse(m.dev_exps.as_ref().is_some_and(|d| d.rp), "moe_ffn_dev")?; let moe = cfg.moe.as_ref().unwrap();
14918 let n_embd = cfg.n_embd as usize;
14919 let n_expert = moe.expert_count as usize;
14920 let n_used = moe.expert_used_count as usize;
14921 let n_ff_exp = moe.expert_ff_length as usize;
14922 debug_assert!(
14926 cfg.sigmoid_router().is_none(),
14927 "moe_ffn_dev routes SOFTMAX: a sigmoid-router arch would pick wrong experts"
14928 );
14929 debug_assert!(
14930 !cfg.swiglu_clamped_at(il as u32),
14931 "moe_ffn_dev's fused epilogue is plain SiLU: no clamped form"
14932 );
14933
14934 let (sel_d, mut w_d) = e.moe_router_topk(logits, t, n_expert, n_used)?;
14936 crate::moe_sel_dump::refuse_device_only(
14941 "the softmax device-routed decode arm (moe_ffn_dev)",
14942 )?;
14943 if m.has_macros {
14946 e.moe_w_scale_by_expert(&mut w_d, &sel_d, &m.dev_macros, n_expert, t * n_used)?;
14947 }
14948
14949 let mut moe_out = e.uninit(t * n_embd)?;
14951
14952 if let Some(dev) = m.dev_exps.as_ref() {
14955 let (rbg_d, rbu_d) = if dev.gu_il {
14958 let sxx = m.gate_exps.row_bytes + m.up_exps.row_bytes;
14959 (sxx, sxx)
14960 } else {
14961 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
14962 };
14963 let q8 = moe_q8_enabled_for_model(cfg, m);
14964 let rows_arm = q8
14973 && t > 1
14974 && crate::spec::spec_m2()
14975 && n_ff_exp == 512
14976 && n_used <= 8
14977 && std::env::var("MEMRA_MOE_DEVQ8_GU")
14978 .map(|v| v.is_empty() || v == "v")
14979 .unwrap_or(true)
14980 && std::env::var("MEMRA_MOE_DEVQ8_DOWN")
14981 .map(|v| v.is_empty() || v == "w8h2v")
14982 .unwrap_or(true);
14983 let csr_mode = std::env::var("MEMRA_MOE_CSR")
14992 .ok()
14993 .and_then(|v| v.parse::<i32>().ok())
14994 .unwrap_or(1);
14995 let csr_nvfp4_probe = std::env::var("MEMRA_MOE_CSR_NVFP4").as_deref() == Ok("1");
15012 let csr_qt = |qt: i32| {
15013 qt == crate::QT_IQ4_XS
15014 || qt == crate::QT_IQ3_S
15015 || (csr_nvfp4_probe && qt == crate::QT_NVFP4)
15016 };
15017 let csr_t_max = if csr_nvfp4_probe { MOE_DEV_MAX_T } else { 10 };
15018 let csr_uniform = m.gate_exps.qtype == m.up_exps.qtype;
15019 let csr_arm = rows_arm
15020 && csr_mode > 0
15021 && t <= csr_t_max
15022 && csr_uniform
15023 && csr_qt(m.gate_exps.qtype)
15024 && csr_qt(m.up_exps.qtype)
15025 && csr_qt(m.down_exps.qtype);
15026 if csr_arm {
15027 if csr_mode == 2 {
15028 static ENGAGED: std::sync::Once = std::sync::Once::new();
15029 ENGAGED.call_once(|| eprintln!("[memra] moe CSR byte-compare mode ON (t={t})"));
15030 }
15031 let n_pairs = t * n_used;
15032 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
15033 let act = e.moe_gate_up_silu8_dev_q8_csr(
15034 &dev.ptr_row,
15035 &sel_d,
15036 &zq,
15037 &zd,
15038 n_pairs,
15039 n_embd,
15040 n_ff_exp,
15041 n_used,
15042 n_expert,
15043 m.gate_exps.qtype,
15044 m.up_exps.qtype,
15045 rbg_d,
15046 rbu_d,
15047 )?;
15048 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
15049 e.moe_down8_fma_dev_q8_rows(
15053 &dev.ptr_row,
15054 &sel_d,
15055 &w_d,
15056 &aq2,
15057 &ad2,
15058 &mut moe_out,
15059 t,
15060 n_ff_exp,
15061 n_embd,
15062 n_used,
15063 n_expert,
15064 m.down_exps.qtype,
15065 m.down_exps.row_bytes,
15066 )?;
15067 if csr_mode == 2 {
15068 let act_r = e.moe_gate_up_silu8_dev_q8_rows(
15070 &dev.ptr_row,
15071 &sel_d,
15072 &zq,
15073 &zd,
15074 t,
15075 n_embd,
15076 n_ff_exp,
15077 n_used,
15078 n_expert,
15079 m.gate_exps.qtype,
15080 m.up_exps.qtype,
15081 rbg_d,
15082 rbu_d,
15083 &m.dev_macros,
15084 )?;
15085 let mut out_r = e.uninit(t * n_embd)?;
15086 let (aq2r, ad2r) = e.quantize_q8_1(&act_r, n_pairs, n_ff_exp)?;
15087 e.moe_down8_fma_dev_q8_rows(
15088 &dev.ptr_row,
15089 &sel_d,
15090 &w_d,
15091 &aq2r,
15092 &ad2r,
15093 &mut out_r,
15094 t,
15095 n_ff_exp,
15096 n_embd,
15097 n_used,
15098 n_expert,
15099 m.down_exps.qtype,
15100 m.down_exps.row_bytes,
15101 )?;
15102 let (a1, a2) = (e.dtoh(&act)?, e.dtoh(&act_r)?);
15103 let (o1, o2) = (e.dtoh(&moe_out)?, e.dtoh(&out_r)?);
15104 let ba = a1
15105 .iter()
15106 .zip(&a2)
15107 .filter(|(x, y)| x.to_bits() != y.to_bits())
15108 .count();
15109 let bo = o1
15110 .iter()
15111 .zip(&o2)
15112 .filter(|(x, y)| x.to_bits() != y.to_bits())
15113 .count();
15114 if ba + bo > 0 {
15115 eprintln!(
15116 "[csr-check] il={il} t={t} ACT diffs={ba}/{} OUT diffs={bo}/{}",
15117 a1.len(),
15118 o1.len()
15119 );
15120 let sel_h = e.dtoh_i32(&sel_d)?;
15122 let mut shown = 0;
15123 for (i, (x, y)) in a1.iter().zip(&a2).enumerate() {
15124 if x.to_bits() != y.to_bits() && shown < 4 {
15125 let (p, o) = (i / n_ff_exp, i % n_ff_exp);
15126 let ex = sel_h[p];
15127 let npx = sel_h.iter().filter(|&&v| v == ex).count();
15128 eprintln!(
15129 " ACT p={p} ex={ex} np={npx} o={o} csr={x:e} rows={y:e}"
15130 );
15131 shown += 1;
15132 }
15133 }
15134 std::process::exit(3);
15135 }
15136 }
15137 } else if rows_arm {
15138 if std::env::var("MEMRA_MOE_OVERLAP").as_deref() == Ok("1") {
15141 use std::sync::atomic::{AtomicU64, Ordering};
15142 static PAIRS: AtomicU64 = AtomicU64::new(0);
15143 static UNIQ: AtomicU64 = AtomicU64::new(0);
15144 static CALLS: AtomicU64 = AtomicU64::new(0);
15145 let sel_h = e.dtoh_i32(&sel_d)?;
15146 let mut u: Vec<i32> = sel_h.clone();
15147 u.sort_unstable();
15148 u.dedup();
15149 PAIRS.fetch_add(sel_h.len() as u64, Ordering::Relaxed);
15150 UNIQ.fetch_add(u.len() as u64, Ordering::Relaxed);
15151 let c = CALLS.fetch_add(1, Ordering::Relaxed) + 1;
15152 if c.is_multiple_of(480) {
15153 let p = PAIRS.load(Ordering::Relaxed);
15154 let q = UNIQ.load(Ordering::Relaxed);
15155 eprintln!(
15156 "[overlap] calls={c} pairs={p} unique={q} ratio={:.3} (t={t})",
15157 q as f64 / p as f64
15158 );
15159 }
15160 }
15161 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
15162 let act = e.moe_gate_up_silu8_dev_q8_rows(
15163 &dev.ptr_row,
15164 &sel_d,
15165 &zq,
15166 &zd,
15167 t,
15168 n_embd,
15169 n_ff_exp,
15170 n_used,
15171 n_expert,
15172 m.gate_exps.qtype,
15173 m.up_exps.qtype,
15174 rbg_d,
15175 rbu_d,
15176 &m.dev_macros,
15177 )?;
15178 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
15179 e.moe_down8_fma_dev_q8_rows(
15180 &dev.ptr_row,
15181 &sel_d,
15182 &w_d,
15183 &aq2,
15184 &ad2,
15185 &mut moe_out,
15186 t,
15187 n_ff_exp,
15188 n_embd,
15189 n_used,
15190 n_expert,
15191 m.down_exps.qtype,
15192 m.down_exps.row_bytes,
15193 )?;
15194 } else {
15195 for tok in 0..t {
15196 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
15197 let selt = sel_d.slice(tok * n_used..(tok + 1) * n_used);
15198 let wt = w_d.slice(tok * n_used..(tok + 1) * n_used);
15199 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
15200 if q8 {
15201 let (zq, zd) = match (t, zq8) {
15202 (1, Some((q, d))) => (q.clone(), d.clone()),
15203 _ => e.quantize_q8_1_view(&zt, 1, n_embd)?,
15204 };
15205 let act = e.moe_gate_up_silu8_dev_q8(
15206 &dev.ptr_row,
15207 &selt,
15208 &zq,
15209 &zd,
15210 n_embd,
15211 n_ff_exp,
15212 n_used,
15213 n_expert,
15214 m.gate_exps.qtype,
15215 m.up_exps.qtype,
15216 rbg_d,
15217 rbu_d,
15218 &m.dev_macros,
15219 )?;
15220 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
15221 e.moe_down8_fma_dev_q8(
15222 &dev.ptr_row,
15223 &selt,
15224 &wt,
15225 &aq2,
15226 &ad2,
15227 &mut dst,
15228 n_ff_exp,
15229 n_embd,
15230 n_used,
15231 n_expert,
15232 m.down_exps.qtype,
15233 m.down_exps.row_bytes,
15234 )?;
15235 } else {
15236 let act = e.moe_gate_up_silu8_dev(
15237 &dev.ptr_row,
15238 &selt,
15239 &zt,
15240 n_embd,
15241 n_ff_exp,
15242 n_used,
15243 n_expert,
15244 m.gate_exps.qtype,
15245 m.up_exps.qtype,
15246 rbg_d,
15247 rbu_d,
15248 &m.dev_macros,
15249 )?;
15250 e.moe_down8_fma_dev(
15251 &dev.ptr_row,
15252 &selt,
15253 &wt,
15254 &act,
15255 &mut dst,
15256 n_ff_exp,
15257 n_embd,
15258 n_used,
15259 n_expert,
15260 m.down_exps.qtype,
15261 m.down_exps.row_bytes,
15262 )?;
15263 }
15264 }
15265 }
15266 } else {
15267 let q8 = moe_q8_enabled_for_model(cfg, m);
15274 e.with_moe_cache(max_block, |c, eng| {
15275 let row = c
15276 .layer_dev_row(il, n_expert, eng)?
15277 .ok_or("moe_ffn_dev: layer row vanished under the lock")?;
15278 for tok in 0..t {
15279 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
15280 let selt = sel_d.slice(tok * n_used..(tok + 1) * n_used);
15281 let wt = w_d.slice(tok * n_used..(tok + 1) * n_used);
15282 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
15283 if q8 {
15284 let (zq, zd) = match (t, zq8) {
15285 (1, Some((q, d))) => (q.clone(), d.clone()),
15286 _ => eng.quantize_q8_1_view(&zt, 1, n_embd)?,
15287 };
15288 let act = eng.moe_gate_up_silu8_dev_q8(
15289 row,
15290 &selt,
15291 &zq,
15292 &zd,
15293 n_embd,
15294 n_ff_exp,
15295 n_used,
15296 n_expert,
15297 m.gate_exps.qtype,
15298 m.up_exps.qtype,
15299 m.gate_exps.row_bytes,
15300 m.up_exps.row_bytes,
15301 &m.dev_macros,
15302 )?;
15303 let (aq2, ad2) = eng.quantize_q8_1(&act, n_used, n_ff_exp)?;
15304 eng.moe_down8_fma_dev_q8(
15305 row,
15306 &selt,
15307 &wt,
15308 &aq2,
15309 &ad2,
15310 &mut dst,
15311 n_ff_exp,
15312 n_embd,
15313 n_used,
15314 n_expert,
15315 m.down_exps.qtype,
15316 m.down_exps.row_bytes,
15317 )?;
15318 } else {
15319 let act = eng.moe_gate_up_silu8_dev(
15320 row,
15321 &selt,
15322 &zt,
15323 n_embd,
15324 n_ff_exp,
15325 n_used,
15326 n_expert,
15327 m.gate_exps.qtype,
15328 m.up_exps.qtype,
15329 m.gate_exps.row_bytes,
15330 m.up_exps.row_bytes,
15331 &m.dev_macros,
15332 )?;
15333 eng.moe_down8_fma_dev(
15334 row,
15335 &selt,
15336 &wt,
15337 &act,
15338 &mut dst,
15339 n_ff_exp,
15340 n_embd,
15341 n_used,
15342 n_expert,
15343 m.down_exps.qtype,
15344 m.down_exps.row_bytes,
15345 )?;
15346 }
15347 }
15348 c.hits += (t * 3 * n_used) as u64;
15350 Ok(())
15351 })?;
15352 }
15353
15354 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
15359 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
15360 {
15361 let n_ff_sh = gate_shexp.out_features();
15362 let verify_t = t > 1 && t < PRIME_MIN_T;
15365 let (sg_gate, sg_up) = if t == 1 {
15366 shexp_gate_up_t1(e, gate_shexp, up_shexp, z, zq8)?
15367 } else if verify_t {
15368 let mut fused = None;
15372 if crate::spec::spec_fused_t()
15373 && (2..=4).contains(&t)
15374 && e.uses_q8_1_fast(gate_shexp)
15375 && e.uses_q8_1_fast(up_shexp)
15376 {
15377 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
15378 fused = e.matmul_q8_fused2_t(gate_shexp, up_shexp, &zq, &zd, t)?;
15379 }
15380 match fused {
15381 Some(pair) => pair,
15382 None => (
15383 e.matmul_decode_exact(gate_shexp, z, t)?,
15384 e.matmul_decode_exact(up_shexp, z, t)?,
15385 ),
15386 }
15387 } else {
15388 (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?)
15389 };
15390 let mut sa = e.uninit(t * n_ff_sh)?; e.silu_mul(&sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
15392 let sh = if verify_t {
15393 e.matmul_decode_exact(down_shexp, &sa, t)?
15394 } else {
15395 e.matmul(down_shexp, &sa, t)?
15396 };
15397 let g = match &m.gate_inp_shexp {
15401 Some(gate_inp_shexp) => {
15402 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
15405 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
15406 } else {
15407 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
15408 let mut g = e.uninit(t)?;
15409 e.sigmoid(&gs, &mut g, t)?;
15410 g
15411 }
15412 }
15413 None => e.htod(&vec![1.0f32; t])?,
15414 };
15415 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
15416 }
15417
15418 Ok(moe_out)
15419 }
15420
15421 #[allow(clippy::too_many_arguments)]
15431 #[allow(clippy::too_many_arguments)]
15434 fn moe_gdec_token_q8(
15435 e: &Engine,
15436 m: &MoeWeights,
15437 il: u16,
15438 max_block: usize,
15439 zq: &CudaSlice<i8>,
15440 zd: &CudaSlice<f32>,
15441 sel: &[u32],
15442 w: &[f32],
15443 moe_out: &mut CudaSlice<f32>,
15444 tok: usize,
15445 n_embd: usize,
15446 n_ff_exp: usize,
15447 n_used: usize,
15448 ) -> Result<bool, Box<dyn std::error::Error>> {
15449 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
15450 use cudarc::driver::DevicePtr;
15451 let ptrs = e.with_moe_cache(max_block, |c, eng| {
15452 let mut g = [0u64; 8];
15453 let mut u = [0u64; 8];
15454 let mut d = [0u64; 8];
15455 for (j, &ex) in sel.iter().enumerate() {
15456 let ex = ex as u16;
15457 let (Some(sg), Some(su), Some(sd)) = (
15458 c.resident(BlockId::new(il, PROJ_GATE, ex)),
15459 c.resident(BlockId::new(il, PROJ_UP, ex)),
15460 c.resident(BlockId::new(il, PROJ_DOWN, ex)),
15461 ) else {
15462 return Ok(None);
15463 };
15464 let __s = eng.stream();
15465 let (pg, _e0) = c.slot(sg).device_ptr(&__s);
15466 let (pu, _e1) = c.slot(su).device_ptr(&__s);
15467 let (pd, _e2) = c.slot(sd).device_ptr(&__s);
15468 g[j] = pg;
15469 u[j] = pu;
15470 d[j] = pd;
15471 }
15472 if cpu_expert_profile_admit_enabled() && !c.is_frozen() {
15473 for &ex in sel {
15474 let ex = ex as u16;
15475 for proj in [PROJ_GATE, PROJ_UP, PROJ_DOWN] {
15476 c.note_profile_hit(BlockId::new(il, proj, ex));
15477 }
15478 }
15479 }
15480 c.hits += (3 * n_used) as u64;
15481 Ok(Some((g, u, d)))
15482 })?;
15483 let Some((g, u, d)) = ptrs else {
15484 return Ok(false);
15485 };
15486 let mut wv = [0f32; 8];
15487 wv[..n_used].copy_from_slice(w);
15488 let act = e.moe_gate_up_silu8_q8(
15489 crate::WPtr8(g),
15490 crate::WPtr8(u),
15491 zq,
15492 zd,
15493 n_embd,
15494 n_ff_exp,
15495 n_used,
15496 m.gate_exps.qtype,
15497 m.up_exps.qtype,
15498 m.gate_exps.row_bytes,
15499 m.up_exps.row_bytes,
15500 )?;
15501 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
15503 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
15504 e.moe_down8_fma_q8(
15505 crate::WPtr8(d),
15506 crate::F32x8(wv),
15507 &aq2,
15508 &ad2,
15509 &mut dst,
15510 n_ff_exp,
15511 n_embd,
15512 n_used,
15513 m.down_exps.qtype,
15514 m.down_exps.row_bytes,
15515 )?;
15516 Ok(true)
15517 }
15518
15519 #[allow(clippy::too_many_arguments)]
15549 fn moe_fused_epi_token_q8(
15550 e: &Engine,
15551 m: &MoeWeights,
15552 il: u16,
15553 max_block: usize,
15554 zq: &CudaSlice<i8>,
15555 zd: &CudaSlice<f32>,
15556 sel: &[u32],
15557 w: &[f32],
15558 moe_out: &mut CudaSlice<f32>,
15559 tok: usize,
15560 n_embd: usize,
15561 n_ff_exp: usize,
15562 n_used: usize,
15563 limit: f32,
15564 ) -> Result<bool, Box<dyn std::error::Error>> {
15565 use crate::moe_cache::{BlockId, DispatchSlot, PROJ_DOWN, PROJ_GATE, PROJ_UP};
15566 use cudarc::driver::DevicePtr;
15567 debug_assert!(
15568 limit > 1e-6,
15569 "the fused epilogue's kernel collapses every gate to silu(0) at limit 0"
15570 );
15571 debug_assert_eq!(sel.len(), n_used);
15572 debug_assert_eq!(w.len(), n_used);
15573
15574 let ptrs = e.with_moe_cache(max_block, |c, eng| {
15575 if c.n_slots() < 3 * n_used {
15578 return Ok(None);
15579 }
15580 for &ex in sel.iter() {
15583 let ex_usize = ex as usize;
15584 for (proj, exps) in [
15585 (PROJ_GATE, &m.gate_exps),
15586 (PROJ_UP, &m.up_exps),
15587 (PROJ_DOWN, &m.down_exps),
15588 ] {
15589 let id = BlockId::new(il, proj, ex as u16);
15590 let DispatchSlot::Resident(_) =
15591 c.dispatch_source(id, exps.expert_source(ex_usize), eng)?;
15592 }
15593 }
15594 let mut g = [0u64; 8];
15596 let mut u = [0u64; 8];
15597 let mut d = [0u64; 8];
15598 for (j, &ex) in sel.iter().enumerate() {
15599 let ex = ex as u16;
15600 let (Some(sg), Some(su), Some(sd)) = (
15601 c.resident(BlockId::new(il, PROJ_GATE, ex)),
15602 c.resident(BlockId::new(il, PROJ_UP, ex)),
15603 c.resident(BlockId::new(il, PROJ_DOWN, ex)),
15604 ) else {
15605 return Ok(None);
15606 };
15607 let __s = eng.stream();
15608 let (pg, _e0) = c.slot(sg).device_ptr(&__s);
15609 let (pu, _e1) = c.slot(su).device_ptr(&__s);
15610 let (pd, _e2) = c.slot(sd).device_ptr(&__s);
15611 g[j] = pg;
15612 u[j] = pu;
15613 d[j] = pd;
15614 }
15615 Ok(Some((g, u, d)))
15616 })?;
15617 let Some((g, u, d)) = ptrs else {
15618 return Ok(false);
15619 };
15620
15621 Self::moe_fused_epi_launch(
15622 e, m, zq, zd, sel, w, g, u, d, moe_out, tok, n_embd, n_ff_exp, n_used, limit,
15623 false, )?;
15625 Ok(true)
15626 }
15627
15628 #[allow(clippy::too_many_arguments)]
15639 fn moe_fused_epi_launch(
15640 e: &Engine,
15641 m: &MoeWeights,
15642 zq: &CudaSlice<i8>,
15643 zd: &CudaSlice<f32>,
15644 sel: &[u32],
15645 w: &[f32],
15646 g: [u64; 8],
15647 u: [u64; 8],
15648 d: [u64; 8],
15649 moe_out: &mut CudaSlice<f32>,
15650 tok: usize,
15651 n_embd: usize,
15652 n_ff_exp: usize,
15653 n_used: usize,
15654 limit: f32,
15655 rp: bool,
15656 ) -> Result<(), Box<dyn std::error::Error>> {
15657 let mut gs = [0f32; 8];
15658 let mut us = [0f32; 8];
15659 let mut wv = [0f32; 8];
15660 let (qt_g, qt_u, qt_d) = (
15662 crate::rp_qt(rp, m.gate_exps.qtype),
15663 crate::rp_qt(rp, m.up_exps.qtype),
15664 crate::rp_qt(rp, m.down_exps.qtype),
15665 );
15666 for (j, &ex) in sel.iter().enumerate() {
15667 let ex = ex as usize;
15668 gs[j] = m.gate_exps.macro_scale(ex);
15669 us[j] = m.up_exps.macro_scale(ex);
15670 wv[j] = w[j] * m.down_exps.macro_scale(ex);
15671 }
15672 let use_w4 = crate::b200_matvec_arm_on();
15694 let act = if use_w4 {
15695 e.moe_gate_up_preclamp8_q8_w4(
15696 crate::WPtr8(g),
15697 crate::WPtr8(u),
15698 zq,
15699 zd,
15700 crate::F32x8(gs),
15701 crate::F32x8(us),
15702 limit,
15703 n_embd,
15704 n_ff_exp,
15705 n_used,
15706 qt_g,
15707 qt_u,
15708 m.gate_exps.row_bytes,
15709 m.up_exps.row_bytes,
15710 )?
15711 } else {
15712 e.moe_gate_up_preclamp8_q8(
15713 crate::WPtr8(g),
15714 crate::WPtr8(u),
15715 zq,
15716 zd,
15717 crate::F32x8(gs),
15718 crate::F32x8(us),
15719 limit,
15720 n_embd,
15721 n_ff_exp,
15722 n_used,
15723 qt_g,
15724 qt_u,
15725 m.gate_exps.row_bytes,
15726 m.up_exps.row_bytes,
15727 )?
15728 };
15729 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
15731 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
15732 if use_w4 {
15733 e.moe_down8_fma_q8_w4(
15734 crate::WPtr8(d),
15735 crate::F32x8(wv),
15736 &aq2,
15737 &ad2,
15738 &mut dst,
15739 n_ff_exp,
15740 n_embd,
15741 n_used,
15742 qt_d,
15743 m.down_exps.row_bytes,
15744 )?;
15745 } else {
15746 e.moe_down8_fma_q8(
15747 crate::WPtr8(d),
15748 crate::F32x8(wv),
15749 &aq2,
15750 &ad2,
15751 &mut dst,
15752 n_ff_exp,
15753 n_embd,
15754 n_used,
15755 qt_d,
15756 m.down_exps.row_bytes,
15757 )?;
15758 }
15759 crate::MOE_FUSED_EPI_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
15760 Ok(())
15761 }
15762
15763 #[allow(clippy::too_many_arguments)]
15770 fn moe_vrows_pairs_q8(
15772 e: &Engine,
15773 m: &MoeWeights,
15774 z: &CudaSlice<f32>,
15775 sel: VrowsSel<'_>,
15776 il: u16,
15777 (pg, pu, pd): (u64, u64, u64),
15778 rp: bool,
15779 t: usize,
15780 n_embd: usize,
15781 n_ff_exp: usize,
15782 n_used: usize,
15783 limit: f32,
15784 moe_out: &mut CudaSlice<f32>,
15785 ) -> Result<(), Box<dyn std::error::Error>> {
15786 let n_pairs = t * n_used;
15787 let order_on = crate::moe_vrows_dedup_order_on() && !crate::moe_vrows_pack_on();
15795 let n_planes = if order_on { 4 } else { 3 };
15796 let mut ptrs_d = e.vws_uninit_u64(n_planes * n_pairs)?;
15802 let mut scl_d = e.vws_uninit(3 * n_pairs)?;
15803 match sel {
15809 VrowsSel::Host(sel_all, w_all) => {
15810 debug_assert_eq!(sel_all.len(), n_pairs);
15811 debug_assert_eq!(w_all.len(), n_pairs);
15812 let mut ptrs = vec![0u64; n_planes * n_pairs];
15813 let mut scl = vec![0f32; 3 * n_pairs];
15814 for (p, (&ex, &w)) in sel_all.iter().zip(w_all).enumerate() {
15815 let ex = ex as usize;
15816 ptrs[p] = pg + (ex * m.gate_exps.expert_stride) as u64;
15817 ptrs[n_pairs + p] = pu + (ex * m.up_exps.expert_stride) as u64;
15818 ptrs[2 * n_pairs + p] = pd + (ex * m.down_exps.expert_stride) as u64;
15819 scl[p] = m.gate_exps.macro_scale(ex);
15820 scl[n_pairs + p] = m.up_exps.macro_scale(ex);
15821 scl[2 * n_pairs + p] = w * m.down_exps.macro_scale(ex);
15824 }
15825 if order_on {
15826 ptrs[3 * n_pairs..].copy_from_slice(&crate::vrows_expert_major_order(sel_all));
15828 let (visits, distinct) = crate::vrows_overlap_counts(sel_all);
15831 crate::MOE_VROWS_SLAB_READS_AVOIDED
15832 .fetch_add(visits - distinct, std::sync::atomic::Ordering::Relaxed);
15833 }
15834 e.htod_u64_into(&ptrs, &mut ptrs_d)?;
15835 e.htod_f32_into(&scl, &mut scl_d)?;
15836 if crate::moe_vrows_dedup_stat_on() {
15841 let (visits, distinct) = crate::vrows_overlap_counts(sel_all);
15842 debug_assert_eq!(visits, n_pairs as u64);
15843 crate::MOE_VROWS_PAIR_VISITS
15844 .fetch_add(visits, std::sync::atomic::Ordering::Relaxed);
15845 crate::MOE_VROWS_PAIR_DISTINCT
15846 .fetch_add(distinct, std::sync::atomic::Ordering::Relaxed);
15847 crate::moe_vrows_dedup_report();
15848 }
15849 }
15850 VrowsSel::Dev(sel_d, selw_d) => {
15851 let macros = match (
15852 m.gate_exps.macros.as_deref(),
15853 m.up_exps.macros.as_deref(),
15854 m.down_exps.macros.as_deref(),
15855 ) {
15856 (Some(g), Some(u), Some(d)) => Some((g, u, d)),
15857 (None, None, None) => None,
15861 _ => {
15862 return Err("vrows device tables: expert macro planes are not uniform \
15863 across gate/up/down"
15864 .into());
15865 }
15866 };
15867 e.moe_vrows_tables_from_sel(
15868 sel_d,
15869 selw_d,
15870 il,
15871 macros,
15872 (pg, pu, pd),
15873 (
15874 m.gate_exps.expert_stride,
15875 m.up_exps.expert_stride,
15876 m.down_exps.expert_stride,
15877 ),
15878 n_pairs,
15879 &mut ptrs_d,
15880 &mut scl_d,
15881 )?;
15882 if order_on {
15883 e.moe_vrows_order_from_sel(sel_d, n_pairs, &mut ptrs_d)?;
15887 }
15888 if crate::MOE_VROWS_DEV_TABLES_DISPATCHES
15889 .fetch_add(1, std::sync::atomic::Ordering::Relaxed)
15890 == 0
15891 {
15892 eprintln!(
15893 "[moe-vrows-dev-tables] engaged: pointer/scale tables built on device \
15894 from the router's own sel/w; the per-layer pinned readback and its \
15895 cuStreamSynchronize are skipped (MEMRA_MOE_VROWS_DEV_TABLES=1)"
15896 );
15897 }
15898 }
15899 }
15900 let (mut zq, mut zd) = (
15903 e.vws_uninit_i8(t * n_embd)?,
15904 e.vws_uninit(t * (n_embd / 32))?,
15905 );
15906 e.quantize_q8_1_into(z, t, n_embd, &mut zq, &mut zd)?;
15907 let vrows_arm = match sel {
15912 VrowsSel::Dev(..) => "device",
15913 VrowsSel::Host(..) => "vrows-host",
15914 };
15915 Self::trace_moe_act(e, vrows_arm, il, t, z, &zq, &zd);
15916 let act = e.moe_gate_up_preclamp8_q8_rows(
15917 &ptrs_d,
15918 &scl_d,
15919 &zq,
15920 &zd,
15921 limit,
15922 n_embd,
15923 n_ff_exp,
15924 n_used,
15925 n_pairs,
15926 crate::rp_qt(rp, m.gate_exps.qtype),
15927 crate::rp_qt(rp, m.up_exps.qtype),
15928 m.gate_exps.row_bytes,
15929 m.up_exps.row_bytes,
15930 )?;
15931 let (mut aq2, mut ad2) = (
15933 e.vws_uninit_i8(n_pairs * n_ff_exp)?,
15934 e.vws_uninit(n_pairs * (n_ff_exp / 32))?,
15935 );
15936 e.quantize_q8_1_into(&act, n_pairs, n_ff_exp, &mut aq2, &mut ad2)?;
15937 e.moe_down8_fma_q8_rows(
15938 &ptrs_d,
15939 &scl_d,
15940 &aq2,
15941 &ad2,
15942 moe_out,
15943 n_ff_exp,
15944 n_embd,
15945 n_used,
15946 n_pairs,
15947 crate::rp_qt(rp, m.down_exps.qtype),
15948 m.down_exps.row_bytes,
15949 )?;
15950 Self::trace_moe_out(e, vrows_arm, il, moe_out);
15951 e.vws_recycle_u64(ptrs_d);
15954 e.vws_recycle(scl_d);
15955 e.vws_recycle_i8(zq);
15956 e.vws_recycle(zd);
15957 e.vws_recycle(act);
15958 e.vws_recycle_i8(aq2);
15959 e.vws_recycle(ad2);
15960 if crate::MOE_VROWS_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed) == 0 {
15961 eprintln!(
15962 "[glm5-vrows] verify MoE batched across rows: pairs={n_pairs} (t={t} x \
15963 {n_used}), one gate/up+preclamp launch + one down/FMA launch per layer-call \
15964 (rides MEMRA_GLM5_VERIFY_BATCH); arm doors: pack={} dedup_order={} \
15965 b200_matvec={} dev_tables={} (env MEMRA_MOE_VROWS_DEV_TABLES={}) — the `_rows` \
15966 pair has several dispatch twins and only the plain one has ever run at t=1, so a \
15967 box log has to say which it took. `dev_tables` is THIS CALL's provenance, not the \
15968 env: the decode-graph door owns both halves and routes the tables through the \
15969 device build whatever `MEMRA_MOE_VROWS_DEV_TABLES` says (run 7B printed the env \
15970 and read false while the door was building on device, which cost a box window)",
15971 crate::moe_vrows_pack_on(),
15972 crate::moe_vrows_dedup_order_on(),
15973 crate::b200_matvec_arm_on(),
15974 matches!(sel, VrowsSel::Dev(..)),
15975 crate::moe_vrows_dev_tables_on(),
15976 );
15977 }
15978 Ok(())
15979 }
15980
15981 fn moe_ffn_grouped_prefill_sigmoid(
16019 e: &Engine,
16020 m: &MoeWeights,
16021 z: &CudaSlice<f32>,
16022 t: usize,
16023 cfg: &ModelConfig,
16024 il: u16,
16025 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
16026 fn decline_once(reason: &str, t: usize, il: u16) {
16036 static DECLINED: std::sync::atomic::AtomicBool =
16037 std::sync::atomic::AtomicBool::new(false);
16038 if !DECLINED.swap(true, std::sync::atomic::Ordering::Relaxed) {
16039 eprintln!(
16040 "[moe-grouped-prefill] DECLINED t={t} il={il}: {reason} -> the sequential \
16041 per-token dispatch serves this prime (logged once per process)"
16042 );
16043 }
16044 }
16045 let Some(dev) = m
16046 .dev_exps
16047 .as_ref()
16048 .filter(|d| moe_slab_enabled() && d.dev == e.ctx().ordinal())
16049 else {
16050 decline_once(
16051 "no LOCAL resident expert slab (dev_exps is None: the resident-experts decision \
16052 did not select RESIDENT on this device, e.g. MEMRA_ST_PINNED=1 / \
16053 MEMRA_MOE_RESIDENT=0 / a budget below the bank; or MEMRA_MOE_SLAB=0)",
16054 t,
16055 il,
16056 );
16057 return Ok(None);
16058 };
16059 if crate::moe_f16g_mode() == 0 {
16060 decline_once(
16061 "MEMRA_MOE_F16G=0 (the grouped f16 GEMM class is off)",
16062 t,
16063 il,
16064 );
16065 return Ok(None);
16066 }
16067 if crate::moe_f16g_mode() == 1 {
16089 decline_once(
16090 "MEMRA_MOE_F16G=1 is REFUSED on the glm5 sigmoid grouped prefill: cublasGemmGroupedBatchedEx issues on cuBLAS-internal streams unordered with ours, and on this walk's 283-group shape that race silently destroyed the trunk and killed the worker (2026-09-02 boot D, research/glm5-b200-20260902/ box/prefill/). Use MEMRA_MOE_F16G=2 (the default) or the MEMRA_B200_PRIME_V2 dequant-once arm",
16091 t,
16092 il,
16093 );
16094 return Ok(None);
16095 }
16096 if std::env::var("MEMRA_MOE_GATE").is_ok() {
16100 return Ok(None);
16101 }
16102 let moe = cfg
16103 .moe
16104 .as_ref()
16105 .ok_or("grouped sigmoid prefill requires MoE model metadata")?;
16106 let n_embd = cfg.n_embd as usize;
16107 let n_expert = moe.expert_count as usize;
16108 let n_used = moe.expert_used_count as usize;
16109 let n_ff_exp = moe.expert_ff_length as usize;
16110 if !(f16g_proj_ok(m.gate_exps.qtype, n_embd)
16111 && f16g_proj_ok(m.up_exps.qtype, n_embd)
16112 && f16g_proj_ok(m.down_exps.qtype, n_ff_exp))
16113 {
16114 decline_once(
16115 "an expert projection qtype/shape is outside the grouped f16 GEMM class",
16116 t,
16117 il,
16118 );
16119 return Ok(None);
16120 }
16121 if n_expert > 512 || n_used == 0 || n_used > 8 {
16124 return Ok(None);
16125 }
16126 let sigmoid = cfg
16127 .sigmoid_router()
16128 .ok_or("grouped sigmoid prefill requires the sigmoid router")?;
16129 let lim_exp = cfg.clamp_exp_at(il as u32);
16133 if matches!(lim_exp, Some(SwigluClamp::Post(_))) {
16134 return Err(
16135 "grouped sigmoid prefill is qualified for the PRE-clamped SwiGLU form only; \
16136 a POST-clamp layer must ride the sequential arm"
16137 .into(),
16138 );
16139 }
16140
16141 let logits = Self::moe_router_logits(e, m, z, t, cfg)?;
16146 Self::trace_sigmoid_router_logits(e, il, t, n_expert, n_used, &logits, m, sigmoid)?;
16147 let (sel_all, w_all) =
16148 Self::moe_route_sigmoid_cfg(e, &logits, t, n_expert, n_used, m, sigmoid)?;
16149 crate::moesd::record_host_routes(il, n_expert, n_used, &sel_all)?;
16150 Self::trace_moe_routes(il, t, &sel_all, &w_all)?;
16151 Self::trace_moe_input(e, il, t, n_embd, z)?;
16152
16153 let mprof = std::env::var("MEMRA_PRIME_PROF").as_deref() == Ok("1");
16154 let mut mt = std::time::Instant::now();
16155 let mut phase = |on: bool| -> f64 {
16156 if on {
16157 let _ = e.stream().synchronize();
16158 let v = mt.elapsed().as_secs_f64() * 1e3;
16159 mt = std::time::Instant::now();
16160 v
16161 } else {
16162 0.0
16163 }
16164 };
16165 let d_router = phase(mprof);
16166
16167 let n_pairs = t * n_used;
16169 if sel_all.len() < n_pairs || w_all.len() < n_pairs || z.len() < t * n_embd {
16170 return Err("grouped sigmoid prefill geometry".into());
16171 }
16172 let mut buckets: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
16173 for (p, &s_id) in sel_all.iter().take(n_pairs).enumerate() {
16174 let s_id = s_id as usize;
16175 if s_id >= n_expert {
16176 return Err(format!("grouped prefill selection {s_id} >= {n_expert}").into());
16177 }
16178 buckets[s_id].push(p as i32);
16179 }
16180 let mut ex_ids: Vec<i32> = Vec::new();
16181 let mut ex_off: Vec<i32> = vec![0];
16182 let mut ex_pairs: Vec<i32> = Vec::with_capacity(n_pairs);
16183 for (e_id, b) in buckets.iter().enumerate() {
16184 if !b.is_empty() {
16185 ex_ids.push(e_id as i32);
16186 ex_pairs.extend_from_slice(b);
16187 ex_off.push(ex_pairs.len() as i32);
16188 }
16189 }
16190 let n_active = ex_ids.len();
16191 if n_active == 0 || n_active > 512 {
16192 return Err(format!("grouped prefill n_active {n_active} outside 1..=512").into());
16193 }
16194 let csr_tok: Vec<i32> = ex_pairs.iter().map(|&p| p / n_used as i32).collect();
16195
16196 let wd: Vec<f32> = (0..n_pairs)
16200 .map(|p| w_all[p] * m.down_exps.macro_scale(sel_all[p] as usize))
16201 .collect();
16202
16203 let (rbg_d, rbu_d) = if dev.gu_il {
16205 let sxx = m.gate_exps.row_bytes + m.up_exps.row_bytes;
16206 (sxx, sxx)
16207 } else {
16208 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
16209 };
16210
16211 let exi = e.htod_i32(&ex_ids)?;
16212 let exo = e.htod_i32(&ex_off)?;
16213 let exp_d = e.htod_i32(&ex_pairs)?;
16214 let csr_tok_d = e.htod_i32(&csr_tok)?;
16215 let pw = e.htod(&wd)?;
16216
16217 let (z16, zs) = e.moe_f16g_act(z, Some(&csr_tok_d), n_embd, n_pairs)?;
16219 let mut g = e.moe_f16_grouped(
16220 &dev.ptr_row,
16221 0,
16222 n_expert,
16223 &exi,
16224 &ex_off,
16225 &exo,
16226 &z16,
16227 &zs,
16228 n_embd,
16229 n_ff_exp,
16230 n_active,
16231 n_pairs,
16232 crate::rp_qt(dev.rp, m.gate_exps.qtype), rbg_d,
16234 )?;
16235 if m.gate_exps.macros.is_some() {
16236 let mg: Vec<f32> = ex_pairs
16237 .iter()
16238 .map(|&p| m.gate_exps.macro_scale(sel_all[p as usize] as usize))
16239 .collect();
16240 let mg_d = e.htod(&mg)?;
16241 e.scale_rows(&mut g, &mg_d, n_ff_exp, n_pairs)?;
16242 }
16243 let mut u = e.moe_f16_grouped(
16244 &dev.ptr_row,
16245 1,
16246 n_expert,
16247 &exi,
16248 &ex_off,
16249 &exo,
16250 &z16,
16251 &zs,
16252 n_embd,
16253 n_ff_exp,
16254 n_active,
16255 n_pairs,
16256 crate::rp_qt(dev.rp, m.up_exps.qtype),
16257 rbu_d,
16258 )?;
16259 if m.up_exps.macros.is_some() {
16260 let mu: Vec<f32> = ex_pairs
16261 .iter()
16262 .map(|&p| m.up_exps.macro_scale(sel_all[p as usize] as usize))
16263 .collect();
16264 let mu_d = e.htod(&mu)?;
16265 e.scale_rows(&mut u, &mu_d, n_ff_exp, n_pairs)?;
16266 }
16267
16268 let act = match lim_exp {
16271 Some(SwigluClamp::Pre(limit)) => {
16272 let mut a = e.uninit(n_pairs * n_ff_exp)?;
16273 e.swiglu_preclamped_mul_scaled(
16276 &g,
16277 &u,
16278 1.0,
16279 1.0,
16280 limit,
16281 &mut a,
16282 n_pairs * n_ff_exp,
16283 )?;
16284 a
16285 }
16286 None => e.moe_pairs_silu_mul(&g, &u, n_pairs * n_ff_exp)?,
16287 Some(SwigluClamp::Post(_)) => unreachable!("refused before any launch"),
16288 };
16289 let d_gemm_gu = phase(mprof);
16290
16291 let (a16, a_s) = e.moe_f16g_act(&act, None, n_ff_exp, n_pairs)?;
16293 let d_csr = e.moe_f16_grouped(
16294 &dev.ptr_row,
16295 2,
16296 n_expert,
16297 &exi,
16298 &ex_off,
16299 &exo,
16300 &a16,
16301 &a_s,
16302 n_ff_exp,
16303 n_embd,
16304 n_active,
16305 n_pairs,
16306 crate::rp_qt(dev.rp, m.down_exps.qtype),
16307 m.down_exps.row_bytes,
16308 )?;
16309 let y_pair = e.rows_permute(&d_csr, &exp_d, n_pairs, n_embd)?;
16310 let toff: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
16311 let tids: Vec<i32> = (0..n_pairs as i32).collect();
16312 let toff_d = e.htod_i32(&toff)?;
16313 let tids_d = e.htod_i32(&tids)?;
16314 let mut moe_out = e.uninit(t * n_embd)?;
16317 e.moe_pairs_scatter(&y_pair, &pw, &toff_d, &tids_d, &mut moe_out, t, n_embd)?;
16318 let d_down = phase(mprof);
16319
16320 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut moe_out)?;
16323 if mprof {
16324 let d_shared = phase(true);
16325 eprintln!(
16326 "[moe-grouped-prefill-prof] il={il} t={t} router={d_router:.1}ms \
16327 gemm_gu={d_gemm_gu:.1}ms down_scatter={d_down:.1}ms shared={d_shared:.1}ms"
16328 );
16329 }
16330
16331 crate::MOE_GROUPED_PREFILL_DISPATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
16332 static GPF_LOGGED: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
16335 let layer_bit = 1u64 << (il as u64 % 64);
16336 if GPF_LOGGED.fetch_or(layer_bit, std::sync::atomic::Ordering::Relaxed) & layer_bit == 0 {
16337 eprintln!(
16338 "[moe-grouped-prefill] execute layer={il} tokens={t} n_active={n_active} \
16339 provenance=resident-slab router=sigmoid-host-oracle epilogue=pre-clamped \
16340 macro_fold=gate-up-rows+down-weight performance_claim=false \
16341 (logged once per layer)"
16342 );
16343 }
16344 Ok(Some(moe_out))
16345 }
16346
16347 #[allow(clippy::too_many_arguments)] fn moe_gdec_token(
16349 e: &Engine,
16350 m: &MoeWeights,
16351 il: u16,
16352 max_block: usize,
16353 zt: &cudarc::driver::CudaView<f32>,
16354 sel: &[u32],
16355 w: &[f32],
16356 moe_out: &mut CudaSlice<f32>,
16357 tok: usize,
16358 n_embd: usize,
16359 n_ff_exp: usize,
16360 n_used: usize,
16361 ) -> Result<bool, Box<dyn std::error::Error>> {
16362 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
16363 use cudarc::driver::DevicePtr;
16364 let ptrs = e.with_moe_cache(max_block, |c, eng| {
16366 let mut g = [0u64; 8];
16367 let mut u = [0u64; 8];
16368 let mut d = [0u64; 8];
16369 for (j, &ex) in sel.iter().enumerate() {
16370 let ex = ex as u16;
16371 let (Some(sg), Some(su), Some(sd)) = (
16372 c.resident(BlockId::new(il, PROJ_GATE, ex)),
16373 c.resident(BlockId::new(il, PROJ_UP, ex)),
16374 c.resident(BlockId::new(il, PROJ_DOWN, ex)),
16375 ) else {
16376 return Ok(None);
16377 };
16378 let __s = eng.stream();
16379 let (pg, _e0) = c.slot(sg).device_ptr(&__s);
16380 let (pu, _e1) = c.slot(su).device_ptr(&__s);
16381 let (pd, _e2) = c.slot(sd).device_ptr(&__s);
16382 g[j] = pg;
16383 u[j] = pu;
16384 d[j] = pd;
16385 }
16386 if cpu_expert_profile_admit_enabled() && !c.is_frozen() {
16387 for &ex in sel {
16388 let ex = ex as u16;
16389 for proj in [PROJ_GATE, PROJ_UP, PROJ_DOWN] {
16390 c.note_profile_hit(BlockId::new(il, proj, ex));
16391 }
16392 }
16393 }
16394 c.hits += (3 * n_used) as u64; Ok(Some((g, u, d)))
16396 })?;
16397 let Some((g, u, d)) = ptrs else {
16398 return Ok(false);
16399 };
16400 let mut wv = [0f32; 8];
16401 wv[..n_used].copy_from_slice(w);
16402 let act = e.moe_gate_up_silu8(
16404 crate::WPtr8(g),
16405 crate::WPtr8(u),
16406 zt,
16407 n_embd,
16408 n_ff_exp,
16409 n_used,
16410 m.gate_exps.qtype,
16411 m.up_exps.qtype,
16412 m.gate_exps.row_bytes,
16413 m.up_exps.row_bytes,
16414 )?;
16415 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
16416 e.moe_down8_fma_into(
16417 crate::WPtr8(d),
16418 crate::F32x8(wv),
16419 &act,
16420 &mut dst,
16421 n_ff_exp,
16422 n_embd,
16423 n_used,
16424 m.down_exps.qtype,
16425 m.down_exps.row_bytes,
16426 )?;
16427 Ok(true)
16428 }
16429
16430 #[allow(clippy::too_many_arguments)] fn moe_cached_gemm_q8(
16436 e: &Engine,
16437 il: u16,
16438 proj: u8,
16439 ex: usize,
16440 m: &MoeWeights,
16441 max_block: usize,
16442 aq: &CudaSlice<i8>,
16443 ad: &CudaSlice<f32>,
16444 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
16445 use crate::moe_cache::{BlockId, DispatchSlot, PROJ_GATE, PROJ_UP};
16446 let exps = match proj {
16447 PROJ_GATE => &m.gate_exps,
16448 PROJ_UP => &m.up_exps,
16449 _ => &m.down_exps,
16450 };
16451 let layout = exps.expert_layout(ex);
16452 let id = BlockId::new(il, proj, ex as u16);
16453 let source = exps.expert_source(ex);
16454 e.with_moe_cache(max_block, |c, eng| {
16455 let slot = c.dispatch_source(id, source, eng)?;
16456 let DispatchSlot::Resident(sl) = slot;
16457 let buf = c.slot(sl);
16458 eng.qmatvec_expert_q8(
16459 buf,
16460 0..layout.len,
16461 aq,
16462 ad,
16463 1,
16464 exps.in_f,
16465 exps.out_f,
16466 layout.qtype,
16467 layout.row_bytes,
16468 )
16469 })
16470 }
16471
16472 fn moe_cached_gemm(
16473 e: &Engine,
16474 il: u16,
16475 proj: u8,
16476 ex: usize,
16477 m: &MoeWeights,
16478 max_block: usize,
16479 x: &cudarc::driver::CudaView<f32>,
16480 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
16481 use crate::moe_cache::{BlockId, DispatchSlot, PROJ_GATE, PROJ_UP};
16482 let exps = match proj {
16483 PROJ_GATE => &m.gate_exps,
16484 PROJ_UP => &m.up_exps,
16485 _ => &m.down_exps,
16486 };
16487 let layout = exps.expert_layout(ex);
16488 let id = BlockId::new(il, proj, ex as u16);
16489 let source = exps.expert_source(ex);
16490 e.with_moe_cache(max_block, |c, eng| {
16492 let slot = c.dispatch_source(id, source, eng)?;
16493 let DispatchSlot::Resident(sl) = slot;
16496 let buf = c.slot(sl);
16497 m.qmatvec_view(
16498 eng,
16499 buf,
16500 0..layout.len,
16501 x,
16502 1,
16503 exps.in_f,
16504 exps.out_f,
16505 layout.qtype,
16506 layout.row_bytes,
16507 )
16508 })
16509 }
16510
16511 fn moe_profile_admit_expert(
16515 e: &Engine,
16516 il: u16,
16517 ex: usize,
16518 m: &MoeWeights,
16519 max_block: usize,
16520 ) -> Result<(), Box<dyn std::error::Error>> {
16521 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
16522 e.with_moe_cache(max_block, |cache, eng| {
16523 for (proj, exps) in [
16524 (PROJ_GATE, &m.gate_exps),
16525 (PROJ_UP, &m.up_exps),
16526 (PROJ_DOWN, &m.down_exps),
16527 ] {
16528 let id = BlockId::new(il, proj, ex as u16);
16529 let _ = cache.dispatch_source(id, exps.expert_source(ex), eng)?;
16530 }
16531 Ok(())
16532 })
16533 }
16534
16535 #[allow(clippy::too_many_arguments)]
16538 fn moe_frozen_gemm(
16539 e: &Engine,
16540 il: u16,
16541 proj: u8,
16542 ex: usize,
16543 m: &MoeWeights,
16544 max_block: usize,
16545 x: &cudarc::driver::CudaView<f32>,
16546 scratch: &mut Option<CudaSlice<u8>>,
16547 scratch_len: usize,
16548 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
16549 use crate::moe_cache::{BlockId, PROJ_GATE, PROJ_UP};
16550 let exps = match proj {
16551 PROJ_GATE => &m.gate_exps,
16552 PROJ_UP => &m.up_exps,
16553 _ => &m.down_exps,
16554 };
16555 let layout = exps.expert_layout(ex);
16556 let id = BlockId::new(il, proj, ex as u16);
16557 if let Some(output) = e.with_moe_cache(max_block, |cache, eng| {
16558 let Some(slot) = cache.resident(id) else {
16559 return Ok(None);
16560 };
16561 let buf = cache.slot(slot);
16562 Ok(Some(m.qmatvec_view(
16563 eng,
16564 buf,
16565 0..layout.len,
16566 x,
16567 1,
16568 exps.in_f,
16569 exps.out_f,
16570 layout.qtype,
16571 layout.row_bytes,
16572 )?))
16573 })? {
16574 return Ok(output);
16575 }
16576 if scratch.is_none() {
16577 *scratch = Some(e.alloc_u8_uninit(scratch_len)?);
16578 }
16579 let scratch = scratch.as_mut().unwrap();
16580 e.stage_expert(exps.expert_bytes(ex), scratch, 0)?;
16581 m.qmatvec_view(
16582 e,
16583 scratch,
16584 0..layout.len,
16585 x,
16586 1,
16587 exps.in_f,
16588 exps.out_f,
16589 layout.qtype,
16590 layout.row_bytes,
16591 )
16592 }
16593
16594 fn moe_prefetch_expert(
16595 e: &Engine,
16596 il: u16,
16597 ex: usize,
16598 m: &MoeWeights,
16599 max_block: usize,
16600 keep: &[crate::moe_cache::BlockId],
16601 ) -> Result<(), Box<dyn std::error::Error>> {
16602 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
16603 e.with_moe_cache(max_block, |c, eng| {
16604 for (proj, exps) in [
16605 (PROJ_GATE, &m.gate_exps),
16606 (PROJ_UP, &m.up_exps),
16607 (PROJ_DOWN, &m.down_exps),
16608 ] {
16609 let id = BlockId::new(il, proj, ex as u16);
16610 let _ = c.prefetch_source(id, exps.expert_source(ex), keep, eng)?;
16611 }
16612 Ok(())
16613 })
16614 }
16615
16616 fn moe_prefetch_disk_expert(
16619 e: &Engine,
16620 il: u16,
16621 ex: usize,
16622 m: &MoeWeights,
16623 max_block: usize,
16624 keep: &[crate::moe_cache::BlockId],
16625 ) -> Result<(), Box<dyn std::error::Error>> {
16626 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
16627 e.with_moe_cache(max_block, |c, eng| {
16628 for (proj, exps) in [
16629 (PROJ_GATE, &m.gate_exps),
16630 (PROJ_UP, &m.up_exps),
16631 (PROJ_DOWN, &m.down_exps),
16632 ] {
16633 let source = exps.expert_source(ex);
16634 if let crate::model::ExpertSource::Disk { .. } = &source {
16635 let id = BlockId::new(il, proj, ex as u16);
16636 let _ = c.prefetch_source(id, source, keep, eng)?;
16637 }
16638 }
16639 Ok(())
16640 })
16641 }
16642
16643 #[inline]
16644 fn moe_prefetch_host_expert(ex: usize, m: &MoeWeights) {
16645 let _ = m.gate_exps.prefetch_expert_pages(ex);
16646 let _ = m.up_exps.prefetch_expert_pages(ex);
16647 let _ = m.down_exps.prefetch_expert_pages(ex);
16648 }
16649}
16650
16651impl HybridModel {
16668 #[allow(clippy::too_many_arguments)]
16672 fn moe_ffn_grouped_resident_q8(
16673 e: &Engine,
16674 m: &MoeWeights,
16675 z: &CudaSlice<f32>,
16676 t: usize,
16677 cfg: &ModelConfig,
16678 il: u16,
16679 sel_all: &[u32],
16680 w_all: &[f32],
16681 table: &CudaSlice<u64>,
16682 gu_il: bool,
16683 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
16684 crate::moe_rp_refuse(
16685 m.dev_exps.as_ref().is_some_and(|d| d.rp),
16686 "moe_ffn_grouped_resident_q8",
16687 )?; let moe = cfg.moe.as_ref().unwrap();
16689 let n_embd = cfg.n_embd as usize;
16690 let n_expert = moe.expert_count as usize;
16691 let n_used = moe.expert_used_count as usize;
16692 let n_ff_exp = moe.expert_ff_length as usize;
16693 let n_pairs = t * n_used;
16694 debug_assert_eq!(sel_all.len(), n_pairs);
16695 debug_assert_eq!(w_all.len(), n_pairs);
16696 debug_assert!(
16697 m.gate_exps.macros.is_none()
16698 && m.up_exps.macros.is_none()
16699 && m.down_exps.macros.is_none(),
16700 "resident grouped q8 does not fold per-expert macro scales",
16701 );
16702
16703 if !cfg.swiglu_clamped_at(il as u32) {
16709 let sel: Vec<i32> = sel_all.iter().map(|&expert| expert as i32).collect();
16710 let sel_d = e.htod_i32(&sel)?;
16711 let w_d = e.htod(w_all)?;
16712 let (gate_row_bytes, up_row_bytes) = if gu_il {
16713 let combined = m.gate_exps.row_bytes + m.up_exps.row_bytes;
16714 (combined, combined)
16715 } else {
16716 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
16717 };
16718 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
16719 let act = e.moe_gate_up_silu8_dev_q8_rows(
16720 table,
16721 &sel_d,
16722 &zq,
16723 &zd,
16724 t,
16725 n_embd,
16726 n_ff_exp,
16727 n_used,
16728 n_expert,
16729 m.gate_exps.qtype,
16730 m.up_exps.qtype,
16731 gate_row_bytes,
16732 up_row_bytes,
16733 &m.dev_macros,
16734 )?;
16735 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
16736 let mut moe_out = e.uninit(t * n_embd)?;
16737 e.moe_down8_fma_dev_q8_rows_g(
16738 table,
16739 &sel_d,
16740 &w_d,
16741 &aq2,
16742 &ad2,
16743 &mut moe_out,
16744 t,
16745 n_ff_exp,
16746 n_embd,
16747 n_used,
16748 n_expert,
16749 m.down_exps.qtype,
16750 m.down_exps.row_bytes,
16751 )?;
16752
16753 if std::env::var("MEMRA_MOE_STATS").is_ok() {
16754 let mut counts = vec![0usize; n_expert];
16755 for &expert in sel_all {
16756 counts[expert as usize] += 1;
16757 }
16758 let mut sizes: Vec<usize> =
16759 counts.into_iter().filter(|&count| count != 0).collect();
16760 sizes.sort_unstable();
16761 let mean = sizes.iter().sum::<usize>() as f64 / sizes.len().max(1) as f64;
16762 println!(
16763 "moe-grouped il={il} t={t} dispatch=resident-q8-rows active={}/{} \
16764 m_e: min={} median={} mean={mean:.1} max={}",
16765 sizes.len(),
16766 n_expert,
16767 sizes.first().copied().unwrap_or(0),
16768 sizes.get(sizes.len() / 2).copied().unwrap_or(0),
16769 sizes.last().copied().unwrap_or(0),
16770 );
16771 }
16772 return Ok(moe_out);
16773 }
16774
16775 let pair_tok: Vec<i32> = (0..n_pairs).map(|pair| (pair / n_used) as i32).collect();
16779 let pair_ex: Vec<i32> = sel_all.iter().map(|&expert| expert as i32).collect();
16780 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
16781 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
16782
16783 let mut by_expert: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
16784 for (pair, &expert) in pair_ex.iter().enumerate() {
16785 by_expert[expert as usize].push(pair as i32);
16786 }
16787
16788 let pair_tok_d = e.htod_i32(&pair_tok)?;
16789 let pair_ex_d = e.htod_i32(&pair_ex)?;
16790 let pair_w_d = e.htod(w_all)?;
16791 let tok_off_d = e.htod_i32(&tok_off)?;
16792 let tok_ids_d = e.htod_i32(&tok_ids)?;
16793
16794 let matvec = |proj: i32,
16795 pair_rows: &CudaSlice<i32>,
16796 aq: &CudaSlice<i8>,
16797 ad: &CudaSlice<f32>,
16798 in_f: usize,
16799 out_f: usize,
16800 qtype: i32,
16801 row_bytes: usize|
16802 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
16803 e.moe_pairs_matvec_q8(
16804 table, proj, pair_rows, &pair_ex_d, aq, ad, in_f, out_f, n_expert, n_pairs, qtype,
16805 row_bytes,
16806 )
16807 };
16808
16809 let (gate_row_bytes, up_row_bytes) = if gu_il {
16810 let combined = m.gate_exps.row_bytes + m.up_exps.row_bytes;
16811 (combined, combined)
16812 } else {
16813 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
16814 };
16815 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
16816 let gate = matvec(
16817 0,
16818 &pair_tok_d,
16819 &zq,
16820 &zd,
16821 n_embd,
16822 n_ff_exp,
16823 m.gate_exps.qtype,
16824 gate_row_bytes,
16825 )?;
16826 let up = matvec(
16827 1,
16828 &pair_tok_d,
16829 &zq,
16830 &zd,
16831 n_embd,
16832 n_ff_exp,
16833 m.up_exps.qtype,
16834 up_row_bytes,
16835 )?;
16836 let mut act = e.uninit(n_pairs * n_ff_exp)?;
16837 Self::ffn_act_lim(
16838 e,
16839 cfg,
16840 &gate,
16841 &up,
16842 1.0,
16843 1.0,
16844 cfg.clamp_exp_at(il as u32),
16845 &mut act,
16846 n_pairs * n_ff_exp,
16847 )?;
16848 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
16849 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
16850 let pair_self_d = e.htod_i32(&pair_self)?;
16851 let down = matvec(
16852 2,
16853 &pair_self_d,
16854 &aq2,
16855 &ad2,
16856 n_ff_exp,
16857 n_embd,
16858 m.down_exps.qtype,
16859 m.down_exps.row_bytes,
16860 )?;
16861 let mut moe_out = e.uninit(t * n_embd)?;
16862 e.moe_pairs_scatter(
16863 &down,
16864 &pair_w_d,
16865 &tok_off_d,
16866 &tok_ids_d,
16867 &mut moe_out,
16868 t,
16869 n_embd,
16870 )?;
16871
16872 if std::env::var("MEMRA_MOE_STATS").is_ok() {
16873 let mut sizes: Vec<usize> = by_expert
16874 .iter()
16875 .filter_map(|pairs| (!pairs.is_empty()).then_some(pairs.len()))
16876 .collect();
16877 sizes.sort_unstable();
16878 let mean = sizes.iter().sum::<usize>() as f64 / sizes.len().max(1) as f64;
16879 println!(
16880 "moe-grouped il={il} t={t} dispatch=resident-q8-clamped-pairs active={}/{} \
16881 m_e: min={} median={} mean={mean:.1} max={}",
16882 sizes.len(),
16883 n_expert,
16884 sizes.first().copied().unwrap_or(0),
16885 sizes.get(sizes.len() / 2).copied().unwrap_or(0),
16886 sizes.last().copied().unwrap_or(0),
16887 );
16888 }
16889 Ok(moe_out)
16890 }
16891
16892 #[allow(clippy::too_many_arguments)]
16898 #[allow(clippy::map_entry)] fn shexp_split_matvec(
16900 e: &Engine,
16901 rank1: &Engine,
16902 wg: &CudaSlice<u8>,
16903 wu: &CudaSlice<u8>,
16904 wd: &CudaSlice<u8>,
16905 z: &CudaSlice<f32>,
16906 lim: Option<SwigluClamp>,
16907 cfg: &ModelConfig,
16908 il: u16,
16909 n_embd: usize,
16910 n_ff_sh: usize,
16911 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
16912 use cudarc::driver::DevicePtr;
16913 if !n_ff_sh.is_multiple_of(2) || !n_embd.is_multiple_of(2) {
16914 return Ok(None);
16915 }
16916 let hf = n_ff_sh / 2;
16917 let nd = n_embd / 2;
16918 struct Rep {
16919 wg1: CudaSlice<u8>,
16920 wu1: CudaSlice<u8>,
16921 wd1: CudaSlice<u8>,
16922 }
16923 struct SplitWs {
16924 pin_dev: usize,
16925 gate0: CudaSlice<f32>,
16927 up0: CudaSlice<f32>,
16928 act: CudaSlice<f32>,
16929 sh_buf: CudaSlice<f32>,
16930 ev_z: cudarc::driver::CudaEvent,
16931 ev_act0: cudarc::driver::CudaEvent,
16932 z1: CudaSlice<f32>,
16934 g1: CudaSlice<f32>,
16935 u1: CudaSlice<f32>,
16936 a1h: CudaSlice<f32>,
16937 act1: CudaSlice<f32>,
16938 y1: CudaSlice<f32>,
16939 ev_act1: cudarc::driver::CudaEvent,
16940 ev_y1: cudarc::driver::CudaEvent,
16941 raw_act_e: u64,
16942 raw_sh_e: u64,
16943 raw_z1: u64,
16944 raw_a1h: u64,
16945 raw_act1: u64,
16946 raw_y1: u64,
16947 }
16948 static WS: std::sync::Mutex<Option<SplitWs>> = std::sync::Mutex::new(None);
16949 static REPS: std::sync::Mutex<Option<std::collections::HashMap<u64, Rep>>> =
16950 std::sync::Mutex::new(None);
16951 let mut guard = WS.lock().map_err(|_| "shexp split lock is poisoned")?;
16952 let mut reps_guard = REPS.lock().map_err(|_| "shexp reps lock is poisoned")?;
16953 let reps = reps_guard.get_or_insert_with(std::collections::HashMap::new);
16954 let pins = e.ctx().ordinal();
16955 if guard.as_ref().is_none_or(|w| w.pin_dev != pins) {
16956 let (gate0, up0, act, sh_buf, ev_z, ev_act0) = {
16957 let _m = e.gpu.enter_main()?;
16958 (
16959 e.htod(&vec![0.0f32; hf])?,
16960 e.htod(&vec![0.0f32; hf])?,
16961 e.htod(&vec![0.0f32; n_ff_sh])?,
16962 e.htod(&vec![0.0f32; n_embd])?,
16963 e.ctx().new_event(None)?,
16964 e.ctx().new_event(None)?,
16965 )
16966 };
16967 let (z1, g1, u1, a1h, act1, y1, ev_act1, ev_y1) = {
16968 let _r = rank1.gpu.enter_main()?;
16969 (
16970 rank1.htod(&vec![0.0f32; n_embd])?,
16971 rank1.htod(&vec![0.0f32; hf])?,
16972 rank1.htod(&vec![0.0f32; hf])?,
16973 rank1.htod(&vec![0.0f32; hf])?,
16974 rank1.htod(&vec![0.0f32; n_ff_sh])?,
16975 rank1.htod(&vec![0.0f32; nd])?,
16976 rank1.ctx().new_event(None)?,
16977 rank1.ctx().new_event(None)?,
16978 )
16979 };
16980 let (raw_act_e, raw_sh_e) = {
16981 let _m = e.gpu.enter_main()?;
16982 let stream = e.stream();
16983 let (a, _g0) = act.device_ptr(&stream);
16984 let (b, _g1) = sh_buf.device_ptr(&stream);
16985 (a, b)
16986 };
16987 let (raw_z1, raw_a1h, raw_act1, raw_y1) = {
16988 let _r = rank1.gpu.enter_main()?;
16989 let rs = rank1.stream();
16990 let (a, _g0) = z1.device_ptr(&rs);
16991 let (b, _g1) = a1h.device_ptr(&rs);
16992 let (c, _g2) = act1.device_ptr(&rs);
16993 let (d, _g3) = y1.device_ptr(&rs);
16994 (a, b, c, d)
16995 };
16996 *guard = Some(SplitWs {
16997 pin_dev: pins,
16998 gate0,
16999 up0,
17000 act,
17001 sh_buf,
17002 ev_z,
17003 ev_act0,
17004 z1,
17005 g1,
17006 u1,
17007 a1h,
17008 act1,
17009 y1,
17010 ev_act1,
17011 ev_y1,
17012 raw_act_e,
17013 raw_sh_e,
17014 raw_z1,
17015 raw_a1h,
17016 raw_act1,
17017 raw_y1,
17018 });
17019 }
17020 let ws = guard.as_mut().expect("armed above");
17021 let wg_pin = {
17022 let _m = e.gpu.enter_main()?;
17023 let stream = e.stream();
17024 let (p, _g) = wg.device_ptr(&stream);
17025 p
17026 };
17027 if !reps.contains_key(&wg_pin) {
17028 let up = |src: &CudaSlice<u8>,
17030 off_bytes: usize,
17031 len: usize|
17032 -> Result<CudaSlice<u8>, Box<dyn std::error::Error>> {
17033 use cudarc::driver::sys;
17034 let sptr = {
17035 let _m = e.gpu.enter_main()?;
17036 let stream = e.stream();
17037 let (p, _g) = src.device_ptr(&stream);
17038 p + off_bytes as u64
17039 };
17040 let dst = {
17041 let _r = rank1.gpu.enter_main()?;
17042 rank1.alloc_u8_uninit(len)?
17043 };
17044 let dptr = {
17045 let _r = rank1.gpu.enter_main()?;
17046 let rs = rank1.stream();
17047 let (p, _g) = dst.device_ptr(&rs);
17048 p
17049 };
17050 let _r = rank1.gpu.enter_main()?;
17051 let r = unsafe {
17052 sys::cuMemcpyAsync(
17053 dptr as sys::CUdeviceptr,
17054 sptr as sys::CUdeviceptr,
17055 len,
17056 rank1.stream().cu_stream() as sys::CUstream,
17057 )
17058 };
17059 if r != sys::CUresult::CUDA_SUCCESS {
17060 return Err(format!("shexp split replica upload: {r:?}").into());
17061 }
17062 rank1.stream().synchronize()?;
17063 Ok(dst)
17064 };
17065 let wg1 = up(wg, hf * n_embd * 2, hf * n_embd * 2)?;
17066 let wu1 = up(wu, hf * n_embd * 2, hf * n_embd * 2)?;
17067 let wd1 = up(wd, nd * n_ff_sh * 2, nd * n_ff_sh * 2)?;
17068 reps.insert(wg_pin, Rep { wg1, wu1, wd1 });
17069 }
17070 let _ = il;
17071 let raw_z = {
17073 let _m = e.gpu.enter_main()?;
17074 let stream = e.stream();
17075 let (p, _g) = z.device_ptr(&stream);
17076 ws.ev_z.record(&stream)?;
17077 p
17078 };
17079 {
17081 let rep = reps.get(&wg_pin).expect("uploaded above");
17082 let _r = rank1.gpu.enter_main()?;
17083 rank1.stream().wait(&ws.ev_z)?;
17084 crate::tp::raw_copy_bytes(ws.raw_z1, raw_z, n_embd * 4, rank1)?;
17085 let SplitWs {
17086 z1, g1, u1, a1h, ..
17087 } = &mut *ws;
17088 rank1.matvec_bf16_dual_into(&rep.wg1, &rep.wu1, z1, g1, u1, n_embd, hf)?;
17089 Self::ffn_act_lim(rank1, cfg, g1, u1, 1.0, 1.0, lim, a1h, hf)?;
17090 crate::tp::raw_copy_bytes(ws.raw_act1 + (hf * 4) as u64, ws.raw_a1h, hf * 4, rank1)?;
17092 crate::tp::raw_copy_bytes(ws.raw_act_e + (hf * 4) as u64, ws.raw_a1h, hf * 4, rank1)?;
17093 ws.ev_act1.record(&rank1.stream())?;
17094 }
17095 {
17097 let _m = e.gpu.enter_main()?;
17098 let SplitWs {
17099 gate0, up0, act, ..
17100 } = &mut *ws;
17101 let wg_lo = wg.slice(0..hf * n_embd * 2);
17102 let wu_lo = wu.slice(0..hf * n_embd * 2);
17103 e.matvec_bf16_dual_view_into(&wg_lo, &wu_lo, z, gate0, up0, n_embd, hf)?;
17104 Self::ffn_act_lim(e, cfg, gate0, up0, 1.0, 1.0, lim, act, hf)?;
17105 ws.ev_act0.record(&e.stream())?;
17106 }
17107 {
17109 let rep = reps.get(&wg_pin).expect("uploaded above");
17110 let _r = rank1.gpu.enter_main()?;
17111 rank1.stream().wait(&ws.ev_act0)?;
17112 crate::tp::raw_copy_bytes(ws.raw_act1, ws.raw_act_e, hf * 4, rank1)?;
17113 let SplitWs { act1, y1, .. } = &mut *ws;
17114 rank1.matvec_bf16_into(&rep.wd1, act1, y1, n_ff_sh, nd)?;
17115 crate::tp::raw_copy_bytes(ws.raw_sh_e + (nd * 4) as u64, ws.raw_y1, nd * 4, rank1)?;
17116 ws.ev_y1.record(&rank1.stream())?;
17117 }
17118 {
17120 let _m = e.gpu.enter_main()?;
17121 e.stream().wait(&ws.ev_act1)?;
17122 let SplitWs { act, sh_buf, .. } = &mut *ws;
17123 let wd_lo = wd.slice(0..nd * n_ff_sh * 2);
17124 e.matvec_bf16_view_into(&wd_lo, act, sh_buf, n_ff_sh, nd)?;
17125 e.stream().wait(&ws.ev_y1)?;
17126 let mut sh = e.uninit(n_embd)?;
17127 {
17128 let mut dst = sh.slice_mut(0..n_embd);
17129 e.stream()
17130 .memcpy_dtod(&ws.sh_buf.slice(0..n_embd), &mut dst)?;
17131 }
17132 Ok(Some(sh))
17133 }
17134 }
17135
17136 fn shexp_overlap_issue(
17143 e: &Engine,
17144 m: &MoeWeights,
17145 z: &CudaSlice<f32>,
17146 cfg: &ModelConfig,
17147 il: u16,
17148 n_embd: usize,
17149 ) -> Result<bool, Box<dyn std::error::Error>> {
17150 if cfg.m3.is_some() || m.gate_inp_shexp.is_some() {
17151 return Ok(false);
17152 }
17153 let (
17154 Some(crate::model::GpuTensor::FloatBf16 { data: wg, .. }),
17155 Some(crate::model::GpuTensor::FloatBf16 { data: wu, .. }),
17156 Some(crate::model::GpuTensor::FloatBf16 { data: wd, .. }),
17157 ) = (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
17158 else {
17159 return Ok(false);
17160 };
17161 let n_ff_sh = m
17162 .gate_shexp
17163 .as_ref()
17164 .expect("matched Some above")
17165 .out_features();
17166 let Ok(lim) = Self::fused_post_limit(cfg.clamp_shexp_at(il as u32)) else {
17169 return Ok(false);
17170 };
17171 let mut guard = SHEXP_OV_WS
17172 .lock()
17173 .map_err(|_| "shexp overlap workspace lock is poisoned")?;
17174 let pins = (e.ctx().ordinal(), n_embd, n_ff_sh);
17175 if guard
17176 .as_ref()
17177 .is_none_or(|(d, ne, nf, ..)| (*d, *ne, *nf) != pins)
17178 {
17179 *guard = Some((
17180 pins.0,
17181 pins.1,
17182 pins.2,
17183 e.uninit(n_ff_sh)?,
17184 e.uninit(n_embd)?,
17185 ));
17186 }
17187 let (_, _, _, act, sh) = guard.as_mut().expect("armed above");
17188 e.matvec_bf16_dual_silu_into(wg, wu, z, act, n_embd, n_ff_sh, lim)?;
17189 e.matvec_bf16_into(wd, act, sh, n_ff_sh, n_embd)?;
17190 drop(guard);
17191 Ok(true)
17192 }
17193
17194 #[allow(clippy::too_many_arguments)]
17200 #[allow(clippy::map_entry)] fn shexp_dev1_issue(
17202 e: &Engine,
17203 rank1: &Engine,
17204 m: &MoeWeights,
17205 z: &CudaSlice<f32>,
17206 cfg: &ModelConfig,
17207 il: u16,
17208 n_embd: usize,
17209 ) -> Result<bool, Box<dyn std::error::Error>> {
17210 use cudarc::driver::DevicePtr;
17211 if cfg.m3.is_some() || m.gate_inp_shexp.is_some() {
17212 return Ok(false);
17213 }
17214 let (
17215 Some(crate::model::GpuTensor::FloatBf16 { data: wg, .. }),
17216 Some(crate::model::GpuTensor::FloatBf16 { data: wu, .. }),
17217 Some(crate::model::GpuTensor::FloatBf16 { data: wd, .. }),
17218 ) = (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
17219 else {
17220 return Ok(false);
17221 };
17222 let n_ff_sh = m
17223 .gate_shexp
17224 .as_ref()
17225 .expect("matched Some above")
17226 .out_features();
17227 let Ok(lim) = Self::fused_post_limit(cfg.clamp_shexp_at(il as u32)) else {
17230 return Ok(false);
17231 };
17232 let mut ws_guard = SHEXP_D1_WS
17234 .lock()
17235 .map_err(|_| "shexp dev1 workspace lock is poisoned")?;
17236 if ws_guard
17237 .as_ref()
17238 .is_none_or(|(k, ..)| *k != (n_embd, n_ff_sh))
17239 {
17240 let (act1, z1, ev_done) = {
17241 let _r1 = rank1.gpu.enter_main()?;
17242 (
17243 rank1.htod(&vec![0.0f32; n_ff_sh])?,
17244 rank1.htod(&vec![0.0f32; n_embd])?,
17245 rank1.ctx().new_event(None)?,
17246 )
17247 };
17248 let (sh_root, ev_z) = {
17249 let _main = e.gpu.enter_main()?;
17250 (e.htod(&vec![0.0f32; n_embd])?, e.ctx().new_event(None)?)
17251 };
17252 *ws_guard = Some(((n_embd, n_ff_sh), act1, z1, sh_root, ev_z, ev_done));
17253 }
17254 let mut reps_guard = SHEXP_D1_REPS
17256 .lock()
17257 .map_err(|_| "shexp dev1 replica lock is poisoned")?;
17258 let reps = reps_guard.get_or_insert_with(Default::default);
17259 if !reps.contains_key(&il) {
17260 let (wg1, wu1, wd1) = {
17261 let _r1 = rank1.gpu.enter_main()?;
17262 (
17263 rank1.alloc_u8_uninit(n_ff_sh * n_embd * 2)?,
17264 rank1.alloc_u8_uninit(n_ff_sh * n_embd * 2)?,
17265 rank1.alloc_u8_uninit(n_embd * n_ff_sh * 2)?,
17266 )
17267 };
17268 for (src, dst) in [(wg, &wg1), (wu, &wu1), (wd, &wd1)] {
17269 let s_ptr = {
17270 let _main = e.gpu.enter_main()?;
17271 let stream = e.stream();
17272 let (p, _g) = src.device_ptr(&stream);
17273 p
17274 };
17275 let d_ptr = {
17276 let _r1 = rank1.gpu.enter_main()?;
17277 let stream = rank1.stream();
17278 let (p, _g) = dst.device_ptr(&stream);
17279 p
17280 };
17281 let _r1 = rank1.gpu.enter_main()?;
17282 crate::tp::raw_copy_bytes(d_ptr, s_ptr, src.len(), rank1)?;
17283 }
17284 {
17285 let _r1 = rank1.gpu.enter_main()?;
17286 rank1.stream().synchronize()?;
17287 }
17288 reps.insert(il, (wg1, wu1, wd1));
17289 }
17290 let (wg1, wu1, wd1) = reps.get(&il).expect("armed above");
17291 let (_, act1, z1, sh_root, ev_z, ev_done) = ws_guard.as_mut().expect("armed above");
17292 let (raw_z, raw_sh) = {
17295 let _main = e.gpu.enter_main()?;
17296 let stream = e.stream();
17297 let (a, _g0) = z.device_ptr(&stream);
17298 let (b, _g1) = sh_root.device_ptr(&stream);
17299 ev_z.record(&stream)?;
17300 (a, b)
17301 };
17302 {
17303 let _r1 = rank1.gpu.enter_main()?;
17304 rank1.stream().wait(ev_z)?;
17305 let raw_z1 = {
17306 let stream = rank1.stream();
17307 let (p, _g) = z1.device_ptr(&stream);
17308 p
17309 };
17310 crate::tp::raw_copy_bytes(raw_z1, raw_z, n_embd * 4, rank1)?;
17311 rank1.matvec_bf16_dual_silu_into(wg1, wu1, z1, act1, n_embd, n_ff_sh, lim)?;
17312 rank1.matvec_bf16_raw_out(wd1, act1, raw_sh, n_ff_sh, n_embd)?;
17316 ev_done.record(&rank1.stream())?;
17317 }
17318 Ok(true)
17319 }
17320
17321 fn shexp_dev1_apply(
17323 e: &Engine,
17324 output: &mut CudaSlice<f32>,
17325 n_embd: usize,
17326 ) -> Result<(), Box<dyn std::error::Error>> {
17327 let guard = SHEXP_D1_WS
17328 .lock()
17329 .map_err(|_| "shexp dev1 workspace lock is poisoned")?;
17330 let (pin, _, _, sh_root, _, ev_done) =
17331 guard.as_ref().ok_or("shexp dev1 apply without issue")?;
17332 if pin.0 != n_embd {
17333 return Err("shexp dev1 width drifted".into());
17334 }
17335 let _main = e.gpu.enter_main()?;
17336 e.stream().wait(ev_done)?;
17337 static ONES_D1: std::sync::Mutex<Option<(usize, CudaSlice<f32>)>> =
17338 std::sync::Mutex::new(None);
17339 let mut og = ONES_D1.lock().map_err(|_| "ones lock is poisoned")?;
17340 if og.as_ref().is_none_or(|(d, _)| *d != e.ctx().ordinal()) {
17341 *og = Some((e.ctx().ordinal(), e.htod(&[1.0f32])?));
17342 }
17343 let ones = &og.as_ref().expect("armed above").1;
17344 e.add_scaled_rows(sh_root, ones, output, n_embd, 1)?;
17345 Ok(())
17346 }
17347
17348 fn shexp_overlap_tail_ptrs(
17352 e: &Engine,
17353 m: &MoeWeights,
17354 cfg: &ModelConfig,
17355 n_embd: usize,
17356 ) -> Result<Option<(u64, u64)>, Box<dyn std::error::Error>> {
17357 use cudarc::driver::DevicePtr;
17358 if cfg.m3.is_some() || m.gate_inp_shexp.is_some() {
17359 return Ok(None);
17360 }
17361 let (
17362 Some(crate::model::GpuTensor::FloatBf16 { .. }),
17363 Some(crate::model::GpuTensor::FloatBf16 { .. }),
17364 Some(crate::model::GpuTensor::FloatBf16 { .. }),
17365 ) = (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
17366 else {
17367 return Ok(None);
17368 };
17369 let n_ff_sh = m
17370 .gate_shexp
17371 .as_ref()
17372 .expect("matched Some above")
17373 .out_features();
17374 let mut guard = SHEXP_OV_WS
17375 .lock()
17376 .map_err(|_| "shexp overlap workspace lock is poisoned")?;
17377 let pins = (e.ctx().ordinal(), n_embd, n_ff_sh);
17378 if guard
17379 .as_ref()
17380 .is_none_or(|(d, ne, nf, ..)| (*d, *ne, *nf) != pins)
17381 {
17382 *guard = Some((
17383 pins.0,
17384 pins.1,
17385 pins.2,
17386 e.uninit(n_ff_sh)?,
17387 e.uninit(n_embd)?,
17388 ));
17389 }
17390 let sh_raw = {
17391 let (_, _, _, _, sh) = guard.as_ref().expect("armed above");
17392 let stream = e.stream();
17393 let (p, _g) = sh.device_ptr(&stream);
17394 p
17395 };
17396 static ONES_T3: std::sync::Mutex<Option<(usize, CudaSlice<f32>)>> =
17397 std::sync::Mutex::new(None);
17398 let mut og = ONES_T3.lock().map_err(|_| "ones lock is poisoned")?;
17399 if og.as_ref().is_none_or(|(d, _)| *d != e.ctx().ordinal()) {
17400 *og = Some((e.ctx().ordinal(), e.htod(&[1.0f32])?));
17401 }
17402 let ones_raw = {
17403 let stream = e.stream();
17404 let (p, _g) = og.as_ref().expect("armed above").1.device_ptr(&stream);
17405 p
17406 };
17407 Ok(Some((sh_raw, ones_raw)))
17408 }
17409
17410 fn shexp_overlap_apply(
17413 e: &Engine,
17414 output: &mut CudaSlice<f32>,
17415 n_embd: usize,
17416 ) -> Result<(), Box<dyn std::error::Error>> {
17417 let guard = SHEXP_OV_WS
17418 .lock()
17419 .map_err(|_| "shexp overlap workspace lock is poisoned")?;
17420 let (_, ne, _, _, sh) = guard.as_ref().ok_or("shexp overlap apply without issue")?;
17421 if *ne != n_embd {
17422 return Err("shexp overlap width drifted".into());
17423 }
17424 static ONES_OV: std::sync::Mutex<Option<(usize, CudaSlice<f32>)>> =
17425 std::sync::Mutex::new(None);
17426 let mut og = ONES_OV.lock().map_err(|_| "ones lock is poisoned")?;
17427 if og.as_ref().is_none_or(|(d, _)| *d != e.ctx().ordinal()) {
17428 *og = Some((e.ctx().ordinal(), e.htod(&[1.0f32])?));
17429 }
17430 let ones = &og.as_ref().expect("armed above").1;
17431 e.add_scaled_rows(sh, ones, output, n_embd, 1)?;
17432 Ok(())
17433 }
17434
17435 fn moe_ffn_grouped_add_shared(
17436 e: &Engine,
17437 m: &MoeWeights,
17438 z: &CudaSlice<f32>,
17439 t: usize,
17440 cfg: &ModelConfig,
17441 il: u16,
17442 moe_out: &mut CudaSlice<f32>,
17443 ) -> Result<(), Box<dyn std::error::Error>> {
17444 static SHEXP_NS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
17447 static SHEXP_CALLS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
17448 let shexp_timing = std::env::var("MEMRA_STEP_TP_TIMING").as_deref() == Ok("1");
17449 let shexp_started = shexp_timing.then(std::time::Instant::now);
17450 let result = Self::moe_ffn_grouped_add_shared_inner(e, m, z, t, cfg, il, moe_out);
17451 if let Some(started) = shexp_started {
17452 use std::sync::atomic::Ordering;
17453 e.stream().synchronize()?;
17454 let ns = SHEXP_NS.fetch_add(started.elapsed().as_nanos() as u64, Ordering::Relaxed)
17455 + started.elapsed().as_nanos() as u64;
17456 let calls = SHEXP_CALLS.fetch_add(1, Ordering::Relaxed) + 1;
17457 if calls.is_multiple_of(430) {
17458 eprintln!(
17459 "[moe-shexp-timing] calls={calls} total_ms={:.1} avg_us={:.1}",
17460 ns as f64 / 1.0e6,
17461 ns as f64 / calls as f64 / 1.0e3,
17462 );
17463 }
17464 }
17465 result
17466 }
17467
17468 #[allow(clippy::too_many_arguments)]
17469 fn moe_ffn_grouped_add_shared_inner(
17470 e: &Engine,
17471 m: &MoeWeights,
17472 z: &CudaSlice<f32>,
17473 t: usize,
17474 cfg: &ModelConfig,
17475 il: u16,
17476 moe_out: &mut CudaSlice<f32>,
17477 ) -> Result<(), Box<dyn std::error::Error>> {
17478 let n_embd = cfg.n_embd as usize;
17479 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
17480 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
17481 {
17482 let n_ff_sh = gate_shexp.out_features();
17483 let lim = cfg.clamp_shexp_at(il as u32);
17484 let fused = t == 1
17491 && lim.is_none()
17492 && cfg.m3.is_none()
17493 && e.uses_q8_1_fast(gate_shexp)
17494 && e.uses_q8_1_fast(up_shexp);
17495 let canonical_w4a16_rows =
17496 t <= 32 && m.step_ep.as_ref().is_some_and(|ep| ep.nvfp4_device_routes);
17497 let bf16_dual = if (t == 1 || canonical_w4a16_rows)
17502 && crate::Engine::bf16_mmv_on()
17503 && n_embd.is_multiple_of(8)
17504 {
17505 match (gate_shexp, up_shexp) {
17506 (
17507 crate::model::GpuTensor::FloatBf16 { data: wg, .. },
17508 crate::model::GpuTensor::FloatBf16 { data: wu, .. },
17509 ) => Some((wg, wu)),
17510 _ => None,
17511 }
17512 } else {
17513 None
17514 };
17515 let sh = if let Some((wg, wu)) = bf16_dual {
17516 type SharedExpertWorkspace = (
17520 usize,
17521 usize,
17522 usize,
17523 CudaSlice<f32>,
17524 CudaSlice<f32>,
17525 CudaSlice<f32>,
17526 CudaSlice<f32>,
17527 );
17528 static SHEXP_WS: std::sync::Mutex<
17529 Option<std::collections::HashMap<usize, SharedExpertWorkspace>>,
17530 > = std::sync::Mutex::new(None);
17531 let down_bf16 = match down_shexp {
17532 crate::model::GpuTensor::FloatBf16 { data, .. } => Some(data),
17533 _ => None,
17534 };
17535 let mut guard = SHEXP_WS
17536 .lock()
17537 .map_err(|_| "shexp workspace lock is poisoned")?;
17538 let capacity = if canonical_w4a16_rows { 32 } else { 1 };
17539 let device = e.ctx().ordinal();
17540 let workspaces = guard.get_or_insert_with(Default::default);
17541 if workspaces
17542 .get(&device)
17543 .is_none_or(|(ne, nf, cap, ..)| (*ne, *nf, *cap) != (n_embd, n_ff_sh, capacity))
17544 {
17545 workspaces.insert(
17546 device,
17547 (
17548 n_embd,
17549 n_ff_sh,
17550 capacity,
17551 e.uninit(capacity * n_ff_sh)?,
17552 e.uninit(capacity * n_ff_sh)?,
17553 e.uninit(capacity * n_ff_sh)?,
17554 e.uninit(capacity * n_embd)?,
17555 ),
17556 );
17557 }
17558 {
17561 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
17562 let split_on = *ON
17563 .get_or_init(|| std::env::var("MEMRA_SHEXP_SPLIT").as_deref() == Ok("1"));
17564 if split_on
17565 && t == 1
17566 && let (Some(wd), Some(rank1)) = (
17567 match down_shexp {
17568 crate::model::GpuTensor::FloatBf16 { data, .. } => Some(data),
17569 _ => None,
17570 },
17571 m.step_tp.as_ref().and_then(|st| st.runtime.rank_engine(1)),
17572 )
17573 && let Some(sh) = Self::shexp_split_matvec(
17574 e, rank1, wg, wu, wd, z, lim, cfg, il, n_embd, n_ff_sh,
17575 )?
17576 {
17577 drop(guard);
17578 let gate = match &m.gate_inp_shexp {
17579 Some(gate_inp_shexp) => {
17580 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
17581 }
17582 None => e.htod(&vec![1.0f32; t])?,
17583 };
17584 e.add_scaled_rows(&sh, &gate, moe_out, n_embd, t)?;
17585 return Ok(());
17586 }
17587 }
17588 let (_, _, _, gate, up, act, sh_buf) = workspaces
17589 .get_mut(&device)
17590 .expect("shexp workspace initialized above");
17591 if let (true, Ok(lim_post)) = (cfg.m3.is_none(), Self::fused_post_limit(lim)) {
17594 if canonical_w4a16_rows {
17597 e.matvec_bf16_dual_silu_rows_into(
17598 wg, wu, z, act, n_embd, n_ff_sh, lim_post, t,
17599 )?;
17600 } else {
17601 e.matvec_bf16_dual_silu_into(wg, wu, z, act, n_embd, n_ff_sh, lim_post)?;
17602 }
17603 let _ = (&gate, &up);
17604 } else {
17605 e.matvec_bf16_dual_into(wg, wu, z, gate, up, n_embd, n_ff_sh)?;
17606 Self::ffn_act_lim(e, cfg, gate, up, 1.0, 1.0, lim, act, n_ff_sh)?;
17607 }
17608 if let Some(down) = down_bf16 {
17609 if canonical_w4a16_rows {
17610 e.matvec_bf16_rows_into(down, act, sh_buf, n_ff_sh, n_embd, t)?;
17611 let mut sh = e.uninit(t * n_embd)?;
17612 {
17613 let mut dst = sh.slice_mut(0..t * n_embd);
17614 e.stream()
17615 .memcpy_dtod(&sh_buf.slice(0..t * n_embd), &mut dst)?;
17616 }
17617 sh
17618 } else {
17619 static FUSE_DA: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
17626 let fuse_da = *FUSE_DA.get_or_init(|| {
17627 std::env::var("MEMRA_FUSE_DOWN_ADDSCALE").as_deref() != Ok("0")
17628 });
17629 if fuse_da && m.gate_inp_shexp.is_none() {
17630 static ONES1: std::sync::Mutex<Option<(usize, CudaSlice<f32>)>> =
17631 std::sync::Mutex::new(None);
17632 let mut og = ONES1.lock().map_err(|_| "shexp ones lock is poisoned")?;
17633 if og.as_ref().is_none_or(|(d, _)| *d != e.ctx().ordinal()) {
17634 *og = Some((e.ctx().ordinal(), e.htod(&[1.0f32])?));
17635 }
17636 let ones = &og.as_ref().expect("armed above").1;
17637 e.matvec_bf16_down_addscale_into(
17638 down, act, ones, moe_out, n_ff_sh, n_embd,
17639 )?;
17640 return Ok(());
17641 }
17642 e.matvec_bf16_into(down, act, sh_buf, n_ff_sh, n_embd)?;
17643 let sh = e.uninit(n_embd)?;
17644 let mut sh = sh;
17646 {
17647 let mut dst = sh.slice_mut(0..n_embd);
17648 e.stream().memcpy_dtod(&sh_buf.slice(0..n_embd), &mut dst)?;
17649 }
17650 sh
17651 }
17652 } else {
17653 e.matmul(down_shexp, act, t)?
17654 }
17655 } else if fused {
17656 let (zq, zd) = e.quantize_q8_1(z, 1, n_embd)?;
17657 let pair = match e.matmul_pre_dual_noscale(gate_shexp, up_shexp, &zq, &zd, 1)? {
17658 Some((gate, up)) => Some((gate, up)),
17659 None => {
17660 match (
17661 e.matmul_pre_noscale(gate_shexp, &zq, &zd, 1)?,
17662 e.matmul_pre_noscale(up_shexp, &zq, &zd, 1)?,
17663 ) {
17664 (Some(gate), Some(up)) => Some((gate, up)),
17665 _ => None,
17666 }
17667 }
17668 };
17669 match pair {
17670 Some(((gate, gs), (up, us))) => {
17671 if e.uses_q8_1_fast(down_shexp) {
17672 let (aq, ad) = e.silu_mul_scaled_q8_1(&gate, &up, gs, us, n_ff_sh)?;
17673 e.matmul_pre(down_shexp, &aq, &ad, &gate, 1)?
17674 } else {
17675 let mut act = e.uninit(n_ff_sh)?;
17676 e.silu_mul_scaled(&gate, &up, gs, us, &mut act, n_ff_sh)?;
17677 e.matmul(down_shexp, &act, 1)?
17678 }
17679 }
17680 None => {
17681 let gate = e.matmul_pre(gate_shexp, &zq, &zd, z, 1)?;
17682 let up = e.matmul_pre(up_shexp, &zq, &zd, z, 1)?;
17683 let mut act = e.uninit(n_ff_sh)?;
17684 Self::ffn_act(e, cfg, &gate, &up, &mut act, n_ff_sh)?;
17685 e.matmul(down_shexp, &act, 1)?
17686 }
17687 }
17688 } else {
17689 let sg_gate = e.matmul(gate_shexp, z, t)?;
17690 let sg_up = e.matmul(up_shexp, z, t)?;
17691 let mut sa = e.uninit(t * n_ff_sh)?;
17692 Self::ffn_act_lim(
17693 e,
17694 cfg,
17695 &sg_gate,
17696 &sg_up,
17697 1.0,
17698 1.0,
17699 lim,
17700 &mut sa,
17701 t * n_ff_sh,
17702 )?;
17703 e.matmul(down_shexp, &sa, t)?
17704 };
17705 let gate = match &m.gate_inp_shexp {
17706 Some(gate_inp_shexp) => {
17707 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
17708 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
17709 } else {
17710 let raw = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
17711 let mut gate = e.uninit(t)?;
17712 e.sigmoid(&raw, &mut gate, t)?;
17713 gate
17714 }
17715 }
17716 None if t == 1 => {
17721 static ONES: std::sync::Mutex<
17722 Option<std::collections::HashMap<usize, CudaSlice<f32>>>,
17723 > = std::sync::Mutex::new(None);
17724 let mut guard = ONES.lock().map_err(|_| "shexp ones lock is poisoned")?;
17725 let device = e.ctx().ordinal();
17726 let rows = guard.get_or_insert_with(Default::default);
17727 use std::collections::hash_map::Entry;
17730 let ones = match rows.entry(device) {
17731 Entry::Occupied(occupied) => occupied.into_mut(),
17732 Entry::Vacant(vacant) => vacant.insert(e.htod(&[1.0f32])?),
17733 };
17734 e.add_scaled_rows(&sh, ones, moe_out, n_embd, t)?;
17735 return Ok(());
17736 }
17737 None => e.htod(&vec![1.0f32; t])?,
17738 };
17739 e.add_scaled_rows(&sh, &gate, moe_out, n_embd, t)?;
17740 }
17741 Ok(())
17742 }
17743
17744 pub(crate) fn moe_ffn_grouped(
17747 e: &Engine,
17748 m: &MoeWeights,
17749 z: &CudaSlice<f32>,
17750 t: usize,
17751 cfg: &ModelConfig,
17752 il: u16,
17753 max_block: usize,
17754 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
17755 let moe = cfg.moe.as_ref().unwrap();
17756 let n_embd = cfg.n_embd as usize;
17757 let n_expert = moe.expert_count as usize;
17758 let n_used = moe.expert_used_count as usize;
17759 let n_ff_exp = moe.expert_ff_length as usize;
17760 let lim_exp = cfg.clamp_exp_at(il as u32);
17762
17763 let logits = Self::moe_router_logits(e, m, z, t, cfg)?;
17767 if let Some(sig) = cfg.sigmoid_router() {
17768 Self::trace_sigmoid_router_logits(e, il, t, n_expert, n_used, &logits, m, sig)?;
17769 }
17770 let (sel_all, w_all) = if let Some(sig) = cfg.sigmoid_router() {
17771 Self::moe_route_sigmoid_cfg(e, &logits, t, n_expert, n_used, m, sig)?
17772 } else {
17773 Self::moe_route_cfg(e, &logits, t, n_expert, n_used, m.active_experts.as_deref())?
17774 };
17775 crate::moesd::record_host_routes(il, n_expert, n_used, &sel_all)?;
17776 Self::trace_moe_routes(il, t, &sel_all, &w_all)?;
17777 Self::trace_moe_input(e, il, t, n_embd, z)?;
17778
17779 let no_exp_macros = m.gate_exps.macros.is_none()
17784 && m.up_exps.macros.is_none()
17785 && m.down_exps.macros.is_none();
17786 let resident_q8 = m.dev_exps.as_ref().filter(|dev| {
17787 m.has_uniform_expert_layout()
17788 && no_exp_macros
17789 && moe_q8_enabled_for_model(cfg, m)
17790 && moe_slab_enabled()
17791 && dev.dev == e.ctx().ordinal()
17792 });
17793 if let Some(dev) = resident_q8 {
17794 let mut moe_out = Self::moe_ffn_grouped_resident_q8(
17795 e,
17796 m,
17797 z,
17798 t,
17799 cfg,
17800 il,
17801 &sel_all,
17802 &w_all,
17803 &dev.ptr_row,
17804 dev.gu_il,
17805 )?;
17806 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut moe_out)?;
17807 return Ok(moe_out);
17808 }
17809
17810 struct ExpertGroup {
17814 tok_indices: Vec<i32>, slot_indices: Vec<i32>, weights: Vec<f32>, }
17818 let mut groups: Vec<ExpertGroup> = (0..n_expert)
17819 .map(|_| ExpertGroup {
17820 tok_indices: Vec::new(),
17821 slot_indices: Vec::new(),
17822 weights: Vec::new(),
17823 })
17824 .collect();
17825
17826 for tok in 0..t {
17827 for j in 0..n_used {
17828 let ex = sel_all[tok * n_used + j] as usize;
17829 let w = w_all[tok * n_used + j];
17830 groups[ex].tok_indices.push(tok as i32);
17831 groups[ex].slot_indices.push(j as i32);
17832 groups[ex].weights.push(w);
17833 }
17834 }
17835
17836 let mut slot_buf = e.zeros(t * n_used * n_embd)?;
17839 let mut wbuf = e.zeros(t * n_used)?; let g_len = m.gate_exps.max_expert_bytes();
17843 let u_len = m.up_exps.max_expert_bytes();
17844 let d_len = m.down_exps.max_expert_bytes();
17845 let moe_q8 = moe_q8_enabled_for_model(cfg, m);
17846 let slab_local = m
17849 .dev_exps
17850 .as_ref()
17851 .filter(|dev| !dev.gu_il && moe_slab_enabled() && dev.dev == e.ctx().ordinal());
17852 let use_cache =
17853 slab_local.is_none() && Engine::moe_cache_enabled() && !e.moe_cache_frozen();
17854 let grouped_q8 = moe_q8 && (slab_local.is_some() || use_cache);
17857
17858 let (mut scratch_g, mut scratch_u, mut scratch_d) = if slab_local.is_none() && !use_cache {
17860 (
17861 Some(e.alloc_u8(g_len)?),
17862 Some(e.alloc_u8(u_len)?),
17863 Some(e.alloc_u8(d_len)?),
17864 )
17865 } else {
17866 (None, None, None)
17867 };
17868
17869 let mut order: Vec<usize> = (0..n_expert)
17880 .filter(|&ex| !groups[ex].tok_indices.is_empty())
17881 .collect();
17882 order.sort_by(|&a, &b| {
17883 groups[b]
17884 .tok_indices
17885 .len()
17886 .cmp(&groups[a].tok_indices.len())
17887 .then(a.cmp(&b))
17888 });
17889 let mut m_dist: Vec<usize> = Vec::new(); let page_window = moe_page_prefetch_window();
17891 let worker_disk_prefetch = use_cache && crate::spill_pread::worker_enabled();
17892 if worker_disk_prefetch
17893 && let Some(first) = grouped_worker_prefetch_position(order.len(), None)
17894 {
17895 Self::moe_prefetch_disk_expert(e, il, order[first], m, max_block, &[])?;
17896 }
17897 for (order_pos, &ex) in order.iter().enumerate() {
17898 for next in page_prefetch_positions(order_pos, order.len(), page_window) {
17899 Self::moe_prefetch_host_expert(order[next], m);
17900 }
17901 if worker_disk_prefetch
17902 && let Some(next) = grouped_worker_prefetch_position(order.len(), Some(order_pos))
17903 {
17904 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
17905 let keep = [
17906 BlockId::new(il, PROJ_GATE, ex as u16),
17907 BlockId::new(il, PROJ_UP, ex as u16),
17908 BlockId::new(il, PROJ_DOWN, ex as u16),
17909 ];
17910 Self::moe_prefetch_disk_expert(e, il, order[next], m, max_block, &keep)?;
17911 }
17912 let grp = &groups[ex];
17913 let m_e = grp.tok_indices.len();
17914 m_dist.push(m_e);
17915 let gl = m.gate_exps.expert_layout(ex);
17916 let ul = m.up_exps.expert_layout(ex);
17917 let dl = m.down_exps.expert_layout(ex);
17918
17919 let tok_idx_d = e.htod_i32(&grp.tok_indices)?;
17923 let slot_idx_d = e.htod_i32(&grp.slot_indices)?;
17924 let dmac = m.down_exps.macro_scale(ex);
17925 let weight_d = if dmac == 1.0 {
17926 e.htod(&grp.weights)?
17927 } else {
17928 let scaled: Vec<f32> = grp.weights.iter().map(|&w| w * dmac).collect();
17929 e.htod(&scaled)?
17930 };
17931
17932 let mut gathered = e.zeros(m_e * n_embd)?;
17934 e.gather_rows(z, &tok_idx_d, &mut gathered, n_embd, m_e)?;
17935 let gv = gathered.slice(0..m_e * n_embd);
17936
17937 let y = if let Some(dev) = slab_local {
17940 let gate_start = ex * m.gate_exps.expert_stride;
17941 let up_start = ex * m.up_exps.expert_stride;
17942 let down_start = ex * m.down_exps.expert_stride;
17943 if grouped_q8 {
17944 let (zq, zd) = e.quantize_q8_1(&gathered, m_e, n_embd)?;
17945 let gate = e.qmatvec_expert_q8(
17946 &dev.gate,
17947 gate_start..gate_start + gl.len,
17948 &zq,
17949 &zd,
17950 m_e,
17951 m.gate_exps.in_f,
17952 m.gate_exps.out_f,
17953 gl.qtype,
17954 gl.row_bytes,
17955 )?;
17956 let up = e.qmatvec_expert_q8(
17957 &dev.up,
17958 up_start..up_start + ul.len,
17959 &zq,
17960 &zd,
17961 m_e,
17962 m.up_exps.in_f,
17963 m.up_exps.out_f,
17964 ul.qtype,
17965 ul.row_bytes,
17966 )?;
17967 let mut act = e.uninit(m_e * n_ff_exp)?;
17968 Self::ffn_act_lim(
17969 e,
17970 cfg,
17971 &gate,
17972 &up,
17973 m.gate_exps.macro_scale(ex),
17974 m.up_exps.macro_scale(ex),
17975 lim_exp,
17976 &mut act,
17977 m_e * n_ff_exp,
17978 )?;
17979 let (aq2, ad2) = e.quantize_q8_1(&act, m_e, n_ff_exp)?;
17980 e.qmatvec_expert_q8(
17981 &dev.down,
17982 down_start..down_start + dl.len,
17983 &aq2,
17984 &ad2,
17985 m_e,
17986 m.down_exps.in_f,
17987 m.down_exps.out_f,
17988 dl.qtype,
17989 dl.row_bytes,
17990 )?
17991 } else {
17992 let gate = m.qmatvec_view(
17993 e,
17994 &dev.gate,
17995 gate_start..gate_start + gl.len,
17996 &gv,
17997 m_e,
17998 m.gate_exps.in_f,
17999 m.gate_exps.out_f,
18000 gl.qtype,
18001 gl.row_bytes,
18002 )?;
18003 let up = m.qmatvec_view(
18004 e,
18005 &dev.up,
18006 up_start..up_start + ul.len,
18007 &gv,
18008 m_e,
18009 m.up_exps.in_f,
18010 m.up_exps.out_f,
18011 ul.qtype,
18012 ul.row_bytes,
18013 )?;
18014 let mut act = e.uninit(m_e * n_ff_exp)?;
18015 Self::ffn_act_lim(
18016 e,
18017 cfg,
18018 &gate,
18019 &up,
18020 m.gate_exps.macro_scale(ex),
18021 m.up_exps.macro_scale(ex),
18022 lim_exp,
18023 &mut act,
18024 m_e * n_ff_exp,
18025 )?;
18026 let actv = act.slice(0..m_e * n_ff_exp);
18027 m.qmatvec_view(
18028 e,
18029 &dev.down,
18030 down_start..down_start + dl.len,
18031 &actv,
18032 m_e,
18033 m.down_exps.in_f,
18034 m.down_exps.out_f,
18035 dl.qtype,
18036 dl.row_bytes,
18037 )?
18038 }
18039 } else if use_cache {
18040 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
18041 if grouped_q8 {
18042 let (zq, zd) = e.quantize_q8_1(&gathered, m_e, n_embd)?;
18043 let gate = e.with_moe_cache(max_block, |cache, eng| {
18044 let id = BlockId::new(il, PROJ_GATE, ex as u16);
18045 let slot = cache.dispatch_source(id, m.gate_exps.expert_source(ex), eng)?;
18046 eng.qmatvec_expert_q8(
18047 cache.buf(slot),
18048 0..gl.len,
18049 &zq,
18050 &zd,
18051 m_e,
18052 m.gate_exps.in_f,
18053 m.gate_exps.out_f,
18054 gl.qtype,
18055 gl.row_bytes,
18056 )
18057 })?;
18058 let up = e.with_moe_cache(max_block, |cache, eng| {
18059 let id = BlockId::new(il, PROJ_UP, ex as u16);
18060 let slot = cache.dispatch_source(id, m.up_exps.expert_source(ex), eng)?;
18061 eng.qmatvec_expert_q8(
18062 cache.buf(slot),
18063 0..ul.len,
18064 &zq,
18065 &zd,
18066 m_e,
18067 m.up_exps.in_f,
18068 m.up_exps.out_f,
18069 ul.qtype,
18070 ul.row_bytes,
18071 )
18072 })?;
18073 let mut act = e.uninit(m_e * n_ff_exp)?;
18074 Self::ffn_act_lim(
18075 e,
18076 cfg,
18077 &gate,
18078 &up,
18079 m.gate_exps.macro_scale(ex),
18080 m.up_exps.macro_scale(ex),
18081 lim_exp,
18082 &mut act,
18083 m_e * n_ff_exp,
18084 )?;
18085 let (aq2, ad2) = e.quantize_q8_1(&act, m_e, n_ff_exp)?;
18086 e.with_moe_cache(max_block, |cache, eng| {
18087 let id = BlockId::new(il, PROJ_DOWN, ex as u16);
18088 let slot = cache.dispatch_source(id, m.down_exps.expert_source(ex), eng)?;
18089 eng.qmatvec_expert_q8(
18090 cache.buf(slot),
18091 0..dl.len,
18092 &aq2,
18093 &ad2,
18094 m_e,
18095 m.down_exps.in_f,
18096 m.down_exps.out_f,
18097 dl.qtype,
18098 dl.row_bytes,
18099 )
18100 })?
18101 } else {
18102 let gate = e.with_moe_cache(max_block, |cache, eng| {
18103 let id = BlockId::new(il, PROJ_GATE, ex as u16);
18104 let slot = cache.dispatch_source(id, m.gate_exps.expert_source(ex), eng)?;
18105 m.qmatvec_view(
18106 eng,
18107 cache.buf(slot),
18108 0..gl.len,
18109 &gv,
18110 m_e,
18111 m.gate_exps.in_f,
18112 m.gate_exps.out_f,
18113 gl.qtype,
18114 gl.row_bytes,
18115 )
18116 })?;
18117 let up = e.with_moe_cache(max_block, |cache, eng| {
18118 let id = BlockId::new(il, PROJ_UP, ex as u16);
18119 let slot = cache.dispatch_source(id, m.up_exps.expert_source(ex), eng)?;
18120 m.qmatvec_view(
18121 eng,
18122 cache.buf(slot),
18123 0..ul.len,
18124 &gv,
18125 m_e,
18126 m.up_exps.in_f,
18127 m.up_exps.out_f,
18128 ul.qtype,
18129 ul.row_bytes,
18130 )
18131 })?;
18132 let mut act = e.uninit(m_e * n_ff_exp)?;
18133 Self::ffn_act_lim(
18134 e,
18135 cfg,
18136 &gate,
18137 &up,
18138 m.gate_exps.macro_scale(ex),
18139 m.up_exps.macro_scale(ex),
18140 lim_exp,
18141 &mut act,
18142 m_e * n_ff_exp,
18143 )?;
18144 let actv = act.slice(0..m_e * n_ff_exp);
18145 e.with_moe_cache(max_block, |cache, eng| {
18146 let id = BlockId::new(il, PROJ_DOWN, ex as u16);
18147 let slot = cache.dispatch_source(id, m.down_exps.expert_source(ex), eng)?;
18148 m.qmatvec_view(
18149 eng,
18150 cache.buf(slot),
18151 0..dl.len,
18152 &actv,
18153 m_e,
18154 m.down_exps.in_f,
18155 m.down_exps.out_f,
18156 dl.qtype,
18157 dl.row_bytes,
18158 )
18159 })?
18160 }
18161 } else {
18162 let sg = scratch_g.as_mut().unwrap();
18163 let su = scratch_u.as_mut().unwrap();
18164 let sd = scratch_d.as_mut().unwrap();
18165 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
18166 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
18167 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
18168 if grouped_q8 {
18169 let (zq, zd) = e.quantize_q8_1(&gathered, m_e, n_embd)?;
18170 let gate = e.qmatvec_expert_q8(
18171 sg,
18172 0..gl.len,
18173 &zq,
18174 &zd,
18175 m_e,
18176 m.gate_exps.in_f,
18177 m.gate_exps.out_f,
18178 gl.qtype,
18179 gl.row_bytes,
18180 )?;
18181 let up = e.qmatvec_expert_q8(
18182 su,
18183 0..ul.len,
18184 &zq,
18185 &zd,
18186 m_e,
18187 m.up_exps.in_f,
18188 m.up_exps.out_f,
18189 ul.qtype,
18190 ul.row_bytes,
18191 )?;
18192 let mut act = e.uninit(m_e * n_ff_exp)?;
18193 Self::ffn_act_lim(
18194 e,
18195 cfg,
18196 &gate,
18197 &up,
18198 m.gate_exps.macro_scale(ex),
18199 m.up_exps.macro_scale(ex),
18200 lim_exp,
18201 &mut act,
18202 m_e * n_ff_exp,
18203 )?;
18204 let (aq2, ad2) = e.quantize_q8_1(&act, m_e, n_ff_exp)?;
18205 e.qmatvec_expert_q8(
18206 sd,
18207 0..dl.len,
18208 &aq2,
18209 &ad2,
18210 m_e,
18211 m.down_exps.in_f,
18212 m.down_exps.out_f,
18213 dl.qtype,
18214 dl.row_bytes,
18215 )?
18216 } else {
18217 let gate = m.qmatvec_view(
18218 e,
18219 sg,
18220 0..gl.len,
18221 &gv,
18222 m_e,
18223 m.gate_exps.in_f,
18224 m.gate_exps.out_f,
18225 gl.qtype,
18226 gl.row_bytes,
18227 )?;
18228 let up = m.qmatvec_view(
18229 e,
18230 su,
18231 0..ul.len,
18232 &gv,
18233 m_e,
18234 m.up_exps.in_f,
18235 m.up_exps.out_f,
18236 ul.qtype,
18237 ul.row_bytes,
18238 )?;
18239 let mut act = e.uninit(m_e * n_ff_exp)?;
18240 Self::ffn_act_lim(
18241 e,
18242 cfg,
18243 &gate,
18244 &up,
18245 m.gate_exps.macro_scale(ex),
18246 m.up_exps.macro_scale(ex),
18247 lim_exp,
18248 &mut act,
18249 m_e * n_ff_exp,
18250 )?;
18251 let actv = act.slice(0..m_e * n_ff_exp);
18252 m.qmatvec_view(
18253 e,
18254 sd,
18255 0..dl.len,
18256 &actv,
18257 m_e,
18258 m.down_exps.in_f,
18259 m.down_exps.out_f,
18260 dl.qtype,
18261 dl.row_bytes,
18262 )?
18263 }
18264 };
18265
18266 e.scatter_slot(
18268 &y,
18269 &tok_idx_d,
18270 &slot_idx_d,
18271 &weight_d,
18272 &mut slot_buf,
18273 &mut wbuf,
18274 n_embd,
18275 n_used,
18276 m_e,
18277 )?;
18278 }
18279
18280 let mut moe_out = e.zeros(t * n_embd)?;
18282 e.reduce_slots(&slot_buf, &wbuf, &mut moe_out, n_embd, n_used, t)?;
18283
18284 if std::env::var("MEMRA_MOE_STATS").is_ok() && !m_dist.is_empty() {
18286 m_dist.sort_unstable();
18287 let active = m_dist.len();
18288 let mean = m_dist.iter().sum::<usize>() as f64 / active as f64;
18289 let median = m_dist[active / 2];
18290 let max_m = *m_dist.last().unwrap();
18291 let min_m = m_dist[0];
18292 let above16 = m_dist.iter().filter(|&&x| x >= 16).count();
18293 println!(
18294 "moe-grouped il={il} t={t} active={active}/{n_expert} \
18295 m_e: min={min_m} median={median} mean={mean:.1} max={max_m} \
18296 above_gemm_threshold(>=16)={above16}/{active}"
18297 );
18298 }
18299
18300 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut moe_out)?;
18301 Ok(moe_out)
18302 }
18303
18304 pub(crate) fn moe_ffn_lockstep(
18311 &self,
18312 e: &Engine,
18313 m: &MoeWeights,
18314 zbatch: &CudaSlice<f32>,
18315 mrows: usize,
18316 il: u16,
18317 max_block: usize,
18318 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
18319 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
18320 let cfg = &self.cfg;
18321 let moe = cfg.moe.as_ref().unwrap();
18322 let n_embd = cfg.n_embd as usize;
18323 let n_expert = moe.expert_count as usize;
18324 let n_used = moe.expert_used_count as usize;
18325 let n_ff_exp = moe.expert_ff_length as usize;
18326 let lim_exp = cfg.clamp_exp_at(il as u32);
18328 let lim_shexp = cfg.clamp_shexp_at(il as u32);
18329
18330 let logits = e.matmul(&m.gate_inp, zbatch, mrows)?;
18331 if let Some(sig) = cfg.sigmoid_router() {
18332 Self::trace_sigmoid_router_logits(e, il, mrows, n_expert, n_used, &logits, m, sig)?;
18333 }
18334 let (sel_all, w_all) = if let Some(sig) = cfg.sigmoid_router() {
18335 Self::moe_route_sigmoid_cfg(e, &logits, mrows, n_expert, n_used, m, sig)?
18336 } else {
18337 Self::moe_route_cfg(
18338 e,
18339 &logits,
18340 mrows,
18341 n_expert,
18342 n_used,
18343 m.active_experts.as_deref(),
18344 )?
18345 };
18346 Self::trace_moe_routes(il, mrows, &sel_all, &w_all)?;
18347
18348 let resident_expert: Vec<bool> = e.with_moe_cache(max_block, |c, _| {
18350 Ok((0..n_expert)
18351 .map(|ex| {
18352 [PROJ_GATE, PROJ_UP, PROJ_DOWN]
18353 .into_iter()
18354 .all(|p| c.resident(BlockId::new(il, p, ex as u16)).is_some())
18355 })
18356 .collect())
18357 })?;
18358
18359 struct Group {
18360 rows: Vec<i32>,
18361 slots: Vec<i32>,
18362 weights: Vec<f32>,
18363 }
18364 let mut groups: std::collections::HashMap<usize, Group> = Default::default();
18365 let mut cpu_rows: Vec<Vec<(usize, f32)>> = vec![Vec::new(); mrows];
18366 let mut cpu_by_expert: std::collections::HashMap<usize, Vec<(usize, f32)>> =
18367 Default::default();
18368 for row in 0..mrows {
18369 for j in 0..n_used {
18370 let ex = sel_all[row * n_used + j] as usize;
18371 let w = w_all[row * n_used + j];
18372 if resident_expert[ex] {
18373 let group = groups.entry(ex).or_insert_with(|| Group {
18374 rows: Vec::new(),
18375 slots: Vec::new(),
18376 weights: Vec::new(),
18377 });
18378 group.rows.push(row as i32);
18379 group.slots.push(j as i32);
18380 group.weights.push(w);
18381 } else {
18382 crate::cpu_experts::record_incomplete_gpu_residency(0);
18383 cpu_rows[row].push((ex, w));
18384 cpu_by_expert.entry(ex).or_default().push((row, w));
18385 }
18386 }
18387 }
18388
18389 let host_rows = e.dtoh(zbatch)?;
18395 let rows_ok = crate::cpu_experts::rows_supported();
18396 enum CpuPart {
18397 Single { row: usize },
18398 Rows { rows: Vec<usize> },
18399 }
18400 let mut tickets: Vec<(CpuPart, crate::cpu_experts::CpuExpertTicket)> = Vec::new();
18401 let mut rows_served: std::collections::HashSet<(usize, usize)> = Default::default();
18402 if rows_ok {
18403 let mut shared: Vec<(usize, Vec<(usize, f32)>)> = cpu_by_expert
18404 .into_iter()
18405 .filter(|(_, rows)| rows.len() >= 2)
18406 .collect();
18407 shared.sort_by_key(|(ex, _)| *ex);
18408 for (ex, mut row_weights) in shared {
18409 row_weights.sort_by_key(|(row, _)| *row);
18410 let inputs: Vec<(&[f32], f32)> = row_weights
18411 .iter()
18412 .map(|&(row, w)| (&host_rows[row * n_embd..(row + 1) * n_embd], w))
18413 .collect();
18414 let job = crate::cpu_experts::prepare_rows_job(m, ex, &inputs)
18415 .map_err(std::io::Error::other)?;
18416 for &(row, _) in &row_weights {
18417 rows_served.insert((row, ex));
18418 }
18419 tickets.push((
18420 CpuPart::Rows {
18421 rows: row_weights.iter().map(|&(row, _)| row).collect(),
18422 },
18423 crate::cpu_experts::submit_rows(job).map_err(std::io::Error::other)?,
18424 ));
18425 }
18426 }
18427 for (row, selected) in cpu_rows.iter().enumerate() {
18428 let leftover: Vec<(usize, f32)> = selected
18429 .iter()
18430 .copied()
18431 .filter(|&(ex, _)| !rows_served.contains(&(row, ex)))
18432 .collect();
18433 if leftover.is_empty() {
18434 continue;
18435 }
18436 let host_row = &host_rows[row * n_embd..(row + 1) * n_embd];
18437 let job = crate::cpu_experts::prepare_job(m, il, &leftover, host_row)
18438 .map_err(std::io::Error::other)?;
18439 tickets.push((
18440 CpuPart::Single { row },
18441 crate::cpu_experts::submit(job).map_err(std::io::Error::other)?,
18442 ));
18443 }
18444
18445 let mut slot_buf = e.zeros(mrows * n_used * n_embd)?;
18446 let mut wbuf = e.zeros(mrows * n_used)?;
18447 let mut order: Vec<usize> = groups.keys().copied().collect();
18448 order.sort_by(|&a, &b| {
18449 groups[&b]
18450 .rows
18451 .len()
18452 .cmp(&groups[&a].rows.len())
18453 .then(a.cmp(&b))
18454 });
18455 for &ex in &order {
18456 let group = &groups[&ex];
18457 let m_e = group.rows.len();
18458 let gl = m.gate_exps.expert_layout(ex);
18459 let ul = m.up_exps.expert_layout(ex);
18460 let dl = m.down_exps.expert_layout(ex);
18461 let row_idx_d = e.htod_i32(&group.rows)?;
18462 let slot_idx_d = e.htod_i32(&group.slots)?;
18463 let dmac = m.down_exps.macro_scale(ex);
18464 let weight_d = if dmac == 1.0 {
18465 e.htod(&group.weights)?
18466 } else {
18467 let scaled: Vec<f32> = group.weights.iter().map(|&w| w * dmac).collect();
18468 e.htod(&scaled)?
18469 };
18470 let mut gathered = e.zeros(m_e * n_embd)?;
18471 e.gather_rows(zbatch, &row_idx_d, &mut gathered, n_embd, m_e)?;
18472 let gv = gathered.slice(0..m_e * n_embd);
18473 let gate = e.with_moe_cache(max_block, |c, eng| {
18474 let slot = c
18475 .resident(BlockId::new(il, PROJ_GATE, ex as u16))
18476 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
18477 m.qmatvec_view(
18478 eng,
18479 c.buf(crate::moe_cache::DispatchSlot::Resident(slot)),
18480 0..gl.len,
18481 &gv,
18482 m_e,
18483 m.gate_exps.in_f,
18484 m.gate_exps.out_f,
18485 gl.qtype,
18486 gl.row_bytes,
18487 )
18488 })?;
18489 let up = e.with_moe_cache(max_block, |c, eng| {
18490 let slot = c
18491 .resident(BlockId::new(il, PROJ_UP, ex as u16))
18492 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
18493 m.qmatvec_view(
18494 eng,
18495 c.buf(crate::moe_cache::DispatchSlot::Resident(slot)),
18496 0..ul.len,
18497 &gv,
18498 m_e,
18499 m.up_exps.in_f,
18500 m.up_exps.out_f,
18501 ul.qtype,
18502 ul.row_bytes,
18503 )
18504 })?;
18505 let mut act = e.zeros(m_e * n_ff_exp)?;
18506 Self::ffn_act_lim(
18507 e,
18508 cfg,
18509 &gate,
18510 &up,
18511 m.gate_exps.macro_scale(ex),
18512 m.up_exps.macro_scale(ex),
18513 lim_exp,
18514 &mut act,
18515 m_e * n_ff_exp,
18516 )?;
18517 let actv = act.slice(0..m_e * n_ff_exp);
18518 let y = e.with_moe_cache(max_block, |c, eng| {
18519 let slot = c
18520 .resident(BlockId::new(il, PROJ_DOWN, ex as u16))
18521 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
18522 m.qmatvec_view(
18523 eng,
18524 c.buf(crate::moe_cache::DispatchSlot::Resident(slot)),
18525 0..dl.len,
18526 &actv,
18527 m_e,
18528 m.down_exps.in_f,
18529 m.down_exps.out_f,
18530 dl.qtype,
18531 dl.row_bytes,
18532 )
18533 })?;
18534 e.scatter_slot(
18535 &y,
18536 &row_idx_d,
18537 &slot_idx_d,
18538 &weight_d,
18539 &mut slot_buf,
18540 &mut wbuf,
18541 n_embd,
18542 n_used,
18543 m_e,
18544 )?;
18545 }
18546 let mut moe_out = e.zeros(mrows * n_embd)?;
18547 e.reduce_slots(&slot_buf, &wbuf, &mut moe_out, n_embd, n_used, mrows)?;
18548
18549 let mut row_sums: Vec<Option<Vec<f32>>> = vec![None; mrows];
18551 for (part, ticket) in tickets {
18552 let cpu_output = ticket.wait().map_err(std::io::Error::other)?;
18553 let mut add_row = |row: usize, chunk: &[f32]| {
18554 let sum = row_sums[row].get_or_insert_with(|| vec![0.0f32; n_embd]);
18555 for (accumulator, value) in sum.iter_mut().zip(chunk) {
18556 *accumulator += value;
18557 }
18558 };
18559 match part {
18560 CpuPart::Single { row } => add_row(row, &cpu_output),
18561 CpuPart::Rows { rows } => {
18562 for (slot, row) in rows.into_iter().enumerate() {
18563 add_row(row, &cpu_output[slot * n_embd..(slot + 1) * n_embd]);
18564 }
18565 }
18566 }
18567 }
18568 for (row, sum) in row_sums.into_iter().enumerate() {
18569 let Some(sum) = sum else { continue };
18570 let cpu_output = e.htod(&sum)?;
18571 let mut dst = moe_out.slice_mut(row * n_embd..(row + 1) * n_embd);
18572 e.axpy_into(&cpu_output, 1.0, &mut dst, n_embd)?;
18573 }
18574
18575 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
18576 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
18577 {
18578 let n_ff_sh = gate_shexp.out_features();
18579 let sg_gate = e.matmul(gate_shexp, zbatch, mrows)?;
18580 let sg_up = e.matmul(up_shexp, zbatch, mrows)?;
18581 let mut sa = e.zeros(mrows * n_ff_sh)?;
18582 Self::ffn_act_lim(
18583 e,
18584 cfg,
18585 &sg_gate,
18586 &sg_up,
18587 1.0,
18588 1.0,
18589 lim_shexp,
18590 &mut sa,
18591 mrows * n_ff_sh,
18592 )?;
18593 let sh = e.matmul(down_shexp, &sa, mrows)?;
18594 let g = match &m.gate_inp_shexp {
18597 Some(gate_inp_shexp) => {
18598 e.sigmoid_dot_rows(zbatch, gate_inp_shexp.float_data(), n_embd, mrows)?
18599 }
18600 None => e.htod(&vec![1.0f32; mrows])?,
18601 };
18602 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, mrows)?;
18603 }
18604
18605 Ok(moe_out)
18606 }
18607}
18608
18609impl HybridModel {
18615 pub(crate) fn gemma4_rope_dims(&self, il: usize) -> usize {
18629 let g = self
18630 .cfg
18631 .gemma4
18632 .as_ref()
18633 .expect("gemma4_rope_dims on a non-gemma4 config");
18634 if g.swa_pattern[il] {
18635 g.rope_dims_swa as usize
18636 } else {
18637 g.rope_dims_global as usize
18638 }
18639 }
18640
18641 pub(crate) fn gemma4_geom(&self, il: usize) -> (usize, usize, usize, f32, f32, bool) {
18642 let g = self.cfg.gemma4.as_ref().unwrap();
18643 let swa = g.swa_pattern[il];
18644 let hd = if swa {
18645 g.key_length_swa
18646 } else {
18647 g.key_length_global
18648 } as usize;
18649 (
18653 hd,
18654 g.head_count_kv[il] as usize,
18655 self.cfg.n_head as usize,
18656 if swa {
18657 g.rope_base_swa
18658 } else {
18659 g.rope_base_global
18660 },
18661 1.0,
18662 swa,
18663 )
18664 }
18665
18666 pub(crate) fn gemma4_suppress(
18670 &self,
18671 e: &Engine,
18672 ld: &mut CudaSlice<f32>,
18673 t: usize,
18674 ) -> Result<(), Box<dyn std::error::Error>> {
18675 if let Some((ids, n)) = self.gemma4_aux.as_ref().and_then(|a| a.suppress_d.as_ref()) {
18676 #[cfg(debug_assertions)]
18681 crate::debug_assert_tensor_stream_device(
18682 ids,
18683 &e.stream(),
18684 "gemma4_suppress.suppress_d",
18685 );
18686 e.mask_ids_rows(ld, ids, *n, self.output.out_features(), t)?;
18687 }
18688 Ok(())
18689 }
18690
18691 #[allow(clippy::too_many_arguments)]
18696 fn gemma_fa_one_program() -> bool {
18705 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
18706 *ON.get_or_init(|| std::env::var("MEMRA_GEMMA_FA_ONE_PROGRAM").as_deref() == Ok("1"))
18707 }
18708
18709 #[allow(clippy::too_many_arguments)] fn gemma4_attn_prime(
18711 &self,
18712 e: &Engine,
18713 fa: &crate::hybrid::FullAttnLayer,
18714 il: usize,
18715 h: &CudaSlice<f32>,
18716 pos_d: &CudaSlice<i32>,
18717 t: usize,
18718 cache: Option<&mut Cache>,
18719 island: Option<&CudaSlice<i32>>,
18720 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
18721 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
18722 let eps = self.cfg.rms_eps;
18723 let aux = self.gemma4_aux.as_ref().unwrap();
18724 let ones = aux.ones(e);
18725 #[cfg(debug_assertions)]
18726 crate::debug_assert_tensor_stream_device(ones, &e.stream(), "gemma4_attn_prime.ones");
18727
18728 e.mmq_act_begin();
18731 let q0 = e.matmul(&fa.wq, h, t)?; if il == 0 && std::env::var("MEMRA_G4_PRIME_TRACE").as_deref() == Ok("1") {
18733 let v = e.dtoh(&q0)?;
18734 let nan = v.iter().filter(|x| x.is_nan()).count();
18735 let amax = v.iter().fold(0f32, |a, x| a.max(x.abs()));
18736 eprintln!(
18737 "[g4-prime-trace] L0 q0: nan={nan}/{} amax={amax:.3}",
18738 v.len()
18739 );
18740 }
18741 let k0 = e.matmul(&fa.wk, h, t)?; let v0 = if swa {
18745 e.matmul(&fa.wv, h, t)?
18746 } else {
18747 e.clone_dtod(&k0)?
18748 };
18749 if il == 0 && std::env::var("MEMRA_G4_PRIME_TRACE").as_deref() == Ok("1") {
18750 for (tag, buf) in [("k0", &k0), ("v0", &v0)] {
18751 let v = e.dtoh(buf)?;
18752 let nan = v.iter().filter(|x| x.is_nan()).count();
18753 let amax = v.iter().fold(0f32, |a, x| a.max(x.abs()));
18754 eprintln!(
18755 "[g4-prime-trace] L0 {tag}: nan={nan}/{} amax={amax:.3}",
18756 v.len()
18757 );
18758 }
18759 }
18760
18761 let mut q = e.uninit(t * nh * hd)?;
18762 let mut k = e.uninit(t * nkv * hd)?;
18763 let mut v = e.uninit(t * nkv * hd)?;
18765 static EMIT: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
18769 let emit = island.is_none()
18772 && t >= 16
18773 && crate::Engine::qkvnorm_w_on_prefill(nh * t + 2 * nkv * t, hd)
18774 && *EMIT.get_or_init(|| {
18775 std::env::var("MEMRA_FA_EMIT")
18776 .map(|s| s != "0")
18777 .unwrap_or(true)
18778 });
18779 let mut qb = e.alloc_uninit::<u8>(if emit { t * nh * hd * 2 } else { 1 })?;
18780 let mut kb = e.alloc_uninit::<u8>(if emit { t * nkv * hd * 2 } else { 1 })?;
18781 let mut vb = e.alloc_uninit::<u8>(if emit { t * nkv * hd * 2 } else { 1 })?;
18782 let v_f16 = emit
18785 && crate::fa_f16pv_on()
18786 && match hd {
18787 512 => true,
18788 256 => swa && crate::faw_hp_on() && nh % 2 == 0 && (nh / nkv) % 2 == 0,
18789 _ => false,
18790 };
18791 if emit {
18792 e.rms_norm_qkv_w4b(
18793 &q0,
18794 &k0,
18795 &v0,
18796 fa.q_norm.float_data(),
18797 fa.k_norm.float_data(),
18798 ones,
18799 &mut q,
18800 &mut k,
18801 &mut v,
18802 &mut vb,
18803 hd,
18804 nh * t,
18805 nkv * t,
18806 eps,
18807 v_f16,
18808 )?;
18809 } else {
18810 e.rms_norm_qkv(
18811 &q0,
18812 &k0,
18813 &v0,
18814 fa.q_norm.float_data(),
18815 fa.k_norm.float_data(),
18816 ones,
18817 &mut q,
18818 &mut k,
18819 &mut v,
18820 hd,
18821 nh * t,
18822 nkv * t,
18823 eps,
18824 )?;
18825 }
18826
18827 let ff = if swa {
18828 None
18829 } else {
18830 Some(
18831 aux.rope_freqs(e)
18832 .expect("gemma4 global rope needs rope_freqs.weight"),
18833 )
18834 };
18835 #[cfg(debug_assertions)]
18836 if let Some(ff) = ff {
18837 crate::debug_assert_tensor_stream_device(
18838 ff,
18839 &e.stream(),
18840 "gemma4_attn_prime.rope_freqs",
18841 );
18842 }
18843 if emit {
18844 e.rope_neox2_bf16e(
18845 &mut q, &mut k, &mut qb, &mut kb, pos_d, hd, hd, nh, nkv, t, base, 1.0, ff,
18846 )?;
18847 } else {
18848 e.rope_neox2(&mut q, &mut k, pos_d, hd, hd, nh, nkv, t, base, 1.0, ff)?;
18849 }
18850
18851 if let Some(cache) = cache {
18852 let kvl = cache.kv[il].as_mut().unwrap();
18853 assert_eq!(kvl.len, 0, "gemma4 prime is fresh-prompt only (v0)");
18854 e.append_kv_quantized_rows(
18855 &k,
18856 &v,
18857 &mut kvl.k,
18858 &mut kvl.v,
18859 kvl.len,
18860 t,
18861 kvl.kv_dim_k,
18862 kvl.kv_dim_v,
18863 kvl.k_tok_bytes,
18864 kvl.v_tok_bytes,
18865 (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on()),
18866 )?;
18867 kvl.len += t;
18868 }
18869 let mut attn = e.zeros(t * nh * hd)?;
18870 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
18874 if let Some(span) = island {
18875 let w = if swa && t > win { win } else { 0 };
18880 e.sdpa_naive_island(&q, &k, &v, &mut attn, span, hd, nh, nkv, t, t, scale, w)?;
18881 } else if swa && (t > win || Self::gemma_fa_one_program()) {
18882 if hd == 256 && std::env::var("MEMRA_NOFA").is_err() {
18883 if emit {
18884 e.fa_prefill_w_pre(
18885 &qb, &kb, &vb, &mut attn, hd, nh, nkv, t, t, scale, true, win, v_f16,
18886 )?;
18887 } else {
18888 e.fa_prefill_w(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
18889 }
18890 } else {
18891 e.sdpa_naive_w(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
18892 }
18893 } else if hd == 256 && std::env::var("MEMRA_NOFA").is_err() {
18894 e.fa_prefill(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
18895 } else if hd == 512 && std::env::var("MEMRA_NOFA").is_err() {
18896 if emit {
18897 e.fa_prefill_hd512_pre(
18898 &qb, &kb, &vb, &mut attn, hd, nh, nkv, t, t, scale, true, v_f16,
18899 )?;
18900 } else {
18901 e.fa_prefill_hd512(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
18902 }
18903 } else {
18904 e.sdpa_naive(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
18905 }
18906 e.matmul(&fa.wo, &attn, t)
18907 }
18908
18909 fn gemma4_attn(
18911 &self,
18912 e: &Engine,
18913 fa: &crate::hybrid::FullAttnLayer,
18914 il: usize,
18915 h: &CudaSlice<f32>,
18916 pos_d: &CudaSlice<i32>,
18917 t: usize,
18918 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
18919 self.gemma4_attn_prime(e, fa, il, h, pos_d, t, None, None)
18920 }
18921
18922 fn gemma4_moe_q8(
18927 &self,
18928 e: &Engine,
18929 m: &crate::hybrid::MoeWeights,
18930 bits: &crate::hybrid::Gemma4MoeBits,
18931 mq: &(CudaSlice<i8>, CudaSlice<f32>),
18932 router_in: &CudaSlice<f32>,
18933 t: usize,
18934 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
18935 let cfg = &self.cfg;
18936 let moe = cfg.moe.as_ref().unwrap();
18937 let n_embd = cfg.n_embd as usize;
18938 let n_expert = moe.expert_count as usize;
18939 let n_used = moe.expert_used_count as usize;
18940 let n_ff_exp = moe.expert_ff_length as usize;
18941 let logits = if crate::router_kernel_on() {
18945 e.router_gemv(m.gate_inp.float_data(), router_in, n_embd, n_expert, t)?
18946 } else {
18947 e.matmul(&m.gate_inp, router_in, t)?
18948 };
18949 let dev = m.dev_exps.as_ref().unwrap();
18950 let (sel_d, w_d) =
18951 e.moe_router_topk_scaled(&logits, t, n_expert, n_used, &bits.per_expert_scale_d)?;
18952 let (zq, zd) = mq;
18953 if t == 1 {
18954 let selv = sel_d.slice(0..n_used);
18955 let wv = w_d.slice(0..n_used);
18956 let act = e.moe_gate_up_gelu8_dev_q8(
18957 &dev.ptr_row,
18958 &selv,
18959 zq,
18960 zd,
18961 n_embd,
18962 n_ff_exp,
18963 n_used,
18964 n_expert,
18965 m.gate_exps.qtype,
18966 m.up_exps.qtype,
18967 m.gate_exps.row_bytes,
18968 m.up_exps.row_bytes,
18969 )?;
18970 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
18971 let mut moe_out = e.uninit(n_embd)?;
18972 e.moe_down8_fma_dev_q8(
18973 &dev.ptr_row,
18974 &selv,
18975 &wv,
18976 &aq2,
18977 &ad2,
18978 &mut moe_out.slice_mut(0..n_embd),
18979 n_ff_exp,
18980 n_embd,
18981 n_used,
18982 n_expert,
18983 m.down_exps.qtype,
18984 m.down_exps.row_bytes,
18985 )?;
18986 return Ok(moe_out);
18987 }
18988 let csr = t <= 10 && std::env::var("MEMRA_GEMMA_CSR").as_deref() != Ok("0");
18989 let act = if csr {
18990 e.moe_gate_up_gelu8_dev_q8_csr(
18991 &dev.ptr_row,
18992 &sel_d,
18993 zq,
18994 zd,
18995 t * n_used,
18996 n_embd,
18997 n_ff_exp,
18998 n_used,
18999 n_expert,
19000 m.gate_exps.qtype,
19001 m.up_exps.qtype,
19002 m.gate_exps.row_bytes,
19003 m.up_exps.row_bytes,
19004 )?
19005 } else {
19006 e.moe_gate_up_gelu8_dev_q8_rows(
19007 &dev.ptr_row,
19008 &sel_d,
19009 zq,
19010 zd,
19011 t,
19012 n_embd,
19013 n_ff_exp,
19014 n_used,
19015 n_expert,
19016 m.gate_exps.qtype,
19017 m.up_exps.qtype,
19018 m.gate_exps.row_bytes,
19019 m.up_exps.row_bytes,
19020 )?
19021 };
19022 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
19023 let mut moe_out = e.uninit(t * n_embd)?;
19024 e.moe_down8_fma_dev_q8_rows_g(
19027 &dev.ptr_row,
19028 &sel_d,
19029 &w_d,
19030 &aq2,
19031 &ad2,
19032 &mut moe_out,
19033 t,
19034 n_ff_exp,
19035 n_embd,
19036 n_used,
19037 n_expert,
19038 m.down_exps.qtype,
19039 m.down_exps.row_bytes,
19040 )?;
19041 Ok(moe_out)
19042 }
19043
19044 fn gemma4_moe(
19048 &self,
19049 e: &Engine,
19050 m: &crate::hybrid::MoeWeights,
19051 bits: &crate::hybrid::Gemma4MoeBits,
19052 moe_in: &CudaSlice<f32>,
19053 router_in: &CudaSlice<f32>,
19054 t: usize,
19055 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
19056 crate::moe_rp_refuse(m.dev_exps.as_ref().is_some_and(|d| d.rp), "gemma4_moe")?; let cfg = &self.cfg;
19058 let moe = cfg.moe.as_ref().unwrap();
19059 let n_embd = cfg.n_embd as usize;
19060 let n_expert = moe.expert_count as usize;
19061 let n_used = moe.expert_used_count as usize;
19062 let n_ff_exp = moe.expert_ff_length as usize;
19063
19064 let logits = if t < PRIME_MIN_T && crate::router_kernel_on() {
19068 e.router_gemv(m.gate_inp.float_data(), router_in, n_embd, n_expert, t)?
19069 } else {
19070 e.matmul(&m.gate_inp, router_in, t)?
19071 };
19072
19073 if t < PRIME_MIN_T
19078 && m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
19079 && expert_dp4a_supported(m.gate_exps.qtype)
19080 && expert_dp4a_supported(m.up_exps.qtype)
19081 && expert_dp4a_supported(m.down_exps.qtype)
19082 && std::env::var("MEMRA_GEMMA_MOE_FAST").as_deref() != Ok("0")
19083 {
19084 let dev = m.dev_exps.as_ref().unwrap();
19085 let (sel_d, w_d) =
19086 e.moe_router_topk_scaled(&logits, t, n_expert, n_used, &bits.per_expert_scale_d)?;
19087 if t == 1 {
19088 let (zq, zd) = e.quantize_q8_1(moe_in, 1, n_embd)?;
19089 let selv = sel_d.slice(0..n_used);
19090 let wv = w_d.slice(0..n_used);
19091 let act = e.moe_gate_up_gelu8_dev_q8(
19092 &dev.ptr_row,
19093 &selv,
19094 &zq,
19095 &zd,
19096 n_embd,
19097 n_ff_exp,
19098 n_used,
19099 n_expert,
19100 m.gate_exps.qtype,
19101 m.up_exps.qtype,
19102 m.gate_exps.row_bytes,
19103 m.up_exps.row_bytes,
19104 )?;
19105 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
19106 let mut moe_out = e.uninit(n_embd)?;
19107 e.moe_down8_fma_dev_q8(
19108 &dev.ptr_row,
19109 &selv,
19110 &wv,
19111 &aq2,
19112 &ad2,
19113 &mut moe_out.slice_mut(0..n_embd),
19114 n_ff_exp,
19115 n_embd,
19116 n_used,
19117 n_expert,
19118 m.down_exps.qtype,
19119 m.down_exps.row_bytes,
19120 )?;
19121 return Ok(moe_out);
19122 }
19123 let (zq, zd) = e.quantize_q8_1(moe_in, t, n_embd)?;
19128 let csr = t <= 10 && std::env::var("MEMRA_GEMMA_CSR").as_deref() != Ok("0");
19129 let act = if csr {
19130 e.moe_gate_up_gelu8_dev_q8_csr(
19131 &dev.ptr_row,
19132 &sel_d,
19133 &zq,
19134 &zd,
19135 t * n_used,
19136 n_embd,
19137 n_ff_exp,
19138 n_used,
19139 n_expert,
19140 m.gate_exps.qtype,
19141 m.up_exps.qtype,
19142 m.gate_exps.row_bytes,
19143 m.up_exps.row_bytes,
19144 )?
19145 } else {
19146 e.moe_gate_up_gelu8_dev_q8_rows(
19147 &dev.ptr_row,
19148 &sel_d,
19149 &zq,
19150 &zd,
19151 t,
19152 n_embd,
19153 n_ff_exp,
19154 n_used,
19155 n_expert,
19156 m.gate_exps.qtype,
19157 m.up_exps.qtype,
19158 m.gate_exps.row_bytes,
19159 m.up_exps.row_bytes,
19160 )?
19161 };
19162 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
19163 let mut moe_out = e.uninit(t * n_embd)?;
19164 e.moe_down8_fma_dev_q8_rows_g(
19165 &dev.ptr_row,
19166 &sel_d,
19167 &w_d,
19168 &aq2,
19169 &ad2,
19170 &mut moe_out,
19171 t,
19172 n_ff_exp,
19173 n_embd,
19174 n_used,
19175 n_expert,
19176 m.down_exps.qtype,
19177 m.down_exps.row_bytes,
19178 )?;
19179 return Ok(moe_out);
19180 }
19181
19182 let (sel_all, mut w_all) = Self::moe_route(e, &logits, t, n_expert, n_used)?;
19183 for (i, &sx) in sel_all.iter().enumerate() {
19184 w_all[i] *= bits.per_expert_scale[sx as usize];
19185 }
19186
19187 if t >= PRIME_MIN_T
19191 && m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
19192 && expert_dp4a_supported(m.gate_exps.qtype)
19193 && expert_dp4a_supported(m.up_exps.qtype)
19194 && expert_dp4a_supported(m.down_exps.qtype)
19195 && std::env::var("MEMRA_GEMMA_MOE_PAIRS").as_deref() != Ok("0")
19196 {
19197 let dev = m.dev_exps.as_ref().unwrap();
19198 let n_pairs = t * n_used;
19199 let pair_ex: Vec<i32> = sel_all.iter().map(|&x| x as i32).collect();
19200 let pair_tok: Vec<i32> = (0..n_pairs).map(|p| (p / n_used) as i32).collect();
19201 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
19202 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
19203 let pt = e.htod_i32(&pair_tok)?;
19204 let pw = e.htod(&w_all)?;
19205 let toff = e.htod_i32(&tok_off)?;
19206 let tids = e.htod_i32(&tok_ids)?;
19207 let mut by_ex: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
19208 for p in 0..n_pairs {
19209 by_ex[pair_ex[p] as usize].push(p as i32);
19210 }
19211 let mut ex_ids: Vec<i32> = Vec::new();
19212 let mut ex_off: Vec<i32> = vec![0];
19213 let mut ex_pairs: Vec<i32> = Vec::with_capacity(n_pairs);
19214 for (ex, list) in by_ex.iter().enumerate() {
19215 if list.is_empty() {
19216 continue;
19217 }
19218 ex_ids.push(ex as i32);
19219 ex_pairs.extend_from_slice(list);
19220 ex_off.push(ex_pairs.len() as i32);
19221 }
19222 let n_active = ex_ids.len();
19223 let exi = e.htod_i32(&ex_ids)?;
19224 let exo = e.htod_i32(&ex_off)?;
19225 let exp_d = e.htod_i32(&ex_pairs)?;
19226 if crate::moe_f16g_gemma_on()
19234 && f16g_proj_ok(m.gate_exps.qtype, n_embd)
19235 && f16g_proj_ok(m.up_exps.qtype, n_embd)
19236 && f16g_proj_ok(m.down_exps.qtype, n_ff_exp)
19237 {
19238 let csr_tok: Vec<i32> = ex_pairs.iter().map(|&p| p / n_used as i32).collect();
19239 let csr_tok_d = e.htod_i32(&csr_tok)?;
19240 let (z_f16, z_s) = e.moe_f16g_act(moe_in, Some(&csr_tok_d), n_embd, n_pairs)?;
19241 let g_csr = e.moe_f16_grouped(
19242 &dev.ptr_row,
19243 0,
19244 n_expert,
19245 &exi,
19246 &ex_off,
19247 &exo,
19248 &z_f16,
19249 &z_s,
19250 n_embd,
19251 n_ff_exp,
19252 n_active,
19253 n_pairs,
19254 m.gate_exps.qtype,
19255 m.gate_exps.row_bytes,
19256 )?;
19257 let u_csr = e.moe_f16_grouped(
19258 &dev.ptr_row,
19259 1,
19260 n_expert,
19261 &exi,
19262 &ex_off,
19263 &exo,
19264 &z_f16,
19265 &z_s,
19266 n_embd,
19267 n_ff_exp,
19268 n_active,
19269 n_pairs,
19270 m.up_exps.qtype,
19271 m.up_exps.row_bytes,
19272 )?;
19273 let act_csr = e.moe_pairs_gelu_mul(&g_csr, &u_csr, n_pairs * n_ff_exp)?;
19274 let (a_f16, a_s) = e.moe_f16g_act(&act_csr, None, n_ff_exp, n_pairs)?;
19275 let d_csr = e.moe_f16_grouped(
19276 &dev.ptr_row,
19277 2,
19278 n_expert,
19279 &exi,
19280 &ex_off,
19281 &exo,
19282 &a_f16,
19283 &a_s,
19284 n_ff_exp,
19285 n_embd,
19286 n_active,
19287 n_pairs,
19288 m.down_exps.qtype,
19289 m.down_exps.row_bytes,
19290 )?;
19291 let y_down = e.rows_permute(&d_csr, &exp_d, n_pairs, n_embd)?;
19292 let mut moe_out = e.uninit(t * n_embd)?;
19293 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
19294 if std::env::var("MEMRA_F16G_DEBUG").is_ok() {
19295 let scan = |v: &[f32]| v.iter().filter(|x| !x.is_finite()).count();
19296 let (yd, mo) = (e.dtoh(&y_down)?, e.dtoh(&moe_out)?);
19297 eprintln!(
19298 "[f16g-debug] post-permute bad={} post-scatter bad={}",
19299 scan(&yd),
19300 scan(&mo)
19301 );
19302 }
19303 return Ok(moe_out);
19304 }
19305 let mma = n_embd.is_multiple_of(256)
19308 && std::env::var("MEMRA_GEMMA_MOE_MMA").as_deref() != Ok("0");
19309 let (gate, up) = if mma {
19310 let z_scr = e.mmq_iq_quantize_act(moe_in, n_embd, t)?;
19311 (
19312 e.mmq_iq_experts(
19313 &dev.ptr_row,
19314 0,
19315 n_expert,
19316 &exi,
19317 &exo,
19318 &exp_d,
19319 &pt,
19320 &z_scr,
19321 n_embd,
19322 n_ff_exp,
19323 n_active,
19324 n_pairs,
19325 t,
19326 m.gate_exps.qtype,
19327 m.gate_exps.row_bytes,
19328 )?,
19329 e.mmq_iq_experts(
19330 &dev.ptr_row,
19331 1,
19332 n_expert,
19333 &exi,
19334 &exo,
19335 &exp_d,
19336 &pt,
19337 &z_scr,
19338 n_embd,
19339 n_ff_exp,
19340 n_active,
19341 n_pairs,
19342 t,
19343 m.up_exps.qtype,
19344 m.up_exps.row_bytes,
19345 )?,
19346 )
19347 } else {
19348 let (zq, zd) = e.quantize_q8_1(moe_in, t, n_embd)?;
19349 (
19350 e.moe_pairs_matvec_q8_dec(
19351 &dev.ptr_row,
19352 0,
19353 &exi,
19354 &exo,
19355 &exp_d,
19356 &pt,
19357 &zq,
19358 &zd,
19359 n_embd,
19360 n_ff_exp,
19361 n_expert,
19362 n_active,
19363 n_pairs,
19364 m.gate_exps.qtype,
19365 m.gate_exps.row_bytes,
19366 )?,
19367 e.moe_pairs_matvec_q8_dec(
19368 &dev.ptr_row,
19369 1,
19370 &exi,
19371 &exo,
19372 &exp_d,
19373 &pt,
19374 &zq,
19375 &zd,
19376 n_embd,
19377 n_ff_exp,
19378 n_expert,
19379 n_active,
19380 n_pairs,
19381 m.up_exps.qtype,
19382 m.up_exps.row_bytes,
19383 )?,
19384 )
19385 };
19386 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
19387 let pself = e.htod_i32(&pair_self)?;
19388 let y_down = if mma {
19400 let in_pad = n_ff_exp.div_ceil(256) * 256;
19401 let a_scr = if crate::moe_fuse_actq_on() {
19402 e.mmq_iq_fused_act_quant(&gate, &up, n_ff_exp, n_pairs, 1)?
19403 } else {
19404 let act = e.moe_pairs_gelu_mul(&gate, &up, n_pairs * n_ff_exp)?;
19405 e.mmq_iq_quantize_act(&act, n_ff_exp, n_pairs)?
19406 };
19407 e.mmq_iq_experts(
19408 &dev.ptr_row,
19409 2,
19410 n_expert,
19411 &exi,
19412 &exo,
19413 &exp_d,
19414 &pself,
19415 &a_scr,
19416 in_pad,
19417 n_embd,
19418 n_active,
19419 n_pairs,
19420 n_pairs,
19421 m.down_exps.qtype,
19422 m.down_exps.row_bytes,
19423 )?
19424 } else {
19425 let act = e.moe_pairs_gelu_mul(&gate, &up, n_pairs * n_ff_exp)?;
19426 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
19427 e.moe_pairs_matvec_q8_dec(
19428 &dev.ptr_row,
19429 2,
19430 &exi,
19431 &exo,
19432 &exp_d,
19433 &pself,
19434 &aq2,
19435 &ad2,
19436 n_ff_exp,
19437 n_embd,
19438 n_expert,
19439 n_active,
19440 n_pairs,
19441 m.down_exps.qtype,
19442 m.down_exps.row_bytes,
19443 )?
19444 };
19445 let mut moe_out = e.uninit(t * n_embd)?;
19446 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
19447 return Ok(moe_out);
19448 }
19449
19450 let g_len = m.gate_exps.expert_stride;
19451 let u_len = m.up_exps.expert_stride;
19452 let d_len = m.down_exps.expert_stride;
19453 let dev = m.dev_exps.as_ref().filter(|d| !d.gu_il);
19457 let (mut sg, mut su, mut sd) = if dev.is_some() {
19458 (None, None, None)
19459 } else {
19460 (
19461 Some(e.alloc_u8_uninit(g_len)?),
19462 Some(e.alloc_u8_uninit(u_len)?),
19463 Some(e.alloc_u8_uninit(d_len)?),
19464 )
19465 };
19466 let mut moe_out = e.zeros(t * n_embd)?;
19467 for tok in 0..t {
19468 let sel = &sel_all[tok * n_used..(tok + 1) * n_used];
19469 let w = &w_all[tok * n_used..(tok + 1) * n_used];
19470 let zt = moe_in.slice(tok * n_embd..(tok + 1) * n_embd);
19471 for (j, &ex) in sel.iter().enumerate() {
19472 let ex = ex as usize;
19473 let gate = match dev {
19474 Some(d) => m.qmatvec_view(
19475 e,
19476 &d.gate,
19477 ex * g_len..(ex + 1) * g_len,
19478 &zt,
19479 1,
19480 m.gate_exps.in_f,
19481 m.gate_exps.out_f,
19482 m.gate_exps.qtype,
19483 m.gate_exps.row_bytes,
19484 )?,
19485 None => {
19486 let sg = sg.as_mut().unwrap();
19487 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
19488 m.qmatvec_view(
19489 e,
19490 sg,
19491 0..g_len,
19492 &zt,
19493 1,
19494 m.gate_exps.in_f,
19495 m.gate_exps.out_f,
19496 m.gate_exps.qtype,
19497 m.gate_exps.row_bytes,
19498 )?
19499 }
19500 };
19501 let up = match dev {
19502 Some(d) => m.qmatvec_view(
19503 e,
19504 &d.up,
19505 ex * u_len..(ex + 1) * u_len,
19506 &zt,
19507 1,
19508 m.up_exps.in_f,
19509 m.up_exps.out_f,
19510 m.up_exps.qtype,
19511 m.up_exps.row_bytes,
19512 )?,
19513 None => {
19514 let su = su.as_mut().unwrap();
19515 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
19516 m.qmatvec_view(
19517 e,
19518 su,
19519 0..u_len,
19520 &zt,
19521 1,
19522 m.up_exps.in_f,
19523 m.up_exps.out_f,
19524 m.up_exps.qtype,
19525 m.up_exps.row_bytes,
19526 )?
19527 }
19528 };
19529 let mut act = e.uninit(n_ff_exp)?;
19530 e.gelu_tanh_mul(&gate, &up, &mut act, n_ff_exp)?;
19531 let actv = act.slice(0..n_ff_exp);
19532 let y = match dev {
19533 Some(d) => m.qmatvec_view(
19534 e,
19535 &d.down,
19536 ex * d_len..(ex + 1) * d_len,
19537 &actv,
19538 1,
19539 m.down_exps.in_f,
19540 m.down_exps.out_f,
19541 m.down_exps.qtype,
19542 m.down_exps.row_bytes,
19543 )?,
19544 None => {
19545 let sd = sd.as_mut().unwrap();
19546 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
19547 m.qmatvec_view(
19548 e,
19549 sd,
19550 0..d_len,
19551 &actv,
19552 1,
19553 m.down_exps.in_f,
19554 m.down_exps.out_f,
19555 m.down_exps.qtype,
19556 m.down_exps.row_bytes,
19557 )?
19558 }
19559 };
19560 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
19561 e.axpy_into(&y, w[j], &mut dst, n_embd)?;
19562 }
19563 }
19564 Ok(moe_out)
19565 }
19566
19567 fn gemma4_layer(
19569 &self,
19570 e: &Engine,
19571 il: usize,
19572 layer: &crate::hybrid::HybridLayer,
19573 x: &CudaSlice<f32>,
19574 pos_d: &CudaSlice<i32>,
19575 t: usize,
19576 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
19577 let n_embd = self.cfg.n_embd as usize;
19578 let eps = self.cfg.rms_eps;
19579
19580 let mut h = e.zeros(t * n_embd)?;
19581 e.rms_norm(x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
19582 let Mixer::Full(fa) = &layer.mixer else {
19583 panic!("gemma4 layer {il} not full-attn")
19584 };
19585 let o = self.gemma4_attn(e, fa, il, &h, pos_d, t)?;
19586 let mut cur = e.zeros(t * n_embd)?;
19588 e.rms_norm(
19589 &o,
19590 layer.post_attn_norm.float_data(),
19591 &mut cur,
19592 n_embd,
19593 t,
19594 eps,
19595 )?;
19596 self.gemma4_layer_tail_add(e, layer, &cur, x, t)
19597 }
19598
19599 fn gemma4_layer_tail_add(
19603 &self,
19604 e: &Engine,
19605 layer: &crate::hybrid::HybridLayer,
19606 cur: &CudaSlice<f32>,
19607 x: &CudaSlice<f32>,
19608 t: usize,
19609 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
19610 Ok(self.gemma4_layer_tail_add_n(e, layer, cur, x, t, None)?.0)
19611 }
19612
19613 #[allow(clippy::type_complexity)] fn gemma4_layer_tail_add_n(
19617 &self,
19618 e: &Engine,
19619 layer: &crate::hybrid::HybridLayer,
19620 cur: &CudaSlice<f32>,
19621 x: &CudaSlice<f32>,
19622 t: usize,
19623 next_norm: Option<&CudaSlice<f32>>,
19624 ) -> Result<(CudaSlice<f32>, Option<CudaSlice<f32>>), Box<dyn std::error::Error>> {
19625 let n_embd = self.cfg.n_embd as usize;
19626 let bits = layer.gemma4.as_ref().unwrap();
19627 let (sn, attn_out) = self.gemma4_layer_tail_core(e, layer, cur, x, t)?;
19628 let mut xn = e.uninit(t * n_embd)?;
19629 match next_norm {
19630 Some(w) => {
19631 let mut hn = e.uninit(t * n_embd)?;
19632 e.add_scale_rms_norm(
19633 &sn,
19634 &attn_out,
19635 bits.layer_scale,
19636 w,
19637 &mut xn,
19638 &mut hn,
19639 n_embd,
19640 t,
19641 self.cfg.rms_eps,
19642 )?;
19643 Ok((xn, Some(hn)))
19644 }
19645 None => {
19646 e.add_scale(&sn, &attn_out, bits.layer_scale, &mut xn, t * n_embd)?;
19647 Ok((xn, None))
19648 }
19649 }
19650 }
19651
19652 fn gemma4_layer_tail_core(
19655 &self,
19656 e: &Engine,
19657 layer: &crate::hybrid::HybridLayer,
19658 cur: &CudaSlice<f32>,
19659 x: &CudaSlice<f32>,
19660 t: usize,
19661 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
19662 self.gemma4_layer_tail_core_pn(e, layer, cur, x, t, None, false)
19663 }
19664
19665 #[allow(clippy::too_many_arguments)] fn gemma4_layer_tail_core_pn(
19673 &self,
19674 e: &Engine,
19675 layer: &crate::hybrid::HybridLayer,
19676 cur: &CudaSlice<f32>,
19677 x: &CudaSlice<f32>,
19678 t: usize,
19679 pre_norm: Option<&CudaSlice<f32>>,
19680 defer_post_norm: bool,
19681 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
19682 let n_embd = self.cfg.n_embd as usize;
19683 let eps = self.cfg.rms_eps;
19684 let bits = layer.gemma4.as_ref().unwrap();
19685
19686 let Some(mbits) = bits.moe_bits.as_ref() else {
19689 let crate::hybrid::Ffn::Dense {
19690 ffn_gate,
19691 ffn_up,
19692 ffn_down,
19693 } = &layer.ffn
19694 else {
19695 panic!("gemma4 dense layer without Dense ffn")
19696 };
19697 let mut attn_out = e.uninit(t * n_embd)?;
19698 let mut zsh = e.uninit(t * n_embd)?;
19699 let mut zpair: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
19702 match pre_norm {
19703 Some(wa) if t == 1 => {
19704 zpair = Some(e.rms_pre_add_rms_norm_q8z(
19705 cur,
19706 wa,
19707 x,
19708 bits.ffn_norm.float_data(),
19709 &mut attn_out,
19710 &mut zsh,
19711 n_embd,
19712 t,
19713 eps,
19714 )?);
19715 }
19716 Some(wa) => e.rms_pre_add_rms_norm(
19717 cur,
19718 wa,
19719 x,
19720 bits.ffn_norm.float_data(),
19721 &mut attn_out,
19722 &mut zsh,
19723 n_embd,
19724 t,
19725 eps,
19726 )?,
19727 None => e.add_rms_norm(
19728 cur,
19729 x,
19730 bits.ffn_norm.float_data(),
19731 &mut attn_out,
19732 &mut zsh,
19733 n_embd,
19734 t,
19735 eps,
19736 )?,
19737 }
19738 let n_ff = ffn_gate.out_features();
19739 let (gate, up) = if t == 1 {
19745 let (zq, zd) = match zpair {
19746 Some(p) => p,
19747 None => e.quantize_q8_1(&zsh, 1, n_embd)?,
19748 };
19749 match e.matmul_q4_fused2(ffn_gate, ffn_up, &zq, &zd)? {
19750 Some(p) => p,
19751 None => match e.matmul_nvfp4_fused2(ffn_gate, ffn_up, &zq, &zd, 1)? {
19753 Some(p) => p,
19754 None => (
19755 e.matmul_pre(ffn_gate, &zq, &zd, &zsh, 1)?,
19756 e.matmul_pre(ffn_up, &zq, &zd, &zsh, 1)?,
19757 ),
19758 },
19759 }
19760 } else {
19761 static F2B: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
19766 let f2b = *F2B.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
19767 let fused = if f2b {
19768 let (zq, zd) = e.quantize_q8_1(&zsh, t, n_embd)?;
19769 e.matmul_q4_fused2_batched(ffn_gate, ffn_up, &zq, &zd, t)?
19770 } else {
19771 None
19772 };
19773 match fused {
19774 Some(p) => p,
19775 None => {
19776 e.mmq_act_begin();
19778 (e.matmul(ffn_gate, &zsh, t)?, e.matmul(ffn_up, &zsh, t)?)
19779 }
19780 }
19781 };
19782 let mut act = e.uninit(t * n_ff)?;
19783 let f0 = if e.uses_q8_1_fast(ffn_down) {
19786 let upv = e.view(&up, t * n_ff);
19787 let up_all = upv.slice(0..t * n_ff);
19788 let (aq, ad) = e.gelu_tanh_mul_q8_1(&gate, &up_all, &mut act, n_ff, t)?;
19789 e.matmul_pre(ffn_down, &aq, &ad, &act, t)?
19790 } else {
19791 e.gelu_tanh_mul(&gate, &up, &mut act, t * n_ff)?;
19792 e.matmul(ffn_down, &act, t)?
19793 };
19794 if defer_post_norm {
19795 return Ok((f0, attn_out));
19796 }
19797 let mut sn = e.uninit(t * n_embd)?;
19798 e.rms_norm(
19799 &f0,
19800 bits.post_ffw_norm.float_data(),
19801 &mut sn,
19802 n_embd,
19803 t,
19804 eps,
19805 )?;
19806 return Ok((sn, attn_out));
19807 };
19808
19809 assert!(pre_norm.is_none(), "pre-norm fold is dense-entry only");
19810 let mut attn_out = e.uninit(t * n_embd)?;
19815 let mut router_in = e.uninit(t * n_embd)?;
19816 let fast_moe = match &layer.ffn {
19817 crate::hybrid::Ffn::Moe(m) => {
19818 m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
19819 && expert_dp4a_supported(m.gate_exps.qtype)
19820 && expert_dp4a_supported(m.up_exps.qtype)
19821 && expert_dp4a_supported(m.down_exps.qtype)
19822 && std::env::var("MEMRA_GEMMA_MOE_FAST").as_deref() != Ok("0")
19823 }
19824 _ => false,
19825 };
19826 let q8z = t < PRIME_MIN_T && fast_moe;
19827 let (zsh_f32, zsh_q8, moe_q8) = if q8z {
19828 let (z0, m2) = e.add_rms_norm3_q8z(
19829 cur,
19830 x,
19831 bits.ffn_norm.float_data(),
19832 &mbits.router_scale_pre,
19833 mbits.pre_ffw_norm_2.float_data(),
19834 &mut attn_out,
19835 &mut router_in,
19836 n_embd,
19837 t,
19838 eps,
19839 )?;
19840 (None, Some(z0), Some(m2))
19841 } else {
19842 let mut zsh = e.uninit(t * n_embd)?;
19843 let mut moe_in = e.uninit(t * n_embd)?;
19844 e.add_rms_norm3(
19845 cur,
19846 x,
19847 bits.ffn_norm.float_data(),
19848 &mbits.router_scale_pre,
19849 mbits.pre_ffw_norm_2.float_data(),
19850 &mut attn_out,
19851 &mut zsh,
19852 &mut router_in,
19853 &mut moe_in,
19854 n_embd,
19855 t,
19856 eps,
19857 )?;
19858 (Some((zsh, moe_in)), None, None)
19859 };
19860 let attn_out2 = attn_out;
19861 #[allow(unused_variables)]
19862 let attn_out = &attn_out2;
19863 let n_ff = mbits.shared_gate.out_features();
19864 let (gate, up) = if let Some((zq, zd)) = zsh_q8.as_ref() {
19865 if t == 1 {
19866 match e.matmul_q4_fused2(&mbits.shared_gate, &mbits.shared_up, zq, zd)? {
19867 Some(p) => p,
19868 None => match e.matmul_nvfp4_fused2(
19869 &mbits.shared_gate,
19870 &mbits.shared_up,
19871 zq,
19872 zd,
19873 1,
19874 )? {
19875 Some(p) => p,
19876 None => {
19877 let h0 = e.zeros(0)?;
19878 (
19879 e.matmul_pre(&mbits.shared_gate, zq, zd, &h0, 1)?,
19880 e.matmul_pre(&mbits.shared_up, zq, zd, &h0, 1)?,
19881 )
19882 }
19883 },
19884 }
19885 } else {
19886 let h0 = e.zeros(0)?;
19888 (
19889 e.matmul_pre(&mbits.shared_gate, zq, zd, &h0, t)?,
19890 e.matmul_pre(&mbits.shared_up, zq, zd, &h0, t)?,
19891 )
19892 }
19893 } else {
19894 let (zsh, _) = zsh_f32.as_ref().unwrap();
19895 (
19896 e.matmul(&mbits.shared_gate, zsh, t)?,
19897 e.matmul(&mbits.shared_up, zsh, t)?,
19898 )
19899 };
19900 let mut act = e.uninit(t * n_ff)?;
19901 e.gelu_tanh_mul(&gate, &up, &mut act, t * n_ff)?;
19902 let mlp0 = e.matmul(&mbits.shared_down, &act, t)?;
19903 let crate::hybrid::Ffn::Moe(m) = &layer.ffn else {
19904 panic!("gemma4 layer not MoE")
19905 };
19906 let moe0 = match (&moe_q8, &zsh_f32) {
19907 (Some(mq), _) => self.gemma4_moe_q8(e, m, mbits, mq, &router_in, t)?,
19908 (None, Some((_, moe_in))) => self.gemma4_moe(e, m, mbits, moe_in, &router_in, t)?,
19909 _ => unreachable!(),
19910 };
19911 let mut mlp = e.uninit(t * n_embd)?;
19913 let mut moe = e.uninit(t * n_embd)?;
19914 e.rms_norm2x(
19915 &mlp0,
19916 &moe0,
19917 mbits.post_ffw_norm_1.float_data(),
19918 mbits.post_ffw_norm_2.float_data(),
19919 &mut mlp,
19920 &mut moe,
19921 n_embd,
19922 t,
19923 eps,
19924 )?;
19925
19926 let mut sum = e.uninit(t * n_embd)?;
19929 let mut sn = e.uninit(t * n_embd)?;
19930 e.add_rms_norm(
19931 &mlp,
19932 &moe,
19933 bits.post_ffw_norm.float_data(),
19934 &mut sum,
19935 &mut sn,
19936 n_embd,
19937 t,
19938 eps,
19939 )?;
19940 Ok((sn, attn_out2))
19941 }
19942
19943 #[allow(clippy::type_complexity)] pub(crate) fn gemma4_layer_tail_add_nq_pn(
19954 &self,
19955 e: &Engine,
19956 layer: &crate::hybrid::HybridLayer,
19957 o: &CudaSlice<f32>,
19958 x: &CudaSlice<f32>,
19959 t: usize,
19960 next_norm: Option<&CudaSlice<f32>>,
19961 ) -> Result<(CudaSlice<f32>, Option<(CudaSlice<i8>, CudaSlice<f32>)>), Box<dyn std::error::Error>>
19962 {
19963 let n_embd = self.cfg.n_embd as usize;
19964 let eps = self.cfg.rms_eps;
19965 let bits = layer.gemma4.as_ref().unwrap();
19966 if Engine::g4_pnfold_on() && matches!(layer.ffn, crate::hybrid::Ffn::Dense { .. }) {
19967 let (f0, attn_out) = self.gemma4_layer_tail_core_pn(
19968 e,
19969 layer,
19970 o,
19971 x,
19972 t,
19973 Some(layer.post_attn_norm.float_data()),
19974 true,
19975 )?;
19976 let mut xn = e.uninit(t * n_embd)?;
19977 return match next_norm {
19978 Some(w) => {
19979 let pair = e.rms_pre_add_scale_rms_norm_q8_1(
19980 &f0,
19981 bits.post_ffw_norm.float_data(),
19982 &attn_out,
19983 bits.layer_scale,
19984 w,
19985 &mut xn,
19986 n_embd,
19987 t,
19988 eps,
19989 )?;
19990 Ok((xn, Some(pair)))
19991 }
19992 None => {
19993 let mut sn = e.uninit(t * n_embd)?;
19994 e.rms_norm(
19995 &f0,
19996 bits.post_ffw_norm.float_data(),
19997 &mut sn,
19998 n_embd,
19999 t,
20000 eps,
20001 )?;
20002 e.add_scale(&sn, &attn_out, bits.layer_scale, &mut xn, t * n_embd)?;
20003 Ok((xn, None))
20004 }
20005 };
20006 }
20007 let mut cur = e.uninit(t * n_embd)?;
20008 e.rms_norm(
20009 o,
20010 layer.post_attn_norm.float_data(),
20011 &mut cur,
20012 n_embd,
20013 t,
20014 eps,
20015 )?;
20016 self.gemma4_layer_tail_add_nq(e, layer, &cur, x, t, next_norm)
20017 }
20018
20019 #[allow(clippy::type_complexity)] pub(crate) fn gemma4_layer_tail_add_nq(
20021 &self,
20022 e: &Engine,
20023 layer: &crate::hybrid::HybridLayer,
20024 cur: &CudaSlice<f32>,
20025 x: &CudaSlice<f32>,
20026 t: usize,
20027 next_norm: Option<&CudaSlice<f32>>,
20028 ) -> Result<(CudaSlice<f32>, Option<(CudaSlice<i8>, CudaSlice<f32>)>), Box<dyn std::error::Error>>
20029 {
20030 let n_embd = self.cfg.n_embd as usize;
20031 let bits = layer.gemma4.as_ref().unwrap();
20032 let (sn, attn_out) = self.gemma4_layer_tail_core(e, layer, cur, x, t)?;
20033 let mut xn = e.uninit(t * n_embd)?;
20034 match next_norm {
20035 Some(w) => {
20036 let pair = e.add_scale_rms_norm_q8_1(
20037 &sn,
20038 &attn_out,
20039 bits.layer_scale,
20040 w,
20041 &mut xn,
20042 n_embd,
20043 t,
20044 self.cfg.rms_eps,
20045 )?;
20046 Ok((xn, Some(pair)))
20047 }
20048 None => {
20049 e.add_scale(&sn, &attn_out, bits.layer_scale, &mut xn, t * n_embd)?;
20050 Ok((xn, None))
20051 }
20052 }
20053 }
20054
20055 fn gemma4_forward(
20058 &self,
20059 e: &Engine,
20060 tokens: &[u32],
20061 last_only: bool,
20062 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
20063 if self.is_gemma4_e4b() {
20066 return self.gemma4_e4b_forward(e, tokens, last_only);
20067 }
20068 let n_embd = self.cfg.n_embd as usize;
20069 let t = tokens.len();
20070 let pos: Vec<i32> = (0..t as i32).collect();
20071 let pos_d = e.htod_i32(&pos)?;
20072
20073 let mut x = self.embed(e, tokens)?;
20074 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
20075 let probe = std::env::var("MEMRA_GEMMA_PROBE").is_ok();
20078 let stat =
20079 |e: &Engine, x: &CudaSlice<f32>, tag: &str| -> Result<(), Box<dyn std::error::Error>> {
20080 let h = e.dtoh(x)?;
20081 let bad = h.iter().filter(|v| !v.is_finite()).count();
20082 let mx = h
20083 .iter()
20084 .filter(|v| v.is_finite())
20085 .fold(0.0f32, |m, v| m.max(v.abs()));
20086 eprintln!(
20087 "[gemma-probe] {tag}: tok0_first3={:?} bad={bad} max={mx:.3e}",
20088 &h[..3]
20089 );
20090 Ok(())
20091 };
20092 if probe {
20093 stat(e, &x, "embed")?;
20094 }
20095 for (il, layer) in self.layers.iter().enumerate() {
20096 x = self.gemma4_layer(e, il, layer, &x, &pos_d, t)?;
20097 if probe {
20098 stat(e, &x, &format!("L{il}"))?;
20099 }
20100 }
20101 let mut hn = e.zeros(t * n_embd)?;
20102 e.rms_norm(
20103 &x,
20104 self.output_norm.float_data(),
20105 &mut hn,
20106 n_embd,
20107 t,
20108 self.cfg.rms_eps,
20109 )?;
20110 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
20111 let n_vocab = self.output.out_features();
20112 let logits = if last_only {
20113 let hv = e.view(&hn, t * n_embd);
20114 let last_row = hv.slice((t - 1) * n_embd..t * n_embd);
20115 let mut hlast = e.zeros(n_embd)?;
20116 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
20117 let mut ld = e.matmul(&self.output, &hlast, 1)?;
20118 e.softcap(&mut ld, cap, n_vocab)?;
20119 self.gemma4_suppress(e, &mut ld, 1)?;
20120 e.dtoh(&ld)?
20121 } else {
20122 let mut ld = e.matmul(&self.output, &hn, t)?;
20123 e.softcap(&mut ld, cap, t * n_vocab)?;
20124 self.gemma4_suppress(e, &mut ld, t)?;
20125 e.dtoh(&ld)?
20126 };
20127 Ok(logits)
20128 }
20129
20130 #[allow(clippy::type_complexity)] pub(crate) fn gemma4_prime(
20136 &self,
20137 e: &Engine,
20138 tokens: &[u32],
20139 cache: &mut Cache,
20140 overlay: Option<&crate::vision::EmbedOverlay>,
20141 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
20142 if cache.pos != 0 {
20147 return Err(
20148 "gemma4 prime v0 is fresh-prompt only (no continuation/chunked prime) \
20149 — prime the full prompt in one call or decode tokenwise"
20150 .into(),
20151 );
20152 }
20153 let n_embd = self.cfg.n_embd as usize;
20154 let eps = self.cfg.rms_eps;
20155 let t = tokens.len();
20156 let pos: Vec<i32> = (0..t as i32).collect();
20157 let pos_d = e.htod_i32(&pos)?;
20158 let mut x = self.embed(e, tokens)?;
20159 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
20160 let island: Option<CudaSlice<i32>> = match overlay {
20167 Some(ov) => {
20168 ov.require_resident(e)?;
20172 let mut span_id = vec![-1i32; t];
20173 for (i, &(pos, row_off, n_rows)) in ov.spans.iter().enumerate() {
20174 if pos + n_rows > t {
20175 return Err(format!(
20176 "gemma4 overlay span {i} [{pos}, {}) exceeds the prompt ({t})",
20177 pos + n_rows
20178 )
20179 .into());
20180 }
20181 let view = ov.rows.slice(row_off * n_embd..(row_off + n_rows) * n_embd);
20182 e.copy_view_into(&mut x, pos * n_embd, &view, n_rows * n_embd)?;
20183 for s in span_id.iter_mut().skip(pos).take(n_rows) {
20184 *s = i as i32;
20185 }
20186 }
20187 if std::env::var("MEMRA_GV_FORCE_CAUSAL").as_deref() == Ok("1") {
20191 None
20192 } else {
20193 Some(e.htod_i32(&span_id)?)
20194 }
20195 }
20196 None => None,
20197 };
20198 for (il, layer) in self.layers.iter().enumerate() {
20199 let mut h = e.zeros(t * n_embd)?;
20200 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
20201 let Mixer::Full(fa) = &layer.mixer else {
20202 panic!("gemma4 layer not full-attn")
20203 };
20204 let trace = il == 0 && std::env::var("MEMRA_G4_PRIME_TRACE").as_deref() == Ok("1");
20205 if trace {
20206 let v = e.dtoh(&h)?;
20207 let nan = v.iter().filter(|x| x.is_nan()).count();
20208 eprintln!("[g4-prime-trace] L0 post-attn_norm: nan={nan}/{}", v.len());
20209 }
20210 let o =
20211 self.gemma4_attn_prime(e, fa, il, &h, &pos_d, t, Some(cache), island.as_ref())?;
20212 if trace {
20213 let v = e.dtoh(&o)?;
20214 let nan = v.iter().filter(|x| x.is_nan()).count();
20215 eprintln!("[g4-prime-trace] L0 post-attn: nan={nan}/{}", v.len());
20216 }
20217 let mut cur = e.zeros(t * n_embd)?;
20218 e.rms_norm(
20219 &o,
20220 layer.post_attn_norm.float_data(),
20221 &mut cur,
20222 n_embd,
20223 t,
20224 eps,
20225 )?;
20226 x = self.gemma4_layer_tail_add(e, layer, &cur, &x, t)?;
20227 self.dflash_tap(e, cache, il, &x, t)?;
20228 if std::env::var("MEMRA_G4_PRIME_TRACE").as_deref() == Ok("1") {
20230 let h = e.dtoh(&x)?;
20231 let nan = h.iter().filter(|v| v.is_nan()).count();
20232 let amax = h.iter().fold(0f32, |a, v| a.max(v.abs()));
20233 eprintln!(
20234 "[g4-prime-trace] layer {il}: nan={nan}/{} amax={amax:.3}",
20235 h.len()
20236 );
20237 if nan > 0 {
20238 return Err(format!("g4-prime-trace: first NaN at layer {il}").into());
20239 }
20240 }
20241 }
20242 cache.pos += t;
20243 let hiddens = e.clone_dtod(&x)?;
20244 let xv = e.view(&x, t * n_embd);
20245 let last_row = xv.slice((t - 1) * n_embd..t * n_embd);
20246 let mut h_seed = e.zeros(n_embd)?;
20247 e.copy_view_into(&mut h_seed, 0, &last_row, n_embd)?;
20248 let mut hn = e.uninit(n_embd)?;
20249 e.rms_norm(
20250 &h_seed,
20251 self.output_norm.float_data(),
20252 &mut hn,
20253 n_embd,
20254 1,
20255 eps,
20256 )?;
20257 let mut ld = e.matmul(&self.output, &hn, 1)?;
20258 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
20259 e.softcap(&mut ld, cap, self.output.out_features())?;
20260 self.gemma4_suppress(e, &mut ld, 1)?;
20261 let logits = e.dtoh(&ld)?;
20262 Ok((logits, h_seed, hiddens))
20263 }
20264
20265 #[allow(clippy::too_many_arguments)] fn gemma4_decode_attn(
20271 &self,
20272 e: &Engine,
20273 fa: &crate::hybrid::FullAttnLayer,
20274 il: usize,
20275 hq: &CudaSlice<i8>,
20276 hdq: &CudaSlice<f32>,
20277 pos_d: &CudaSlice<i32>,
20278 cache: &mut Cache,
20279 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
20280 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
20281 let eps = self.cfg.rms_eps;
20282 let aux = self.gemma4_aux.as_ref().unwrap();
20283 let ones = aux.ones(e);
20284 #[cfg(debug_assertions)]
20285 crate::debug_assert_tensor_stream_device(ones, &e.stream(), "gemma4_decode_attn.ones");
20286 let (hq, hdq) = (hq, hdq);
20287 let h0 = e.zeros(0)?;
20288 let h = &h0;
20289 let (q0, k0, v0) = if swa {
20290 match e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)? {
20291 Some(t3) => t3,
20292 None => match e.matmul_nvfp4_fused2(&fa.wq, &fa.wk, hq, hdq, 1)? {
20295 Some((q0, k0)) => {
20296 let v0 = e.matmul_pre(&fa.wv, hq, hdq, h, 1)?;
20297 (q0, k0, v0)
20298 }
20299 None => (
20300 e.matmul_pre(&fa.wq, hq, hdq, h, 1)?,
20301 e.matmul_pre(&fa.wk, hq, hdq, h, 1)?,
20302 e.matmul_pre(&fa.wv, hq, hdq, h, 1)?,
20303 ),
20304 },
20305 }
20306 } else {
20307 let (q0, k0) = match e.matmul_q4_fused2(&fa.wq, &fa.wk, hq, hdq)? {
20308 Some(p) => p,
20309 None => match e.matmul_nvfp4_fused2(&fa.wq, &fa.wk, hq, hdq, 1)? {
20310 Some(p) => p,
20311 None => (
20312 e.matmul_pre(&fa.wq, hq, hdq, h, 1)?,
20313 e.matmul_pre(&fa.wk, hq, hdq, h, 1)?,
20314 ),
20315 },
20316 };
20317 let v0 = e.clone_dtod(&k0)?;
20318 (q0, k0, v0)
20319 };
20320 let mut q = e.uninit(nh * hd)?;
20321 let mut k = e.uninit(nkv * hd)?;
20322 let mut v = e.uninit(nkv * hd)?;
20323 let ff = if swa {
20326 None
20327 } else {
20328 Some(
20329 aux.rope_freqs(e)
20330 .expect("gemma4 global rope needs rope_freqs.weight"),
20331 )
20332 };
20333 #[cfg(debug_assertions)]
20334 if let Some(ff) = ff {
20335 crate::debug_assert_tensor_stream_device(
20336 ff,
20337 &e.stream(),
20338 "gemma4_decode_attn.rope_freqs",
20339 );
20340 }
20341 let kvl = cache.kv[il].as_mut().unwrap();
20342 let kv_fp8 = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
20343 if crate::Engine::qkv_append_on() {
20344 e.rms_norm_qkv_rope_append(
20348 &q0,
20349 &k0,
20350 &v0,
20351 fa.q_norm.float_data(),
20352 fa.k_norm.float_data(),
20353 ones,
20354 &mut q,
20355 &mut k,
20356 &mut v,
20357 hd,
20358 self.gemma4_rope_dims(il),
20359 nh,
20360 nkv,
20361 pos_d,
20362 nh,
20363 nkv,
20364 base,
20365 1.0,
20366 ff,
20367 eps,
20368 &mut kvl.k,
20369 &mut kvl.v,
20370 kvl.len,
20371 kvl.k_tok_bytes,
20372 kvl.v_tok_bytes,
20373 kv_fp8,
20374 )?;
20375 } else {
20376 e.rms_norm_qkv_rope(
20377 &q0,
20378 &k0,
20379 &v0,
20380 fa.q_norm.float_data(),
20381 fa.k_norm.float_data(),
20382 ones,
20383 &mut q,
20384 &mut k,
20385 &mut v,
20386 hd,
20387 self.gemma4_rope_dims(il),
20388 nh,
20389 nkv,
20390 pos_d,
20391 nh,
20392 nkv,
20393 base,
20394 1.0,
20395 ff,
20396 eps,
20397 )?;
20398 e.append_kv_quantized(
20399 &k,
20400 &v,
20401 &mut kvl.k,
20402 &mut kvl.v,
20403 kvl.len,
20404 kvl.kv_dim_k,
20405 kvl.kv_dim_v,
20406 kvl.k_tok_bytes,
20407 kvl.v_tok_bytes,
20408 kv_fp8,
20409 )?;
20410 }
20411 kvl.len += 1;
20412 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
20416 let mut attn = e.uninit(nh * hd)?;
20417 if !swa
20419 && hd == 512
20420 && kvl.len >= crate::fa512_min_tkv()
20421 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
20422 {
20423 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
20424 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
20425 let base = kvl.len as i32;
20427 e.i32_set_k(&mut kvl.len_d, base)?;
20428 e.fa_decode_rows(
20429 &q,
20430 &kp,
20431 &vp,
20432 &mut attn,
20433 hd,
20434 nh,
20435 nkv,
20436 kvl.len - 1,
20437 1,
20438 scale,
20439 kvl.k_tok_bytes,
20440 kvl.v_tok_bytes,
20441 Some((&kvl.len_d, -1)),
20442 false,
20443 false,
20444 None,
20445 )?;
20446 return e.matmul(&fa.wo, &attn, 1);
20447 }
20448 if swa
20450 && kvl.len > win
20451 && hd == 256
20452 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
20453 {
20454 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
20455 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
20456 let base = kvl.len as i32;
20457 e.i32_set_k(&mut kvl.len_d, base)?;
20458 e.fa_decode_rows_w(
20459 &q,
20460 &kp,
20461 &vp,
20462 &mut attn,
20463 hd,
20464 nh,
20465 nkv,
20466 &kvl.len_d,
20467 -1,
20468 1,
20469 scale,
20470 win,
20471 kvl.k_tok_bytes,
20472 kvl.v_tok_bytes,
20473 None,
20474 )?;
20475 return e.matmul(&fa.wo, &attn, 1);
20476 }
20477 let (off_tok, t_kv) = if swa && kvl.len > win {
20478 (kvl.len - win, win)
20479 } else {
20480 (0, kvl.len)
20481 };
20482 let k_view = e.view_u8_range(
20483 &kvl.k,
20484 off_tok * kvl.k_tok_bytes,
20485 (off_tok + t_kv) * kvl.k_tok_bytes,
20486 );
20487 let v_view = e.view_u8_range(
20488 &kvl.v,
20489 off_tok * kvl.v_tok_bytes,
20490 (off_tok + t_kv) * kvl.v_tok_bytes,
20491 );
20492 e.fa_decode_kvmod(
20493 &q,
20494 &k_view,
20495 &v_view,
20496 &mut attn,
20497 hd,
20498 nh,
20499 nkv,
20500 t_kv,
20501 scale,
20502 kvl.k_tok_bytes,
20503 kvl.v_tok_bytes,
20504 swa && crate::Engine::wkv_on(),
20505 )?;
20506 e.matmul(&fa.wo, &attn, 1)
20507 }
20508
20509 #[allow(clippy::too_many_arguments)]
20516 pub fn gemma4_decode_step_dc(
20517 &self,
20518 e: &Engine,
20519 token_d: &CudaSlice<u32>,
20520 pos_d: &mut CudaSlice<i32>,
20521 embd_gpu: &CudaSlice<u8>,
20522 embd_qt: i32,
20523 embd_rb: usize,
20524 cache: &mut Cache,
20525 n_vocab: usize,
20526 cap_bucket_max: Option<(usize, usize)>,
20527 ) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
20528 let mut tok_out = e.stream().alloc_zeros::<u32>(1)?;
20529 self.gemma4_decode_step_dc_into(
20530 e,
20531 token_d,
20532 pos_d,
20533 embd_gpu,
20534 embd_qt,
20535 embd_rb,
20536 cache,
20537 n_vocab,
20538 cap_bucket_max,
20539 &mut tok_out,
20540 )?;
20541 Ok(tok_out)
20542 }
20543
20544 #[allow(clippy::too_many_arguments)]
20547 pub fn gemma4_decode_step_dc_into(
20548 &self,
20549 e: &Engine,
20550 token_d: &CudaSlice<u32>,
20551 pos_d: &mut CudaSlice<i32>,
20552 embd_gpu: &CudaSlice<u8>,
20553 embd_qt: i32,
20554 embd_rb: usize,
20555 cache: &mut Cache,
20556 n_vocab: usize,
20557 cap_bucket_max: Option<(usize, usize)>,
20558 tok_out: &mut CudaSlice<u32>,
20559 ) -> Result<(), Box<dyn std::error::Error>> {
20560 let n_embd = self.cfg.n_embd as usize;
20561 let eps = self.cfg.rms_eps;
20562 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
20563 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
20564 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
20565 let n_layers = self.layers.len();
20566 for (il, layer) in self.layers.iter().enumerate() {
20567 let (hq, hdq) = match h_carry.take() {
20568 Some(p) => p,
20569 None => {
20570 e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, 1, eps)?
20571 }
20572 };
20573 let Mixer::Full(fa) = &layer.mixer else {
20574 panic!("gemma4 layer {il} not full-attn")
20575 };
20576 let o =
20577 self.gemma4_decode_attn_dc(e, fa, il, &hq, &hdq, pos_d, cache, cap_bucket_max)?;
20578 let next_norm = if il + 1 < n_layers {
20579 Some(self.layers[il + 1].attn_norm.float_data())
20580 } else {
20581 None
20582 };
20583 let (xn, hn) = self.gemma4_layer_tail_add_nq_pn(e, layer, &o, &x, 1, next_norm)?;
20584 x = xn;
20585 h_carry = hn;
20586 }
20587 let mut hn = e.uninit(n_embd)?;
20588 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
20589 let mut logits = e.matmul(&self.output, &hn, 1)?;
20590 self.gemma4_suppress(e, &mut logits, 1)?; e.argmax_token_device_into(&logits, tok_out, n_vocab)?;
20592 e.inc_seqlen(pos_d)?;
20593 if cap_bucket_max.is_none() {
20594 cache.pos += 1;
20595 }
20596 Ok(())
20597 }
20598
20599 pub fn g4_dc_slots(&self, e: &Engine) -> Result<G4DcSlots, Box<dyn std::error::Error>> {
20606 let n_embd = self.cfg.n_embd as usize;
20607 let n_vocab = self.output.out_features();
20608 let n_layers = self.layers.len();
20609 let (mut qmax, mut kvmax, mut ffmax) = (0usize, 0usize, 0usize);
20610 for il in 0..n_layers {
20611 let (hd, nkv, nh, _b, _s, _w) = self.gemma4_geom(il);
20612 qmax = qmax.max(nh * hd);
20613 kvmax = kvmax.max(nkv * hd);
20614 if let crate::hybrid::Ffn::Dense { ffn_gate, .. } = &self.layers[il].ffn {
20615 ffmax = ffmax.max(ffn_gate.out_features());
20616 }
20617 }
20618 Ok(G4DcSlots {
20619 x: e.uninit(n_embd)?,
20620 xn: e.uninit(n_embd)?,
20621 cur: e.uninit(n_embd)?,
20622 hq: e.alloc_i8_uninit(n_embd)?,
20623 hd_: e.uninit(n_embd / 32)?,
20624 q0: e.uninit(qmax)?,
20625 k0: e.uninit(kvmax)?,
20626 v0: e.uninit(kvmax)?,
20627 q: e.uninit(qmax)?,
20628 k: e.uninit(kvmax)?,
20629 v: e.uninit(kvmax)?,
20630 attn: e.uninit(qmax)?,
20631 o: e.uninit(n_embd)?,
20632 attn_out: e.uninit(n_embd)?,
20633 zsh: e.uninit(n_embd)?,
20634 zq: e.alloc_i8_uninit(n_embd.max(qmax))?,
20637 zd: e.uninit(n_embd.max(qmax) / 32)?,
20638 gate: e.uninit(ffmax)?,
20639 up: e.uninit(ffmax)?,
20640 act: e.uninit(ffmax)?,
20641 actq: e.alloc_i8_uninit(ffmax)?,
20642 actd: e.uninit(ffmax / 32)?,
20643 f0: e.uninit(n_embd)?,
20644 sn: e.uninit(n_embd)?,
20645 hn: e.uninit(n_embd)?,
20646 logits: e.uninit(n_vocab)?,
20647 })
20648 }
20649
20650 fn g4_matvec_m1_into(
20653 &self,
20654 e: &Engine,
20655 w: &crate::model::GpuTensor,
20656 aq: &CudaSlice<i8>,
20657 ad: &CudaSlice<f32>,
20658 y: &mut CudaSlice<f32>,
20659 ) -> Result<(), Box<dyn std::error::Error>> {
20660 use crate::model::GpuTensor;
20661 let (bytes, qtype, row_bytes, scale, rp) = match w {
20662 GpuTensor::Quant {
20663 bytes,
20664 qtype,
20665 row_bytes,
20666 scale,
20667 rp,
20668 ..
20669 } => (bytes, *qtype, *row_bytes, *scale, *rp),
20670 _ => return Err("g4_matvec_m1_into: non-quant tensor".into()),
20671 };
20672 let (mbytes, mrp) = match w {
20673 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
20674 _ => (bytes, rp),
20675 };
20676 e.qmatvec_mmvq_into(
20677 mbytes,
20678 aq,
20679 ad,
20680 1,
20681 w.in_features(),
20682 w.out_features(),
20683 qtype,
20684 row_bytes,
20685 scale,
20686 mrp,
20687 y,
20688 )
20689 }
20690
20691 #[allow(clippy::too_many_arguments)]
20695 pub fn gemma4_decode_step_dc_slotted(
20696 &self,
20697 e: &Engine,
20698 token_d: &CudaSlice<u32>,
20699 pos_d: &mut CudaSlice<i32>,
20700 embd_gpu: &CudaSlice<u8>,
20701 embd_qt: i32,
20702 embd_rb: usize,
20703 cache: &mut Cache,
20704 n_vocab: usize,
20705 cap_bucket_max: Option<(usize, usize)>,
20706 sl: &mut G4DcSlots,
20707 tok_out: &mut CudaSlice<u32>,
20708 ring: Option<(&mut CudaSlice<u32>, usize)>,
20709 ) -> Result<(), Box<dyn std::error::Error>> {
20710 let n_embd = self.cfg.n_embd as usize;
20711 let eps = self.cfg.rms_eps;
20712 e.embed_gather_device_into(embd_gpu, token_d, &mut sl.x, n_embd, embd_qt, embd_rb)?;
20713 e.scale_inplace(&mut sl.x, (n_embd as f32).sqrt(), n_embd)?;
20714 let n_layers = self.layers.len();
20715 let mut has_carry = false;
20716 for il in 0..n_layers {
20717 if !has_carry {
20718 e.rms_norm_q8_1_into(
20719 &sl.x,
20720 self.layers[il].attn_norm.float_data(),
20721 n_embd,
20722 1,
20723 eps,
20724 &mut sl.hq,
20725 &mut sl.hd_,
20726 )?;
20727 }
20728 has_carry = true;
20729 let layer = &self.layers[il];
20730 let Mixer::Full(fa) = &layer.mixer else {
20731 panic!("gemma4 layer {il} not full-attn")
20732 };
20733 self.gemma4_decode_attn_dc_slotted(e, fa, il, pos_d, cache, cap_bucket_max, sl)?;
20734 if !Engine::g4_pnfold_on() {
20737 e.rms_norm(
20738 &sl.o,
20739 layer.post_attn_norm.float_data(),
20740 &mut sl.cur,
20741 n_embd,
20742 1,
20743 eps,
20744 )?;
20745 }
20746 let next_norm = if il + 1 < n_layers {
20747 Some(self.layers[il + 1].attn_norm.float_data())
20748 } else {
20749 None
20750 };
20751 self.gemma4_layer_tail_slotted(e, layer, next_norm, sl)?;
20752 std::mem::swap(&mut sl.x, &mut sl.xn);
20753 }
20754 e.rms_norm(
20755 &sl.x,
20756 self.output_norm.float_data(),
20757 &mut sl.hn,
20758 n_embd,
20759 1,
20760 eps,
20761 )?;
20762 e.quantize_q8_1_into(&sl.hn, 1, n_embd, &mut sl.zq, &mut sl.zd)?;
20763 {
20765 let (zq, zd) = (&sl.zq, &sl.zd);
20766 let zq = unsafe { &*(zq as *const CudaSlice<i8>) };
20767 let zd = unsafe { &*(zd as *const CudaSlice<f32>) };
20768 self.g4_matvec_m1_into(e, &self.output, zq, zd, &mut sl.logits)?;
20769 }
20770 self.gemma4_suppress(e, &mut sl.logits, 1)?;
20771 e.argmax_token_device_into(&sl.logits, tok_out, n_vocab)?;
20772 if let Some((ring, base)) = ring {
20773 e.plain_tok_ring(tok_out, pos_d, base, ring)?;
20777 }
20778 e.inc_seqlen(pos_d)?;
20779 if cap_bucket_max.is_none() {
20780 cache.pos += 1;
20781 }
20782 Ok(())
20783 }
20784
20785 #[allow(clippy::too_many_arguments)]
20787 fn gemma4_decode_attn_dc_slotted(
20788 &self,
20789 e: &Engine,
20790 fa: &crate::hybrid::FullAttnLayer,
20791 il: usize,
20792 pos_d: &CudaSlice<i32>,
20793 cache: &mut Cache,
20794 cap_bucket_max: Option<(usize, usize)>,
20795 sl: &mut G4DcSlots,
20796 ) -> Result<(), Box<dyn std::error::Error>> {
20797 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
20798 let eps = self.cfg.rms_eps;
20799 let aux = self.gemma4_aux.as_ref().unwrap();
20800 let ones = aux.ones(e);
20801 #[cfg(debug_assertions)]
20802 crate::debug_assert_tensor_stream_device(
20803 ones,
20804 &e.stream(),
20805 "gemma4_decode_attn_dc_slotted.ones",
20806 );
20807 {
20808 let hq = unsafe { &*(&sl.hq as *const CudaSlice<i8>) };
20809 let hdq = unsafe { &*(&sl.hd_ as *const CudaSlice<f32>) };
20810 if swa {
20811 if !e.matmul_q4_fused3_into(
20812 &fa.wq, &fa.wk, &fa.wv, hq, hdq, &mut sl.q0, &mut sl.k0, &mut sl.v0,
20813 )? {
20814 if e.matmul_nvfp4_fused2_into(&fa.wq, &fa.wk, hq, hdq, &mut sl.q0, &mut sl.k0)?
20818 {
20819 self.g4_matvec_m1_into(e, &fa.wv, hq, hdq, &mut sl.v0)?;
20820 } else {
20821 return Err("slotted step: fused3 unavailable (non-uniform trunk)".into());
20822 }
20823 }
20824 } else {
20825 if !e.matmul_q4_fused2_into(&fa.wq, &fa.wk, hq, hdq, &mut sl.q0, &mut sl.k0)?
20826 && !e
20827 .matmul_nvfp4_fused2_into(&fa.wq, &fa.wk, hq, hdq, &mut sl.q0, &mut sl.k0)?
20828 {
20829 return Err("slotted step: fused2 unavailable".into());
20830 }
20831 let k0r = unsafe { &*(&sl.k0 as *const CudaSlice<f32>) };
20832 e.copy_into(&mut sl.v0, 0, k0r, nkv * hd)?;
20833 }
20834 }
20835 let ff = if swa {
20838 None
20839 } else {
20840 Some(
20841 aux.rope_freqs(e)
20842 .expect("gemma4 global rope needs rope_freqs.weight"),
20843 )
20844 };
20845 #[cfg(debug_assertions)]
20846 if let Some(ff) = ff {
20847 crate::debug_assert_tensor_stream_device(
20848 ff,
20849 &e.stream(),
20850 "gemma4_decode_attn_dc_slotted.rope_freqs",
20851 );
20852 }
20853 let kvl = cache.kv[il].as_mut().unwrap();
20854 let kv_fp8 = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
20855 if crate::Engine::qkv_append_on() {
20856 e.rms_norm_qkv_rope_append_dc(
20858 &sl.q0,
20859 &sl.k0,
20860 &sl.v0,
20861 fa.q_norm.float_data(),
20862 fa.k_norm.float_data(),
20863 ones,
20864 &mut sl.q,
20865 &mut sl.k,
20866 &mut sl.v,
20867 hd,
20868 self.gemma4_rope_dims(il),
20869 nh,
20870 nkv,
20871 pos_d,
20872 nh,
20873 nkv,
20874 base,
20875 1.0,
20876 ff,
20877 eps,
20878 &mut kvl.k,
20879 &mut kvl.v,
20880 &kvl.len_d,
20881 kvl.k_tok_bytes,
20882 kvl.v_tok_bytes,
20883 kv_fp8,
20884 )?;
20885 } else {
20886 e.rms_norm_qkv_rope(
20887 &sl.q0,
20888 &sl.k0,
20889 &sl.v0,
20890 fa.q_norm.float_data(),
20891 fa.k_norm.float_data(),
20892 ones,
20893 &mut sl.q,
20894 &mut sl.k,
20895 &mut sl.v,
20896 hd,
20897 self.gemma4_rope_dims(il),
20898 nh,
20899 nkv,
20900 pos_d,
20901 nh,
20902 nkv,
20903 base,
20904 1.0,
20905 ff,
20906 eps,
20907 )?;
20908 e.append_kv_quantized_dc(
20909 &sl.k,
20910 &sl.v,
20911 &mut kvl.k,
20912 &mut kvl.v,
20913 &kvl.len_d,
20914 kvl.kv_dim_k,
20915 kvl.kv_dim_v,
20916 kvl.k_tok_bytes,
20917 kvl.v_tok_bytes,
20918 kv_fp8,
20919 )?;
20920 }
20921 e.inc_seqlen(&mut kvl.len_d)?;
20922 let (b_swa, b_glob) = cap_bucket_max.expect("slotted step is capture-only");
20923 let k_view = e.view_u8(&kvl.k, kvl.k.len());
20924 let v_view = e.view_u8(&kvl.v, kvl.v.len());
20925 let rows_on = std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0");
20926 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
20927 let mut fa_q8 = false;
20931 if !swa && hd == 512 && b_glob >= crate::fa512_min_tkv() && rows_on {
20932 e.fa_decode_rows(
20933 &sl.q,
20934 &k_view,
20935 &v_view,
20936 &mut sl.attn,
20937 hd,
20938 nh,
20939 nkv,
20940 b_glob - 1,
20941 1,
20942 scale,
20943 kvl.k_tok_bytes,
20944 kvl.v_tok_bytes,
20945 Some((&kvl.len_d, -1)),
20946 false,
20947 false,
20948 Some((&mut sl.zq, &mut sl.zd)),
20949 )?;
20950 fa_q8 = true;
20951 } else if swa && b_swa > win && hd == 256 && rows_on {
20952 e.fa_decode_rows_w(
20953 &sl.q,
20954 &k_view,
20955 &v_view,
20956 &mut sl.attn,
20957 hd,
20958 nh,
20959 nkv,
20960 &kvl.len_d,
20961 -1,
20962 1,
20963 scale,
20964 win,
20965 kvl.k_tok_bytes,
20966 kvl.v_tok_bytes,
20967 Some((&mut sl.zq, &mut sl.zd)),
20968 )?;
20969 fa_q8 = true;
20970 } else {
20971 let b = if swa { b_swa } else { b_glob };
20972 e.fa_decode_dc(
20973 &sl.q,
20974 &k_view,
20975 &v_view,
20976 &mut sl.attn,
20977 hd,
20978 nh,
20979 nkv,
20980 &kvl.len_d,
20981 b,
20982 scale,
20983 kvl.k_tok_bytes,
20984 kvl.v_tok_bytes,
20985 swa && crate::Engine::wkv_on(),
20986 )?;
20987 }
20988 if !fa_q8 {
20989 let aq = unsafe { &*(&sl.attn as *const CudaSlice<f32>) };
20990 e.quantize_q8_1_into(aq, 1, nh * hd, &mut sl.zq, &mut sl.zd)?;
20991 }
20992 {
20993 let zq = unsafe { &*(&sl.zq as *const CudaSlice<i8>) };
20994 let zd = unsafe { &*(&sl.zd as *const CudaSlice<f32>) };
20995 self.g4_matvec_m1_into(e, &fa.wo, zq, zd, &mut sl.o)?;
20996 }
20997 Ok(())
20998 }
20999
21000 fn gemma4_layer_tail_slotted(
21003 &self,
21004 e: &Engine,
21005 layer: &crate::hybrid::HybridLayer,
21006 next_norm: Option<&CudaSlice<f32>>,
21007 sl: &mut G4DcSlots,
21008 ) -> Result<(), Box<dyn std::error::Error>> {
21009 let n_embd = self.cfg.n_embd as usize;
21010 let eps = self.cfg.rms_eps;
21011 let bits = layer.gemma4.as_ref().unwrap();
21012 let crate::hybrid::Ffn::Dense {
21013 ffn_gate,
21014 ffn_up,
21015 ffn_down,
21016 } = &layer.ffn
21017 else {
21018 return Err("slotted tail: dense ffn only".into());
21019 };
21020 let pnfold = Engine::g4_pnfold_on();
21021 if pnfold {
21022 let or = unsafe { &*(&sl.o as *const CudaSlice<f32>) };
21025 let xr = unsafe { &*(&sl.x as *const CudaSlice<f32>) };
21026 e.rms_pre_add_rms_norm_q8z_into(
21027 or,
21028 layer.post_attn_norm.float_data(),
21029 xr,
21030 bits.ffn_norm.float_data(),
21031 &mut sl.attn_out,
21032 &mut sl.zsh,
21033 n_embd,
21034 1,
21035 eps,
21036 &mut sl.zq,
21037 &mut sl.zd,
21038 )?;
21039 } else {
21040 e.add_rms_norm(
21041 &sl.cur,
21042 &sl.x,
21043 bits.ffn_norm.float_data(),
21044 &mut sl.attn_out,
21045 &mut sl.zsh,
21046 n_embd,
21047 1,
21048 eps,
21049 )?;
21050 }
21051 let n_ff = ffn_gate.out_features();
21052 if !pnfold {
21053 let zshr = unsafe { &*(&sl.zsh as *const CudaSlice<f32>) };
21054 e.quantize_q8_1_into(zshr, 1, n_embd, &mut sl.zq, &mut sl.zd)?;
21055 }
21056 {
21057 let zq = unsafe { &*(&sl.zq as *const CudaSlice<i8>) };
21058 let zd = unsafe { &*(&sl.zd as *const CudaSlice<f32>) };
21059 if !e.matmul_q4_fused2_into(ffn_gate, ffn_up, zq, zd, &mut sl.gate, &mut sl.up)?
21060 && !e.matmul_nvfp4_fused2_into(
21061 ffn_gate,
21062 ffn_up,
21063 zq,
21064 zd,
21065 &mut sl.gate,
21066 &mut sl.up,
21067 )?
21068 {
21069 return Err("slotted tail: ffn fused2 unavailable".into());
21070 }
21071 }
21072 debug_assert!(e.uses_q8_1_fast(ffn_down));
21073 {
21074 let upr = unsafe { &*(&sl.up as *const CudaSlice<f32>) };
21075 let upv = e.view(upr, n_ff);
21076 let up_all = upv.slice(0..n_ff);
21077 let gr = unsafe { &*(&sl.gate as *const CudaSlice<f32>) };
21078 e.gelu_tanh_mul_q8_1_into(
21079 gr,
21080 &up_all,
21081 &mut sl.act,
21082 n_ff,
21083 1,
21084 &mut sl.actq,
21085 &mut sl.actd,
21086 )?;
21087 }
21088 {
21089 let aq = unsafe { &*(&sl.actq as *const CudaSlice<i8>) };
21090 let ad = unsafe { &*(&sl.actd as *const CudaSlice<f32>) };
21091 self.g4_matvec_m1_into(e, ffn_down, aq, ad, &mut sl.f0)?;
21092 }
21093 if pnfold {
21094 if let Some(w) = next_norm {
21097 let f0r = unsafe { &*(&sl.f0 as *const CudaSlice<f32>) };
21098 let aor = unsafe { &*(&sl.attn_out as *const CudaSlice<f32>) };
21099 e.rms_pre_add_scale_rms_norm_q8_1_into(
21100 f0r,
21101 bits.post_ffw_norm.float_data(),
21102 aor,
21103 bits.layer_scale,
21104 w,
21105 &mut sl.xn,
21106 n_embd,
21107 1,
21108 eps,
21109 &mut sl.hq,
21110 &mut sl.hd_,
21111 )?;
21112 return Ok(());
21113 }
21114 }
21115 e.rms_norm(
21116 &sl.f0,
21117 bits.post_ffw_norm.float_data(),
21118 &mut sl.sn,
21119 n_embd,
21120 1,
21121 eps,
21122 )?;
21123 match next_norm {
21124 Some(w) => {
21125 e.add_scale_rms_norm_q8_1_into(
21126 &sl.sn,
21127 &sl.attn_out,
21128 bits.layer_scale,
21129 w,
21130 &mut sl.xn,
21131 n_embd,
21132 1,
21133 eps,
21134 &mut sl.hq,
21135 &mut sl.hd_,
21136 )?;
21137 }
21138 None => {
21139 e.add_scale(&sl.sn, &sl.attn_out, bits.layer_scale, &mut sl.xn, n_embd)?;
21140 }
21141 }
21142 Ok(())
21143 }
21144
21145 #[allow(clippy::too_many_arguments)]
21147 fn gemma4_decode_attn_dc(
21148 &self,
21149 e: &Engine,
21150 fa: &crate::hybrid::FullAttnLayer,
21151 il: usize,
21152 hq: &CudaSlice<i8>,
21153 hdq: &CudaSlice<f32>,
21154 pos_d: &CudaSlice<i32>,
21155 cache: &mut Cache,
21156 cap_bucket_max: Option<(usize, usize)>,
21157 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
21158 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
21159 let eps = self.cfg.rms_eps;
21160 let aux = self.gemma4_aux.as_ref().unwrap();
21161 let ones = aux.ones(e);
21162 #[cfg(debug_assertions)]
21163 crate::debug_assert_tensor_stream_device(ones, &e.stream(), "gemma4_decode_attn_dc.ones");
21164 let (q0, k0, v0) = if swa {
21165 match e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)? {
21166 Some(t3) => t3,
21167 None => match e.matmul_nvfp4_fused2(&fa.wq, &fa.wk, hq, hdq, 1)? {
21169 Some((q0, k0)) => {
21170 let h0 = e.zeros(0)?;
21171 let v0 = e.matmul_pre(&fa.wv, hq, hdq, &h0, 1)?;
21172 (q0, k0, v0)
21173 }
21174 None => {
21175 let h0 = e.zeros(0)?;
21176 (
21177 e.matmul_pre(&fa.wq, hq, hdq, &h0, 1)?,
21178 e.matmul_pre(&fa.wk, hq, hdq, &h0, 1)?,
21179 e.matmul_pre(&fa.wv, hq, hdq, &h0, 1)?,
21180 )
21181 }
21182 },
21183 }
21184 } else {
21185 let (q0, k0) = match e.matmul_q4_fused2(&fa.wq, &fa.wk, hq, hdq)? {
21186 Some(p) => p,
21187 None => match e.matmul_nvfp4_fused2(&fa.wq, &fa.wk, hq, hdq, 1)? {
21188 Some(p) => p,
21189 None => {
21190 let h0 = e.zeros(0)?;
21191 (
21192 e.matmul_pre(&fa.wq, hq, hdq, &h0, 1)?,
21193 e.matmul_pre(&fa.wk, hq, hdq, &h0, 1)?,
21194 )
21195 }
21196 },
21197 };
21198 let v0 = e.clone_dtod(&k0)?;
21199 (q0, k0, v0)
21200 };
21201 let mut q = e.uninit(nh * hd)?;
21202 let mut k = e.uninit(nkv * hd)?;
21203 let mut v = e.uninit(nkv * hd)?;
21204 let ff = if swa {
21206 None
21207 } else {
21208 Some(
21209 aux.rope_freqs(e)
21210 .expect("gemma4 global rope needs rope_freqs.weight"),
21211 )
21212 };
21213 #[cfg(debug_assertions)]
21214 if let Some(ff) = ff {
21215 crate::debug_assert_tensor_stream_device(
21216 ff,
21217 &e.stream(),
21218 "gemma4_decode_attn_dc.rope_freqs",
21219 );
21220 }
21221 let kvl = cache.kv[il].as_mut().unwrap();
21222 let kv_fp8 = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
21223 if crate::Engine::qkv_append_on() {
21224 e.rms_norm_qkv_rope_append_dc(
21226 &q0,
21227 &k0,
21228 &v0,
21229 fa.q_norm.float_data(),
21230 fa.k_norm.float_data(),
21231 ones,
21232 &mut q,
21233 &mut k,
21234 &mut v,
21235 hd,
21236 self.gemma4_rope_dims(il),
21237 nh,
21238 nkv,
21239 pos_d,
21240 nh,
21241 nkv,
21242 base,
21243 1.0,
21244 ff,
21245 eps,
21246 &mut kvl.k,
21247 &mut kvl.v,
21248 &kvl.len_d,
21249 kvl.k_tok_bytes,
21250 kvl.v_tok_bytes,
21251 kv_fp8,
21252 )?;
21253 } else {
21254 e.rms_norm_qkv_rope(
21255 &q0,
21256 &k0,
21257 &v0,
21258 fa.q_norm.float_data(),
21259 fa.k_norm.float_data(),
21260 ones,
21261 &mut q,
21262 &mut k,
21263 &mut v,
21264 hd,
21265 self.gemma4_rope_dims(il),
21266 nh,
21267 nkv,
21268 pos_d,
21269 nh,
21270 nkv,
21271 base,
21272 1.0,
21273 ff,
21274 eps,
21275 )?;
21276 e.append_kv_quantized_dc(
21277 &k,
21278 &v,
21279 &mut kvl.k,
21280 &mut kvl.v,
21281 &kvl.len_d,
21282 kvl.kv_dim_k,
21283 kvl.kv_dim_v,
21284 kvl.k_tok_bytes,
21285 kvl.v_tok_bytes,
21286 kv_fp8,
21287 )?;
21288 }
21289 e.inc_seqlen(&mut kvl.len_d)?;
21290 let mut attn = e.uninit(nh * hd)?;
21291 let mut fa_q8: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
21294 match cap_bucket_max {
21299 None => {
21300 kvl.len += 1;
21304 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
21305 if !swa
21306 && hd == 512
21307 && kvl.len >= crate::fa512_min_tkv()
21308 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
21309 {
21310 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
21313 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
21314 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
21315 e.fa_decode_rows(
21316 &q,
21317 &kp,
21318 &vp,
21319 &mut attn,
21320 hd,
21321 nh,
21322 nkv,
21323 kvl.len - 1,
21324 1,
21325 scale,
21326 kvl.k_tok_bytes,
21327 kvl.v_tok_bytes,
21328 Some((&kvl.len_d, -1)),
21329 false,
21330 false,
21331 Some((&mut aq8, &mut ad8)),
21332 )?;
21333 fa_q8 = Some((aq8, ad8));
21334 } else if swa
21335 && kvl.len > win
21336 && hd == 256
21337 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
21338 {
21339 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
21341 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
21342 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
21343 e.fa_decode_rows_w(
21344 &q,
21345 &kp,
21346 &vp,
21347 &mut attn,
21348 hd,
21349 nh,
21350 nkv,
21351 &kvl.len_d,
21352 -1,
21353 1,
21354 scale,
21355 win,
21356 kvl.k_tok_bytes,
21357 kvl.v_tok_bytes,
21358 Some((&mut aq8, &mut ad8)),
21359 )?;
21360 fa_q8 = Some((aq8, ad8));
21361 } else {
21362 let (off_tok, t_kv) = if swa && kvl.len > win {
21363 (kvl.len - win, win)
21364 } else {
21365 (0, kvl.len)
21366 };
21367 let k_view = e.view_u8_range(
21368 &kvl.k,
21369 off_tok * kvl.k_tok_bytes,
21370 (off_tok + t_kv) * kvl.k_tok_bytes,
21371 );
21372 let v_view = e.view_u8_range(
21373 &kvl.v,
21374 off_tok * kvl.v_tok_bytes,
21375 (off_tok + t_kv) * kvl.v_tok_bytes,
21376 );
21377 e.fa_decode_kvmod(
21378 &q,
21379 &k_view,
21380 &v_view,
21381 &mut attn,
21382 hd,
21383 nh,
21384 nkv,
21385 t_kv,
21386 scale,
21387 kvl.k_tok_bytes,
21388 kvl.v_tok_bytes,
21389 swa && crate::Engine::wkv_on(),
21390 )?;
21391 }
21392 }
21393 Some((b_swa, b_glob)) => {
21394 let k_view = e.view_u8(&kvl.k, kvl.k.len());
21400 let v_view = e.view_u8(&kvl.v, kvl.v.len());
21401 let rows_on = std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0");
21402 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
21403 if !swa && hd == 512 && b_glob >= crate::fa512_min_tkv() && rows_on {
21404 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
21405 e.fa_decode_rows(
21406 &q,
21407 &k_view,
21408 &v_view,
21409 &mut attn,
21410 hd,
21411 nh,
21412 nkv,
21413 b_glob - 1,
21414 1,
21415 scale,
21416 kvl.k_tok_bytes,
21417 kvl.v_tok_bytes,
21418 Some((&kvl.len_d, -1)),
21419 false,
21420 false,
21421 Some((&mut aq8, &mut ad8)),
21422 )?;
21423 fa_q8 = Some((aq8, ad8));
21424 } else if swa && b_swa > win && hd == 256 && rows_on {
21425 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
21426 e.fa_decode_rows_w(
21427 &q,
21428 &k_view,
21429 &v_view,
21430 &mut attn,
21431 hd,
21432 nh,
21433 nkv,
21434 &kvl.len_d,
21435 -1,
21436 1,
21437 scale,
21438 win,
21439 kvl.k_tok_bytes,
21440 kvl.v_tok_bytes,
21441 Some((&mut aq8, &mut ad8)),
21442 )?;
21443 fa_q8 = Some((aq8, ad8));
21444 } else {
21445 let b = if swa { b_swa } else { b_glob };
21446 e.fa_decode_dc(
21447 &q,
21448 &k_view,
21449 &v_view,
21450 &mut attn,
21451 hd,
21452 nh,
21453 nkv,
21454 &kvl.len_d,
21455 b,
21456 scale,
21457 kvl.k_tok_bytes,
21458 kvl.v_tok_bytes,
21459 swa && crate::Engine::wkv_on(),
21460 )?;
21461 }
21462 }
21463 }
21464 if let Some((aq8, ad8)) = fa_q8 {
21467 let mut y = e.uninit(fa.wo.out_features())?;
21468 self.g4_matvec_m1_into(e, &fa.wo, &aq8, &ad8, &mut y)?;
21469 return Ok(y);
21470 }
21471 e.matmul(&fa.wo, &attn, 1)
21472 }
21473
21474 #[allow(clippy::too_many_arguments)]
21479 #[allow(clippy::map_entry)] pub fn gemma4_generate_graph(
21482 &self,
21483 e: &Engine,
21484 prompt_pos: usize,
21485 first_token: u32,
21486 cache: &mut Cache,
21487 max_new: usize,
21488 eos: &[u32],
21489 mut on_token: impl FnMut(u32) -> bool,
21490 ) -> Result<(Vec<u32>, crate::decode::StopReason), Box<dyn std::error::Error>> {
21491 if self.is_gemma4_e4b() {
21492 return Err(
21493 "E4B graph serving is unwired (HANDOVER-E4B.md) — dc-eager is the serving arm"
21494 .into(),
21495 );
21496 }
21497 use crate::decode::StopReason;
21498 let n_vocab = self.output.out_features();
21499 let n_embd = self.cfg.n_embd as usize;
21500 let embd_gpu = self
21501 .embd_gpu
21502 .get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload"));
21503 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
21504 for kvl in cache.kv.iter_mut().flatten() {
21505 e.set_i32_one(&mut kvl.len_d, kvl.len as i32)?;
21506 }
21507 let mut token_d = e.stream().clone_htod(&[first_token])?;
21508 let mut pos_d = e.htod_i32(&[prompt_pos as i32])?;
21509 let g4 = self.cfg.gemma4.as_ref().unwrap();
21510 let (hd_s, hd_g) = (g4.key_length_swa as usize, g4.key_length_global as usize);
21511 let nkv_s = g4
21513 .head_count_kv
21514 .iter()
21515 .zip(g4.swa_pattern.iter())
21516 .find(|p| *p.1)
21517 .map(|p| *p.0 as usize)
21518 .unwrap_or(8);
21519 let nkv_g = g4
21520 .head_count_kv
21521 .iter()
21522 .zip(g4.swa_pattern.iter())
21523 .find(|p| !*p.1)
21524 .map(|p| *p.0 as usize)
21525 .unwrap_or(2);
21526 #[allow(clippy::type_complexity)]
21527 let mut graphs: std::collections::HashMap<
21529 ((bool, usize), (bool, usize), bool, bool),
21530 (
21531 cudarc::driver::CudaGraph,
21532 Vec<Box<dyn std::any::Any + Send>>,
21533 ),
21534 > = Default::default();
21535 let mut slots = self.g4_dc_slots(e)?;
21538 const RING: usize = 64;
21541 const DRAIN: usize = 1;
21547 let mut ring = e.stream().alloc_zeros::<u32>(RING)?;
21548 let ring_base = prompt_pos;
21549 let mut out = Vec::with_capacity(max_new);
21550 let mut reason = StopReason::MaxNew;
21551 let mut next = first_token;
21552 let mut captures = 0usize;
21553 for _ in 0..max_new {
21554 out.push(next);
21555 if eos.contains(&next) {
21556 reason = StopReason::Eos;
21557 break;
21558 }
21559 if !on_token(next) {
21560 reason = StopReason::Callback;
21561 break;
21562 }
21563 let t_kv = cache.pos + 1;
21564 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
21572 let f512 = crate::fa512_min_tkv();
21573 let key_s = if t_kv > win {
21574 (true, usize::MAX)
21575 } else {
21576 e.fa_bucket_key(t_kv, hd_s, nkv_s, crate::Engine::wkv_on())
21577 };
21578 let (key_g, rung_end) = if t_kv >= f512 {
21579 let end = (t_kv + 1).next_power_of_two().max(f512 * 2);
21582 ((true, end), end)
21583 } else {
21584 (e.fa_bucket_key(t_kv, hd_g, nkv_g, false), t_kv)
21585 };
21586 let key = (key_s, key_g, t_kv >= f512, t_kv > win);
21587 if !graphs.contains_key(&key) {
21588 let bucket_max = (t_kv, rung_end);
21589 let snap = cache.snapshot(e)?;
21591 let pos_save = e.dtoh_i32_one(&pos_d)?;
21592 let len_save: Vec<Option<i32>> = cache
21593 .kv
21594 .iter()
21595 .map(|k| k.as_ref().map(|kvl| e.dtoh_i32_one(&kvl.len_d).unwrap()))
21596 .collect();
21597 let tok_save = e.dtoh_u32_one(&token_d)?;
21598 let graph = {
21603 let tok_ref = &mut token_d;
21604 let pos_ref = &mut pos_d;
21605 let cache_ref = &mut *cache;
21606 let slots_ref = &mut slots;
21607 let ring_ref = &mut ring;
21608 e.capture_graph_retained_flags(
21609 cudarc::driver::sys::CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_USE_NODE_PRIORITY,
21610 |e| {
21611 let tok_in = unsafe { &*(tok_ref as *const CudaSlice<u32>) };
21613 let sl = unsafe { &mut *(slots_ref as *mut G4DcSlots) };
21614 let rg = unsafe { &mut *(ring_ref as *mut CudaSlice<u32>) };
21615 self.gemma4_decode_step_dc_slotted(e, tok_in, pos_ref, embd_gpu, qt, rb,
21616 cache_ref, n_vocab, Some(bucket_max),
21617 sl, tok_ref, Some((rg, ring_base)))
21618 })?
21619 };
21620 cache.rollback(e, &snap, 0)?;
21621 e.set_i32_one(&mut pos_d, pos_save)?;
21622 for (il, ls) in len_save.iter().enumerate() {
21623 if let (Some(kvl), Some(v)) = (cache.kv[il].as_mut(), ls) {
21624 e.set_i32_one(&mut kvl.len_d, *v)?;
21625 }
21626 }
21627 e.set_u32_one(&mut token_d, tok_save)?;
21628 if std::env::var("MEMRA_GRAPH_CENSUS").as_deref() == Ok("1")
21629 && let Ok(c) = crate::graph_update::node_census(&graph.0)
21630 {
21631 eprintln!("[graph-census] {c:?}");
21632 }
21633 graphs.insert(key, graph);
21634 captures += 1;
21635 }
21636 let mut chunk = 1usize;
21641 let drain_cap: usize = std::env::var("MEMRA_GRAPH_DRAIN")
21642 .ok()
21643 .and_then(|v| v.parse().ok())
21644 .unwrap_or(DRAIN);
21645 while chunk < drain_cap && out.len() + chunk < max_new {
21646 let t_next = cache.pos + 1 + chunk;
21647 let key_s2 = if t_next > win {
21648 (true, usize::MAX)
21649 } else {
21650 e.fa_bucket_key(t_next, hd_s, nkv_s, crate::Engine::wkv_on())
21651 };
21652 let key_g2 = if t_next >= f512 {
21653 (true, (t_next + 1).next_power_of_two().max(f512 * 2))
21654 } else {
21655 e.fa_bucket_key(t_next, hd_g, nkv_g, false)
21656 };
21657 if (key_s2, key_g2, t_next >= f512, t_next > win) != key {
21658 break;
21659 }
21660 chunk += 1;
21661 }
21662 let g = &graphs.get(&key).unwrap().0;
21663 for _ in 0..chunk {
21664 g.launch()?;
21665 }
21666 e.stream().synchronize()?;
21667 let ringh = e.dtoh_u32(&ring)?;
21668 for j in 0..chunk {
21669 let pos_j = cache.pos + j;
21670 let tok_j = ringh[(pos_j - ring_base) % RING];
21671 cache.pos += 0; if j + 1 == chunk {
21673 next = tok_j;
21674 } else {
21675 out.push(tok_j);
21676 if eos.contains(&tok_j) || !on_token(tok_j) {
21677 reason = if eos.contains(&tok_j) {
21678 StopReason::Eos
21679 } else {
21680 StopReason::Callback
21681 };
21682 let keep = cache.pos + j + 1;
21684 e.set_i32_one(&mut pos_d, keep as i32)?;
21685 for kvl in cache.kv.iter_mut().filter_map(|k| k.as_mut()) {
21686 e.set_i32_one(&mut kvl.len_d, keep as i32)?;
21687 kvl.len = keep;
21688 }
21689 cache.pos = keep;
21690 if std::env::var("MEMRA_GRAPH_STATS").is_ok() {
21691 eprintln!("[gemma-graph] captures={captures} buckets={}", graphs.len());
21692 }
21693 return Ok((out, reason));
21694 }
21695 }
21696 }
21697 cache.pos += chunk;
21698 for kvl in cache.kv.iter_mut().filter_map(|k| k.as_mut()) {
21699 kvl.len += chunk;
21700 }
21701 }
21702 if std::env::var("MEMRA_GRAPH_STATS").is_ok() {
21703 eprintln!("[gemma-graph] captures={captures} buckets={}", graphs.len());
21704 }
21705 Ok((out, reason))
21706 }
21707
21708 pub(crate) fn gemma4_decode_step_t(
21714 &self,
21715 e: &Engine,
21716 tokens: &[u32],
21717 pos0: usize,
21718 cache: &mut Cache,
21719 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
21720 Ok(self.gemma4_decode_step_t_h(e, tokens, pos0, cache)?.0)
21721 }
21722
21723 pub(crate) fn gemma4_decode_step_t_am(
21727 &self,
21728 e: &Engine,
21729 tokens: &[u32],
21730 pos0: usize,
21731 cache: &mut Cache,
21732 ) -> Result<(Vec<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
21733 let (ld, hn) = self.gemma4_verify_trunk(e, tokens, pos0, cache, None)?;
21734 let t = tokens.len();
21735 let n_vocab = self.output.out_features();
21736 let mut toks = e.stream().alloc_zeros::<u32>(t)?;
21737 for i in 0..t {
21738 e.argmax_token_device_col(&ld, i, n_vocab, &mut toks, i)?;
21739 }
21740 Ok((e.dtoh_u32(&toks)?, hn))
21741 }
21742
21743 pub(crate) fn gemma4_decode_step_t_am_dev(
21746 &self,
21747 e: &Engine,
21748 tok_d: &CudaSlice<u32>,
21749 t: usize,
21750 pos0: usize,
21751 cache: &mut Cache,
21752 ) -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
21753 let (ld, hn) = self.gemma4_verify_trunk(e, &vec![0u32; t], pos0, cache, Some(tok_d))?;
21754 let n_vocab = self.output.out_features();
21755 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
21756 for i in 0..t {
21757 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
21758 }
21759 Ok((vam, hn))
21760 }
21761
21762 pub(crate) fn gemma4_decode_step_t_h(
21765 &self,
21766 e: &Engine,
21767 tokens: &[u32],
21768 pos0: usize,
21769 cache: &mut Cache,
21770 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
21771 let (mut ld, hn) = self.gemma4_verify_trunk(e, tokens, pos0, cache, None)?;
21772 let t = tokens.len();
21773 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
21774 e.softcap(&mut ld, cap, t * self.output.out_features())?;
21775 Ok((e.dtoh(&ld)?, hn))
21776 }
21777
21778 pub(crate) fn verify_stream_scratch(
21781 &self,
21782 e: &Engine,
21783 cap: usize,
21784 ) -> Result<VerifyStreamScratch, Box<dyn std::error::Error>> {
21785 Ok(VerifyStreamScratch {
21786 pos_d: e.htod_i32(&vec![0i32; cap])?,
21787 row_ctrs: (0..cap)
21788 .map(|_| e.htod_i32(&[0]))
21789 .collect::<Result<_, _>>()?,
21790 })
21791 }
21792
21793 #[allow(clippy::too_many_arguments)] pub(crate) fn gemma4_verify_t_am_stream(
21802 &self,
21803 e: &Engine,
21804 tok_d: &CudaSlice<u32>,
21805 t: usize,
21806 ctr: &CudaSlice<i32>,
21807 hint: usize,
21808 cache: &mut Cache,
21809 scr: &mut VerifyStreamScratch,
21810 ) -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
21811 let n_embd = self.cfg.n_embd as usize;
21812 let eps = self.cfg.rms_eps;
21813 assert!(t <= scr.row_ctrs.len() && t <= 64);
21814 e.i32_iota_from(ctr, &mut scr.pos_d, t)?;
21815 for i in 0..t {
21816 e.i32_copy_add(ctr, &mut scr.row_ctrs[i], (i + 1) as i32)?;
21817 }
21818 let (pos_d, row_ctrs) = (&scr.pos_d, &scr.row_ctrs);
21819 let embd_gpu = self
21820 .embd_gpu
21821 .get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload"));
21822 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
21823 let mut x = e.embed_gather_device_td(embd_gpu, tok_d, t, n_embd, qt, rb)?;
21824 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
21825 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
21826 let n_layers = self.layers.len();
21827 for (il, layer) in self.layers.iter().enumerate() {
21828 let (hq, hdq) = match h_carry.take() {
21829 Some(p) => p,
21830 None => {
21831 e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, t, eps)?
21832 }
21833 };
21834 let Mixer::Full(fa) = &layer.mixer else {
21835 panic!("gemma4 layer {il} not full-attn")
21836 };
21837 let o = self
21838 .gemma4_verify_attn_stream(e, fa, il, &hq, &hdq, pos_d, t, cache, hint, row_ctrs)?;
21839 let next_norm = if il + 1 < n_layers {
21840 Some(self.layers[il + 1].attn_norm.float_data())
21841 } else {
21842 None
21843 };
21844 let (xn, hn) = self.gemma4_layer_tail_add_nq_pn(e, layer, &o, &x, t, next_norm)?;
21845 x = xn;
21846 h_carry = hn;
21847 self.dflash_tap(e, cache, il, &x, t)?;
21848 }
21849 let mut hn = e.uninit(t * n_embd)?;
21850 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
21851 let ld = e.matmul(&self.output, &hn, t)?;
21852 let n_vocab = self.output.out_features();
21853 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
21854 for i in 0..t {
21855 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
21856 }
21857 Ok((vam, hn))
21858 }
21859
21860 pub(crate) fn dflash_tap(
21867 &self,
21868 e: &Engine,
21869 cache: &mut Cache,
21870 il: usize,
21871 x: &CudaSlice<f32>,
21872 t: usize,
21873 ) -> Result<(), Box<dyn std::error::Error>> {
21874 let Some(taps) = cache.dflash_taps.as_mut() else {
21875 return Ok(());
21876 };
21877 let Some(slot) = taps.layer_ids.iter().position(|&l| l == il) else {
21878 return Ok(());
21879 };
21880 let h = taps.hidden;
21881 let n_taps = taps.layer_ids.len();
21882 let base = taps.base;
21883 debug_assert!(
21884 base + t <= taps.t,
21885 "tap window {base}+{t} exceeds sink {}",
21886 taps.t
21887 );
21888 let xv = e.view(x, t * h);
21889 for r in 0..t {
21890 let row = xv.slice(r * h..(r + 1) * h);
21891 e.copy_view_into(&mut taps.buf, (base + r) * n_taps * h + slot * h, &row, h)?;
21892 }
21893 Ok(())
21894 }
21895
21896 fn gemma4_verify_trunk(
21897 &self,
21898 e: &Engine,
21899 tokens: &[u32],
21900 pos0: usize,
21901 cache: &mut Cache,
21902 tok_dev: Option<&CudaSlice<u32>>,
21903 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
21904 let n_embd = self.cfg.n_embd as usize;
21905 let eps = self.cfg.rms_eps;
21906 let t = tokens.len();
21907 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
21908 let pos_d = e.htod_i32(&pos)?;
21909 let mut x = match tok_dev {
21910 Some(td) => {
21911 let embd_gpu = self
21912 .embd_gpu
21913 .get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload"));
21914 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
21915 e.embed_gather_device_td(embd_gpu, td, t, n_embd, qt, rb)?
21916 }
21917 None => e.htod(&self.embd.try_gather(n_embd, tokens)?)?,
21918 };
21919 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
21920 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
21921 let n_layers = self.layers.len();
21922 for (il, layer) in self.layers.iter().enumerate() {
21923 let (hq, hdq) = match h_carry.take() {
21924 Some(p) => p,
21925 None => {
21926 e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, t, eps)?
21927 }
21928 };
21929 let Mixer::Full(fa) = &layer.mixer else {
21930 panic!("gemma4 layer {il} not full-attn")
21931 };
21932 let o = self.gemma4_verify_attn(e, fa, il, &hq, &hdq, &pos_d, t, cache)?;
21933 let next_norm = if il + 1 < n_layers {
21934 Some(self.layers[il + 1].attn_norm.float_data())
21935 } else {
21936 None
21937 };
21938 let (xn, hn) = self.gemma4_layer_tail_add_nq_pn(e, layer, &o, &x, t, next_norm)?;
21939 x = xn;
21940 h_carry = hn;
21941 self.dflash_tap(e, cache, il, &x, t)?;
21942 }
21943 let mut hn = e.uninit(t * n_embd)?;
21944 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
21945 let mut ld = e.matmul(&self.output, &hn, t)?;
21946 self.gemma4_suppress(e, &mut ld, t)?; cache.pos += t;
21948 Ok((ld, hn))
21949 }
21950
21951 #[allow(clippy::too_many_arguments)]
21959 fn gemma4_verify_attn_stream(
21960 &self,
21961 e: &Engine,
21962 fa: &crate::hybrid::FullAttnLayer,
21963 il: usize,
21964 hq: &CudaSlice<i8>,
21965 hdq: &CudaSlice<f32>,
21966 pos_d: &CudaSlice<i32>,
21967 t: usize,
21968 cache: &mut Cache,
21969 hint: usize,
21970 row_ctrs: &[CudaSlice<i32>],
21971 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
21972 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
21973 let eps = self.cfg.rms_eps;
21974 let aux = self.gemma4_aux.as_ref().unwrap();
21975 let ones = aux.ones(e);
21976 #[cfg(debug_assertions)]
21977 crate::debug_assert_tensor_stream_device(
21978 ones,
21979 &e.stream(),
21980 "gemma4_verify_attn_stream.ones",
21981 );
21982 let h0 = e.zeros(0)?;
21983 let h = &h0;
21984 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
21987 let f2b = *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
21988 let fused_qkv = if f2b {
21989 if swa {
21990 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
21991 .map(|(a, b, c)| (a, b, Some(c)))
21992 } else {
21993 e.matmul_q4_fused2_batched(&fa.wq, &fa.wk, hq, hdq, t)?
21994 .map(|(a, b)| (a, b, None))
21995 }
21996 } else {
21997 None
21998 };
21999 let (q0, k0, v0) = match fused_qkv {
22000 Some((a, b, cv)) => {
22001 let v = match cv {
22002 Some(c) => c,
22003 None => e.clone_dtod(&b)?,
22004 };
22005 (a, b, v)
22006 }
22007 None => {
22008 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
22009 let k0 = e.matmul_pre(&fa.wk, hq, hdq, h, t)?;
22010 let v0 = if swa {
22011 e.matmul_pre(&fa.wv, hq, hdq, h, t)?
22012 } else {
22013 e.clone_dtod(&k0)?
22014 };
22015 (q0, k0, v0)
22016 }
22017 };
22018 let mut q = e.uninit(t * nh * hd)?;
22019 let mut k = e.uninit(t * nkv * hd)?;
22020 let mut v = e.uninit(t * nkv * hd)?;
22021 let ff = if swa {
22024 None
22025 } else {
22026 Some(
22027 aux.rope_freqs(e)
22028 .expect("gemma4 global rope needs rope_freqs.weight"),
22029 )
22030 };
22031 #[cfg(debug_assertions)]
22032 if let Some(ff) = ff {
22033 crate::debug_assert_tensor_stream_device(
22034 ff,
22035 &e.stream(),
22036 "gemma4_verify_attn_stream.rope_freqs",
22037 );
22038 }
22039 e.rms_norm_qkv_rope(
22040 &q0,
22041 &k0,
22042 &v0,
22043 fa.q_norm.float_data(),
22044 fa.k_norm.float_data(),
22045 ones,
22046 &mut q,
22047 &mut k,
22048 &mut v,
22049 hd,
22050 self.gemma4_rope_dims(il),
22051 nh * t,
22052 nkv * t,
22053 pos_d,
22054 nh,
22055 nkv,
22056 base,
22057 1.0,
22058 ff,
22059 eps,
22060 )?;
22061 let kvl = cache.kv[il].as_mut().unwrap();
22062 e.append_kv_quantized_rows_dc(
22064 &k,
22065 &v,
22066 &mut kvl.k,
22067 &mut kvl.v,
22068 &kvl.len_d,
22069 t,
22070 kvl.kv_dim_k,
22071 kvl.kv_dim_v,
22072 kvl.k_tok_bytes,
22073 kvl.v_tok_bytes,
22074 (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on()),
22075 )?;
22076 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
22079 let mut attn = e.uninit(t * nh * hd)?;
22080 let k_view = e.view_u8(&kvl.k, kvl.k.len());
22081 let v_view = e.view_u8(&kvl.v, kvl.v.len());
22082 if swa && hint + 1 >= win {
22085 e.fa_decode_rows_w(
22088 &q,
22089 &k_view,
22090 &v_view,
22091 &mut attn,
22092 hd,
22093 nh,
22094 nkv,
22095 &kvl.len_d,
22096 0,
22097 t,
22098 scale,
22099 win,
22100 kvl.k_tok_bytes,
22101 kvl.v_tok_bytes,
22102 None,
22103 )?;
22104 } else if hd == 512 && hint + t < crate::fa512_min_tkv() {
22105 let bucket = (hint + t + 2)
22118 .next_power_of_two()
22119 .min(crate::fa512_min_tkv().saturating_sub(1));
22120 let qv = e.view(&q, t * nh * hd);
22121 #[allow(clippy::needless_range_loop)]
22122 for i in 0..t {
22124 let q_row = qv.slice(i * nh * hd..(i + 1) * nh * hd);
22125 let mut q_one = e.uninit(nh * hd)?;
22126 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
22127 let mut a_one = e.uninit(nh * hd)?;
22128 e.fa_decode_dc(
22129 &q_one,
22130 &k_view,
22131 &v_view,
22132 &mut a_one,
22133 hd,
22134 nh,
22135 nkv,
22136 &row_ctrs[i],
22137 bucket,
22138 scale,
22139 kvl.k_tok_bytes,
22140 kvl.v_tok_bytes,
22141 false,
22142 )?;
22143 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
22144 }
22145 } else if hd == 512 {
22146 e.fa_decode_rows(
22149 &q,
22150 &k_view,
22151 &v_view,
22152 &mut attn,
22153 hd,
22154 nh,
22155 nkv,
22156 hint,
22157 t,
22158 scale,
22159 kvl.k_tok_bytes,
22160 kvl.v_tok_bytes,
22161 Some((&kvl.len_d, 0)),
22162 false,
22163 false,
22164 None,
22165 )?;
22166 } else {
22167 e.fa_decode_rows_dc(
22169 &q,
22170 &k_view,
22171 &v_view,
22172 &mut attn,
22173 hd,
22174 nh,
22175 nkv,
22176 &kvl.len_d,
22177 hint + t,
22178 t,
22179 scale,
22180 kvl.k_tok_bytes,
22181 kvl.v_tok_bytes,
22182 0,
22183 swa && crate::Engine::wkv_on(),
22184 )?;
22185 }
22186 e.matmul(&fa.wo, &attn, t)
22187 }
22188
22189 #[allow(clippy::too_many_arguments)] fn gemma4_verify_attn(
22191 &self,
22192 e: &Engine,
22193 fa: &crate::hybrid::FullAttnLayer,
22194 il: usize,
22195 hq: &CudaSlice<i8>,
22196 hdq: &CudaSlice<f32>,
22197 pos_d: &CudaSlice<i32>,
22198 t: usize,
22199 cache: &mut Cache,
22200 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
22201 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
22202 let eps = self.cfg.rms_eps;
22203 let aux = self.gemma4_aux.as_ref().unwrap();
22204 let ones = aux.ones(e);
22205 #[cfg(debug_assertions)]
22206 crate::debug_assert_tensor_stream_device(ones, &e.stream(), "gemma4_verify_attn.ones");
22207 let n_embd = self.cfg.n_embd as usize;
22208 let _ = n_embd;
22209
22210 let h0 = e.zeros(0)?;
22211 let h = &h0;
22212 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
22215 let f2b = *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
22216 let fused_qkv = if f2b {
22217 if swa {
22218 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
22219 .map(|(a, b, c)| (a, b, Some(c)))
22220 } else {
22221 e.matmul_q4_fused2_batched(&fa.wq, &fa.wk, hq, hdq, t)?
22222 .map(|(a, b)| (a, b, None))
22223 }
22224 } else {
22225 None
22226 };
22227 let (q0, k0, v0) = match fused_qkv {
22228 Some((a, b, cv)) => {
22229 let v = match cv {
22230 Some(c) => c,
22231 None => e.clone_dtod(&b)?,
22232 };
22233 (a, b, v)
22234 }
22235 None => {
22236 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
22237 let k0 = e.matmul_pre(&fa.wk, hq, hdq, h, t)?;
22238 let v0 = if swa {
22239 e.matmul_pre(&fa.wv, hq, hdq, h, t)?
22240 } else {
22241 e.clone_dtod(&k0)?
22242 };
22243 (q0, k0, v0)
22244 }
22245 };
22246 let mut q = e.uninit(t * nh * hd)?;
22247 let mut k = e.uninit(t * nkv * hd)?;
22248 let mut v = e.uninit(t * nkv * hd)?;
22249 let ff = if swa {
22252 None
22253 } else {
22254 Some(
22255 aux.rope_freqs(e)
22256 .expect("gemma4 global rope needs rope_freqs.weight"),
22257 )
22258 };
22259 #[cfg(debug_assertions)]
22260 if let Some(ff) = ff {
22261 crate::debug_assert_tensor_stream_device(
22262 ff,
22263 &e.stream(),
22264 "gemma4_verify_attn.rope_freqs",
22265 );
22266 }
22267 e.rms_norm_qkv_rope(
22268 &q0,
22269 &k0,
22270 &v0,
22271 fa.q_norm.float_data(),
22272 fa.k_norm.float_data(),
22273 ones,
22274 &mut q,
22275 &mut k,
22276 &mut v,
22277 hd,
22278 self.gemma4_rope_dims(il),
22279 nh * t,
22280 nkv * t,
22281 pos_d,
22282 nh,
22283 nkv,
22284 base,
22285 1.0,
22286 ff,
22287 eps,
22288 )?;
22289 let kvl = cache.kv[il].as_mut().unwrap();
22290 let base_len = kvl.len;
22291 e.append_kv_quantized_rows(
22292 &k,
22293 &v,
22294 &mut kvl.k,
22295 &mut kvl.v,
22296 base_len,
22297 t,
22298 kvl.kv_dim_k,
22299 kvl.kv_dim_v,
22300 kvl.k_tok_bytes,
22301 kvl.v_tok_bytes,
22302 (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on()),
22303 )?;
22304 kvl.len += t;
22305 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
22306 let mut attn = e.uninit(t * nh * hd)?;
22307 let rows_ok = (hd == 256 && base_len + 1 >= crate::fa_vec_min_tkv())
22310 || (hd == 512 && !swa && base_len + 1 >= crate::fa512_min_tkv());
22313 if rows_ok && (!swa || base_len + t <= win) {
22314 let k_view = e.view_u8(&kvl.k, (base_len + t) * kvl.k_tok_bytes);
22315 let v_view = e.view_u8(&kvl.v, (base_len + t) * kvl.v_tok_bytes);
22316 if hd == 512 {
22317 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
22319 e.fa_decode_rows(
22320 &q,
22321 &k_view,
22322 &v_view,
22323 &mut attn,
22324 hd,
22325 nh,
22326 nkv,
22327 base_len,
22328 t,
22329 scale,
22330 kvl.k_tok_bytes,
22331 kvl.v_tok_bytes,
22332 Some((&kvl.len_d, 0)),
22333 false,
22334 swa && crate::Engine::wkv_on(),
22335 None,
22336 )?;
22337 } else {
22338 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
22342 e.fa_decode_rows_dc(
22343 &q,
22344 &k_view,
22345 &v_view,
22346 &mut attn,
22347 hd,
22348 nh,
22349 nkv,
22350 &kvl.len_d,
22351 base_len + t,
22352 t,
22353 scale,
22354 kvl.k_tok_bytes,
22355 kvl.v_tok_bytes,
22356 0,
22357 swa && crate::Engine::wkv_on(),
22358 )?;
22359 }
22360 return e.matmul(&fa.wo, &attn, t);
22361 }
22362 if hd == 256
22370 && swa
22371 && base_len + 1 >= win
22372 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
22373 {
22374 let k_view = e.view_u8(&kvl.k, (base_len + t) * kvl.k_tok_bytes);
22375 let v_view = e.view_u8(&kvl.v, (base_len + t) * kvl.v_tok_bytes);
22376 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
22377 e.fa_decode_rows_w(
22378 &q,
22379 &k_view,
22380 &v_view,
22381 &mut attn,
22382 hd,
22383 nh,
22384 nkv,
22385 &kvl.len_d,
22386 0,
22387 t,
22388 scale,
22389 win,
22390 kvl.k_tok_bytes,
22391 kvl.v_tok_bytes,
22392 None,
22393 )?;
22394 return e.matmul(&fa.wo, &attn, t);
22395 }
22396 for i in 0..t {
22397 let avail = base_len + i + 1;
22398 let (off_tok, t_kv) = if swa && avail > win {
22399 (avail - win, win)
22400 } else {
22401 (0, avail)
22402 };
22403 let k_view = e.view_u8_range(
22404 &kvl.k,
22405 off_tok * kvl.k_tok_bytes,
22406 (off_tok + t_kv) * kvl.k_tok_bytes,
22407 );
22408 let v_view = e.view_u8_range(
22409 &kvl.v,
22410 off_tok * kvl.v_tok_bytes,
22411 (off_tok + t_kv) * kvl.v_tok_bytes,
22412 );
22413 let qi = e.view(&q, t * nh * hd);
22414 let q_row = qi.slice(i * nh * hd..(i + 1) * nh * hd);
22415 let mut q_one = e.uninit(nh * hd)?;
22416 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
22417 let mut a_one = e.uninit(nh * hd)?;
22418 if swa
22422 && avail > win
22423 && hd == 256
22424 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
22425 {
22426 let kp = e.view_u8(&kvl.k, avail * kvl.k_tok_bytes);
22427 let vp = e.view_u8(&kvl.v, avail * kvl.v_tok_bytes);
22428 e.i32_set_k(&mut kvl.len_d, (avail - 1) as i32)?;
22429 e.fa_decode_rows_w(
22430 &q_one,
22431 &kp,
22432 &vp,
22433 &mut a_one,
22434 hd,
22435 nh,
22436 nkv,
22437 &kvl.len_d,
22438 0,
22439 1,
22440 scale,
22441 win,
22442 kvl.k_tok_bytes,
22443 kvl.v_tok_bytes,
22444 None,
22445 )?;
22446 } else if !swa
22447 && hd == 512
22448 && avail >= crate::fa512_min_tkv()
22449 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0")
22450 {
22451 let kp = e.view_u8(&kvl.k, avail * kvl.k_tok_bytes);
22452 let vp = e.view_u8(&kvl.v, avail * kvl.v_tok_bytes);
22453 e.i32_set_k(&mut kvl.len_d, (avail - 1) as i32)?;
22454 e.fa_decode_rows(
22455 &q_one,
22456 &kp,
22457 &vp,
22458 &mut a_one,
22459 hd,
22460 nh,
22461 nkv,
22462 avail - 1,
22463 1,
22464 scale,
22465 kvl.k_tok_bytes,
22466 kvl.v_tok_bytes,
22467 Some((&kvl.len_d, 0)),
22468 false,
22469 false,
22470 None,
22471 )?;
22472 } else {
22473 e.fa_decode_kvmod(
22474 &q_one,
22475 &k_view,
22476 &v_view,
22477 &mut a_one,
22478 hd,
22479 nh,
22480 nkv,
22481 t_kv,
22482 scale,
22483 kvl.k_tok_bytes,
22484 kvl.v_tok_bytes,
22485 swa && crate::Engine::wkv_on(),
22486 )?;
22487 }
22488 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
22489 }
22490 e.matmul(&fa.wo, &attn, t)
22491 }
22492
22493 pub(crate) fn gemma4_decode_step_h(
22496 &self,
22497 e: &Engine,
22498 token: u32,
22499 cache: &mut Cache,
22500 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
22501 if let Some(split) = crate::pp::pp2_split(self.layers.len()) {
22506 let rt = crate::pp::Pp2Rt::get(e)?;
22507 let _walk = rt.acquire_walk("gemma4_decode_step_h_pp2")?;
22508 return self.gemma4_decode_step_h_pp2(e, token, cache, split);
22509 }
22510 if crate::pp::pp_cuts(self.layers.len()).is_some() {
22511 crate::pp::warn_unwired_once("gemma4 eager decode (N>2)");
22512 }
22513 let n_embd = self.cfg.n_embd as usize;
22514 let eps = self.cfg.rms_eps;
22515 let pos_d = e.htod_i32(&[cache.pos as i32])?;
22516 let mut x = e.htod(&self.embd.try_gather(n_embd, &[token])?)?;
22517 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
22518 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
22521 let n_layers = self.layers.len();
22522 for (il, layer) in self.layers.iter().enumerate() {
22523 let (hq, hdq) = match h_carry.take() {
22524 Some(p) => p,
22525 None => {
22526 e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, 1, eps)?
22527 }
22528 };
22529 let Mixer::Full(fa) = &layer.mixer else {
22530 panic!("gemma4 layer {il} not full-attn")
22531 };
22532 let o = self.gemma4_decode_attn(e, fa, il, &hq, &hdq, &pos_d, cache)?;
22533 let next_norm = if il + 1 < n_layers {
22534 Some(self.layers[il + 1].attn_norm.float_data())
22535 } else {
22536 None
22537 };
22538 let (xn, hn) = self.gemma4_layer_tail_add_nq_pn(e, layer, &o, &x, 1, next_norm)?;
22539 x = xn;
22540 h_carry = hn;
22541 }
22542 let mut hn = e.uninit(n_embd)?;
22543 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
22544 let h_seed = e.clone_dtod(&x)?;
22545 let mut ld = e.matmul(&self.output, &hn, 1)?;
22546 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
22547 e.softcap(&mut ld, cap, self.output.out_features())?; self.gemma4_suppress(e, &mut ld, 1)?;
22549 let logits = e.dtoh(&ld)?;
22550 cache.pos += 1;
22551 Ok((logits, h_seed))
22552 }
22553
22554 fn gemma4_decode_layers(
22562 &self,
22563 e: &Engine,
22564 mut x: CudaSlice<f32>,
22565 lo: usize,
22566 hi: usize,
22567 pos_d: &CudaSlice<i32>,
22568 cache: &mut Cache,
22569 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
22570 let n_embd = self.cfg.n_embd as usize;
22571 let eps = self.cfg.rms_eps;
22572 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
22573 for il in lo..hi {
22574 let layer = &self.layers[il];
22575 let (hq, hdq) = match h_carry.take() {
22576 Some(p) => p,
22577 None => {
22579 e.rms_norm_q8_1(&x, self.layers[il].attn_norm.float_data(), n_embd, 1, eps)?
22580 }
22581 };
22582 let Mixer::Full(fa) = &layer.mixer else {
22583 panic!("gemma4 layer {il} not full-attn")
22584 };
22585 let o = self.gemma4_decode_attn(e, fa, il, &hq, &hdq, pos_d, cache)?;
22586 let next_norm = if il + 1 < hi {
22587 Some(self.layers[il + 1].attn_norm.float_data())
22588 } else {
22589 None
22590 };
22591 let (xn, hn) = self.gemma4_layer_tail_add_nq_pn(e, layer, &o, &x, 1, next_norm)?;
22592 x = xn;
22593 h_carry = hn;
22594 }
22595 Ok(x)
22596 }
22597
22598 fn gemma4_decode_step_h_pp2(
22606 &self,
22607 e: &Engine,
22608 token: u32,
22609 cache: &mut Cache,
22610 split: usize,
22611 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
22612 if crate::pp::pp2_streams_off() {
22613 return self.gemma4_decode_step_h_pp2_samestream(e, token, cache, split);
22614 }
22615 let rt = crate::pp::Pp2Rt::get(e)?;
22616 let e0 = rt.engine(0, e);
22617 let e1 = rt.engine(1, e);
22618 let n_embd = self.cfg.n_embd as usize;
22619 let eps = self.cfg.rms_eps;
22620 let pos = cache.pos as i32;
22621
22622 let slot = {
22624 let _st0 = rt.enter(0);
22625 let pos_d = e0.htod_i32(&[pos])?;
22626 #[cfg(debug_assertions)]
22627 crate::debug_assert_tensor_stream_device(
22628 &pos_d,
22629 &e0.stream(),
22630 "gemma4_decode_step_h_pp2.stage0.pos_d",
22631 );
22632 let mut x = e0.htod(&self.embd.try_gather(n_embd, &[token])?)?;
22633 e0.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
22634 let x = self.gemma4_decode_layers(e0, x, 0, split, &pos_d, cache)?;
22635 rt.tx(0, &x, n_embd)?
22636 };
22637
22638 let _st1 = rt.enter(1);
22640 let pos_d = e1.htod_i32(&[pos])?;
22641 #[cfg(debug_assertions)]
22642 crate::debug_assert_tensor_stream_device(
22643 &pos_d,
22644 &e1.stream(),
22645 "gemma4_decode_step_h_pp2.stage1.pos_d",
22646 );
22647 let x = rt.rx(0, slot, n_embd)?;
22648 let x = self.gemma4_decode_layers(e1, x, split, self.layers.len(), &pos_d, cache)?;
22649
22650 let mut hn = e1.uninit(n_embd)?;
22651 e1.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
22652 let h_seed = e1.clone_dtod(&x)?;
22653 let mut ld = e1.matmul(&self.output, &hn, 1)?;
22654 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
22655 e1.softcap(&mut ld, cap, self.output.out_features())?;
22656 self.gemma4_suppress(e1, &mut ld, 1)?;
22657 let logits = e1.dtoh(&ld)?;
22658 cache.pos += 1;
22659 Ok((logits, h_seed))
22660 }
22661
22662 fn gemma4_decode_step_h_pp2_samestream(
22665 &self,
22666 e: &Engine,
22667 token: u32,
22668 cache: &mut Cache,
22669 split: usize,
22670 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
22671 let n_embd = self.cfg.n_embd as usize;
22672 let eps = self.cfg.rms_eps;
22673 let pos_d = e.htod_i32(&[cache.pos as i32])?;
22674
22675 let mut x = e.htod(&self.embd.try_gather(n_embd, &[token])?)?;
22677 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
22678 let x = self.gemma4_decode_layers(e, x, 0, split, &pos_d, cache)?;
22679
22680 let boundary_tx = e.clone_dtod(&x)?;
22682 let boundary_rx = e.clone_dtod(&boundary_tx)?;
22683
22684 let x =
22686 self.gemma4_decode_layers(e, boundary_rx, split, self.layers.len(), &pos_d, cache)?;
22687
22688 let mut hn = e.uninit(n_embd)?;
22689 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
22690 let h_seed = e.clone_dtod(&x)?;
22691 let mut ld = e.matmul(&self.output, &hn, 1)?;
22692 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
22693 e.softcap(&mut ld, cap, self.output.out_features())?;
22694 self.gemma4_suppress(e, &mut ld, 1)?;
22695 let logits = e.dtoh(&ld)?;
22696 cache.pos += 1;
22697 Ok((logits, h_seed))
22698 }
22699}
22700
22701impl HybridModel {
22720 pub(crate) fn step35_geom(&self, il: usize) -> memra_gguf::config::LayerGeometry {
22723 let geometry = self
22724 .cfg
22725 .layer_geometry(il as u32)
22726 .unwrap_or_else(|| panic!("step35 layer {il} has no geometry-table row"));
22727 debug_assert_eq!(
22728 geometry.attention_gate,
22729 memra_gguf::config::AttentionGateKind::SeparateHead
22730 );
22731 geometry
22732 }
22733
22734 #[allow(clippy::too_many_arguments)]
22794 fn step35_attn_pre_wo(
22795 &self,
22796 e: &Engine,
22797 fa: &FullAttnLayer,
22798 mut g3: Vec<CudaSlice<f32>>,
22799 hg: Option<&CudaSlice<f32>>,
22800 gt_pre: Option<&CudaSlice<f32>>,
22801 pos_d: &CudaSlice<i32>,
22802 t: usize,
22803 cache: Option<&mut Cache>,
22804 il: usize,
22805 seq_end: usize,
22806 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
22807 let geometry = self.cfg.full_attention_geometry_at(il as u32);
22808 let hd = geometry.head_dim_k as usize;
22809 let nkv = geometry.n_head_kv as usize;
22810 let nh = geometry.n_head as usize;
22811 let rbase = geometry.rope_base;
22812 let scale = geometry.attention_scale();
22813 let swa = geometry.window.is_some();
22814 let eps = self.cfg.rms_eps;
22815 let win = geometry.window.unwrap_or(0) as usize;
22816 let n_rot = geometry.n_rot as usize;
22817
22818 let v = g3.pop().unwrap();
22819 let k0 = g3.pop().unwrap();
22820 let q0 = g3.pop().unwrap();
22821
22822 let mut q = e.uninit(t * nh * hd)?;
22826 e.rms_norm(&q0, fa.q_norm.float_data(), &mut q, hd, nh * t, eps)?;
22827 let mut k = e.uninit(t * nkv * hd)?;
22828 e.rms_norm(&k0, fa.k_norm.float_data(), &mut k, hd, nkv * t, eps)?;
22829 let ff = if geometry.rope_factors {
22830 self.step35_aux.as_ref().and_then(|a| a.rope_freqs(e))
22831 } else {
22832 None
22833 };
22834 #[cfg(debug_assertions)]
22835 if let Some(ff) = ff {
22836 crate::debug_assert_tensor_stream_device(
22837 ff,
22838 &e.stream(),
22839 "step35_attn_pre_wo.rope_freqs",
22840 );
22841 }
22842 e.rope_neox2(&mut q, &mut k, pos_d, hd, n_rot, nh, nkv, t, rbase, 1.0, ff)?;
22843
22844 let mut attn = e.uninit(t * nh * hd)?;
22845 match cache {
22846 Some(cache) => {
22847 let base_len = cache.kv[il].as_ref().unwrap().len;
22848 let legacy_tkv = std::env::var("MEMRA_STEP35_SWA_TKV").as_deref() == Ok("1");
22850 let legacy_calllocal = std::env::var("MEMRA_PRIME_CALLLOCAL").as_deref() == Ok("1");
22851 let off = if swa {
22852 let raw = base_len.saturating_sub(win - 1);
22853 if legacy_tkv || legacy_calllocal {
22854 raw
22855 } else {
22856 raw & !31usize
22857 }
22858 } else {
22859 0
22860 };
22861 {
22862 let kvl = cache.kv[il].as_mut().unwrap();
22863 assert!(kvl.len + t <= cache.max_ctx, "step35 prime: KV overflow");
22864 let write_row = e.prepare_kv_append(kvl, off, t)?;
22865 e.append_kv_quantized_rows(
22866 &k,
22867 &v,
22868 &mut kvl.k,
22869 &mut kvl.v,
22870 write_row,
22871 t,
22872 kvl.kv_dim_k,
22873 kvl.kv_dim_v,
22874 kvl.k_tok_bytes,
22875 kvl.v_tok_bytes,
22876 crate::Engine::kv_fp8_on(),
22877 )?;
22878 kvl.len += t;
22879 let new_len = kvl.len as i32;
22880 e.set_i32_one(&mut kvl.len_d, new_len)?;
22881 }
22882 let kvl = cache.kv[il].as_ref().unwrap();
22883 let t_kv = base_len + t - off;
22906 let physical = kvl.physical_rows(off, off + t_kv)?;
22907 let k_view = e.view_u8_range(
22908 &kvl.k,
22909 physical.start * kvl.k_tok_bytes,
22910 physical.end * kvl.k_tok_bytes,
22911 );
22912 let v_view = e.view_u8_range(
22913 &kvl.v,
22914 physical.start * kvl.v_tok_bytes,
22915 physical.end * kvl.v_tok_bytes,
22916 );
22917 let swa_naive = if legacy_tkv {
22929 t_kv > win
22930 } else {
22931 seq_end > win
22932 };
22933 if swa && swa_naive {
22934 if std::env::var("MEMRA_STEP35_SWA_FA").as_deref() == Ok("0") {
22947 e.sdpa_naive_w_quantized_view(
22948 &q,
22949 &k_view,
22950 &v_view,
22951 &mut attn,
22952 hd,
22953 nh,
22954 nkv,
22955 t,
22956 t_kv,
22957 scale,
22958 true,
22959 win,
22960 kvl.k_tok_bytes,
22961 kvl.v_tok_bytes,
22962 )?;
22963 } else {
22964 e.fa_prefill_view_ws_w_hd128(
22965 &q,
22966 &k_view,
22967 &v_view,
22968 &mut attn,
22969 hd,
22970 nh,
22971 nkv,
22972 t,
22973 t_kv,
22974 scale,
22975 true,
22976 win,
22977 kvl.k_tok_bytes,
22978 kvl.v_tok_bytes,
22979 )?;
22980 }
22981 } else if std::env::var("MEMRA_NOFA").is_ok() {
22982 e.sdpa_naive_quantized_view(
22983 &q,
22984 &k_view,
22985 &v_view,
22986 &mut attn,
22987 hd,
22988 nh,
22989 nkv,
22990 t,
22991 t_kv,
22992 scale,
22993 true,
22994 kvl.k_tok_bytes,
22995 kvl.v_tok_bytes,
22996 )?;
22997 } else {
22998 e.fa_prefill_view_ws(
23003 &q,
23004 &k_view,
23005 &v_view,
23006 &mut attn,
23007 hd,
23008 nh,
23009 nkv,
23010 t,
23011 t_kv,
23012 scale,
23013 true,
23014 kvl.k_tok_bytes,
23015 kvl.v_tok_bytes,
23016 crate::Engine::kv_fp8_on(),
23017 )?;
23018 }
23019 }
23020 None => {
23021 debug_assert_eq!(
23026 seq_end, t,
23027 "step35 cacheless prefill is monolithic (seq_end == t)"
23028 );
23029 if swa && seq_end > win {
23030 e.sdpa_naive_w(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
23031 } else if std::env::var("MEMRA_NOFA").is_ok() {
23032 e.sdpa_naive(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
23033 } else {
23034 e.fa_prefill(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
23035 }
23036 }
23037 }
23038
23039 let gw = fa
23042 .attn_gate
23043 .as_ref()
23044 .ok_or("step35 layer is missing attn_gate.weight (head-wise attention gate)")?;
23045 let gt_owned = if gt_pre.is_none() {
23046 Some(e.matmul(
23047 gw,
23048 hg.ok_or("step35 attention needs hg when gt_pre is absent")?,
23049 t,
23050 )?)
23051 } else {
23052 None
23053 };
23054 let gt = gt_pre.or(gt_owned.as_ref()).unwrap();
23055 let mut ag = e.uninit(t * nh * hd)?;
23056 e.attn_head_gate(&attn, gt, &mut ag, None, hd, nh, t)?;
23057 Ok(ag)
23058 }
23059
23060 pub(crate) fn step35_attn(
23063 &self,
23064 e: &Engine,
23065 fa: &FullAttnLayer,
23066 h: &CudaSlice<f32>,
23067 pos_d: &CudaSlice<i32>,
23068 t: usize,
23069 il: usize,
23070 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
23071 let g3 = match self.full_attn_tp_qkv(e, fa, h, t)? {
23072 Some(g3) => g3,
23073 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
23074 };
23075 let ag = self.step35_attn_pre_wo(e, fa, g3, Some(h), None, pos_d, t, None, il, t)?;
23077 self.full_attn_o(e, fa, &ag, t)
23078 }
23079
23080 #[allow(clippy::too_many_arguments)]
23087 pub(crate) fn step35_attn_prime(
23088 &self,
23089 e: &Engine,
23090 fa: &FullAttnLayer,
23091 h: &CudaSlice<f32>,
23092 hx: Option<&CudaSlice<u8>>,
23093 pos_d: &CudaSlice<i32>,
23094 t: usize,
23095 cache: &mut Cache,
23096 il: usize,
23097 seq_end: usize,
23098 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
23099 if step_tp_prefill_enabled()? && fa.step_tp_qkv.is_some() {
23100 if hx.is_some() {
23101 return Err(
23102 "rank-local Step prefill preserves BF16 activations and refuses the q8_1 \
23103 pre-quantized prime path"
23104 .into(),
23105 );
23106 }
23107 return self.step35_tp_prefill_attn_resident(e, fa, il, h, pos_d, t, cache, seq_end);
23108 }
23109 let g3 = if fa.step_tp_qkv.is_some() {
23110 if hx.is_some() {
23111 return Err(
23112 "Step Q/K/V TP preserves BF16 activations and refuses the q8_1 \
23113 pre-quantized prime path"
23114 .into(),
23115 );
23116 }
23117 self.full_attn_tp_qkv(e, fa, h, t)?
23118 .expect("Step Q/K/V TP disappeared after the presence check")
23119 } else {
23120 match hx {
23121 Some(xh) => e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], h, xh, t)?,
23122 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
23123 }
23124 };
23125 let ag =
23126 self.step35_attn_pre_wo(e, fa, g3, Some(h), None, pos_d, t, Some(cache), il, seq_end)?;
23127 self.full_attn_o(e, fa, &ag, t)
23128 }
23129
23130 fn ensure_step_tp_kv_cache(
23131 &self,
23132 e: &Engine,
23133 fa: &FullAttnLayer,
23134 il: usize,
23135 cache: &mut Cache,
23136 ) -> Result<bool, Box<dyn std::error::Error>> {
23137 let tp = fa
23138 .step_tp_qkv
23139 .as_ref()
23140 .ok_or("Step TP cache hydration lost its resident projections")?;
23141 let geometry = self.cfg.full_attention_geometry_at(il as u32);
23142 let window = geometry.window.map(|window| window as usize);
23143 let ranks = tp.runtime.devices().len();
23144 let head_dim = geometry.head_dim_k as usize;
23145 let kv_heads = geometry.n_head_kv as usize;
23146 let max_ctx = cache.max_ctx;
23147
23148 if cache.tp_kv[il].is_some() {
23149 return Ok(false);
23150 }
23151 let local = cache.kv[il]
23152 .as_ref()
23153 .ok_or_else(|| format!("Step TP layer {il} has no owning-stage KV cache"))?;
23154 if local.kv_dim_k != kv_heads * head_dim || local.kv_dim_v != kv_heads * head_dim {
23155 return Err(format!(
23156 "Step TP layer {il} local KV geometry k={} v={} != {}",
23157 local.kv_dim_k,
23158 local.kv_dim_v,
23159 kv_heads * head_dim
23160 )
23161 .into());
23162 }
23163 let resident_start = window
23164 .map(|window| local.len.saturating_sub(window.saturating_sub(1)) & !31usize)
23165 .unwrap_or(0);
23166 let resident_rows = local.len - resident_start;
23167 let physical = local.physical_rows(resident_start, local.len)?;
23168 let k_rows = if resident_rows == 0 {
23169 Vec::new()
23170 } else {
23171 e.dtoh_u8_view(&e.view_u8_range(
23172 &local.k,
23173 physical.start * local.k_tok_bytes,
23174 physical.end * local.k_tok_bytes,
23175 ))?
23176 };
23177 let v_rows = if resident_rows == 0 {
23178 Vec::new()
23179 } else {
23180 e.dtoh_u8_view(&e.view_u8_range(
23181 &local.v,
23182 physical.start * local.v_tok_bytes,
23183 physical.end * local.v_tok_bytes,
23184 ))?
23185 };
23186 let mut distributed = match window {
23187 Some(window) => tp.runtime.allocate_tp_swa_kv_cache(
23188 kv_heads * head_dim,
23189 kv_heads * head_dim,
23190 max_ctx,
23191 window,
23192 )?,
23193 None => tp.runtime.allocate_tp_kv_cache(
23194 kv_heads * head_dim,
23195 kv_heads * head_dim,
23196 max_ctx,
23197 )?,
23198 };
23199 if distributed.k_tok_bytes() * ranks != local.k_tok_bytes
23200 || distributed.v_tok_bytes() * ranks != local.v_tok_bytes
23201 {
23202 return Err(format!(
23203 "Step TP layer {il} distributed/local KV token bytes disagree: \
23204 k={}x{ranks}/{} v={}x{ranks}/{}",
23205 distributed.k_tok_bytes(),
23206 local.k_tok_bytes,
23207 distributed.v_tok_bytes(),
23208 local.v_tok_bytes,
23209 )
23210 .into());
23211 }
23212 tp.runtime.hydrate_tp_kv_cache_from(
23213 &mut distributed,
23214 local.len,
23215 resident_start,
23216 &k_rows,
23217 &v_rows,
23218 )?;
23219 cache.tp_kv[il] = Some(distributed);
23220 Ok(true)
23221 }
23222
23223 #[allow(clippy::too_many_arguments)]
23224 #[allow(clippy::manual_is_multiple_of)] fn step35_tp_prefill_attn_resident(
23226 &self,
23227 e: &Engine,
23228 fa: &FullAttnLayer,
23229 il: usize,
23230 h: &CudaSlice<f32>,
23231 pos_d: &CudaSlice<i32>,
23232 tokens: usize,
23233 cache: &mut Cache,
23234 seq_end: usize,
23235 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
23236 let tp = fa
23237 .step_tp_qkv
23238 .as_ref()
23239 .ok_or("Step TP prefill lost its resident projections")?;
23240 let attention = tp
23241 .attention
23242 .as_ref()
23243 .ok_or("Step TP prefill lost its resident attention auxiliaries")?;
23244 let ranks = tp.runtime.devices().len();
23245 if !step_tp_prefill_shape(
23246 true,
23247 tokens,
23248 ranks,
23249 tp.runtime.native_p2p(),
23250 true,
23251 crate::Engine::kv_fp8_on(),
23252 ) {
23253 return Err(format!(
23254 "rank-local Step prefill requires tokens>={PRIME_MIN_T}, TP2/TP4 native P2P, \
23255 rank-local attention, and q8_0/q5_1 KV; got tokens={tokens} ranks={ranks} \
23256 native_p2p={} fp8_kv={}",
23257 tp.runtime.native_p2p(),
23258 crate::Engine::kv_fp8_on(),
23259 )
23260 .into());
23261 }
23262 for seam in [
23263 "MEMRA_STEP35_SWA_TKV",
23264 "MEMRA_PRIME_CALLLOCAL",
23265 "MEMRA_PRIME_F32CHUNK0",
23266 ] {
23267 if std::env::var(seam).as_deref() == Ok("1") {
23268 return Err(format!(
23269 "rank-local Step prefill has not qualified the legacy seam {seam}=1"
23270 )
23271 .into());
23272 }
23273 }
23274
23275 let geometry = self.cfg.full_attention_geometry_at(il as u32);
23276 let window = geometry.window.map(|window| window as usize);
23277 let head_dim = geometry.head_dim_k as usize;
23278 let heads = geometry.n_head as usize;
23279 let kv_heads = geometry.n_head_kv as usize;
23280 if heads % ranks != 0 || kv_heads % ranks != 0 {
23281 return Err(format!(
23282 "Step attention heads q={heads} kv={kv_heads} are not divisible by TP={ranks}"
23283 )
23284 .into());
23285 }
23286 let local_heads = heads / ranks;
23287 let local_kv_heads = kv_heads / ranks;
23288 let local_kv_dim = local_kv_heads * head_dim;
23289 let hidden = self.cfg.n_embd as usize;
23290 let expected_input = tokens
23291 .checked_mul(hidden)
23292 .ok_or("Step TP prefill input size overflow")?;
23293 if h.len() < expected_input {
23294 return Err(format!(
23295 "Step TP prefill input {} is shorter than {tokens}x{hidden}",
23296 h.len()
23297 )
23298 .into());
23299 }
23300 let positions = e.dtoh_i32(pos_d)?;
23301 if positions.len() != tokens {
23302 return Err(format!(
23303 "rank-local Step prefill positions {} != tokens {tokens}",
23304 positions.len()
23305 )
23306 .into());
23307 }
23308
23309 let mut active_input = e.uninit(expected_input)?;
23310 e.copy_view_into(
23311 &mut active_input,
23312 0,
23313 &h.slice(0..expected_input),
23314 expected_input,
23315 )?;
23316 let mut input = tp.runtime.allocate_replicated_device_rows(tokens, hidden)?;
23317 e.stream().synchronize()?;
23322 tp.runtime
23323 .refresh_replicated_device_rows_from_root(&mut input, &active_input)?;
23324 let q_raw = tp
23325 .runtime
23326 .bf16_column_parallel_resident_replicated_device_shards(&tp.q, &input)?;
23327 let k_raw = tp
23328 .runtime
23329 .bf16_column_parallel_resident_replicated_device_shards(&tp.k, &input)?;
23330 let v_raw = tp
23331 .runtime
23332 .bf16_column_parallel_resident_replicated_device_shards(&tp.v, &input)?;
23333 let mut q = Vec::with_capacity(ranks);
23334 let mut k = Vec::with_capacity(ranks);
23335 for rank in 0..ranks {
23336 let engine = tp
23337 .runtime
23338 .rank_engine(rank)
23339 .ok_or_else(|| format!("Step TP layer {il} has no engine for rank {rank}"))?;
23340 let _main = engine.gpu.enter_main()?;
23341 let mut q_rank = engine.uninit(tokens * local_heads * head_dim)?;
23342 engine.rms_norm(
23343 &q_raw[rank],
23344 &attention.q_norm[rank],
23345 &mut q_rank,
23346 head_dim,
23347 tokens * local_heads,
23348 self.cfg.rms_eps,
23349 )?;
23350 let mut k_rank = engine.uninit(tokens * local_kv_dim)?;
23351 engine.rms_norm(
23352 &k_raw[rank],
23353 &attention.k_norm[rank],
23354 &mut k_rank,
23355 head_dim,
23356 tokens * local_kv_heads,
23357 self.cfg.rms_eps,
23358 )?;
23359 let position = engine.htod_i32(&positions)?;
23360 let rope_freqs = if geometry.rope_factors {
23361 self.step35_aux
23362 .as_ref()
23363 .and_then(|aux| aux.rope_freqs(engine))
23364 } else {
23365 None
23366 };
23367 engine.rope_neox2(
23368 &mut q_rank,
23369 &mut k_rank,
23370 &position,
23371 head_dim,
23372 geometry.n_rot as usize,
23373 local_heads,
23374 local_kv_heads,
23375 tokens,
23376 geometry.rope_base,
23377 1.0,
23378 rope_freqs,
23379 )?;
23380 q.push(q_rank);
23381 k.push(k_rank);
23382 }
23383
23384 let gate_weight = fa
23385 .attn_gate
23386 .as_ref()
23387 .ok_or("step35 layer is missing attn_gate.weight")?;
23388 let gate = e.dtoh(&e.matmul(gate_weight, h, tokens)?)?;
23389 if gate.len() != tokens * heads {
23390 return Err(format!(
23391 "Step TP layer {il} gate output {} != {tokens}x{heads}",
23392 gate.len()
23393 )
23394 .into());
23395 }
23396
23397 let hydrated = self.ensure_step_tp_kv_cache(e, fa, il, cache)?;
23398 let base_len = cache.kv[il]
23399 .as_ref()
23400 .ok_or_else(|| format!("Step TP layer {il} lost its local KV cache"))?
23401 .len;
23402 let distributed = cache.tp_kv[il]
23403 .as_ref()
23404 .ok_or_else(|| format!("Step TP layer {il} lost its distributed KV cache"))?;
23405 if distributed.committed_len() != base_len || distributed.staged_len() != base_len {
23406 return Err(format!(
23407 "Step TP layer {il} cache lengths diverged before prefill: \
23408 local={base_len} distributed={}/{}",
23409 distributed.committed_len(),
23410 distributed.staged_len()
23411 )
23412 .into());
23413 }
23414 let target_len = base_len
23415 .checked_add(tokens)
23416 .ok_or("Step TP prefill cache length overflow")?;
23417 if target_len > cache.max_ctx {
23418 return Err(format!(
23419 "Step TP layer {il} prefill exceeds cache: {base_len}+{tokens}>{}",
23420 cache.max_ctx
23421 )
23422 .into());
23423 }
23424 if seq_end < target_len {
23425 return Err(format!(
23426 "Step TP layer {il} request end {seq_end} precedes chunk end {target_len}"
23427 )
23428 .into());
23429 }
23430
23431 let transaction = cache.tp_kv[il]
23432 .as_mut()
23433 .expect("distributed cache checked above")
23434 .begin_transaction()?;
23435 if let Err(error) = tp.runtime.append_tp_kv_transaction(
23436 cache.tp_kv[il]
23437 .as_mut()
23438 .expect("distributed cache checked above"),
23439 transaction,
23440 &k,
23441 &v_raw,
23442 tokens,
23443 ) {
23444 let _ = tp.runtime.rollback_tp_kv_transaction(
23445 cache.tp_kv[il]
23446 .as_mut()
23447 .expect("distributed cache checked above"),
23448 transaction,
23449 );
23450 return Err(error);
23451 }
23452
23453 let staged = (|| -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
23454 let distributed = cache.tp_kv[il]
23455 .as_ref()
23456 .expect("distributed cache checked above");
23457 let staged_len = distributed.staged_len();
23458 let view_start = window
23459 .map(|window| base_len.saturating_sub(window.saturating_sub(1)) & !31usize)
23460 .unwrap_or(0);
23461 let physical = distributed.physical_range(view_start, staged_len)?;
23462 let t_kv = staged_len - view_start;
23463 let swa_naive = window.is_some_and(|window| seq_end > window);
23464 let mut gated = Vec::with_capacity(ranks);
23465 #[allow(clippy::needless_range_loop)]
23466 for rank in 0..ranks {
23468 let engine = tp
23469 .runtime
23470 .rank_engine(rank)
23471 .ok_or_else(|| format!("Step TP layer {il} has no engine for rank {rank}"))?;
23472 let _main = engine.gpu.enter_main()?;
23473 let rank_cache = distributed
23474 .rank(rank)
23475 .ok_or_else(|| format!("Step TP layer {il} has no KV cache rank {rank}"))?;
23476 let k_view = engine.view_u8_range(
23477 rank_cache.k(),
23478 physical.start * distributed.k_tok_bytes(),
23479 physical.end * distributed.k_tok_bytes(),
23480 );
23481 let v_view = engine.view_u8_range(
23482 rank_cache.v(),
23483 physical.start * distributed.v_tok_bytes(),
23484 physical.end * distributed.v_tok_bytes(),
23485 );
23486 let mut attention_out = engine.uninit(tokens * local_heads * head_dim)?;
23487 if swa_naive {
23488 let window = window.expect("SWA predicate requires a window");
23489 if std::env::var("MEMRA_STEP35_SWA_FA").as_deref() == Ok("0") {
23490 engine.sdpa_naive_w_quantized_view(
23491 &q[rank],
23492 &k_view,
23493 &v_view,
23494 &mut attention_out,
23495 head_dim,
23496 local_heads,
23497 local_kv_heads,
23498 tokens,
23499 t_kv,
23500 geometry.attention_scale(),
23501 true,
23502 window,
23503 distributed.k_tok_bytes(),
23504 distributed.v_tok_bytes(),
23505 )?;
23506 } else {
23507 engine.fa_prefill_view_ws_w_hd128(
23508 &q[rank],
23509 &k_view,
23510 &v_view,
23511 &mut attention_out,
23512 head_dim,
23513 local_heads,
23514 local_kv_heads,
23515 tokens,
23516 t_kv,
23517 geometry.attention_scale(),
23518 true,
23519 window,
23520 distributed.k_tok_bytes(),
23521 distributed.v_tok_bytes(),
23522 )?;
23523 }
23524 } else if std::env::var("MEMRA_NOFA").is_ok() {
23525 engine.sdpa_naive_quantized_view(
23526 &q[rank],
23527 &k_view,
23528 &v_view,
23529 &mut attention_out,
23530 head_dim,
23531 local_heads,
23532 local_kv_heads,
23533 tokens,
23534 t_kv,
23535 geometry.attention_scale(),
23536 true,
23537 distributed.k_tok_bytes(),
23538 distributed.v_tok_bytes(),
23539 )?;
23540 } else {
23541 engine.fa_prefill_view_ws(
23542 &q[rank],
23543 &k_view,
23544 &v_view,
23545 &mut attention_out,
23546 head_dim,
23547 local_heads,
23548 local_kv_heads,
23549 tokens,
23550 t_kv,
23551 geometry.attention_scale(),
23552 true,
23553 distributed.k_tok_bytes(),
23554 distributed.v_tok_bytes(),
23555 false,
23556 )?;
23557 }
23558
23559 let gate_start = rank * local_heads;
23560 let mut gate_rank = Vec::with_capacity(tokens * local_heads);
23561 for token in 0..tokens {
23562 let start = token * heads + gate_start;
23563 gate_rank.extend_from_slice(&gate[start..start + local_heads]);
23564 }
23565 let gate_rank = engine.htod(&gate_rank)?;
23566 let mut gated_rank = engine.uninit(tokens * local_heads * head_dim)?;
23567 engine.attn_head_gate(
23568 &attention_out,
23569 &gate_rank,
23570 &mut gated_rank,
23571 None,
23572 head_dim,
23573 local_heads,
23574 tokens,
23575 )?;
23576 gated.push(gated_rank);
23577 }
23578 for rank in 1..ranks {
23579 let engine = tp
23580 .runtime
23581 .rank_engine(rank)
23582 .ok_or_else(|| format!("Step TP layer {il} has no engine for rank {rank}"))?;
23583 let _main = engine.gpu.enter_main()?;
23584 engine.stream().synchronize()?;
23585 }
23586
23587 let (output, k_shadow, v_shadow) = if tp.runtime.bulk_p2p() {
23588 let output = tp
23589 .runtime
23590 .step_bf16_row_parallel_resident_root_device(&tp.o, &gated, tokens)?;
23591 let k_shadow =
23592 tp.runtime
23593 .gather_native_column_shards_device(&k, tokens, local_kv_dim)?;
23594 let v_shadow =
23595 tp.runtime
23596 .gather_native_column_shards_device(&v_raw, tokens, local_kv_dim)?;
23597 let root = tp
23598 .runtime
23599 .rank_engine(0)
23600 .ok_or("Step TP prefill lost its root engine")?;
23601 let _main = root.gpu.enter_main()?;
23602 root.stream().synchronize()?;
23603 (output, k_shadow, v_shadow)
23604 } else {
23605 let attention = tp.runtime.gather_native_column_shards(
23606 &gated,
23607 tokens,
23608 local_heads * head_dim,
23609 )?;
23610 let output = tp
23611 .runtime
23612 .step_bf16_row_parallel_resident_native(&tp.o, &attention, tokens)?;
23613 let k_shadow = tp
23614 .runtime
23615 .gather_native_column_shards(&k, tokens, local_kv_dim)?;
23616 let v_shadow =
23617 tp.runtime
23618 .gather_native_column_shards(&v_raw, tokens, local_kv_dim)?;
23619 (e.htod(&output)?, e.htod(&k_shadow)?, e.htod(&v_shadow)?)
23620 };
23621 let local = cache.kv[il]
23622 .as_mut()
23623 .ok_or_else(|| format!("Step TP layer {il} lost its local KV cache"))?;
23624 if local.len != base_len {
23625 return Err(format!(
23626 "Step TP layer {il} local cache changed during prefill: \
23627 len={} base={base_len}",
23628 local.len
23629 )
23630 .into());
23631 }
23632 let retain_from = window
23633 .map(|window| {
23634 let staged_retain = staged_len.saturating_sub(window) & !31usize;
23635 let rollback_retain =
23636 base_len.saturating_sub(window.saturating_sub(1)) & !31usize;
23637 staged_retain.min(rollback_retain)
23638 })
23639 .unwrap_or(0);
23640 let write_row = e.prepare_kv_append(local, retain_from, tokens)?;
23641 e.append_kv_quantized_rows(
23642 &k_shadow,
23643 &v_shadow,
23644 &mut local.k,
23645 &mut local.v,
23646 write_row,
23647 tokens,
23648 local.kv_dim_k,
23649 local.kv_dim_v,
23650 local.k_tok_bytes,
23651 local.v_tok_bytes,
23652 false,
23653 )?;
23654 local.len = staged_len;
23655 e.set_i32_one(&mut local.len_d, staged_len as i32)?;
23656 Ok(output)
23657 })();
23658
23659 let output = match staged {
23660 Ok(output) => output,
23661 Err(error) => {
23662 let _ = tp.runtime.rollback_tp_kv_transaction(
23663 cache.tp_kv[il]
23664 .as_mut()
23665 .expect("distributed cache checked above"),
23666 transaction,
23667 );
23668 if let Some(local) = cache.kv[il].as_mut() {
23669 local.len = base_len;
23670 let _ = e.set_i32_one(&mut local.len_d, base_len as i32);
23671 }
23672 return Err(error);
23673 }
23674 };
23675 if let Err(error) = tp.runtime.commit_tp_kv_transaction(
23676 cache.tp_kv[il]
23677 .as_mut()
23678 .expect("distributed cache checked above"),
23679 transaction,
23680 tokens,
23681 ) {
23682 let _ = tp.runtime.rollback_tp_kv_transaction(
23683 cache.tp_kv[il]
23684 .as_mut()
23685 .expect("distributed cache checked above"),
23686 transaction,
23687 );
23688 let local = cache.kv[il].as_mut().expect("local cache checked above");
23689 local.len = base_len;
23690 e.set_i32_one(&mut local.len_d, base_len as i32)?;
23691 return Err(error);
23692 }
23693
23694 let committed = cache.tp_kv[il]
23695 .as_ref()
23696 .expect("distributed cache checked above")
23697 .committed_len();
23698 let local_len = cache.kv[il]
23699 .as_ref()
23700 .expect("local cache checked above")
23701 .len;
23702 if committed != local_len {
23703 return Err(format!(
23704 "Step TP layer {il} committed cache length {committed} != local shadow {local_len}"
23705 )
23706 .into());
23707 }
23708 eprintln!(
23709 "[step-tp-prefill-attn] execute layer={} devices={:?} tokens={tokens} \
23710 qkv_tensor_parallel=true qk_norm_rank_local=true rope_rank_local=true \
23711 kv_cache_distributed=true kv_cache_hydrated={} attention_tensor_parallel=true \
23712 attention_scope={} input_path=root-device-replicated gate_tensor_parallel=false \
23713 gate_shards=host-canonical o_tensor_parallel=true local_cache_shadow=true \
23714 cache_commit=chunk transport={} native_p2p=true bulk_p2p={} \
23715 output={} performance_claim=false",
23716 tp.layer,
23717 tp.devices,
23718 hydrated,
23719 if window.is_some() {
23720 "rank-local-swa-ring"
23721 } else {
23722 "rank-local-global"
23723 },
23724 tp.runtime.transport_label(),
23725 tp.runtime.bulk_p2p(),
23726 if tp.runtime.bulk_p2p() {
23727 "root-device"
23728 } else {
23729 "root-readback"
23730 },
23731 );
23732 Ok(output)
23733 }
23734
23735 fn step35_tp_decode_attn_resident(
23736 &self,
23737 e: &Engine,
23738 fa: &FullAttnLayer,
23739 il: usize,
23740 h: &CudaSlice<f32>,
23741 pos_d: &CudaSlice<i32>,
23742 cache: &mut Cache,
23743 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
23744 static ATTN_NS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
23748 static ATTN_CALLS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
23749 let timing = std::env::var("MEMRA_STEP_TP_TIMING").as_deref() == Ok("1");
23750 let started = timing.then(std::time::Instant::now);
23751 let result = if crate::tp::step_tp_decode_v2_enabled()? {
23752 self.step35_tp_decode_attn_resident_v2(e, fa, il, h, pos_d, cache)
23753 } else {
23754 self.step35_tp_decode_attn_resident_inner(e, fa, il, h, pos_d, cache)
23755 };
23756 if let Some(started) = started {
23757 use std::sync::atomic::Ordering;
23758 let ns = ATTN_NS.fetch_add(started.elapsed().as_nanos() as u64, Ordering::Relaxed)
23759 + started.elapsed().as_nanos() as u64;
23760 let calls = ATTN_CALLS.fetch_add(1, Ordering::Relaxed) + 1;
23761 if calls.is_multiple_of(430) {
23762 eprintln!(
23763 "[step-tp-attn-timing] calls={calls} total_ms={:.1} avg_us={:.1}",
23764 ns as f64 / 1.0e6,
23765 ns as f64 / calls as f64 / 1.0e3,
23766 );
23767 }
23768 }
23769 result
23770 }
23771
23772 #[allow(clippy::too_many_arguments)]
23773 fn step35_tp_decode_attn_resident_inner(
23774 &self,
23775 e: &Engine,
23776 fa: &FullAttnLayer,
23777 il: usize,
23778 h: &CudaSlice<f32>,
23779 pos_d: &CudaSlice<i32>,
23780 cache: &mut Cache,
23781 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
23782 static T_POS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
23787 static T_QKV: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
23788 static T_NORMROPE: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
23789 static T_GATE: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
23790 static T_APPEND: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
23791 static T_ATTN: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
23792 static T_OPROJ: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
23793 static T_SHADOW: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
23794 static T_PHASE_CALLS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
23795 let timing = std::env::var("MEMRA_STEP_TP_TIMING").as_deref() == Ok("1");
23796 #[allow(clippy::manual_is_multiple_of)] fn lap(
23798 runtime: &crate::tp::TpE4m3HostBounce,
23799 e: &Engine,
23800 timer: &std::sync::atomic::AtomicU64,
23801 started: &mut Option<std::time::Instant>,
23802 ) -> Result<(), Box<dyn std::error::Error>> {
23803 let Some(start) = started.as_mut() else {
23804 return Ok(());
23805 };
23806 for rank in 0..runtime.devices().len() {
23807 if let Some(engine) = runtime.rank_engine(rank) {
23808 let _main = engine.gpu.enter_main()?;
23809 engine.stream().synchronize()?;
23810 }
23811 }
23812 e.stream().synchronize()?;
23813 timer.fetch_add(
23814 start.elapsed().as_nanos() as u64,
23815 std::sync::atomic::Ordering::Relaxed,
23816 );
23817 *start = std::time::Instant::now();
23818 Ok(())
23819 }
23820 let tp = fa
23821 .step_tp_qkv
23822 .as_ref()
23823 .ok_or("Step TP decode lost its resident projections")?;
23824 let attention = tp
23825 .attention
23826 .as_ref()
23827 .ok_or("Step TP decode lost its resident attention auxiliaries")?;
23828 if !tp.runtime.native_p2p() {
23829 return Err("rank-local Step attention requires native P2P".into());
23830 }
23831 if crate::Engine::kv_fp8_on() {
23832 return Err("rank-local Step attention has not qualified the FP8 KV cache".into());
23833 }
23834
23835 let geometry = self.step35_geom(il);
23836 let window = geometry.window.map(|window| window as usize);
23837 let ranks = tp.runtime.devices().len();
23838 let head_dim = geometry.head_dim_k as usize;
23839 let heads = geometry.n_head as usize;
23840 let kv_heads = geometry.n_head_kv as usize;
23841 if !heads.is_multiple_of(ranks) || !kv_heads.is_multiple_of(ranks) {
23842 return Err(format!(
23843 "Step attention heads q={heads} kv={kv_heads} are not divisible by TP={ranks}"
23844 )
23845 .into());
23846 }
23847 let local_heads = heads / ranks;
23848 let local_kv_heads = kv_heads / ranks;
23849 let local_kv_dim = local_kv_heads * head_dim;
23850 let max_ctx = cache.max_ctx;
23851
23852 let hydrated = self.ensure_step_tp_kv_cache(e, fa, il, cache)?;
23853
23854 let base_len = cache.kv[il]
23855 .as_ref()
23856 .ok_or_else(|| format!("Step TP layer {il} lost its local KV cache"))?
23857 .len;
23858 let distributed = cache.tp_kv[il]
23859 .as_ref()
23860 .ok_or_else(|| format!("Step TP layer {il} lost its distributed KV cache"))?;
23861 if distributed.committed_len() != base_len || distributed.staged_len() != base_len {
23862 return Err(format!(
23863 "Step TP layer {il} cache lengths diverged before decode: \
23864 local={base_len} distributed={}/{}",
23865 distributed.committed_len(),
23866 distributed.staged_len()
23867 )
23868 .into());
23869 }
23870
23871 let mut lap_start = timing.then(std::time::Instant::now);
23872 let positions = e.dtoh_i32(pos_d)?;
23873 if positions.len() != 1 {
23874 return Err(format!(
23875 "rank-local Step decode requires one position, got {}",
23876 positions.len()
23877 )
23878 .into());
23879 }
23880 lap(&tp.runtime, e, &T_POS, &mut lap_start)?;
23881 let (q_raw, k_raw, v_raw, input_path) = if let Some(decode_input) =
23882 attention.decode_input.as_ref()
23883 {
23884 let mut decode_input = decode_input
23885 .lock()
23886 .map_err(|_| "Step TP replicated decode input lock is poisoned")?;
23887 e.stream().synchronize()?;
23891 tp.runtime
23892 .refresh_replicated_device_rows_from_root(&mut decode_input, h)?;
23893 let q_raw = tp
23894 .runtime
23895 .bf16_column_parallel_resident_replicated_device_shards(&tp.q, &decode_input)?;
23896 let k_raw = tp
23897 .runtime
23898 .bf16_column_parallel_resident_replicated_device_shards(&tp.k, &decode_input)?;
23899 let v_raw = tp
23900 .runtime
23901 .bf16_column_parallel_resident_replicated_device_shards(&tp.v, &decode_input)?;
23902 (q_raw, k_raw, v_raw, "root-device-replicated")
23903 } else {
23904 let activation = e.dtoh(h)?;
23905 let q_raw =
23906 tp.runtime
23907 .bf16_column_parallel_resident_device_shards(&tp.q, &activation, 1)?;
23908 let k_raw =
23909 tp.runtime
23910 .bf16_column_parallel_resident_device_shards(&tp.k, &activation, 1)?;
23911 let v_raw =
23912 tp.runtime
23913 .bf16_column_parallel_resident_device_shards(&tp.v, &activation, 1)?;
23914 (q_raw, k_raw, v_raw, "host-replicated")
23915 };
23916 lap(&tp.runtime, e, &T_QKV, &mut lap_start)?;
23917 let mut q = Vec::with_capacity(ranks);
23918 let mut k = Vec::with_capacity(ranks);
23919 for rank in 0..ranks {
23920 let engine = tp
23921 .runtime
23922 .rank_engine(rank)
23923 .ok_or_else(|| format!("Step TP layer {il} has no engine for rank {rank}"))?;
23924 let _main = engine.gpu.enter_main()?;
23925 let mut q_rank = engine.uninit(local_heads * head_dim)?;
23926 engine.rms_norm(
23927 &q_raw[rank],
23928 &attention.q_norm[rank],
23929 &mut q_rank,
23930 head_dim,
23931 local_heads,
23932 self.cfg.rms_eps,
23933 )?;
23934 let mut k_rank = engine.uninit(local_kv_dim)?;
23935 engine.rms_norm(
23936 &k_raw[rank],
23937 &attention.k_norm[rank],
23938 &mut k_rank,
23939 head_dim,
23940 local_kv_heads,
23941 self.cfg.rms_eps,
23942 )?;
23943 let position = engine.htod_i32(&positions)?;
23944 let rope_freqs = if geometry.rope_factors {
23945 self.step35_aux
23946 .as_ref()
23947 .and_then(|aux| aux.rope_freqs(engine))
23948 } else {
23949 None
23950 };
23951 engine.rope_neox2(
23952 &mut q_rank,
23953 &mut k_rank,
23954 &position,
23955 head_dim,
23956 geometry.n_rot as usize,
23957 local_heads,
23958 local_kv_heads,
23959 1,
23960 geometry.rope_base,
23961 1.0,
23962 rope_freqs,
23963 )?;
23964 q.push(q_rank);
23965 k.push(k_rank);
23966 }
23967 lap(&tp.runtime, e, &T_NORMROPE, &mut lap_start)?;
23968
23969 let gate_weight = fa
23970 .attn_gate
23971 .as_ref()
23972 .ok_or("step35 layer is missing attn_gate.weight")?;
23973 let gate = e.matmul(gate_weight, h, 1)?;
23974 let gate = e.dtoh(&gate)?;
23975 if gate.len() != heads {
23976 return Err(format!("Step TP layer {il} gate output {} != {heads}", gate.len()).into());
23977 }
23978 lap(&tp.runtime, e, &T_GATE, &mut lap_start)?;
23979
23980 let transaction = cache.tp_kv[il]
23981 .as_mut()
23982 .expect("distributed cache checked above")
23983 .begin_transaction()?;
23984 if let Err(error) = tp.runtime.append_tp_kv_transaction(
23985 cache.tp_kv[il]
23986 .as_mut()
23987 .expect("distributed cache checked above"),
23988 transaction,
23989 &k,
23990 &v_raw,
23991 1,
23992 ) {
23993 let _ = tp.runtime.rollback_tp_kv_transaction(
23994 cache.tp_kv[il]
23995 .as_mut()
23996 .expect("distributed cache checked above"),
23997 transaction,
23998 );
23999 return Err(error);
24000 }
24001 lap(&tp.runtime, e, &T_APPEND, &mut lap_start)?;
24002
24003 let staged = (|| -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
24004 let distributed = cache.tp_kv[il]
24005 .as_ref()
24006 .expect("distributed cache checked above");
24007 let staged_len = distributed.staged_len();
24008 let view_start = window
24009 .map(|window| staged_len.saturating_sub(window))
24010 .unwrap_or(0);
24011 let physical = distributed.physical_range(view_start, staged_len)?;
24012 let t_kv = staged_len - view_start;
24013 let mut gated = Vec::with_capacity(ranks);
24014 #[allow(clippy::needless_range_loop)]
24015 for rank in 0..ranks {
24017 let engine = tp
24018 .runtime
24019 .rank_engine(rank)
24020 .ok_or_else(|| format!("Step TP layer {il} has no engine for rank {rank}"))?;
24021 let _main = engine.gpu.enter_main()?;
24022 let rank_cache = distributed
24023 .rank(rank)
24024 .ok_or_else(|| format!("Step TP layer {il} has no KV cache rank {rank}"))?;
24025 let k_view = engine.view_u8_range(
24026 rank_cache.k(),
24027 physical.start * distributed.k_tok_bytes(),
24028 physical.end * distributed.k_tok_bytes(),
24029 );
24030 let v_view = engine.view_u8_range(
24031 rank_cache.v(),
24032 physical.start * distributed.v_tok_bytes(),
24033 physical.end * distributed.v_tok_bytes(),
24034 );
24035 let mut attention_out = engine.uninit(local_heads * head_dim)?;
24036 engine.fa_decode_kvmod(
24037 &q[rank],
24038 &k_view,
24039 &v_view,
24040 &mut attention_out,
24041 head_dim,
24042 local_heads,
24043 local_kv_heads,
24044 t_kv,
24045 geometry.attention_scale(),
24046 distributed.k_tok_bytes(),
24047 distributed.v_tok_bytes(),
24048 false,
24049 )?;
24050 let gate_start = rank * local_heads;
24051 let gate_rank = engine.htod(&gate[gate_start..gate_start + local_heads])?;
24052 let mut gated_rank = engine.uninit(local_heads * head_dim)?;
24053 engine.attn_head_gate(
24054 &attention_out,
24055 &gate_rank,
24056 &mut gated_rank,
24057 None,
24058 head_dim,
24059 local_heads,
24060 1,
24061 )?;
24062 gated.push(gated_rank);
24063 }
24064 lap(&tp.runtime, e, &T_ATTN, &mut lap_start)?;
24065
24066 let gathered =
24067 tp.runtime
24068 .gather_native_column_shards(&gated, 1, local_heads * head_dim)?;
24069 let output = tp
24070 .runtime
24071 .step_bf16_row_parallel_resident_native(&tp.o, &gathered, 1)?;
24072 let output = e.htod(&output)?;
24073 lap(&tp.runtime, e, &T_OPROJ, &mut lap_start)?;
24074
24075 let k_shadow = tp
24076 .runtime
24077 .gather_native_column_shards(&k, 1, local_kv_dim)?;
24078 let v_shadow = tp
24079 .runtime
24080 .gather_native_column_shards(&v_raw, 1, local_kv_dim)?;
24081 let k_shadow = e.htod(&k_shadow)?;
24082 let v_shadow = e.htod(&v_shadow)?;
24083 let local = cache.kv[il]
24084 .as_mut()
24085 .ok_or_else(|| format!("Step TP layer {il} lost its local KV cache"))?;
24086 if local.len != base_len || base_len + 1 > max_ctx {
24087 return Err(format!(
24088 "Step TP layer {il} local cache changed during decode: \
24089 len={} base={base_len} max={max_ctx}",
24090 local.len
24091 )
24092 .into());
24093 }
24094 let retain_from = window
24095 .map(|window| {
24096 let staged_retain = (base_len + 1).saturating_sub(window) & !31usize;
24097 let rollback_retain =
24098 base_len.saturating_sub(window.saturating_sub(1)) & !31usize;
24099 staged_retain.min(rollback_retain)
24100 })
24101 .unwrap_or(0);
24102 let write_row = e.prepare_kv_append(local, retain_from, 1)?;
24103 e.append_kv_quantized(
24104 &k_shadow,
24105 &v_shadow,
24106 &mut local.k,
24107 &mut local.v,
24108 write_row,
24109 local.kv_dim_k,
24110 local.kv_dim_v,
24111 local.k_tok_bytes,
24112 local.v_tok_bytes,
24113 false,
24114 )?;
24115 local.len = base_len + 1;
24116 e.set_i32_one(&mut local.len_d, local.len as i32)?;
24117 Ok(output)
24118 })();
24119
24120 let output = match staged {
24121 Ok(output) => output,
24122 Err(error) => {
24123 let _ = tp.runtime.rollback_tp_kv_transaction(
24124 cache.tp_kv[il]
24125 .as_mut()
24126 .expect("distributed cache checked above"),
24127 transaction,
24128 );
24129 if let Some(local) = cache.kv[il].as_mut() {
24130 local.len = base_len;
24131 let _ = e.set_i32_one(&mut local.len_d, base_len as i32);
24132 }
24133 return Err(error);
24134 }
24135 };
24136 if let Err(error) = tp.runtime.commit_tp_kv_transaction(
24137 cache.tp_kv[il]
24138 .as_mut()
24139 .expect("distributed cache checked above"),
24140 transaction,
24141 1,
24142 ) {
24143 let _ = tp.runtime.rollback_tp_kv_transaction(
24144 cache.tp_kv[il]
24145 .as_mut()
24146 .expect("distributed cache checked above"),
24147 transaction,
24148 );
24149 let local = cache.kv[il].as_mut().expect("local cache checked above");
24150 local.len = base_len;
24151 e.set_i32_one(&mut local.len_d, base_len as i32)?;
24152 return Err(error);
24153 }
24154
24155 let committed = cache.tp_kv[il]
24156 .as_ref()
24157 .expect("distributed cache checked above")
24158 .committed_len();
24159 let local_len = cache.kv[il]
24160 .as_ref()
24161 .expect("local cache checked above")
24162 .len;
24163 if committed != local_len {
24164 return Err(format!(
24165 "Step TP layer {il} committed cache length {committed} != local shadow {local_len}"
24166 )
24167 .into());
24168 }
24169 lap(&tp.runtime, e, &T_SHADOW, &mut lap_start)?;
24170 if timing {
24171 use std::sync::atomic::Ordering;
24172 let calls = T_PHASE_CALLS.fetch_add(1, Ordering::Relaxed) + 1;
24173 if calls.is_multiple_of(430) {
24174 let avg = |t: &std::sync::atomic::AtomicU64| {
24175 t.load(Ordering::Relaxed) as f64 / calls as f64 / 1.0e3
24176 };
24177 eprintln!(
24178 "[step-tp-attn-phase] calls={calls} avg_us pos={:.1} qkv={:.1} \
24179 normrope={:.1} gate={:.1} append={:.1} attn={:.1} oproj={:.1} shadow={:.1}",
24180 avg(&T_POS),
24181 avg(&T_QKV),
24182 avg(&T_NORMROPE),
24183 avg(&T_GATE),
24184 avg(&T_APPEND),
24185 avg(&T_ATTN),
24186 avg(&T_OPROJ),
24187 avg(&T_SHADOW),
24188 );
24189 }
24190 }
24191 eprintln!(
24192 "[step-tp-attn] execute layer={} devices={:?} tokens=1 \
24193 qkv_tensor_parallel=true qk_norm_rank_local=true rope_rank_local=true \
24194 kv_cache_distributed=true kv_cache_hydrated={} attention_tensor_parallel=true \
24195 attention_scope={} input_path={} kv_physical_rows={} \
24196 gate_tensor_parallel=false gate_shards=host-canonical o_tensor_parallel=true \
24197 local_cache_shadow=true cache_commit=immediate transport={} native_p2p=true \
24198 bulk_p2p={} output=root-readback performance_claim=false",
24199 tp.layer,
24200 tp.devices,
24201 hydrated,
24202 if window.is_some() {
24203 "rank-local-swa-ring"
24204 } else {
24205 "rank-local-global"
24206 },
24207 input_path,
24208 cache.tp_kv[il]
24209 .as_ref()
24210 .expect("distributed cache checked above")
24211 .physical_capacity(),
24212 tp.runtime.transport_label(),
24213 tp.runtime.bulk_p2p(),
24214 );
24215 Ok(output)
24216 }
24217
24218 #[allow(clippy::too_many_arguments)]
24225 pub(crate) fn step35_verify_qkv_precompute(
24230 &self,
24231 e: &Engine,
24232 il: usize,
24233 h_t: &CudaSlice<f32>,
24234 t: usize,
24235 ) -> Result<bool, Box<dyn std::error::Error>> {
24236 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
24237 return Ok(false);
24238 };
24239 let Some(tp) = fa.step_tp_qkv.as_ref() else {
24240 return Ok(false);
24241 };
24242 let Some(attention) = tp.attention.as_ref() else {
24243 return Ok(false);
24244 };
24245 if !tp.runtime.native_p2p() || !crate::tp::step_tp_qkv_fused_enabled()? {
24246 return Ok(false);
24247 }
24248 let geometry = self.step35_geom(il);
24249 let heads = geometry.n_head as usize;
24250 let ws_index = tp
24251 .runtime
24252 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
24253 let gate_shards = attention
24254 .gate_shards_bf16
24255 .as_deref()
24256 .map(crate::tp::StepTpGateShards::Bf16);
24257 tp.runtime.decode_v2_input_qkv_tcol(
24258 ws_index,
24259 e,
24260 h_t,
24261 t,
24262 &tp.q,
24263 &tp.k,
24264 &tp.v,
24265 gate_shards,
24266 )?;
24267 Ok(true)
24268 }
24269
24270 pub(crate) fn step35_verify_oproj_tcol(
24275 &self,
24276 e: &Engine,
24277 il: usize,
24278 t: usize,
24279 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
24280 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
24281 return Err("tcol o_proj join expects full attention".into());
24282 };
24283 let tp = fa
24284 .step_tp_qkv
24285 .as_ref()
24286 .ok_or("tcol o_proj join lost its resident projections")?;
24287 let heads = self.step35_geom(il).n_head as usize;
24288 let ws_index = tp
24289 .runtime
24290 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
24291 tp.runtime.decode_v2_oproj_tcol(ws_index, e, &tp.o, t)
24292 }
24293
24294 #[allow(dead_code)] pub(crate) fn step35_spec_fa2_precheck(
24302 &self,
24303 cache: &Cache,
24304 il: usize,
24305 pos0: usize,
24306 ) -> Result<bool, Box<dyn std::error::Error>> {
24307 fn nope(clause: &str, il: usize, pos0: usize) -> bool {
24310 static DBG: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
24311 static SEEN: std::sync::Mutex<Vec<&'static str>> = std::sync::Mutex::new(Vec::new());
24312 if *DBG.get_or_init(|| std::env::var("MEMRA_SPEC_FA2_DEBUG").as_deref() == Ok("1")) {
24313 let mut seen = SEEN.lock().unwrap();
24314 if !seen.contains(&clause) {
24315 seen.push(Box::leak(clause.to_string().into_boxed_str()));
24317 eprintln!("[spec-fa2] precheck FAIL clause={clause} il={il} pos0={pos0}");
24318 }
24319 }
24320 false
24321 }
24322 static ONLY: std::sync::OnceLock<Option<usize>> = std::sync::OnceLock::new();
24324 if let Some(only) =
24325 ONLY.get_or_init(|| std::env::var("MEMRA_SPEC_FA2_LAYER").ok()?.parse().ok())
24326 && *only != il
24327 {
24328 return Ok(false);
24329 }
24330 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
24331 return Ok(nope("mixer", il, pos0));
24332 };
24333 let Some(tp) = fa.step_tp_qkv.as_ref() else {
24334 return Ok(nope("step_tp", il, pos0));
24335 };
24336 let Some(attention) = tp.attention.as_ref() else {
24337 return Ok(nope("attention", il, pos0));
24338 };
24339 if !tp.runtime.native_p2p()
24340 || crate::Engine::kv_fp8_on()
24341 || !crate::tp::step_tp_dcw_enabled()?
24342 || !crate::tp::step_tp_qkv_fused_enabled()?
24343 || (attention.gate_shards.is_none() && attention.gate_shards_bf16.is_none())
24344 {
24345 return Ok(nope("runtime-doors", il, pos0));
24346 }
24347 let geometry = self.step35_geom(il);
24348 let head_dim = geometry.head_dim_k as usize;
24349 if head_dim > 256 || !head_dim.is_multiple_of(32) || !crate::fa_v3_on() {
24350 return Ok(nope("fa-class", il, pos0));
24351 }
24352 let Some(distributed) = cache.tp_kv[il].as_ref() else {
24353 return Ok(nope("tp-kv", il, pos0));
24354 };
24355 if distributed.staged_len() != pos0 {
24356 return Ok(nope("staged-len", il, pos0));
24357 }
24358 let (_, would_rebase) = distributed.peek_append_ring(2)?;
24361 if would_rebase {
24362 return Ok(nope("rebase", il, pos0));
24363 }
24364 let window = geometry.window.map(|w| w as usize);
24365 if let Some(w) = window
24373 && pos0 + 2 > w
24374 {
24375 return Ok(nope("swa-capped", il, pos0));
24376 }
24377 let (t0, t1) = (pos0 + 1, pos0 + 2);
24381 if t0 < 96 {
24382 return Ok(nope("dcw-floor", il, pos0));
24383 }
24384 if std::env::var("MEMRA_NO_FA_VEC").is_ok() || t0 < crate::fa_vec_min_tkv() {
24385 return Ok(nope("vec-floor", il, pos0));
24386 }
24387 let ranks = tp.runtime.devices().len();
24393 let local_kv_heads = (geometry.n_head_kv as usize / ranks).max(1);
24394 let sp0 = crate::fa_split_keys_pub(t0, local_kv_heads);
24395 let sp1 = crate::fa_split_keys_pub(t1, local_kv_heads);
24396 if sp0 != sp1 {
24397 return Ok(nope("partition-sp", il, pos0));
24398 }
24399 let (ns0, ns1) = (t0.div_ceil(sp0), t1.div_ceil(sp1));
24400 if ns0 != ns1 {
24401 return Ok(nope("partition-ns", il, pos0));
24402 }
24403 if t0.div_ceil(ns0) != t1.div_ceil(ns1) {
24404 return Ok(nope("partition-per", il, pos0));
24405 }
24406 Ok(true)
24407 }
24408
24409 pub(crate) fn step35_fa_rows_precheck(
24414 &self,
24415 cache: &Cache,
24416 il: usize,
24417 pos0: usize,
24418 t: usize,
24419 ) -> Result<bool, Box<dyn std::error::Error>> {
24420 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
24421 return Ok(false);
24422 };
24423 let Some(tp) = fa.step_tp_qkv.as_ref() else {
24424 return Ok(false);
24425 };
24426 let Some(attention) = tp.attention.as_ref() else {
24427 return Ok(false);
24428 };
24429 if !tp.runtime.native_p2p()
24430 || crate::Engine::kv_fp8_on()
24431 || !crate::tp::step_tp_dcw_enabled()?
24432 || !crate::tp::step_tp_qkv_fused_enabled()?
24433 || (attention.gate_shards.is_none() && attention.gate_shards_bf16.is_none())
24434 {
24435 return Ok(false);
24436 }
24437 let geometry = self.step35_geom(il);
24438 let head_dim = geometry.head_dim_k as usize;
24439 if head_dim > 256 || !head_dim.is_multiple_of(32) || !crate::fa_v3_on() {
24440 return Ok(false);
24441 }
24442 if crate::fa_sm_count() < 128
24443 || std::env::var("MEMRA_FA_SPLIT").is_ok()
24444 || std::env::var("MEMRA_FA_SP_SHORT").is_ok()
24445 || std::env::var("MEMRA_FA_SP16").is_ok()
24446 || std::env::var("MEMRA_NO_FA_VEC").is_ok()
24447 {
24448 return Ok(false);
24449 }
24450 let Some(distributed) = cache.tp_kv[il].as_ref() else {
24451 return Ok(false);
24452 };
24453 if distributed.staged_len() != pos0 {
24454 return Ok(false);
24455 }
24456 if distributed.peek_append_ring(t).is_err() {
24457 return Ok(false);
24458 }
24459 if distributed.ring_base().is_none() && pos0 + t > distributed.physical_capacity() {
24460 return Ok(false);
24461 }
24462 let window = geometry.window.map(|w| w as usize);
24465 let t0 = window.map(|w| (pos0 + 1).min(w)).unwrap_or(pos0 + 1);
24466 if t0 < 96 || t0 < crate::fa_vec_min_tkv() {
24467 return Ok(false);
24468 }
24469 Ok(true)
24470 }
24471
24472 pub(crate) fn step35_verify_fa_rows_join(
24476 &self,
24477 e: &Engine,
24478 il: usize,
24479 cache: &Cache,
24480 pos0: usize,
24481 t: usize,
24482 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
24483 use cudarc::driver::DevicePtr;
24484 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
24485 return Err("fa rows join expects full attention".into());
24486 };
24487 let tp = fa
24488 .step_tp_qkv
24489 .as_ref()
24490 .ok_or("fa rows join lost its resident projections")?;
24491 let geometry = self.step35_geom(il);
24492 let heads = geometry.n_head as usize;
24493 let head_dim = geometry.head_dim_k as usize;
24494 let window = geometry.window.map(|w| w as usize);
24495 let distributed = cache.tp_kv[il]
24496 .as_ref()
24497 .ok_or("fa rows join lost its distributed KV cache")?;
24498 let (k_tok_bytes, v_tok_bytes) = (distributed.k_tok_bytes(), distributed.v_tok_bytes());
24499 let ladder = |t_kv: usize| -> usize {
24501 if t_kv <= 2048 {
24502 16
24503 } else if t_kv <= 16384 {
24504 64
24505 } else {
24506 128
24507 }
24508 };
24509 let mut max_ns = 1usize;
24510 for r in 0..t {
24511 let t_kv = window
24512 .map(|w| (pos0 + r + 1).min(w))
24513 .unwrap_or(pos0 + r + 1);
24514 max_ns = max_ns.max(t_kv.div_ceil(ladder(t_kv)));
24515 }
24516 let ranks = tp.runtime.devices().len();
24521 let mut tables = Vec::with_capacity(ranks);
24522 for rank in 0..ranks {
24523 let engine = tp
24524 .runtime
24525 .rank_engine(rank)
24526 .ok_or("fa rows join lost a rank engine")?;
24527 let rank_cache = distributed
24528 .rank(rank)
24529 .ok_or("fa rows join lost a KV cache rank")?;
24530 let _main = engine.gpu.enter_main()?;
24531 let s = engine.stream();
24532 let (kp, _g0) = rank_cache.k().device_ptr(&s);
24533 let (vp, _g1) = rank_cache.v().device_ptr(&s);
24534 let (lp, _g2) = rank_cache.len_d().device_ptr(&s);
24535 let bp = match rank_cache.base_d() {
24536 Some(b) => {
24537 let (p, _g) = b.device_ptr(&s);
24538 p
24539 }
24540 None => 0u64,
24541 };
24542 let mut host = Vec::with_capacity(t * 6);
24543 for r in 0..t {
24544 host.extend_from_slice(&[kp, vp, lp, bp, 0u64, (t - 1 - r) as u64]);
24545 }
24546 tables.push(engine.stream().clone_htod(&host)?);
24547 }
24548 let tabs: Vec<&CudaSlice<u64>> = tables.iter().collect();
24549 let ws_index = tp
24550 .runtime
24551 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
24552 tp.runtime.decode_v2_fa_rows_join(
24553 ws_index,
24554 e,
24555 &tp.o,
24556 &tabs,
24557 t,
24558 head_dim,
24559 window.unwrap_or(0),
24560 max_ns,
24561 geometry.attention_scale(),
24562 k_tok_bytes,
24563 v_tok_bytes,
24564 )
24565 }
24566
24567 pub(crate) fn step35_batch_fa_rows_precheck(
24571 &self,
24572 caches: &[&mut Cache],
24573 row_to_cache: impl Fn(usize) -> usize,
24574 positions: &[i32],
24575 il: usize,
24576 ) -> Result<bool, Box<dyn std::error::Error>> {
24577 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
24578 return Ok(false);
24579 };
24580 let Some(tp) = fa.step_tp_qkv.as_ref() else {
24581 return Ok(false);
24582 };
24583 let Some(attention) = tp.attention.as_ref() else {
24584 return Ok(false);
24585 };
24586 if !tp.runtime.native_p2p()
24587 || crate::Engine::kv_fp8_on()
24588 || !crate::tp::step_tp_dcw_enabled()?
24589 || !crate::tp::step_tp_qkv_fused_enabled()?
24590 || (attention.gate_shards.is_none() && attention.gate_shards_bf16.is_none())
24591 {
24592 return Ok(false);
24593 }
24594 let geometry = self.step35_geom(il);
24595 let head_dim = geometry.head_dim_k as usize;
24596 if head_dim > 256 || !head_dim.is_multiple_of(32) || !crate::fa_v3_on() {
24597 return Ok(false);
24598 }
24599 if crate::fa_sm_count() < 128
24600 || std::env::var("MEMRA_FA_SPLIT").is_ok()
24601 || std::env::var("MEMRA_FA_SP_SHORT").is_ok()
24602 || std::env::var("MEMRA_FA_SP16").is_ok()
24603 || std::env::var("MEMRA_NO_FA_VEC").is_ok()
24604 {
24605 return Ok(false);
24606 }
24607 let window = geometry.window.map(|w| w as usize);
24608 for (r, &pos) in positions.iter().enumerate() {
24609 let cache = &caches[row_to_cache(r)];
24610 let Some(distributed) = cache.tp_kv[il].as_ref() else {
24611 return Ok(false);
24612 };
24613 if distributed.staged_len() != pos as usize {
24614 return Ok(false);
24615 }
24616 if distributed.peek_append_ring(1)?.1 {
24617 return Ok(false);
24618 }
24619 let t0 = window
24620 .map(|w| (pos as usize + 1).min(w))
24621 .unwrap_or(pos as usize + 1);
24622 if t0 < 96 || t0 < crate::fa_vec_min_tkv() {
24623 return Ok(false);
24624 }
24625 }
24626 Ok(true)
24627 }
24628
24629 pub(crate) fn step35_verify_rope_fa_pass(
24635 &self,
24636 e: &Engine,
24637 il: usize,
24638 cache: &Cache,
24639 pos0: usize,
24640 t: usize,
24641 stage_pos: bool,
24642 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
24643 use cudarc::driver::DevicePtr;
24644 if !crate::tp::fuse_rope_append_on() {
24645 return Ok(None);
24646 }
24647 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
24648 return Ok(None);
24649 };
24650 let Some(tp) = fa.step_tp_qkv.as_ref() else {
24651 return Ok(None);
24652 };
24653 let Some(attention) = tp.attention.as_ref() else {
24654 return Ok(None);
24655 };
24656 let geometry = self.step35_geom(il);
24657 let head_dim = geometry.head_dim_k as usize;
24658 if head_dim != 128 {
24659 return Ok(None);
24660 }
24661 let heads = geometry.n_head as usize;
24662 let window = geometry.window.map(|w| w as usize);
24663 let ranks = tp.runtime.devices().len();
24664 let Some(distributed) = cache.tp_kv[il].as_ref() else {
24665 return Ok(None);
24666 };
24667 if distributed.kv_dim_k() != distributed.kv_dim_v() {
24668 return Ok(None);
24669 }
24670 {
24671 let rank0 = distributed.rank(0).ok_or("verify rope pass lost rank 0")?;
24672 if rank0.base_d().is_none()
24673 && distributed.staged_len() > distributed.physical_capacity()
24674 {
24675 return Ok(None);
24676 }
24677 }
24678 let mut rope_freqs = Vec::with_capacity(ranks);
24679 for rank in 0..ranks {
24680 let engine = tp
24681 .runtime
24682 .rank_engine(rank)
24683 .ok_or("verify rope pass lost a rank engine")?;
24684 rope_freqs.push(if geometry.rope_factors {
24685 match self
24686 .step35_aux
24687 .as_ref()
24688 .and_then(|aux| aux.rope_freqs(engine))
24689 {
24690 Some(f) => Some(f),
24691 None => return Ok(None),
24692 }
24693 } else {
24694 None
24695 });
24696 }
24697 let ladder = |t_kv: usize| -> usize {
24698 if t_kv <= 2048 {
24699 16
24700 } else if t_kv <= 16384 {
24701 64
24702 } else {
24703 128
24704 }
24705 };
24706 let (k_tok_bytes, v_tok_bytes) = (distributed.k_tok_bytes(), distributed.v_tok_bytes());
24707 let mut max_ns = 1usize;
24708 let mut positions = Vec::with_capacity(t);
24709 for r in 0..t {
24710 positions.push((pos0 + r) as i32);
24711 let t_kv = window
24712 .map(|w| (pos0 + r + 1).min(w))
24713 .unwrap_or(pos0 + r + 1);
24714 max_ns = max_ns.max(t_kv.div_ceil(ladder(t_kv)));
24715 }
24716 let mut session_parts: Vec<Vec<[u64; 4]>> = vec![Vec::with_capacity(t); ranks];
24717 let mut tab_keys = vec![0u64; ranks];
24718 for rank in 0..ranks {
24719 let engine = tp
24720 .runtime
24721 .rank_engine(rank)
24722 .ok_or("verify rope pass lost a rank engine")?;
24723 let rank_cache = distributed
24724 .rank(rank)
24725 .ok_or("verify rope pass lost a KV cache rank")?;
24726 let _main = engine.gpu.enter_main()?;
24727 let s = engine.stream();
24728 let (kp, _g0) = rank_cache.k().device_ptr(&s);
24729 let (vp, _g1) = rank_cache.v().device_ptr(&s);
24730 let (lp, _g2) = rank_cache.len_d().device_ptr(&s);
24731 let bp = match rank_cache.base_d() {
24732 Some(b) => {
24733 let (p, _g) = b.device_ptr(&s);
24734 p
24735 }
24736 None => 0u64,
24737 };
24738 tab_keys[rank] = kp
24739 .rotate_left(17)
24740 .wrapping_add(bp)
24741 .wrapping_add((il as u64) << 32)
24742 .wrapping_add(t as u64)
24743 .wrapping_add(1 << 63);
24744 for _r in 0..t {
24745 session_parts[rank].push([kp, vp, lp, bp]);
24746 }
24747 }
24748 let ws_index = tp
24749 .runtime
24750 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
24751 tp.runtime
24752 .decode_v2_rope_fa_rows(
24753 ws_index,
24754 e,
24755 &tp.o,
24756 &session_parts,
24757 &tab_keys,
24758 &positions,
24759 stage_pos,
24760 true,
24761 &attention.q_norm,
24762 &attention.k_norm,
24763 &rope_freqs,
24764 t,
24765 head_dim,
24766 geometry.n_rot as usize,
24767 window.unwrap_or(0),
24768 max_ns,
24769 geometry.attention_scale(),
24770 k_tok_bytes,
24771 v_tok_bytes,
24772 self.cfg.rms_eps,
24773 geometry.rope_base,
24774 )
24775 .map(Some)
24776 }
24777
24778 #[allow(clippy::too_many_arguments)]
24783 pub(crate) fn step35_batch_rope_fa_pass(
24784 &self,
24785 e: &Engine,
24786 il: usize,
24787 caches: &[&mut Cache],
24788 row_to_cache: impl Fn(usize) -> usize,
24789 positions: &[i32],
24790 t: usize,
24791 stage_pos: bool,
24792 ) -> Result<Option<CudaSlice<f32>>, Box<dyn std::error::Error>> {
24793 use cudarc::driver::DevicePtr;
24794 if !crate::tp::fuse_rope_append_on() {
24795 return Ok(None);
24796 }
24797 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
24798 return Ok(None);
24799 };
24800 let Some(tp) = fa.step_tp_qkv.as_ref() else {
24801 return Ok(None);
24802 };
24803 let Some(attention) = tp.attention.as_ref() else {
24804 return Ok(None);
24805 };
24806 let geometry = self.step35_geom(il);
24807 let head_dim = geometry.head_dim_k as usize;
24808 if head_dim != 128 {
24809 return Ok(None);
24810 }
24811 let heads = geometry.n_head as usize;
24812 let window = geometry.window.map(|w| w as usize);
24813 let ranks = tp.runtime.devices().len();
24814 for r in 0..t {
24817 let cache = &caches[row_to_cache(r)];
24818 let Some(distributed) = cache.tp_kv[il].as_ref() else {
24819 return Ok(None);
24820 };
24821 if distributed.kv_dim_k() != distributed.kv_dim_v() {
24822 return Ok(None);
24823 }
24824 let rank0 = distributed.rank(0).ok_or("rope fa pass lost rank 0")?;
24825 if rank0.base_d().is_none()
24826 && distributed.staged_len() + t > distributed.physical_capacity()
24827 {
24828 return Ok(None);
24829 }
24830 }
24831 let mut rope_freqs = Vec::with_capacity(ranks);
24832 for rank in 0..ranks {
24833 let engine = tp
24834 .runtime
24835 .rank_engine(rank)
24836 .ok_or("rope fa pass lost a rank engine")?;
24837 rope_freqs.push(if geometry.rope_factors {
24838 match self
24839 .step35_aux
24840 .as_ref()
24841 .and_then(|aux| aux.rope_freqs(engine))
24842 {
24843 Some(f) => Some(f),
24844 None => return Ok(None),
24845 }
24846 } else {
24847 None
24848 });
24849 }
24850 let ladder = |t_kv: usize| -> usize {
24851 if t_kv <= 2048 {
24852 16
24853 } else if t_kv <= 16384 {
24854 64
24855 } else {
24856 128
24857 }
24858 };
24859 let (mut max_ns, mut k_tok_bytes, mut v_tok_bytes) = (1usize, 0usize, 0usize);
24860 let mut session_parts: Vec<Vec<[u64; 4]>> = vec![Vec::with_capacity(t); ranks];
24861 let mut tab_keys = vec![0u64; ranks];
24862 for (r, &pos) in positions.iter().enumerate().take(t) {
24863 let cache = &caches[row_to_cache(r)];
24864 let distributed = cache.tp_kv[il]
24865 .as_ref()
24866 .ok_or("rope fa pass lost a distributed KV cache")?;
24867 (k_tok_bytes, v_tok_bytes) = (distributed.k_tok_bytes(), distributed.v_tok_bytes());
24868 let t_kv = window
24869 .map(|w| (pos as usize + 1).min(w))
24870 .unwrap_or(pos as usize + 1);
24871 max_ns = max_ns.max(t_kv.div_ceil(ladder(t_kv)));
24872 for rank in 0..ranks {
24873 let engine = tp
24874 .runtime
24875 .rank_engine(rank)
24876 .ok_or("rope fa pass lost a rank engine")?;
24877 let rank_cache = distributed
24878 .rank(rank)
24879 .ok_or("rope fa pass lost a KV cache rank")?;
24880 let _main = engine.gpu.enter_main()?;
24881 let s = engine.stream();
24882 let (kp, _g0) = rank_cache.k().device_ptr(&s);
24883 let (vp, _g1) = rank_cache.v().device_ptr(&s);
24884 let (lp, _g2) = rank_cache.len_d().device_ptr(&s);
24885 let bp = match rank_cache.base_d() {
24886 Some(b) => {
24887 let (p, _g) = b.device_ptr(&s);
24888 p
24889 }
24890 None => 0u64,
24891 };
24892 tab_keys[rank] = tab_keys[rank]
24893 .rotate_left(9)
24894 .wrapping_add(kp)
24895 .wrapping_add(bp)
24896 .wrapping_add(il as u64);
24897 session_parts[rank].push([kp, vp, lp, bp]);
24898 }
24899 }
24900 let ws_index = tp
24901 .runtime
24902 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
24903 tp.runtime
24904 .decode_v2_rope_fa_rows(
24905 ws_index,
24906 e,
24907 &tp.o,
24908 &session_parts,
24909 &tab_keys,
24910 positions,
24911 stage_pos,
24912 false,
24913 &attention.q_norm,
24914 &attention.k_norm,
24915 &rope_freqs,
24916 t,
24917 head_dim,
24918 geometry.n_rot as usize,
24919 window.unwrap_or(0),
24920 max_ns,
24921 geometry.attention_scale(),
24922 k_tok_bytes,
24923 v_tok_bytes,
24924 self.cfg.rms_eps,
24925 geometry.rope_base,
24926 )
24927 .map(Some)
24928 }
24929
24930 #[allow(clippy::too_many_arguments)]
24934 pub(crate) fn step35_batch_fa_rows_join(
24935 &self,
24936 e: &Engine,
24937 il: usize,
24938 caches: &[&mut Cache],
24939 row_to_cache: impl Fn(usize) -> usize,
24940 positions: &[i32],
24941 t: usize,
24942 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
24943 use cudarc::driver::DevicePtr;
24944 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
24945 return Err("batch fa rows join expects full attention".into());
24946 };
24947 let tp = fa
24948 .step_tp_qkv
24949 .as_ref()
24950 .ok_or("batch fa rows join lost its resident projections")?;
24951 let geometry = self.step35_geom(il);
24952 let heads = geometry.n_head as usize;
24953 let head_dim = geometry.head_dim_k as usize;
24954 let window = geometry.window.map(|w| w as usize);
24955 let ladder = |t_kv: usize| -> usize {
24956 if t_kv <= 2048 {
24957 16
24958 } else if t_kv <= 16384 {
24959 64
24960 } else {
24961 128
24962 }
24963 };
24964 let (mut max_ns, mut k_tok_bytes, mut v_tok_bytes) = (1usize, 0usize, 0usize);
24965 for (r, &pos) in positions.iter().enumerate() {
24966 let cache = &caches[row_to_cache(r)];
24967 let distributed = cache.tp_kv[il]
24968 .as_ref()
24969 .ok_or("batch fa rows join lost a distributed KV cache")?;
24970 (k_tok_bytes, v_tok_bytes) = (distributed.k_tok_bytes(), distributed.v_tok_bytes());
24971 let t_kv = window
24972 .map(|w| (pos as usize + 1).min(w))
24973 .unwrap_or(pos as usize + 1);
24974 max_ns = max_ns.max(t_kv.div_ceil(ladder(t_kv)));
24975 }
24976 let ranks = tp.runtime.devices().len();
24980 let mut tables = Vec::with_capacity(ranks);
24981 for rank in 0..ranks {
24982 let engine = tp
24983 .runtime
24984 .rank_engine(rank)
24985 .ok_or("batch fa rows join lost a rank engine")?;
24986 let _main = engine.gpu.enter_main()?;
24987 let s = engine.stream();
24988 let mut host = Vec::with_capacity(t * 6);
24989 for r in 0..t {
24990 let cache = &caches[row_to_cache(r)];
24991 let distributed = cache.tp_kv[il]
24992 .as_ref()
24993 .ok_or("batch fa rows join lost a distributed KV cache")?;
24994 let rank_cache = distributed
24995 .rank(rank)
24996 .ok_or("batch fa rows join lost a KV cache rank")?;
24997 let (kp, _g0) = rank_cache.k().device_ptr(&s);
24998 let (vp, _g1) = rank_cache.v().device_ptr(&s);
24999 let (lp, _g2) = rank_cache.len_d().device_ptr(&s);
25000 let bp = match rank_cache.base_d() {
25001 Some(b) => {
25002 let (p, _g) = b.device_ptr(&s);
25003 p
25004 }
25005 None => 0u64,
25006 };
25007 host.extend_from_slice(&[kp, vp, lp, bp, 0u64, 0u64]);
25008 }
25009 tables.push(engine.stream().clone_htod(&host)?);
25010 }
25011 let tabs: Vec<&CudaSlice<u64>> = tables.iter().collect();
25012 let ws_index = tp
25013 .runtime
25014 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
25015 tp.runtime.decode_v2_fa_rows_join(
25016 ws_index,
25017 e,
25018 &tp.o,
25019 &tabs,
25020 t,
25021 head_dim,
25022 window.unwrap_or(0),
25023 max_ns,
25024 geometry.attention_scale(),
25025 k_tok_bytes,
25026 v_tok_bytes,
25027 )
25028 }
25029
25030 #[allow(dead_code)] pub(crate) fn step35_verify_spec_fa2_join(
25035 &self,
25036 e: &Engine,
25037 il: usize,
25038 cache: &Cache,
25039 pos0: usize,
25040 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
25041 let crate::hybrid::Mixer::Full(fa) = &self.layers[il].mixer else {
25042 return Err("spec fa2 join expects full attention".into());
25043 };
25044 let tp = fa
25045 .step_tp_qkv
25046 .as_ref()
25047 .ok_or("spec fa2 join lost its resident projections")?;
25048 let geometry = self.step35_geom(il);
25049 let heads = geometry.n_head as usize;
25050 let head_dim = geometry.head_dim_k as usize;
25051 let window = geometry.window.map(|w| w as usize);
25052 let bucket = window.map(|w| (pos0 + 2).min(w)).unwrap_or(pos0 + 2);
25055 let distributed = cache.tp_kv[il]
25056 .as_ref()
25057 .ok_or("spec fa2 join lost its distributed KV cache")?;
25058 let ws_index = tp
25059 .runtime
25060 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
25061 tp.runtime.decode_v2_spec_fa2_join(
25062 ws_index,
25063 e,
25064 &tp.o,
25065 distributed,
25066 head_dim,
25067 window.unwrap_or(0),
25068 bucket,
25069 geometry.attention_scale(),
25070 )
25071 }
25072
25073 pub(crate) fn step35_tp_decode_attn_resident_v2(
25074 &self,
25075 e: &Engine,
25076 fa: &FullAttnLayer,
25077 il: usize,
25078 h: &CudaSlice<f32>,
25079 pos_d: &CudaSlice<i32>,
25080 cache: &mut Cache,
25081 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
25082 let tp = fa
25083 .step_tp_qkv
25084 .as_ref()
25085 .ok_or("Step TP decode lost its resident projections")?;
25086 let attention = tp
25087 .attention
25088 .as_ref()
25089 .ok_or("Step TP decode lost its resident attention auxiliaries")?;
25090 if !tp.runtime.native_p2p() {
25091 return Err("rank-local Step attention requires native P2P".into());
25092 }
25093 if crate::Engine::kv_fp8_on() {
25094 return Err("rank-local Step attention has not qualified the FP8 KV cache".into());
25095 }
25096
25097 let geometry = self.cfg.full_attention_geometry_at(il as u32);
25098 let window = geometry.window.map(|window| window as usize);
25099 let ranks = tp.runtime.devices().len();
25100 let head_dim = geometry.head_dim_k as usize;
25101 let heads = geometry.n_head as usize;
25102 let kv_heads = geometry.n_head_kv as usize;
25103 if !heads.is_multiple_of(ranks) || !kv_heads.is_multiple_of(ranks) {
25104 return Err(format!(
25105 "Step attention heads q={heads} kv={kv_heads} are not divisible by TP={ranks}"
25106 )
25107 .into());
25108 }
25109 let local_heads = heads / ranks;
25110 let local_kv_heads = kv_heads / ranks;
25111 let max_ctx = cache.max_ctx;
25112
25113 let hydrated = self.ensure_step_tp_kv_cache(e, fa, il, cache)?;
25114
25115 let base_len = cache.kv[il]
25116 .as_ref()
25117 .ok_or_else(|| format!("Step TP layer {il} lost its local KV cache"))?
25118 .len;
25119 {
25120 let distributed = cache.tp_kv[il]
25121 .as_ref()
25122 .ok_or_else(|| format!("Step TP layer {il} lost its distributed KV cache"))?;
25123 if distributed.committed_len() != base_len || distributed.staged_len() != base_len {
25124 return Err(format!(
25125 "Step TP layer {il} cache lengths diverged before decode: \
25126 local={base_len} distributed={}/{}",
25127 distributed.committed_len(),
25128 distributed.staged_len()
25129 )
25130 .into());
25131 }
25132 }
25133 if pos_d.len() != 1 {
25134 return Err(format!(
25135 "rank-local Step decode requires one position, got {}",
25136 pos_d.len()
25137 )
25138 .into());
25139 }
25140
25141 let decode_input = attention
25142 .decode_input
25143 .as_ref()
25144 .ok_or("Step TP decode v2 requires the replicated decode input")?;
25145 let mut decode_input = decode_input
25146 .lock()
25147 .map_err(|_| "Step TP replicated decode input lock is poisoned")?;
25148
25149 let has_gate = fa.attn_gate.is_some();
25150 let use_gate_shards = has_gate
25155 && (attention.gate_shards.is_some() || attention.gate_shards_bf16.is_some())
25156 && crate::tp::step_tp_qkv_fused_enabled()?;
25157 let gate_raw = if !has_gate || use_gate_shards {
25158 None
25159 } else {
25160 let gate_weight = fa
25161 .attn_gate
25162 .as_ref()
25163 .ok_or("step35 layer is missing attn_gate.weight")?;
25164 let gate_raw = e.matmul(gate_weight, h, 1)?;
25165 if gate_raw.len() != heads {
25166 return Err(format!(
25167 "Step TP layer {il} gate output {} != {heads}",
25168 gate_raw.len()
25169 )
25170 .into());
25171 }
25172 Some(gate_raw)
25173 };
25174
25175 let ws_index = tp
25176 .runtime
25177 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
25178 let mut ws_guard = tp
25179 .runtime
25180 .decode_v2_workspace()
25181 .lock()
25182 .map_err(|_| "Step TP decode v2 workspace lock is poisoned")?;
25183 let ws = ws_guard
25184 .get_mut(ws_index)
25185 .ok_or("Step TP decode v2 workspace missing after ensure")?;
25186
25187 let mut rope_freqs = Vec::with_capacity(ranks);
25188 for rank in 0..ranks {
25189 let engine = tp
25190 .runtime
25191 .rank_engine(rank)
25192 .ok_or_else(|| format!("Step TP layer {il} has no engine for rank {rank}"))?;
25193 rope_freqs.push(if geometry.rope_factors {
25194 self.step35_aux
25195 .as_ref()
25196 .and_then(|aux| aux.rope_freqs(engine))
25197 } else {
25198 None
25199 });
25200 }
25201 let staged_next = base_len + 1;
25208 let t_kv_eff = window
25209 .map(|window| staged_next.min(window))
25210 .unwrap_or(staged_next);
25211 let dcw = crate::tp::step_tp_dcw_enabled()?
25212 && (use_gate_shards || (!has_gate && crate::tp::step_tp_qkv_fused_enabled()?))
25213 && t_kv_eff >= 96
25214 && {
25215 let (write_row, would_rebase) = cache.tp_kv[il]
25216 .as_ref()
25217 .expect("distributed cache checked above")
25218 .peek_append_ring(1)?;
25219 if !would_rebase {
25220 let base = (base_len - write_row) as i32;
25222 let distributed = cache.tp_kv[il]
25223 .as_mut()
25224 .expect("distributed cache checked above");
25225 for rank in 0..ranks {
25226 let engine = tp.runtime.rank_engine(rank).ok_or_else(|| {
25227 format!("Step TP layer {il} has no engine for rank {rank}")
25228 })?;
25229 let _main = engine.gpu.enter_main()?;
25230 let rank_cache = distributed.rank_mut(rank).ok_or_else(|| {
25231 format!("Step TP layer {il} has no KV cache rank {rank}")
25232 })?;
25233 if rank_cache.base_d().is_none() {
25234 rank_cache.arm_base_d(engine.htod_i32(&[base])?);
25235 }
25236 }
25237 }
25238 !would_rebase
25239 };
25240 let fuse_rope = dcw
25241 && crate::tp::fuse_rope_append_on()
25242 && head_dim == 128
25243 && cache.tp_kv[il]
25244 .as_ref()
25245 .map(|d| d.kv_dim_k() == d.kv_dim_v() && d.kv_dim_k() == local_kv_heads * head_dim)
25246 .unwrap_or(false);
25247
25248 let tcol_col = crate::tp::take_verify_tcol();
25249 let fa2_col = crate::tp::take_spec_fa2_defer();
25256 tp.runtime.decode_v2_input_qkv(
25257 ws,
25258 e,
25259 h,
25260 pos_d,
25261 gate_raw.as_ref(),
25262 if !use_gate_shards {
25263 None
25264 } else if let Some(shards) = attention.gate_shards.as_deref() {
25265 Some(crate::tp::StepTpGateShards::F32(shards))
25266 } else {
25267 attention
25268 .gate_shards_bf16
25269 .as_deref()
25270 .map(crate::tp::StepTpGateShards::Bf16)
25271 },
25272 &mut decode_input,
25273 &tp.q,
25274 &tp.k,
25275 &tp.v,
25276 &attention.q_norm,
25277 &attention.k_norm,
25278 head_dim,
25279 geometry.n_rot as usize,
25280 geometry.rope_base,
25281 &rope_freqs,
25282 self.cfg.rms_eps,
25283 has_gate,
25284 fuse_rope,
25285 tcol_col,
25286 )?;
25287
25288 let transaction = cache.tp_kv[il]
25289 .as_mut()
25290 .expect("distributed cache checked above")
25291 .begin_transaction()?;
25292 let append_result = tp.runtime.append_tp_kv_transaction_inner(
25293 cache.tp_kv[il]
25294 .as_mut()
25295 .expect("distributed cache checked above"),
25296 transaction,
25297 &ws.k,
25298 &ws.v_raw,
25299 1,
25300 dcw,
25301 );
25302 if let Err(error) = append_result {
25303 let _ = tp.runtime.rollback_tp_kv_transaction(
25304 cache.tp_kv[il]
25305 .as_mut()
25306 .expect("distributed cache checked above"),
25307 transaction,
25308 );
25309 return Err(error);
25310 }
25311
25312 let staged = (|| -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
25313 let (staged_len, physical, k_tok_bytes_c, v_tok_bytes_c, capacity) = {
25316 let distributed = cache.tp_kv[il]
25317 .as_ref()
25318 .expect("distributed cache checked above");
25319 let staged_len = distributed.staged_len();
25320 let view_start = window
25321 .map(|window| staged_len.saturating_sub(window))
25322 .unwrap_or(0);
25323 (
25324 staged_len,
25325 distributed.physical_range(view_start, staged_len)?,
25326 distributed.k_tok_bytes(),
25327 distributed.v_tok_bytes(),
25328 distributed.physical_capacity(),
25329 )
25330 };
25331 let view_start = window
25332 .map(|window| staged_len.saturating_sub(window))
25333 .unwrap_or(0);
25334 let t_kv = staged_len - view_start;
25335 for rank in 0..ranks {
25336 let engine = tp
25337 .runtime
25338 .rank_engine(rank)
25339 .ok_or_else(|| format!("Step TP layer {il} has no engine for rank {rank}"))?;
25340 let _main = engine.gpu.enter_main()?;
25341 if dcw {
25342 {
25346 let distributed_mut = cache.tp_kv[il]
25347 .as_mut()
25348 .expect("distributed cache checked above");
25349 let (kv_dim_k, kv_dim_v) =
25350 (distributed_mut.kv_dim_k(), distributed_mut.kv_dim_v());
25351 let (k_tok_bytes, v_tok_bytes) =
25352 (distributed_mut.k_tok_bytes(), distributed_mut.v_tok_bytes());
25353 let rank_cache = distributed_mut.rank_mut(rank).ok_or_else(|| {
25354 format!("Step TP layer {il} has no KV cache rank {rank}")
25355 })?;
25356 let (k_plane, v_plane, len_d, base_d) =
25357 rank_cache.planes_and_counters_mut();
25358 if fuse_rope {
25359 let same_dev = engine.ctx().ordinal() == e.ctx().ordinal();
25362 let crate::tp::StepTpDecodeV2Ws {
25363 q_raw,
25364 k_raw,
25365 v_raw,
25366 q,
25367 k,
25368 pos,
25369 pos_stage,
25370 fuse_ctr,
25371 ..
25372 } = &mut *ws;
25373 let pos_ref: &CudaSlice<i32> = if same_dev {
25377 pos_stage
25378 .as_ref()
25379 .ok_or("step TP decode v2 pos stage not armed")?
25380 } else {
25381 &pos[rank]
25382 };
25383 engine.qk_norm_rope_append_inc_dcw(
25384 &q_raw[rank],
25385 &k_raw[rank],
25386 &v_raw[rank],
25387 &attention.q_norm[rank],
25388 &attention.k_norm[rank],
25389 &mut q[rank],
25390 &mut k[rank],
25391 pos_ref,
25392 k_plane,
25393 v_plane,
25394 len_d,
25395 base_d,
25396 &mut fuse_ctr[rank],
25397 kv_dim_k,
25398 kv_dim_v,
25399 k_tok_bytes,
25400 v_tok_bytes,
25401 head_dim,
25402 geometry.n_rot as usize,
25403 local_heads,
25404 local_kv_heads,
25405 self.cfg.rms_eps,
25406 geometry.rope_base,
25407 1.0,
25408 rope_freqs[rank],
25409 )?;
25410 } else {
25411 engine.append_kv_quantized_dcw(
25412 &ws.k[rank],
25413 &ws.v_raw[rank],
25414 k_plane,
25415 v_plane,
25416 len_d,
25417 base_d,
25418 kv_dim_k,
25419 kv_dim_v,
25420 k_tok_bytes,
25421 v_tok_bytes,
25422 )?;
25423 }
25424 if !fuse_rope {
25425 let rank_cache = distributed_mut.rank_mut(rank).ok_or_else(|| {
25426 format!("Step TP layer {il} has no KV cache rank {rank}")
25427 })?;
25428 engine.inc_i32(rank_cache.len_d_mut())?;
25429 }
25430 }
25431 let distributed = cache.tp_kv[il]
25432 .as_ref()
25433 .expect("distributed cache checked above");
25434 let rank_cache = distributed
25435 .rank(rank)
25436 .ok_or_else(|| format!("Step TP layer {il} has no KV cache rank {rank}"))?;
25437 let k_ring = engine.view_u8_range(rank_cache.k(), 0, capacity * k_tok_bytes_c);
25438 let v_ring = engine.view_u8_range(rank_cache.v(), 0, capacity * v_tok_bytes_c);
25439 if fa2_col.is_some() {
25440 continue;
25443 }
25444 {
25445 let crate::tp::StepTpDecodeV2Ws { q, gate, gated, .. } = &mut *ws;
25448 engine.fa_decode_dcw(
25449 &q[rank],
25450 &k_ring,
25451 &v_ring,
25452 &mut gated[rank],
25453 head_dim,
25454 local_heads,
25455 local_kv_heads,
25456 rank_cache.len_d(),
25457 rank_cache.base_d(),
25458 window.unwrap_or(0),
25459 t_kv,
25460 geometry.attention_scale(),
25461 k_tok_bytes_c,
25462 v_tok_bytes_c,
25463 has_gate.then_some(&gate[rank]),
25464 )?;
25465 }
25466 continue;
25467 }
25468 let distributed = cache.tp_kv[il]
25469 .as_ref()
25470 .expect("distributed cache checked above");
25471 let rank_cache = distributed
25472 .rank(rank)
25473 .ok_or_else(|| format!("Step TP layer {il} has no KV cache rank {rank}"))?;
25474 let k_view = engine.view_u8_range(
25475 rank_cache.k(),
25476 physical.start * k_tok_bytes_c,
25477 physical.end * k_tok_bytes_c,
25478 );
25479 let v_view = engine.view_u8_range(
25480 rank_cache.v(),
25481 physical.start * v_tok_bytes_c,
25482 physical.end * v_tok_bytes_c,
25483 );
25484 if has_gate {
25485 engine.fa_decode_kvmod(
25486 &ws.q[rank],
25487 &k_view,
25488 &v_view,
25489 &mut ws.attn_out[rank],
25490 head_dim,
25491 local_heads,
25492 local_kv_heads,
25493 t_kv,
25494 geometry.attention_scale(),
25495 k_tok_bytes_c,
25496 v_tok_bytes_c,
25497 false,
25498 )?;
25499 engine.attn_head_gate(
25500 &ws.attn_out[rank],
25501 &ws.gate[rank],
25502 &mut ws.gated[rank],
25503 None,
25504 head_dim,
25505 local_heads,
25506 1,
25507 )?;
25508 } else {
25509 engine.fa_decode_kvmod(
25510 &ws.q[rank],
25511 &k_view,
25512 &v_view,
25513 &mut ws.gated[rank],
25514 head_dim,
25515 local_heads,
25516 local_kv_heads,
25517 t_kv,
25518 geometry.attention_scale(),
25519 k_tok_bytes_c,
25520 v_tok_bytes_c,
25521 false,
25522 )?;
25523 }
25524 }
25525
25526 let output = if let Some(col) = fa2_col.filter(|_| dcw) {
25533 tp.runtime.decode_v2_stash_fa2(ws, e, col)?;
25537 crate::tp::set_spec_fa2_stashed();
25538 e.uninit(ws.o_out)?
25539 } else if let Some(col) = crate::tp::take_tcol_oproj_defer() {
25540 if tp.runtime.decode_v2_oproj_tcol_eligible(ws, &tp.o) {
25541 tp.runtime.decode_v2_stash_gated(ws, e, col)?;
25542 crate::tp::set_tcol_oproj_stashed();
25543 e.uninit(ws.o_out)?
25544 } else {
25545 tp.runtime.decode_v2_finish(ws, e, &tp.o)?
25546 }
25547 } else {
25548 tp.runtime.decode_v2_finish(ws, e, &tp.o)?
25549 };
25550
25551 let local = cache.kv[il]
25555 .as_mut()
25556 .ok_or_else(|| format!("Step TP layer {il} lost its local KV cache"))?;
25557 if local.len != base_len || base_len + 1 > max_ctx {
25558 return Err(format!(
25559 "Step TP layer {il} local cache changed during decode: \
25560 len={} base={base_len} max={max_ctx}",
25561 local.len
25562 )
25563 .into());
25564 }
25565 if crate::tp::no_local_shadow_on() {
25566 local.len = base_len + 1;
25569 if !crate::tp::len_mirror_lazy_on() {
25573 e.set_i32_one(&mut local.len_d, local.len as i32)?;
25574 }
25575 } else {
25576 let retain_from = window
25577 .map(|window| {
25578 let staged_retain = (base_len + 1).saturating_sub(window) & !31usize;
25579 let rollback_retain =
25580 base_len.saturating_sub(window.saturating_sub(1)) & !31usize;
25581 staged_retain.min(rollback_retain)
25582 })
25583 .unwrap_or(0);
25584 let write_row = e.prepare_kv_append(local, retain_from, 1)?;
25585 e.append_kv_quantized(
25586 &ws.k_shadow,
25587 &ws.v_shadow,
25588 &mut local.k,
25589 &mut local.v,
25590 write_row,
25591 local.kv_dim_k,
25592 local.kv_dim_v,
25593 local.k_tok_bytes,
25594 local.v_tok_bytes,
25595 false,
25596 )?;
25597 local.len = base_len + 1;
25598 e.set_i32_one(&mut local.len_d, local.len as i32)?;
25599 }
25600 Ok(output)
25601 })();
25602
25603 let output = match staged {
25604 Ok(output) => output,
25605 Err(error) => {
25606 let _ = tp.runtime.rollback_tp_kv_transaction(
25607 cache.tp_kv[il]
25608 .as_mut()
25609 .expect("distributed cache checked above"),
25610 transaction,
25611 );
25612 if let Some(local) = cache.kv[il].as_mut() {
25613 local.len = base_len;
25614 let _ = e.set_i32_one(&mut local.len_d, base_len as i32);
25615 }
25616 return Err(error);
25617 }
25618 };
25619 let lazy_commit = fuse_rope && crate::tp::len_mirror_lazy_on();
25624 if lazy_commit {
25625 if let Err(error) = tp.runtime.commit_tp_kv_transaction_external(
25626 cache.tp_kv[il]
25627 .as_mut()
25628 .expect("distributed cache checked above"),
25629 transaction,
25630 1,
25631 ) {
25632 let _ = tp.runtime.rollback_tp_kv_transaction(
25633 cache.tp_kv[il]
25634 .as_mut()
25635 .expect("distributed cache checked above"),
25636 transaction,
25637 );
25638 let local = cache.kv[il].as_mut().expect("local cache checked above");
25639 local.len = base_len;
25640 e.set_i32_one(&mut local.len_d, base_len as i32)?;
25641 return Err(error);
25642 }
25643 } else if let Err(error) = tp.runtime.commit_tp_kv_transaction(
25644 cache.tp_kv[il]
25645 .as_mut()
25646 .expect("distributed cache checked above"),
25647 transaction,
25648 1,
25649 ) {
25650 let _ = tp.runtime.rollback_tp_kv_transaction(
25651 cache.tp_kv[il]
25652 .as_mut()
25653 .expect("distributed cache checked above"),
25654 transaction,
25655 );
25656 let local = cache.kv[il].as_mut().expect("local cache checked above");
25657 local.len = base_len;
25658 e.set_i32_one(&mut local.len_d, base_len as i32)?;
25659 return Err(error);
25660 }
25661
25662 let committed = cache.tp_kv[il]
25663 .as_ref()
25664 .expect("distributed cache checked above")
25665 .committed_len();
25666 let local_len = cache.kv[il]
25667 .as_ref()
25668 .expect("local cache checked above")
25669 .len;
25670 if committed != local_len {
25671 return Err(format!(
25672 "Step TP layer {il} committed cache length {committed} != local shadow {local_len}"
25673 )
25674 .into());
25675 }
25676 static V2_LOGGED: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false);
25677 if !V2_LOGGED.swap(true, std::sync::atomic::Ordering::Relaxed) {
25678 eprintln!(
25679 "[step-tp-attn-v2] execute layer={} devices={:?} tokens=1 driver=v2 \
25680 qkv_tensor_parallel=true qk_norm_rank_local=true rope_rank_local=true \
25681 kv_cache_distributed=true kv_cache_hydrated={hydrated} \
25682 attention_tensor_parallel=true attention_scope={} \
25683 input_path=root-device-replicated gate={} gate_tensor_parallel={} \
25684 gate_shards={} o_tensor_parallel=true o_reduce=root-device \
25685 local_cache_shadow=true cache_commit=immediate transport={} native_p2p=true \
25686 bulk_p2p={} workspace=persistent ordering=evented output=e-device \
25687 performance_claim=false (logged once; every decode layer runs this driver)",
25688 tp.layer,
25689 tp.devices,
25690 if window.is_some() {
25691 "rank-local-swa-ring"
25692 } else {
25693 "rank-local-global"
25694 },
25695 has_gate,
25696 use_gate_shards,
25697 if use_gate_shards {
25698 "device-staged"
25699 } else if has_gate {
25700 "root-staged"
25701 } else {
25702 "none"
25703 },
25704 tp.runtime.transport_label(),
25705 tp.runtime.bulk_p2p(),
25706 );
25707 }
25708 Ok(output)
25709 }
25710
25711 #[allow(clippy::too_many_arguments)]
25721 pub(crate) fn step35_decode_attn(
25722 &self,
25723 e: &Engine,
25724 fa: &FullAttnLayer,
25725 il: usize,
25726 h: &CudaSlice<f32>,
25727 pre_q: Option<(&CudaSlice<i8>, &CudaSlice<f32>)>,
25728 pos_d: &CudaSlice<i32>,
25729 cache: &mut Cache,
25730 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
25731 if fa
25732 .step_tp_qkv
25733 .as_ref()
25734 .is_some_and(|tp| tp.attention.is_some())
25735 {
25736 if pre_q.is_some() {
25737 return Err(
25738 "rank-local Step attention preserves BF16 activations and refuses the q8_1 \
25739 pre-quantized decode path"
25740 .into(),
25741 );
25742 }
25743 return self.step35_tp_decode_attn_resident(e, fa, il, h, pos_d, cache);
25744 }
25745
25746 let geometry = self.step35_geom(il);
25747 let hd = geometry.head_dim_k as usize;
25748 let nkv = geometry.n_head_kv as usize;
25749 let nh = geometry.n_head as usize;
25750 let rbase = geometry.rope_base;
25751 let scale = geometry.attention_scale();
25752 let swa = geometry.window.is_some();
25753 let eps = self.cfg.rms_eps;
25754 let win = geometry.window.unwrap_or(0) as usize;
25755 let n_rot = geometry.n_rot as usize;
25756 let n_embd = self.cfg.n_embd as usize;
25757 let gw = fa
25758 .attn_gate
25759 .as_ref()
25760 .ok_or("step35 layer is missing attn_gate.weight (head-wise attention gate)")?;
25761
25762 let tp_qkv = if fa.step_tp_qkv.is_some() {
25763 if pre_q.is_some() {
25764 return Err(
25765 "Step Q/K/V TP preserves BF16 activations and refuses the q8_1 \
25766 pre-quantized decode path"
25767 .into(),
25768 );
25769 }
25770 self.full_attn_tp_qkv(e, fa, h, 1)?
25771 } else {
25772 None
25773 };
25774
25775 let (q0, k0, v0, gt) = match tp_qkv {
25776 Some(mut g3) => {
25777 let v = g3.pop().unwrap();
25778 let k = g3.pop().unwrap();
25779 let q = g3.pop().unwrap();
25780 let gt = e.matmul(gw, h, 1)?;
25781 (q, k, v, gt)
25782 }
25783 None => match pre_q {
25784 Some((hq, hdq)) => {
25785 debug_assert!(
25786 e.uses_q8_1_fast(gw),
25787 "step35 pre-quantized decode requires attn_gate on the q8_1 fast path \
25788 (h is a zero-length placeholder here) — see mixer_in_q8_1_fast"
25789 );
25790 let (a, b, c) = match e.matmul_q8_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)? {
25791 Some(t3) => t3,
25792 None => (
25793 e.matmul_pre(&fa.wq, hq, hdq, h, 1)?,
25794 e.matmul_pre(&fa.wk, hq, hdq, h, 1)?,
25795 e.matmul_pre(&fa.wv, hq, hdq, h, 1)?,
25796 ),
25797 };
25798 let gt = e.matmul_pre(gw, hq, hdq, h, 1)?;
25799 (a, b, c, gt)
25800 }
25801 None => {
25802 if e.uses_q8_1_fast(&fa.wq)
25803 && e.uses_q8_1_fast(&fa.wk)
25804 && e.uses_q8_1_fast(&fa.wv)
25805 && e.uses_q8_1_fast(gw)
25806 {
25807 let (hq, hdq) = e.quantize_q8_1(h, 1, n_embd)?;
25808 let (a, b, c) =
25809 match e.matmul_q8_fused3(&fa.wq, &fa.wk, &fa.wv, &hq, &hdq)? {
25810 Some(t3) => t3,
25811 None => (
25812 e.matmul_pre(&fa.wq, &hq, &hdq, h, 1)?,
25813 e.matmul_pre(&fa.wk, &hq, &hdq, h, 1)?,
25814 e.matmul_pre(&fa.wv, &hq, &hdq, h, 1)?,
25815 ),
25816 };
25817 let gt = e.matmul_pre(gw, &hq, &hdq, h, 1)?;
25818 (a, b, c, gt)
25819 } else {
25820 (
25821 e.matmul(&fa.wq, h, 1)?,
25822 e.matmul(&fa.wk, h, 1)?,
25823 e.matmul(&fa.wv, h, 1)?,
25824 e.matmul(gw, h, 1)?,
25825 )
25826 }
25827 }
25828 },
25829 };
25830
25831 let mut q = e.uninit(nh * hd)?;
25832 e.rms_norm(&q0, fa.q_norm.float_data(), &mut q, hd, nh, eps)?;
25833 let mut k = e.uninit(nkv * hd)?;
25834 e.rms_norm(&k0, fa.k_norm.float_data(), &mut k, hd, nkv, eps)?;
25835 let ff = if swa {
25836 None
25837 } else {
25838 self.step35_aux.as_ref().and_then(|a| a.rope_freqs(e))
25839 };
25840 #[cfg(debug_assertions)]
25841 if let Some(ff) = ff {
25842 crate::debug_assert_tensor_stream_device(
25843 ff,
25844 &e.stream(),
25845 "step35_decode_attn.rope_freqs",
25846 );
25847 }
25848 e.rope_neox2(&mut q, &mut k, pos_d, hd, n_rot, nh, nkv, 1, rbase, 1.0, ff)?;
25849
25850 if std::env::var("MEMRA_NOFA").is_ok() {
25851 return Err(
25852 "MEMRA_NOFA (naive f32 SDPA) is incompatible with the quantized KV \
25853 cache; unset MEMRA_NOFA to use fa_decode"
25854 .into(),
25855 );
25856 }
25857 let kvl = cache.kv[il].as_mut().unwrap();
25858 let next_len = kvl.len + 1;
25859 let (off, t_kv) = if swa && next_len > win {
25860 (next_len - win, win)
25861 } else {
25862 (0, next_len)
25863 };
25864 let write_row = e.prepare_kv_append(kvl, off & !31usize, 1)?;
25865 e.append_kv_quantized(
25866 &k,
25867 &v0,
25868 &mut kvl.k,
25869 &mut kvl.v,
25870 write_row,
25871 kvl.kv_dim_k,
25872 kvl.kv_dim_v,
25873 kvl.k_tok_bytes,
25874 kvl.v_tok_bytes,
25875 crate::Engine::kv_fp8_on(),
25876 )?;
25877 kvl.len = next_len;
25878 let physical = kvl.physical_rows(off, off + t_kv)?;
25879 let k_view = e.view_u8_range(
25880 &kvl.k,
25881 physical.start * kvl.k_tok_bytes,
25882 physical.end * kvl.k_tok_bytes,
25883 );
25884 let v_view = e.view_u8_range(
25885 &kvl.v,
25886 physical.start * kvl.v_tok_bytes,
25887 physical.end * kvl.v_tok_bytes,
25888 );
25889 let mut attn = e.uninit(nh * hd)?;
25890 e.fa_decode_kvmod(
25891 &q,
25892 &k_view,
25893 &v_view,
25894 &mut attn,
25895 hd,
25896 nh,
25897 nkv,
25898 t_kv,
25899 scale,
25900 kvl.k_tok_bytes,
25901 kvl.v_tok_bytes,
25902 crate::Engine::kv_fp8_on(),
25903 )?;
25904
25905 let mut ag = e.uninit(nh * hd)?;
25906 e.attn_head_gate(&attn, >, &mut ag, None, hd, nh, 1)?;
25907 self.full_attn_o(e, fa, &ag, 1)
25908 }
25909}
25910
25911impl HybridModel {
25920 pub fn is_gemma4_e4b(&self) -> bool {
25921 self.gemma4_aux.as_ref().is_some_and(|a| a.e4b.is_some())
25922 }
25923
25924 fn gemma4_e4b_geom(&self, il: usize) -> (usize, usize, usize, f32, f32, bool) {
25928 let g = self.cfg.gemma4.as_ref().unwrap();
25929 let swa = g.swa_pattern[il];
25930 let hd = if swa {
25931 g.key_length_swa
25932 } else {
25933 g.key_length_global
25934 } as usize;
25935 let Mixer::Full(fa) = &self.layers[il].mixer else {
25936 panic!("e4b layer {il} not full-attn")
25937 };
25938 let nh = fa.wq.out_features() / hd;
25939 let nkv = fa.wk.out_features() / hd;
25940 (
25941 hd,
25942 nkv,
25943 nh,
25944 if swa {
25945 g.rope_base_swa
25946 } else {
25947 g.rope_base_global
25948 },
25949 1.0,
25950 swa,
25951 )
25952 }
25953
25954 fn gemma4_e4b_kv_target(&self, il: usize) -> Option<usize> {
25956 self.layers[il]
25957 .gemma4
25958 .as_ref()
25959 .and_then(|b| b.e4b.as_ref())
25960 .and_then(|e4| e4.kv_share.map(|t| t as usize))
25961 }
25962
25963 fn gemma4_e4b_inp_pl(
25968 &self,
25969 e: &Engine,
25970 tokens: &[u32],
25971 x_scaled: &CudaSlice<f32>,
25972 t: usize,
25973 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
25974 let tok_d = e.stream().clone_htod(&tokens.to_vec())?;
25975 self.gemma4_e4b_inp_pl_dev(e, &tok_d, x_scaled, t)
25976 }
25977
25978 fn gemma4_e4b_inp_pl_dev(
25980 &self,
25981 e: &Engine,
25982 tok_d: &CudaSlice<u32>,
25983 x_scaled: &CudaSlice<f32>,
25984 t: usize,
25985 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
25986 let aux = self.gemma4_aux.as_ref().unwrap();
25987 let m = aux.e4b.as_ref().unwrap();
25988 let n_embd = self.cfg.n_embd as usize;
25989 let n_layer = self.layers.len();
25990 let width = m.n_epl * n_layer;
25991 let tbl = m.tok_tbl_gpu.get_or_init(|| {
25992 e.upload_u8(&m.tok_embd_bytes)
25993 .expect("e4b per-layer token table upload")
25994 });
25995 let mut a =
25996 e.embed_gather_device_td(tbl, tok_d, t, width, m.tok_embd_qt, m.tok_embd_row_bytes)?;
25997 e.scale_inplace(&mut a, (m.n_epl as f32).sqrt(), t * width)?;
25998 let mut p = e.matmul(&m.model_proj, x_scaled, t)?;
25999 e.scale_inplace(&mut p, 1.0 / (n_embd as f32).sqrt(), t * width)?;
26000 let mut pn = e.uninit(t * width)?;
26001 e.rms_norm(
26002 &p,
26003 m.proj_norm.float_data(),
26004 &mut pn,
26005 m.n_epl,
26006 t * n_layer,
26007 self.cfg.rms_eps,
26008 )?;
26009 let mut out = e.uninit(t * width)?;
26010 e.add_scale(&a, &pn, 1.0 / 2f32.sqrt(), &mut out, t * width)?;
26011 Ok(out)
26012 }
26013
26014 #[allow(clippy::too_many_arguments)]
26019 fn gemma4_e4b_attn(
26020 &self,
26021 e: &Engine,
26022 il: usize,
26023 hq: &CudaSlice<i8>,
26024 hdq: &CudaSlice<f32>,
26025 pos_d: &CudaSlice<i32>,
26026 t: usize,
26027 cache: &mut Cache,
26028 dc_bucket: Option<usize>,
26029 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
26030 let (hd, nkv, nh, base, scale, swa) = self.gemma4_e4b_geom(il);
26031 let eps = self.cfg.rms_eps;
26032 let aux = self.gemma4_aux.as_ref().unwrap();
26033 let ones = aux.ones(e);
26034 #[cfg(debug_assertions)]
26035 crate::debug_assert_tensor_stream_device(ones, &e.stream(), "gemma4_e4b_attn.ones");
26036 let Mixer::Full(fa) = &self.layers[il].mixer else {
26037 unreachable!()
26038 };
26039 let h0 = e.zeros(0)?;
26043 let h = &h0;
26044
26045 let ff = if swa {
26046 None
26047 } else {
26048 Some(
26049 aux.rope_freqs(e)
26050 .expect("e4b global rope needs rope_freqs.weight"),
26051 )
26052 };
26053 #[cfg(debug_assertions)]
26054 if let Some(ff) = ff {
26055 crate::debug_assert_tensor_stream_device(ff, &e.stream(), "gemma4_e4b_attn.rope_freqs");
26056 }
26057 let share = self.gemma4_e4b_kv_target(il);
26058 let mut kv_f32: Option<(CudaSlice<f32>, CudaSlice<f32>)> = None;
26060 let mut q;
26061 if let Some(_tgt) = share {
26062 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
26063 q = e.uninit(t * nh * hd)?;
26064 let mut kdummy = e.uninit(1)?;
26067 let mut vdummy = e.uninit(1)?;
26068 e.rms_norm_qkv_rope(
26069 &q0,
26070 &q0,
26071 &q0,
26072 fa.q_norm.float_data(),
26073 fa.q_norm.float_data(),
26074 ones,
26075 &mut q,
26076 &mut kdummy,
26077 &mut vdummy,
26078 hd,
26079 self.gemma4_rope_dims(il),
26080 nh * t,
26081 0,
26082 pos_d,
26083 nh,
26084 1,
26085 base,
26086 1.0,
26087 ff,
26088 eps,
26089 )?;
26090 } else {
26091 let e4bits = self.layers[il].gemma4.as_ref().and_then(|g| g.e4b.as_ref());
26095 let cat = e4bits.and_then(|e4| e4.qkv_cat.as_ref());
26096 q = e.uninit(t * nh * hd)?;
26097 let mut k = e.uninit(t * nkv * hd)?;
26098 let mut v = e.uninit(t * nkv * hd)?;
26099 if t == 1 && cat.is_some() {
26100 #[allow(clippy::unnecessary_unwrap)]
26101 let qkv0 = e.matmul_pre(cat.unwrap(), hq, hdq, h, 1)?;
26103 e.rms_norm_qkv_rope_cat(
26104 &qkv0,
26105 fa.q_norm.float_data(),
26106 fa.k_norm.float_data(),
26107 ones,
26108 &mut q,
26109 &mut k,
26110 &mut v,
26111 hd,
26112 self.gemma4_rope_dims(il),
26113 nh,
26114 nkv,
26115 pos_d,
26116 nh,
26117 nkv,
26118 base,
26119 1.0,
26120 ff,
26121 eps,
26122 )?;
26123 } else {
26124 let (q0, k0, v0) = match if t == 1 {
26125 e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)?
26126 } else {
26127 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
26130 if *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0")) {
26131 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
26132 } else {
26133 None
26134 }
26135 } {
26136 Some(triple) => triple,
26137 None => (
26138 e.matmul_pre(&fa.wq, hq, hdq, h, t)?,
26139 e.matmul_pre(&fa.wk, hq, hdq, h, t)?,
26140 e.matmul_pre(&fa.wv, hq, hdq, h, t)?,
26141 ), };
26143 e.rms_norm_qkv_rope(
26146 &q0,
26147 &k0,
26148 &v0,
26149 fa.q_norm.float_data(),
26150 fa.k_norm.float_data(),
26151 ones,
26152 &mut q,
26153 &mut k,
26154 &mut v,
26155 hd,
26156 self.gemma4_rope_dims(il),
26157 nh * t,
26158 nkv * t,
26159 pos_d,
26160 nh,
26161 nkv,
26162 base,
26163 1.0,
26164 ff,
26165 eps,
26166 )?;
26167 }
26168 let kvl = cache.kv[il].as_mut().unwrap();
26169 let cls = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
26173 if dc_bucket.is_some() {
26174 debug_assert!(t == 1);
26179 e.append_kv_quantized_row_dc_inc(
26181 &k,
26182 &v,
26183 &mut kvl.k,
26184 &mut kvl.v,
26185 &mut kvl.len_d,
26186 kvl.kv_dim_k,
26187 kvl.kv_dim_v,
26188 kvl.k_tok_bytes,
26189 kvl.v_tok_bytes,
26190 cls,
26191 )?;
26192 } else {
26193 e.append_kv_quantized_rows(
26194 &k,
26195 &v,
26196 &mut kvl.k,
26197 &mut kvl.v,
26198 kvl.len,
26199 t,
26200 kvl.kv_dim_k,
26201 kvl.kv_dim_v,
26202 kvl.k_tok_bytes,
26203 kvl.v_tok_bytes,
26204 cls,
26205 )?;
26206 kvl.len += t;
26207 }
26208 kv_f32 = Some((k, v));
26209 }
26210 let kvl_idx = share.unwrap_or(il);
26213 let kvl = cache.kv[kvl_idx].as_ref().unwrap();
26214 let base_len = kvl.len - t; let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
26216 let mut attn = e.uninit(t * nh * hd)?;
26217 if t > 1 && base_len == 0 && std::env::var("MEMRA_NOFA").is_err() {
26229 if let Some((kf, vf)) = &kv_f32 {
26230 if hd == 256 && t <= win {
26231 e.fa_prefill(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale, true)?;
26232 return e.matmul(&fa.wo, &attn, t);
26233 }
26234 if hd == 256 && swa && t > win {
26235 e.fa_prefill_w(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
26236 return e.matmul(&fa.wo, &attn, t);
26237 }
26238 if hd == 512 && !swa {
26239 e.fa_prefill_hd512(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale, true)?;
26240 return e.matmul(&fa.wo, &attn, t);
26241 }
26242 } else if share.is_some() {
26243 let g = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
26244 let k_view = e.view_u8(&kvl.k, kvl.k.len());
26245 let v_view = e.view_u8(&kvl.v, kvl.v.len());
26246 if hd == 256 && (!swa || t <= win) {
26247 e.fa_prefill_view(
26249 &q,
26250 &k_view,
26251 &v_view,
26252 &mut attn,
26253 hd,
26254 nh,
26255 nkv,
26256 t,
26257 t,
26258 scale,
26259 true,
26260 kvl.k_tok_bytes,
26261 kvl.v_tok_bytes,
26262 g,
26263 )?;
26264 return e.matmul(&fa.wo, &attn, t);
26265 }
26266 let kv_dim = nkv * hd;
26269 let mut kf = e.uninit(t * kv_dim)?;
26270 let mut vf = e.uninit(t * kv_dim)?;
26271 e.fa_dequant_kv_view_f32(
26272 &k_view,
26273 &v_view,
26274 &mut kf,
26275 &mut vf,
26276 kv_dim,
26277 kv_dim,
26278 t,
26279 kvl.k_tok_bytes,
26280 kvl.v_tok_bytes,
26281 g,
26282 )?;
26283 if hd == 512 {
26284 e.fa_prefill_hd512(&q, &kf, &vf, &mut attn, hd, nh, nkv, t, t, scale, true)?;
26285 } else {
26286 e.fa_prefill_w(&q, &kf, &vf, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
26287 }
26288 return e.matmul(&fa.wo, &attn, t);
26289 }
26290 }
26291 if let Some(bucket) = dc_bucket {
26292 assert!(t == 1);
26297 let bucket = if hd == 512 && win <= crate::fa512_min_tkv() {
26303 bucket.min(crate::fa512_min_tkv().saturating_sub(1))
26304 } else {
26305 bucket
26306 };
26307 let k_view = e.view_u8(&kvl.k, kvl.k.len());
26308 let v_view = e.view_u8(&kvl.v, kvl.v.len());
26309 let g = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
26310 if crate::Engine::wpf_level() >= 1 {
26318 e.prefetch_weight_l2(&fa.wo)?;
26319 }
26320 if e.uses_q8_1_fast(&fa.wo) {
26323 let mut oq = e.alloc_i8_uninit(nh * hd)?;
26324 let mut od = e.zeros(nh * hd / 32)?;
26325 e.fa_decode_dc_q8(
26326 &q,
26327 &k_view,
26328 &v_view,
26329 &mut attn,
26330 hd,
26331 nh,
26332 nkv,
26333 &kvl.len_d,
26334 bucket,
26335 scale,
26336 kvl.k_tok_bytes,
26337 kvl.v_tok_bytes,
26338 g,
26339 Some((&mut oq, &mut od)),
26340 )?;
26341 return e.matmul_pre(&fa.wo, &oq, &od, &attn, t);
26342 }
26343 e.fa_decode_dc(
26344 &q,
26345 &k_view,
26346 &v_view,
26347 &mut attn,
26348 hd,
26349 nh,
26350 nkv,
26351 &kvl.len_d,
26352 bucket,
26353 scale,
26354 kvl.k_tok_bytes,
26355 kvl.v_tok_bytes,
26356 g,
26357 )?;
26358 return e.matmul(&fa.wo, &attn, t);
26359 }
26360 for i in 0..t {
26361 let avail = base_len + i + 1;
26362 let (off_tok, t_kv) = if swa && avail > win {
26363 (avail - win, win)
26364 } else {
26365 (0, avail)
26366 };
26367 let k_view = e.view_u8_range(
26368 &kvl.k,
26369 off_tok * kvl.k_tok_bytes,
26370 (off_tok + t_kv) * kvl.k_tok_bytes,
26371 );
26372 let v_view = e.view_u8_range(
26373 &kvl.v,
26374 off_tok * kvl.v_tok_bytes,
26375 (off_tok + t_kv) * kvl.v_tok_bytes,
26376 );
26377 let qv = e.view(&q, t * nh * hd);
26378 let q_row = qv.slice(i * nh * hd..(i + 1) * nh * hd);
26379 let mut q_one = e.uninit(nh * hd)?;
26380 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
26381 let mut a_one = e.uninit(nh * hd)?;
26382 e.fa_decode_kvmod(
26386 &q_one,
26387 &k_view,
26388 &v_view,
26389 &mut a_one,
26390 hd,
26391 nh,
26392 nkv,
26393 t_kv,
26394 scale,
26395 kvl.k_tok_bytes,
26396 kvl.v_tok_bytes,
26397 (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on()),
26398 )?;
26399 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
26400 }
26401 e.matmul(&fa.wo, &attn, t)
26402 }
26403
26404 fn gemma4_e4b_trunk(
26409 &self,
26410 e: &Engine,
26411 tokens: &[u32],
26412 pos0: usize,
26413 cache: &mut Cache,
26414 head_last: bool,
26415 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
26416 let n_embd = self.cfg.n_embd as usize;
26417 let t = tokens.len();
26418 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
26419 let pos_d = e.htod_i32(&pos)?;
26420 let mut x = e.htod(&self.embd.try_gather(n_embd, tokens)?)?;
26421 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
26422 let inp_pl = self.gemma4_e4b_inp_pl(e, tokens, &x, t)?;
26423 self.gemma4_e4b_trunk_core(e, x, inp_pl, &pos_d, t, cache, None, true, head_last)
26424 }
26425
26426 #[allow(clippy::too_many_arguments)] fn gemma4_e4b_trunk_core(
26431 &self,
26432 e: &Engine,
26433 x_in: CudaSlice<f32>,
26434 inp_pl: CudaSlice<f32>,
26435 pos_d: &CudaSlice<i32>,
26436 t: usize,
26437 cache: &mut Cache,
26438 dc_bucket: Option<usize>,
26439 cap_logits: bool,
26440 head_last: bool,
26441 ) -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
26442 let n_embd = self.cfg.n_embd as usize;
26443 let eps = self.cfg.rms_eps;
26444 let n_layer = self.layers.len();
26445 let mut x = x_in;
26446 let aux_e4b = self.gemma4_aux.as_ref().unwrap().e4b.as_ref().unwrap();
26447 let n_epl = aux_e4b.n_epl;
26448
26449 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
26455 for il in 0..n_layer {
26456 let layer = &self.layers[il];
26457 let (hq, hdq) = match h_carry.take() {
26458 Some(p) => p,
26459 None => e.rms_norm_q8_1(&x, layer.attn_norm.float_data(), n_embd, t, eps)?,
26460 };
26461 let o = self.gemma4_e4b_attn(e, il, &hq, &hdq, pos_d, t, cache, dc_bucket)?;
26462 let bits = layer.gemma4.as_ref().unwrap();
26465 let e4b = bits.e4b.as_ref().expect("e4b layer bits");
26466 let fuse_exit = e.uses_q8_1_fast(&e4b.inp_gate);
26477 let (sn, attn_out) = self.gemma4_layer_tail_core_pn(
26478 e,
26479 layer,
26480 &o,
26481 &x,
26482 t,
26483 Some(layer.post_attn_norm.float_data()),
26484 fuse_exit,
26485 )?;
26486 let mut resid = e.uninit(t * n_embd)?;
26487 let g = if fuse_exit {
26493 let (rq, rd) = e.rms_pre_add_q8_1(
26495 &sn,
26496 bits.post_ffw_norm.float_data(),
26497 &attn_out,
26498 &mut resid,
26499 n_embd,
26500 t,
26501 self.cfg.rms_eps,
26502 )?;
26503 e.matmul_pre(&e4b.inp_gate, &rq, &rd, &resid, t)?
26504 } else {
26505 e.add(&sn, &attn_out, &mut resid, t * n_embd)?;
26506 e.matmul(&e4b.inp_gate, &resid, t)?
26507 };
26508 let mut act = e.uninit(t * n_epl)?;
26509 let y = if t == 1 && e.uses_q8_1_fast(&e4b.proj) {
26510 let ipv = e.view(&inp_pl, n_epl * n_layer);
26511 let row = ipv.slice(il * n_epl..(il + 1) * n_epl);
26512 let (aq, ad) = e.gelu_tanh_mul_q8_1(&g, &row, &mut act, n_epl, 1)?;
26513 e.matmul_pre(&e4b.proj, &aq, &ad, &act, t)?
26514 } else {
26515 let mut inp_this = e.uninit(t * n_epl)?;
26516 e.copy_rows_strided(
26517 &inp_pl,
26518 &mut inp_this,
26519 n_epl,
26520 t,
26521 n_epl * n_layer,
26522 il * n_epl,
26523 )?;
26524 e.gelu_tanh_mul(&g, &inp_this, &mut act, t * n_epl)?;
26525 e.matmul(&e4b.proj, &act, t)?
26526 };
26527 let next_norm = if il + 1 < n_layer {
26530 self.layers[il + 1].attn_norm.float_data()
26531 } else {
26532 self.output_norm.float_data()
26533 };
26534 let mut xn = e.uninit(t * n_embd)?;
26535 let pair = e.rms_pre_add_scale_rms_norm_q8_1(
26536 &y,
26537 e4b.post_norm.float_data(),
26538 &resid,
26539 bits.layer_scale,
26540 next_norm,
26541 &mut xn,
26542 n_embd,
26543 t,
26544 eps,
26545 )?;
26546 h_carry = Some(pair);
26547 x = xn;
26548 }
26549 let (oq, odq) = h_carry.take().unwrap();
26553 let h0 = e.zeros(0)?;
26554 let hm = if head_last { 1 } else { t };
26555 let (hq, hd) = if head_last && t > 1 {
26556 let mut q1 = e.uninit_i8(n_embd)?;
26557 e.dtod_copy_view_i8(&oq.slice((t - 1) * n_embd..t * n_embd), &mut q1)?;
26558 let nb = n_embd / 32;
26559 let mut d1 = e.uninit(nb)?;
26560 e.dtod_copy_view(&odq.slice((t - 1) * nb..t * nb), &mut d1)?;
26561 (q1, d1)
26562 } else {
26563 (oq, odq)
26564 };
26565 let mut ld = e.matmul_pre(&self.output, &hq, &hd, &h0, hm)?;
26566 if cap_logits {
26570 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
26571 e.softcap(&mut ld, cap, hm * self.output.out_features())?;
26572 }
26573 self.gemma4_suppress(e, &mut ld, hm)?; Ok((ld, x))
26575 }
26576
26577 pub fn gemma4_e4b_decode_step_t_am_dev(
26584 &self,
26585 e: &Engine,
26586 tok_d: &CudaSlice<u32>,
26587 t: usize,
26588 pos0: usize,
26589 cache: &mut Cache,
26590 ) -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
26591 let n_embd = self.cfg.n_embd as usize;
26592 let eps = self.cfg.rms_eps;
26593 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
26594 let pos_d = e.htod_i32(&pos)?;
26595 let embd_gpu = self
26596 .embd_gpu
26597 .get_or_init(|| e.upload_u8(&self.embd.raw).expect("embed table upload"));
26598 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
26599 let mut x = e.embed_gather_device_td(embd_gpu, tok_d, t, n_embd, qt, rb)?;
26600 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
26601 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, tok_d, &x, t)?;
26602 let (ld, xp) =
26603 self.gemma4_e4b_trunk_core(e, x, inp_pl, &pos_d, t, cache, None, true, false)?;
26604 let n_vocab = self.output.out_features();
26607 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
26608 for i in 0..t {
26609 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
26610 }
26611 let mut hn = e.uninit(t * n_embd)?;
26612 e.rms_norm(&xp, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
26613 cache.pos += t;
26614 Ok((vam, hn))
26615 }
26616
26617 pub(crate) fn gemma4_e4b_decode_step_t_h(
26620 &self,
26621 e: &Engine,
26622 tokens: &[u32],
26623 pos0: usize,
26624 cache: &mut Cache,
26625 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
26626 let n_embd = self.cfg.n_embd as usize;
26627 let eps = self.cfg.rms_eps;
26628 let t = tokens.len();
26629 let (ld, xp) = self.gemma4_e4b_trunk(e, tokens, pos0, cache, false)?;
26630 let mut hn = e.uninit(t * n_embd)?;
26631 e.rms_norm(&xp, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
26632 cache.pos += t;
26633 Ok((e.dtoh(&ld)?, hn))
26634 }
26635
26636 #[allow(clippy::too_many_arguments)] pub fn gemma4_e4b_decode_step_dcg(
26643 &self,
26644 e: &Engine,
26645 token_d: &mut CudaSlice<u32>,
26646 pos_d: &mut CudaSlice<i32>,
26647 embd_gpu: &CudaSlice<u8>,
26648 embd_qt: i32,
26649 embd_rb: usize,
26650 cache: &mut Cache,
26651 n_vocab: usize,
26652 bucket: usize,
26653 ) -> Result<(), Box<dyn std::error::Error>> {
26654 let n_embd = self.cfg.n_embd as usize;
26655 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
26656 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
26657 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, token_d, &x, 1)?;
26658 let (ld, _x) =
26659 self.gemma4_e4b_trunk_core(e, x, inp_pl, pos_d, 1, cache, Some(bucket), false, false)?;
26660 e.argmax_token_device_into(&ld, token_d, n_vocab)?;
26661 e.inc_seqlen(pos_d)?;
26662 Ok(())
26663 }
26664
26665 #[allow(clippy::too_many_arguments)]
26673 pub fn gemma4_e4b_decode_step_dc(
26674 &self,
26675 e: &Engine,
26676 token_d: &CudaSlice<u32>,
26677 pos_d: &mut CudaSlice<i32>,
26678 embd_gpu: &CudaSlice<u8>,
26679 embd_qt: i32,
26680 embd_rb: usize,
26681 cache: &mut Cache,
26682 n_vocab: usize,
26683 ) -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
26684 let n_embd = self.cfg.n_embd as usize;
26685 let eps = self.cfg.rms_eps;
26686 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
26687 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
26688 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, token_d, &x, 1)?;
26689 let (ld, _x) =
26690 self.gemma4_e4b_trunk_core(e, x, inp_pl, pos_d, 1, cache, None, false, false)?;
26691 let mut tok_out = e.stream().alloc_zeros::<u32>(1)?;
26692 e.argmax_token_device_into(&ld, &mut tok_out, n_vocab)?;
26693 e.inc_seqlen(pos_d)?;
26694 cache.pos += 1;
26695 let _ = eps;
26696 Ok(tok_out)
26697 }
26698
26699 pub(crate) fn gemma4_e4b_decode_step_h(
26702 &self,
26703 e: &Engine,
26704 token: u32,
26705 cache: &mut Cache,
26706 ) -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
26707 let (ld, x) = self.gemma4_e4b_trunk(e, &[token], cache.pos, cache, false)?;
26708 let logits = e.dtoh(&ld)?;
26709 cache.pos += 1;
26710 Ok((logits, x))
26711 }
26712
26713 #[allow(clippy::type_complexity)] pub(crate) fn gemma4_e4b_prime(
26718 &self,
26719 e: &Engine,
26720 tokens: &[u32],
26721 cache: &mut Cache,
26722 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
26723 if cache.pos != 0 {
26726 return Err(
26727 "e4b prime is fresh-prompt only (v0) — prime the full prompt in one \
26728 call or decode tokenwise"
26729 .into(),
26730 );
26731 }
26732 let n_embd = self.cfg.n_embd as usize;
26733 let t = tokens.len();
26734 let (ld, x) = self.gemma4_e4b_trunk(e, tokens, 0, cache, true)?;
26735 cache.pos += t;
26736 let last = e.dtoh(&ld)?; let xv = e.view(&x, t * n_embd);
26738 let row = xv.slice((t - 1) * n_embd..t * n_embd);
26739 let mut h_seed = e.uninit(n_embd)?;
26740 e.copy_view_into(&mut h_seed, 0, &row, n_embd)?;
26741 Ok((last, h_seed, x))
26742 }
26743
26744 pub(crate) fn gemma4_e4b_forward(
26746 &self,
26747 e: &Engine,
26748 tokens: &[u32],
26749 last_only: bool,
26750 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
26751 let mut cache = Cache::new(e, &self.cfg, tokens.len() + 8)?;
26752 let (ld, _x) = self.gemma4_e4b_trunk(e, tokens, 0, &mut cache, last_only)?;
26753 e.dtoh(&ld) }
26755}
26756
26757#[cfg(test)]
26758mod prime_chunk_schedule_tests {
26759 use super::{
26760 CUDA_GRID_YZ_MAX, PRIME_CHUNK_LAUNCH_CAP, PRIME_MIN_T, PRIME_PIPE_MIN_CHUNK, PrimePpSignal,
26761 PrimePpStageChannels, PrimePpWaveCredits, PrimePpWaveSlot, active_matrix_values,
26762 align_prime_ranges_to_gdn, dynamic_prime_chunk_ranges, explicit_prime_chunk,
26763 fixed_prime_chunk_ranges, fixed_prime_chunk_ranges_for_ring, move_prime_cache_layers,
26764 parse_step_ep_grouped_prefill, parse_step_tp_prefill, prime_cache_stage_for_layer,
26765 recv_prime_pp_signal, restore_prime_cache_layers, step_grouped_decode_shape,
26766 step_grouped_prefill_shape, step_tp_prefill_shape, validate_step_prime_batch_modes,
26767 };
26768
26769 fn sizes(ranges: &[(usize, usize)]) -> Vec<usize> {
26770 ranges.iter().map(|(start, end)| end - start).collect()
26771 }
26772
26773 #[allow(clippy::manual_clamp)] fn auto_chunk(t: usize) -> usize {
26775 t.div_ceil(8).max(PRIME_PIPE_MIN_CHUNK).min(4096)
26776 }
26777
26778 #[test]
26785 fn prime_workspace_shape_is_chunk_bounded_and_prompt_scaled_when_monolithic() {
26786 let shape = crate::hybrid_forward::PrimeWorkspaceShape {
26787 call_row_bytes: 42 * 2048 + 14 * 2048,
26788 prompt_row_bytes: 4 * 2048,
26789 n_layers: 41,
26790 };
26791 let chunked_60k = shape.admission_bytes_with_call_rows(60_000, 4096);
26792 let chunked_257k = shape.admission_bytes_with_call_rows(257_000, 4096);
26793 let mono_60k = shape.admission_bytes_with_call_rows(60_000, 60_000);
26794 let mono_257k = shape.admission_bytes_with_call_rows(257_000, 257_000);
26795 assert_eq!(
26797 chunked_257k - chunked_60k,
26798 (257_000 - 60_000) * shape.prompt_row_bytes
26799 );
26800 assert_eq!(
26802 mono_257k - mono_60k,
26803 (257_000 - 60_000) * (shape.call_row_bytes + shape.prompt_row_bytes)
26804 );
26805 assert!(
26806 mono_257k > chunked_257k * 4,
26807 "monolithic 257k must dwarf the chunked charge"
26808 );
26809 assert_eq!(
26811 shape.admission_bytes_with_call_rows(100, 4096),
26812 100 * (shape.call_row_bytes + shape.prompt_row_bytes)
26813 );
26814 }
26815
26816 #[test]
26817 #[allow(clippy::assertions_on_constants)]
26820 fn monolithic_prime_chunk_caps_at_the_cuda_launch_wall() {
26821 assert_eq!(explicit_prime_chunk(0, false), PRIME_CHUNK_LAUNCH_CAP);
26824 assert_eq!(explicit_prime_chunk(100_000, false), PRIME_CHUNK_LAUNCH_CAP);
26825 assert_eq!(explicit_prime_chunk(4096, false), 4096);
26826 assert_eq!(
26827 explicit_prime_chunk(PRIME_CHUNK_LAUNCH_CAP, false),
26828 PRIME_CHUNK_LAUNCH_CAP
26829 );
26830 assert_eq!(
26832 explicit_prime_chunk(0, true),
26833 crate::cache::PRIME_CHUNK_MAX_TOKENS
26834 );
26835 assert_eq!(
26836 explicit_prime_chunk(100_000, true),
26837 crate::cache::PRIME_CHUNK_MAX_TOKENS
26838 );
26839 assert_eq!(explicit_prime_chunk(512, true), 512);
26840 assert!(PRIME_CHUNK_LAUNCH_CAP + PRIME_MIN_T - 1 <= CUDA_GRID_YZ_MAX);
26842 }
26843
26844 #[test]
26845 fn capped_monolithic_ranges_are_identical_below_the_wall_and_legal_above() {
26846 let chunk = explicit_prime_chunk(0, false);
26847 for t in [
26851 PRIME_MIN_T,
26852 4096,
26853 61_000,
26854 64_984,
26855 PRIME_CHUNK_LAUNCH_CAP,
26856 PRIME_CHUNK_LAUNCH_CAP + 1,
26857 CUDA_GRID_YZ_MAX,
26858 ] {
26859 assert_eq!(
26860 fixed_prime_chunk_ranges_for_ring(t, chunk, false),
26861 vec![(0, t)],
26862 "t={t} must stay a single monolithic range"
26863 );
26864 }
26865 for t in [65_536, 65_643, 66_045, 79_717, 82_440, 262_144] {
26869 let ranges = fixed_prime_chunk_ranges_for_ring(t, chunk, false);
26870 assert!(ranges.len() >= 2, "t={t} must chunk");
26871 let mut cursor = 0usize;
26872 for &(start, end) in &ranges {
26873 assert_eq!(start, cursor, "t={t}: ranges must be contiguous");
26874 assert!(
26875 end - start <= CUDA_GRID_YZ_MAX,
26876 "t={t}: range width {} exceeds the CUDA grid.y limit",
26877 end - start
26878 );
26879 assert!(
26880 end - start >= PRIME_MIN_T,
26881 "t={t}: range width {} below PRIME_MIN_T",
26882 end - start
26883 );
26884 cursor = end;
26885 }
26886 assert_eq!(cursor, t, "t={t}: ranges must cover the prompt");
26887 }
26888 for t in (CUDA_GRID_YZ_MAX - 64)..=(CUDA_GRID_YZ_MAX + 2 * PRIME_MIN_T + 64) {
26890 for &(start, end) in &fixed_prime_chunk_ranges_for_ring(t, chunk, false) {
26891 assert!(
26892 end - start <= CUDA_GRID_YZ_MAX,
26893 "t={t} width {}",
26894 end - start
26895 );
26896 }
26897 }
26898 }
26899
26900 #[test]
26901 fn ppn_prime_cache_partition_moves_and_restores_every_layer() {
26902 let round_trip = |fence: &[usize], layers: usize| {
26903 let original: Vec<Option<usize>> = (0..layers).map(Some).collect();
26904 let mut parent = original.clone();
26905 let mut stages: Vec<Vec<Option<usize>>> =
26906 (0..fence.len() - 1).map(|_| vec![None; layers]).collect();
26907
26908 move_prime_cache_layers(&mut parent, &mut stages, fence);
26909 assert!(parent.iter().all(Option::is_none));
26910 for layer in 0..layers {
26911 let owner = prime_cache_stage_for_layer(fence, layer);
26912 for (stage, values) in stages.iter().enumerate() {
26913 assert_eq!(values[layer], (stage == owner).then_some(layer));
26914 }
26915 }
26916
26917 restore_prime_cache_layers(&mut parent, &mut stages, fence);
26918 assert_eq!(parent, original);
26919 assert!(stages.iter().flatten().all(Option::is_none));
26920 };
26921
26922 round_trip(&[0, 5, 8], 10);
26924 round_trip(&[0, 2, 5, 8], 10);
26925 round_trip(&[0, 1, 3, 6, 8], 10);
26926 }
26927
26928 #[test]
26929 fn ppn_prime_wave_credit_requires_the_exact_oldest_wave_and_slot() {
26930 let mut credits = PrimePpWaveCredits::default();
26931 let wave0 = PrimePpWaveSlot { wave: 0, slot: 1 };
26932 let wave1 = PrimePpWaveSlot { wave: 1, slot: 0 };
26933 credits.record_send(wave0).unwrap();
26934 assert_eq!(credits.release_required(), None);
26935 credits.record_send(wave1).unwrap();
26936 assert_eq!(credits.release_required(), Some(wave0));
26937
26938 assert!(
26939 credits
26940 .record_release(PrimePpWaveSlot { wave: 0, slot: 0 })
26941 .unwrap_err()
26942 .contains("does not match oldest pending")
26943 );
26944 assert_eq!(credits.release_required(), Some(wave0));
26945 credits.record_release(wave0).unwrap();
26946 credits
26947 .record_send(PrimePpWaveSlot { wave: 2, slot: 1 })
26948 .unwrap();
26949 assert!(
26950 credits
26951 .record_send(PrimePpWaveSlot { wave: 4, slot: 0 })
26952 .unwrap_err()
26953 .contains("while wave 3 was next")
26954 );
26955 assert!(
26956 credits
26957 .record_send(PrimePpWaveSlot { wave: 3, slot: 1 })
26958 .unwrap_err()
26959 .contains("reused slot 1")
26960 );
26961 }
26962
26963 #[test]
26964 fn ppn_prime_wave_signal_reports_order_error_injected_error_and_closure() {
26965 let expected = PrimePpWaveSlot { wave: 2, slot: 1 };
26966
26967 let (sender, receiver) = std::sync::mpsc::channel();
26968 sender.send(PrimePpSignal::Slot(expected)).unwrap();
26969 assert_eq!(
26970 recv_prime_pp_signal(&receiver, expected, true, "test").unwrap(),
26971 expected
26972 );
26973
26974 let (sender, receiver) = std::sync::mpsc::channel();
26975 sender
26976 .send(PrimePpSignal::Slot(PrimePpWaveSlot { wave: 3, slot: 1 }))
26977 .unwrap();
26978 assert!(
26979 recv_prime_pp_signal(&receiver, expected, true, "test")
26980 .unwrap_err()
26981 .contains("expected wave/slot")
26982 );
26983
26984 let (sender, receiver) = std::sync::mpsc::channel();
26985 sender
26986 .send(PrimePpSignal::Error("injected stage failure".into()))
26987 .unwrap();
26988 assert_eq!(
26989 recv_prime_pp_signal(&receiver, expected, true, "test").unwrap_err(),
26990 "injected stage failure"
26991 );
26992
26993 let (upstream_sender, upstream_receiver) = std::sync::mpsc::channel();
26994 let (outgoing_sender, outgoing_receiver) = std::sync::mpsc::channel();
26995 let (_release_sender, released_downstream) = std::sync::mpsc::channel();
26996 PrimePpStageChannels {
26997 incoming: None,
26998 release_upstream: Some(upstream_sender),
26999 outgoing: outgoing_sender,
27000 released_downstream,
27001 }
27002 .notify_failure("injected worker error");
27003 assert_eq!(
27004 recv_prime_pp_signal(&upstream_receiver, expected, false, "test").unwrap_err(),
27005 "injected worker error"
27006 );
27007 assert_eq!(
27008 recv_prime_pp_signal(&outgoing_receiver, expected, false, "test").unwrap_err(),
27009 "injected worker error"
27010 );
27011
27012 let (sender, receiver) = std::sync::mpsc::channel::<PrimePpSignal>();
27013 drop(sender);
27014 assert!(
27015 recv_prime_pp_signal(&receiver, expected, true, "test")
27016 .unwrap_err()
27017 .contains("channel closed while waiting for wave 2")
27018 );
27019 }
27020
27021 #[test]
27026 fn auto_prime_ranges_align_to_the_gdn_grid() {
27027 let c = 32usize; let assert_covers = |ranges: &[(usize, usize)], t: usize| {
27029 assert_eq!(ranges.first().map(|&(s, _)| s), Some(0));
27030 assert_eq!(ranges.last().map(|&(_, e)| e), Some(t));
27031 for w in ranges.windows(2) {
27032 assert_eq!(w[0].1, w[1].0, "ranges must stay contiguous");
27033 }
27034 assert!(ranges.iter().all(|&(s, e)| e > s), "no empty range");
27035 };
27036
27037 let t = 9510usize;
27040 let fill = auto_chunk(t);
27041 let fixed = fixed_prime_chunk_ranges(t, fill);
27042 assert!(
27043 fixed[..fixed.len() - 1].iter().any(|&(_, e)| e % c != 0),
27044 "broken arm vanished: fixed auto boundaries all landed on-grid"
27045 );
27046 let dynamic = dynamic_prime_chunk_ranges(t, fill, &fixed);
27047 assert!(
27048 dynamic[..dynamic.len() - 1]
27049 .iter()
27050 .any(|&(_, e)| e % c != 0),
27051 "broken arm vanished: dynamic auto boundaries all landed on-grid"
27052 );
27053
27054 for ranges in [&fixed, &dynamic] {
27055 let aligned = align_prime_ranges_to_gdn(ranges, t, c);
27056 assert_covers(&aligned, t);
27057 for &(_, e) in &aligned[..aligned.len() - 1] {
27058 assert_eq!(e % c, 0, "internal boundary {e} off the {c}-grid");
27059 }
27060 for (&(_, a), &(_, b)) in aligned.iter().zip(ranges.iter()) {
27062 assert!(a <= b && b - a < c);
27063 }
27064 }
27065
27066 let tight = vec![(0usize, 33usize), (33, 40), (40, 200)];
27069 let aligned = align_prime_ranges_to_gdn(&tight, 200, c);
27070 assert_covers(&aligned, 200);
27071 assert_eq!(aligned, vec![(0, 32), (32, 200)]);
27072
27073 assert_eq!(align_prime_ranges_to_gdn(&[(0, 200)], 200, c), [(0, 200)]);
27075 assert_eq!(align_prime_ranges_to_gdn(&tight, 200, 0), tight.as_slice());
27076 let on_grid = vec![(0usize, 128usize), (128, 256), (256, 300)];
27077 assert_eq!(
27078 align_prime_ranges_to_gdn(&on_grid, 300, c),
27079 on_grid.as_slice()
27080 );
27081 }
27082
27083 #[test]
27084 fn active_matrix_prefix_scopes_reused_prime_slabs() {
27085 assert_eq!(
27086 active_matrix_values(40 * 4096, 29, 4096, "activation").unwrap(),
27087 29 * 4096
27088 );
27089 assert_eq!(
27090 active_matrix_values(29 * 4096, 29, 4096, "activation").unwrap(),
27091 29 * 4096
27092 );
27093 assert_eq!(
27094 active_matrix_values(29 * 4096, 24, 4096, "activation").unwrap(),
27095 24 * 4096
27096 );
27097 assert!(active_matrix_values(28 * 4096, 29, 4096, "activation").is_err());
27098 assert!(active_matrix_values(usize::MAX, usize::MAX, 2, "activation").is_err());
27099 }
27100
27101 #[test]
27102 fn step_tp_prefill_batch_refuses_before_scheduler_fallback() {
27103 assert!(validate_step_prime_batch_modes(false, false).is_ok());
27104
27105 let grouped_without_tp = validate_step_prime_batch_modes(false, true).unwrap_err();
27106 assert!(grouped_without_tp.contains("requires MEMRA_STEP_TP_PREFILL=1"));
27107
27108 for grouped in [false, true] {
27109 let err = validate_step_prime_batch_modes(true, grouped).unwrap_err();
27110 assert!(err.contains("did not clear the live-server performance gate"));
27111 assert!(err.contains("per-session grouped prefill"));
27112 }
27113 }
27114
27115 #[test]
27116 fn step_grouped_path_is_eager_single_token_only() {
27117 assert!(step_grouped_decode_shape(false, 1));
27118 assert!(!step_grouped_decode_shape(true, 1));
27119 assert!(!step_grouped_decode_shape(false, 2));
27120 assert!(!step_grouped_decode_shape(true, 2));
27121 }
27122
27123 #[test]
27124 fn step_grouped_prefill_door_is_strict_and_capacity_bounded() {
27125 assert!(!parse_step_ep_grouped_prefill(None).unwrap());
27126 assert!(!parse_step_ep_grouped_prefill(Some("")).unwrap());
27127 assert!(!parse_step_ep_grouped_prefill(Some("0")).unwrap());
27128 assert!(parse_step_ep_grouped_prefill(Some("1")).unwrap());
27129 assert!(parse_step_ep_grouped_prefill(Some("true")).is_err());
27130 assert!(parse_step_ep_grouped_prefill(Some("2")).is_err());
27131
27132 assert!(step_grouped_prefill_shape(true, true, PRIME_MIN_T));
27133 assert!(step_grouped_prefill_shape(
27134 true,
27135 true,
27136 crate::cache::PRIME_CHUNK_MAX_TOKENS,
27137 ));
27138 assert!(!step_grouped_prefill_shape(true, true, PRIME_MIN_T - 1,));
27139 assert!(!step_grouped_prefill_shape(
27140 true,
27141 true,
27142 crate::cache::PRIME_CHUNK_MAX_TOKENS + 1,
27143 ));
27144 assert!(!step_grouped_prefill_shape(false, true, PRIME_MIN_T));
27145 assert!(!step_grouped_prefill_shape(true, false, PRIME_MIN_T));
27146 }
27147
27148 #[test]
27149 fn step_tp_prefill_door_is_strict_and_default_off() {
27150 assert!(!parse_step_tp_prefill(None).unwrap());
27151 assert!(!parse_step_tp_prefill(Some("")).unwrap());
27152 assert!(!parse_step_tp_prefill(Some("0")).unwrap());
27153 assert!(parse_step_tp_prefill(Some("1")).unwrap());
27154 assert!(parse_step_tp_prefill(Some("true")).is_err());
27155 assert!(parse_step_tp_prefill(Some("2")).is_err());
27156 }
27157
27158 #[test]
27159 fn step_tp_prefill_requires_a_qualified_even_rank_shape() {
27160 assert!(step_tp_prefill_shape(
27161 true,
27162 PRIME_MIN_T,
27163 4,
27164 true,
27165 true,
27166 false,
27167 ));
27168 assert!(!step_tp_prefill_shape(
27169 false,
27170 PRIME_MIN_T,
27171 4,
27172 true,
27173 true,
27174 false,
27175 ));
27176 assert!(!step_tp_prefill_shape(
27177 true,
27178 PRIME_MIN_T - 1,
27179 4,
27180 true,
27181 true,
27182 false,
27183 ));
27184 assert!(step_tp_prefill_shape(
27186 true,
27187 PRIME_MIN_T,
27188 2,
27189 true,
27190 true,
27191 false
27192 ));
27193 assert!(!step_tp_prefill_shape(
27194 true,
27195 PRIME_MIN_T,
27196 1,
27197 true,
27198 true,
27199 false
27200 ));
27201 assert!(!step_tp_prefill_shape(
27202 true,
27203 PRIME_MIN_T,
27204 3,
27205 true,
27206 true,
27207 false
27208 ));
27209 assert!(!step_tp_prefill_shape(
27210 true,
27211 PRIME_MIN_T,
27212 4,
27213 false,
27214 true,
27215 false,
27216 ));
27217 assert!(!step_tp_prefill_shape(
27218 true,
27219 PRIME_MIN_T,
27220 4,
27221 true,
27222 false,
27223 false,
27224 ));
27225 assert!(!step_tp_prefill_shape(
27226 true,
27227 PRIME_MIN_T,
27228 4,
27229 true,
27230 true,
27231 true,
27232 ));
27233 }
27234
27235 #[test]
27236 fn fixed_schedule_retains_measured_geometry() {
27237 assert_eq!(
27238 sizes(&fixed_prime_chunk_ranges(461, 128)),
27239 vec![128, 128, 128, 77]
27240 );
27241 assert_eq!(
27242 sizes(&fixed_prime_chunk_ranges(1833, 230)),
27243 vec![230, 230, 230, 230, 230, 230, 230, 223]
27244 );
27245 assert_eq!(sizes(&fixed_prime_chunk_ranges(4096, 512)), vec![512; 8]);
27246 let capped = sizes(&fixed_prime_chunk_ranges_for_ring(8200, 4096, true));
27247 assert_eq!(capped, vec![4096, 4088, 16]);
27248 assert!(capped.iter().all(|&rows| rows <= 4096));
27249 assert_eq!(
27250 sizes(&fixed_prime_chunk_ranges_for_ring(4100, 4096, false)),
27251 vec![4100],
27252 "flag-off schedule remains byte-for-byte the legacy monolithic tail",
27253 );
27254 }
27255
27256 #[test]
27257 fn dynamic_schedule_matches_registered_shapes() {
27258 let cases = [
27259 (461, vec![64, 141, 132, 124]),
27260 (1833, vec![115, 269, 260, 252, 244, 237, 231, 225]),
27261 (4096, vec![256, 602, 580, 563, 545, 531, 516, 503]),
27262 ];
27263 for (t, expected) in cases {
27264 let chunk = auto_chunk(t);
27265 let fixed = fixed_prime_chunk_ranges(t, chunk);
27266 assert_eq!(
27267 sizes(&dynamic_prime_chunk_ranges(t, chunk, &fixed)),
27268 expected
27269 );
27270 }
27271 }
27272
27273 #[test]
27274 fn dynamic_schedule_covers_exactly_and_shrinks_after_fill() {
27275 for t in 256..=8192 {
27276 let chunk = auto_chunk(t);
27277 let fixed = fixed_prime_chunk_ranges(t, chunk);
27278 let dynamic = dynamic_prime_chunk_ranges(t, chunk, &fixed);
27279 assert_eq!(dynamic.len(), fixed.len(), "T={t}");
27280 assert_eq!(dynamic.first().unwrap().0, 0, "T={t}");
27281 assert_eq!(dynamic.last().unwrap().1, t, "T={t}");
27282 for pair in dynamic.windows(2) {
27283 assert_eq!(pair[0].1, pair[1].0, "T={t}");
27284 }
27285 assert!(
27286 dynamic
27287 .iter()
27288 .all(|(start, end)| end - start >= PRIME_MIN_T),
27289 "T={t} sizes={:?}",
27290 sizes(&dynamic)
27291 );
27292 if dynamic.len() >= 3 {
27293 let chunk_sizes = sizes(&dynamic);
27294 assert!(
27295 chunk_sizes[0] < chunk_sizes[1],
27296 "T={t} sizes={chunk_sizes:?}"
27297 );
27298 assert!(
27299 chunk_sizes[1..].windows(2).all(|pair| pair[0] >= pair[1]),
27300 "T={t} sizes={chunk_sizes:?}"
27301 );
27302 }
27303 }
27304 }
27305}
27306
27307#[cfg(test)]
27308mod page_prefetch_tests {
27309 use super::{
27310 grouped_worker_prefetch_position, page_prefetch_positions,
27311 page_prefetch_window_from_values, worker_prefetch_positions,
27312 };
27313
27314 #[test]
27315 fn page_prefetch_window_keeps_existing_opt_in_default() {
27316 assert_eq!(page_prefetch_window_from_values(false, None), 0);
27317 assert_eq!(page_prefetch_window_from_values(false, Some("8")), 0);
27318 assert_eq!(page_prefetch_window_from_values(true, None), 1);
27319 assert_eq!(page_prefetch_window_from_values(true, Some("bad")), 1);
27320 assert_eq!(page_prefetch_window_from_values(true, Some("0")), 0);
27321 assert_eq!(page_prefetch_window_from_values(true, Some("8")), 8);
27322 }
27323
27324 #[test]
27325 fn rolling_page_prefetch_advises_each_future_expert_once() {
27326 let advised: Vec<_> = (0..7)
27327 .flat_map(|position| page_prefetch_positions(position, 7, 3))
27328 .collect();
27329 assert_eq!(advised, vec![1, 2, 3, 4, 5, 6]);
27330
27331 let one_ahead: Vec<_> = (0..4)
27332 .flat_map(|position| page_prefetch_positions(position, 4, 1))
27333 .collect();
27334 assert_eq!(one_ahead, vec![1, 2, 3]);
27335 assert!(page_prefetch_positions(0, 4, 0).is_empty());
27336 }
27337
27338 #[test]
27339 fn grouped_worker_prefetch_primes_first_then_each_known_next_once() {
27340 assert_eq!(grouped_worker_prefetch_position(0, None), None);
27341 let positions: Vec<_> = std::iter::once(grouped_worker_prefetch_position(4, None).unwrap())
27342 .chain(
27343 (0..4).filter_map(|position| grouped_worker_prefetch_position(4, Some(position))),
27344 )
27345 .collect();
27346 assert_eq!(positions, vec![0, 1, 2, 3]);
27347 assert_eq!(grouped_worker_prefetch_position(1, Some(0)), None);
27348 }
27349
27350 #[test]
27351 fn rolling_worker_prefetch_primes_current_and_each_future_expert_once() {
27352 let queued: Vec<_> = (0..8)
27353 .flat_map(|position| worker_prefetch_positions(position, 8, 5))
27354 .collect();
27355 assert_eq!(queued, (0..8).collect::<Vec<_>>());
27356
27357 let one_at_a_time: Vec<_> = (0..4)
27358 .flat_map(|position| worker_prefetch_positions(position, 4, 1))
27359 .collect();
27360 assert_eq!(one_at_a_time, vec![0, 1, 2, 3]);
27361 assert!(worker_prefetch_positions(0, 4, 0).is_empty());
27362 }
27363}
27364
27365pub struct G4DcSlots {
27366 x: CudaSlice<f32>,
27367 xn: CudaSlice<f32>,
27368 cur: CudaSlice<f32>,
27369 hq: CudaSlice<i8>,
27370 hd_: CudaSlice<f32>,
27371 q0: CudaSlice<f32>,
27372 k0: CudaSlice<f32>,
27373 v0: CudaSlice<f32>,
27374 q: CudaSlice<f32>,
27375 k: CudaSlice<f32>,
27376 v: CudaSlice<f32>,
27377 attn: CudaSlice<f32>,
27378 o: CudaSlice<f32>,
27379 attn_out: CudaSlice<f32>,
27380 zsh: CudaSlice<f32>,
27381 zq: CudaSlice<i8>,
27382 zd: CudaSlice<f32>,
27383 gate: CudaSlice<f32>,
27384 up: CudaSlice<f32>,
27385 act: CudaSlice<f32>,
27386 actq: CudaSlice<i8>,
27387 actd: CudaSlice<f32>,
27388 f0: CudaSlice<f32>,
27389 sn: CudaSlice<f32>,
27390 hn: CudaSlice<f32>,
27391 logits: CudaSlice<f32>,
27392}
27393
27394pub struct Step35TokenGraphState {
27399 pub graphs: Vec<(usize, crate::tp::TokenGraph)>,
27401 pub token_d: cudarc::driver::CudaSlice<u32>,
27402 pub pos_d: cudarc::driver::CudaSlice<i32>,
27403 pub logits_stage: cudarc::driver::CudaSlice<f32>,
27404 pub x: cudarc::driver::CudaSlice<f32>,
27409 pub x1: cudarc::driver::CudaSlice<f32>,
27410 pub mixed_stage: cudarc::driver::CudaSlice<f32>,
27411 pub sh_stage: cudarc::driver::CudaSlice<f32>,
27412 pub k_shadow_stage: cudarc::driver::CudaSlice<f32>,
27413 pub v_shadow_stage: cudarc::driver::CudaSlice<f32>,
27414 pub router_logits: cudarc::driver::CudaSlice<f32>,
27417 pub shexp_gate: cudarc::driver::CudaSlice<f32>,
27418 pub shexp_up: cudarc::driver::CudaSlice<f32>,
27419 pub shexp_act: cudarc::driver::CudaSlice<f32>,
27420 pub gate_sig: cudarc::driver::CudaSlice<f32>,
27421 pub dense_z: cudarc::driver::CudaSlice<f32>,
27422 pub dense_gate: cudarc::driver::CudaSlice<f32>,
27423 pub dense_up: cudarc::driver::CudaSlice<f32>,
27424 pub dense_act: cudarc::driver::CudaSlice<f32>,
27425 pub hn: cudarc::driver::CudaSlice<f32>,
27426 pub probe_mixed: cudarc::driver::CudaSlice<f32>,
27429 pub probe_x: cudarc::driver::CudaSlice<f32>,
27430 pub token_hist: cudarc::driver::CudaSlice<u32>,
27433 pub hist_idx: cudarc::driver::CudaSlice<i32>,
27434}
27435
27436impl HybridModel {
27437 #[allow(clippy::type_complexity)] pub(crate) fn step35_token_graph_step(
27449 &self,
27450 e: &Engine,
27451 token: u32,
27452 cache: &mut Cache,
27453 ) -> Result<Option<(Vec<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
27454 if !self.uses_sliding_gated_moe_program()
27455 || !crate::tp::step_tp_graph_enabled()?
27456 || !crate::tp::step_tp_dcw_enabled()?
27457 || !crate::tp::step_tp_qkv_fused_enabled()?
27458 || !crate::tp::step_tp_dev_router_enabled()?
27459 || !crate::tp::step_nvfp4_dev_routes_enabled()?
27460 {
27461 return Ok(None);
27462 }
27463 if !crate::spec::graph_launch_headroom_ok(e) {
27469 static NOTED: std::sync::Once = std::sync::Once::new();
27470 NOTED.call_once(|| crate::spec::graph_replay_suspended_note("step-tp-token"));
27471 return Ok(None);
27472 }
27473 let n_embd = self.cfg.n_embd as usize;
27474 let n_vocab = self.cfg.n_vocab as usize;
27475 let n_layers = self.layers.len();
27476 let pos = cache.pos;
27477 let staged_next = pos + 1;
27478 if staged_next < 96 {
27479 return Ok(None); }
27481
27482 for il in 0..n_layers {
27485 let Some(tp_kv) = cache.tp_kv[il].as_ref() else {
27486 return Ok(None); };
27488 if tp_kv.peek_append_ring(1)?.1 {
27489 return Ok(None);
27490 }
27491 }
27492
27493 let (fa_vec, n_splits) = e.fa_geom_eager(staged_next, 128, 8, false);
27496 if !fa_vec {
27497 return Ok(None);
27498 }
27499 let sp = crate::fa_split_keys(staged_next, 8);
27500 let bucket_max = (n_splits * sp).max(staged_next);
27501
27502 let mut state_guard = self
27503 .step35_token_graph
27504 .lock()
27505 .map_err(|_| "step35 token graph lock is poisoned")?;
27506 if state_guard.is_none() {
27507 let _main = e.gpu.enter_main()?;
27508 let n_expert = self
27509 .cfg
27510 .moe
27511 .as_ref()
27512 .map(|m| m.expert_count as usize)
27513 .unwrap_or(0);
27514 let n_ff_sh = self
27515 .layers
27516 .iter()
27517 .find_map(|l| match &l.ffn {
27518 crate::hybrid::Ffn::Moe(m) => m.gate_shexp.as_ref().map(|g| g.out_features()),
27519 _ => None,
27520 })
27521 .unwrap_or(0);
27522 let n_ff_dense = self
27523 .layers
27524 .iter()
27525 .find_map(|l| match &l.ffn {
27526 crate::hybrid::Ffn::Dense { ffn_gate, .. } => Some(ffn_gate.out_features()),
27527 _ => None,
27528 })
27529 .unwrap_or(0);
27530 *state_guard = Some(Step35TokenGraphState {
27531 graphs: Vec::new(),
27532 token_d: e.stream().clone_htod(&[0u32])?,
27533 pos_d: e.htod_i32(&[pos as i32])?,
27534 logits_stage: e.htod(&vec![0.0f32; n_vocab])?,
27535 x: e.htod(&vec![0.0f32; n_embd])?,
27536 x1: e.htod(&vec![0.0f32; n_embd])?,
27537 mixed_stage: e.htod(&vec![0.0f32; n_embd])?,
27538 sh_stage: e.htod(&vec![0.0f32; n_embd])?,
27539 k_shadow_stage: e.htod(&vec![0.0f32; 2048])?,
27540 v_shadow_stage: e.htod(&vec![0.0f32; 2048])?,
27541 router_logits: e.htod(&vec![0.0f32; n_expert.max(1)])?,
27542 shexp_gate: e.htod(&vec![0.0f32; n_ff_sh.max(1)])?,
27543 shexp_up: e.htod(&vec![0.0f32; n_ff_sh.max(1)])?,
27544 shexp_act: e.htod(&vec![0.0f32; n_ff_sh.max(1)])?,
27545 gate_sig: e.htod(&[1.0f32; 1])?,
27546 dense_z: e.htod(&vec![0.0f32; n_embd])?,
27547 dense_gate: e.htod(&vec![0.0f32; n_ff_dense.max(1)])?,
27548 dense_up: e.htod(&vec![0.0f32; n_ff_dense.max(1)])?,
27549 dense_act: e.htod(&vec![0.0f32; n_ff_dense.max(1)])?,
27550 hn: e.htod(&vec![0.0f32; n_embd])?,
27551 probe_mixed: e.htod(&vec![0.0f32; n_embd])?,
27552 probe_x: e.htod(&vec![0.0f32; n_embd])?,
27553 token_hist: e.stream().clone_htod(&[0u32; 16])?,
27554 hist_idx: e.htod_i32(&[0])?,
27555 });
27556 }
27557 let state = state_guard.as_mut().expect("state armed above");
27558 {
27562 let _main = e.gpu.enter_main()?;
27563 let Step35TokenGraphState {
27564 logits_stage,
27565 token_d,
27566 ..
27567 } = &mut *state;
27568 e.argmax_token_device_into(logits_stage, token_d, n_vocab)?;
27569 }
27570
27571 if state.graphs.is_empty() {
27576 self.step35_token_graph_build(e, cache, state, bucket_max)?;
27579 }
27580 {
27581 let (b, g) = state.graphs.first_mut().expect("graph built above");
27582 if *b != bucket_max {
27583 g.retarget_bucket(bucket_max)?;
27584 *b = bucket_max;
27585 }
27586 }
27587 let graph = state
27588 .graphs
27589 .first()
27590 .map(|(_, g)| g)
27591 .expect("graph built above");
27592
27593 let tg_timing = std::env::var("MEMRA_STEP_TP_TIMING").as_deref() == Ok("1");
27594 let t_fence = tg_timing.then(std::time::Instant::now);
27595 {
27600 let fa0 = match &self.layers[0].mixer {
27601 Mixer::Full(fa) => fa,
27602 _ => return Err("step35 token graph expects full-attention layers".into()),
27603 };
27604 let tp0 = fa0
27605 .step_tp_qkv
27606 .as_ref()
27607 .ok_or("step35 token graph lost its TP state")?;
27608 for rank in 0..tp0.runtime.devices().len() {
27609 let engine = tp0
27610 .runtime
27611 .rank_engine(rank)
27612 .ok_or("step35 token graph lost a rank engine")?;
27613 let _main = engine.gpu.enter_main()?;
27614 engine.stream().synchronize()?;
27615 }
27616 }
27617
27618 {
27620 let _main = e.gpu.enter_main()?;
27621 e.set_u32_one(&mut state.token_d, token)?;
27622 e.set_i32_one(&mut state.pos_d, pos as i32)?;
27623 }
27624 let t_launch = tg_timing.then(std::time::Instant::now);
27625 graph.launch(e)?;
27626 let t_book = tg_timing.then(std::time::Instant::now);
27627 for il in 0..n_layers {
27632 let tp_kv = cache.tp_kv[il].as_mut().expect("eligibility checked above");
27633 let transaction = tp_kv.begin_transaction()?;
27634 let fa = match &self.layers[il].mixer {
27635 Mixer::Full(fa) => fa,
27636 _ => return Err("step35 token graph expects full-attention layers".into()),
27637 };
27638 let tp = fa
27639 .step_tp_qkv
27640 .as_ref()
27641 .ok_or("step35 token graph lost its TP state")?;
27642 let empty: [CudaSlice<f32>; 0] = [];
27645 tp.runtime.append_tp_kv_transaction_inner(
27646 tp_kv,
27647 transaction,
27648 &empty,
27649 &empty,
27650 1,
27651 true,
27652 )?;
27653 tp.runtime
27654 .commit_tp_kv_transaction_external(tp_kv, transaction, 1)?;
27655 if let Some(local) = cache.kv[il].as_mut() {
27657 local.len = pos + 1;
27658 let _main = e.gpu.enter_main()?;
27659 e.set_i32_one(&mut local.len_d, (pos + 1) as i32)?;
27660 }
27661 }
27662 cache.pos = pos + 1;
27663 let t_sync = tg_timing.then(std::time::Instant::now);
27664 let (logits, h_seed) = {
27665 let _main = e.gpu.enter_main()?;
27666 e.stream().synchronize()?;
27667 (e.dtoh(&state.logits_stage)?, e.clone_dtod(&state.x)?)
27668 };
27669 if let (Some(f), Some(l), Some(b), Some(sy)) = (t_fence, t_launch, t_book, t_sync) {
27670 use std::sync::atomic::{AtomicU64, Ordering};
27671 static NS: [AtomicU64; 5] = [
27672 AtomicU64::new(0),
27673 AtomicU64::new(0),
27674 AtomicU64::new(0),
27675 AtomicU64::new(0),
27676 AtomicU64::new(0),
27677 ];
27678 static CALLS: AtomicU64 = AtomicU64::new(0);
27679 let now = std::time::Instant::now();
27680 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;
27686 if calls.is_multiple_of(100) {
27687 let avg = |i: usize| NS[i].load(Ordering::Relaxed) as f64 / calls as f64 / 1e3;
27688 eprintln!(
27689 "[tg-timing] calls={calls} fence_us={:.0} launch_us={:.0} book_us={:.0} \
27690 syncdtoh_us={:.0} total_us={:.0}",
27691 avg(0),
27692 avg(1),
27693 avg(2),
27694 avg(3),
27695 avg(4)
27696 );
27697 }
27698 }
27699 if std::env::var("MEMRA_TG_PROBE_LAYER").is_ok() {
27701 use std::io::Write;
27702 let (pm, px) = {
27703 let _main = e.gpu.enter_main()?;
27704 (e.dtoh(&state.probe_mixed)?, e.dtoh(&state.probe_x)?)
27705 };
27706 for (path, data) in [
27707 ("/root/tg-probe-mixed.bin", &pm),
27708 ("/root/tg-probe-x.bin", &px),
27709 ] {
27710 let mut fo = std::fs::OpenOptions::new()
27711 .create(true)
27712 .append(true)
27713 .open(path)?;
27714 for v in data {
27715 fo.write_all(&v.to_le_bytes())?;
27716 }
27717 }
27718 }
27719 if let Ok(path) = std::env::var("MEMRA_DUMP_HN") {
27722 let hh = {
27723 let _main = e.gpu.enter_main()?;
27724 e.dtoh(&state.hn)?
27725 };
27726 use std::io::Write;
27727 let mut fo = std::fs::OpenOptions::new()
27728 .create(true)
27729 .append(true)
27730 .open(path)?;
27731 for v in &hh {
27732 fo.write_all(&v.to_le_bytes())?;
27733 }
27734 }
27735 if std::env::var("MEMRA_STEP_TP_GRAPH_DEBUG").as_deref() == Ok("1") {
27738 for il in [0usize, 1, 44] {
27739 let tp_kv = cache.tp_kv[il].as_ref().expect("eligibility checked above");
27740 let host_len = tp_kv.staged_len();
27741 let fa = match &self.layers[il].mixer {
27742 Mixer::Full(fa) => fa,
27743 _ => continue,
27744 };
27745 let tp = fa
27746 .step_tp_qkv
27747 .as_ref()
27748 .ok_or("step35 token graph lost its TP state")?;
27749 for rank in 0..tp.runtime.devices().len() {
27750 let engine = tp
27751 .runtime
27752 .rank_engine(rank)
27753 .ok_or("step35 token graph lost a rank engine")?;
27754 let rank_cache = tp_kv.rank(rank).ok_or("debug rank cache missing")?;
27755 let _main = engine.gpu.enter_main()?;
27756 engine.stream().synchronize()?;
27757 let len_d = engine.dtoh_i32_one(rank_cache.len_d())?;
27758 let base_d = match rank_cache.base_d() {
27759 Some(b) => engine.dtoh_i32_one(b)?,
27760 None => -1,
27761 };
27762 eprintln!(
27763 "[graph-debug] pos={pos} il={il} rank={rank} host_len={host_len} \
27764 len_d={len_d} base_d={base_d}"
27765 );
27766 }
27767 }
27768 }
27769 Ok(Some((logits, h_seed)))
27770 }
27771
27772 pub(crate) fn head_split_matvec(
27778 &self,
27779 e: &Engine,
27780 hn: &CudaSlice<f32>,
27781 ) -> Result<Option<Vec<f32>>, Box<dyn std::error::Error>> {
27782 if self.head_split_fill_device(e, hn)?.is_none() {
27783 return Ok(None);
27784 }
27785 let guard = HEAD_SPLIT_WS
27786 .lock()
27787 .map_err(|_| "head split lock is poisoned")?;
27788 let ws = guard.as_ref().expect("filled above");
27789 let _main = e.gpu.enter_main()?;
27790 Ok(Some(e.dtoh(&ws.logits_e)?))
27791 }
27792
27793 fn head_split_fill_device(
27797 &self,
27798 e: &Engine,
27799 hn: &CudaSlice<f32>,
27800 ) -> Result<Option<()>, Box<dyn std::error::Error>> {
27801 use cudarc::driver::DevicePtr;
27802 let crate::model::GpuTensor::FloatBf16 { data: head, .. } = &self.output else {
27803 return Ok(None);
27804 };
27805 let Some(rank1) = self.layers.first().and_then(|l| match &l.mixer {
27806 Mixer::Full(fa) => fa
27807 .step_tp_qkv
27808 .as_ref()
27809 .and_then(|tp| tp.runtime.rank_engine(1)),
27810 _ => None,
27811 }) else {
27812 return Ok(None);
27813 };
27814 let n_embd = self.cfg.n_embd as usize;
27815 let n_vocab = self.cfg.n_vocab as usize;
27816 let half = n_vocab / 2;
27817 let mut guard = HEAD_SPLIT_WS
27818 .lock()
27819 .map_err(|_| "head split lock is poisoned")?;
27820 let pin = {
27821 let _main = e.gpu.enter_main()?;
27822 let stream = e.stream();
27823 let (ptr, _g) = head.device_ptr(&stream);
27824 ptr
27825 };
27826 if guard.as_ref().is_none_or(|ws| ws.pin != pin) {
27827 let hi_rows = n_vocab - half;
27829 let (w1, hn1, y1, ev_done) = {
27830 let _r1 = rank1.gpu.enter_main()?;
27831 (
27832 rank1.alloc_u8_uninit(hi_rows * n_embd * 2)?,
27833 rank1.htod(&vec![0.0f32; n_embd])?,
27834 rank1.htod(&vec![0.0f32; hi_rows])?,
27835 rank1.ctx().new_event(None)?,
27836 )
27837 };
27838 {
27839 use cudarc::driver::sys;
27840 let src = pin + (half * n_embd * 2) as u64;
27841 let dst = {
27842 let _r1 = rank1.gpu.enter_main()?;
27843 let rstream = rank1.stream();
27844 let (d, _g) = w1.device_ptr(&rstream);
27845 d
27846 };
27847 let _r1 = rank1.gpu.enter_main()?;
27848 let r = unsafe {
27849 sys::cuMemcpyAsync(
27850 dst as sys::CUdeviceptr,
27851 src as sys::CUdeviceptr,
27852 hi_rows * n_embd * 2,
27853 rank1.stream().cu_stream() as sys::CUstream,
27854 )
27855 };
27856 if r != sys::CUresult::CUDA_SUCCESS {
27857 return Err(format!("head split replica upload: {r:?}").into());
27858 }
27859 rank1.stream().synchronize()?;
27860 }
27861 let (logits_e, ev_hn) = {
27862 let _main = e.gpu.enter_main()?;
27863 (e.htod(&vec![0.0f32; n_vocab])?, e.ctx().new_event(None)?)
27864 };
27865 let (raw_hn1, raw_y1) = {
27866 let _r1 = rank1.gpu.enter_main()?;
27867 let rstream = rank1.stream();
27868 let (a, _g0) = hn1.device_ptr(&rstream);
27869 let (b, _g1) = y1.device_ptr(&rstream);
27870 (a, b)
27871 };
27872 let raw_logits_hi = {
27873 let _main = e.gpu.enter_main()?;
27874 let stream = e.stream();
27875 let (l, _g) = logits_e.device_ptr(&stream);
27876 l + (half * 4) as u64
27877 };
27878 *guard = Some(HeadSplit {
27879 pin,
27880 w1,
27881 hn1,
27882 y1,
27883 logits_e,
27884 ev_hn,
27885 ev_done,
27886 raw_hn1,
27887 raw_y1,
27888 raw_logits_hi,
27889 samp: None,
27890 });
27891 }
27892 let ws = guard.as_mut().expect("armed above");
27893 let hi_rows = n_vocab - half;
27894 let raw_hn = {
27896 let _main = e.gpu.enter_main()?;
27897 let stream = e.stream();
27898 let (h, _g) = hn.device_ptr(&stream);
27899 ws.ev_hn.record(&stream)?;
27900 h
27901 };
27902 {
27903 let _r1 = rank1.gpu.enter_main()?;
27904 rank1.stream().wait(&ws.ev_hn)?;
27905 crate::tp::raw_copy_bytes(ws.raw_hn1, raw_hn, n_embd * 4, rank1)?;
27906 let HeadSplit { w1, hn1, y1, .. } = &mut *ws;
27907 rank1.matvec_bf16_into(w1, hn1, y1, n_embd, hi_rows)?;
27908 crate::tp::raw_copy_bytes(ws.raw_logits_hi, ws.raw_y1, hi_rows * 4, rank1)?;
27909 ws.ev_done.record(&rank1.stream())?;
27910 }
27911 {
27912 let _main = e.gpu.enter_main()?;
27913 let head_lo = head.slice(0..half * n_embd * 2);
27914 let HeadSplit { logits_e, .. } = &mut *ws;
27915 e.matvec_bf16_view_into(&head_lo, hn, logits_e, n_embd, half)?;
27917 e.stream().wait(&ws.ev_done)?;
27918 Ok(Some(()))
27919 }
27920 }
27921
27922 pub(crate) fn head_split_argmax_device(
27927 &self,
27928 e: &Engine,
27929 hn: &CudaSlice<f32>,
27930 token_d: &mut CudaSlice<u32>,
27931 ) -> Result<bool, Box<dyn std::error::Error>> {
27932 if self.head_split_fill_device(e, hn)?.is_none() {
27933 return Ok(false);
27934 }
27935 let n_vocab = self.cfg.n_vocab as usize;
27936 let guard = HEAD_SPLIT_WS
27937 .lock()
27938 .map_err(|_| "head split lock is poisoned")?;
27939 let ws = guard.as_ref().expect("filled above");
27940 let _main = e.gpu.enter_main()?;
27941 e.argmax_token_device_into(&ws.logits_e, token_d, n_vocab)?;
27942 Ok(true)
27943 }
27944
27945 pub(crate) fn head_split_sample_device(
27951 &self,
27952 e: &Engine,
27953 hn: &CudaSlice<f32>,
27954 token_d: &mut CudaSlice<u32>,
27955 samp: &crate::decode_batch::DevSamp,
27956 ctr: u32,
27957 ) -> Result<bool, Box<dyn std::error::Error>> {
27958 if self.head_split_fill_device(e, hn)?.is_none() {
27959 return Ok(false);
27960 }
27961 let n_vocab = self.cfg.n_vocab as usize;
27962 let guard = HEAD_SPLIT_WS
27963 .lock()
27964 .map_err(|_| "head split lock is poisoned")?;
27965 let mut guard = guard;
27966 let ws = guard.as_mut().expect("filled above");
27967 let _main = e.gpu.enter_main()?;
27968 if ws.samp.is_none() {
27969 ws.samp = Some(SampScratch {
27970 pb: e.zeros(n_vocab)?,
27971 th: e.zeros(1)?,
27972 z: e.zeros(1)?,
27973 mx: e.zeros(1)?,
27974 rows: e.htod_i32(&[0i32])?,
27975 });
27976 }
27977 let filtered = samp.top_k > 0 || samp.top_p < 1.0 || samp.min_p > 0.0;
27978 let HeadSplit {
27979 logits_e,
27980 samp: scratch,
27981 ..
27982 } = &mut *ws;
27983 let sc = scratch.as_mut().expect("armed above");
27984 if filtered {
27985 e.filter_stats(
27986 logits_e, n_vocab, &sc.rows, &mut sc.th, &mut sc.z, &mut sc.mx, n_vocab, 1,
27987 samp.temp, samp.top_k, samp.top_p, samp.min_p,
27988 )?;
27989 let SampScratch { pb, th, mx, .. } = sc;
27990 e.gumbel_perturb_filtered_col(
27991 logits_e, 0, pb, n_vocab, samp.seed, ctr, samp.temp, mx, th, 0,
27992 )?;
27993 } else {
27994 e.gumbel_perturb_col(logits_e, 0, &mut sc.pb, n_vocab, samp.seed, ctr, samp.temp)?;
27995 }
27996 e.argmax_token_device_col(&sc.pb, 0, n_vocab, token_d, 0)?;
27997 Ok(true)
27998 }
27999
28000 pub(crate) fn head_split_logits_dtoh(
28003 &self,
28004 e: &Engine,
28005 ) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
28006 let guard = HEAD_SPLIT_WS
28007 .lock()
28008 .map_err(|_| "head split lock is poisoned")?;
28009 let ws = guard.as_ref().ok_or("head split logits not armed")?;
28010 let _main = e.gpu.enter_main()?;
28011 e.dtoh(&ws.logits_e)
28012 }
28013
28014 #[allow(clippy::type_complexity)] pub fn step35_token_graph_chunk(
28024 &self,
28025 e: &Engine,
28026 token: u32,
28027 k_target: usize,
28028 cache: &mut Cache,
28029 ) -> Result<Option<(Vec<u32>, Vec<f32>)>, Box<dyn std::error::Error>> {
28030 if !self.uses_sliding_gated_moe_program()
28031 || !crate::tp::step_tp_graph_enabled()?
28032 || !crate::tp::step_tp_dcw_enabled()?
28033 || !crate::tp::step_tp_qkv_fused_enabled()?
28034 || !crate::tp::step_tp_dev_router_enabled()?
28035 || !crate::tp::step_nvfp4_dev_routes_enabled()?
28036 {
28037 return Ok(None);
28038 }
28039 if !crate::spec::graph_launch_headroom_ok(e) {
28042 static NOTED: std::sync::Once = std::sync::Once::new();
28043 NOTED.call_once(|| crate::spec::graph_replay_suspended_note("step-tp-token"));
28044 return Ok(None);
28045 }
28046 let n_layers = self.layers.len();
28047 let pos = cache.pos;
28048 let staged_next = pos + 1;
28049 if staged_next < 96 {
28050 return Ok(None);
28051 }
28052 let (fa_vec, n_splits) = e.fa_geom_eager(staged_next, 128, 8, false);
28055 if !fa_vec {
28056 return Ok(None);
28057 }
28058 let sp = crate::fa_split_keys(staged_next, 8);
28059 let bucket_max = (n_splits * sp).max(staged_next);
28060 let to_boundary = bucket_max.saturating_sub(staged_next) + 1;
28061 let mut k = k_target.min(to_boundary).min(16);
28062 if k < 2 {
28063 return Ok(None);
28064 }
28065 for il in 0..n_layers {
28067 let Some(tp_kv) = cache.tp_kv[il].as_ref() else {
28068 return Ok(None);
28069 };
28070 while k >= 2 && tp_kv.peek_append_ring(k)?.1 {
28071 k -= 1;
28072 }
28073 if k < 2 {
28074 return Ok(None);
28075 }
28076 }
28077
28078 let mut state_guard = self
28079 .step35_token_graph
28080 .lock()
28081 .map_err(|_| "step35 token graph lock is poisoned")?;
28082 let Some(state) = state_guard.as_mut() else {
28083 return Ok(None); };
28085 if state.graphs.is_empty() {
28086 return Ok(None);
28087 }
28088 {
28089 let (b, g) = state.graphs.first_mut().expect("checked above");
28090 if *b != bucket_max {
28091 g.retarget_bucket(bucket_max)?;
28092 *b = bucket_max;
28093 }
28094 }
28095 let graph = state.graphs.first().map(|(_, g)| g).expect("checked above");
28096
28097 {
28099 let fa0 = match &self.layers[0].mixer {
28100 Mixer::Full(fa) => fa,
28101 _ => return Err("step35 token graph expects full-attention layers".into()),
28102 };
28103 let tp0 = fa0
28104 .step_tp_qkv
28105 .as_ref()
28106 .ok_or("step35 token graph lost its TP state")?;
28107 for rank in 0..tp0.runtime.devices().len() {
28108 let engine = tp0
28109 .runtime
28110 .rank_engine(rank)
28111 .ok_or("step35 token graph lost a rank engine")?;
28112 let _main = engine.gpu.enter_main()?;
28113 engine.stream().synchronize()?;
28114 }
28115 }
28116
28117 {
28120 let _main = e.gpu.enter_main()?;
28121 e.set_u32_one(&mut state.token_d, token)?;
28122 e.set_i32_one(&mut state.pos_d, pos as i32)?;
28123 e.set_i32_one(&mut state.hist_idx, 0)?;
28124 }
28125 for _ in 0..k {
28126 graph.launch(e)?;
28127 }
28128 for il in 0..n_layers {
28130 let tp_kv = cache.tp_kv[il].as_mut().expect("eligibility checked above");
28131 let transaction = tp_kv.begin_transaction()?;
28132 let fa = match &self.layers[il].mixer {
28133 Mixer::Full(fa) => fa,
28134 _ => return Err("step35 token graph expects full-attention layers".into()),
28135 };
28136 let tp = fa
28137 .step_tp_qkv
28138 .as_ref()
28139 .ok_or("step35 token graph lost its TP state")?;
28140 let empty: [CudaSlice<f32>; 0] = [];
28141 tp.runtime.append_tp_kv_transaction_inner(
28142 tp_kv,
28143 transaction,
28144 &empty,
28145 &empty,
28146 k,
28147 true,
28148 )?;
28149 tp.runtime
28150 .commit_tp_kv_transaction_external(tp_kv, transaction, k)?;
28151 if let Some(local) = cache.kv[il].as_mut() {
28152 local.len = pos + k;
28153 let _main = e.gpu.enter_main()?;
28154 e.set_i32_one(&mut local.len_d, (pos + k) as i32)?;
28155 }
28156 }
28157 cache.pos = pos + k;
28158 let (hist, logits) = {
28159 let _main = e.gpu.enter_main()?;
28160 e.stream().synchronize()?;
28161 (e.dtoh_u32(&state.token_hist)?, e.dtoh(&state.logits_stage)?)
28162 };
28163 Ok(Some((hist[..k].to_vec(), logits)))
28164 }
28165}
28166
28167impl HybridModel {
28168 #[allow(clippy::too_many_arguments)]
28173 fn step35_token_graph_build(
28174 &self,
28175 e: &Engine,
28176 cache: &mut Cache,
28177 state: &mut Step35TokenGraphState,
28178 bucket_max: usize,
28179 ) -> Result<(), Box<dyn std::error::Error>> {
28180 use cudarc::driver::DevicePtr;
28181 let n_embd = self.cfg.n_embd as usize;
28182 let eps = self.cfg.rms_eps;
28183 let n_layers = self.layers.len();
28184 let started = std::time::Instant::now();
28185 if !crate::router_kernel_on() {
28186 return Err(
28187 "step35 token graph requires the router kernel (MEMRA_ROUTER_KERNEL=0)".into(),
28188 );
28189 }
28190 if !Engine::bf16_mmv_on() || !n_embd.is_multiple_of(8) {
28191 return Err("step35 token graph requires MEMRA_BF16_MMV bf16-resident matvecs".into());
28192 }
28193
28194 let embd_gpu = self
28196 .embd_gpu_try(e)
28197 .ok_or("step35 token graph could not upload the device embed table")?;
28198 let embd_qtype = match self.embd.ggml_type {
28199 memra_gguf::GgmlType::BF16 => crate::QT_BF16,
28200 memra_gguf::GgmlType::Q8_0 => crate::QT_Q8_0,
28201 other => return Err(format!("token graph embed dtype {other:?} unhandled").into()),
28202 };
28203 let embd_row_bytes = self.embd.raw.len() / self.cfg.n_vocab as usize;
28204
28205 let (p_mixed, p_kshadow, p_vshadow) = {
28207 let _main = e.gpu.enter_main()?;
28208 let stream = e.stream();
28209 let (a, _g) = state.mixed_stage.device_ptr(&stream);
28210 let (b, _g) = state.k_shadow_stage.device_ptr(&stream);
28211 let (c, _g) = state.v_shadow_stage.device_ptr(&stream);
28212 (a, b, c)
28213 };
28214
28215 crate::tp::token_graph_build_begin()?;
28216 let mut group_id: u32 = 0;
28217 for il in 0..n_layers {
28218 let layer = &self.layers[il];
28219 let fa = match &layer.mixer {
28220 Mixer::Full(fa) => fa,
28221 _ => return Err("step35 token graph expects full-attention layers".into()),
28222 };
28223 let tp = fa
28224 .step_tp_qkv
28225 .as_ref()
28226 .ok_or("step35 token graph lost its TP state")?;
28227 let attention = tp
28228 .attention
28229 .as_ref()
28230 .ok_or("step35 token graph lost its attention aux")?;
28231 let geometry = self.step35_geom(il);
28232 let window = geometry.window.map(|w| w as usize);
28233 let head_dim = geometry.head_dim_k as usize;
28234 let heads = geometry.n_head as usize;
28235 let kv_heads = geometry.n_head_kv as usize;
28236 let ranks = tp.runtime.devices().len();
28237 let local_heads = heads / ranks;
28238 let local_kv_heads = kv_heads / ranks;
28239 let layer_bucket = window.map(|w| bucket_max.min(w)).unwrap_or(bucket_max);
28240 let use_gate_shards =
28241 attention.gate_shards.is_some() || attention.gate_shards_bf16.is_some();
28242 if !use_gate_shards {
28243 return Err("step35 token graph requires the fused gate shards".into());
28244 }
28245
28246 let ws_index = tp
28247 .runtime
28248 .decode_v2_ensure(e, &tp.q, &tp.k, &tp.v, &tp.o, heads)?;
28249 let ws_mutex = tp.runtime.decode_v2_workspace();
28250 let mut ws_guard = ws_mutex
28251 .lock()
28252 .map_err(|_| "step TP decode v2 workspace lock is poisoned")?;
28253 let ws = ws_guard
28254 .get_mut(ws_index)
28255 .ok_or("step TP decode v2 workspace missing after ensure")?;
28256 tp.runtime
28257 .decode_v2_arm_token_mirrors(ws, p_mixed, (p_kshadow, p_vshadow))?;
28258 let mut rope_freqs = Vec::with_capacity(ranks);
28259 for rank in 0..ranks {
28260 let engine = tp
28261 .runtime
28262 .rank_engine(rank)
28263 .ok_or("step35 token graph lost a rank engine")?;
28264 rope_freqs.push(if geometry.rope_factors {
28265 self.step35_aux
28266 .as_ref()
28267 .and_then(|aux| aux.rope_freqs(engine))
28268 } else {
28269 None
28270 });
28271 }
28272 let gate_shards_arg = if let Some(shards) = attention.gate_shards.as_deref() {
28273 Some(crate::tp::StepTpGateShards::F32(shards))
28274 } else {
28275 attention
28276 .gate_shards_bf16
28277 .as_deref()
28278 .map(crate::tp::StepTpGateShards::Bf16)
28279 };
28280
28281 let decode_input = attention
28283 .decode_input
28284 .as_ref()
28285 .ok_or("step35 token graph requires the replicated decode input")?;
28286 let mut decode_input = decode_input
28287 .lock()
28288 .map_err(|_| "replicated decode input lock is poisoned")?;
28289 if ws.h_stage.is_none() {
28291 return Err(
28292 "step35 token graph requires the stage flow armed (run eager dcw first)".into(),
28293 );
28294 }
28295 {
28296 let state_x = &mut state.x;
28297 let token_d = &state.token_d;
28298 let pos_d = &state.pos_d;
28299 crate::tp::graph_section(e, None, || {
28300 let _main = e.gpu.enter_main()?;
28301 if il == 0 {
28302 e.embed_gather_device_into(
28303 embd_gpu,
28304 token_d,
28305 state_x,
28306 n_embd,
28307 embd_qtype,
28308 embd_row_bytes,
28309 )?;
28310 }
28311 {
28312 let h_stage = ws.h_stage.as_mut().expect("stage armed checked above");
28313 e.rms_norm(
28314 state_x,
28315 layer.attn_norm.float_data(),
28316 h_stage,
28317 n_embd,
28318 1,
28319 eps,
28320 )?;
28321 }
28322 {
28323 let pos_stage = ws.pos_stage.as_mut().expect("stage armed above");
28324 let mut dst = pos_stage.slice_mut(0..1);
28325 e.stream().memcpy_dtod(&pos_d.slice(0..1), &mut dst)?;
28326 }
28327 Ok(())
28328 })?;
28329 }
28330
28331 group_id += 1;
28333 for rank in 0..ranks {
28334 let engine = tp
28335 .runtime
28336 .rank_engine(rank)
28337 .ok_or("step35 token graph lost a rank engine")?;
28338 {
28339 let ceiling = window
28344 .map(|w| cache.max_ctx.min(w))
28345 .unwrap_or(cache.max_ctx);
28346 let _main = engine.gpu.enter_main()?;
28347 engine.fa_dcw_pool_ensure(
28348 head_dim,
28349 local_heads,
28350 local_kv_heads,
28351 ceiling.min(2048),
28352 )?;
28353 engine.fa_dcw_pool_ensure(head_dim, local_heads, local_kv_heads, ceiling)?;
28354 engine.fa_dcw_pool_ensure(
28355 head_dim,
28356 local_heads,
28357 local_kv_heads,
28358 layer_bucket,
28359 )?;
28360 }
28361 let runtime = &tp.runtime;
28362 let q_norm = &attention.q_norm;
28363 let k_norm = &attention.k_norm;
28364 let gate_ref = gate_shards_arg.as_ref();
28365 crate::tp::graph_section(engine, Some(group_id), || {
28366 runtime.decode_v2_input_qkv_rank(
28367 ws,
28368 &state.pos_d,
28369 &mut decode_input,
28370 &tp.q,
28371 &tp.k,
28372 &tp.v,
28373 q_norm,
28374 k_norm,
28375 head_dim,
28376 geometry.n_rot as usize,
28377 geometry.rope_base,
28378 &rope_freqs,
28379 eps,
28380 gate_ref,
28381 true,
28382 true,
28383 false,
28384 rank,
28385 None,
28386 )?;
28387 let distributed = cache.tp_kv[il]
28390 .as_mut()
28391 .ok_or("step35 token graph lost a TP cache")?;
28392 let (kv_dim_k, kv_dim_v) = (distributed.kv_dim_k(), distributed.kv_dim_v());
28393 let (ktb, vtb) = (distributed.k_tok_bytes(), distributed.v_tok_bytes());
28394 let capacity = distributed.physical_capacity();
28395 {
28396 let rank_cache = distributed
28397 .rank_mut(rank)
28398 .ok_or("step35 token graph lost a rank cache")?;
28399 let (k_plane, v_plane, len_d, base_d) =
28400 rank_cache.planes_and_counters_mut();
28401 engine.append_kv_quantized_dcw(
28402 &ws.k[rank],
28403 &ws.v_raw[rank],
28404 k_plane,
28405 v_plane,
28406 len_d,
28407 base_d,
28408 kv_dim_k,
28409 kv_dim_v,
28410 ktb,
28411 vtb,
28412 )?;
28413 }
28414 {
28415 let rank_cache = distributed
28416 .rank_mut(rank)
28417 .ok_or("step35 token graph lost a rank cache")?;
28418 engine.inc_i32(rank_cache.len_d_mut())?;
28419 }
28420 let rank_cache = distributed
28421 .rank(rank)
28422 .ok_or("step35 token graph lost a rank cache")?;
28423 let k_ring = engine.view_u8_range(rank_cache.k(), 0, capacity * ktb);
28424 let v_ring = engine.view_u8_range(rank_cache.v(), 0, capacity * vtb);
28425 engine.fa_decode_dcw(
28430 &ws.q[rank],
28431 &k_ring,
28432 &v_ring,
28433 &mut ws.attn_out[rank],
28434 head_dim,
28435 local_heads,
28436 local_kv_heads,
28437 rank_cache.len_d(),
28438 rank_cache.base_d(),
28439 window.unwrap_or(0),
28440 layer_bucket,
28441 geometry.attention_scale(),
28442 ktb,
28443 vtb,
28444 None,
28445 )?;
28446 engine.attn_head_gate(
28447 &ws.attn_out[rank],
28448 &ws.gate[rank],
28449 &mut ws.gated[rank],
28450 None,
28451 head_dim,
28452 local_heads,
28453 1,
28454 )?;
28455 runtime.decode_v2_finish_rank_partial(ws, &tp.o, true, rank)?;
28456 Ok(())
28457 })?;
28458 }
28459
28460 {
28462 let root = tp
28463 .runtime
28464 .rank_engine(0)
28465 .ok_or("step35 token graph lost the root engine")?;
28466 let runtime = &tp.runtime;
28467 crate::tp::graph_section(root, None, || runtime.decode_v2_finish_root_fused(ws))?;
28468 }
28469 drop(ws_guard);
28470 drop(decode_input);
28471
28472 let probe_layer: Option<usize> = std::env::var("MEMRA_TG_PROBE_LAYER")
28473 .ok()
28474 .and_then(|v| v.parse().ok());
28475 if probe_layer == Some(il) {
28476 let Step35TokenGraphState {
28477 mixed_stage,
28478 probe_mixed,
28479 ..
28480 } = &mut *state;
28481 crate::tp::graph_section(e, None, || {
28482 let _main = e.gpu.enter_main()?;
28483 let mut dst = probe_mixed.slice_mut(0..n_embd);
28484 e.stream()
28485 .memcpy_dtod(&mixed_stage.slice(0..n_embd), &mut dst)?;
28486 Ok(())
28487 })?;
28488 }
28489
28490 match &layer.ffn {
28492 crate::hybrid::Ffn::Dense {
28493 ffn_gate,
28494 ffn_up,
28495 ffn_down,
28496 } => {
28497 let n_ff = ffn_gate.out_features();
28498 let lim = self.cfg.clamp_shexp_at(il as u32);
28499 if lim.is_some() {
28503 return Err("step35 token graph dense FFN with clamp unsupported".into());
28504 }
28505 let (wg_d, wu_d, wd_d) = match (ffn_gate, ffn_up, ffn_down) {
28506 (
28507 crate::model::GpuTensor::FloatBf16 { data: wg, .. },
28508 crate::model::GpuTensor::FloatBf16 { data: wu, .. },
28509 crate::model::GpuTensor::FloatBf16 { data: wd, .. },
28510 ) => (wg, wu, wd),
28511 _ => {
28512 return Err(
28513 "step35 token graph dense FFN requires bf16-resident weights"
28514 .into(),
28515 );
28516 }
28517 };
28518 crate::tp::graph_section(e, None, || {
28519 let _main = e.gpu.enter_main()?;
28520 let Step35TokenGraphState {
28521 x,
28522 x1,
28523 mixed_stage,
28524 dense_z,
28525 dense_gate,
28526 dense_up,
28527 dense_act,
28528 sh_stage,
28529 ..
28530 } = &mut *state;
28531 e.add_rms_norm(
28532 x,
28533 mixed_stage,
28534 layer.post_attn_norm.float_data(),
28535 x1,
28536 dense_z,
28537 n_embd,
28538 1,
28539 eps,
28540 )?;
28541 e.matvec_bf16_into(wg_d, dense_z, dense_gate, n_embd, n_ff)?;
28545 e.matvec_bf16_into(wu_d, dense_z, dense_up, n_embd, n_ff)?;
28546 Self::ffn_act_lim(
28547 e, &self.cfg, dense_gate, dense_up, 1.0, 1.0, lim, dense_act, n_ff,
28548 )?;
28549 e.matvec_bf16_into(wd_d, dense_act, sh_stage, n_ff, n_embd)?;
28550 e.add(x1, sh_stage, x, n_embd)?;
28551 Ok(())
28552 })?;
28553 }
28554 crate::hybrid::Ffn::Moe(m) => {
28555 let moe = self
28556 .cfg
28557 .moe
28558 .as_ref()
28559 .ok_or("step35 token graph needs moe cfg")?;
28560 let n_expert = moe.expert_count as usize;
28561 let n_used = moe.expert_used_count as usize;
28562 let sigmoid = self
28563 .cfg
28564 .sigmoid_router()
28565 .ok_or("step35 token graph needs the sigmoid router")?;
28566 let step_tp = m
28567 .step_tp
28568 .as_ref()
28569 .ok_or("step35 token graph needs TP experts")?;
28570 let bank = match &step_tp.experts {
28571 crate::hybrid::StepTpExpertBank::Nvfp4(bank) => bank,
28572 _ => return Err("step35 token graph needs the NVFP4 bank".into()),
28573 };
28574 let routes_ws_mutex = bank.device_workspace_handle();
28575 let mut routes_guard = routes_ws_mutex
28576 .lock()
28577 .map_err(|_| "routes workspace lock is poisoned")?;
28578 let routes_ws = routes_guard
28579 .as_mut()
28580 .ok_or("step35 token graph requires the routes workspace warmed")?;
28581 routes_ws.arm_stages(e, bank.input_width, n_used)?;
28582 step_tp.runtime.routes_arm_raw(bank, routes_ws)?;
28583 let p_z = {
28584 let root = step_tp
28585 .runtime
28586 .rank_engine(0)
28587 .ok_or("routes root engine missing")?;
28588 let _main = root.gpu.enter_main()?;
28589 let stream = root.stream();
28590 let in_stage = routes_ws
28591 .in_stage_handle()
28592 .ok_or("routes in stage not armed")?;
28593 let (a, _g) = in_stage.device_ptr(&stream);
28594 a
28595 };
28596 let local_out = bank.expert_width / ranks;
28597
28598 crate::tp::graph_section(e, None, || {
28600 let _main = e.gpu.enter_main()?;
28601 {
28602 let in_stage = routes_ws
28603 .in_stage_mut()
28604 .ok_or("routes in stage not armed")?;
28605 let Step35TokenGraphState {
28606 x, x1, mixed_stage, ..
28607 } = &mut *state;
28608 e.add_rms_norm(
28609 x,
28610 mixed_stage,
28611 layer.post_attn_norm.float_data(),
28612 x1,
28613 in_stage,
28614 n_embd,
28615 1,
28616 eps,
28617 )?;
28618 }
28619 {
28620 let z_ref = routes_ws
28621 .in_stage_handle()
28622 .ok_or("routes in stage not armed")?;
28623 e.router_gemv_into(
28624 m.gate_inp.float_data(),
28625 z_ref,
28626 &mut state.router_logits,
28627 n_embd,
28628 n_expert,
28629 1,
28630 )?;
28631 }
28632 let (sel_e, w_e) = routes_ws
28633 .dev_route_e_mut()
28634 .ok_or("routes staging not armed")?;
28635 e.moe_router_sigmoid_topk_into(
28636 &state.router_logits,
28637 1,
28638 n_expert,
28639 n_used,
28640 m.active_count(),
28641 &m.exp_probs_b_dev,
28642 &m.active_experts_dev,
28643 sigmoid.0,
28644 sigmoid.1,
28645 sel_e,
28646 w_e,
28647 )?;
28648 Ok(())
28649 })?;
28650
28651 group_id += 1;
28653 for rank in 0..ranks {
28654 let engine = step_tp
28655 .runtime
28656 .rank_engine(rank)
28657 .ok_or("routes rank engine missing")?;
28658 let runtime = &step_tp.runtime;
28659 crate::tp::graph_section(engine, Some(group_id), || {
28660 runtime.routes_rank_section(
28661 bank,
28662 routes_ws,
28663 p_z,
28664 local_out,
28665 n_used,
28666 step_tp.activation_limit,
28667 rank,
28668 )
28669 })?;
28670 }
28671
28672 {
28674 let root = step_tp
28675 .runtime
28676 .rank_engine(0)
28677 .ok_or("routes root engine missing")?;
28678 let runtime = &step_tp.runtime;
28679 crate::tp::graph_section(root, None, || {
28680 runtime.routes_root_section(bank, routes_ws)
28681 })?;
28682 }
28683
28684 let lim_sh = self.cfg.clamp_shexp_at(il as u32);
28688 let (wg_sh, wu_sh, wd_sh) = match (&m.gate_shexp, &m.up_shexp, &m.down_shexp) {
28689 (
28690 Some(crate::model::GpuTensor::FloatBf16 { data: wg, .. }),
28691 Some(crate::model::GpuTensor::FloatBf16 { data: wu, .. }),
28692 Some(crate::model::GpuTensor::FloatBf16 { data: wd, .. }),
28693 ) => (wg, wu, wd),
28694 _ => {
28695 return Err(
28696 "step35 token graph shexp requires bf16-resident weights".into()
28697 );
28698 }
28699 };
28700 let n_ff_sh = m
28701 .gate_shexp
28702 .as_ref()
28703 .expect("matched Some above")
28704 .out_features();
28705 let gate_inp_shexp = m.gate_inp_shexp.as_ref();
28708 crate::tp::graph_section(e, None, || {
28709 let _main = e.gpu.enter_main()?;
28710 let (z_ref, out_stage) = routes_ws
28711 .in_and_out_stages_mut()
28712 .ok_or("routes stages not armed")?;
28713 let Step35TokenGraphState {
28714 x,
28715 x1,
28716 sh_stage,
28717 shexp_gate,
28718 shexp_up,
28719 shexp_act,
28720 gate_sig,
28721 ..
28722 } = &mut *state;
28723 e.matvec_bf16_dual_into(
28724 wg_sh, wu_sh, z_ref, shexp_gate, shexp_up, n_embd, n_ff_sh,
28725 )?;
28726 Self::ffn_act_lim(
28727 e, &self.cfg, shexp_gate, shexp_up, 1.0, 1.0, lim_sh, shexp_act,
28728 n_ff_sh,
28729 )?;
28730 e.matvec_bf16_into(wd_sh, shexp_act, sh_stage, n_ff_sh, n_embd)?;
28731 if let Some(gate_w) = gate_inp_shexp {
28732 e.sigmoid_dot_rows_into(
28733 z_ref,
28734 gate_w.float_data(),
28735 gate_sig,
28736 n_embd,
28737 1,
28738 )?;
28739 }
28740 e.add_scaled_rows(sh_stage, gate_sig, out_stage, n_embd, 1)?;
28741 e.add(x1, out_stage, x, n_embd)?;
28742 Ok(())
28743 })?;
28744 }
28745 }
28746 if probe_layer == Some(il) {
28747 let Step35TokenGraphState { x, probe_x, .. } = &mut *state;
28748 crate::tp::graph_section(e, None, || {
28749 let _main = e.gpu.enter_main()?;
28750 let mut dst = probe_x.slice_mut(0..n_embd);
28751 e.stream().memcpy_dtod(&x.slice(0..n_embd), &mut dst)?;
28752 Ok(())
28753 })?;
28754 }
28755 }
28756
28757 let head = match &self.output {
28759 crate::model::GpuTensor::FloatBf16 { data, .. } => data,
28760 _ => return Err("step35 token graph head requires the bf16-resident output".into()),
28761 };
28762 crate::tp::graph_section(e, None, || {
28763 let _main = e.gpu.enter_main()?;
28764 let Step35TokenGraphState {
28765 x,
28766 hn,
28767 logits_stage,
28768 token_d,
28769 pos_d,
28770 token_hist,
28771 hist_idx,
28772 ..
28773 } = &mut *state;
28774 e.rms_norm(x, self.output_norm.float_data(), hn, n_embd, 1, eps)?;
28775 e.matvec_bf16_into(head, hn, logits_stage, n_embd, self.cfg.n_vocab as usize)?;
28776 e.argmax_token_device_into(logits_stage, token_d, self.cfg.n_vocab as usize)?;
28782 e.u32_hist_append(token_d, token_hist, hist_idx)?;
28783 e.inc_i32(pos_d)?;
28784 Ok(())
28785 })?;
28786
28787 let graph = crate::tp::token_graph_build_finish()?;
28788 state.graphs.push((bucket_max, graph));
28789 eprintln!(
28790 "[step35-token-graph] built bucket={bucket_max} layers={n_layers} \
28791 build_ms={:.0} performance_claim=false",
28792 started.elapsed().as_secs_f64() * 1e3
28793 );
28794 Ok(())
28795 }
28796}