1use cudarc::driver::CudaSlice;
6use memra_gguf::config::ModelConfig;
7use crate::Engine;
8use crate::cache::Cache;
9
10pub struct PrimeSlabs {
13 pub t_cap: usize,
14 pub h: CudaSlice<f32>,
15 pub x1: CudaSlice<f32>,
16 pub z: CudaSlice<f32>,
17 pub act: CudaSlice<f32>,
18 pub xa: CudaSlice<f32>,
19 pub xb: CudaSlice<f32>,
20 pub h16: CudaSlice<u8>,
21 pub z16: CudaSlice<u8>,
22 pub gate: CudaSlice<f32>, pub up: CudaSlice<f32>, pub ffn_out: CudaSlice<f32>, pub seg_glue: Vec<Option<cudarc::driver::CudaGraph>>,
32 pub mixed: CudaSlice<f32>,
36 pub seg_mid: Vec<Option<cudarc::driver::CudaGraph>>,
37 pub seg_t: usize,
38}
39
40unsafe impl Send for PrimeSlabs {}
43
44fn empty_cache_layers<T>(n: usize) -> Vec<Option<T>> {
45 std::iter::repeat_with(|| None).take(n).collect()
46}
47
48struct PrimeCacheStages<'a> {
53 parent: &'a mut Cache,
54 cut: usize,
55 stage0: Cache,
56 stage1: Cache,
57}
58
59impl<'a> PrimeCacheStages<'a> {
60 fn new(parent: &'a mut Cache, cut: usize) -> Self {
61 let n = parent.kv.len();
62 assert_eq!(parent.recur.len(), n, "cache layer vectors disagree");
63 assert!(cut <= n, "PP-2 cache cut {cut} exceeds {n} layers");
64 let mut kv0 = empty_cache_layers(n);
65 let mut kv1 = empty_cache_layers(n);
66 let mut recur0 = empty_cache_layers(n);
67 let mut recur1 = empty_cache_layers(n);
68 for i in 0..cut {
69 kv0[i] = parent.kv[i].take();
70 recur0[i] = parent.recur[i].take();
71 }
72 for i in cut..n {
73 kv1[i] = parent.kv[i].take();
74 recur1[i] = parent.recur[i].take();
75 }
76 let pos = parent.pos;
77 let max_ctx = parent.max_ctx;
78 Self {
79 parent,
80 cut,
81 stage0: Cache {
82 kv: kv0,
83 recur: recur0,
84 pos,
85 max_ctx,
86 last_logits_dev: None,
87 dflash_taps: None,
88 },
89 stage1: Cache {
90 kv: kv1,
91 recur: recur1,
92 pos,
93 max_ctx,
94 last_logits_dev: None,
95 dflash_taps: None,
96 },
97 }
98 }
99
100 fn parts(&mut self) -> (&mut Cache, &mut Cache) {
101 (&mut self.stage0, &mut self.stage1)
102 }
103}
104
105impl Drop for PrimeCacheStages<'_> {
106 fn drop(&mut self) {
107 let n = self.parent.kv.len();
108 for i in 0..n {
109 let source = if i < self.cut {
110 &mut self.stage0
111 } else {
112 &mut self.stage1
113 };
114 debug_assert!(self.parent.kv[i].is_none());
115 debug_assert!(self.parent.recur[i].is_none());
116 self.parent.kv[i] = source.kv[i].take();
117 self.parent.recur[i] = source.recur[i].take();
118 }
119 self.parent.pos = self.stage0.pos.min(self.stage1.pos);
120 }
121}
122
123
124pub(crate) struct AttnPre {
126 pub q: cudarc::driver::CudaSlice<f32>,
127 pub k: cudarc::driver::CudaSlice<f32>,
128 pub v: cudarc::driver::CudaSlice<f32>,
129 pub gate: Option<cudarc::driver::CudaSlice<f32>>,
130}
131
132pub(crate) struct GdnPrep {
134 pub hk: usize,
135 pub q_l2: cudarc::driver::CudaSlice<f32>,
136 pub k_l2: cudarc::driver::CudaSlice<f32>,
137 pub v_g: cudarc::driver::CudaSlice<f32>,
138 pub beta: cudarc::driver::CudaSlice<f32>,
139 pub g_log: cudarc::driver::CudaSlice<f32>,
140 pub kb16: Option<cudarc::driver::CudaSlice<u8>>,
141 pub qb16: Option<cudarc::driver::CudaSlice<u8>>,
142}
143
144pub(crate) struct VerifyStreamScratch {
146 pub pos_d: CudaSlice<i32>,
147 pub row_ctrs: Vec<CudaSlice<i32>>,
148}
149use crate::hybrid::{HybridModel, Mixer, FullAttnLayer, LinearAttnLayer, MoeWeights};
150
151struct MoeInputTraceWriter {
152 dir: std::path::PathBuf,
153 index: std::fs::File,
154 payloads: std::collections::HashMap<u16, (std::fs::File, u64)>,
155}
156
157static MOE_INPUT_TRACE_WRITER: std::sync::OnceLock<
158 std::sync::Mutex<Option<MoeInputTraceWriter>>,
159> = std::sync::OnceLock::new();
160
161fn gdec_enabled() -> bool {
164 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
165 *E.get_or_init(|| std::env::var("MEMRA_MOE_GDEC").map(|v| v != "0").unwrap_or(true))
166}
167
168fn moe_slab_enabled() -> bool {
179 std::env::var("MEMRA_MOE_SLAB").as_deref() != Ok("0")
180}
181
182fn moe_grouped_enabled(_cfg: &ModelConfig, _prefill: bool) -> bool {
186 std::env::var("MEMRA_MOE_GROUPED")
187 .map(|value| value != "0")
188 .unwrap_or(false)
189}
190
191fn moe_prefetch_enabled() -> bool {
194 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
195 *E.get_or_init(|| std::env::var("MEMRA_MOE_PREFETCH").as_deref() == Ok("1")
196 || crate::spill_pread::worker_enabled())
197}
198
199fn moe_page_prefetch_window() -> usize {
204 static W: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
205 *W.get_or_init(|| page_prefetch_window_from_values(
206 std::env::var("MEMRA_MOE_PAGE_PREFETCH").as_deref() == Ok("1"),
207 std::env::var("MEMRA_MOE_PAGE_PREFETCH_WINDOW").ok().as_deref(),
208 ))
209}
210
211fn page_prefetch_window_from_values(enabled: bool, raw_window: Option<&str>) -> usize {
212 if !enabled {
213 return 0;
214 }
215 raw_window
216 .and_then(|value| value.parse().ok())
217 .unwrap_or(1)
218}
219
220fn page_prefetch_positions(
224 position: usize,
225 len: usize,
226 window: usize,
227) -> std::ops::Range<usize> {
228 if window == 0 || position >= len {
229 return len..len;
230 }
231 let (start, count) = if position == 0 {
232 (1, window)
233 } else {
234 (position.saturating_add(window), 1)
235 };
236 let start = start.min(len);
237 start..start.saturating_add(count).min(len)
238}
239
240fn grouped_worker_prefetch_position(order_len: usize, current: Option<usize>) -> Option<usize> {
243 let position = current.map_or(0, |position| position.saturating_add(1));
244 (position < order_len).then_some(position)
245}
246
247fn worker_prefetch_window() -> usize {
252 static WINDOW: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
253 *WINDOW.get_or_init(|| {
254 let automatic = crate::spill_pread::configured_depth().saturating_sub(1) / 3;
255 std::env::var("MEMRA_SPILL_WORKER_EXPERT_WINDOW")
256 .ok()
257 .and_then(|value| value.parse::<usize>().ok())
258 .unwrap_or(automatic.max(1))
259 })
260}
261
262fn worker_prefetch_positions(position: usize, len: usize, window: usize) -> std::ops::Range<usize> {
266 if window == 0 || position >= len {
267 return len..len;
268 }
269 let (start, count) = if position == 0 {
270 (0, window)
271 } else {
272 (position.saturating_add(window).saturating_sub(1), 1)
273 };
274 let start = start.min(len);
275 start..start.saturating_add(count).min(len)
276}
277
278fn moe_dev_enabled() -> bool {
283 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
284 *E.get_or_init(|| std::env::var("MEMRA_MOE_DEV").map(|v| v != "0").unwrap_or(true)
285 && !matches!(std::env::var("MEMRA_FUSED_ROUTER").as_deref(), Ok("0")))
286}
287
288fn moe_q8_enabled() -> bool {
293 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
294 *E.get_or_init(|| std::env::var("MEMRA_MOE_Q8").map(|v| v != "0").unwrap_or(true))
295}
296
297fn expert_dp4a_supported(qt: i32) -> bool {
300 qt == crate::QT_Q4_0 || qt == crate::QT_IQ3_S || qt == crate::QT_IQ4_XS
301 || qt == crate::QT_Q3_K || qt == crate::QT_Q4_K || qt == crate::QT_Q6_K
302}
303
304fn q8_expert_supported(qt: i32) -> bool {
305 static KQ: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
311 let kq = *KQ.get_or_init(|| {
312 std::env::var("MEMRA_MOE_Q8_KQ").map(|v| v != "0").unwrap_or(true)
313 });
314 let nvfp4_q8 = std::env::var("MEMRA_MOE_Q8_NVFP4").map(|v| v != "0").unwrap_or(true);
321 qt == crate::QT_IQ3_S || qt == crate::QT_IQ4_XS || (nvfp4_q8 && qt == crate::QT_NVFP4)
322 || (kq && (qt == crate::QT_Q3_K || qt == crate::QT_Q4_K || qt == crate::QT_Q6_K))
323}
324
325fn q8_expert_dec_supported(qt: i32) -> bool {
328 qt == crate::QT_IQ3_S || qt == crate::QT_IQ4_XS || qt == crate::QT_Q4_0
329}
330
331fn f16g_proj_ok(qt: i32, in_f: usize) -> bool {
337 match qt {
338 crate::QT_Q4_0 => in_f % 32 == 0,
339 crate::QT_IQ4_XS | crate::QT_IQ3_S | crate::QT_Q3_K | crate::QT_Q4_K
340 | crate::QT_Q6_K => in_f % 256 == 0,
341 _ => false,
342 }
343}
344
345fn moe_prewarm_enabled() -> bool {
348 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
349 *E.get_or_init(|| std::env::var("MEMRA_MOE_PREWARM").map(|v| v != "0").unwrap_or(true))
350}
351
352fn cpu_expert_profile_admit_enabled() -> bool {
356 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
357 *E.get_or_init(|| std::env::var("MEMRA_CPU_EXPERT_FREEZE_PROFILE_ADMIT").as_deref() == Ok("1"))
358}
359
360pub const PRIME_MIN_T: usize = 16;
364const PRIME_PIPE_MICROBATCHES: usize = 8;
365const PRIME_PIPE_MIN_CHUNK: usize = 128;
366const PRIME_PIPE_EDGE_MIN_CHUNK: usize = 64;
367const PRIME_PIPE_LINEAR_WORK: usize = 8;
368
369fn prime_pp2_auto_geometry(n_layers: usize) -> bool {
370 crate::pp::prime_pp_on()
371 && !crate::pp::pp2_streams_off()
372 && crate::pp::pp_cuts(n_layers).is_some_and(|cuts| cuts.len() == 3)
373}
374
375pub fn prime_chunk_tokens(t: usize, n_layers: usize) -> usize {
379 if let Ok(value) = std::env::var("MEMRA_PRIME_CHUNK") {
380 let parsed = value.parse::<usize>().unwrap_or(crate::cache::PRIME_CHUNK_MAX_TOKENS);
381 return if crate::cache::swa_ring_on() {
382 if parsed == 0 {
383 crate::cache::PRIME_CHUNK_MAX_TOKENS
384 } else {
385 parsed.min(crate::cache::PRIME_CHUNK_MAX_TOKENS)
386 }
387 } else {
388 parsed
389 };
390 }
391 let chunk = crate::cache::PRIME_CHUNK_MAX_TOKENS;
392 if prime_pp2_auto_geometry(n_layers) && t >= 2 * PRIME_PIPE_MIN_CHUNK {
393 chunk.min(
394 t.div_ceil(PRIME_PIPE_MICROBATCHES)
395 .max(PRIME_PIPE_MIN_CHUNK),
396 )
397 } else {
398 chunk
399 }
400}
401
402fn fixed_prime_chunk_ranges(t: usize, chunk: usize) -> Vec<(usize, usize)> {
403 fixed_prime_chunk_ranges_for_ring(t, chunk, crate::cache::swa_ring_on())
404}
405
406fn fixed_prime_chunk_ranges_for_ring(t: usize, chunk: usize, ring_on: bool) -> Vec<(usize, usize)> {
407 if chunk == 0 || t <= chunk {
408 return vec![(0, t)];
409 }
410 let mut ranges = Vec::with_capacity(t.div_ceil(chunk));
411 let mut start = 0usize;
412 while start < t {
413 let mut end = (start + chunk).min(t);
414 if t - end > 0 && t - end < PRIME_MIN_T {
415 if ring_on {
416 let shifted = t - PRIME_MIN_T;
417 end = if shifted > start { shifted } else { t };
418 } else {
419 end = t;
420 }
421 }
422 ranges.push((start, end));
423 start = end;
424 }
425 ranges
426}
427
428fn prime_chunk_work(prefix: usize, total: usize) -> u128 {
429 let prefix = prefix as u128;
430 prefix * (prefix + (PRIME_PIPE_LINEAR_WORK as u128) * (total as u128))
431}
432
433fn dynamic_prime_chunk_ranges(
434 t: usize,
435 fixed_chunk: usize,
436 fixed: &[(usize, usize)],
437) -> Vec<(usize, usize)> {
438 let n = fixed.len();
439 if n < 3 {
440 return fixed.to_vec();
441 }
442
443 let max_first = t - (n - 1) * PRIME_MIN_T;
444 let first = fixed_chunk
445 .div_ceil(2)
446 .max(PRIME_PIPE_EDGE_MIN_CHUNK)
447 .min(max_first);
448 let mut ranges = Vec::with_capacity(n);
449 ranges.push((0, first));
450
451 let first_work = prime_chunk_work(first, t);
452 let work_span = prime_chunk_work(t, t) - first_work;
453 let denominator = (n - 1) as u128;
454 let mut previous = first;
455 for boundary in 1..n - 1 {
456 let target = first_work * denominator + work_span * (boundary as u128);
457 let remaining = n - 1 - boundary;
458 let mut low = previous + PRIME_MIN_T;
459 let mut high = t - remaining * PRIME_MIN_T;
460 while low < high {
461 let mid = low + (high - low) / 2;
462 if prime_chunk_work(mid, t) * denominator >= target {
463 high = mid;
464 } else {
465 low = mid + 1;
466 }
467 }
468 ranges.push((previous, low));
469 previous = low;
470 }
471 ranges.push((previous, t));
472 ranges
473}
474
475pub fn prime_chunk_ranges(t: usize, n_layers: usize) -> Vec<(usize, usize)> {
479 let explicit_chunk = std::env::var_os("MEMRA_PRIME_CHUNK").is_some();
480 let chunk = prime_chunk_tokens(t, n_layers);
481 let fixed = fixed_prime_chunk_ranges(t, chunk);
482 let dynamic = match std::env::var("MEMRA_PRIME_CHUNK_SCHED") {
483 Ok(value) => value == "dynamic",
484 Err(_) => true,
485 };
486 if explicit_chunk || !dynamic || !prime_pp2_auto_geometry(n_layers) {
487 fixed
488 } else {
489 dynamic_prime_chunk_ranges(t, chunk, &fixed)
490 }
491}
492
493impl HybridModel {
494 fn prime_trace_path() -> Option<&'static str> {
499 static P: std::sync::OnceLock<Option<String>> = std::sync::OnceLock::new();
500 P.get_or_init(|| std::env::var("MEMRA_PRIME_TRACE").ok())
501 .as_deref()
502 }
503
504 pub fn forward(&self, e: &Engine, tokens: &[u32]) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
506 if self.is_gemma4_e4b() { return self.gemma4_e4b_forward(e, tokens, false); }
507 if self.cfg.gemma4.is_some() { return self.gemma4_forward(e, tokens, false); }
508 let cfg = &self.cfg;
509 let n_embd = cfg.n_embd as usize;
510 let t = tokens.len();
511 let eps = cfg.rms_eps;
512 let pos: Vec<i32> = (0..t as i32).collect();
513 let pos_d = e.htod_i32(&pos)?;
514
515 let mut x = self.embed(e, tokens)?; for (il, layer) in self.layers.iter().enumerate() {
518 let mut h = e.uninit(t * n_embd)?;
520 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
521
522 let mixed = match &layer.mixer {
523 Mixer::Full(fa) => self.full_attn(e, fa, &h, &pos_d, t, il)?,
524 Mixer::Linear(la) => self.linear_attn(e, la, &h, t)?,
525 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
526 };
527
528 let mut x1 = e.uninit(t * n_embd)?;
530 e.add(&x, &mixed, &mut x1, t * n_embd)?;
531
532 let mut z = e.uninit(t * n_embd)?;
534 e.rms_norm(&x1, layer.post_attn_norm.float_data(), &mut z, n_embd, t, eps)?;
535 let ffn_out = match &layer.ffn {
536 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
537 let n_ff = ffn_gate.out_features();
538 let mut g2 = e.matmul_group(&[ffn_gate, ffn_up], &z, t)?;
539 let up = g2.pop().unwrap();
540 let gate = g2.pop().unwrap();
541 let mut act = e.uninit(t * n_ff)?;
542 Self::ffn_act_lim(e, &self.cfg, &gate, &up, 1.0, 1.0,
547 self.cfg.clamp_shexp_at(il as u32), &mut act, t * n_ff)?;
548 e.matmul(ffn_down, &act, t)?
549 }
550 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il_prefill(e, m, &z, t, il as u16)?,
551 };
552 let mut x2 = e.uninit(t * n_embd)?;
553 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
554 x = x2;
555 }
556
557 let mut hn = e.uninit(t * n_embd)?;
558 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
559 let logits = e.matmul(&self.output, &hn, t)?;
560 Ok(e.dtoh(&logits)?)
561 }
562
563 pub fn forward_last(&self, e: &Engine, tokens: &[u32]) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
569 if self.cfg.gemma4.is_some() { return self.gemma4_forward(e, tokens, true); }
570 let cfg = &self.cfg;
571 let n_embd = cfg.n_embd as usize;
572 let t = tokens.len();
573 let eps = cfg.rms_eps;
574 let pos: Vec<i32> = (0..t as i32).collect();
575 let pos_d = e.htod_i32(&pos)?;
576
577 let mut x = self.embed(e, tokens)?; let probe = std::env::var("MEMRA_LAYER_PROBE").is_ok();
581 for (il, layer) in self.layers.iter().enumerate() {
582 let mut h = e.uninit(t * n_embd)?;
583 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
584 if probe { e.stream().synchronize()?; eprintln!("[probe] L{il} norm ok"); }
585 let mixed = match &layer.mixer {
586 Mixer::Full(fa) => self.full_attn(e, fa, &h, &pos_d, t, il)?,
587 Mixer::Linear(la) => self.linear_attn(e, la, &h, t)?,
588 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
589 };
590 if probe { e.stream().synchronize()?; eprintln!("[probe] L{il} mixer ok"); }
591 let mut x1 = e.uninit(t * n_embd)?;
592 e.add(&x, &mixed, &mut x1, t * n_embd)?;
593 let mut z = e.uninit(t * n_embd)?;
594 e.rms_norm(&x1, layer.post_attn_norm.float_data(), &mut z, n_embd, t, eps)?;
595 let ffn_out = match &layer.ffn {
596 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
597 let n_ff = ffn_gate.out_features();
598 let mut g2 = e.matmul_group(&[ffn_gate, ffn_up], &z, t)?;
599 let up = g2.pop().unwrap();
600 let gate = g2.pop().unwrap();
601 let mut act = e.uninit(t * n_ff)?;
602 Self::ffn_act_lim(e, &self.cfg, &gate, &up, 1.0, 1.0,
604 self.cfg.clamp_shexp_at(il as u32), &mut act, t * n_ff)?;
605 e.matmul(ffn_down, &act, t)?
606 }
607 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il_prefill(e, m, &z, t, il as u16)?,
608 };
609 if probe { e.stream().synchronize()?; eprintln!("[probe] L{il} ffn ok"); }
610 let mut x2 = e.uninit(t * n_embd)?;
611 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
612 x = x2;
613 }
614 let mut hn = e.uninit(t * n_embd)?;
616 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
617 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)?;
620 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
621 let logits = e.matmul(&self.output, &hlast, 1)?; Ok(e.dtoh(&logits)?)
623 }
624
625 pub fn prime_cache(&self, e: &Engine, tokens: &[u32], cache: &mut Cache, queued_after: usize)
657 -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
658 let n_embd = self.cfg.n_embd as usize;
659 let t = tokens.len();
660 assert!(t >= PRIME_MIN_T, "prime_cache needs T >= {PRIME_MIN_T} (caller gates)");
664 assert!(cache.pos + t <= cache.max_ctx, "prime_cache: prompt exceeds cache max_ctx");
665
666 if self.is_gemma4_e4b() {
678 return self.gemma4_e4b_prime(e, tokens, cache);
679 }
680 if self.cfg.gemma4.is_some() {
681 return self.gemma4_prime(e, tokens, cache);
683 }
684 let ranges = prime_chunk_ranges(t, self.layers.len());
685 let legacy_calllocal =
721 std::env::var("MEMRA_PRIME_CALLLOCAL").as_deref() == Ok("1");
722 let seq_end = if legacy_calllocal {
723 cache.pos + t
724 } else {
725 cache.pos + t + queued_after
726 };
727 if ranges.len() == 1 {
728 return self.prime_chunk(e, tokens, cache, seq_end);
729 }
730 if crate::pp::prime_pipe_on()
735 && crate::pp::prime_pp_on()
736 && !crate::pp::pp2_streams_off()
737 {
738 if let Some(fence) = crate::pp::pp_cuts(self.layers.len()).filter(|f| f.len() == 3) {
739 if crate::pp::pp_multi_stream_same_device() {
740 return Err(
741 "prime chunk pipeline refused with 2 stage streams on one device — \
742 that concurrent-stream placement remains quarantined by the deferred \
743 pp flake record. Use one device per stage or MEMRA_PRIME_PIPE=0 for \
744 the serial split."
745 .into(),
746 );
747 }
748 return self.prime_cache_pp2_pipelined(
749 e, tokens, cache, seq_end, &ranges, &fence,
750 );
751 }
752 }
753 let mut hiddens = e.uninit(t * n_embd)?;
754 let mut last: Option<(Vec<f32>, CudaSlice<f32>)> = None;
755 for &(start, end) in &ranges {
756 let (l, hs, x) = self.prime_chunk(e, &tokens[start..end], cache, seq_end)?;
757 e.copy_into(&mut hiddens, start * n_embd, &x, (end - start) * n_embd)?;
758 last = Some((l, hs));
759 }
760 let (logits, h_seed) = last.unwrap();
761 Ok((logits, h_seed, hiddens))
762 }
763
764 fn prime_cache_pp2_pipelined(
769 &self,
770 e: &Engine,
771 tokens: &[u32],
772 cache: &mut Cache,
773 seq_end: usize,
774 ranges: &[(usize, usize)],
775 fence: &[usize],
776 ) -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
777 debug_assert_eq!(fence.len(), 3);
778 debug_assert!(ranges.len() >= 2);
779 let rt = crate::pp::PpNRt::get(e)?;
780 assert_eq!(rt.n_stages(), 2, "prime pipeline requires exactly two PP stages");
781 let n_embd = self.cfg.n_embd as usize;
782 let t = tokens.len();
783 let initial_base = cache.pos;
784 let caller_stream = e.stream();
785
786 rt.fence_stages_behind(&caller_stream)?;
791 let max_payload = ranges
792 .iter()
793 .map(|(s, e)| (e - s) * n_embd)
794 .max()
795 .unwrap();
796 rt.prepare_overlap_slots(0, max_payload)?;
797
798 let mut hiddens = e.uninit(t * n_embd)?;
799 let mut last: Option<(Vec<f32>, CudaSlice<f32>)> = None;
800 let mut stage_caches = PrimeCacheStages::new(cache, fence[1]);
801 let (cache0, cache1) = stage_caches.parts();
802 let (first_start, first_end) = ranges[0];
803 let mut slot = self.prime_pp2_stage0_enqueue(
804 e,
805 rt,
806 &tokens[first_start..first_end],
807 cache0,
808 seq_end,
809 fence,
810 initial_base + first_start,
811 true,
812 )?;
813 cache0.pos = initial_base + first_end;
814
815 for (i, &(start, end)) in ranges.iter().enumerate() {
816 let base = initial_base + start;
817 debug_assert_eq!(
818 cache1.pos, base,
819 "stage 1 must drain chunks in original position order"
820 );
821 let (out, next_slot) = if let Some(&(next_start, next_end)) = ranges.get(i + 1) {
822 let next_base = initial_base + next_start;
823 debug_assert_eq!(
824 cache0.pos, next_base,
825 "stage 0 must issue chunks in original position order"
826 );
827 let cache0_stage = &mut *cache0;
828 std::thread::scope(
833 |scope| -> Result<_, Box<dyn std::error::Error>> {
834 let stage0 = scope.spawn(move || -> Result<usize, String> {
835 let next = self
836 .prime_pp2_stage0_enqueue(
837 e,
838 rt,
839 &tokens[next_start..next_end],
840 cache0_stage,
841 seq_end,
842 fence,
843 next_base,
844 true,
845 )
846 .map_err(|err| err.to_string())?;
847 cache0_stage.pos = initial_base + next_end;
848 Ok(next)
849 });
850 let x = self.prime_pp2_stage1_enqueue(
851 e,
852 rt,
853 slot,
854 end - start,
855 cache1,
856 seq_end,
857 fence,
858 base,
859 true,
860 )?;
861 let out = {
862 rt.bind_stage(1)?;
863 let _st1 = rt.enter(1);
864 let e1 = rt.engine(1, e);
865 self.prime_chunk_epilogue(e1, x, end - start, cache1)?
866 };
867 let next = stage0
868 .join()
869 .map_err(|_| "pipeprime stage-0 host walker panicked")?
870 .map_err(|err| -> Box<dyn std::error::Error> { err.into() })?;
871 Ok((out, Some(next)))
872 },
873 )?
874 } else {
875 let x = self.prime_pp2_stage1_enqueue(
876 e,
877 rt,
878 slot,
879 end - start,
880 cache1,
881 seq_end,
882 fence,
883 base,
884 true,
885 )?;
886 let out = {
887 rt.bind_stage(1)?;
888 let _st1 = rt.enter(1);
889 let e1 = rt.engine(1, e);
890 self.prime_chunk_epilogue(e1, x, end - start, cache1)?
891 };
892 (out, None)
893 };
894
895 rt.publish_to(1, &caller_stream)?;
896 e.copy_into(
897 &mut hiddens,
898 start * n_embd,
899 &out.2,
900 (end - start) * n_embd,
901 )?;
902 last = Some((out.0, out.1));
903 crate::pp::PRIME_SPLIT_CHUNKS
904 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
905
906 if let Some(next) = next_slot {
907 rt.fence_stages_behind(&caller_stream)?;
912 slot = next;
913 }
914 }
915
916 debug_assert_eq!(cache0.pos, initial_base + t);
917 debug_assert_eq!(cache1.pos, initial_base + t);
918 let (logits, h_seed) = last.unwrap();
919 Ok((logits, h_seed, hiddens))
920 }
921
922 fn gdn_hk(e: &Engine, t: usize, num_v: usize, num_k: usize) -> usize {
929 if Engine::gdn_db_on()
930 && Engine::gdn_chunked_enabled() && t >= 16
931 && e.gdn_mma_enabled(Engine::gdn_chunk_size())
932 && num_k * 2 == num_v
933 {
934 num_k
935 } else {
936 num_v
937 }
938 }
939
940 fn f16out_on(e: &Engine, t: usize) -> bool {
945 crate::f16_ffi::pp_f16_enabled() && t >= 16 && !e.verify_exact_on()
946 && std::env::var("MEMRA_F16OUT").as_deref() != Ok("0")
947 }
948
949 pub fn prime_slabs_get(
957 &self,
958 e: &Engine,
959 t: usize,
960 n_embd: usize,
961 n_ff_max: usize,
962 ) -> Result<std::sync::Arc<std::sync::Mutex<PrimeSlabs>>, Box<dyn std::error::Error>> {
963 let mut slabs = self.prime_slabs.lock().unwrap();
964 let dev = e.ctx().ordinal();
965 let need_new = match slabs.get(&dev) {
966 None => true,
967 Some(sl) => sl.lock().unwrap().t_cap < t,
968 };
969 if need_new {
970 slabs.insert(dev, std::sync::Arc::new(std::sync::Mutex::new(PrimeSlabs {
971 t_cap: t,
972 h: e.uninit(t * n_embd)?,
973 x1: e.uninit(t * n_embd)?,
974 z: e.uninit(t * n_embd)?,
975 act: e.uninit(t * n_ff_max)?,
976 xa: e.uninit(t * n_embd)?,
977 xb: e.uninit(t * n_embd)?,
978 h16: e.alloc_u8_uninit(t * n_embd * 2)?,
979 z16: e.alloc_u8_uninit(t * n_embd * 2)?,
980 gate: e.uninit(t * n_ff_max)?,
981 up: e.uninit(t * n_ff_max)?,
982 ffn_out: e.uninit(t * n_embd)?,
983 seg_glue: Vec::new(),
984 mixed: e.uninit(t * n_embd)?,
985 seg_mid: Vec::new(),
986 seg_t: 0,
987 })));
988 }
989 Ok(slabs.get(&dev).expect("prime slab inserted").clone())
990 }
991
992 fn prime_chunk(&self, e: &Engine, tokens: &[u32], cache: &mut Cache, seq_end: usize)
996 -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
997 if self.cfg.gemma4.is_none()
1006 && !crate::pp::pp2_streams_off()
1007 && crate::pp::prime_pp_on()
1008 {
1009 if let Some(fence) = crate::pp::pp_cuts(self.layers.len()) {
1010 return self.prime_chunk_ppn(e, tokens, cache, seq_end, &fence);
1011 }
1012 }
1013 let t = tokens.len();
1014 let base = cache.pos;
1015 debug_assert!(seq_end >= base + t, "prime_chunk: seq_end must cover this chunk");
1016 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
1017 let pos_d = e.htod_i32(&pos)?;
1018
1019 let x_embed = self.embed(e, tokens)?; let x = self.prime_layers(
1021 e, x_embed, 0, self.layers.len(), &pos_d, t, base, cache, seq_end,
1022 )?;
1023 self.prime_chunk_epilogue(e, x, t, cache)
1024 }
1025
1026 #[allow(clippy::too_many_arguments)]
1042 fn prime_layers(&self, e: &Engine, x_in: CudaSlice<f32>, lo: usize, hi: usize,
1043 pos_d: &CudaSlice<i32>, t: usize, base: usize, cache: &mut Cache,
1044 seq_end: usize)
1045 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1046 let cfg = &self.cfg;
1047 let n_embd = cfg.n_embd as usize;
1048 let eps = cfg.rms_eps;
1049 let f16fuse = crate::f16_ffi::pp_f16_enabled() && t >= 16;
1053 let n_ff_max = self.layers.iter().map(|l| match &l.ffn {
1059 crate::hybrid::Ffn::Dense { ffn_gate, .. } => ffn_gate.out_features(),
1060 _ => n_embd,
1061 }).max().unwrap_or(n_embd).max(n_embd);
1062 let use_slabs = std::env::var("MEMRA_PRIME_SLABS").as_deref() != Ok("0");
1063 let slab = if use_slabs {
1064 Some(self.prime_slabs_get(e, t, n_embd, n_ff_max)?)
1065 } else {
1066 None
1067 };
1068 let mut slab_guard = slab.as_ref().map(|sl| sl.lock().unwrap());
1069 let mut x_own; type SlabRefs<'a> = (&'a mut CudaSlice<f32>, &'a mut CudaSlice<f32>, &'a mut CudaSlice<f32>, &'a mut CudaSlice<f32>, &'a mut CudaSlice<u8>, &'a mut CudaSlice<u8>, &'a mut CudaSlice<f32>, &'a mut CudaSlice<f32>, &'a mut CudaSlice<f32>);
1071 let (mut x_cur, mut x_nxt, sl): (&mut CudaSlice<f32>, &mut CudaSlice<f32>, Option<SlabRefs>);
1072 let mut seg: Option<(&mut Vec<Option<cudarc::driver::CudaGraph>>, &mut Vec<Option<cudarc::driver::CudaGraph>>, &mut CudaSlice<f32>, &mut usize)> = None;
1073 let mut x_own2;
1074 match slab_guard.as_mut() {
1075 Some(g) => {
1076 let slabs = &mut **g;
1077 e.copy_into(&mut slabs.xa, 0, &x_in, t * n_embd)?;
1078 let PrimeSlabs { xa, xb, h, x1, z, act, h16, z16, gate, up, ffn_out, seg_glue, mixed, seg_mid, seg_t, .. } = slabs;
1079 x_cur = xa;
1080 x_nxt = xb;
1081 seg = Some((seg_glue, seg_mid, mixed, seg_t));
1082 sl = Some((h, x1, z, act, h16, z16, gate, up, ffn_out));
1083 }
1084 None => {
1085 x_own = x_in;
1086 x_own2 = e.uninit(t * n_embd)?;
1087 x_cur = &mut x_own;
1088 x_nxt = &mut x_own2;
1089 sl = None;
1090 }
1091 }
1092 let mut alloc_h; let mut alloc_x1; let mut alloc_z; let mut alloc_act;
1093 let mut alloc_h16; let mut alloc_z16;
1094 let mut alloc_gate; let mut alloc_up; let mut alloc_fo;
1095 let (h, x1, z, act): (&mut CudaSlice<f32>, &mut CudaSlice<f32>, &mut CudaSlice<f32>, &mut CudaSlice<f32>);
1096 let (h16, z16): (&mut CudaSlice<u8>, &mut CudaSlice<u8>);
1097 let (sl_gate, sl_up, sl_fo): (&mut CudaSlice<f32>, &mut CudaSlice<f32>, &mut CudaSlice<f32>);
1098 match sl {
1099 Some((a, b, c, d, e16, f16b, g, u, fo)) => {
1100 h = a; x1 = b; z = c; act = d; h16 = e16; z16 = f16b;
1101 sl_gate = g; sl_up = u; sl_fo = fo;
1102 }
1103 None => {
1104 alloc_h = e.uninit(t * n_embd)?;
1105 alloc_x1 = e.uninit(t * n_embd)?;
1106 alloc_z = e.uninit(t * n_embd)?;
1107 alloc_act = e.uninit(t * n_ff_max)?;
1108 alloc_h16 = e.alloc_u8_uninit(t * n_embd * 2)?;
1109 alloc_z16 = e.alloc_u8_uninit(t * n_embd * 2)?;
1110 alloc_gate = e.uninit(t * n_ff_max)?;
1111 alloc_up = e.uninit(t * n_ff_max)?;
1112 alloc_fo = e.uninit(t * n_embd)?;
1113 h = &mut alloc_h; x1 = &mut alloc_x1; z = &mut alloc_z; act = &mut alloc_act;
1114 h16 = &mut alloc_h16; z16 = &mut alloc_z16;
1115 sl_gate = &mut alloc_gate; sl_up = &mut alloc_up; sl_fo = &mut alloc_fo;
1116 }
1117 }
1118 let n_layers = self.layers.len();
1123 let use_seg = f16fuse && seg.is_some() && self.cfg.step35.is_none()
1133 && lo == 0 && hi == n_layers
1134 && std::env::var("MEMRA_PRIME_SEG").as_deref() == Ok("1");
1135 if let Some((sg, sm, _, st)) = seg.as_mut() {
1136 if **st != t {
1137 sg.clear();
1138 sg.extend((0..n_layers).map(|_| None));
1139 sm.clear();
1140 sm.extend((0..n_layers).map(|_| None));
1141 **st = t;
1142 }
1143 }
1144 {
1145 let layer_lo = &self.layers[lo];
1146 if f16fuse {
1147 e.rms_norm_f16out(x_cur, layer_lo.attn_norm.float_data(), h, h16, n_embd, t, eps)?;
1148 } else {
1149 e.rms_norm(x_cur, layer_lo.attn_norm.float_data(), h, n_embd, t, eps)?;
1150 }
1151 }
1152 for il in lo..hi {
1153 let layer = &self.layers[il];
1154 let hx16 = if f16fuse { Some(&*h16) } else { None };
1155 if use_seg {
1156 let (pre, pre16, w_out) = match &layer.mixer {
1159 Mixer::Full(fa) => {
1160 let g3 = match hx16 {
1161 Some(xh) => e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], h, xh, t)?,
1162 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
1163 };
1164 let (pre, pre16) = self.full_attn_prime_core_inner(e, fa, g3, &pos_d, t, cache, il)?;
1165 (pre, pre16, &fa.wo)
1166 }
1167 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
1168 Mixer::Linear(la) => {
1169 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
1170 let g4 = match hx16 {
1171 Some(xh) => e.matmul_group_xh(&ws, h, xh, t)?,
1172 None => e.matmul_group(&ws, h, t)?,
1173 };
1174 let (pre, pre16) = self.linear_attn_prime_core_pad_inner(e, la, g4, t, cache, il, None)?;
1175 (pre, pre16, &la.ssm_out)
1176 }
1177 };
1178 {
1179 let (_, sm, mslab, _) = seg.as_mut().unwrap();
1180 let pre_n = pre.len() / t;
1181 let xh_pre = match pre16 {
1182 Some(x) => x,
1183 None => e.f16_act(&pre, t * pre_n, pre_n)?,
1184 };
1185 if !e.try_f16_gemm_pre_into(w_out, &xh_pre, t, mslab)? {
1186 let y = e.matmul(w_out, &pre, t)?;
1187 e.copy_into(mslab, 0, &y, t * n_embd)?;
1188 }
1189 if sm[il].is_none() {
1190 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
1191 let w_post = layer.post_attn_norm.float_data();
1192 e.stream().synchronize()?;
1193 e.stream().begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
1194 let r = (|| -> Result<(), Box<dyn std::error::Error>> {
1195 e.add(x_cur, mslab, x1, t * n_embd)?;
1196 e.rms_norm_f16out(x1, w_post, z, z16, n_embd, t, eps)?;
1197 Ok(())
1198 })();
1199 let g = e.stream().end_capture(
1200 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH);
1201 r?;
1202 sm[il] = Some(g?.ok_or("S-mid capture produced no graph")?);
1203 }
1204 sm[il].as_ref().unwrap().launch()?;
1205 }
1206 } else {
1207 let mixed = match &layer.mixer {
1208 Mixer::Full(fa) => self.full_attn_prime(e, fa, h, hx16, &pos_d, t, cache, il,
1209 seq_end)?,
1210 Mixer::Linear(la) => self.linear_attn_prime(e, la, h, hx16, t, cache, il)?,
1211 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
1212 };
1213 if f16fuse {
1214 e.add_rms_norm_f16out(x_cur, &mixed, layer.post_attn_norm.float_data(),
1217 x1, z, z16, n_embd, t, eps)?;
1218 } else {
1219 e.add(x_cur, &mixed, x1, t * n_embd)?;
1220 e.rms_norm(x1, layer.post_attn_norm.float_data(), z, n_embd, t, eps)?;
1221 }
1222 }
1223 let zx16 = if f16fuse { Some(&*z16) } else { None };
1224 match &layer.ffn {
1225 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
1226 let n_ff = ffn_gate.out_features();
1227 let mut into_ok = false;
1230 if let Some(xh) = zx16 {
1231 into_ok = e.try_f16_gemm_pre_into(ffn_gate, xh, t, sl_gate)?
1232 && e.try_f16_gemm_pre_into(ffn_up, xh, t, sl_up)?;
1233 }
1234 if !into_ok {
1235 let mut g2 = match zx16 {
1236 Some(xh) => e.matmul_group_xh(&[ffn_gate, ffn_up], z, xh, t)?,
1237 None => e.matmul_group(&[ffn_gate, ffn_up], z, t)?,
1238 };
1239 let up_y = g2.pop().unwrap();
1240 let gate_y = g2.pop().unwrap();
1241 e.copy_into(sl_gate, 0, &gate_y, t * n_ff)?;
1242 e.copy_into(sl_up, 0, &up_y, t * n_ff)?;
1243 }
1244 let d_lim = self.cfg.clamp_shexp_at(il as u32);
1249 let act16 = if Self::f16out_on(e, t) && self.cfg.m3.is_none()
1250 && d_lim.is_none() {
1251 let mut a16 = e.alloc_u8_uninit(t * n_ff * 2)?;
1252 e.silu_mul_f16out(sl_gate, sl_up, act, &mut a16, t * n_ff)?;
1253 Some(a16)
1254 } else {
1255 Self::ffn_act_lim(e, &self.cfg, sl_gate, sl_up, 1.0, 1.0, d_lim,
1256 act, t * n_ff)?;
1257 None
1258 };
1259 let xh_act = match act16 {
1261 Some(x) => x,
1262 None => e.f16_act(act, t * n_ff, n_ff)?,
1263 };
1264 if !e.try_f16_gemm_pre_into(ffn_down, &xh_act, t, sl_fo)? {
1265 let y = e.matmul(ffn_down, &*act, t)?;
1266 e.copy_into(sl_fo, 0, &y, t * n_embd)?;
1267 }
1268 }
1269 crate::hybrid::Ffn::Moe(m) => {
1270 let y = self.moe_ffn_il_prefill(e, m, z, t, il as u16)?;
1271 e.copy_into(sl_fo, 0, &y, t * n_embd)?;
1272 }
1273 }
1274 if use_seg && il + 1 < hi {
1275 let w_next = self.layers[il + 1].attn_norm.float_data();
1277 let (sg, _, _, _) = seg.as_mut().unwrap();
1278 if sg[il].is_none() {
1279 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
1280 e.stream().synchronize()?;
1281 e.stream().begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
1282 let r = (|| -> Result<(), Box<dyn std::error::Error>> {
1283 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
1284 e.rms_norm_f16out(x_nxt, w_next, h, h16, n_embd, t, eps)?;
1285 Ok(())
1286 })();
1287 let g = e.stream().end_capture(
1288 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH);
1289 r?;
1290 sg[il] = Some(g?.ok_or("S-glue capture produced no graph")?);
1291 }
1292 sg[il].as_ref().unwrap().launch()?;
1293 } else {
1294 if il + 1 < hi {
1295 let w_next = self.layers[il + 1].attn_norm.float_data();
1296 if f16fuse {
1297 e.add_rms_norm_f16out(x1, sl_fo, w_next, x_nxt, h, h16, n_embd, t, eps)?;
1298 } else {
1299 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
1300 e.rms_norm(x_nxt, w_next, h, n_embd, t, eps)?;
1301 }
1302 } else {
1303 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
1304 }
1305 }
1306 if let Some(path) = Self::prime_trace_path() {
1312 let row = (base + t - 1) as usize;
1313 let host = e.dtoh(x_nxt)?;
1314 let last = &host[(t - 1) * n_embd..t * n_embd];
1315 use std::io::Write as _;
1316 let mut f = std::fs::OpenOptions::new().create(true).append(true).open(path)?;
1317 let mut h64: u64 = 0xcbf29ce484222325;
1318 for v in last {
1319 h64 ^= v.to_bits() as u64;
1320 h64 = h64.wrapping_mul(0x100000001b3);
1321 }
1322 writeln!(f, "{{\"pos\":{row},\"layer\":{il},\"t\":{t},\"base\":{base},\
1323 \"hash\":\"{h64:016x}\",\"v0\":{:.9e},\"v1\":{:.9e},\"v2\":{:.9e}}}",
1324 last[0], last[1], last[2])?;
1325 }
1326 std::mem::swap(&mut x_cur, &mut x_nxt);
1327 }
1328 let mut x = e.uninit(t * n_embd)?;
1330 e.copy_into(&mut x, 0, x_cur, t * n_embd)?;
1331 drop(slab_guard);
1332 Ok(x)
1333 }
1334
1335 fn prime_chunk_epilogue(&self, e: &Engine, x: CudaSlice<f32>, t: usize, cache: &mut Cache)
1340 -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
1341 let n_embd = self.cfg.n_embd as usize;
1342 let eps = self.cfg.rms_eps;
1343 let mut h_seed = e.uninit(n_embd)?;
1347 if !crate::spec::spec_hpost() {
1348 e.copy_view_into(&mut h_seed, 0, &x.slice((t - 1) * n_embd..t * n_embd), n_embd)?;
1349 }
1350 let mut hn = e.uninit(t * n_embd)?;
1352 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
1353 if crate::spec::spec_hpost() {
1354 e.copy_view_into(&mut h_seed, 0, &hn.slice((t - 1) * n_embd..t * n_embd), n_embd)?;
1355 }
1356 let last = e.view(&hn, t * n_embd);
1357 let last_row = last.slice((t - 1) * n_embd..t * n_embd);
1358 let mut hlast = e.uninit(n_embd)?;
1359 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
1360 let logits = e.matmul(&self.output, &hlast, 1)?;
1361 cache.pos += t;
1362 Ok((e.dtoh(&logits)?, h_seed, if crate::spec::spec_hpost() { hn } else { x }))
1365 }
1366
1367 fn prime_chunk_ppn(&self, e: &Engine, tokens: &[u32], cache: &mut Cache, seq_end: usize,
1391 fence: &[usize])
1392 -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
1393 let rt = crate::pp::PpNRt::get(e)?;
1394 let n_st = fence.len() - 1;
1395 assert_eq!(
1396 rt.n_stages(), n_st,
1397 "PpNRt stage count {} != fence stages {n_st}", rt.n_stages()
1398 );
1399 let n_embd = self.cfg.n_embd as usize;
1400 let t = tokens.len();
1401 let base = cache.pos;
1402 debug_assert!(seq_end >= base + t, "prime_chunk_ppn: seq_end must cover this chunk");
1403 let payload = t * n_embd;
1404 let caller_stream = e.stream();
1408 rt.fence_stages_behind(&caller_stream)?;
1409
1410 if n_st == 2 {
1411 let slot = self.prime_pp2_stage0_enqueue(
1412 e, rt, tokens, cache, seq_end, fence, base, false,
1413 )?;
1414 let x = self.prime_pp2_stage1_enqueue(
1415 e, rt, slot, t, cache, seq_end, fence, base, false,
1416 )?;
1417 let out = {
1418 rt.bind_stage(1)?;
1419 let _st1 = rt.enter(1);
1420 let e1 = rt.engine(1, e);
1421 self.prime_chunk_epilogue(e1, x, t, cache)?
1422 };
1423 rt.publish_to(1, &caller_stream)?;
1424 crate::pp::PRIME_SPLIT_CHUNKS
1425 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
1426 return Ok(out);
1427 }
1428
1429 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
1430
1431 let mut slot = {
1433 let _st0 = rt.enter(0);
1434 let e0 = rt.engine(0, e);
1435 let pos_d = e0.htod_i32(&pos)?;
1436 let x = self.embed(e0, tokens)?;
1437 let x = self.prime_layers(
1438 e0, x, fence[0], fence[1], &pos_d, t, base, cache, seq_end,
1439 )?;
1440 rt.tx(0, &x, payload)?
1441 };
1443
1444 for s in 1..n_st - 1 {
1446 let _st = rt.enter(s);
1447 let es = rt.engine(s, e);
1448 let pos_d = es.htod_i32(&pos)?;
1449 let x = rt.rx(s - 1, slot, payload)?;
1450 let x = self.prime_layers(
1451 es, x, fence[s], fence[s + 1], &pos_d, t, base, cache, seq_end,
1452 )?;
1453 slot = rt.tx(s, &x, payload)?;
1454 }
1455
1456 let _stl = rt.enter(n_st - 1);
1458 let el = rt.engine(n_st - 1, e);
1459 let pos_d = el.htod_i32(&pos)?;
1460 let x = rt.rx(n_st - 2, slot, payload)?;
1461 let x = self.prime_layers(
1462 el, x, fence[n_st - 1], fence[n_st], &pos_d, t, base, cache, seq_end,
1463 )?;
1464 let out = self.prime_chunk_epilogue(el, x, t, cache)?;
1465 rt.publish_to(n_st - 1, &caller_stream)?;
1471 crate::pp::PRIME_SPLIT_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
1472 Ok(out)
1473 }
1474
1475 fn prime_pp2_stage0_enqueue(
1476 &self,
1477 e: &Engine,
1478 rt: &crate::pp::PpNRt,
1479 tokens: &[u32],
1480 cache: &mut Cache,
1481 seq_end: usize,
1482 fence: &[usize],
1483 base: usize,
1484 pipelined: bool,
1485 ) -> Result<usize, Box<dyn std::error::Error>> {
1486 let t = tokens.len();
1487 let n_embd = self.cfg.n_embd as usize;
1488 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
1489 rt.bind_stage(0)?;
1490 let _st0 = rt.enter(0);
1491 let e0 = rt.engine(0, e);
1492 let pos_d = e0.htod_i32(&pos)?;
1493 let x = self.embed(e0, tokens)?;
1494 let _overlap = pipelined.then(crate::pp::enter_prime_pipe_stage);
1495 let x = self.prime_layers(
1496 e0, x, fence[0], fence[1], &pos_d, t, base, cache, seq_end,
1497 )?;
1498 if pipelined {
1499 rt.tx_pipelined(0, &x, t * n_embd)
1500 } else {
1501 rt.tx(0, &x, t * n_embd)
1502 }
1503 }
1504
1505 fn prime_pp2_stage1_enqueue(
1506 &self,
1507 e: &Engine,
1508 rt: &crate::pp::PpNRt,
1509 slot: usize,
1510 t: usize,
1511 cache: &mut Cache,
1512 seq_end: usize,
1513 fence: &[usize],
1514 base: usize,
1515 pipelined: bool,
1516 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1517 let n_embd = self.cfg.n_embd as usize;
1518 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
1519 rt.bind_stage(1)?;
1520 let _st1 = rt.enter(1);
1521 let e1 = rt.engine(1, e);
1522 let pos_d = e1.htod_i32(&pos)?;
1523 let x = rt.rx(0, slot, t * n_embd)?;
1524 let _overlap = pipelined.then(crate::pp::enter_prime_pipe_stage);
1525 self.prime_layers(
1526 e1, x, fence[1], fence[2], &pos_d, t, base, cache, seq_end,
1527 )
1528 }
1529
1530 pub fn prime_chunk_captured(&self, e: &Engine, x_in: &CudaSlice<f32>, pos_d: &CudaSlice<i32>,
1546 t: usize, cache: &mut Cache,
1547 len_d: &CudaSlice<i32>,
1548 logits_out: &mut CudaSlice<f32>, h_seed_out: &mut CudaSlice<f32>)
1549 -> Result<(), Box<dyn std::error::Error>> {
1550 let cfg = &self.cfg;
1551 let n_embd = cfg.n_embd as usize;
1552 let eps = cfg.rms_eps;
1553 let f16fuse = crate::f16_ffi::pp_f16_enabled() && t >= 16;
1554 let mut x = e.uninit(t * n_embd)?;
1555 e.copy_into(&mut x, 0, x_in, t * n_embd)?;
1556 for (il, layer) in self.layers.iter().enumerate() {
1557 let mut h = e.uninit(t * n_embd)?;
1558 let mut hx16: Option<CudaSlice<u8>> = None;
1559 if f16fuse {
1560 let mut b16 = e.alloc_u8_uninit(t * n_embd * 2)?;
1561 e.rms_norm_f16out(&x, layer.attn_norm.float_data(), &mut h, &mut b16, n_embd, t, eps)?;
1562 hx16 = Some(b16);
1563 } else {
1564 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
1565 }
1566 let mixed = match &layer.mixer {
1567 Mixer::Full(fa) => self.full_attn_prime(e, fa, &h, hx16.as_ref(), pos_d, t, cache,
1571 il, t)?,
1572 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
1573 Mixer::Linear(la) => {
1574 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
1575 let g4 = match hx16.as_ref() {
1576 Some(xh) => e.matmul_group_xh(&ws, &h, xh, t)?,
1577 None => e.matmul_group(&ws, &h, t)?,
1578 };
1579 self.linear_attn_prime_core_pad(e, la, g4, t, cache, il, Some(len_d))?
1580 }
1581 };
1582 let mut x1 = e.uninit(t * n_embd)?;
1583 e.add(&x, &mixed, &mut x1, t * n_embd)?;
1584 let mut z = e.uninit(t * n_embd)?;
1585 let mut zx16: Option<CudaSlice<u8>> = None;
1586 if f16fuse {
1587 let mut b16 = e.alloc_u8_uninit(t * n_embd * 2)?;
1588 e.rms_norm_f16out(&x1, layer.post_attn_norm.float_data(), &mut z, &mut b16, n_embd, t, eps)?;
1589 zx16 = Some(b16);
1590 } else {
1591 e.rms_norm(&x1, layer.post_attn_norm.float_data(), &mut z, n_embd, t, eps)?;
1592 }
1593 let ffn_out = match &layer.ffn {
1594 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
1595 let n_ff = ffn_gate.out_features();
1596 let mut g2 = match &zx16 {
1597 Some(xh) => e.matmul_group_xh(&[ffn_gate, ffn_up], &z, xh, t)?,
1598 None => e.matmul_group(&[ffn_gate, ffn_up], &z, t)?,
1599 };
1600 let up = g2.pop().unwrap();
1601 let gate = g2.pop().unwrap();
1602 let mut act = e.uninit(t * n_ff)?;
1603 Self::ffn_act_lim(e, &self.cfg, &gate, &up, 1.0, 1.0,
1605 self.cfg.clamp_shexp_at(il as u32), &mut act, t * n_ff)?;
1606 e.matmul(ffn_down, &act, t)?
1607 }
1608 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il_prefill(e, m, &z, t, il as u16)?,
1609 };
1610 let mut x2 = e.uninit(t * n_embd)?;
1611 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
1612 x = x2;
1613 }
1614 if !crate::spec::spec_hpost() {
1616 e.row_gather_dev(&x, h_seed_out, len_d, n_embd)?;
1617 }
1618 let mut hn = e.uninit(t * n_embd)?;
1619 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
1620 if crate::spec::spec_hpost() {
1621 e.row_gather_dev(&hn, h_seed_out, len_d, n_embd)?;
1622 }
1623 let mut hlast = e.uninit(n_embd)?;
1624 e.row_gather_dev(&hn, &mut hlast, len_d, n_embd)?;
1625 let logits = e.matmul(&self.output, &hlast, 1)?;
1626 let nv = logits.len();
1627 e.copy_into(logits_out, 0, &logits, nv)?;
1628 Ok(())
1629 }
1630
1631 fn step35_prime_batch_on() -> bool {
1632 std::env::var("MEMRA_STEP35_PRIME_BATCH").as_deref() != Ok("0")
1633 }
1634
1635 #[allow(clippy::too_many_arguments)]
1638 fn step35_prime_batch_layers(
1639 &self,
1640 e: &Engine,
1641 mut x: CudaSlice<f32>,
1642 lo: usize,
1643 hi: usize,
1644 ts: &[usize],
1645 offs: &[usize],
1646 pos_ds: &[CudaSlice<i32>],
1647 caches: &mut [&mut Cache],
1648 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1649 let cfg = &self.cfg;
1650 let n_embd = cfg.n_embd as usize;
1651 let eps = cfg.rms_eps;
1652 let b = ts.len();
1653 let total: usize = ts.iter().sum();
1654 let f16fuse = crate::f16_ffi::pp_f16_enabled() && total >= 16;
1655
1656 let split = |e: &Engine, y: &CudaSlice<f32>, dim: usize|
1657 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
1658 let mut out = Vec::with_capacity(b);
1659 for s in 0..b {
1660 let mut ys = e.uninit(ts[s] * dim)?;
1661 e.copy_view_into(
1662 &mut ys,
1663 0,
1664 &y.slice(offs[s] * dim..(offs[s] + ts[s]) * dim),
1665 ts[s] * dim,
1666 )?;
1667 out.push(ys);
1668 }
1669 Ok(out)
1670 };
1671
1672 for il in lo..hi {
1673 let layer = &self.layers[il];
1674 let Mixer::Full(fa) = &layer.mixer else {
1675 return Err(format!("step35 layer {il} is not full-attn — corrupt config").into());
1676 };
1677
1678 let mut h = e.uninit(total * n_embd)?;
1679 let mut hx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
1680 if f16fuse {
1681 e.rms_norm_f16out(
1682 &x,
1683 layer.attn_norm.float_data(),
1684 &mut h,
1685 &mut hx16,
1686 n_embd,
1687 total,
1688 eps,
1689 )?;
1690 } else {
1691 e.rms_norm(
1692 &x,
1693 layer.attn_norm.float_data(),
1694 &mut h,
1695 n_embd,
1696 total,
1697 eps,
1698 )?;
1699 }
1700
1701 let gate_w = fa
1705 .attn_gate
1706 .as_ref()
1707 .ok_or("step35 layer is missing attn_gate.weight (head-wise attention gate)")?;
1708 let mut g4 = if f16fuse {
1709 e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv, gate_w], &h, &hx16, total)?
1710 } else {
1711 e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv, gate_w], &h, total)?
1712 };
1713 let gate = g4.pop().unwrap();
1714 let mut parts: Vec<Vec<CudaSlice<f32>>> =
1715 (0..b).map(|_| Vec::with_capacity(3)).collect();
1716 for (w, y) in [&fa.wq, &fa.wk, &fa.wv].iter().zip(g4) {
1717 for (s, ys) in split(e, &y, w.out_features())?.into_iter().enumerate() {
1718 parts[s].push(ys);
1719 }
1720 }
1721 let gates = split(e, &gate, gate_w.out_features())?;
1722 let geometry = self.step35_geom(il);
1723 let hd = geometry.head_dim_k as usize;
1724 let nh = geometry.n_head as usize;
1725 let mut ag_cat = e.uninit(total * nh * hd)?;
1726 for (s, (g3s, gate)) in parts.into_iter().zip(gates).enumerate() {
1727 let ag = self.step35_attn_pre_wo(
1728 e,
1729 fa,
1730 g3s,
1731 None,
1732 Some(&gate),
1733 &pos_ds[s],
1734 ts[s],
1735 Some(&mut *caches[s]),
1736 il,
1737 ts[s],
1738 )?;
1739 e.copy_into(
1740 &mut ag_cat,
1741 offs[s] * nh * hd,
1742 &ag,
1743 ts[s] * nh * hd,
1744 )?;
1745 }
1746 let mixed = e.matmul(&fa.wo, &ag_cat, total)?;
1747
1748 let mut x1 = e.uninit(total * n_embd)?;
1749 let mut z = e.uninit(total * n_embd)?;
1750 let mut zx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
1751 if f16fuse {
1752 e.add_rms_norm_f16out(
1753 &x,
1754 &mixed,
1755 layer.post_attn_norm.float_data(),
1756 &mut x1,
1757 &mut z,
1758 &mut zx16,
1759 n_embd,
1760 total,
1761 eps,
1762 )?;
1763 } else {
1764 e.add(&x, &mixed, &mut x1, total * n_embd)?;
1765 e.rms_norm(
1766 &x1,
1767 layer.post_attn_norm.float_data(),
1768 &mut z,
1769 n_embd,
1770 total,
1771 eps,
1772 )?;
1773 }
1774
1775 let ffn_out = match &layer.ffn {
1776 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
1777 let n_ff = ffn_gate.out_features();
1778 let mut g2 = if f16fuse {
1779 e.matmul_group_xh(&[ffn_gate, ffn_up], &z, &zx16, total)?
1780 } else {
1781 e.matmul_group(&[ffn_gate, ffn_up], &z, total)?
1782 };
1783 let up = g2.pop().unwrap();
1784 let gate = g2.pop().unwrap();
1785 let mut act = e.uninit(total * n_ff)?;
1786 let d_lim = cfg.clamp_shexp_at(il as u32);
1787 if Self::f16out_on(e, total) && cfg.m3.is_none() && d_lim.is_none() {
1788 let mut a16 = e.alloc_u8_uninit(total * n_ff * 2)?;
1789 e.silu_mul_f16out(&gate, &up, &mut act, &mut a16, total * n_ff)?;
1790 match e.try_f16_gemm_pre(ffn_down, &a16, total)? {
1791 Some(y) => y,
1792 None => e.matmul(ffn_down, &act, total)?,
1793 }
1794 } else {
1795 Self::ffn_act_lim(
1796 e,
1797 cfg,
1798 &gate,
1799 &up,
1800 1.0,
1801 1.0,
1802 d_lim,
1803 &mut act,
1804 total * n_ff,
1805 )?;
1806 e.matmul(ffn_down, &act, total)?
1807 }
1808 }
1809 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il(e, m, &z, total, il as u16)?,
1810 };
1811 let mut x2 = e.uninit(total * n_embd)?;
1812 e.add(&x1, &ffn_out, &mut x2, total * n_embd)?;
1813 x = x2;
1814 }
1815 Ok(x)
1816 }
1817
1818 fn step35_prime_batch_epilogue(
1819 &self,
1820 e: &Engine,
1821 x: CudaSlice<f32>,
1822 ts: &[usize],
1823 offs: &[usize],
1824 caches: &mut [&mut Cache],
1825 ) -> Result<Vec<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
1826 let n_embd = self.cfg.n_embd as usize;
1827 let total: usize = ts.iter().sum();
1828 let mut hn = e.uninit(total * n_embd)?;
1829 e.rms_norm(
1830 &x,
1831 self.output_norm.float_data(),
1832 &mut hn,
1833 n_embd,
1834 total,
1835 self.cfg.rms_eps,
1836 )?;
1837
1838 let hidden_src = if crate::spec::spec_hpost() { &hn } else { &x };
1839 let mut out = Vec::with_capacity(ts.len());
1840 for s in 0..ts.len() {
1841 let mut hidden = e.uninit(ts[s] * n_embd)?;
1842 e.copy_view_into(
1843 &mut hidden,
1844 0,
1845 &hidden_src.slice(offs[s] * n_embd..(offs[s] + ts[s]) * n_embd),
1846 ts[s] * n_embd,
1847 )?;
1848 let last0 = (offs[s] + ts[s] - 1) * n_embd;
1849 let mut h_seed = e.uninit(n_embd)?;
1850 e.copy_view_into(
1851 &mut h_seed,
1852 0,
1853 &hidden_src.slice(last0..last0 + n_embd),
1854 n_embd,
1855 )?;
1856 let mut hlast = e.uninit(n_embd)?;
1858 e.copy_view_into(
1859 &mut hlast,
1860 0,
1861 &hn.slice(last0..last0 + n_embd),
1862 n_embd,
1863 )?;
1864 let logits = e.dtoh(&e.matmul(&self.output, &hlast, 1)?)?;
1865 caches[s].pos += ts[s];
1866 out.push((logits, h_seed, hidden));
1867 }
1868 Ok(out)
1869 }
1870
1871 fn step35_prime_cache_batch(
1872 &self,
1873 e: &Engine,
1874 prompts: &[&[u32]],
1875 caches: &mut [&mut Cache],
1876 ) -> Result<Vec<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
1877 if !Self::step35_prime_batch_on() {
1878 return Err("step35 batched prime is disabled (MEMRA_STEP35_PRIME_BATCH=0)".into());
1879 }
1880 if caches.iter().any(|c| c.pos != 0) {
1881 return Err(
1882 "step35 batched prime currently supports complete fresh prompts only; \
1883 continuation/tick chunks require per-request queued_after"
1884 .into(),
1885 );
1886 }
1887
1888 let ts: Vec<usize> = prompts.iter().map(|p| p.len()).collect();
1889 for &t in &ts {
1890 assert!(t >= PRIME_MIN_T, "step35 batched prime needs T >= {PRIME_MIN_T}");
1891 }
1892 for (s, c) in caches.iter().enumerate() {
1893 assert!(ts[s] <= c.max_ctx, "step35 batched prime exceeds cache max_ctx");
1894 }
1895 let offs: Vec<usize> = ts
1896 .iter()
1897 .scan(0usize, |a, &t| {
1898 let o = *a;
1899 *a += t;
1900 Some(o)
1901 })
1902 .collect();
1903 let total: usize = ts.iter().sum();
1904 let payload = total * self.cfg.n_embd as usize;
1905 let cat_tokens: Vec<u32> = prompts.iter().flat_map(|p| p.iter().copied()).collect();
1906 let positions: Vec<Vec<i32>> = ts
1907 .iter()
1908 .map(|&t| (0..t as i32).collect())
1909 .collect();
1910 let upload_positions = |e: &Engine|
1911 -> Result<Vec<CudaSlice<i32>>, Box<dyn std::error::Error>> {
1912 positions
1913 .iter()
1914 .map(|p| e.htod_i32(p))
1915 .collect::<Result<_, _>>()
1916 };
1917
1918 static ONCE: std::sync::Once = std::sync::Once::new();
1919 ONCE.call_once(|| {
1920 eprintln!(
1921 "[step35-prime-batch] first concat prime: B={} tokens={total}",
1922 prompts.len()
1923 );
1924 });
1925
1926 let out = if !crate::pp::pp2_streams_off() && crate::pp::prime_pp_on() {
1927 if let Some(fence) = crate::pp::pp_cuts(self.layers.len()) {
1928 let rt = crate::pp::PpNRt::get(e)?;
1929 let n_st = fence.len() - 1;
1930 assert_eq!(rt.n_stages(), n_st, "step35 prime batch stage count mismatch");
1931 let caller_stream = e.stream();
1932 rt.fence_stages_behind(&caller_stream)?;
1933
1934 let mut slot = {
1935 let _st0 = rt.enter(0);
1936 let e0 = rt.engine(0, e);
1937 let pos_ds = upload_positions(e0)?;
1938 let x = self.embed(e0, &cat_tokens)?;
1939 let x = self.step35_prime_batch_layers(
1940 e0,
1941 x,
1942 fence[0],
1943 fence[1],
1944 &ts,
1945 &offs,
1946 &pos_ds,
1947 caches,
1948 )?;
1949 rt.tx(0, &x, payload)?
1950 };
1951 for s in 1..n_st - 1 {
1952 let _st = rt.enter(s);
1953 let es = rt.engine(s, e);
1954 let pos_ds = upload_positions(es)?;
1955 let x = rt.rx(s - 1, slot, payload)?;
1956 let x = self.step35_prime_batch_layers(
1957 es,
1958 x,
1959 fence[s],
1960 fence[s + 1],
1961 &ts,
1962 &offs,
1963 &pos_ds,
1964 caches,
1965 )?;
1966 slot = rt.tx(s, &x, payload)?;
1967 }
1968
1969 let _stl = rt.enter(n_st - 1);
1970 let el = rt.engine(n_st - 1, e);
1971 let pos_ds = upload_positions(el)?;
1972 let x = rt.rx(n_st - 2, slot, payload)?;
1973 let x = self.step35_prime_batch_layers(
1974 el,
1975 x,
1976 fence[n_st - 1],
1977 fence[n_st],
1978 &ts,
1979 &offs,
1980 &pos_ds,
1981 caches,
1982 )?;
1983 let out = self.step35_prime_batch_epilogue(el, x, &ts, &offs, caches)?;
1984 rt.publish_to(n_st - 1, &caller_stream)?;
1985 crate::pp::STEP35_PRIME_BATCH_SPLITS.fetch_add(
1986 1,
1987 std::sync::atomic::Ordering::Relaxed,
1988 );
1989 out
1990 } else {
1991 let pos_ds = upload_positions(e)?;
1992 let x = self.embed(e, &cat_tokens)?;
1993 let x = self.step35_prime_batch_layers(
1994 e,
1995 x,
1996 0,
1997 self.layers.len(),
1998 &ts,
1999 &offs,
2000 &pos_ds,
2001 caches,
2002 )?;
2003 self.step35_prime_batch_epilogue(e, x, &ts, &offs, caches)?
2004 }
2005 } else {
2006 let pos_ds = upload_positions(e)?;
2007 let x = self.embed(e, &cat_tokens)?;
2008 let x = self.step35_prime_batch_layers(
2009 e,
2010 x,
2011 0,
2012 self.layers.len(),
2013 &ts,
2014 &offs,
2015 &pos_ds,
2016 caches,
2017 )?;
2018 self.step35_prime_batch_epilogue(e, x, &ts, &offs, caches)?
2019 };
2020 crate::pp::STEP35_PRIME_BATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2021 Ok(out)
2022 }
2023
2024 pub fn prime_cache_batch(&self, e: &Engine, prompts: &[&[u32]], caches: &mut [&mut Cache])
2041 -> Result<Vec<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
2042 let cfg = &self.cfg;
2043 let n_embd = cfg.n_embd as usize;
2044 let eps = cfg.rms_eps;
2045 let b = prompts.len();
2046 assert!(b >= 1 && b == caches.len());
2047 let pos0s: Vec<usize> = caches.iter().map(|c| c.pos).collect();
2048 let carried = pos0s.iter().any(|&p| p > 0);
2049 if cfg.gemma4.is_some() {
2055 return Err("prime_cache_batch: gemma4 has no batched prime core (per-layer \
2056 swa/global geometry, softcapped head) — use gemma4_prime per sequence".into());
2057 }
2058 if cfg.step35.is_some() {
2061 return self.step35_prime_cache_batch(e, prompts, caches);
2062 }
2063 let ts: Vec<usize> = prompts.iter().map(|p| p.len()).collect();
2064 for &t in &ts { assert!(t >= PRIME_MIN_T, "prime_cache_batch needs T >= {PRIME_MIN_T}"); }
2065 for (s, c) in caches.iter().enumerate() {
2066 assert!(c.pos + ts[s] <= c.max_ctx, "prime_cache_batch: prompt exceeds cache max_ctx");
2067 }
2068 let total: usize = ts.iter().sum();
2069 let offs: Vec<usize> = ts.iter().scan(0usize, |a, &t| { let o = *a; *a += t; Some(o) }).collect();
2070 let pos_ds: Vec<CudaSlice<i32>> = ts.iter().zip(&pos0s)
2072 .map(|(&t, &p0)| e.htod_i32(&(p0 as i32..(p0 + t) as i32).collect::<Vec<_>>()))
2073 .collect::<Result<_, _>>()?;
2074 let split = |e: &Engine, y: &CudaSlice<f32>, dim: usize|
2076 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
2077 let mut out = Vec::with_capacity(b);
2078 for s in 0..b {
2079 let mut ys = e.uninit(ts[s] * dim)?;
2080 e.copy_view_into(&mut ys, 0, &y.slice(offs[s] * dim..(offs[s] + ts[s]) * dim), ts[s] * dim)?;
2081 out.push(ys);
2082 }
2083 Ok(out)
2084 };
2085
2086 let cat_tokens: Vec<u32> = prompts.iter().flat_map(|p| p.iter().copied()).collect();
2087 let mut x = self.embed(e, &cat_tokens)?; for (il, layer) in self.layers.iter().enumerate() {
2089 let mut h = e.uninit(total * n_embd)?;
2090 let mut hx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
2091 e.rms_norm_f16out(&x, layer.attn_norm.float_data(), &mut h, &mut hx16, n_embd, total, eps)?;
2092 let mut mixed = e.uninit(total * n_embd)?;
2094 match &layer.mixer {
2095 Mixer::Full(fa) => {
2096 let g3 = e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], &h, &hx16, total)?;
2097 let geometry = self.cfg.full_attention_geometry_at(il as u32);
2103 let (n_head, n_head_kv, head_dim) = (
2104 geometry.n_head as usize,
2105 geometry.n_head_kv as usize,
2106 geometry.head_dim_k as usize,
2107 );
2108 let fa_scale = geometry.attention_scale();
2109 let use_favl = !carried
2110 && (2..=8).contains(&b)
2111 && (head_dim == 256 || head_dim == 128)
2112 && geometry.attention_gate
2113 == memra_gguf::config::AttentionGateKind::FusedQ
2114 && std::env::var("MEMRA_NOFA").is_err()
2115 && std::env::var("MEMRA_FA_FLOOR").is_err()
2116 && std::env::var("MEMRA_FA_PP_W2").as_deref() != Ok("1")
2117 && std::env::var("MEMRA_FA_BF16KV").as_deref() != Ok("0")
2118 && std::env::var("MEMRA_FA_VL").as_deref() != Ok("0");
2119 if use_favl {
2120 let (qf_w, kf_w, vf_w) =
2121 (fa.wq.out_features(), fa.wk.out_features(), fa.wv.out_features());
2122 struct APre {
2123 q: CudaSlice<f32>, gate: Option<CudaSlice<f32>>,
2124 qn: CudaSlice<f32>, kn: CudaSlice<f32>,
2125 }
2126 let mut aps = Vec::with_capacity(b);
2127 for &t in ts.iter().take(b) {
2128 aps.push(APre {
2129 q: e.uninit(t * n_head * head_dim)?,
2130 gate: Some(e.uninit(t * n_head * head_dim)?),
2131 qn: e.uninit(t * n_head * head_dim)?,
2132 kn: e.uninit(t * n_head_kv * head_dim)?,
2133 });
2134 }
2135 let (kv_dim_k, kv_dim_v, ktb, vtb) = {
2136 let kvl = caches[0].kv[il].as_ref().unwrap();
2137 (kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes)
2138 };
2139 let pargs: Vec<crate::AttnPreVl> = (0..b).map(|s| {
2140 let (o, t) = (offs[s], ts[s]);
2141 let kvl = caches[s].kv[il].as_ref().unwrap();
2142 assert!(kvl.len == 0 && kvl.len + t <= caches[s].max_ctx,
2143 "prime_cache_batch attn vl: fresh + capacity");
2144 crate::AttnPreVl {
2145 qf: e.addr_f32v(&g3[0].slice(o * qf_w..(o + t) * qf_w)),
2146 kf: e.addr_f32v(&g3[1].slice(o * kf_w..(o + t) * kf_w)),
2147 vf: e.addr_f32v(&g3[2].slice(o * vf_w..(o + t) * vf_w)),
2148 q: e.addr_f32(&aps[s].q),
2149 gate: e.addr_f32(aps[s].gate.as_ref().unwrap()),
2150 qn: e.addr_f32(&aps[s].qn), kn: e.addr_f32(&aps[s].kn),
2151 kc: e.addr_u8(&kvl.k), vc: e.addr_u8(&kvl.v),
2152 t: t as i32, pad: 0,
2153 }
2154 }).collect();
2155 e.attn_pre_vl8(&pargs, fa.q_norm.float_data(), fa.k_norm.float_data(),
2156 head_dim, geometry.n_rot as usize, n_head, n_head_kv,
2157 self.cfg.rms_eps, geometry.rope_base, 1.0,
2158 kv_dim_k, kv_dim_v, ktb, vtb)?;
2159 for s in 0..b {
2160 let kvl = caches[s].kv[il].as_mut().unwrap();
2161 kvl.len += ts[s];
2162 let new_len = kvl.len as i32;
2163 e.set_i32_one(&mut kvl.len_d, new_len)?;
2164 }
2165 let mut attns = Vec::with_capacity(b);
2166 let mut mirrors = Vec::with_capacity(b);
2167 for &t in ts.iter().take(b) {
2168 attns.push(e.uninit(t * n_head * head_dim)?);
2169 let n = t * n_head_kv * head_dim;
2170 mirrors.push((e.alloc_u8_uninit(n * 2)?, e.alloc_u8_uninit(n * 2)?));
2171 }
2172 let fa3_on = match std::env::var("MEMRA_FA3").as_deref() {
2175 Ok("0") => false,
2176 Ok("1") => true,
2177 _ => cfg!(memra_hopper_mma),
2178 };
2179 if fa3_on {
2180 let mut q16s = Vec::with_capacity(b);
2181 let mut v16s = Vec::with_capacity(b);
2182 for s in 0..b {
2183 let t = ts[s];
2184 let mut q16 = e.alloc_u8_uninit(t * n_head * head_dim * 2)?;
2185 e.f32_to_bf16_into(&aps[s].qn, &mut q16, t * n_head * head_dim)?;
2186 let mut k16 = e.alloc_u8_uninit(t * n_head_kv * head_dim * 2)?;
2187 e.f32_to_bf16_into(&aps[s].kn, &mut k16, t * n_head_kv * head_dim)?;
2188 let mut v16 = e.alloc_u8_uninit(t * n_head_kv * head_dim * 2)?;
2189 e.f32_to_bf16_v(&g3[2].slice(offs[s] * vf_w..(offs[s] + t) * vf_w),
2190 &mut v16, t * n_head_kv * head_dim)?;
2191 q16s.push(q16);
2192 v16s.push((k16, v16));
2193 }
2194 let mut qp = [core::ptr::null::<core::ffi::c_void>(); 8];
2195 let mut kp = qp;
2196 let mut vp = qp;
2197 let mut op = [core::ptr::null_mut::<f32>(); 8];
2198 let mut tsv = [0i32; 8];
2199 for s in 0..b {
2200 qp[s] = e.addr_u8(&q16s[s]) as *const core::ffi::c_void;
2201 kp[s] = e.addr_u8(&v16s[s].0) as *const core::ffi::c_void;
2202 vp[s] = e.addr_u8(&v16s[s].1) as *const core::ffi::c_void;
2203 op[s] = e.addr_f32(&attns[s]) as *mut f32;
2204 tsv[s] = ts[s] as i32;
2205 }
2206 let rc = unsafe {
2207 crate::fa3_vl_raw(qp.as_ptr(), kp.as_ptr(), vp.as_ptr(), op.as_ptr(),
2208 tsv.as_ptr(), b as i32, n_head as i32,
2209 n_head_kv as i32, head_dim as i32, fa_scale,
2210 e.stream().cu_stream() as *mut core::ffi::c_void)
2211 };
2212 if rc != 0 {
2213 return Err(format!("memra_fa3_vl rc={rc}").into());
2214 }
2215 } else {
2216 let fargs: Vec<crate::FaSeqVl> = (0..b).map(|s| crate::FaSeqVl {
2217 q: e.addr_f32(&aps[s].qn), k16: e.addr_u8(&mirrors[s].0),
2218 v16: e.addr_u8(&mirrors[s].1), o: e.addr_f32(&attns[s]),
2219 kf: e.addr_f32(&aps[s].kn),
2220 vf: e.addr_f32v(&g3[2].slice(offs[s] * vf_w..(offs[s] + ts[s]) * vf_w)),
2221 t: ts[s] as i32, pad: 0,
2222 }).collect();
2223 e.fa_prefill_vl8(&fargs, head_dim, n_head, n_head_kv, fa_scale)?;
2224 }
2225 for (s, attn) in attns.into_iter().enumerate() {
2226 let (attn_g, ag16) = self.full_attn_prime_post_fa(
2227 e, attn, &aps[s].gate, ts[s], n_head, head_dim)?;
2228 let mut done = false;
2229 if let Some(xh) = &ag16 {
2230 done = e.try_f16_gemm_pre_into_off(&fa.wo, xh, ts[s], &mut mixed, offs[s] * n_embd)?;
2231 }
2232 if !done {
2233 let m = e.matmul(&fa.wo, &attn_g, ts[s])?;
2234 e.copy_into(&mut mixed, offs[s] * n_embd, &m, ts[s] * n_embd)?;
2235 }
2236 }
2237 } else {
2238 let mut parts: Vec<Vec<CudaSlice<f32>>> = (0..b).map(|_| Vec::new()).collect();
2239 for (w, y) in [&fa.wq, &fa.wk, &fa.wv].iter().zip(g3) {
2240 for (s, ys) in split(e, &y, w.out_features())?.into_iter().enumerate() {
2241 parts[s].push(ys);
2242 }
2243 }
2244 for (s, g3s) in parts.into_iter().enumerate() {
2245 let (attn_g, ag16) = self.full_attn_prime_core_inner(
2247 e, fa, g3s, &pos_ds[s], ts[s], caches[s], il)?;
2248 let mut done = false;
2249 if let Some(xh) = &ag16 {
2250 done = e.try_f16_gemm_pre_into_off(&fa.wo, xh, ts[s], &mut mixed, offs[s] * n_embd)?;
2251 }
2252 if !done {
2253 let m = e.matmul(&fa.wo, &attn_g, ts[s])?;
2254 e.copy_into(&mut mixed, offs[s] * n_embd, &m, ts[s] * n_embd)?;
2255 }
2256 }
2257 }
2258 }
2259 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
2260 Mixer::Linear(la) => {
2261 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
2266 let g4 = e.matmul_group_xh(&ws, &h, &hx16, total)?;
2267 let outs = self.linear_attn_prime_core_batch(e, la, &g4, &offs, &ts, caches, il)?;
2268 for (s, (gn, gn16)) in outs.into_iter().enumerate() {
2269 let (o, t) = (offs[s], ts[s]);
2270 let mut done = false;
2271 if let Some(xh) = &gn16 {
2272 done = e.try_f16_gemm_pre_into_off(&la.ssm_out, xh, t, &mut mixed, o * n_embd)?;
2273 }
2274 if !done {
2275 let m = e.matmul(&la.ssm_out, &gn, t)?;
2276 e.copy_into(&mut mixed, o * n_embd, &m, t * n_embd)?;
2277 }
2278 }
2279 }
2280 }
2281 let mut x1 = e.uninit(total * n_embd)?;
2282 let mut z = e.uninit(total * n_embd)?;
2283 let mut zx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
2284 e.add_rms_norm_f16out(&x, &mixed, layer.post_attn_norm.float_data(),
2285 &mut x1, &mut z, &mut zx16, n_embd, total, eps)?;
2286 let ffn_out = match &layer.ffn {
2287 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
2288 let n_ff = ffn_gate.out_features();
2289 let mut g2 = e.matmul_group_xh(&[ffn_gate, ffn_up], &z, &zx16, total)?;
2290 let up = g2.pop().unwrap();
2291 let gate = g2.pop().unwrap();
2292 let mut act = e.uninit(total * n_ff)?;
2293 let d_lim = self.cfg.clamp_shexp_at(il as u32);
2297 if Self::f16out_on(e, total) && self.cfg.m3.is_none() && d_lim.is_none() {
2298 let mut a16 = e.alloc_u8_uninit(total * n_ff * 2)?;
2299 e.silu_mul_f16out(&gate, &up, &mut act, &mut a16, total * n_ff)?;
2300 match e.try_f16_gemm_pre(ffn_down, &a16, total)? {
2301 Some(y) => y,
2302 None => e.matmul(ffn_down, &act, total)?,
2303 }
2304 } else {
2305 Self::ffn_act_lim(e, &self.cfg, &gate, &up, 1.0, 1.0, d_lim,
2306 &mut act, total * n_ff)?;
2307 e.matmul(ffn_down, &act, total)?
2308 }
2309 }
2310 crate::hybrid::Ffn::Moe(m) => {
2311 self.moe_ffn_il_prefill(e, m, &z, total, il as u16)?
2312 }
2313 };
2314 let mut x2 = e.uninit(total * n_embd)?;
2315 e.add(&x1, &ffn_out, &mut x2, total * n_embd)?;
2316 x = x2;
2317 }
2318 let mut hn = e.uninit(total * n_embd)?;
2320 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, total, eps)?;
2321 let mut hcat = e.uninit(b * n_embd)?;
2327 for s in 0..b {
2328 let last0 = (offs[s] + ts[s] - 1) * n_embd;
2329 e.copy_view_into(&mut hcat, s * n_embd, &hn.slice(last0..last0 + n_embd), n_embd)?;
2330 }
2331 let logits_cat = if b >= 2 { e.try_f16_gemm(&self.output, &hcat, b)? } else { None };
2332 let logits_host: Option<Vec<f32>> = match &logits_cat {
2333 Some(lc) => Some(e.dtoh(lc)?),
2334 None => None,
2335 };
2336 let n_vocab = self.output.out_features();
2337 let mut hidden_all = if crate::spec::spec_hpost() {
2338 split(e, &hn, n_embd)?
2339 } else {
2340 split(e, &x, n_embd)?
2341 };
2342 let mut out = Vec::with_capacity(b);
2343 for s in 0..b {
2344 let last0 = (offs[s] + ts[s] - 1) * n_embd;
2345 let mut h_seed = e.uninit(n_embd)?;
2346 if !crate::spec::spec_hpost() {
2347 e.copy_view_into(&mut h_seed, 0, &x.slice(last0..last0 + n_embd), n_embd)?;
2348 } else {
2349 e.copy_view_into(&mut h_seed, 0, &hn.slice(last0..last0 + n_embd), n_embd)?;
2350 }
2351 let logits = match &logits_host {
2352 Some(lh) => lh[s * n_vocab..(s + 1) * n_vocab].to_vec(),
2353 None => {
2354 let mut hlast = e.uninit(n_embd)?;
2355 e.copy_view_into(&mut hlast, 0, &hn.slice(last0..last0 + n_embd), n_embd)?;
2356 e.dtoh(&e.matmul(&self.output, &hlast, 1)?)?
2357 }
2358 };
2359 caches[s].pos += ts[s];
2360 out.push((logits, h_seed, hidden_all.remove(0)));
2361 }
2362 Ok(out)
2363 }
2364
2365 #[allow(clippy::too_many_arguments)]
2376 fn full_attn_prime(&self, e: &Engine, fa: &FullAttnLayer, h: &CudaSlice<f32>,
2377 hx: Option<&CudaSlice<u8>>,
2378 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache, il: usize,
2379 seq_end: usize)
2380 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2381 if self.cfg.step35.is_some() {
2382 return self.step35_attn_prime(e, fa, h, hx, pos_d, t, cache, il, seq_end);
2383 }
2384 let g3 = match hx {
2389 Some(xh) => e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], h, xh, t)?,
2390 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
2391 };
2392 self.full_attn_prime_core(e, fa, g3, pos_d, t, cache, il)
2393 }
2394
2395 fn full_attn_prime_core(&self, e: &Engine, fa: &FullAttnLayer, g3: Vec<CudaSlice<f32>>,
2399 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache, il: usize)
2400 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2401 let (attn_g, ag16) = self.full_attn_prime_core_inner(e, fa, g3, pos_d, t, cache, il)?;
2402 if let Some(xh) = &ag16 {
2403 if let Some(y) = e.try_f16_gemm_pre(&fa.wo, xh, t)? {
2404 return Ok(y);
2405 }
2406 }
2407 Ok(e.matmul(&fa.wo, &attn_g, t)?)
2408 }
2409
2410 fn full_attn_prime_core_inner(&self, e: &Engine, fa: &FullAttnLayer, g3: Vec<CudaSlice<f32>>,
2411 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache, il: usize)
2412 -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
2413 let cfg = &self.cfg;
2414 let geometry = cfg.full_attention_geometry_at(il as u32);
2415 let n_head = geometry.n_head as usize;
2416 let n_head_kv = geometry.n_head_kv as usize;
2417 let head_dim = geometry.head_dim_k as usize;
2418 let scale = geometry.attention_scale();
2419 let (pre, base_len) = self.full_attn_prime_pre_fa(e, fa, g3, pos_d, t, cache, il)?;
2420 let AttnPre { q, k, v, gate } = pre;
2421 let mut attn = e.uninit(t * n_head * head_dim)?;
2422 self.full_attn_prime_fa_dispatch(e, &q, &k, &v, &mut attn, base_len, t, cache, il,
2423 head_dim, n_head, n_head_kv, scale)?;
2424 self.full_attn_prime_post_fa(e, attn, &gate, t, n_head, head_dim)
2425 }
2426
2427 #[allow(clippy::type_complexity)]
2431 fn full_attn_prime_pre_fa(&self, e: &Engine, fa: &FullAttnLayer, mut g3: Vec<CudaSlice<f32>>,
2432 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache, il: usize)
2433 -> Result<(AttnPre, usize), Box<dyn std::error::Error>> {
2434 let cfg = &self.cfg;
2435 let geometry = cfg.full_attention_geometry_at(il as u32);
2436 let n_head = geometry.n_head as usize;
2437 let n_head_kv = geometry.n_head_kv as usize;
2438 let head_dim = geometry.head_dim_k as usize;
2439 let eps = cfg.rms_eps;
2440
2441 let gated = geometry.attention_gate
2445 == memra_gguf::config::AttentionGateKind::FusedQ;
2446 let v = g3.pop().unwrap();
2447 let mut k = g3.pop().unwrap();
2448 let qf = g3.pop().unwrap();
2449 let (mut q, gate) = if gated {
2450 let mut q = e.uninit(t * n_head * head_dim)?;
2451 let mut gate = e.uninit(t * n_head * head_dim)?;
2452 e.q_gate_split(&qf, &mut q, &mut gate, head_dim, n_head, t)?;
2453 (q, Some(gate))
2454 } else {
2455 (qf, None)
2456 };
2457
2458 let mut qn = e.uninit(t * n_head * head_dim)?;
2459 e.rms_norm(&q, fa.q_norm.float_data(), &mut qn, head_dim, n_head * t, eps)?;
2460 q = qn;
2461 let mut kn = e.uninit(t * n_head_kv * head_dim)?;
2462 e.rms_norm(&k, fa.k_norm.float_data(), &mut kn, head_dim, n_head_kv * t, eps)?;
2463 k = kn;
2464 let rope_dims = geometry.n_rot as usize;
2465 e.rope_neox(&mut q, pos_d, head_dim, rope_dims, n_head, t, geometry.rope_base, 1.0)?;
2466 e.rope_neox(&mut k, pos_d, head_dim, rope_dims, n_head_kv, t, geometry.rope_base, 1.0)?;
2467
2468 {
2471 let kvl = cache.kv[il].as_mut().unwrap();
2472 assert!(kvl.len + t <= cache.max_ctx, "prime_cache: KV overflow");
2473 e.append_kv_quantized_rows(&k, &v, &mut kvl.k, &mut kvl.v, kvl.len, t,
2474 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes,
2475 crate::Engine::kv_fp8_on())?;
2476 kvl.len += t;
2477 let new_len = kvl.len as i32;
2478 e.set_i32_one(&mut kvl.len_d, new_len)?;
2479 }
2480
2481 let base_len = {
2482 let kvl = cache.kv[il].as_ref().unwrap();
2483 kvl.len - t };
2485 Ok((AttnPre { q, k, v, gate }, base_len))
2486 }
2487
2488 #[allow(clippy::too_many_arguments)]
2495 fn full_attn_prime_fa_dispatch(&self, e: &Engine, q: &CudaSlice<f32>, k: &CudaSlice<f32>,
2496 v: &CudaSlice<f32>, attn: &mut CudaSlice<f32>, base_len: usize,
2497 t: usize, cache: &mut Cache, il: usize,
2498 head_dim: usize, n_head: usize, n_head_kv: usize, scale: f32)
2499 -> Result<(), Box<dyn std::error::Error>> {
2500 if base_len == 0 && std::env::var("MEMRA_PRIME_F32CHUNK0").as_deref() == Ok("1") {
2513 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
2514 e.sdpa_naive(q, k, v, attn, head_dim, n_head, n_head_kv, t, t, scale, true)?;
2515 } else {
2516 e.fa_prefill(q, k, v, attn, head_dim, n_head, n_head_kv, t, t, scale, true)?;
2517 }
2518 return Ok(());
2519 }
2520 let kvl = cache.kv[il].as_ref().unwrap();
2521 let t_kv = base_len + t;
2522 let k_view = e.view_u8(&kvl.k, t_kv * kvl.k_tok_bytes);
2523 let v_view = e.view_u8(&kvl.v, t_kv * kvl.v_tok_bytes);
2524 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
2528 e.sdpa_naive_quantized_view(q, &k_view, &v_view, attn, head_dim, n_head,
2529 n_head_kv, t, t_kv, scale, true,
2530 kvl.k_tok_bytes, kvl.v_tok_bytes)?;
2531 return Ok(());
2532 }
2533 let deqw = std::env::var("MEMRA_PRIME_DEQW").map(|v| v != "0").unwrap_or(true);
2541 if deqw {
2542 e.fa_prefill_view_ws(q, &k_view, &v_view, attn, head_dim, n_head, n_head_kv,
2543 t, t_kv, scale, true, kvl.k_tok_bytes, kvl.v_tok_bytes,
2544 crate::Engine::kv_fp8_on())?;
2545 } else {
2546 e.fa_prefill_view(q, &k_view, &v_view, attn, head_dim, n_head, n_head_kv,
2547 t, t_kv, scale, true, kvl.k_tok_bytes, kvl.v_tok_bytes,
2548 crate::Engine::kv_fp8_on())?;
2549 }
2550 Ok(())
2551 }
2552
2553 fn full_attn_prime_post_fa(&self, e: &Engine, attn: CudaSlice<f32>,
2556 gate: &Option<CudaSlice<f32>>, t: usize,
2557 n_head: usize, head_dim: usize)
2558 -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
2559 let (attn_g, ag16) = match gate {
2560 Some(gate) => {
2561 let n = t * n_head * head_dim;
2562 let mut ag = e.uninit(n)?;
2563 if Self::f16out_on(e, t) {
2564 let mut a16 = e.alloc_u8_uninit(n * 2)?;
2565 e.sig_mul_f16out(&attn, gate, &mut ag, &mut a16, n)?;
2566 (ag, Some(a16))
2567 } else {
2568 let mut gsig = e.uninit(n)?;
2569 e.sigmoid(gate, &mut gsig, n)?;
2570 e.mul(&attn, &gsig, &mut ag, n)?;
2571 (ag, None)
2572 }
2573 }
2574 None => (attn, None),
2575 };
2576 Ok((attn_g, ag16))
2577 }
2578
2579 fn linear_attn_prime(&self, e: &Engine, la: &LinearAttnLayer, h: &CudaSlice<f32>,
2586 hx: Option<&CudaSlice<u8>>, t: usize,
2587 cache: &mut Cache, il: usize)
2588 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2589 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
2591 let g4 = match hx {
2592 Some(xh) => e.matmul_group_xh(&ws, h, xh, t)?,
2593 None => e.matmul_group(&ws, h, t)?,
2594 };
2595 self.linear_attn_prime_core(e, la, g4, t, cache, il)
2596 }
2597
2598 fn linear_attn_prime_core(&self, e: &Engine, la: &LinearAttnLayer, mut g4: Vec<CudaSlice<f32>>,
2600 t: usize, cache: &mut Cache, il: usize)
2601 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2602 self.linear_attn_prime_core_pad(e, la, g4.drain(..).collect(), t, cache, il, None)
2603 }
2604
2605 #[allow(clippy::too_many_arguments)]
2609 fn linear_attn_prime_core_pad_inner(&self, e: &Engine, la: &LinearAttnLayer, mut g4: Vec<CudaSlice<f32>>,
2610 t: usize, cache: &mut Cache, il: usize,
2611 pad_len: Option<&CudaSlice<i32>>)
2612 -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
2613 let ssm = self.cfg.ssm.as_ref().unwrap();
2615 let d_state = ssm.state_size as usize;
2616 let num_k = ssm.group_count as usize;
2617 let num_v = ssm.time_step_rank as usize;
2618 let key_dim = d_state * num_k;
2619 let value_dim = d_state * num_v;
2620 let conv_dim = key_dim * 2 + value_dim;
2621 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(
2626 e, la,
2627 &qkv_mixed.slice(0..t * conv_dim), &z.slice(0..t * value_dim),
2628 &beta_raw.slice(0..t * num_v), &alpha.slice(0..t * num_v),
2629 t, cache, il, pad_len)
2630 }
2631
2632 #[allow(clippy::too_many_arguments)]
2635 fn linear_attn_gdn_prep(&self, e: &Engine, la: &LinearAttnLayer,
2636 qkv_mixed: &cudarc::driver::CudaView<f32>,
2637 beta_raw: &cudarc::driver::CudaView<f32>,
2638 alpha: &cudarc::driver::CudaView<f32>,
2639 t: usize, cache: &mut Cache, il: usize,
2640 pad_len: Option<&CudaSlice<i32>>)
2641 -> Result<GdnPrep, Box<dyn std::error::Error>> {
2642 let cfg = &self.cfg;
2643 let ssm = cfg.ssm.as_ref().unwrap();
2644 let d_state = ssm.state_size as usize; let num_k = ssm.group_count as usize; let num_v = ssm.time_step_rank as usize; let d_conv = ssm.conv_kernel as usize; let key_dim = d_state * num_k; let value_dim = d_state * num_v; let conv_dim = key_dim * 2 + value_dim; let eps = cfg.rms_eps;
2652 debug_assert!(t >= d_conv - 1, "stateful conv needs T >= pad (PRIME_MIN_T gates)");
2653
2654 let rl = cache.recur[il].as_mut().unwrap();
2659 let hk = Self::gdn_hk(e, t, num_v, num_k);
2660 let conv_fuse = std::env::var("MEMRA_CONV_FUSE").as_deref() != Ok("0");
2661 let hk = if conv_fuse { hk } else { num_v }; let mut q_g = e.uninit(d_state * hk * t)?;
2663 let mut k_g = e.uninit(d_state * hk * t)?;
2664 let mut v_g = e.uninit(d_state * num_v * t)?;
2665 if conv_fuse {
2666 e.ssm_conv1d_gdn_state_pad(qkv_mixed, &mut rl.conv_state, la.ssm_conv1d.float_data(),
2667 &mut q_g, &mut k_g, &mut v_g,
2668 conv_dim, t, d_conv, d_state, num_v, num_k, key_dim, hk, pad_len)?;
2669 } else {
2670 let mut conv_out = e.uninit(conv_dim * t)?; e.ssm_conv1d_tm_state_pad_v(qkv_mixed, &mut rl.conv_state, la.ssm_conv1d.float_data(),
2672 &mut conv_out, conv_dim, t, d_conv, pad_len)?;
2673 e.qkv_to_gdn_repack(&conv_out, &mut q_g, &mut k_g, &mut v_g, d_state, num_v, num_k, key_dim, t)?;
2674 }
2675 let mut q_l2 = e.uninit(d_state * hk * t)?;
2676 let qb16 = if Engine::l2_v2_on(d_state) && e.gdn_wgmma_on(32) {
2680 let mut qb = e.alloc_u8_uninit(d_state * hk * t * 2)?;
2681 e.l2_norm_pp(&q_g, &mut q_l2, Some(&mut qb), d_state, hk * t, eps)?;
2682 Some(qb)
2683 } else {
2684 e.l2_norm_pp(&q_g, &mut q_l2, None, d_state, hk * t, eps)?;
2685 None
2686 };
2687 let mut k_l2 = e.uninit(d_state * hk * t)?;
2688 let kb16 = if Engine::l2_v2_on(d_state) {
2690 let mut kb = e.alloc_u8_uninit(d_state * hk * t * 2)?;
2691 e.l2_norm_pp(&k_g, &mut k_l2, Some(&mut kb), d_state, hk * t, eps)?;
2692 Some(kb)
2693 } else {
2694 e.l2_norm_pp(&k_g, &mut k_l2, None, d_state, hk * t, eps)?;
2695 None
2696 };
2697 let mut beta = e.uninit(t * num_v)?;
2698 e.sigmoid_v(beta_raw, &mut beta, t * num_v)?;
2699 let mut g_log = e.uninit(t * num_v)?;
2700 e.gdn_glog_v(alpha, la.ssm_dt.float_data(), la.ssm_a.float_data(), &mut g_log, num_v, t)?;
2701 if let Some(len_d) = pad_len {
2702 e.gdn_pad_mask(&mut beta, &mut g_log, len_d, num_v, t)?;
2703 }
2704 Ok(GdnPrep { hk, q_l2, k_l2, v_g, beta, g_log, kb16, qb16 })
2705 }
2706
2707 #[allow(clippy::too_many_arguments)]
2712 fn linear_attn_prime_core_batch(&self, e: &Engine, la: &LinearAttnLayer,
2713 g4: &[CudaSlice<f32>], offs: &[usize], ts: &[usize],
2714 caches: &mut [&mut Cache], il: usize)
2715 -> Result<Vec<(CudaSlice<f32>, Option<CudaSlice<u8>>)>, Box<dyn std::error::Error>> {
2716 let ssm = self.cfg.ssm.as_ref().unwrap();
2717 let d_state = ssm.state_size as usize;
2718 let num_k = ssm.group_count as usize;
2719 let num_v = ssm.time_step_rank as usize;
2720 let key_dim = d_state * num_k;
2721 let value_dim = d_state * num_v;
2722 let conv_dim = key_dim * 2 + value_dim;
2723 let eps = self.cfg.rms_eps;
2724 let scale = 1.0 / (d_state as f32).sqrt();
2725 let b = ts.len();
2726 let c = Engine::gdn_chunk_size();
2727 let carried = caches.iter().any(|c| c.pos > 0);
2730 let use_vl = !carried
2731 && (2..=8).contains(&b)
2732 && Engine::gdn_chunked_enabled() && ts.iter().all(|&t| t >= 16)
2733 && e.gdn_mma_enabled(c)
2734 && std::env::var("MEMRA_GDN_VL").as_deref() != Ok("0");
2735 if !use_vl {
2736 return (0..b).map(|s| {
2737 let (o, t) = (offs[s], ts[s]);
2738 self.linear_attn_prime_core_pad_view(
2739 e, la,
2740 &g4[0].slice(o * conv_dim..(o + t) * conv_dim),
2741 &g4[1].slice(o * value_dim..(o + t) * value_dim),
2742 &g4[2].slice(o * num_v..(o + t) * num_v),
2743 &g4[3].slice(o * num_v..(o + t) * num_v),
2744 t, caches[s], il, None)
2745 }).collect();
2746 }
2747 struct SeqBufs {
2751 conv_out: CudaSlice<f32>, q_g: CudaSlice<f32>, k_g: CudaSlice<f32>, v_g: CudaSlice<f32>,
2752 q_l2: CudaSlice<f32>, k_l2: CudaSlice<f32>, beta: CudaSlice<f32>, g_log: CudaSlice<f32>,
2753 gn: CudaSlice<f32>, gn16: CudaSlice<u8>,
2754 }
2755 let d_conv = ssm.conv_kernel as usize;
2756 let f16o = Self::f16out_on(e, 16);
2757 let hk = Self::gdn_hk(e, 16, num_v, num_k); let mut sb = Vec::with_capacity(b);
2759 let mut pres = Vec::with_capacity(b);
2760 for &t in ts.iter().take(b) {
2761 sb.push(SeqBufs {
2762 conv_out: e.uninit(conv_dim * t)?,
2763 q_g: e.uninit(d_state * hk * t)?,
2764 k_g: e.uninit(d_state * hk * t)?,
2765 v_g: e.uninit(d_state * num_v * t)?,
2766 q_l2: e.uninit(d_state * hk * t)?,
2767 k_l2: e.uninit(d_state * hk * t)?,
2768 beta: e.uninit(t * num_v)?,
2769 g_log: e.uninit(t * num_v)?,
2770 gn: e.uninit(d_state * num_v * t)?,
2771 gn16: e.alloc_u8_uninit(d_state * num_v * t * 2)?,
2772 });
2773 pres.push(e.gdn_chunk_alloc(num_v, t, c, hk)?);
2774 }
2775 let prep_args: Vec<crate::GdnPrepVl> = (0..b).map(|s| {
2776 let (o, t) = (offs[s], ts[s]);
2777 let rl = caches[s].recur[il].as_ref().unwrap();
2778 crate::GdnPrepVl {
2779 qkv: e.addr_f32v(&g4[0].slice(o * conv_dim..(o + t) * conv_dim)),
2780 conv_state: e.addr_f32(&rl.conv_state),
2781 conv_out: e.addr_f32(&sb[s].conv_out),
2782 q_g: e.addr_f32(&sb[s].q_g), k_g: e.addr_f32(&sb[s].k_g), v_g: e.addr_f32(&sb[s].v_g),
2783 q_l2: e.addr_f32(&sb[s].q_l2), k_l2: e.addr_f32(&sb[s].k_l2),
2784 beta_raw: e.addr_f32v(&g4[2].slice(o * num_v..(o + t) * num_v)),
2785 alpha: e.addr_f32v(&g4[3].slice(o * num_v..(o + t) * num_v)),
2786 beta: e.addr_f32(&sb[s].beta), g_log: e.addr_f32(&sb[s].g_log),
2787 o: e.addr_f32(&pres[s].o),
2788 z: e.addr_f32v(&g4[1].slice(o * value_dim..(o + t) * value_dim)),
2789 gn: e.addr_f32(&sb[s].gn), gn16: e.addr_u8(&sb[s].gn16),
2790 kb16: if Engine::l2_v2_on(d_state) { e.addr_u8(&pres[s].kb16) } else { 0 },
2791 qb16: if Engine::l2_v2_on(d_state) && e.gdn_wgmma_on(c) { e.addr_u8(&pres[s].qb16) } else { 0 },
2792 t: t as i32, pad: 0,
2793 }
2794 }).collect();
2795 let args: Vec<crate::GdnSeqVl> = (0..b).map(|s| {
2796 let rl = caches[s].recur[il].as_ref().unwrap();
2797 crate::GdnSeqVl {
2798 kb16: e.addr_u8(&pres[s].kb16), gcum: e.addr_f32(&pres[s].gcum),
2799 beta: e.addr_f32(&sb[s].beta), u: e.addr_f32(&pres[s].u),
2800 wb16: e.addr_u8(&pres[s].wb16), y: e.addr_u8(&pres[s].y16),
2801 ssnap: e.addr_u8(&pres[s].ssnap16),
2802 state_in: e.addr_f32(&rl.ssm_state), state_out: e.addr_f32(&rl.ssm_state_alt),
2803 q: e.addr_f32(&sb[s].q_l2), p: e.addr_f32(&pres[s].p),
2804 o: e.addr_f32(&pres[s].o),
2805 k: e.addr_f32(&sb[s].k_l2), v: e.addr_f32(&sb[s].v_g),
2806 g: e.addr_f32(&sb[s].g_log), a: e.addr_f32(&pres[s].a),
2807 w: e.addr_f32(&pres[s].w),
2808 t: ts[s] as i32, nc: pres[s].nc as i32,
2809 }
2810 }).collect();
2811 e.gdn_prep_vl8(&prep_args, la.ssm_conv1d.float_data(), la.ssm_dt.float_data(),
2812 la.ssm_a.float_data(), conv_dim, d_conv, d_state, num_v, num_k, key_dim, hk, eps)?;
2813 if !Engine::l2_v2_on(d_state) {
2816 e.gdn_mirror_vl8(&args, num_v, 0, hk)?;
2817 }
2818 let wq8: Option<crate::GdnWVl8> = if e.gdn_wgmma_on(c) {
2820 if !Engine::l2_v2_on(d_state) {
2822 for s in 0..b {
2823 e.f32_to_bf16_into(&sb[s].q_l2, &mut pres[s].qb16, d_state * hk * ts[s])?;
2824 }
2825 }
2826 let mut wa = [crate::GdnWVl::default(); 8];
2827 for s in 0..b {
2828 wa[s] = crate::GdnWVl { qb16: e.addr_u8(&pres[s].qb16), pb16: e.addr_u8(&pres[s].pb16) };
2829 }
2830 Some(crate::GdnWVl8(wa))
2831 } else { None };
2832 e.gdn_chunk_k123_vl8(&args, num_v, hk, wq8.as_ref())?;
2833 e.gdn_chunk_vl8(&args, num_v, scale, hk, wq8.as_ref())?;
2834 if f16o {
2835 e.gdn_tail_vl8(&prep_args, la.ssm_norm.float_data(), d_state, num_v, eps)?;
2836 }
2837 let mut out = Vec::with_capacity(b);
2839 for (s, bufs) in sb.into_iter().enumerate() {
2840 let rl = caches[s].recur[il].as_mut().unwrap();
2841 std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
2842 let (o, t) = (offs[s], ts[s]);
2843 let SeqBufs { mut gn, gn16, .. } = bufs;
2844 if f16o {
2845 out.push((gn, Some(gn16)));
2846 } else {
2847 let z_v = g4[1].slice(o * value_dim..(o + t) * value_dim);
2848 e.gated_rmsnorm_zv(&pres[s].o, la.ssm_norm.float_data(), &z_v, &mut gn,
2849 d_state, num_v * t, eps)?;
2850 out.push((gn, None));
2851 }
2852 }
2853 Ok(out)
2854 }
2855
2856 #[allow(clippy::too_many_arguments)]
2860 fn linear_attn_prime_core_pad_view(&self, e: &Engine, la: &LinearAttnLayer,
2861 qkv_mixed: &cudarc::driver::CudaView<f32>,
2862 z: &cudarc::driver::CudaView<f32>,
2863 beta_raw: &cudarc::driver::CudaView<f32>,
2864 alpha: &cudarc::driver::CudaView<f32>,
2865 t: usize, cache: &mut Cache, il: usize,
2866 pad_len: Option<&CudaSlice<i32>>)
2867 -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
2868 let cfg = &self.cfg;
2869 let ssm = cfg.ssm.as_ref().unwrap();
2870 let d_state = ssm.state_size as usize; let num_v = ssm.time_step_rank as usize; let eps = cfg.rms_eps;
2873 let scale = 1.0 / (d_state as f32).sqrt();
2874
2875 let prep = self.linear_attn_gdn_prep(e, la, qkv_mixed, beta_raw, alpha, t, cache, il, pad_len)?;
2876
2877 let mut o = e.uninit(d_state * num_v * t)?;
2883 let rl = cache.recur[il].as_mut().unwrap();
2884 {
2885 let crate::cache::RecurLayer { ssm_state, ssm_state_alt, .. } = rl;
2886 e.gdn_scan_prefill(&prep.q_l2, &prep.k_l2, &prep.v_g, &prep.g_log, &prep.beta,
2887 prep.kb16.as_ref(), prep.qb16.as_ref(), ssm_state, ssm_state_alt, &mut o, num_v, t, scale,
2888 prep.hk)?;
2889 }
2890 std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
2891
2892 let mut gn = e.uninit(d_state * num_v * t)?;
2895 let gn16 = if Self::f16out_on(e, t) {
2896 let mut g16 = e.alloc_u8_uninit(d_state * num_v * t * 2)?;
2897 e.gated_rmsnorm_f16out_zv(&o, la.ssm_norm.float_data(), z, &mut gn, &mut g16,
2898 d_state, num_v * t, eps)?;
2899 Some(g16)
2900 } else {
2901 e.gated_rmsnorm_zv(&o, la.ssm_norm.float_data(), z, &mut gn, d_state, num_v * t, eps)?;
2902 None
2903 };
2904 Ok((gn, gn16))
2905 }
2906
2907 #[allow(clippy::too_many_arguments)]
2909 fn linear_attn_prime_core_pad(&self, e: &Engine, la: &LinearAttnLayer, g4: Vec<CudaSlice<f32>>,
2910 t: usize, cache: &mut Cache, il: usize,
2911 pad_len: Option<&CudaSlice<i32>>)
2912 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2913 let (gn, gn16) = self.linear_attn_prime_core_pad_inner(e, la, g4, t, cache, il, pad_len)?;
2914 if let Some(xh) = &gn16 {
2915 if let Some(y) = e.try_f16_gemm_pre(&la.ssm_out, xh, t)? {
2916 return Ok(y);
2917 }
2918 }
2919 Ok(e.matmul(&la.ssm_out, &gn, t)?)
2920 }
2921
2922 pub fn full_attn(&self, e: &Engine, fa: &FullAttnLayer, h: &CudaSlice<f32>, pos_d: &CudaSlice<i32>, t: usize, il: usize)
2927 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2928 if self.cfg.step35.is_some() {
2929 return self.step35_attn(e, fa, h, pos_d, t, il);
2930 }
2931 let cfg = &self.cfg;
2932 let _n_embd = cfg.n_embd as usize;
2933 let geometry = cfg.full_attention_geometry_at(il as u32);
2934 let n_head = geometry.n_head as usize;
2935 let n_head_kv = geometry.n_head_kv as usize;
2936 let head_dim = geometry.head_dim_k as usize;
2937 let eps = cfg.rms_eps;
2938 let scale = geometry.attention_scale();
2939
2940 let gated = geometry.attention_gate
2943 == memra_gguf::config::AttentionGateKind::FusedQ;
2944 let mut g3 = e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?;
2946 let v = g3.pop().unwrap();
2947 let mut k = g3.pop().unwrap();
2948 let qf = g3.pop().unwrap();
2949 let (mut q, gate) = if gated {
2950 let mut q = e.uninit(t * n_head * head_dim)?;
2951 let mut gate = e.uninit(t * n_head * head_dim)?;
2952 e.q_gate_split(&qf, &mut q, &mut gate, head_dim, n_head, t)?;
2953 (q, Some(gate))
2954 } else {
2955 (qf, None)
2956 };
2957
2958 let mut qn = e.uninit(t * n_head * head_dim)?;
2960 e.rms_norm(&q, fa.q_norm.float_data(), &mut qn, head_dim, n_head * t, eps)?;
2961 q = qn;
2962 let mut kn = e.uninit(t * n_head_kv * head_dim)?;
2963 e.rms_norm(&k, fa.k_norm.float_data(), &mut kn, head_dim, n_head_kv * t, eps)?;
2964 k = kn;
2965 let rope_dims = geometry.n_rot as usize;
2966 e.rope_neox(&mut q, pos_d, head_dim, rope_dims, n_head, t, geometry.rope_base, 1.0)?;
2967 e.rope_neox(&mut k, pos_d, head_dim, rope_dims, n_head_kv, t, geometry.rope_base, 1.0)?;
2968
2969 let mut attn = e.uninit(t * n_head * head_dim)?;
2971 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
2974 e.sdpa_naive(&q, &k, &v, &mut attn, head_dim, n_head, n_head_kv, t, t, scale, true)?;
2976 } else {
2977 e.fa_prefill(&q, &k, &v, &mut attn, head_dim, n_head, n_head_kv, t, t, scale, true)?;
2978 }
2979
2980 let attn_g = match &gate {
2982 Some(gate) => {
2983 let mut gsig = e.uninit(t * n_head * head_dim)?;
2984 e.sigmoid(gate, &mut gsig, t * n_head * head_dim)?;
2985 let mut ag = e.uninit(t * n_head * head_dim)?;
2986 e.mul(&attn, &gsig, &mut ag, t * n_head * head_dim)?;
2987 ag
2988 }
2989 None => attn,
2990 };
2991
2992 let o = e.matmul(&fa.wo, &attn_g, t)?;
2994 Ok(o)
2995 }
2996
2997 pub fn linear_attn(&self, e: &Engine, la: &LinearAttnLayer, h: &CudaSlice<f32>, t: usize)
2999 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3000 let cfg = &self.cfg;
3001 let _n_embd = cfg.n_embd as usize;
3002 let ssm = cfg.ssm.as_ref().unwrap();
3003 let d_state = ssm.state_size as usize; let num_k = ssm.group_count as usize; let num_v = ssm.time_step_rank as usize; let d_conv = ssm.conv_kernel as usize; let head_k = d_state; let head_v = d_state;
3008 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;
3012 let scale = 1.0 / (d_state as f32).sqrt();
3013
3014 let mut g4 = e.matmul_group(&[&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha], h, t)?;
3017 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);
3029 let mut q_g = e.uninit(d_state * num_v * t)?;
3030 let mut k_g = e.uninit(d_state * num_v * t)?;
3031 let mut v_g = e.uninit(d_state * num_v * t)?;
3032 e.ssm_conv1d_gdn(&qkv_mixed, la.ssm_conv1d.float_data(), &mut q_g, &mut k_g, &mut v_g,
3033 conv_dim, t, d_conv, d_state, num_v, num_k, key_dim)?;
3034 let mut q_l2 = e.uninit(d_state * num_v * t)?;
3036 e.l2_norm(&q_g, &mut q_l2, d_state, num_v * t, eps)?;
3037 let mut k_l2 = e.uninit(d_state * num_v * t)?;
3038 e.l2_norm(&k_g, &mut k_l2, d_state, num_v * t, eps)?;
3039 let v_gd = v_g;
3040
3041 let mut beta = e.uninit(t * num_v)?;
3044 e.sigmoid(&beta_raw, &mut beta, t * num_v)?;
3045 let mut g_log = e.uninit(t * num_v)?;
3047 e.gdn_glog(&alpha, la.ssm_dt.float_data(), la.ssm_a.float_data(), &mut g_log, num_v, t)?;
3048
3049 let state_in = e.zeros(d_state * d_state * num_v)?; let mut state_out = e.zeros(d_state * d_state * num_v)?;
3052 let mut o = e.uninit(d_state * num_v * t)?;
3053 e.gdn_scan_prefill(&q_l2, &k_l2, &v_gd, &g_log, &beta, None, None, &state_in, &mut state_out, &mut o, num_v, t, scale, num_v)?;
3054
3055 let mut gn = e.uninit(d_state * num_v * t)?;
3060 e.gated_rmsnorm(&o, la.ssm_norm.float_data(), &z, &mut gn, d_state, num_v * t, eps)?;
3061
3062 let out = e.matmul(&la.ssm_out, &gn, t)?;
3066 Ok(out)
3067 }
3068}
3069
3070impl HybridModel {
3071 pub fn moe_ffn_il(&self, e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, t: usize, il: u16)
3082 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3083 Self::moe_ffn_inner(e, m, z, None, t, &self.cfg, il, self.max_moe_block(), false)
3084 }
3085
3086 pub fn moe_ffn_il_prefill(
3089 &self,
3090 e: &Engine,
3091 m: &MoeWeights,
3092 z: &CudaSlice<f32>,
3093 t: usize,
3094 il: u16,
3095 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3096 Self::moe_ffn_inner(e, m, z, None, t, &self.cfg, il, self.max_moe_block(), true)
3097 }
3098
3099 pub fn moe_ffn_il_zq8(&self, e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>,
3103 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>, t: usize, il: u16)
3104 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3105 Self::moe_ffn_inner(
3106 e, m, z, zq8, t, &self.cfg, il, self.max_moe_block(), false,
3107 )
3108 }
3109
3110 pub(crate) fn moe_ffn(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, t: usize,
3118 cfg: &ModelConfig, il: u16, max_block: usize)
3119 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3120 Self::moe_ffn_inner(e, m, z, None, t, cfg, il, max_block, false)
3121 }
3122
3123 #[allow(clippy::too_many_arguments)]
3124 pub(crate) fn moe_ffn_inner(
3125 e: &Engine,
3126 m: &MoeWeights,
3127 z: &CudaSlice<f32>,
3128 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
3129 t: usize,
3130 cfg: &ModelConfig,
3131 il: u16,
3132 max_block: usize,
3133 prefill: bool,
3134 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3135 let worker_io = crate::spill_pread::worker_enabled();
3136 let epoch_lfu = std::env::var_os("MEMRA_MOE_LFU_DECAY").is_some();
3137 if Engine::moe_cache_enabled() && (worker_io || epoch_lfu) {
3138 e.with_moe_cache(max_block, |cache, _| {
3139 cache.begin_forward_epoch(il, t);
3140 if worker_io {
3141 cache.begin_worker_scope();
3142 }
3143 Ok(())
3144 })?;
3145 }
3146 if t > 1 && moe_grouped_enabled(cfg, prefill) {
3149 let grouped_out = Self::moe_ffn_grouped(e, m, z, t, cfg, il, max_block)?;
3150 if std::env::var("MEMRA_MOE_GATE").is_ok() {
3155 let seq_out = Self::moe_ffn_sequential(e, m, z, t, cfg, il, max_block)?;
3156 let g_host = e.dtoh(&grouped_out)?;
3157 let s_host = e.dtoh(&seq_out)?;
3158 let g_bytes: &[u8] = unsafe { std::slice::from_raw_parts(g_host.as_ptr() as *const u8, g_host.len() * 4) };
3159 let s_bytes: &[u8] = unsafe { std::slice::from_raw_parts(s_host.as_ptr() as *const u8, s_host.len() * 4) };
3160 if g_bytes == s_bytes {
3161 println!("moe-gate il={il} t={t} BYTE-IDENTICAL");
3162 } else {
3163 let diffs = g_host.iter().zip(s_host.iter()).enumerate()
3164 .filter(|(_, (a, b))| a != b).count();
3165 let maxdiff = g_host.iter().zip(s_host.iter())
3166 .map(|(a, b)| (a - b).abs()).fold(0.0f32, f32::max);
3167 panic!("moe-gate il={il} t={t} MISMATCH: {diffs}/{} elems differ, maxdiff={maxdiff:.6e}", g_host.len());
3168 }
3169 }
3170 return Ok(grouped_out);
3171 }
3172 Self::moe_ffn_sequential_zq8(e, m, z, zq8, t, cfg, il, max_block)
3173 }
3174
3175 pub(crate) fn moe_ffn_sequential(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, t: usize,
3177 cfg: &ModelConfig, il: u16, max_block: usize)
3178 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3179 Self::moe_ffn_sequential_zq8(e, m, z, None, t, cfg, il, max_block)
3180 }
3181
3182 fn moe_router_logits(
3186 e: &Engine,
3187 m: &MoeWeights,
3188 z: &CudaSlice<f32>,
3189 t: usize,
3190 cfg: &ModelConfig,
3191 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3192 if t < PRIME_MIN_T {
3193 if crate::router_kernel_on() {
3195 e.router_gemv(
3196 m.gate_inp.float_data(),
3197 z,
3198 cfg.n_embd as usize,
3199 m.gate_exps.n_expert,
3200 t,
3201 )
3202 } else {
3203 e.matmul_decode_exact(&m.gate_inp, z, t)
3204 }
3205 } else if crate::router_prefill_exact_on() && crate::router_kernel_on() {
3206 e.router_gemv(
3207 m.gate_inp.float_data(),
3208 z,
3209 cfg.n_embd as usize,
3210 m.gate_exps.n_expert,
3211 t,
3212 )
3213 } else {
3214 e.matmul(&m.gate_inp, z, t)
3215 }
3216 }
3217
3218 fn trace_moe_routes(il: u16, t: usize, sel_all: &[u32], weights: &[f32])
3222 -> Result<(), Box<dyn std::error::Error>> {
3223 use std::io::Write as _;
3224 if let Ok(path) = std::env::var("MEMRA_MOE_TRACE") {
3225 let mut f = std::fs::OpenOptions::new().create(true).append(true).open(path)?;
3226 let ids: Vec<String> = sel_all.iter().map(|s| s.to_string()).collect();
3227 writeln!(f, "{} {} {}", il, t, ids.join(","))?;
3228 }
3229 if let Ok(path) = std::env::var("MEMRA_MOE_WEIGHT_TRACE") {
3230 let mut f = std::fs::OpenOptions::new().create(true).append(true).open(path)?;
3231 let pairs: Vec<String> = sel_all.iter().zip(weights)
3232 .map(|(expert, weight)| format!("{expert}:{weight:.9}"))
3233 .collect();
3234 writeln!(f, "{} {} {}", il, t, pairs.join(","))?;
3235 }
3236 Ok(())
3237 }
3238
3239 fn trace_moe_input(e: &Engine, il: u16, t: usize, n_embd: usize, z: &CudaSlice<f32>)
3244 -> Result<(), Box<dyn std::error::Error>> {
3245 use std::io::Write as _;
3246 let Ok(dir) = std::env::var("MEMRA_MOE_INPUT_TRACE_DIR") else { return Ok(()) };
3247 let host = e.dtoh(z)?;
3248 if host.len() != t * n_embd {
3249 return Err(format!(
3250 "MoE input trace shape mismatch at layer {il}: got {} values, expected {}x{}",
3251 host.len(), t, n_embd
3252 ).into());
3253 }
3254 let bytes = unsafe {
3255 std::slice::from_raw_parts(
3256 host.as_ptr().cast::<u8>(), host.len() * std::mem::size_of::<f32>()
3257 )
3258 };
3259 let state = MOE_INPUT_TRACE_WRITER.get_or_init(|| std::sync::Mutex::new(None));
3260 let mut state = state.lock().map_err(|_| "MoE input trace writer lock is poisoned")?;
3261 if state.is_none() {
3262 let dir = std::path::PathBuf::from(&dir);
3263 std::fs::create_dir_all(&dir)?;
3264 let index = std::fs::OpenOptions::new().create(true).append(true)
3265 .open(dir.join("index.jsonl"))?;
3266 *state = Some(MoeInputTraceWriter {
3267 dir,
3268 index,
3269 payloads: std::collections::HashMap::new(),
3270 });
3271 }
3272 let writer = state.as_mut().unwrap();
3273 if writer.dir != std::path::Path::new(&dir) {
3274 return Err("MEMRA_MOE_INPUT_TRACE_DIR changed after capture started".into());
3275 }
3276 let file_name = format!("layer-{il:03}.f32");
3277 if !writer.payloads.contains_key(&il) {
3278 let payload = std::fs::OpenOptions::new().create(true).append(true)
3279 .open(writer.dir.join(&file_name))?;
3280 let offset = payload.metadata()?.len();
3281 writer.payloads.insert(il, (payload, offset));
3282 }
3283 let (payload, offset) = writer.payloads.get_mut(&il).unwrap();
3284 let row_offset = *offset;
3285 payload.write_all(bytes)?;
3286 *offset += bytes.len() as u64;
3287 writeln!(
3288 writer.index,
3289 "{{\"format\":\"memra-moe-input-trace-v1\",\"layer\":{il},\"tokens\":{t},\
3290 \"hidden_size\":{n_embd},\"file\":\"{file_name}\",\"offset\":{row_offset},\
3291 \"payload_bytes\":{}}}",
3292 bytes.len()
3293 )?;
3294 Ok(())
3295 }
3296
3297 #[allow(clippy::too_many_arguments)]
3298 pub(crate) fn moe_ffn_sequential_zq8(
3299 e: &Engine,
3300 m: &MoeWeights,
3301 z: &CudaSlice<f32>,
3302 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
3303 t: usize,
3304 cfg: &ModelConfig,
3305 il: u16,
3306 max_block: usize,
3307 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3308 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
3309 let moe = cfg.moe.as_ref().unwrap();
3310 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);
3317 debug_assert_eq!(m.gate_exps.out_f, n_ff_exp);
3318 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);
3321
3322 let lim_exp = cfg.clamp_exp_at(il as u32);
3325 let lim_shexp = cfg.clamp_shexp_at(il as u32);
3326 let use_cache = Engine::moe_cache_enabled();
3327 let uniform_experts = m.has_uniform_expert_layout();
3328 let moe_q8 = uniform_experts && moe_q8_enabled()
3329 && q8_expert_supported(m.gate_exps.qtype) && q8_expert_supported(m.up_exps.qtype)
3330 && q8_expert_supported(m.down_exps.qtype);
3331 let cpu_expert_requested = crate::cpu_experts::configured();
3338 if cpu_expert_requested && (cfg.hy3.is_none() || cfg.m3.is_some()) {
3339 return Err(std::io::Error::other(
3340 "MEMRA_CPU_EXPERT_LIB is experimental and currently gated to Hy3",
3341 )
3342 .into());
3343 }
3344 let cpu_hybrid = cpu_expert_requested && t < PRIME_MIN_T && m.dev_exps.is_none();
3345 let freeze_cpu_residency = cpu_expert_requested
3351 && std::env::var("MEMRA_CPU_EXPERT_FREEZE_CACHE").as_deref() == Ok("1");
3352 let caller_warms_before_freeze = std::env::var("MEMRA_CPU_EXPERT_FREEZE_WARMUP_TOKENS")
3353 .ok()
3354 .and_then(|value| value.parse::<usize>().ok())
3355 .is_some_and(|tokens| tokens > 0);
3356 if cpu_hybrid && freeze_cpu_residency && !caller_warms_before_freeze {
3357 e.freeze_moe_cache();
3358 }
3359 let cache_frozen = use_cache && e.moe_cache_frozen();
3360 let cache_dispatch = use_cache && (!cache_frozen || cpu_hybrid);
3361
3362 let logits = Self::moe_router_logits(e, m, z, t, cfg)?;
3365
3366 let no_exp_macros = m.gate_exps.macros.is_none() && m.up_exps.macros.is_none()
3404 && m.down_exps.macros.is_none();
3405 if cfg.sigmoid_router().is_none() && cfg.m3.is_none() && cfg.hy3.is_none()
3409 && !cfg.swiglu_clamped_at(il as u32)
3410 && no_exp_macros
3411 && t >= PRIME_MIN_T && m.dev_exps.is_some() && moe_q8_enabled()
3412 && q8_expert_supported(m.gate_exps.qtype) && q8_expert_supported(m.up_exps.qtype)
3413 && q8_expert_supported(m.down_exps.qtype)
3414 && std::env::var("MEMRA_MOE_PAIRS").map(|v| v != "0").unwrap_or(true)
3415 && std::env::var("MEMRA_MOE_STATS").is_err() {
3416 return Self::moe_ffn_pairs(e, m, z, &logits, t, cfg);
3417 }
3418
3419 let dev_ok = uniform_experts && cfg.sigmoid_router().is_none()
3436 && cfg.m3.is_none() && cfg.hy3.is_none()
3437 && !cfg.swiglu_clamped_at(il as u32);
3438 let observe_routes = std::env::var("MEMRA_MOE_STATS").is_ok()
3442 || std::env::var("MEMRA_MOE_TRACE").is_ok()
3443 || std::env::var("MEMRA_MOE_WEIGHT_TRACE").is_ok()
3444 || std::env::var("MEMRA_MOE_INPUT_TRACE_DIR").is_ok();
3445 if dev_ok && t < PRIME_MIN_T && m.dev_exps.is_some() && n_used <= 8 && moe_dev_enabled()
3446 && !observe_routes {
3447 return Self::moe_ffn_dev(e, m, z, zq8, &logits, t, cfg, il, max_block);
3448 }
3449 if dev_ok && use_cache && n_used <= 8 && moe_dev_enabled()
3450 && !observe_routes {
3451 let row_ok = e.with_moe_cache(max_block, |c, eng| {
3452 if moe_prewarm_enabled() { c.prewarm_layer(il, m, eng)?; }
3453 Ok(c.layer_dev_row(il, n_expert, eng)?.is_some())
3454 })?;
3455 if row_ok {
3456 return Self::moe_ffn_dev(e, m, z, zq8, &logits, t, cfg, il, max_block);
3457 }
3458 }
3459
3460 let (sel_all, w_all, routed_cpu_input) = if let Some(sig) = cfg.sigmoid_router() {
3462 if cpu_hybrid {
3463 let (sel, w, input) = Self::moe_route_sigmoid_with_input(
3464 e,
3465 &logits,
3466 z,
3467 t,
3468 n_expert,
3469 n_used,
3470 m.exp_probs_b.as_deref(),
3471 sig,
3472 m.active_experts.as_deref(),
3473 )?;
3474 (sel, w, Some(input))
3475 } else {
3476 let (sel, w) = Self::moe_route_cfg(
3477 e,
3478 &logits,
3479 t,
3480 n_expert,
3481 n_used,
3482 m.exp_probs_b.as_deref(),
3483 Some(sig),
3484 m.active_experts.as_deref(),
3485 )?;
3486 (sel, w, None)
3487 }
3488 } else {
3489 let (sel, w) = Self::moe_route_cfg(
3490 e,
3491 &logits,
3492 t,
3493 n_expert,
3494 n_used,
3495 None,
3496 None,
3497 m.active_experts.as_deref(),
3498 )?;
3499 (sel, w, None)
3500 };
3501
3502 Self::trace_moe_routes(il, t, &sel_all, &w_all)?;
3506 Self::trace_moe_input(e, il, t, n_embd, z)?;
3507
3508 let worker_disk_prefetch =
3520 cache_dispatch && crate::spill_pread::worker_enabled() && !cpu_hybrid;
3521 let promote_worker_h2d =
3522 t == 1 && worker_disk_prefetch && crate::spill_pread::copy_h2d_enabled();
3523 if promote_worker_h2d {
3524 let mut selected_blocks = Vec::with_capacity(n_used * 3);
3525 for &ex in sel_all.iter().take(n_used) {
3526 let ex = ex as u16;
3527 selected_blocks.extend([
3528 BlockId::new(il, PROJ_GATE, ex),
3529 BlockId::new(il, PROJ_UP, ex),
3530 BlockId::new(il, PROJ_DOWN, ex),
3531 ]);
3532 }
3533 for &ex in sel_all.iter().take(n_used) {
3534 Self::moe_prefetch_disk_expert(e, il, ex as usize, m, max_block, &selected_blocks)?;
3535 }
3536 e.with_moe_cache(max_block, |cache, eng| {
3537 cache.promote_worker_reads_at_safe_boundary(
3538 &selected_blocks,
3539 &selected_blocks,
3540 eng,
3541 )?;
3542 Ok(())
3543 })?;
3544 }
3545
3546 if t > 1 && std::env::var("MEMRA_MOE_STATS").is_ok() {
3549 let mut cnt = vec![0u32; n_expert];
3550 for &s in sel_all.iter() { cnt[s as usize] += 1; }
3551 let total = sel_all.len() as f64;
3552 let mut h = 0.0f64;
3553 let mut active = 0usize;
3554 for &c in &cnt { if c > 0 { active += 1; let p = c as f64 / total; h -= p * p.log2(); } }
3555 let maxc = cnt.iter().copied().max().unwrap_or(0);
3556 println!("moe-stats il={} t={} assignments={} active={}/{} entropy={:.3}b (max {:.3}b) mean_tok_per_active={:.2} max_tok_per_expert={}",
3557 il, t, sel_all.len(), active, n_expert, h, (n_expert as f64).log2(), total / active.max(1) as f64, maxc);
3558 }
3559
3560 let gdec_may_fire = uniform_experts && use_cache && n_used <= 8 && gdec_enabled()
3573 && !cfg.swiglu_clamped_at(il as u32);
3574 let slab_local = m.dev_exps.as_ref()
3590 .filter(|d| !d.gu_il && moe_slab_enabled() && d.dev == e.ctx().ordinal());
3591 let slab_bases = slab_local.map(|d| {
3592 use cudarc::driver::DevicePtr;
3593 let s = e.stream();
3594 let (pg, _g0) = d.gate.device_ptr(&s);
3595 let (pu, _g1) = d.up.device_ptr(&s);
3596 let (pd, _g2) = d.down.device_ptr(&s);
3597 (pg as u64, pu as u64, pd as u64)
3598 });
3599 let slab_fused_may_fire = slab_bases.is_some() && n_used <= 8 && gdec_enabled()
3609 && !cfg.swiglu_clamped_at(il as u32) && cfg.m3.is_none()
3610 && no_exp_macros && moe_q8;
3611 let mut moe_out = if gdec_may_fire || slab_fused_may_fire {
3614 e.uninit(t * n_embd)?
3615 } else {
3616 e.zeros(t * n_embd)?
3617 };
3618 let cpu_input = if cpu_hybrid {
3621 Some(routed_cpu_input.ok_or("CPU expert routing did not return the MoE input")?)
3622 } else {
3623 None
3624 };
3625
3626 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;
3634 let mut scratch_u: Option<CudaSlice<u8>> = None;
3635 let mut scratch_d: Option<CudaSlice<u8>> = None;
3636 let page_window = moe_page_prefetch_window();
3644
3645 for tok in 0..t {
3648 let sel = &sel_all[tok * n_used..(tok + 1) * n_used];
3649 let w = &w_all[tok * n_used..(tok + 1) * n_used];
3650 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd); let mut tok_q8: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
3652
3653 let no_macros = m.gate_exps.macros.is_none() && m.up_exps.macros.is_none()
3667 && m.down_exps.macros.is_none();
3668 if slab_fused_may_fire {
3678 let (pg, pu, pd) = slab_bases.unwrap();
3679 let mut gp = [0u64; 8];
3680 let mut up = [0u64; 8];
3681 let mut dp = [0u64; 8];
3682 for (j, &ex) in sel.iter().enumerate() {
3683 let ex = ex as usize;
3684 gp[j] = pg + (ex * m.gate_exps.expert_stride) as u64;
3685 up[j] = pu + (ex * m.up_exps.expert_stride) as u64;
3686 dp[j] = pd + (ex * m.down_exps.expert_stride) as u64;
3687 }
3688 let mut wv = [0f32; 8];
3689 wv[..n_used].copy_from_slice(w);
3690 if tok_q8.is_none() {
3691 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
3692 }
3693 let (zq, zd) = tok_q8.as_ref().unwrap();
3694 let act = e.moe_gate_up_silu8_q8(crate::WPtr8(gp), crate::WPtr8(up), zq, zd,
3695 n_embd, n_ff_exp, n_used,
3696 m.gate_exps.qtype, m.up_exps.qtype,
3697 m.gate_exps.row_bytes, m.up_exps.row_bytes)?;
3698 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
3699 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
3700 e.moe_down8_fma_q8(crate::WPtr8(dp), crate::F32x8(wv), &aq2, &ad2, &mut dst,
3701 n_ff_exp, n_embd, n_used,
3702 m.down_exps.qtype, m.down_exps.row_bytes)?;
3703 continue;
3704 }
3705 if gdec_may_fire && moe_q8 && cfg.m3.is_none() && no_macros {
3706 if tok_q8.is_none() {
3707 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
3708 }
3709 let (zq, zd) = tok_q8.as_ref().unwrap();
3710 if Self::moe_gdec_token_q8(e, m, il, max_block, zq, zd, sel, w,
3711 &mut moe_out, tok, n_embd, n_ff_exp, n_used)? {
3712 continue;
3713 }
3714 } else if gdec_may_fire && cfg.m3.is_none() && no_macros
3715 && Self::moe_gdec_token(e, m, il, max_block, &zt, sel, w,
3716 &mut moe_out, tok, n_embd, n_ff_exp, n_used)? {
3717 continue;
3718 }
3719
3720 if gdec_may_fire || slab_fused_may_fire {
3726 let mut row = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
3727 e.memset_zeros_view(&mut row)?;
3728 }
3729
3730 let mut cpu_mask = vec![false; sel.len()];
3736 let cpu_worker = if let Some(host_input) = cpu_input.as_ref() {
3737 let gpu_resident = if use_cache {
3738 e.with_moe_cache(max_block, |cache, _| {
3739 Ok(sel
3740 .iter()
3741 .map(|&expert| {
3742 let expert = expert as u16;
3743 [PROJ_GATE, PROJ_UP, PROJ_DOWN]
3744 .into_iter()
3745 .filter(|&projection| {
3746 cache
3747 .resident(BlockId::new(il, projection, expert))
3748 .is_some()
3749 })
3750 .count()
3751 })
3752 .collect::<Vec<_>>())
3753 })?
3754 } else {
3755 vec![0; sel.len()]
3756 };
3757 let mut cpu_selected = Vec::new();
3758 for (index, (&expert, &route_weight)) in sel.iter().zip(w).enumerate() {
3759 if gpu_resident[index] != 3 {
3760 cpu_mask[index] = true;
3761 crate::cpu_experts::record_incomplete_gpu_residency(gpu_resident[index]);
3762 let expert = expert as usize;
3763 cpu_selected.push((expert, route_weight));
3764 }
3765 }
3766 if crate::cpu_experts::predictor_enabled() {
3767 let row = &host_input[tok * n_embd..(tok + 1) * n_embd];
3771 crate::cpu_experts::predictor_submit(il, row);
3772 }
3773 if cpu_selected.is_empty() {
3774 None
3775 } else {
3776 let row = &host_input[tok * n_embd..(tok + 1) * n_embd];
3777 let job = crate::cpu_experts::prepare_job(m, il, &cpu_selected, row)
3778 .map_err(std::io::Error::other)?;
3779 Some(crate::cpu_experts::submit(job).map_err(std::io::Error::other)?)
3780 }
3781 } else {
3782 None
3783 };
3784
3785 let worker_window = worker_disk_prefetch
3786 .then(worker_prefetch_window)
3787 .unwrap_or(0);
3788 for (j, &ex) in sel.iter().enumerate() {
3789 if cpu_mask[j] {
3790 continue;
3791 }
3792 let ex = ex as usize;
3793 if let Some(d) = slab_local {
3800 let gl = m.gate_exps.expert_layout(ex);
3801 let ul = m.up_exps.expert_layout(ex);
3802 let dl = m.down_exps.expert_layout(ex);
3803 let (g0, u0, d0) = (ex * m.gate_exps.expert_stride,
3804 ex * m.up_exps.expert_stride,
3805 ex * m.down_exps.expert_stride);
3806 let (gate, up) = if moe_q8 {
3807 if tok_q8.is_none() {
3808 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
3809 }
3810 let (zq, zd) = tok_q8.as_ref().unwrap();
3811 (e.qmatvec_expert_q8(&d.gate, g0..g0 + gl.len, zq, zd, 1,
3812 m.gate_exps.in_f, m.gate_exps.out_f,
3813 gl.qtype, gl.row_bytes)?,
3814 e.qmatvec_expert_q8(&d.up, u0..u0 + ul.len, zq, zd, 1,
3815 m.up_exps.in_f, m.up_exps.out_f,
3816 ul.qtype, ul.row_bytes)?)
3817 } else {
3818 (e.qmatvec_view(&d.gate, g0..g0 + gl.len, &zt, 1,
3819 m.gate_exps.in_f, m.gate_exps.out_f,
3820 gl.qtype, gl.row_bytes)?,
3821 e.qmatvec_view(&d.up, u0..u0 + ul.len, &zt, 1,
3822 m.up_exps.in_f, m.up_exps.out_f,
3823 ul.qtype, ul.row_bytes)?)
3824 };
3825 let mut act = e.uninit(n_ff_exp)?;
3826 Self::ffn_act_lim(e, cfg, &gate, &up, m.gate_exps.macro_scale(ex),
3827 m.up_exps.macro_scale(ex), lim_exp, &mut act, n_ff_exp)?;
3828 let y = if moe_q8 {
3829 let (aq2, ad2) = e.quantize_q8_1(&act, 1, n_ff_exp)?;
3830 e.qmatvec_expert_q8(&d.down, d0..d0 + dl.len, &aq2, &ad2, 1,
3831 m.down_exps.in_f, m.down_exps.out_f,
3832 dl.qtype, dl.row_bytes)?
3833 } else {
3834 let actv = act.slice(0..n_ff_exp);
3835 e.qmatvec_view(&d.down, d0..d0 + dl.len, &actv, 1,
3836 m.down_exps.in_f, m.down_exps.out_f,
3837 dl.qtype, dl.row_bytes)?
3838 };
3839 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
3840 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
3841 continue;
3842 }
3843 for next in page_prefetch_positions(j, sel.len(), page_window) {
3844 Self::moe_prefetch_host_expert(sel[next] as usize, m);
3845 }
3846 let keep = [
3847 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_GATE, ex as u16),
3848 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_UP, ex as u16),
3849 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_DOWN, ex as u16),
3850 ];
3851 if worker_disk_prefetch && worker_window > 0 {
3852 for next in worker_prefetch_positions(j, sel.len(), worker_window) {
3853 Self::moe_prefetch_disk_expert(
3854 e,
3855 il,
3856 sel[next] as usize,
3857 m,
3858 max_block,
3859 &keep,
3860 )?;
3861 }
3862 } else if cache_dispatch
3863 && !cpu_hybrid
3864 && moe_prefetch_enabled()
3865 && j + 1 < sel.len()
3866 {
3867 let next = sel[j + 1] as usize;
3868 Self::moe_prefetch_expert(e, il, next, m, max_block, &keep)?;
3869 }
3870 let [gate_q8, up_q8, down_q8] = [moe_q8; 3];
3871 if cache_dispatch && (gate_q8 || up_q8 || down_q8) {
3872 if (gate_q8 || up_q8) && tok_q8.is_none() {
3875 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
3876 }
3877 let gate = if gate_q8 {
3878 let (zq, zd) = tok_q8.as_ref().unwrap();
3879 Self::moe_cached_gemm_q8(e, il, PROJ_GATE, ex, m, max_block, zq, zd)?
3880 } else {
3881 Self::moe_cached_gemm(e, il, PROJ_GATE, ex, m, max_block, &zt)?
3882 };
3883 let up = if up_q8 {
3884 let (zq, zd) = tok_q8.as_ref().unwrap();
3885 Self::moe_cached_gemm_q8(e, il, PROJ_UP, ex, m, max_block, zq, zd)?
3886 } else {
3887 Self::moe_cached_gemm(e, il, PROJ_UP, ex, m, max_block, &zt)?
3888 };
3889 let mut act = e.uninit(n_ff_exp)?;
3890 Self::ffn_act_lim(
3891 e,
3892 cfg,
3893 &gate,
3894 &up,
3895 m.gate_exps.macro_scale(ex),
3896 m.up_exps.macro_scale(ex),
3897 lim_exp,
3898 &mut act,
3899 n_ff_exp,
3900 )?;
3901 let y = if down_q8 {
3902 let (aq2, ad2) = e.quantize_q8_1(&act, 1, n_ff_exp)?;
3903 Self::moe_cached_gemm_q8(e, il, PROJ_DOWN, ex, m, max_block, &aq2, &ad2)?
3904 } else {
3905 let actv = act.slice(0..n_ff_exp);
3906 Self::moe_cached_gemm(e, il, PROJ_DOWN, ex, m, max_block, &actv)?
3907 };
3908 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
3909 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
3911 } else if cache_dispatch {
3912 let gate = Self::moe_cached_gemm(e, il, PROJ_GATE, ex, m, max_block, &zt)?;
3917 let up = Self::moe_cached_gemm(e, il, PROJ_UP, ex, m, max_block, &zt)?;
3918 let mut act = e.uninit(n_ff_exp)?; Self::ffn_act_lim(e, cfg, &gate, &up, m.gate_exps.macro_scale(ex),
3920 m.up_exps.macro_scale(ex), lim_exp, &mut act, n_ff_exp)?;
3921 let actv = act.slice(0..n_ff_exp);
3922 let y = Self::moe_cached_gemm(e, il, PROJ_DOWN, ex, m, max_block, &actv)?;
3923 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
3924 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
3926 } else if cache_frozen {
3927 let gate = Self::moe_frozen_gemm(
3932 e,
3933 il,
3934 PROJ_GATE,
3935 ex,
3936 m,
3937 max_block,
3938 &zt,
3939 &mut scratch_g,
3940 g_len,
3941 )?;
3942 let up = Self::moe_frozen_gemm(
3943 e,
3944 il,
3945 PROJ_UP,
3946 ex,
3947 m,
3948 max_block,
3949 &zt,
3950 &mut scratch_u,
3951 u_len,
3952 )?;
3953 let mut act = e.uninit(n_ff_exp)?;
3954 Self::ffn_act_lim(
3955 e,
3956 cfg,
3957 &gate,
3958 &up,
3959 m.gate_exps.macro_scale(ex),
3960 m.up_exps.macro_scale(ex),
3961 lim_exp,
3962 &mut act,
3963 n_ff_exp,
3964 )?;
3965 let actv = act.slice(0..n_ff_exp);
3966 let y = Self::moe_frozen_gemm(
3967 e,
3968 il,
3969 PROJ_DOWN,
3970 ex,
3971 m,
3972 max_block,
3973 &actv,
3974 &mut scratch_d,
3975 d_len,
3976 )?;
3977 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
3978 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
3979 } else {
3980 if scratch_g.is_none() {
3984 scratch_g = Some(e.alloc_u8_uninit(g_len)?);
3985 scratch_u = Some(e.alloc_u8_uninit(u_len)?);
3986 scratch_d = Some(e.alloc_u8_uninit(d_len)?);
3987 }
3988 let (sg, su, sd) = (scratch_g.as_mut().unwrap(), scratch_u.as_mut().unwrap(),
3989 scratch_d.as_mut().unwrap());
3990 let gl = m.gate_exps.expert_layout(ex);
3991 let ul = m.up_exps.expert_layout(ex);
3992 let dl = m.down_exps.expert_layout(ex);
3993 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
3994 let gate = e.qmatvec_view(sg, 0..gl.len, &zt, 1,
3995 m.gate_exps.in_f, m.gate_exps.out_f, gl.qtype, gl.row_bytes)?;
3996
3997 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
3998 let up = e.qmatvec_view(su, 0..ul.len, &zt, 1,
3999 m.up_exps.in_f, m.up_exps.out_f, ul.qtype, ul.row_bytes)?;
4000
4001 let mut act = e.uninit(n_ff_exp)?; Self::ffn_act_lim(e, cfg, &gate, &up, m.gate_exps.macro_scale(ex),
4003 m.up_exps.macro_scale(ex), lim_exp, &mut act, n_ff_exp)?;
4004
4005 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
4006 let actv = act.slice(0..n_ff_exp);
4007 let y = e.qmatvec_view(sd, 0..dl.len, &actv, 1,
4008 m.down_exps.in_f, m.down_exps.out_f, dl.qtype, dl.row_bytes)?;
4009
4010 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
4011 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
4012 }
4013 }
4014 if let Some(worker) = cpu_worker {
4015 let cpu_output = worker.wait().map_err(std::io::Error::other)?;
4016 let cpu_output = e.htod(&cpu_output)?;
4017 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
4018 e.axpy_into(&cpu_output, 1.0, &mut dst, n_embd)?;
4019 }
4020 if cpu_hybrid && !cache_frozen && cpu_expert_profile_admit_enabled() {
4021 for (j, &ex) in sel.iter().enumerate() {
4022 if cpu_mask[j] {
4023 Self::moe_profile_admit_expert(e, il, ex as usize, m, max_block)?;
4024 }
4025 }
4026 }
4027 }
4028
4029 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
4034 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
4035 {
4036 let n_ff_sh = gate_shexp.out_features(); let verify_t = t > 1 && t < PRIME_MIN_T;
4045 let (sg_gate, sg_up) = if t == 1 {
4046 match e.matmul_q8_fused2_x(gate_shexp, up_shexp, z)? {
4047 Some(pair) => pair,
4048 None => (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?),
4049 }
4050 } else if verify_t {
4051 (e.matmul_decode_exact(gate_shexp, z, t)?, e.matmul_decode_exact(up_shexp, z, t)?)
4052 } else {
4053 (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?) };
4055 let mut sa = e.uninit(t * n_ff_sh)?; Self::ffn_act_lim(e, cfg, &sg_gate, &sg_up, 1.0, 1.0, lim_shexp, &mut sa, t * n_ff_sh)?;
4057 let sh = if verify_t { e.matmul_decode_exact(down_shexp, &sa, t)? }
4058 else { e.matmul(down_shexp, &sa, t)? }; let g = match &m.gate_inp_shexp {
4072 Some(gate_inp_shexp) => {
4073 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
4074 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
4075 } else {
4076 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
4077 let mut g = e.uninit(t)?; e.sigmoid(&gs, &mut g, t)?;
4079 g
4080 }
4081 }
4082 None => e.htod(&vec![1.0f32; t])?,
4083 };
4084 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
4086 }
4087
4088 Ok(moe_out)
4089 }
4090
4091 pub fn stage1_h2d_per_token(&self) -> u64 {
4094 use crate::hybrid::Ffn;
4095 let n_used = self.cfg.moe.as_ref().map(|m| m.expert_used_count as u64).unwrap_or(0);
4096 let mut bytes = 0u64;
4097 for l in self.layers.iter() {
4098 if let Ffn::Moe(m) = &l.ffn {
4099 bytes += n_used * (m.gate_exps.max_expert_bytes() + m.up_exps.max_expert_bytes()
4100 + m.down_exps.max_expert_bytes()) as u64;
4101 }
4102 }
4103 bytes
4104 }
4105
4106 pub(crate) fn max_moe_block(&self) -> usize {
4110 use crate::hybrid::Ffn;
4111 let mut mx = 0usize;
4112 let mut scan = |ffn: &Ffn| {
4113 if let Ffn::Moe(m) = ffn {
4114 mx = mx.max(m.gate_exps.max_expert_bytes())
4115 .max(m.up_exps.max_expert_bytes())
4116 .max(m.down_exps.max_expert_bytes());
4117 }
4118 };
4119 for l in self.layers.iter() { scan(&l.ffn); }
4120 if let Some(mtp) = self.mtp.as_ref() { scan(&mtp.ffn); }
4121 mx
4122 }
4123
4124 pub(crate) fn moe_cache_block_sizes(&self) -> Vec<usize> {
4127 use crate::hybrid::Ffn;
4128 let mut sizes = Vec::new();
4129 let mut scan = |ffn: &Ffn| {
4130 let Ffn::Moe(m) = ffn else { return };
4131 for ex in 0..m.gate_exps.n_expert {
4132 if m.active_experts.as_ref().is_some_and(|active| !active[ex]) {
4133 continue;
4134 }
4135 for exps in [&m.gate_exps, &m.up_exps, &m.down_exps] {
4136 let len = exps.expert_layout(ex).len;
4137 if len > 0 {
4138 sizes.push(len);
4139 }
4140 }
4141 }
4142 };
4143 for layer in &self.layers {
4144 scan(&layer.ffn);
4145 }
4146 if let Some(mtp) = &self.mtp {
4147 scan(&mtp.ffn);
4148 }
4149 sizes
4150 }
4151
4152 pub fn save_cpu_expert_residency_profile(
4158 &self,
4159 e: &Engine,
4160 path: &std::path::Path,
4161 ) -> Result<(), Box<dyn std::error::Error>> {
4162 let Some(ids) = e.export_moe_residency() else {
4163 return Err("no MoE residency cache to persist".into());
4164 };
4165 let mut body = format!(
4166 "memra-freeze-profile v1 max_block={} blocks={}\n",
4167 self.max_moe_block(),
4168 ids.len()
4169 );
4170 for (layer, proj, ex) in &ids {
4171 body.push_str(&format!("{layer} {proj} {ex}\n"));
4172 }
4173 let tmp = path.with_extension("tmp");
4174 std::fs::write(&tmp, body)?;
4175 std::fs::rename(&tmp, path)?;
4176 println!(
4177 "[moe-cache] freeze profile saved: {} blocks -> {}",
4178 ids.len(),
4179 path.display()
4180 );
4181 Ok(())
4182 }
4183
4184 pub fn restore_cpu_expert_residency_profile(
4188 &self,
4189 e: &Engine,
4190 path: &std::path::Path,
4191 ) -> Result<bool, Box<dyn std::error::Error>> {
4192 use crate::hybrid::Ffn;
4193 use crate::moe_cache::BlockId;
4194 let Ok(content) = std::fs::read_to_string(path) else {
4195 return Ok(false);
4196 };
4197 let mut lines = content.lines();
4198 let Some(header) = lines.next() else { return Ok(false) };
4199 let expected = format!("memra-freeze-profile v1 max_block={}", self.max_moe_block());
4200 if !header.starts_with(&expected) {
4201 println!(
4202 "[moe-cache] freeze profile ignored (geometry mismatch): {}",
4203 path.display()
4204 );
4205 return Ok(false);
4206 }
4207 let mut by_layer: std::collections::HashMap<u16, Vec<BlockId>> =
4208 std::collections::HashMap::new();
4209 for line in lines {
4210 let mut fields = line.split_whitespace();
4211 let (Some(layer), Some(proj), Some(ex)) =
4212 (fields.next(), fields.next(), fields.next())
4213 else {
4214 continue;
4215 };
4216 let (Ok(layer), Ok(proj), Ok(ex)) =
4217 (layer.parse::<u16>(), proj.parse::<u8>(), ex.parse::<u16>())
4218 else {
4219 continue;
4220 };
4221 by_layer
4222 .entry(layer)
4223 .or_default()
4224 .push(BlockId::new(layer, proj, ex));
4225 }
4226 let requested: usize = by_layer.values().map(Vec::len).sum();
4227 if requested == 0 {
4228 return Ok(false);
4229 }
4230 let max_block = self.max_moe_block();
4231 let mut restaged = 0usize;
4232 let mut stage_layer = |layer_index: u16,
4233 ffn: &Ffn|
4234 -> Result<(), Box<dyn std::error::Error>> {
4235 let Ffn::Moe(m) = ffn else { return Ok(()) };
4236 let Some(ids) = by_layer.get(&layer_index) else {
4237 return Ok(());
4238 };
4239 e.with_moe_cache(max_block, |cache, eng| {
4240 for id in ids {
4241 if cache.restage_block(*id, m, eng)? {
4242 restaged += 1;
4243 }
4244 }
4245 Ok(())
4246 })
4247 };
4248 for (index, layer) in self.layers.iter().enumerate() {
4249 stage_layer(index as u16, &layer.ffn)?;
4250 }
4251 if let Some(mtp) = self.mtp.as_ref() {
4252 stage_layer(u16::MAX, &mtp.ffn)?;
4253 }
4254 e.freeze_moe_cache();
4255 println!(
4256 "[moe-cache] freeze profile restored: {restaged}/{requested} blocks restaged from {}",
4257 path.display()
4258 );
4259 Ok(true)
4260 }
4261
4262 pub fn freeze_cpu_expert_residency(
4264 &self,
4265 e: &Engine,
4266 ) -> Result<(), Box<dyn std::error::Error>> {
4267 e.freeze_moe_cache();
4268 Ok(())
4269 }
4270
4271 pub fn ffn_act(e: &Engine, cfg: &ModelConfig, gate: &CudaSlice<f32>, up: &CudaSlice<f32>,
4279 act: &mut CudaSlice<f32>, n: usize) -> Result<(), Box<dyn std::error::Error>> {
4280 Self::ffn_act_scaled(e, cfg, gate, up, 1.0, 1.0, act, n)
4281 }
4282
4283 #[allow(clippy::too_many_arguments)]
4287 pub(crate) fn ffn_act_scaled(e: &Engine, cfg: &ModelConfig, gate: &CudaSlice<f32>, up: &CudaSlice<f32>,
4288 gs: f32, us: f32, act: &mut CudaSlice<f32>, n: usize)
4289 -> Result<(), Box<dyn std::error::Error>> {
4290 Self::ffn_act_lim(e, cfg, gate, up, gs, us, None, act, n)
4291 }
4292
4293 #[allow(clippy::too_many_arguments)]
4302 pub(crate) fn ffn_act_lim(e: &Engine, cfg: &ModelConfig, gate: &CudaSlice<f32>, up: &CudaSlice<f32>,
4303 gs: f32, us: f32, limit: Option<f32>, act: &mut CudaSlice<f32>, n: usize)
4304 -> Result<(), Box<dyn std::error::Error>> {
4305 if let Some(m3) = cfg.m3.as_ref() {
4306 debug_assert!(limit.is_none(), "m3 swigluoai and step35 clamp are different archs");
4307 return e.swigluoai_mul_scaled(gate, up, gs, us, m3.swiglu_alpha, m3.swiglu_limit, act, n);
4308 }
4309 if let Some(l) = limit {
4310 return e.swiglu_clamped_mul_scaled(gate, up, gs, us, l, act, n);
4311 }
4312 if gs == 1.0 && us == 1.0 { return e.silu_mul(gate, up, act, n); }
4313 e.silu_mul_scaled(gate, up, gs, us, act, n)
4314 }
4315
4316 fn moe_route(e: &Engine, logits: &CudaSlice<f32>, t: usize, n_expert: usize, n_used: usize)
4322 -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
4323 Self::moe_route_cfg(e, logits, t, n_expert, n_used, None, None, None)
4324 }
4325
4326 fn moe_route_cfg(e: &Engine, logits: &CudaSlice<f32>, t: usize, n_expert: usize, n_used: usize,
4334 bias: Option<&[f32]>, sig: Option<(f32, bool)>, active: Option<&[bool]>)
4335 -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
4336 if let Some((sf, route_norm)) = sig {
4337 let lg = e.dtoh(logits)?;
4339 return Self::moe_route_sigmoid_host(
4340 &lg, t, n_expert, n_used, bias, sf, route_norm, active,
4341 );
4342 }
4343 if active.is_none() && !matches!(std::env::var("MEMRA_FUSED_ROUTER").as_deref(), Ok("0")) {
4347 return e.moe_router_topk_host(logits, t, n_expert, n_used);
4348 }
4349 let lg = e.dtoh(logits)?; let mut sel = vec![0u32; t * n_used];
4352 let mut w_out = vec![0f32; t * n_used];
4353 for tok in 0..t {
4354 let row = &lg[tok * n_expert..(tok + 1) * n_expert];
4355 let maxl = row.iter().enumerate()
4357 .filter(|(i, _)| active.is_none_or(|mask| mask[*i]))
4358 .map(|(_, &x)| x).fold(f32::NEG_INFINITY, f32::max);
4359 let mut probs = vec![0f32; n_expert];
4360 let mut den = 0f32;
4361 for i in 0..n_expert {
4362 if active.is_some_and(|mask| !mask[i]) { continue; }
4363 let x = (row[i] - maxl).exp(); probs[i] = x; den += x;
4364 }
4365 for p in probs.iter_mut() { *p /= den; }
4366 let mut idx: Vec<usize> = (0..n_expert)
4368 .filter(|&i| active.is_none_or(|mask| mask[i])).collect();
4369 idx.sort_by(|&a, &b| probs[b].total_cmp(&probs[a]).then(a.cmp(&b)));
4370 let sl = &idx[..n_used];
4371 let mut wv: Vec<f32> = sl.iter().map(|&i| probs[i]).collect();
4372 let mut ws: f32 = wv.iter().sum();
4373 ws = ws.max(6.103515625e-5_f32); for x in wv.iter_mut() { *x /= ws; }
4375 for j in 0..n_used {
4376 sel[tok * n_used + j] = sl[j] as u32;
4377 w_out[tok * n_used + j] = wv[j];
4378 }
4379 }
4380 Ok((sel, w_out))
4381 }
4382
4383 #[allow(clippy::too_many_arguments)]
4384 fn moe_route_sigmoid_with_input(
4385 e: &Engine,
4386 logits: &CudaSlice<f32>,
4387 input: &CudaSlice<f32>,
4388 t: usize,
4389 n_expert: usize,
4390 n_used: usize,
4391 bias: Option<&[f32]>,
4392 (sf, route_norm): (f32, bool),
4393 active: Option<&[bool]>,
4394 ) -> Result<(Vec<u32>, Vec<f32>, Vec<f32>), Box<dyn std::error::Error>> {
4395 let (lg, input) = e.dtoh_pair(logits, input)?;
4396 let (sel, w) =
4397 Self::moe_route_sigmoid_host(&lg, t, n_expert, n_used, bias, sf, route_norm, active)?;
4398 Ok((sel, w, input))
4399 }
4400
4401 pub fn start_moe_prefetch_predictor(
4406 &self,
4407 e: &Engine,
4408 cfg: &ModelConfig,
4409 ) -> Result<(), Box<dyn std::error::Error>> {
4410 use crate::hybrid::Ffn;
4411 let Some(sig) = cfg.sigmoid_router() else {
4412 return Err("prefetch predictor requires a sigmoid-router arch".into());
4413 };
4414 let resident: std::collections::HashSet<(u16, u8, u16)> = e
4415 .export_moe_residency()
4416 .ok_or("prefetch predictor needs the frozen MoE residency cache")?
4417 .into_iter()
4418 .collect();
4419 let mut layers = Vec::new();
4420 for (index, layer) in self.layers.iter().enumerate() {
4421 let Ffn::Moe(m) = &layer.ffn else { continue };
4422 let crate::model::GpuTensor::Float { data, .. } = &m.gate_inp else { continue };
4423 let router = e.dtoh(data)?;
4424 let n_expert = m.gate_exps.n_expert;
4425 let n_embd = m.gate_exps.in_f;
4426 if router.len() != n_embd * n_expert {
4427 continue;
4428 }
4429 let build = |exps: &crate::model::HostExps| {
4430 (0..n_expert)
4431 .map(|expert| crate::cpu_experts::predictor_projection(exps, expert))
4432 .collect::<Vec<_>>()
4433 };
4434 layers.push((index as u16, crate::cpu_experts::PredictLayerInit {
4435 router,
4436 bias: m.exp_probs_b.clone(),
4437 active: m.active_experts.clone(),
4438 n_embd,
4439 n_used: cfg
4440 .moe
4441 .as_ref()
4442 .map(|moe| moe.expert_used_count as usize)
4443 .ok_or("prefetch predictor requires MoE config")?,
4444 sig,
4445 weights_n_expert: n_expert,
4446 gate: build(&m.gate_exps),
4447 up: build(&m.up_exps),
4448 down: build(&m.down_exps),
4449 }));
4450 }
4451 crate::cpu_experts::start_prefetch_predictor(layers, resident)
4452 .map_err(|error| error.into())
4453 }
4454
4455 #[allow(clippy::too_many_arguments)]
4458 pub(crate) fn moe_route_sigmoid_host_public(
4459 logits: &[f32],
4460 t: usize,
4461 n_expert: usize,
4462 n_used: usize,
4463 bias: Option<&[f32]>,
4464 sf: f32,
4465 route_norm: bool,
4466 active: Option<&[bool]>,
4467 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
4468 Self::moe_route_sigmoid_host(logits, t, n_expert, n_used, bias, sf, route_norm, active)
4469 }
4470
4471 #[allow(clippy::too_many_arguments)]
4472 fn moe_route_sigmoid_host(
4473 lg: &[f32],
4474 t: usize,
4475 n_expert: usize,
4476 n_used: usize,
4477 bias: Option<&[f32]>,
4478 sf: f32,
4479 route_norm: bool,
4480 active: Option<&[bool]>,
4481 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
4482 if lg.len() != t * n_expert {
4483 return Err(format!(
4484 "sigmoid router logits length mismatch: got {}, expected {}",
4485 lg.len(),
4486 t * n_expert,
4487 )
4488 .into());
4489 }
4490 let mut sel = vec![0u32; t * n_used];
4491 let mut w_out = vec![0f32; t * n_used];
4492 for tok in 0..t {
4493 let row = &lg[tok * n_expert..(tok + 1) * n_expert];
4494 let scores: Vec<f32> = row.iter().map(|&x| 1.0 / (1.0 + (-x).exp())).collect();
4495 let selsc: Vec<f32> = match bias {
4497 Some(b) => scores.iter().zip(b).map(|(s, bb)| s + bb).collect(),
4498 None => scores.clone(),
4499 };
4500 let mut idx: Vec<usize> = (0..n_expert)
4501 .filter(|&i| active.is_none_or(|mask| mask[i]))
4502 .collect();
4503 idx.sort_by(|&a, &b| selsc[b].total_cmp(&selsc[a]).then(a.cmp(&b)));
4504 let sl = &idx[..n_used];
4505 let mut wv: Vec<f32> = sl.iter().map(|&i| scores[i]).collect();
4506 if route_norm {
4507 let ws: f32 = wv.iter().sum::<f32>().max(1e-20);
4508 for x in wv.iter_mut() {
4509 *x = *x / ws * sf;
4510 }
4511 } else {
4512 for x in wv.iter_mut() {
4513 *x *= sf;
4514 }
4515 }
4516 for j in 0..n_used {
4517 sel[tok * n_used + j] = sl[j] as u32;
4518 w_out[tok * n_used + j] = wv[j];
4519 }
4520 }
4521 Ok((sel, w_out))
4522 }
4523
4524 fn moe_ffn_pairs(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, logits: &CudaSlice<f32>,
4533 t: usize, cfg: &ModelConfig)
4534 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4535 let moe = cfg.moe.as_ref().unwrap();
4536 let n_embd = cfg.n_embd as usize;
4537 let n_expert = moe.expert_count as usize;
4538 let n_used = moe.expert_used_count as usize;
4539 let n_ff_exp = moe.expert_ff_length as usize;
4540 debug_assert!(!cfg.swiglu_clamped_anywhere(),
4545 "moe_ffn_pairs has no per-layer clamp: fused epilogues are plain SiLU");
4546 let dev = m.dev_exps.as_ref().unwrap();
4547 let (rbg_d, rbu_d) = if dev.gu_il {
4549 let sxx = m.gate_exps.row_bytes + m.up_exps.row_bytes; (sxx, sxx)
4550 } else { (m.gate_exps.row_bytes, m.up_exps.row_bytes) };
4551
4552 let (sel_all, w_all) = Self::moe_route(e, logits, t, n_expert, n_used)?;
4553 let n_pairs = t * n_used;
4554 let pair_tok: Vec<i32> = (0..n_pairs).map(|p| (p / n_used) as i32).collect();
4557 let pair_ex: Vec<i32> = sel_all.iter().map(|&x| x as i32).collect();
4558 let pair_w: Vec<f32> = w_all.clone();
4559 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
4560 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
4561 let pt = e.htod_i32(&pair_tok)?;
4562 let px = e.htod_i32(&pair_ex)?;
4563 let pw = e.htod(&pair_w)?;
4564 let toff = e.htod_i32(&tok_off)?;
4565 let tids = e.htod_i32(&tok_ids)?;
4566
4567 let mut by_ex: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
4571 for p in 0..n_pairs { by_ex[pair_ex[p] as usize].push(p as i32); }
4572 let mut ex_ids: Vec<i32> = Vec::new();
4573 let mut ex_off: Vec<i32> = vec![0];
4574 let mut ex_pairs: Vec<i32> = Vec::with_capacity(n_pairs);
4575 for (ex, list) in by_ex.iter().enumerate() {
4576 if list.is_empty() { continue; }
4577 ex_ids.push(ex as i32);
4578 ex_pairs.extend_from_slice(list);
4579 ex_off.push(ex_pairs.len() as i32);
4580 }
4581 let n_active = ex_ids.len();
4582 let exi = e.htod_i32(&ex_ids)?;
4583 let exo = e.htod_i32(&ex_off)?;
4584 let exp_d = e.htod_i32(&ex_pairs)?;
4585 let _ = &px; static MMA_T: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
4606 let mma_t = *MMA_T.get_or_init(|| {
4607 std::env::var("MEMRA_MOE_MMA_T").ok().and_then(|v| v.parse().ok()).unwrap_or(16)
4608 });
4609 let use_mma = std::env::var("MEMRA_MOE_MMA").map(|v| v != "0").unwrap_or(true)
4610 && t >= mma_t
4611 && q8_expert_dec_supported(m.gate_exps.qtype) && q8_expert_dec_supported(m.up_exps.qtype)
4612 && q8_expert_dec_supported(m.down_exps.qtype)
4613 && n_embd % 256 == 0 && n_ff_exp % 256 == 0;
4614 let mma_capable = q8_expert_dec_supported(m.gate_exps.qtype)
4630 && q8_expert_dec_supported(m.up_exps.qtype)
4631 && q8_expert_dec_supported(m.down_exps.qtype)
4632 && n_embd % 256 == 0 && n_ff_exp % 256 == 0;
4633 let f16g_mode = crate::moe_f16g_mode();
4634 let f16g = f16g_mode != 0 && t >= mma_t
4635 && (f16g_mode != 3 || !mma_capable)
4636 && f16g_proj_ok(m.gate_exps.qtype, n_embd)
4637 && f16g_proj_ok(m.up_exps.qtype, n_embd)
4638 && f16g_proj_ok(m.down_exps.qtype, n_ff_exp);
4639 if use_mma || f16g {
4640 let y_down = if f16g {
4648 let csr_tok: Vec<i32> = ex_pairs.iter().map(|&p| p / n_used as i32).collect();
4652 let csr_tok_d = e.htod_i32(&csr_tok)?;
4653 let (z_f16, z_s) = e.moe_f16g_act(z, Some(&csr_tok_d), n_embd, n_pairs)?;
4654 let g_csr = e.moe_f16_grouped(&dev.ptr_row, 0, n_expert, &exi, &ex_off, &exo,
4655 &z_f16, &z_s, n_embd, n_ff_exp, n_active, n_pairs,
4656 m.gate_exps.qtype, rbg_d)?;
4657 let u_csr = e.moe_f16_grouped(&dev.ptr_row, 1, n_expert, &exi, &ex_off, &exo,
4658 &z_f16, &z_s, n_embd, n_ff_exp, n_active, n_pairs,
4659 m.up_exps.qtype, rbu_d)?;
4660 let act_csr = e.moe_pairs_silu_mul(&g_csr, &u_csr, n_pairs * n_ff_exp)?;
4661 let (a_f16, a_s) = e.moe_f16g_act(&act_csr, None, n_ff_exp, n_pairs)?;
4662 let d_csr = e.moe_f16_grouped(&dev.ptr_row, 2, n_expert, &exi, &ex_off, &exo,
4663 &a_f16, &a_s, n_ff_exp, n_embd, n_active, n_pairs,
4664 m.down_exps.qtype, m.down_exps.row_bytes)?;
4665 e.rows_permute(&d_csr, &exp_d, n_pairs, n_embd)?
4666 } else {
4667 let z_scr = e.mmq_iq_quantize_act(z, n_embd, t)?;
4669 let gate = e.mmq_iq_experts(&dev.ptr_row, 0, n_expert, &exi, &exo, &exp_d, &pt, &z_scr,
4670 n_embd, n_ff_exp, n_active, n_pairs, t,
4671 m.gate_exps.qtype, rbg_d)?;
4672 let up = e.mmq_iq_experts(&dev.ptr_row, 1, n_expert, &exi, &exo, &exp_d, &pt, &z_scr,
4673 n_embd, n_ff_exp, n_active, n_pairs, t,
4674 m.up_exps.qtype, rbu_d)?;
4675 let a_scr = if crate::moe_fuse_actq_on() {
4681 e.mmq_iq_fused_act_quant(&gate, &up, n_ff_exp, n_pairs, 0)?
4682 } else {
4683 let act = e.moe_pairs_silu_mul(&gate, &up, n_pairs * n_ff_exp)?;
4684 e.mmq_iq_quantize_act(&act, n_ff_exp, n_pairs)?
4685 };
4686 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
4687 let pself = e.htod_i32(&pair_self)?;
4688 e.mmq_iq_experts(&dev.ptr_row, 2, n_expert, &exi, &exo, &exp_d, &pself, &a_scr,
4689 n_ff_exp, n_embd, n_active, n_pairs, n_pairs,
4690 m.down_exps.qtype, m.down_exps.row_bytes)?
4691 };
4692 let mut moe_out = e.uninit(t * n_embd)?;
4693 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
4694 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
4695 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
4696 {
4697 let n_ff_sh = gate_shexp.out_features();
4698 let sg_gate = e.matmul(gate_shexp, z, t)?;
4699 let sg_up = e.matmul(up_shexp, z, t)?;
4700 let mut sa = e.uninit(t * n_ff_sh)?;
4701 Self::ffn_act(e, cfg, &sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
4702 let sh = e.matmul(down_shexp, &sa, t)?;
4703 let g = match &m.gate_inp_shexp {
4709 Some(gate_inp_shexp) if crate::router_prefill_exact_on() => {
4710 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
4711 }
4712 Some(gate_inp_shexp) => {
4713 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
4714 let mut g = e.uninit(t)?;
4715 e.sigmoid(&gs, &mut g, t)?;
4716 g
4717 }
4718 None => e.htod(&vec![1.0f32; t])?,
4719 };
4720 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
4721 }
4722 return Ok(moe_out);
4723 }
4724
4725 let dec = std::env::var("MEMRA_MOE_DEC").map(|v| v != "0").unwrap_or(true);
4728 let matvec = |proj, exi: &_, exo: &_, exp_d: &_, pt: &_, aq: &_, ad: &_,
4729 inf, outf, qtype, rb| -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4730 let dec = dec && q8_expert_dec_supported(qtype);
4732 if dec { e.moe_pairs_matvec_q8_dec(&dev.ptr_row, proj, exi, exo, exp_d, pt, aq, ad,
4733 inf, outf, n_expert, n_active, n_pairs, qtype, rb) }
4734 else { e.moe_pairs_matvec_q8_em (&dev.ptr_row, proj, exi, exo, exp_d, pt, aq, ad,
4735 inf, outf, n_expert, n_active, n_pairs, qtype, rb) }
4736 };
4737 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
4738 let gate = matvec(0, &exi, &exo, &exp_d, &pt, &zq, &zd,
4739 n_embd, n_ff_exp, m.gate_exps.qtype, rbg_d)?;
4740 let up = matvec(1, &exi, &exo, &exp_d, &pt, &zq, &zd,
4741 n_embd, n_ff_exp, m.up_exps.qtype, rbu_d)?;
4742 let act = e.moe_pairs_silu_mul(&gate, &up, n_pairs * n_ff_exp)?;
4743 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
4744 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
4746 let pself = e.htod_i32(&pair_self)?;
4747 let y_down = matvec(2, &exi, &exo, &exp_d, &pself, &aq2, &ad2,
4748 n_ff_exp, n_embd, m.down_exps.qtype, m.down_exps.row_bytes)?;
4749 let mut moe_out = e.uninit(t * n_embd)?; e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
4751
4752 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
4756 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
4757 {
4758 let n_ff_sh = gate_shexp.out_features();
4759 let sg_gate = e.matmul(gate_shexp, z, t)?;
4760 let sg_up = e.matmul(up_shexp, z, t)?;
4761 let mut sa = e.uninit(t * n_ff_sh)?;
4762 e.silu_mul(&sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
4763 let sh = e.matmul(down_shexp, &sa, t)?;
4764 let g = match &m.gate_inp_shexp {
4769 Some(gate_inp_shexp) if crate::router_prefill_exact_on() => {
4770 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
4771 }
4772 Some(gate_inp_shexp) => {
4773 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
4774 let mut g = e.uninit(t)?;
4775 e.sigmoid(&gs, &mut g, t)?;
4776 g
4777 }
4778 None => e.htod(&vec![1.0f32; t])?,
4779 };
4780 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
4781 }
4782 Ok(moe_out)
4783 }
4784
4785 #[allow(clippy::too_many_arguments)]
4787 #[allow(clippy::too_many_arguments)]
4788 fn moe_ffn_dev(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>,
4789 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>, logits: &CudaSlice<f32>,
4790 t: usize, cfg: &ModelConfig, il: u16, max_block: usize)
4791 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4792 let moe = cfg.moe.as_ref().unwrap();
4793 let n_embd = cfg.n_embd as usize;
4794 let n_expert = moe.expert_count as usize;
4795 let n_used = moe.expert_used_count as usize;
4796 let n_ff_exp = moe.expert_ff_length as usize;
4797 debug_assert!(cfg.sigmoid_router().is_none(),
4801 "moe_ffn_dev routes SOFTMAX: a sigmoid-router arch would pick wrong experts");
4802 debug_assert!(!cfg.swiglu_clamped_at(il as u32),
4803 "moe_ffn_dev's fused epilogue is plain SiLU: no clamped form");
4804
4805 let (sel_d, mut w_d) = e.moe_router_topk(logits, t, n_expert, n_used)?;
4807 if m.has_macros {
4810 e.moe_w_scale_by_expert(&mut w_d, &sel_d, &m.dev_macros, n_expert, t * n_used)?;
4811 }
4812
4813 let mut moe_out = e.uninit(t * n_embd)?;
4815
4816 if let Some(dev) = m.dev_exps.as_ref() {
4819 let (rbg_d, rbu_d) = if dev.gu_il {
4822 let sxx = m.gate_exps.row_bytes + m.up_exps.row_bytes; (sxx, sxx)
4823 } else { (m.gate_exps.row_bytes, m.up_exps.row_bytes) };
4824 let q8 = moe_q8_enabled()
4825 && q8_expert_supported(m.gate_exps.qtype) && q8_expert_supported(m.up_exps.qtype)
4826 && q8_expert_supported(m.down_exps.qtype);
4827 let rows_arm = q8 && t > 1 && crate::spec::spec_m2()
4836 && n_ff_exp == 512 && n_used <= 8
4837 && std::env::var("MEMRA_MOE_DEVQ8_GU").map(|v| v.is_empty() || v == "v").unwrap_or(true)
4838 && std::env::var("MEMRA_MOE_DEVQ8_DOWN").map(|v| v.is_empty() || v == "w8h2v").unwrap_or(true);
4839 let csr_mode = std::env::var("MEMRA_MOE_CSR").ok()
4848 .and_then(|v| v.parse::<i32>().ok()).unwrap_or(1);
4849 let csr_qt = |qt: i32| qt == crate::QT_IQ4_XS || qt == crate::QT_IQ3_S;
4850 let csr_arm = rows_arm && csr_mode > 0 && t <= 10
4851 && csr_qt(m.gate_exps.qtype) && csr_qt(m.up_exps.qtype)
4852 && csr_qt(m.down_exps.qtype);
4853 if csr_arm {
4854 if csr_mode == 2 {
4855 static ENGAGED: std::sync::Once = std::sync::Once::new();
4856 ENGAGED.call_once(|| eprintln!("[memra] moe CSR byte-compare mode ON (t={t})"));
4857 }
4858 let n_pairs = t * n_used;
4859 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
4860 let act = e.moe_gate_up_silu8_dev_q8_csr(&dev.ptr_row, &sel_d, &zq, &zd, n_pairs,
4861 n_embd, n_ff_exp, n_used, n_expert,
4862 m.gate_exps.qtype, m.up_exps.qtype,
4863 rbg_d, rbu_d)?;
4864 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
4865 e.moe_down8_fma_dev_q8_rows(&dev.ptr_row, &sel_d, &w_d, &aq2, &ad2, &mut moe_out,
4869 t, n_ff_exp, n_embd, n_used, n_expert,
4870 m.down_exps.qtype, m.down_exps.row_bytes)?;
4871 if csr_mode == 2 {
4872 let act_r = e.moe_gate_up_silu8_dev_q8_rows(&dev.ptr_row, &sel_d, &zq, &zd, t,
4874 n_embd, n_ff_exp, n_used, n_expert,
4875 m.gate_exps.qtype, m.up_exps.qtype,
4876 rbg_d, rbu_d, &m.dev_macros)?;
4877 let mut out_r = e.uninit(t * n_embd)?;
4878 let (aq2r, ad2r) = e.quantize_q8_1(&act_r, n_pairs, n_ff_exp)?;
4879 e.moe_down8_fma_dev_q8_rows(&dev.ptr_row, &sel_d, &w_d, &aq2r, &ad2r, &mut out_r,
4880 t, n_ff_exp, n_embd, n_used, n_expert,
4881 m.down_exps.qtype, m.down_exps.row_bytes)?;
4882 let (a1, a2) = (e.dtoh(&act)?, e.dtoh(&act_r)?);
4883 let (o1, o2) = (e.dtoh(&moe_out)?, e.dtoh(&out_r)?);
4884 let ba = a1.iter().zip(&a2).filter(|(x, y)| x.to_bits() != y.to_bits()).count();
4885 let bo = o1.iter().zip(&o2).filter(|(x, y)| x.to_bits() != y.to_bits()).count();
4886 if ba + bo > 0 {
4887 eprintln!("[csr-check] il={il} t={t} ACT diffs={ba}/{} OUT diffs={bo}/{}",
4888 a1.len(), o1.len());
4889 let sel_h = e.dtoh_i32(&sel_d)?;
4891 let mut shown = 0;
4892 for (i, (x, y)) in a1.iter().zip(&a2).enumerate() {
4893 if x.to_bits() != y.to_bits() && shown < 4 {
4894 let (p, o) = (i / n_ff_exp, i % n_ff_exp);
4895 let ex = sel_h[p];
4896 let npx = sel_h.iter().filter(|&&v| v == ex).count();
4897 eprintln!(" ACT p={p} ex={ex} np={npx} o={o} csr={x:e} rows={y:e}");
4898 shown += 1;
4899 }
4900 }
4901 std::process::exit(3);
4902 }
4903 }
4904 } else if rows_arm {
4905 if std::env::var("MEMRA_MOE_OVERLAP").as_deref() == Ok("1") {
4908 use std::sync::atomic::{AtomicU64, Ordering};
4909 static PAIRS: AtomicU64 = AtomicU64::new(0);
4910 static UNIQ: AtomicU64 = AtomicU64::new(0);
4911 static CALLS: AtomicU64 = AtomicU64::new(0);
4912 let sel_h = e.dtoh_i32(&sel_d)?;
4913 let mut u: Vec<i32> = sel_h.clone(); u.sort_unstable(); u.dedup();
4914 PAIRS.fetch_add(sel_h.len() as u64, Ordering::Relaxed);
4915 UNIQ.fetch_add(u.len() as u64, Ordering::Relaxed);
4916 let c = CALLS.fetch_add(1, Ordering::Relaxed) + 1;
4917 if c % 480 == 0 {
4918 let p = PAIRS.load(Ordering::Relaxed); let q = UNIQ.load(Ordering::Relaxed);
4919 eprintln!("[overlap] calls={c} pairs={p} unique={q} ratio={:.3} (t={t})",
4920 q as f64 / p as f64);
4921 }
4922 }
4923 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
4924 let act = e.moe_gate_up_silu8_dev_q8_rows(&dev.ptr_row, &sel_d, &zq, &zd, t,
4925 n_embd, n_ff_exp, n_used, n_expert,
4926 m.gate_exps.qtype, m.up_exps.qtype,
4927 rbg_d, rbu_d, &m.dev_macros)?;
4928 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
4929 e.moe_down8_fma_dev_q8_rows(&dev.ptr_row, &sel_d, &w_d, &aq2, &ad2, &mut moe_out,
4930 t, n_ff_exp, n_embd, n_used, n_expert,
4931 m.down_exps.qtype, m.down_exps.row_bytes)?;
4932 } else {
4933 for tok in 0..t {
4934 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
4935 let selt = sel_d.slice(tok * n_used..(tok + 1) * n_used);
4936 let wt = w_d.slice(tok * n_used..(tok + 1) * n_used);
4937 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
4938 if q8 {
4939 let (zq, zd) = match (t, zq8) {
4940 (1, Some((q, d))) => (q.clone(), d.clone()),
4941 _ => e.quantize_q8_1_view(&zt, 1, n_embd)?,
4942 };
4943 let act = e.moe_gate_up_silu8_dev_q8(&dev.ptr_row, &selt, &zq, &zd,
4944 n_embd, n_ff_exp, n_used, n_expert,
4945 m.gate_exps.qtype, m.up_exps.qtype,
4946 rbg_d, rbu_d, &m.dev_macros)?;
4947 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
4948 e.moe_down8_fma_dev_q8(&dev.ptr_row, &selt, &wt, &aq2, &ad2, &mut dst,
4949 n_ff_exp, n_embd, n_used, n_expert,
4950 m.down_exps.qtype, m.down_exps.row_bytes)?;
4951 } else {
4952 let act = e.moe_gate_up_silu8_dev(&dev.ptr_row, &selt, &zt, n_embd, n_ff_exp,
4953 n_used, n_expert,
4954 m.gate_exps.qtype, m.up_exps.qtype,
4955 rbg_d, rbu_d, &m.dev_macros)?;
4956 e.moe_down8_fma_dev(&dev.ptr_row, &selt, &wt, &act, &mut dst,
4957 n_ff_exp, n_embd, n_used, n_expert,
4958 m.down_exps.qtype, m.down_exps.row_bytes)?;
4959 }
4960 }
4961 }
4962 } else {
4963 let q8 = moe_q8_enabled()
4970 && q8_expert_supported(m.gate_exps.qtype) && q8_expert_supported(m.up_exps.qtype)
4971 && q8_expert_supported(m.down_exps.qtype);
4972 e.with_moe_cache(max_block, |c, eng| {
4973 let row = c.layer_dev_row(il, n_expert, eng)?
4974 .ok_or("moe_ffn_dev: layer row vanished under the lock")?;
4975 for tok in 0..t {
4976 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
4977 let selt = sel_d.slice(tok * n_used..(tok + 1) * n_used);
4978 let wt = w_d.slice(tok * n_used..(tok + 1) * n_used);
4979 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
4980 if q8 {
4981 let (zq, zd) = match (t, zq8) {
4982 (1, Some((q, d))) => (q.clone(), d.clone()),
4983 _ => eng.quantize_q8_1_view(&zt, 1, n_embd)?,
4984 };
4985 let act = eng.moe_gate_up_silu8_dev_q8(row, &selt, &zq, &zd,
4986 n_embd, n_ff_exp, n_used, n_expert,
4987 m.gate_exps.qtype, m.up_exps.qtype,
4988 m.gate_exps.row_bytes, m.up_exps.row_bytes,
4989 &m.dev_macros)?;
4990 let (aq2, ad2) = eng.quantize_q8_1(&act, n_used, n_ff_exp)?;
4991 eng.moe_down8_fma_dev_q8(row, &selt, &wt, &aq2, &ad2, &mut dst,
4992 n_ff_exp, n_embd, n_used, n_expert,
4993 m.down_exps.qtype, m.down_exps.row_bytes)?;
4994 } else {
4995 let act = eng.moe_gate_up_silu8_dev(row, &selt, &zt, n_embd, n_ff_exp,
4996 n_used, n_expert,
4997 m.gate_exps.qtype, m.up_exps.qtype,
4998 m.gate_exps.row_bytes, m.up_exps.row_bytes,
4999 &m.dev_macros)?;
5000 eng.moe_down8_fma_dev(row, &selt, &wt, &act, &mut dst,
5001 n_ff_exp, n_embd, n_used, n_expert,
5002 m.down_exps.qtype, m.down_exps.row_bytes)?;
5003 }
5004 }
5005 c.hits += (t * 3 * n_used) as u64;
5007 Ok(())
5008 })?;
5009 }
5010
5011 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
5016 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
5017 {
5018 let n_ff_sh = gate_shexp.out_features();
5019 let verify_t = t > 1 && t < PRIME_MIN_T;
5022 let (sg_gate, sg_up) = if t == 1 {
5023 match e.matmul_q8_fused2_x(gate_shexp, up_shexp, z)? {
5024 Some(pair) => pair,
5025 None => (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?),
5026 }
5027 } else if verify_t {
5028 let mut fused = None;
5032 if crate::spec::spec_fused_t() && (2..=4).contains(&t)
5033 && e.uses_q8_1_fast(gate_shexp) && e.uses_q8_1_fast(up_shexp) {
5034 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
5035 fused = e.matmul_q8_fused2_t(gate_shexp, up_shexp, &zq, &zd, t)?;
5036 }
5037 match fused {
5038 Some(pair) => pair,
5039 None => (e.matmul_decode_exact(gate_shexp, z, t)?,
5040 e.matmul_decode_exact(up_shexp, z, t)?),
5041 }
5042 } else {
5043 (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?)
5044 };
5045 let mut sa = e.uninit(t * n_ff_sh)?; e.silu_mul(&sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
5047 let sh = if verify_t { e.matmul_decode_exact(down_shexp, &sa, t)? }
5048 else { e.matmul(down_shexp, &sa, t)? };
5049 let g = match &m.gate_inp_shexp {
5053 Some(gate_inp_shexp) => {
5054 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
5057 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
5058 } else {
5059 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
5060 let mut g = e.uninit(t)?;
5061 e.sigmoid(&gs, &mut g, t)?;
5062 g
5063 }
5064 }
5065 None => e.htod(&vec![1.0f32; t])?,
5066 };
5067 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
5068 }
5069
5070 Ok(moe_out)
5071 }
5072
5073 #[allow(clippy::too_many_arguments)]
5083 #[allow(clippy::too_many_arguments)]
5086 fn moe_gdec_token_q8(e: &Engine, m: &MoeWeights, il: u16, max_block: usize,
5087 zq: &CudaSlice<i8>, zd: &CudaSlice<f32>, sel: &[u32], w: &[f32],
5088 moe_out: &mut CudaSlice<f32>, tok: usize,
5089 n_embd: usize, n_ff_exp: usize, n_used: usize)
5090 -> Result<bool, Box<dyn std::error::Error>> {
5091 use crate::moe_cache::{BlockId, PROJ_GATE, PROJ_UP, PROJ_DOWN};
5092 use cudarc::driver::DevicePtr;
5093 let ptrs = e.with_moe_cache(max_block, |c, eng| {
5094 let mut g = [0u64; 8];
5095 let mut u = [0u64; 8];
5096 let mut d = [0u64; 8];
5097 for (j, &ex) in sel.iter().enumerate() {
5098 let ex = ex as u16;
5099 let (Some(sg), Some(su), Some(sd)) = (c.resident(BlockId::new(il, PROJ_GATE, ex)),
5100 c.resident(BlockId::new(il, PROJ_UP, ex)),
5101 c.resident(BlockId::new(il, PROJ_DOWN, ex)))
5102 else { return Ok(None); };
5103 let __s = eng.stream();
5104 let (pg, _e0) = c.slot(sg).device_ptr(&__s);
5105 let (pu, _e1) = c.slot(su).device_ptr(&__s);
5106 let (pd, _e2) = c.slot(sd).device_ptr(&__s);
5107 g[j] = pg as u64; u[j] = pu as u64; d[j] = pd as u64;
5108 }
5109 if cpu_expert_profile_admit_enabled() && !c.is_frozen() {
5110 for &ex in sel {
5111 let ex = ex as u16;
5112 for proj in [PROJ_GATE, PROJ_UP, PROJ_DOWN] {
5113 c.note_profile_hit(BlockId::new(il, proj, ex));
5114 }
5115 }
5116 }
5117 c.hits += (3 * n_used) as u64;
5118 Ok(Some((g, u, d)))
5119 })?;
5120 let Some((g, u, d)) = ptrs else { return Ok(false) };
5121 let mut wv = [0f32; 8];
5122 wv[..n_used].copy_from_slice(w);
5123 let act = e.moe_gate_up_silu8_q8(crate::WPtr8(g), crate::WPtr8(u), zq, zd,
5124 n_embd, n_ff_exp, n_used,
5125 m.gate_exps.qtype, m.up_exps.qtype,
5126 m.gate_exps.row_bytes, m.up_exps.row_bytes)?;
5127 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
5129 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
5130 e.moe_down8_fma_q8(crate::WPtr8(d), crate::F32x8(wv), &aq2, &ad2, &mut dst,
5131 n_ff_exp, n_embd, n_used,
5132 m.down_exps.qtype, m.down_exps.row_bytes)?;
5133 Ok(true)
5134 }
5135
5136 fn moe_gdec_token(e: &Engine, m: &MoeWeights, il: u16, max_block: usize,
5137 zt: &cudarc::driver::CudaView<f32>, sel: &[u32], w: &[f32],
5138 moe_out: &mut CudaSlice<f32>, tok: usize,
5139 n_embd: usize, n_ff_exp: usize, n_used: usize)
5140 -> Result<bool, Box<dyn std::error::Error>> {
5141 use crate::moe_cache::{BlockId, PROJ_GATE, PROJ_UP, PROJ_DOWN};
5142 use cudarc::driver::DevicePtr;
5143 let ptrs = e.with_moe_cache(max_block, |c, eng| {
5145 let mut g = [0u64; 8];
5146 let mut u = [0u64; 8];
5147 let mut d = [0u64; 8];
5148 for (j, &ex) in sel.iter().enumerate() {
5149 let ex = ex as u16;
5150 let (Some(sg), Some(su), Some(sd)) = (c.resident(BlockId::new(il, PROJ_GATE, ex)),
5151 c.resident(BlockId::new(il, PROJ_UP, ex)),
5152 c.resident(BlockId::new(il, PROJ_DOWN, ex)))
5153 else { return Ok(None); };
5154 let __s = eng.stream();
5155 let (pg, _e0) = c.slot(sg).device_ptr(&__s);
5156 let (pu, _e1) = c.slot(su).device_ptr(&__s);
5157 let (pd, _e2) = c.slot(sd).device_ptr(&__s);
5158 g[j] = pg as u64; u[j] = pu as u64; d[j] = pd as u64;
5159 }
5160 if cpu_expert_profile_admit_enabled() && !c.is_frozen() {
5161 for &ex in sel {
5162 let ex = ex as u16;
5163 for proj in [PROJ_GATE, PROJ_UP, PROJ_DOWN] {
5164 c.note_profile_hit(BlockId::new(il, proj, ex));
5165 }
5166 }
5167 }
5168 c.hits += (3 * n_used) as u64; Ok(Some((g, u, d)))
5170 })?;
5171 let Some((g, u, d)) = ptrs else { return Ok(false) };
5172 let mut wv = [0f32; 8];
5173 wv[..n_used].copy_from_slice(w);
5174 let act = e.moe_gate_up_silu8(crate::WPtr8(g), crate::WPtr8(u), zt,
5176 n_embd, n_ff_exp, n_used,
5177 m.gate_exps.qtype, m.up_exps.qtype,
5178 m.gate_exps.row_bytes, m.up_exps.row_bytes)?;
5179 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
5180 e.moe_down8_fma_into(crate::WPtr8(d), crate::F32x8(wv), &act, &mut dst,
5181 n_ff_exp, n_embd, n_used,
5182 m.down_exps.qtype, m.down_exps.row_bytes)?;
5183 Ok(true)
5184 }
5185
5186 fn moe_cached_gemm_q8(e: &Engine, il: u16, proj: u8, ex: usize, m: &MoeWeights,
5191 max_block: usize, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
5192 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5193 use crate::moe_cache::{BlockId, DispatchSlot, PROJ_GATE, PROJ_UP};
5194 let exps = match proj { PROJ_GATE => &m.gate_exps, PROJ_UP => &m.up_exps, _ => &m.down_exps };
5195 let layout = exps.expert_layout(ex);
5196 let id = BlockId::new(il, proj, ex as u16);
5197 let source = exps.expert_source(ex);
5198 e.with_moe_cache(max_block, |c, eng| {
5199 let slot = c.dispatch_source(id, source, eng)?;
5200 let DispatchSlot::Resident(sl) = slot;
5201 let buf = c.slot(sl);
5202 eng.qmatvec_expert_q8(buf, 0..layout.len, aq, ad, 1, exps.in_f, exps.out_f,
5203 layout.qtype, layout.row_bytes)
5204 })
5205 }
5206
5207 fn moe_cached_gemm(e: &Engine, il: u16, proj: u8, ex: usize, m: &MoeWeights,
5208 max_block: usize, x: &cudarc::driver::CudaView<f32>)
5209 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5210 use crate::moe_cache::{BlockId, DispatchSlot, PROJ_GATE, PROJ_UP};
5211 let exps = match proj { PROJ_GATE => &m.gate_exps, PROJ_UP => &m.up_exps, _ => &m.down_exps };
5212 let layout = exps.expert_layout(ex);
5213 let id = BlockId::new(il, proj, ex as u16);
5214 let source = exps.expert_source(ex);
5215 e.with_moe_cache(max_block, |c, eng| {
5217 let slot = c.dispatch_source(id, source, eng)?;
5218 let DispatchSlot::Resident(sl) = slot;
5221 let buf = c.slot(sl);
5222 eng.qmatvec_view(buf, 0..layout.len, x, 1, exps.in_f, exps.out_f,
5223 layout.qtype, layout.row_bytes)
5224 })
5225 }
5226
5227 fn moe_profile_admit_expert(
5231 e: &Engine,
5232 il: u16,
5233 ex: usize,
5234 m: &MoeWeights,
5235 max_block: usize,
5236 ) -> Result<(), Box<dyn std::error::Error>> {
5237 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
5238 e.with_moe_cache(max_block, |cache, eng| {
5239 for (proj, exps) in [
5240 (PROJ_GATE, &m.gate_exps),
5241 (PROJ_UP, &m.up_exps),
5242 (PROJ_DOWN, &m.down_exps),
5243 ] {
5244 let id = BlockId::new(il, proj, ex as u16);
5245 let _ = cache.dispatch_source(id, exps.expert_source(ex), eng)?;
5246 }
5247 Ok(())
5248 })
5249 }
5250
5251 #[allow(clippy::too_many_arguments)]
5254 fn moe_frozen_gemm(
5255 e: &Engine,
5256 il: u16,
5257 proj: u8,
5258 ex: usize,
5259 m: &MoeWeights,
5260 max_block: usize,
5261 x: &cudarc::driver::CudaView<f32>,
5262 scratch: &mut Option<CudaSlice<u8>>,
5263 scratch_len: usize,
5264 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5265 use crate::moe_cache::{BlockId, PROJ_GATE, PROJ_UP};
5266 let exps = match proj {
5267 PROJ_GATE => &m.gate_exps,
5268 PROJ_UP => &m.up_exps,
5269 _ => &m.down_exps,
5270 };
5271 let layout = exps.expert_layout(ex);
5272 let id = BlockId::new(il, proj, ex as u16);
5273 if let Some(output) = e.with_moe_cache(max_block, |cache, eng| {
5274 let Some(slot) = cache.resident(id) else {
5275 return Ok(None);
5276 };
5277 let buf = cache.slot(slot);
5278 Ok(Some(eng.qmatvec_view(
5279 buf,
5280 0..layout.len,
5281 x,
5282 1,
5283 exps.in_f,
5284 exps.out_f,
5285 layout.qtype,
5286 layout.row_bytes,
5287 )?))
5288 })? {
5289 return Ok(output);
5290 }
5291 if scratch.is_none() {
5292 *scratch = Some(e.alloc_u8_uninit(scratch_len)?);
5293 }
5294 let scratch = scratch.as_mut().unwrap();
5295 e.stage_expert(exps.expert_bytes(ex), scratch, 0)?;
5296 e.qmatvec_view(
5297 scratch,
5298 0..layout.len,
5299 x,
5300 1,
5301 exps.in_f,
5302 exps.out_f,
5303 layout.qtype,
5304 layout.row_bytes,
5305 )
5306 }
5307
5308 fn moe_prefetch_expert(
5309 e: &Engine,
5310 il: u16,
5311 ex: usize,
5312 m: &MoeWeights,
5313 max_block: usize,
5314 keep: &[crate::moe_cache::BlockId],
5315 ) -> Result<(), Box<dyn std::error::Error>> {
5316 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
5317 e.with_moe_cache(max_block, |c, eng| {
5318 for (proj, exps) in [(PROJ_GATE, &m.gate_exps), (PROJ_UP, &m.up_exps),
5319 (PROJ_DOWN, &m.down_exps)] {
5320 let id = BlockId::new(il, proj, ex as u16);
5321 let _ = c.prefetch_source(id, exps.expert_source(ex), keep, eng)?;
5322 }
5323 Ok(())
5324 })
5325 }
5326
5327 fn moe_prefetch_disk_expert(e: &Engine, il: u16, ex: usize, m: &MoeWeights,
5330 max_block: usize, keep: &[crate::moe_cache::BlockId])
5331 -> Result<(), Box<dyn std::error::Error>> {
5332 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
5333 e.with_moe_cache(max_block, |c, eng| {
5334 for (proj, exps) in [(PROJ_GATE, &m.gate_exps), (PROJ_UP, &m.up_exps),
5335 (PROJ_DOWN, &m.down_exps)] {
5336 let source = exps.expert_source(ex);
5337 if let crate::model::ExpertSource::Disk { .. } = &source {
5338 let id = BlockId::new(il, proj, ex as u16);
5339 let _ = c.prefetch_source(id, source, keep, eng)?;
5340 }
5341 }
5342 Ok(())
5343 })
5344 }
5345
5346 #[inline]
5347 fn moe_prefetch_host_expert(ex: usize, m: &MoeWeights) {
5348 let _ = m.gate_exps.prefetch_expert_pages(ex);
5349 let _ = m.up_exps.prefetch_expert_pages(ex);
5350 let _ = m.down_exps.prefetch_expert_pages(ex);
5351 }
5352}
5353
5354impl HybridModel {
5371 #[allow(clippy::too_many_arguments)]
5375 fn moe_ffn_grouped_resident_q8(
5376 e: &Engine,
5377 m: &MoeWeights,
5378 z: &CudaSlice<f32>,
5379 t: usize,
5380 cfg: &ModelConfig,
5381 il: u16,
5382 sel_all: &[u32],
5383 w_all: &[f32],
5384 table: &CudaSlice<u64>,
5385 gu_il: bool,
5386 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5387 let moe = cfg.moe.as_ref().unwrap();
5388 let n_embd = cfg.n_embd as usize;
5389 let n_expert = moe.expert_count as usize;
5390 let n_used = moe.expert_used_count as usize;
5391 let n_ff_exp = moe.expert_ff_length as usize;
5392 let n_pairs = t * n_used;
5393 debug_assert_eq!(sel_all.len(), n_pairs);
5394 debug_assert_eq!(w_all.len(), n_pairs);
5395 debug_assert!(
5396 m.gate_exps.macros.is_none()
5397 && m.up_exps.macros.is_none()
5398 && m.down_exps.macros.is_none(),
5399 "resident grouped q8 does not fold per-expert macro scales",
5400 );
5401
5402 if !cfg.swiglu_clamped_at(il as u32) {
5407 let sel: Vec<i32> = sel_all.iter().map(|&expert| expert as i32).collect();
5408 let sel_d = e.htod_i32(&sel)?;
5409 let w_d = e.htod(w_all)?;
5410 let (gate_row_bytes, up_row_bytes) = if gu_il {
5411 let combined = m.gate_exps.row_bytes + m.up_exps.row_bytes;
5412 (combined, combined)
5413 } else {
5414 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
5415 };
5416 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
5417 let act = e.moe_gate_up_silu8_dev_q8_rows(
5418 table,
5419 &sel_d,
5420 &zq,
5421 &zd,
5422 t,
5423 n_embd,
5424 n_ff_exp,
5425 n_used,
5426 n_expert,
5427 m.gate_exps.qtype,
5428 m.up_exps.qtype,
5429 gate_row_bytes,
5430 up_row_bytes,
5431 &m.dev_macros,
5432 )?;
5433 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
5434 let mut moe_out = e.uninit(t * n_embd)?;
5435 e.moe_down8_fma_dev_q8_rows_g(
5436 table,
5437 &sel_d,
5438 &w_d,
5439 &aq2,
5440 &ad2,
5441 &mut moe_out,
5442 t,
5443 n_ff_exp,
5444 n_embd,
5445 n_used,
5446 n_expert,
5447 m.down_exps.qtype,
5448 m.down_exps.row_bytes,
5449 )?;
5450
5451 if std::env::var("MEMRA_MOE_STATS").is_ok() {
5452 let mut counts = vec![0usize; n_expert];
5453 for &expert in sel_all {
5454 counts[expert as usize] += 1;
5455 }
5456 let mut sizes: Vec<usize> =
5457 counts.into_iter().filter(|&count| count != 0).collect();
5458 sizes.sort_unstable();
5459 let mean = sizes.iter().sum::<usize>() as f64 / sizes.len().max(1) as f64;
5460 println!(
5461 "moe-grouped il={il} t={t} dispatch=resident-q8-rows active={}/{} \
5462 m_e: min={} median={} mean={mean:.1} max={}",
5463 sizes.len(),
5464 n_expert,
5465 sizes.first().copied().unwrap_or(0),
5466 sizes.get(sizes.len() / 2).copied().unwrap_or(0),
5467 sizes.last().copied().unwrap_or(0),
5468 );
5469 }
5470 return Ok(moe_out);
5471 }
5472
5473 let pair_tok: Vec<i32> = (0..n_pairs).map(|pair| (pair / n_used) as i32).collect();
5477 let pair_ex: Vec<i32> = sel_all.iter().map(|&expert| expert as i32).collect();
5478 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
5479 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
5480
5481 let mut by_expert: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
5482 for (pair, &expert) in pair_ex.iter().enumerate() {
5483 by_expert[expert as usize].push(pair as i32);
5484 }
5485
5486 let pair_tok_d = e.htod_i32(&pair_tok)?;
5487 let pair_ex_d = e.htod_i32(&pair_ex)?;
5488 let pair_w_d = e.htod(w_all)?;
5489 let tok_off_d = e.htod_i32(&tok_off)?;
5490 let tok_ids_d = e.htod_i32(&tok_ids)?;
5491
5492 let matvec = |
5493 proj: i32,
5494 pair_rows: &CudaSlice<i32>,
5495 aq: &CudaSlice<i8>,
5496 ad: &CudaSlice<f32>,
5497 in_f: usize,
5498 out_f: usize,
5499 qtype: i32,
5500 row_bytes: usize,
5501 | -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5502 e.moe_pairs_matvec_q8(
5503 table,
5504 proj,
5505 pair_rows,
5506 &pair_ex_d,
5507 aq,
5508 ad,
5509 in_f,
5510 out_f,
5511 n_expert,
5512 n_pairs,
5513 qtype,
5514 row_bytes,
5515 )
5516 };
5517
5518 let (gate_row_bytes, up_row_bytes) = if gu_il {
5519 let combined = m.gate_exps.row_bytes + m.up_exps.row_bytes;
5520 (combined, combined)
5521 } else {
5522 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
5523 };
5524 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
5525 let gate = matvec(
5526 0,
5527 &pair_tok_d,
5528 &zq,
5529 &zd,
5530 n_embd,
5531 n_ff_exp,
5532 m.gate_exps.qtype,
5533 gate_row_bytes,
5534 )?;
5535 let up = matvec(
5536 1,
5537 &pair_tok_d,
5538 &zq,
5539 &zd,
5540 n_embd,
5541 n_ff_exp,
5542 m.up_exps.qtype,
5543 up_row_bytes,
5544 )?;
5545 let mut act = e.uninit(n_pairs * n_ff_exp)?;
5546 Self::ffn_act_lim(
5547 e,
5548 cfg,
5549 &gate,
5550 &up,
5551 1.0,
5552 1.0,
5553 cfg.clamp_exp_at(il as u32),
5554 &mut act,
5555 n_pairs * n_ff_exp,
5556 )?;
5557 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
5558 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
5559 let pair_self_d = e.htod_i32(&pair_self)?;
5560 let down = matvec(
5561 2,
5562 &pair_self_d,
5563 &aq2,
5564 &ad2,
5565 n_ff_exp,
5566 n_embd,
5567 m.down_exps.qtype,
5568 m.down_exps.row_bytes,
5569 )?;
5570 let mut moe_out = e.uninit(t * n_embd)?;
5571 e.moe_pairs_scatter(
5572 &down,
5573 &pair_w_d,
5574 &tok_off_d,
5575 &tok_ids_d,
5576 &mut moe_out,
5577 t,
5578 n_embd,
5579 )?;
5580
5581 if std::env::var("MEMRA_MOE_STATS").is_ok() {
5582 let mut sizes: Vec<usize> = by_expert
5583 .iter()
5584 .filter_map(|pairs| (!pairs.is_empty()).then_some(pairs.len()))
5585 .collect();
5586 sizes.sort_unstable();
5587 let mean = sizes.iter().sum::<usize>() as f64 / sizes.len().max(1) as f64;
5588 println!(
5589 "moe-grouped il={il} t={t} dispatch=resident-q8-clamped-pairs active={}/{} \
5590 m_e: min={} median={} mean={mean:.1} max={}",
5591 sizes.len(),
5592 n_expert,
5593 sizes.first().copied().unwrap_or(0),
5594 sizes.get(sizes.len() / 2).copied().unwrap_or(0),
5595 sizes.last().copied().unwrap_or(0),
5596 );
5597 }
5598 Ok(moe_out)
5599 }
5600
5601 fn moe_ffn_grouped_add_shared(
5602 e: &Engine,
5603 m: &MoeWeights,
5604 z: &CudaSlice<f32>,
5605 t: usize,
5606 cfg: &ModelConfig,
5607 il: u16,
5608 moe_out: &mut CudaSlice<f32>,
5609 ) -> Result<(), Box<dyn std::error::Error>> {
5610 let n_embd = cfg.n_embd as usize;
5611 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
5612 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
5613 {
5614 let n_ff_sh = gate_shexp.out_features();
5615 let sg_gate = e.matmul(gate_shexp, z, t)?;
5616 let sg_up = e.matmul(up_shexp, z, t)?;
5617 let mut sa = e.uninit(t * n_ff_sh)?;
5618 Self::ffn_act_lim(
5619 e,
5620 cfg,
5621 &sg_gate,
5622 &sg_up,
5623 1.0,
5624 1.0,
5625 cfg.clamp_shexp_at(il as u32),
5626 &mut sa,
5627 t * n_ff_sh,
5628 )?;
5629 let sh = e.matmul(down_shexp, &sa, t)?;
5630 let gate = match &m.gate_inp_shexp {
5631 Some(gate_inp_shexp) => {
5632 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
5633 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
5634 } else {
5635 let raw = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
5636 let mut gate = e.uninit(t)?;
5637 e.sigmoid(&raw, &mut gate, t)?;
5638 gate
5639 }
5640 }
5641 None => e.htod(&vec![1.0f32; t])?,
5642 };
5643 e.add_scaled_rows(&sh, &gate, moe_out, n_embd, t)?;
5644 }
5645 Ok(())
5646 }
5647
5648 pub(crate) fn moe_ffn_grouped(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, t: usize,
5651 cfg: &ModelConfig, il: u16, max_block: usize)
5652 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5653 let moe = cfg.moe.as_ref().unwrap();
5654 let n_embd = cfg.n_embd as usize;
5655 let n_expert = moe.expert_count as usize;
5656 let n_used = moe.expert_used_count as usize;
5657 let n_ff_exp = moe.expert_ff_length as usize;
5658 let lim_exp = cfg.clamp_exp_at(il as u32);
5660
5661 let logits = Self::moe_router_logits(e, m, z, t, cfg)?;
5664 let (sel_all, w_all) = if let Some(sig) = cfg.sigmoid_router() {
5665 Self::moe_route_cfg(e, &logits, t, n_expert, n_used,
5666 m.exp_probs_b.as_deref(), Some(sig), m.active_experts.as_deref())?
5667 } else {
5668 Self::moe_route_cfg(e, &logits, t, n_expert, n_used,
5669 None, None, m.active_experts.as_deref())?
5670 };
5671 Self::trace_moe_routes(il, t, &sel_all, &w_all)?;
5672 Self::trace_moe_input(e, il, t, n_embd, z)?;
5673
5674 let no_exp_macros = m.gate_exps.macros.is_none()
5679 && m.up_exps.macros.is_none()
5680 && m.down_exps.macros.is_none();
5681 let resident_q8 = m.dev_exps.as_ref().filter(|dev| {
5682 m.has_uniform_expert_layout()
5683 && no_exp_macros
5684 && moe_q8_enabled()
5685 && q8_expert_supported(m.gate_exps.qtype)
5686 && q8_expert_supported(m.up_exps.qtype)
5687 && q8_expert_supported(m.down_exps.qtype)
5688 && moe_slab_enabled()
5689 && dev.dev == e.ctx().ordinal()
5690 });
5691 if let Some(dev) = resident_q8 {
5692 let mut moe_out = Self::moe_ffn_grouped_resident_q8(
5693 e,
5694 m,
5695 z,
5696 t,
5697 cfg,
5698 il,
5699 &sel_all,
5700 &w_all,
5701 &dev.ptr_row,
5702 dev.gu_il,
5703 )?;
5704 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut moe_out)?;
5705 return Ok(moe_out);
5706 }
5707
5708 struct ExpertGroup {
5712 tok_indices: Vec<i32>, slot_indices: Vec<i32>, weights: Vec<f32>, }
5716 let mut groups: Vec<ExpertGroup> = (0..n_expert).map(|_| ExpertGroup {
5717 tok_indices: Vec::new(), slot_indices: Vec::new(), weights: Vec::new(),
5718 }).collect();
5719
5720 for tok in 0..t {
5721 for j in 0..n_used {
5722 let ex = sel_all[tok * n_used + j] as usize;
5723 let w = w_all[tok * n_used + j];
5724 groups[ex].tok_indices.push(tok as i32);
5725 groups[ex].slot_indices.push(j as i32);
5726 groups[ex].weights.push(w);
5727 }
5728 }
5729
5730 let mut slot_buf = e.zeros(t * n_used * n_embd)?;
5733 let mut wbuf = e.zeros(t * n_used)?; let g_len = m.gate_exps.max_expert_bytes();
5737 let u_len = m.up_exps.max_expert_bytes();
5738 let d_len = m.down_exps.max_expert_bytes();
5739 let moe_q8 = m.has_uniform_expert_layout()
5740 && moe_q8_enabled()
5741 && q8_expert_supported(m.gate_exps.qtype)
5742 && q8_expert_supported(m.up_exps.qtype)
5743 && q8_expert_supported(m.down_exps.qtype);
5744 let slab_local = m.dev_exps.as_ref().filter(|dev| {
5747 !dev.gu_il && moe_slab_enabled() && dev.dev == e.ctx().ordinal()
5748 });
5749 let use_cache =
5750 slab_local.is_none() && Engine::moe_cache_enabled() && !e.moe_cache_frozen();
5751 let grouped_q8 = moe_q8 && (slab_local.is_some() || use_cache);
5754
5755 let (mut scratch_g, mut scratch_u, mut scratch_d) = if slab_local.is_none() && !use_cache {
5757 (Some(e.alloc_u8(g_len)?), Some(e.alloc_u8(u_len)?), Some(e.alloc_u8(d_len)?))
5758 } else {
5759 (None, None, None)
5760 };
5761
5762 let mut order: Vec<usize> =
5773 (0..n_expert).filter(|&ex| !groups[ex].tok_indices.is_empty()).collect();
5774 order.sort_by(|&a, &b| groups[b].tok_indices.len()
5775 .cmp(&groups[a].tok_indices.len()).then(a.cmp(&b)));
5776 let mut m_dist: Vec<usize> = Vec::new(); let page_window = moe_page_prefetch_window();
5778 let worker_disk_prefetch = use_cache && crate::spill_pread::worker_enabled();
5779 if worker_disk_prefetch {
5780 if let Some(first) = grouped_worker_prefetch_position(order.len(), None) {
5781 Self::moe_prefetch_disk_expert(e, il, order[first], m, max_block, &[])?;
5782 }
5783 }
5784 for (order_pos, &ex) in order.iter().enumerate() {
5785 for next in page_prefetch_positions(order_pos, order.len(), page_window) {
5786 Self::moe_prefetch_host_expert(order[next], m);
5787 }
5788 if worker_disk_prefetch {
5789 if let Some(next) = grouped_worker_prefetch_position(order.len(), Some(order_pos)) {
5790 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
5791 let keep = [
5792 BlockId::new(il, PROJ_GATE, ex as u16),
5793 BlockId::new(il, PROJ_UP, ex as u16),
5794 BlockId::new(il, PROJ_DOWN, ex as u16),
5795 ];
5796 Self::moe_prefetch_disk_expert(e, il, order[next], m, max_block, &keep)?;
5797 }
5798 }
5799 let grp = &groups[ex];
5800 let m_e = grp.tok_indices.len();
5801 m_dist.push(m_e);
5802 let gl = m.gate_exps.expert_layout(ex);
5803 let ul = m.up_exps.expert_layout(ex);
5804 let dl = m.down_exps.expert_layout(ex);
5805
5806 let tok_idx_d = e.htod_i32(&grp.tok_indices)?;
5810 let slot_idx_d = e.htod_i32(&grp.slot_indices)?;
5811 let dmac = m.down_exps.macro_scale(ex);
5812 let weight_d = if dmac == 1.0 { e.htod(&grp.weights)? } else {
5813 let scaled: Vec<f32> = grp.weights.iter().map(|&w| w * dmac).collect();
5814 e.htod(&scaled)?
5815 };
5816
5817 let mut gathered = e.zeros(m_e * n_embd)?;
5819 e.gather_rows(z, &tok_idx_d, &mut gathered, n_embd, m_e)?;
5820 let gv = gathered.slice(0..m_e * n_embd);
5821
5822 let y = if let Some(dev) = slab_local {
5825 let gate_start = ex * m.gate_exps.expert_stride;
5826 let up_start = ex * m.up_exps.expert_stride;
5827 let down_start = ex * m.down_exps.expert_stride;
5828 if grouped_q8 {
5829 let (zq, zd) = e.quantize_q8_1(&gathered, m_e, n_embd)?;
5830 let gate = e.qmatvec_expert_q8(
5831 &dev.gate,
5832 gate_start..gate_start + gl.len,
5833 &zq,
5834 &zd,
5835 m_e,
5836 m.gate_exps.in_f,
5837 m.gate_exps.out_f,
5838 gl.qtype,
5839 gl.row_bytes,
5840 )?;
5841 let up = e.qmatvec_expert_q8(
5842 &dev.up,
5843 up_start..up_start + ul.len,
5844 &zq,
5845 &zd,
5846 m_e,
5847 m.up_exps.in_f,
5848 m.up_exps.out_f,
5849 ul.qtype,
5850 ul.row_bytes,
5851 )?;
5852 let mut act = e.uninit(m_e * n_ff_exp)?;
5853 Self::ffn_act_lim(
5854 e,
5855 cfg,
5856 &gate,
5857 &up,
5858 m.gate_exps.macro_scale(ex),
5859 m.up_exps.macro_scale(ex),
5860 lim_exp,
5861 &mut act,
5862 m_e * n_ff_exp,
5863 )?;
5864 let (aq2, ad2) = e.quantize_q8_1(&act, m_e, n_ff_exp)?;
5865 e.qmatvec_expert_q8(
5866 &dev.down,
5867 down_start..down_start + dl.len,
5868 &aq2,
5869 &ad2,
5870 m_e,
5871 m.down_exps.in_f,
5872 m.down_exps.out_f,
5873 dl.qtype,
5874 dl.row_bytes,
5875 )?
5876 } else {
5877 let gate = e.qmatvec_view(
5878 &dev.gate,
5879 gate_start..gate_start + gl.len,
5880 &gv,
5881 m_e,
5882 m.gate_exps.in_f,
5883 m.gate_exps.out_f,
5884 gl.qtype,
5885 gl.row_bytes,
5886 )?;
5887 let up = e.qmatvec_view(
5888 &dev.up,
5889 up_start..up_start + ul.len,
5890 &gv,
5891 m_e,
5892 m.up_exps.in_f,
5893 m.up_exps.out_f,
5894 ul.qtype,
5895 ul.row_bytes,
5896 )?;
5897 let mut act = e.uninit(m_e * n_ff_exp)?;
5898 Self::ffn_act_lim(
5899 e,
5900 cfg,
5901 &gate,
5902 &up,
5903 m.gate_exps.macro_scale(ex),
5904 m.up_exps.macro_scale(ex),
5905 lim_exp,
5906 &mut act,
5907 m_e * n_ff_exp,
5908 )?;
5909 let actv = act.slice(0..m_e * n_ff_exp);
5910 e.qmatvec_view(
5911 &dev.down,
5912 down_start..down_start + dl.len,
5913 &actv,
5914 m_e,
5915 m.down_exps.in_f,
5916 m.down_exps.out_f,
5917 dl.qtype,
5918 dl.row_bytes,
5919 )?
5920 }
5921 } else if use_cache {
5922 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
5923 if grouped_q8 {
5924 let (zq, zd) = e.quantize_q8_1(&gathered, m_e, n_embd)?;
5925 let gate = e.with_moe_cache(max_block, |cache, eng| {
5926 let id = BlockId::new(il, PROJ_GATE, ex as u16);
5927 let slot =
5928 cache.dispatch_source(id, m.gate_exps.expert_source(ex), eng)?;
5929 eng.qmatvec_expert_q8(
5930 cache.buf(slot),
5931 0..gl.len,
5932 &zq,
5933 &zd,
5934 m_e,
5935 m.gate_exps.in_f,
5936 m.gate_exps.out_f,
5937 gl.qtype,
5938 gl.row_bytes,
5939 )
5940 })?;
5941 let up = e.with_moe_cache(max_block, |cache, eng| {
5942 let id = BlockId::new(il, PROJ_UP, ex as u16);
5943 let slot =
5944 cache.dispatch_source(id, m.up_exps.expert_source(ex), eng)?;
5945 eng.qmatvec_expert_q8(
5946 cache.buf(slot),
5947 0..ul.len,
5948 &zq,
5949 &zd,
5950 m_e,
5951 m.up_exps.in_f,
5952 m.up_exps.out_f,
5953 ul.qtype,
5954 ul.row_bytes,
5955 )
5956 })?;
5957 let mut act = e.uninit(m_e * n_ff_exp)?;
5958 Self::ffn_act_lim(
5959 e,
5960 cfg,
5961 &gate,
5962 &up,
5963 m.gate_exps.macro_scale(ex),
5964 m.up_exps.macro_scale(ex),
5965 lim_exp,
5966 &mut act,
5967 m_e * n_ff_exp,
5968 )?;
5969 let (aq2, ad2) = e.quantize_q8_1(&act, m_e, n_ff_exp)?;
5970 e.with_moe_cache(max_block, |cache, eng| {
5971 let id = BlockId::new(il, PROJ_DOWN, ex as u16);
5972 let slot =
5973 cache.dispatch_source(id, m.down_exps.expert_source(ex), eng)?;
5974 eng.qmatvec_expert_q8(
5975 cache.buf(slot),
5976 0..dl.len,
5977 &aq2,
5978 &ad2,
5979 m_e,
5980 m.down_exps.in_f,
5981 m.down_exps.out_f,
5982 dl.qtype,
5983 dl.row_bytes,
5984 )
5985 })?
5986 } else {
5987 let gate = e.with_moe_cache(max_block, |cache, eng| {
5988 let id = BlockId::new(il, PROJ_GATE, ex as u16);
5989 let slot =
5990 cache.dispatch_source(id, m.gate_exps.expert_source(ex), eng)?;
5991 eng.qmatvec_view(
5992 cache.buf(slot),
5993 0..gl.len,
5994 &gv,
5995 m_e,
5996 m.gate_exps.in_f,
5997 m.gate_exps.out_f,
5998 gl.qtype,
5999 gl.row_bytes,
6000 )
6001 })?;
6002 let up = e.with_moe_cache(max_block, |cache, eng| {
6003 let id = BlockId::new(il, PROJ_UP, ex as u16);
6004 let slot =
6005 cache.dispatch_source(id, m.up_exps.expert_source(ex), eng)?;
6006 eng.qmatvec_view(
6007 cache.buf(slot),
6008 0..ul.len,
6009 &gv,
6010 m_e,
6011 m.up_exps.in_f,
6012 m.up_exps.out_f,
6013 ul.qtype,
6014 ul.row_bytes,
6015 )
6016 })?;
6017 let mut act = e.uninit(m_e * n_ff_exp)?;
6018 Self::ffn_act_lim(
6019 e,
6020 cfg,
6021 &gate,
6022 &up,
6023 m.gate_exps.macro_scale(ex),
6024 m.up_exps.macro_scale(ex),
6025 lim_exp,
6026 &mut act,
6027 m_e * n_ff_exp,
6028 )?;
6029 let actv = act.slice(0..m_e * n_ff_exp);
6030 e.with_moe_cache(max_block, |cache, eng| {
6031 let id = BlockId::new(il, PROJ_DOWN, ex as u16);
6032 let slot =
6033 cache.dispatch_source(id, m.down_exps.expert_source(ex), eng)?;
6034 eng.qmatvec_view(
6035 cache.buf(slot),
6036 0..dl.len,
6037 &actv,
6038 m_e,
6039 m.down_exps.in_f,
6040 m.down_exps.out_f,
6041 dl.qtype,
6042 dl.row_bytes,
6043 )
6044 })?
6045 }
6046 } else {
6047 let sg = scratch_g.as_mut().unwrap();
6048 let su = scratch_u.as_mut().unwrap();
6049 let sd = scratch_d.as_mut().unwrap();
6050 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
6051 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
6052 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
6053 if grouped_q8 {
6054 let (zq, zd) = e.quantize_q8_1(&gathered, m_e, n_embd)?;
6055 let gate = e.qmatvec_expert_q8(
6056 sg,
6057 0..gl.len,
6058 &zq,
6059 &zd,
6060 m_e,
6061 m.gate_exps.in_f,
6062 m.gate_exps.out_f,
6063 gl.qtype,
6064 gl.row_bytes,
6065 )?;
6066 let up = e.qmatvec_expert_q8(
6067 su,
6068 0..ul.len,
6069 &zq,
6070 &zd,
6071 m_e,
6072 m.up_exps.in_f,
6073 m.up_exps.out_f,
6074 ul.qtype,
6075 ul.row_bytes,
6076 )?;
6077 let mut act = e.uninit(m_e * n_ff_exp)?;
6078 Self::ffn_act_lim(
6079 e,
6080 cfg,
6081 &gate,
6082 &up,
6083 m.gate_exps.macro_scale(ex),
6084 m.up_exps.macro_scale(ex),
6085 lim_exp,
6086 &mut act,
6087 m_e * n_ff_exp,
6088 )?;
6089 let (aq2, ad2) = e.quantize_q8_1(&act, m_e, n_ff_exp)?;
6090 e.qmatvec_expert_q8(
6091 sd,
6092 0..dl.len,
6093 &aq2,
6094 &ad2,
6095 m_e,
6096 m.down_exps.in_f,
6097 m.down_exps.out_f,
6098 dl.qtype,
6099 dl.row_bytes,
6100 )?
6101 } else {
6102 let gate = e.qmatvec_view(
6103 sg,
6104 0..gl.len,
6105 &gv,
6106 m_e,
6107 m.gate_exps.in_f,
6108 m.gate_exps.out_f,
6109 gl.qtype,
6110 gl.row_bytes,
6111 )?;
6112 let up = e.qmatvec_view(
6113 su,
6114 0..ul.len,
6115 &gv,
6116 m_e,
6117 m.up_exps.in_f,
6118 m.up_exps.out_f,
6119 ul.qtype,
6120 ul.row_bytes,
6121 )?;
6122 let mut act = e.uninit(m_e * n_ff_exp)?;
6123 Self::ffn_act_lim(
6124 e,
6125 cfg,
6126 &gate,
6127 &up,
6128 m.gate_exps.macro_scale(ex),
6129 m.up_exps.macro_scale(ex),
6130 lim_exp,
6131 &mut act,
6132 m_e * n_ff_exp,
6133 )?;
6134 let actv = act.slice(0..m_e * n_ff_exp);
6135 e.qmatvec_view(
6136 sd,
6137 0..dl.len,
6138 &actv,
6139 m_e,
6140 m.down_exps.in_f,
6141 m.down_exps.out_f,
6142 dl.qtype,
6143 dl.row_bytes,
6144 )?
6145 }
6146 };
6147
6148 e.scatter_slot(&y, &tok_idx_d, &slot_idx_d, &weight_d,
6150 &mut slot_buf, &mut wbuf, n_embd, n_used, m_e)?;
6151 }
6152
6153 let mut moe_out = e.zeros(t * n_embd)?;
6155 e.reduce_slots(&slot_buf, &wbuf, &mut moe_out, n_embd, n_used, t)?;
6156
6157 if std::env::var("MEMRA_MOE_STATS").is_ok() && !m_dist.is_empty() {
6159 m_dist.sort_unstable();
6160 let active = m_dist.len();
6161 let mean = m_dist.iter().sum::<usize>() as f64 / active as f64;
6162 let median = m_dist[active / 2];
6163 let max_m = *m_dist.last().unwrap();
6164 let min_m = m_dist[0];
6165 let above16 = m_dist.iter().filter(|&&x| x >= 16).count();
6166 println!("moe-grouped il={il} t={t} active={active}/{n_expert} \
6167 m_e: min={min_m} median={median} mean={mean:.1} max={max_m} \
6168 above_gemm_threshold(>=16)={above16}/{active}");
6169 }
6170
6171 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut moe_out)?;
6172 Ok(moe_out)
6173 }
6174
6175 pub(crate) fn moe_ffn_lockstep(
6182 &self,
6183 e: &Engine,
6184 m: &MoeWeights,
6185 zbatch: &CudaSlice<f32>,
6186 mrows: usize,
6187 il: u16,
6188 max_block: usize,
6189 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6190 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
6191 let cfg = &self.cfg;
6192 let moe = cfg.moe.as_ref().unwrap();
6193 let n_embd = cfg.n_embd as usize;
6194 let n_expert = moe.expert_count as usize;
6195 let n_used = moe.expert_used_count as usize;
6196 let n_ff_exp = moe.expert_ff_length as usize;
6197 let lim_exp = cfg.clamp_exp_at(il as u32);
6199 let lim_shexp = cfg.clamp_shexp_at(il as u32);
6200
6201 let logits = e.matmul(&m.gate_inp, zbatch, mrows)?;
6202 let (sel_all, w_all) = if let Some(sig) = cfg.sigmoid_router() {
6203 Self::moe_route_cfg(e, &logits, mrows, n_expert, n_used,
6204 m.exp_probs_b.as_deref(), Some(sig), m.active_experts.as_deref())?
6205 } else {
6206 Self::moe_route_cfg(e, &logits, mrows, n_expert, n_used,
6207 None, None, m.active_experts.as_deref())?
6208 };
6209 Self::trace_moe_routes(il, mrows, &sel_all, &w_all)?;
6210
6211 let resident_expert: Vec<bool> = e.with_moe_cache(max_block, |c, _| {
6213 Ok((0..n_expert)
6214 .map(|ex| {
6215 [PROJ_GATE, PROJ_UP, PROJ_DOWN].into_iter().all(|p| {
6216 c.resident(BlockId::new(il, p, ex as u16)).is_some()
6217 })
6218 })
6219 .collect())
6220 })?;
6221
6222 struct Group {
6223 rows: Vec<i32>,
6224 slots: Vec<i32>,
6225 weights: Vec<f32>,
6226 }
6227 let mut groups: std::collections::HashMap<usize, Group> = Default::default();
6228 let mut cpu_rows: Vec<Vec<(usize, f32)>> = vec![Vec::new(); mrows];
6229 let mut cpu_by_expert: std::collections::HashMap<usize, Vec<(usize, f32)>> =
6230 Default::default();
6231 for row in 0..mrows {
6232 for j in 0..n_used {
6233 let ex = sel_all[row * n_used + j] as usize;
6234 let w = w_all[row * n_used + j];
6235 if resident_expert[ex] {
6236 let group = groups.entry(ex).or_insert_with(|| Group {
6237 rows: Vec::new(),
6238 slots: Vec::new(),
6239 weights: Vec::new(),
6240 });
6241 group.rows.push(row as i32);
6242 group.slots.push(j as i32);
6243 group.weights.push(w);
6244 } else {
6245 crate::cpu_experts::record_incomplete_gpu_residency(0);
6246 cpu_rows[row].push((ex, w));
6247 cpu_by_expert.entry(ex).or_default().push((row, w));
6248 }
6249 }
6250 }
6251
6252 let host_rows = e.dtoh(zbatch)?;
6258 let rows_ok = crate::cpu_experts::rows_supported();
6259 enum CpuPart {
6260 Single { row: usize },
6261 Rows { rows: Vec<usize> },
6262 }
6263 let mut tickets: Vec<(CpuPart, crate::cpu_experts::CpuExpertTicket)> = Vec::new();
6264 let mut rows_served: std::collections::HashSet<(usize, usize)> = Default::default();
6265 if rows_ok {
6266 let mut shared: Vec<(usize, Vec<(usize, f32)>)> = cpu_by_expert
6267 .into_iter()
6268 .filter(|(_, rows)| rows.len() >= 2)
6269 .collect();
6270 shared.sort_by_key(|(ex, _)| *ex);
6271 for (ex, mut row_weights) in shared {
6272 row_weights.sort_by_key(|(row, _)| *row);
6273 let inputs: Vec<(&[f32], f32)> = row_weights
6274 .iter()
6275 .map(|&(row, w)| (&host_rows[row * n_embd..(row + 1) * n_embd], w))
6276 .collect();
6277 let job = crate::cpu_experts::prepare_rows_job(m, ex, &inputs)
6278 .map_err(std::io::Error::other)?;
6279 for &(row, _) in &row_weights {
6280 rows_served.insert((row, ex));
6281 }
6282 tickets.push((
6283 CpuPart::Rows {
6284 rows: row_weights.iter().map(|&(row, _)| row).collect(),
6285 },
6286 crate::cpu_experts::submit_rows(job).map_err(std::io::Error::other)?,
6287 ));
6288 }
6289 }
6290 for (row, selected) in cpu_rows.iter().enumerate() {
6291 let leftover: Vec<(usize, f32)> = selected
6292 .iter()
6293 .copied()
6294 .filter(|&(ex, _)| !rows_served.contains(&(row, ex)))
6295 .collect();
6296 if leftover.is_empty() {
6297 continue;
6298 }
6299 let host_row = &host_rows[row * n_embd..(row + 1) * n_embd];
6300 let job = crate::cpu_experts::prepare_job(m, il, &leftover, host_row)
6301 .map_err(std::io::Error::other)?;
6302 tickets.push((
6303 CpuPart::Single { row },
6304 crate::cpu_experts::submit(job).map_err(std::io::Error::other)?,
6305 ));
6306 }
6307
6308 let mut slot_buf = e.zeros(mrows * n_used * n_embd)?;
6309 let mut wbuf = e.zeros(mrows * n_used)?;
6310 let mut order: Vec<usize> = groups.keys().copied().collect();
6311 order.sort_by(|&a, &b| {
6312 groups[&b].rows.len().cmp(&groups[&a].rows.len()).then(a.cmp(&b))
6313 });
6314 for &ex in &order {
6315 let group = &groups[&ex];
6316 let m_e = group.rows.len();
6317 let gl = m.gate_exps.expert_layout(ex);
6318 let ul = m.up_exps.expert_layout(ex);
6319 let dl = m.down_exps.expert_layout(ex);
6320 let row_idx_d = e.htod_i32(&group.rows)?;
6321 let slot_idx_d = e.htod_i32(&group.slots)?;
6322 let dmac = m.down_exps.macro_scale(ex);
6323 let weight_d = if dmac == 1.0 {
6324 e.htod(&group.weights)?
6325 } else {
6326 let scaled: Vec<f32> = group.weights.iter().map(|&w| w * dmac).collect();
6327 e.htod(&scaled)?
6328 };
6329 let mut gathered = e.zeros(m_e * n_embd)?;
6330 e.gather_rows(zbatch, &row_idx_d, &mut gathered, n_embd, m_e)?;
6331 let gv = gathered.slice(0..m_e * n_embd);
6332 let gate = e.with_moe_cache(max_block, |c, eng| {
6333 let slot = c
6334 .resident(BlockId::new(il, PROJ_GATE, ex as u16))
6335 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
6336 eng.qmatvec_view(c.buf(crate::moe_cache::DispatchSlot::Resident(slot)), 0..gl.len, &gv, m_e,
6337 m.gate_exps.in_f, m.gate_exps.out_f, gl.qtype, gl.row_bytes)
6338 })?;
6339 let up = e.with_moe_cache(max_block, |c, eng| {
6340 let slot = c
6341 .resident(BlockId::new(il, PROJ_UP, ex as u16))
6342 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
6343 eng.qmatvec_view(c.buf(crate::moe_cache::DispatchSlot::Resident(slot)), 0..ul.len, &gv, m_e,
6344 m.up_exps.in_f, m.up_exps.out_f, ul.qtype, ul.row_bytes)
6345 })?;
6346 let mut act = e.zeros(m_e * n_ff_exp)?;
6347 Self::ffn_act_lim(e, cfg, &gate, &up, m.gate_exps.macro_scale(ex),
6348 m.up_exps.macro_scale(ex), lim_exp, &mut act, m_e * n_ff_exp)?;
6349 let actv = act.slice(0..m_e * n_ff_exp);
6350 let y = e.with_moe_cache(max_block, |c, eng| {
6351 let slot = c
6352 .resident(BlockId::new(il, PROJ_DOWN, ex as u16))
6353 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
6354 eng.qmatvec_view(c.buf(crate::moe_cache::DispatchSlot::Resident(slot)), 0..dl.len, &actv, m_e,
6355 m.down_exps.in_f, m.down_exps.out_f, dl.qtype, dl.row_bytes)
6356 })?;
6357 e.scatter_slot(&y, &row_idx_d, &slot_idx_d, &weight_d,
6358 &mut slot_buf, &mut wbuf, n_embd, n_used, m_e)?;
6359 }
6360 let mut moe_out = e.zeros(mrows * n_embd)?;
6361 e.reduce_slots(&slot_buf, &wbuf, &mut moe_out, n_embd, n_used, mrows)?;
6362
6363 let mut row_sums: Vec<Option<Vec<f32>>> = vec![None; mrows];
6365 for (part, ticket) in tickets {
6366 let cpu_output = ticket.wait().map_err(std::io::Error::other)?;
6367 let mut add_row = |row: usize, chunk: &[f32]| {
6368 let sum = row_sums[row].get_or_insert_with(|| vec![0.0f32; n_embd]);
6369 for (accumulator, value) in sum.iter_mut().zip(chunk) {
6370 *accumulator += value;
6371 }
6372 };
6373 match part {
6374 CpuPart::Single { row } => add_row(row, &cpu_output),
6375 CpuPart::Rows { rows } => {
6376 for (slot, row) in rows.into_iter().enumerate() {
6377 add_row(row, &cpu_output[slot * n_embd..(slot + 1) * n_embd]);
6378 }
6379 }
6380 }
6381 }
6382 for (row, sum) in row_sums.into_iter().enumerate() {
6383 let Some(sum) = sum else { continue };
6384 let cpu_output = e.htod(&sum)?;
6385 let mut dst = moe_out.slice_mut(row * n_embd..(row + 1) * n_embd);
6386 e.axpy_into(&cpu_output, 1.0, &mut dst, n_embd)?;
6387 }
6388
6389 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
6390 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
6391 {
6392 let n_ff_sh = gate_shexp.out_features();
6393 let sg_gate = e.matmul(gate_shexp, zbatch, mrows)?;
6394 let sg_up = e.matmul(up_shexp, zbatch, mrows)?;
6395 let mut sa = e.zeros(mrows * n_ff_sh)?;
6396 Self::ffn_act_lim(e, cfg, &sg_gate, &sg_up, 1.0, 1.0, lim_shexp,
6397 &mut sa, mrows * n_ff_sh)?;
6398 let sh = e.matmul(down_shexp, &sa, mrows)?;
6399 let g = match &m.gate_inp_shexp {
6402 Some(gate_inp_shexp) => {
6403 e.sigmoid_dot_rows(zbatch, gate_inp_shexp.float_data(), n_embd, mrows)?
6404 }
6405 None => e.htod(&vec![1.0f32; mrows])?,
6406 };
6407 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, mrows)?;
6408 }
6409
6410 Ok(moe_out)
6411 }
6412}
6413
6414impl HybridModel {
6420 pub(crate) fn gemma4_geom(&self, il: usize) -> (usize, usize, usize, f32, f32, bool) {
6422 let g = self.cfg.gemma4.as_ref().unwrap();
6423 let swa = g.swa_pattern[il];
6424 let hd = if swa { g.key_length_swa } else { g.key_length_global } as usize;
6425 (hd, g.head_count_kv[il] as usize, self.cfg.n_head as usize,
6429 if swa { g.rope_base_swa } else { g.rope_base_global },
6430 1.0, swa)
6431 }
6432
6433 fn gemma4_suppress(&self, e: &Engine, ld: &mut CudaSlice<f32>, t: usize)
6437 -> Result<(), Box<dyn std::error::Error>> {
6438 if let Some((ids, n)) = self.gemma4_aux.as_ref().and_then(|a| a.suppress_d.as_ref()) {
6439 e.mask_ids_rows(ld, ids, *n, self.output.out_features(), t)?;
6440 }
6441 Ok(())
6442 }
6443
6444 fn gemma4_attn_prime(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
6449 h: &CudaSlice<f32>, pos_d: &CudaSlice<i32>, t: usize,
6450 cache: Option<&mut Cache>)
6451 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6452 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
6453 let eps = self.cfg.rms_eps;
6454 let aux = self.gemma4_aux.as_ref().unwrap();
6455
6456 e.mmq_act_begin();
6459 let q0 = e.matmul(&fa.wq, h, t)?; let k0 = e.matmul(&fa.wk, h, t)?; let v0 = if swa { e.matmul(&fa.wv, h, t)? } else { e.clone_dtod(&k0)? };
6464
6465 let mut q = e.uninit(t * nh * hd)?;
6466 let mut k = e.uninit(t * nkv * hd)?;
6467 let mut v = e.uninit(t * nkv * hd)?;
6469 static EMIT: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6473 let emit = t >= 16 && crate::Engine::qkvnorm_w_on_prefill(nh * t + 2 * nkv * t, hd)
6474 && *EMIT.get_or_init(|| std::env::var("MEMRA_FA_EMIT").map(|s| s != "0").unwrap_or(true));
6475 let mut qb = e.alloc_uninit::<u8>(if emit { t * nh * hd * 2 } else { 1 })?;
6476 let mut kb = e.alloc_uninit::<u8>(if emit { t * nkv * hd * 2 } else { 1 })?;
6477 let mut vb = e.alloc_uninit::<u8>(if emit { t * nkv * hd * 2 } else { 1 })?;
6478 let v_f16 = emit && crate::fa_f16pv_on() && match hd {
6481 512 => true,
6482 256 => swa && crate::faw_hp_on() && nh % 2 == 0 && (nh / nkv) % 2 == 0,
6483 _ => false,
6484 };
6485 if emit {
6486 e.rms_norm_qkv_w4b(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
6487 &aux.ones, &mut q, &mut k, &mut v, &mut vb,
6488 hd, nh * t, nkv * t, eps, v_f16)?;
6489 } else {
6490 e.rms_norm_qkv(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
6491 &aux.ones, &mut q, &mut k, &mut v, hd, nh * t, nkv * t, eps)?;
6492 }
6493
6494 let ff = if swa { None } else {
6495 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
6496 };
6497 if emit {
6498 e.rope_neox2_bf16e(&mut q, &mut k, &mut qb, &mut kb, pos_d, hd, hd, nh, nkv, t,
6499 base, 1.0, ff)?;
6500 } else {
6501 e.rope_neox2(&mut q, &mut k, pos_d, hd, hd, nh, nkv, t, base, 1.0, ff)?;
6502 }
6503
6504 if let Some(cache) = cache {
6505 let kvl = cache.kv[il].as_mut().unwrap();
6506 assert_eq!(kvl.len, 0, "gemma4 prime is fresh-prompt only (v0)");
6507 e.append_kv_quantized_rows(&k, &v, &mut kvl.k, &mut kvl.v, kvl.len, t,
6508 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes, (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on()))?;
6509 kvl.len += t;
6510 }
6511 let mut attn = e.zeros(t * nh * hd)?;
6512 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
6516 if swa && t > win {
6517 if hd == 256 && std::env::var("MEMRA_NOFA").is_err() {
6518 if emit { e.fa_prefill_w_pre(&qb, &kb, &vb, &mut attn, hd, nh, nkv, t, t,
6519 scale, true, win, v_f16)?; }
6520 else { e.fa_prefill_w(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true,
6521 win)?; }
6522 } else {
6523 e.sdpa_naive_w(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
6524 }
6525 } else if hd == 256 && std::env::var("MEMRA_NOFA").is_err() {
6526 e.fa_prefill(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
6527 } else if hd == 512 && std::env::var("MEMRA_NOFA").is_err() {
6528 if emit { e.fa_prefill_hd512_pre(&qb, &kb, &vb, &mut attn, hd, nh, nkv, t, t,
6529 scale, true, v_f16)?; }
6530 else { e.fa_prefill_hd512(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?; }
6531 } else {
6532 e.sdpa_naive(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
6533 }
6534 Ok(e.matmul(&fa.wo, &attn, t)?)
6535 }
6536
6537 fn gemma4_attn(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
6539 h: &CudaSlice<f32>, pos_d: &CudaSlice<i32>, t: usize)
6540 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6541 self.gemma4_attn_prime(e, fa, il, h, pos_d, t, None)
6542 }
6543
6544 fn gemma4_moe_q8(&self, e: &Engine, m: &crate::hybrid::MoeWeights,
6549 bits: &crate::hybrid::Gemma4MoeBits,
6550 mq: &(CudaSlice<i8>, CudaSlice<f32>),
6551 router_in: &CudaSlice<f32>, t: usize)
6552 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6553 let cfg = &self.cfg;
6554 let moe = cfg.moe.as_ref().unwrap();
6555 let n_embd = cfg.n_embd as usize;
6556 let n_expert = moe.expert_count as usize;
6557 let n_used = moe.expert_used_count as usize;
6558 let n_ff_exp = moe.expert_ff_length as usize;
6559 let logits = if crate::router_kernel_on() {
6563 e.router_gemv(m.gate_inp.float_data(), router_in, n_embd, n_expert, t)?
6564 } else {
6565 e.matmul(&m.gate_inp, router_in, t)?
6566 };
6567 let dev = m.dev_exps.as_ref().unwrap();
6568 let (sel_d, w_d) = e.moe_router_topk_scaled(&logits, t, n_expert, n_used,
6569 &bits.per_expert_scale_d)?;
6570 let (zq, zd) = mq;
6571 if t == 1 {
6572 let selv = sel_d.slice(0..n_used);
6573 let wv = w_d.slice(0..n_used);
6574 let act = e.moe_gate_up_gelu8_dev_q8(&dev.ptr_row, &selv, zq, zd,
6575 n_embd, n_ff_exp, n_used, n_expert,
6576 m.gate_exps.qtype, m.up_exps.qtype,
6577 m.gate_exps.row_bytes, m.up_exps.row_bytes)?;
6578 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
6579 let mut moe_out = e.uninit(n_embd)?;
6580 e.moe_down8_fma_dev_q8(&dev.ptr_row, &selv, &wv, &aq2, &ad2,
6581 &mut moe_out.slice_mut(0..n_embd), n_ff_exp, n_embd,
6582 n_used, n_expert, m.down_exps.qtype, m.down_exps.row_bytes)?;
6583 return Ok(moe_out);
6584 }
6585 let csr = t <= 10 && std::env::var("MEMRA_GEMMA_CSR").as_deref() != Ok("0");
6586 let act = if csr {
6587 e.moe_gate_up_gelu8_dev_q8_csr(&dev.ptr_row, &sel_d, zq, zd, t * n_used,
6588 n_embd, n_ff_exp, n_used, n_expert,
6589 m.gate_exps.qtype, m.up_exps.qtype,
6590 m.gate_exps.row_bytes, m.up_exps.row_bytes)?
6591 } else {
6592 e.moe_gate_up_gelu8_dev_q8_rows(&dev.ptr_row, &sel_d, zq, zd, t,
6593 n_embd, n_ff_exp, n_used, n_expert,
6594 m.gate_exps.qtype, m.up_exps.qtype,
6595 m.gate_exps.row_bytes, m.up_exps.row_bytes)?
6596 };
6597 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
6598 let mut moe_out = e.uninit(t * n_embd)?;
6599 e.moe_down8_fma_dev_q8_rows_g(&dev.ptr_row, &sel_d, &w_d, &aq2, &ad2, &mut moe_out, t,
6602 n_ff_exp, n_embd, n_used, n_expert,
6603 m.down_exps.qtype, m.down_exps.row_bytes)?;
6604 Ok(moe_out)
6605 }
6606
6607 fn gemma4_moe(&self, e: &Engine, m: &crate::hybrid::MoeWeights,
6611 bits: &crate::hybrid::Gemma4MoeBits, moe_in: &CudaSlice<f32>,
6612 router_in: &CudaSlice<f32>, t: usize)
6613 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6614 let cfg = &self.cfg;
6615 let moe = cfg.moe.as_ref().unwrap();
6616 let n_embd = cfg.n_embd as usize;
6617 let n_expert = moe.expert_count as usize;
6618 let n_used = moe.expert_used_count as usize;
6619 let n_ff_exp = moe.expert_ff_length as usize;
6620
6621 let logits = if t < PRIME_MIN_T && crate::router_kernel_on() {
6625 e.router_gemv(m.gate_inp.float_data(), router_in, n_embd, n_expert, t)?
6626 } else {
6627 e.matmul(&m.gate_inp, router_in, t)?
6628 };
6629
6630 if t < PRIME_MIN_T && m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
6635 && expert_dp4a_supported(m.gate_exps.qtype) && expert_dp4a_supported(m.up_exps.qtype)
6636 && expert_dp4a_supported(m.down_exps.qtype)
6637 && std::env::var("MEMRA_GEMMA_MOE_FAST").as_deref() != Ok("0") {
6638 let dev = m.dev_exps.as_ref().unwrap();
6639 let (sel_d, w_d) = e.moe_router_topk_scaled(&logits, t, n_expert, n_used,
6640 &bits.per_expert_scale_d)?;
6641 if t == 1 {
6642 let (zq, zd) = e.quantize_q8_1(moe_in, 1, n_embd)?;
6643 let selv = sel_d.slice(0..n_used);
6644 let wv = w_d.slice(0..n_used);
6645 let act = e.moe_gate_up_gelu8_dev_q8(&dev.ptr_row, &selv, &zq, &zd,
6646 n_embd, n_ff_exp, n_used, n_expert,
6647 m.gate_exps.qtype, m.up_exps.qtype,
6648 m.gate_exps.row_bytes, m.up_exps.row_bytes)?;
6649 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
6650 let mut moe_out = e.uninit(n_embd)?;
6651 e.moe_down8_fma_dev_q8(&dev.ptr_row, &selv, &wv, &aq2, &ad2,
6652 &mut moe_out.slice_mut(0..n_embd), n_ff_exp, n_embd,
6653 n_used, n_expert, m.down_exps.qtype, m.down_exps.row_bytes)?;
6654 return Ok(moe_out);
6655 }
6656 let (zq, zd) = e.quantize_q8_1(moe_in, t, n_embd)?;
6661 let csr = t <= 10 && std::env::var("MEMRA_GEMMA_CSR").as_deref() != Ok("0");
6662 let act = if csr {
6663 e.moe_gate_up_gelu8_dev_q8_csr(&dev.ptr_row, &sel_d, &zq, &zd, t * n_used,
6664 n_embd, n_ff_exp, n_used, n_expert,
6665 m.gate_exps.qtype, m.up_exps.qtype,
6666 m.gate_exps.row_bytes, m.up_exps.row_bytes)?
6667 } else {
6668 e.moe_gate_up_gelu8_dev_q8_rows(&dev.ptr_row, &sel_d, &zq, &zd, t,
6669 n_embd, n_ff_exp, n_used, n_expert,
6670 m.gate_exps.qtype, m.up_exps.qtype,
6671 m.gate_exps.row_bytes, m.up_exps.row_bytes)?
6672 };
6673 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
6674 let mut moe_out = e.uninit(t * n_embd)?;
6675 e.moe_down8_fma_dev_q8_rows_g(&dev.ptr_row, &sel_d, &w_d, &aq2, &ad2, &mut moe_out, t,
6676 n_ff_exp, n_embd, n_used, n_expert,
6677 m.down_exps.qtype, m.down_exps.row_bytes)?;
6678 return Ok(moe_out);
6679 }
6680
6681 let (sel_all, mut w_all) = Self::moe_route(e, &logits, t, n_expert, n_used)?;
6682 for (i, &sx) in sel_all.iter().enumerate() {
6683 w_all[i] *= bits.per_expert_scale[sx as usize];
6684 }
6685
6686 if t >= PRIME_MIN_T && m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
6690 && expert_dp4a_supported(m.gate_exps.qtype) && expert_dp4a_supported(m.up_exps.qtype)
6691 && expert_dp4a_supported(m.down_exps.qtype)
6692 && std::env::var("MEMRA_GEMMA_MOE_PAIRS").as_deref() != Ok("0") {
6693 let dev = m.dev_exps.as_ref().unwrap();
6694 let n_pairs = t * n_used;
6695 let pair_ex: Vec<i32> = sel_all.iter().map(|&x| x as i32).collect();
6696 let pair_tok: Vec<i32> = (0..n_pairs).map(|p| (p / n_used) as i32).collect();
6697 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
6698 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
6699 let pt = e.htod_i32(&pair_tok)?;
6700 let pw = e.htod(&w_all)?;
6701 let toff = e.htod_i32(&tok_off)?;
6702 let tids = e.htod_i32(&tok_ids)?;
6703 let mut by_ex: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
6704 for p in 0..n_pairs { by_ex[pair_ex[p] as usize].push(p as i32); }
6705 let mut ex_ids: Vec<i32> = Vec::new();
6706 let mut ex_off: Vec<i32> = vec![0];
6707 let mut ex_pairs: Vec<i32> = Vec::with_capacity(n_pairs);
6708 for (ex, list) in by_ex.iter().enumerate() {
6709 if list.is_empty() { continue; }
6710 ex_ids.push(ex as i32);
6711 ex_pairs.extend_from_slice(list);
6712 ex_off.push(ex_pairs.len() as i32);
6713 }
6714 let n_active = ex_ids.len();
6715 let exi = e.htod_i32(&ex_ids)?;
6716 let exo = e.htod_i32(&ex_off)?;
6717 let exp_d = e.htod_i32(&ex_pairs)?;
6718 if crate::moe_f16g_gemma_on()
6726 && f16g_proj_ok(m.gate_exps.qtype, n_embd)
6727 && f16g_proj_ok(m.up_exps.qtype, n_embd)
6728 && f16g_proj_ok(m.down_exps.qtype, n_ff_exp) {
6729 let csr_tok: Vec<i32> = ex_pairs.iter().map(|&p| p / n_used as i32).collect();
6730 let csr_tok_d = e.htod_i32(&csr_tok)?;
6731 let (z_f16, z_s) = e.moe_f16g_act(moe_in, Some(&csr_tok_d), n_embd, n_pairs)?;
6732 let g_csr = e.moe_f16_grouped(&dev.ptr_row, 0, n_expert, &exi, &ex_off, &exo,
6733 &z_f16, &z_s, n_embd, n_ff_exp, n_active, n_pairs,
6734 m.gate_exps.qtype, m.gate_exps.row_bytes)?;
6735 let u_csr = e.moe_f16_grouped(&dev.ptr_row, 1, n_expert, &exi, &ex_off, &exo,
6736 &z_f16, &z_s, n_embd, n_ff_exp, n_active, n_pairs,
6737 m.up_exps.qtype, m.up_exps.row_bytes)?;
6738 let act_csr = e.moe_pairs_gelu_mul(&g_csr, &u_csr, n_pairs * n_ff_exp)?;
6739 let (a_f16, a_s) = e.moe_f16g_act(&act_csr, None, n_ff_exp, n_pairs)?;
6740 let d_csr = e.moe_f16_grouped(&dev.ptr_row, 2, n_expert, &exi, &ex_off, &exo,
6741 &a_f16, &a_s, n_ff_exp, n_embd, n_active, n_pairs,
6742 m.down_exps.qtype, m.down_exps.row_bytes)?;
6743 let y_down = e.rows_permute(&d_csr, &exp_d, n_pairs, n_embd)?;
6744 let mut moe_out = e.uninit(t * n_embd)?;
6745 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
6746 if std::env::var("MEMRA_F16G_DEBUG").is_ok() {
6747 let scan = |v: &[f32]| v.iter().filter(|x| !x.is_finite()).count();
6748 let (yd, mo) = (e.dtoh(&y_down)?, e.dtoh(&moe_out)?);
6749 eprintln!("[f16g-debug] post-permute bad={} post-scatter bad={}",
6750 scan(&yd), scan(&mo));
6751 }
6752 return Ok(moe_out);
6753 }
6754 let mma = n_embd % 256 == 0
6757 && std::env::var("MEMRA_GEMMA_MOE_MMA").as_deref() != Ok("0");
6758 let (gate, up) = if mma {
6759 let z_scr = e.mmq_iq_quantize_act(moe_in, n_embd, t)?;
6760 (e.mmq_iq_experts(&dev.ptr_row, 0, n_expert, &exi, &exo, &exp_d, &pt, &z_scr,
6761 n_embd, n_ff_exp, n_active, n_pairs, t,
6762 m.gate_exps.qtype, m.gate_exps.row_bytes)?,
6763 e.mmq_iq_experts(&dev.ptr_row, 1, n_expert, &exi, &exo, &exp_d, &pt, &z_scr,
6764 n_embd, n_ff_exp, n_active, n_pairs, t,
6765 m.up_exps.qtype, m.up_exps.row_bytes)?)
6766 } else {
6767 let (zq, zd) = e.quantize_q8_1(moe_in, t, n_embd)?;
6768 (e.moe_pairs_matvec_q8_dec(&dev.ptr_row, 0, &exi, &exo, &exp_d, &pt, &zq, &zd,
6769 n_embd, n_ff_exp, n_expert, n_active, n_pairs,
6770 m.gate_exps.qtype, m.gate_exps.row_bytes)?,
6771 e.moe_pairs_matvec_q8_dec(&dev.ptr_row, 1, &exi, &exo, &exp_d, &pt, &zq, &zd,
6772 n_embd, n_ff_exp, n_expert, n_active, n_pairs,
6773 m.up_exps.qtype, m.up_exps.row_bytes)?)
6774 };
6775 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
6776 let pself = e.htod_i32(&pair_self)?;
6777 let y_down = if mma {
6789 let in_pad = n_ff_exp.div_ceil(256) * 256;
6790 let a_scr = if crate::moe_fuse_actq_on() {
6791 e.mmq_iq_fused_act_quant(&gate, &up, n_ff_exp, n_pairs, 1)?
6792 } else {
6793 let act = e.moe_pairs_gelu_mul(&gate, &up, n_pairs * n_ff_exp)?;
6794 e.mmq_iq_quantize_act(&act, n_ff_exp, n_pairs)?
6795 };
6796 e.mmq_iq_experts(&dev.ptr_row, 2, n_expert, &exi, &exo, &exp_d, &pself, &a_scr,
6797 in_pad, n_embd, n_active, n_pairs, n_pairs,
6798 m.down_exps.qtype, m.down_exps.row_bytes)?
6799 } else {
6800 let act = e.moe_pairs_gelu_mul(&gate, &up, n_pairs * n_ff_exp)?;
6801 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
6802 e.moe_pairs_matvec_q8_dec(&dev.ptr_row, 2, &exi, &exo, &exp_d, &pself, &aq2, &ad2,
6803 n_ff_exp, n_embd, n_expert, n_active, n_pairs,
6804 m.down_exps.qtype, m.down_exps.row_bytes)?
6805 };
6806 let mut moe_out = e.uninit(t * n_embd)?;
6807 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
6808 return Ok(moe_out);
6809 }
6810
6811 let g_len = m.gate_exps.expert_stride;
6812 let u_len = m.up_exps.expert_stride;
6813 let d_len = m.down_exps.expert_stride;
6814 let dev = m.dev_exps.as_ref().filter(|d| !d.gu_il);
6818 let (mut sg, mut su, mut sd) = if dev.is_some() { (None, None, None) } else {
6819 (Some(e.alloc_u8_uninit(g_len)?), Some(e.alloc_u8_uninit(u_len)?), Some(e.alloc_u8_uninit(d_len)?))
6820 };
6821 let mut moe_out = e.zeros(t * n_embd)?;
6822 for tok in 0..t {
6823 let sel = &sel_all[tok * n_used..(tok + 1) * n_used];
6824 let w = &w_all[tok * n_used..(tok + 1) * n_used];
6825 let zt = moe_in.slice(tok * n_embd..(tok + 1) * n_embd);
6826 for (j, &ex) in sel.iter().enumerate() {
6827 let ex = ex as usize;
6828 let gate = match dev {
6829 Some(d) => e.qmatvec_view(&d.gate, ex * g_len..(ex + 1) * g_len, &zt, 1,
6830 m.gate_exps.in_f, m.gate_exps.out_f, m.gate_exps.qtype, m.gate_exps.row_bytes)?,
6831 None => {
6832 let sg = sg.as_mut().unwrap();
6833 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
6834 e.qmatvec_view(sg, 0..g_len, &zt, 1,
6835 m.gate_exps.in_f, m.gate_exps.out_f, m.gate_exps.qtype, m.gate_exps.row_bytes)?
6836 }
6837 };
6838 let up = match dev {
6839 Some(d) => e.qmatvec_view(&d.up, ex * u_len..(ex + 1) * u_len, &zt, 1,
6840 m.up_exps.in_f, m.up_exps.out_f, m.up_exps.qtype, m.up_exps.row_bytes)?,
6841 None => {
6842 let su = su.as_mut().unwrap();
6843 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
6844 e.qmatvec_view(su, 0..u_len, &zt, 1,
6845 m.up_exps.in_f, m.up_exps.out_f, m.up_exps.qtype, m.up_exps.row_bytes)?
6846 }
6847 };
6848 let mut act = e.uninit(n_ff_exp)?;
6849 e.gelu_tanh_mul(&gate, &up, &mut act, n_ff_exp)?;
6850 let actv = act.slice(0..n_ff_exp);
6851 let y = match dev {
6852 Some(d) => e.qmatvec_view(&d.down, ex * d_len..(ex + 1) * d_len, &actv, 1,
6853 m.down_exps.in_f, m.down_exps.out_f, m.down_exps.qtype, m.down_exps.row_bytes)?,
6854 None => {
6855 let sd = sd.as_mut().unwrap();
6856 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
6857 e.qmatvec_view(sd, 0..d_len, &actv, 1,
6858 m.down_exps.in_f, m.down_exps.out_f, m.down_exps.qtype, m.down_exps.row_bytes)?
6859 }
6860 };
6861 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
6862 e.axpy_into(&y, w[j], &mut dst, n_embd)?;
6863 }
6864 }
6865 Ok(moe_out)
6866 }
6867
6868 fn gemma4_layer(&self, e: &Engine, il: usize, layer: &crate::hybrid::HybridLayer,
6870 x: &CudaSlice<f32>, pos_d: &CudaSlice<i32>, t: usize)
6871 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6872 let n_embd = self.cfg.n_embd as usize;
6873 let eps = self.cfg.rms_eps;
6874
6875 let mut h = e.zeros(t * n_embd)?;
6876 e.rms_norm(x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
6877 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
6878 let o = self.gemma4_attn(e, fa, il, &h, pos_d, t)?;
6879 let mut cur = e.zeros(t * n_embd)?;
6881 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, t, eps)?;
6882 self.gemma4_layer_tail_add(e, layer, &cur, x, t)
6883 }
6884
6885 fn gemma4_layer_tail_add(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
6889 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize)
6890 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6891 Ok(self.gemma4_layer_tail_add_n(e, layer, cur, x, t, None)?.0)
6892 }
6893
6894 fn gemma4_layer_tail_add_n(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
6897 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize,
6898 next_norm: Option<&CudaSlice<f32>>)
6899 -> Result<(CudaSlice<f32>, Option<CudaSlice<f32>>), Box<dyn std::error::Error>> {
6900 let n_embd = self.cfg.n_embd as usize;
6901 let bits = layer.gemma4.as_ref().unwrap();
6902 let (sn, attn_out) = self.gemma4_layer_tail_core(e, layer, cur, x, t)?;
6903 let mut xn = e.uninit(t * n_embd)?;
6904 match next_norm {
6905 Some(w) => {
6906 let mut hn = e.uninit(t * n_embd)?;
6907 e.add_scale_rms_norm(&sn, &attn_out, bits.layer_scale, w, &mut xn, &mut hn,
6908 n_embd, t, self.cfg.rms_eps)?;
6909 Ok((xn, Some(hn)))
6910 }
6911 None => {
6912 e.add_scale(&sn, &attn_out, bits.layer_scale, &mut xn, t * n_embd)?;
6913 Ok((xn, None))
6914 }
6915 }
6916 }
6917
6918 fn gemma4_layer_tail_core(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
6921 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize)
6922 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6923 self.gemma4_layer_tail_core_pn(e, layer, cur, x, t, None, false)
6924 }
6925
6926 fn gemma4_layer_tail_core_pn(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
6933 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize,
6934 pre_norm: Option<&CudaSlice<f32>>, defer_post_norm: bool)
6935 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6936 let n_embd = self.cfg.n_embd as usize;
6937 let eps = self.cfg.rms_eps;
6938 let bits = layer.gemma4.as_ref().unwrap();
6939
6940 let Some(mbits) = bits.moe_bits.as_ref() else {
6943 let crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } = &layer.ffn
6944 else { panic!("gemma4 dense layer without Dense ffn") };
6945 let mut attn_out = e.uninit(t * n_embd)?;
6946 let mut zsh = e.uninit(t * n_embd)?;
6947 let mut zpair: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
6950 match pre_norm {
6951 Some(wa) if t == 1 => {
6952 zpair = Some(e.rms_pre_add_rms_norm_q8z(cur, wa, x,
6953 bits.ffn_norm.float_data(),
6954 &mut attn_out, &mut zsh,
6955 n_embd, t, eps)?);
6956 }
6957 Some(wa) => e.rms_pre_add_rms_norm(cur, wa, x, bits.ffn_norm.float_data(),
6958 &mut attn_out, &mut zsh, n_embd, t, eps)?,
6959 None => e.add_rms_norm(cur, x, bits.ffn_norm.float_data(), &mut attn_out,
6960 &mut zsh, n_embd, t, eps)?,
6961 }
6962 let n_ff = ffn_gate.out_features();
6963 let (gate, up) = if t == 1 {
6969 let (zq, zd) = match zpair {
6970 Some(p) => p,
6971 None => e.quantize_q8_1(&zsh, 1, n_embd)?,
6972 };
6973 match e.matmul_q4_fused2(ffn_gate, ffn_up, &zq, &zd)? {
6974 Some(p) => p,
6975 None => (e.matmul_pre(ffn_gate, &zq, &zd, &zsh, 1)?,
6976 e.matmul_pre(ffn_up, &zq, &zd, &zsh, 1)?),
6977 }
6978 } else {
6979 static F2B: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6984 let f2b = *F2B.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
6985 let fused = if f2b {
6986 let (zq, zd) = e.quantize_q8_1(&zsh, t, n_embd)?;
6987 e.matmul_q4_fused2_batched(ffn_gate, ffn_up, &zq, &zd, t)?
6988 } else { None };
6989 match fused {
6990 Some(p) => p,
6991 None => {
6992 e.mmq_act_begin();
6994 (e.matmul(ffn_gate, &zsh, t)?, e.matmul(ffn_up, &zsh, t)?)
6995 }
6996 }
6997 };
6998 let mut act = e.uninit(t * n_ff)?;
6999 let f0 = if e.uses_q8_1_fast(ffn_down) {
7002 let upv = e.view(&up, t * n_ff);
7003 let up_all = upv.slice(0..t * n_ff);
7004 let (aq, ad) = e.gelu_tanh_mul_q8_1(&gate, &up_all, &mut act, n_ff, t)?;
7005 e.matmul_pre(ffn_down, &aq, &ad, &act, t)?
7006 } else {
7007 e.gelu_tanh_mul(&gate, &up, &mut act, t * n_ff)?;
7008 e.matmul(ffn_down, &act, t)?
7009 };
7010 if defer_post_norm { return Ok((f0, attn_out)); }
7011 let mut sn = e.uninit(t * n_embd)?;
7012 e.rms_norm(&f0, bits.post_ffw_norm.float_data(), &mut sn, n_embd, t, eps)?;
7013 return Ok((sn, attn_out));
7014 };
7015
7016 assert!(pre_norm.is_none(), "pre-norm fold is dense-entry only");
7017 let mut attn_out = e.uninit(t * n_embd)?;
7022 let mut router_in = e.uninit(t * n_embd)?;
7023 let fast_moe = match &layer.ffn {
7024 crate::hybrid::Ffn::Moe(m) => m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
7025 && expert_dp4a_supported(m.gate_exps.qtype)
7026 && expert_dp4a_supported(m.up_exps.qtype)
7027 && expert_dp4a_supported(m.down_exps.qtype)
7028 && std::env::var("MEMRA_GEMMA_MOE_FAST").as_deref() != Ok("0"),
7029 _ => false,
7030 };
7031 let q8z = t < PRIME_MIN_T && fast_moe;
7032 let (zsh_f32, zsh_q8, moe_q8) = if q8z {
7033 let (z0, m2) = e.add_rms_norm3_q8z(cur, x, bits.ffn_norm.float_data(),
7034 &mbits.router_scale_pre,
7035 mbits.pre_ffw_norm_2.float_data(),
7036 &mut attn_out, &mut router_in, n_embd, t, eps)?;
7037 (None, Some(z0), Some(m2))
7038 } else {
7039 let mut zsh = e.uninit(t * n_embd)?;
7040 let mut moe_in = e.uninit(t * n_embd)?;
7041 e.add_rms_norm3(cur, x, bits.ffn_norm.float_data(), &mbits.router_scale_pre,
7042 mbits.pre_ffw_norm_2.float_data(), &mut attn_out, &mut zsh,
7043 &mut router_in, &mut moe_in, n_embd, t, eps)?;
7044 (Some((zsh, moe_in)), None, None)
7045 };
7046 let attn_out2 = attn_out;
7047 #[allow(unused_variables)]
7048 let attn_out = &attn_out2;
7049 let n_ff = mbits.shared_gate.out_features();
7050 let (gate, up) = if let Some((zq, zd)) = zsh_q8.as_ref() {
7051 if t == 1 {
7052 match e.matmul_q4_fused2(&mbits.shared_gate, &mbits.shared_up, zq, zd)? {
7053 Some(p) => p,
7054 None => {
7055 let h0 = e.zeros(0)?;
7056 (e.matmul_pre(&mbits.shared_gate, zq, zd, &h0, 1)?,
7057 e.matmul_pre(&mbits.shared_up, zq, zd, &h0, 1)?)
7058 }
7059 }
7060 } else {
7061 let h0 = e.zeros(0)?;
7063 (e.matmul_pre(&mbits.shared_gate, zq, zd, &h0, t)?,
7064 e.matmul_pre(&mbits.shared_up, zq, zd, &h0, t)?)
7065 }
7066 } else {
7067 let (zsh, _) = zsh_f32.as_ref().unwrap();
7068 (e.matmul(&mbits.shared_gate, zsh, t)?, e.matmul(&mbits.shared_up, zsh, t)?)
7069 };
7070 let mut act = e.uninit(t * n_ff)?;
7071 e.gelu_tanh_mul(&gate, &up, &mut act, t * n_ff)?;
7072 let mlp0 = e.matmul(&mbits.shared_down, &act, t)?;
7073 let crate::hybrid::Ffn::Moe(m) = &layer.ffn else { panic!("gemma4 layer not MoE") };
7074 let moe0 = match (&moe_q8, &zsh_f32) {
7075 (Some(mq), _) => self.gemma4_moe_q8(e, m, mbits, mq, &router_in, t)?,
7076 (None, Some((_, moe_in))) => self.gemma4_moe(e, m, mbits, moe_in, &router_in, t)?,
7077 _ => unreachable!(),
7078 };
7079 let mut mlp = e.uninit(t * n_embd)?;
7081 let mut moe = e.uninit(t * n_embd)?;
7082 e.rms_norm2x(&mlp0, &moe0, mbits.post_ffw_norm_1.float_data(),
7083 mbits.post_ffw_norm_2.float_data(), &mut mlp, &mut moe, n_embd, t, eps)?;
7084
7085 let mut sum = e.uninit(t * n_embd)?;
7088 let mut sn = e.uninit(t * n_embd)?;
7089 e.add_rms_norm(&mlp, &moe, bits.post_ffw_norm.float_data(), &mut sum, &mut sn,
7090 n_embd, t, eps)?;
7091 Ok((sn, attn_out2))
7092 }
7093
7094 fn gemma4_layer_tail_add_nq(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
7096 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize,
7097 next_norm: Option<&CudaSlice<f32>>)
7098 -> Result<(CudaSlice<f32>, Option<(CudaSlice<i8>, CudaSlice<f32>)>), Box<dyn std::error::Error>> {
7099 let n_embd = self.cfg.n_embd as usize;
7100 let bits = layer.gemma4.as_ref().unwrap();
7101 let (sn, attn_out) = self.gemma4_layer_tail_core(e, layer, cur, x, t)?;
7102 let mut xn = e.uninit(t * n_embd)?;
7103 match next_norm {
7104 Some(w) => {
7105 let pair = e.add_scale_rms_norm_q8_1(&sn, &attn_out, bits.layer_scale, w, &mut xn,
7106 n_embd, t, self.cfg.rms_eps)?;
7107 Ok((xn, Some(pair)))
7108 }
7109 None => {
7110 e.add_scale(&sn, &attn_out, bits.layer_scale, &mut xn, t * n_embd)?;
7111 Ok((xn, None))
7112 }
7113 }
7114 }
7115
7116 fn gemma4_forward(&self, e: &Engine, tokens: &[u32], last_only: bool)
7119 -> Result<Vec<f32>, Box<dyn std::error::Error>> {
7120 if self.is_gemma4_e4b() { return self.gemma4_e4b_forward(e, tokens, last_only); }
7123 let n_embd = self.cfg.n_embd as usize;
7124 let t = tokens.len();
7125 let pos: Vec<i32> = (0..t as i32).collect();
7126 let pos_d = e.htod_i32(&pos)?;
7127
7128 let mut x = self.embed(e, tokens)?;
7129 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
7130 let probe = std::env::var("MEMRA_GEMMA_PROBE").is_ok();
7133 let stat = |e: &Engine, x: &CudaSlice<f32>, tag: &str| -> Result<(), Box<dyn std::error::Error>> {
7134 let h = e.dtoh(x)?;
7135 let bad = h.iter().filter(|v| !v.is_finite()).count();
7136 let mx = h.iter().filter(|v| v.is_finite()).fold(0.0f32, |m, v| m.max(v.abs()));
7137 eprintln!("[gemma-probe] {tag}: tok0_first3={:?} bad={bad} max={mx:.3e}", &h[..3]);
7138 Ok(())
7139 };
7140 if probe { stat(e, &x, "embed")?; }
7141 for (il, layer) in self.layers.iter().enumerate() {
7142 x = self.gemma4_layer(e, il, layer, &x, &pos_d, t)?;
7143 if probe { stat(e, &x, &format!("L{il}"))?; }
7144 }
7145 let mut hn = e.zeros(t * n_embd)?;
7146 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, self.cfg.rms_eps)?;
7147 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
7148 let n_vocab = self.output.out_features();
7149 let logits = if last_only {
7150 let hv = e.view(&hn, t * n_embd);
7151 let last_row = hv.slice((t - 1) * n_embd..t * n_embd);
7152 let mut hlast = e.zeros(n_embd)?;
7153 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
7154 let mut ld = e.matmul(&self.output, &hlast, 1)?;
7155 e.softcap(&mut ld, cap, n_vocab)?;
7156 self.gemma4_suppress(e, &mut ld, 1)?;
7157 e.dtoh(&ld)?
7158 } else {
7159 let mut ld = e.matmul(&self.output, &hn, t)?;
7160 e.softcap(&mut ld, cap, t * n_vocab)?;
7161 self.gemma4_suppress(e, &mut ld, t)?;
7162 e.dtoh(&ld)?
7163 };
7164 Ok(logits)
7165 }
7166
7167 pub(crate) fn gemma4_prime(&self, e: &Engine, tokens: &[u32], cache: &mut Cache)
7172 -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7173 if cache.pos != 0 {
7178 return Err("gemma4 prime v0 is fresh-prompt only (no continuation/chunked prime) \
7179 — prime the full prompt in one call or decode tokenwise".into());
7180 }
7181 let n_embd = self.cfg.n_embd as usize;
7182 let eps = self.cfg.rms_eps;
7183 let t = tokens.len();
7184 let pos: Vec<i32> = (0..t as i32).collect();
7185 let pos_d = e.htod_i32(&pos)?;
7186 let mut x = self.embed(e, tokens)?;
7187 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
7188 for (il, layer) in self.layers.iter().enumerate() {
7189 let mut h = e.zeros(t * n_embd)?;
7190 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
7191 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer not full-attn") };
7192 let o = self.gemma4_attn_prime(e, fa, il, &h, &pos_d, t, Some(cache))?;
7193 let mut cur = e.zeros(t * n_embd)?;
7194 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, t, eps)?;
7195 x = self.gemma4_layer_tail_add(e, layer, &cur, &x, t)?;
7196 self.dflash_tap(e, cache, il, &x, t)?;
7197 }
7198 cache.pos += t;
7199 let hiddens = e.clone_dtod(&x)?;
7200 let xv = e.view(&x, t * n_embd);
7201 let last_row = xv.slice((t - 1) * n_embd..t * n_embd);
7202 let mut h_seed = e.zeros(n_embd)?;
7203 e.copy_view_into(&mut h_seed, 0, &last_row, n_embd)?;
7204 let mut hn = e.uninit(n_embd)?;
7205 e.rms_norm(&h_seed, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
7206 let mut ld = e.matmul(&self.output, &hn, 1)?;
7207 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
7208 e.softcap(&mut ld, cap, self.output.out_features())?;
7209 self.gemma4_suppress(e, &mut ld, 1)?;
7210 let logits = e.dtoh(&ld)?;
7211 Ok((logits, h_seed, hiddens))
7212 }
7213
7214 fn gemma4_decode_attn(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
7219 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
7220 pos_d: &CudaSlice<i32>, cache: &mut Cache)
7221 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7222 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
7223 let eps = self.cfg.rms_eps;
7224 let aux = self.gemma4_aux.as_ref().unwrap();
7225 let (hq, hdq) = (hq, hdq);
7226 let h0 = e.zeros(0)?;
7227 let h = &h0;
7228 let (q0, k0, v0) = if swa {
7229 match e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, &hq, &hdq)? {
7230 Some(t3) => t3,
7231 None => (e.matmul_pre(&fa.wq, &hq, &hdq, h, 1)?,
7232 e.matmul_pre(&fa.wk, &hq, &hdq, h, 1)?,
7233 e.matmul_pre(&fa.wv, &hq, &hdq, h, 1)?),
7234 }
7235 } else {
7236 let (q0, k0) = match e.matmul_q4_fused2(&fa.wq, &fa.wk, &hq, &hdq)? {
7237 Some(p) => p,
7238 None => (e.matmul_pre(&fa.wq, &hq, &hdq, h, 1)?,
7239 e.matmul_pre(&fa.wk, &hq, &hdq, h, 1)?),
7240 };
7241 let v0 = e.clone_dtod(&k0)?;
7242 (q0, k0, v0)
7243 };
7244 let mut q = e.uninit(nh * hd)?;
7245 let mut k = e.uninit(nkv * hd)?;
7246 let mut v = e.uninit(nkv * hd)?;
7247 let ff = if swa { None } else {
7250 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
7251 };
7252 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
7253 &aux.ones, &mut q, &mut k, &mut v, hd, nh, nkv,
7254 pos_d, nh, nkv, base, 1.0, ff, eps)?;
7255 let kvl = cache.kv[il].as_mut().unwrap();
7256 e.append_kv_quantized(&k, &v, &mut kvl.k, &mut kvl.v, kvl.len,
7257 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes, (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on()))?;
7258 kvl.len += 1;
7259 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
7263 let mut attn = e.uninit(nh * hd)?;
7264 if !swa && hd == 512 && kvl.len >= crate::fa512_min_tkv()
7266 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
7267 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
7268 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
7269 let base = kvl.len as i32;
7271 e.i32_set_k(&mut kvl.len_d, base)?;
7272 e.fa_decode_rows(&q, &kp, &vp, &mut attn, hd, nh, nkv, kvl.len - 1, 1, scale,
7273 kvl.k_tok_bytes, kvl.v_tok_bytes, Some((&kvl.len_d, -1)), false,
7274 false, None)?;
7275 return Ok(e.matmul(&fa.wo, &attn, 1)?);
7276 }
7277 if swa && kvl.len > win && hd == 256
7279 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
7280 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
7281 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
7282 let base = kvl.len as i32;
7283 e.i32_set_k(&mut kvl.len_d, base)?;
7284 e.fa_decode_rows_w(&q, &kp, &vp, &mut attn, hd, nh, nkv, &kvl.len_d, -1, 1, scale,
7285 win, kvl.k_tok_bytes, kvl.v_tok_bytes, None)?;
7286 return Ok(e.matmul(&fa.wo, &attn, 1)?);
7287 }
7288 let (off_tok, t_kv) = if swa && kvl.len > win { (kvl.len - win, win) } else { (0, kvl.len) };
7289 let k_view = e.view_u8_range(&kvl.k, off_tok * kvl.k_tok_bytes,
7290 (off_tok + t_kv) * kvl.k_tok_bytes);
7291 let v_view = e.view_u8_range(&kvl.v, off_tok * kvl.v_tok_bytes,
7292 (off_tok + t_kv) * kvl.v_tok_bytes);
7293 e.fa_decode_kvmod(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, t_kv, scale,
7294 kvl.k_tok_bytes, kvl.v_tok_bytes, swa && crate::Engine::wkv_on())?;
7295 Ok(e.matmul(&fa.wo, &attn, 1)?)
7296 }
7297
7298 #[allow(clippy::too_many_arguments)]
7305 pub fn gemma4_decode_step_dc(&self, e: &Engine, token_d: &CudaSlice<u32>,
7306 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
7307 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
7308 n_vocab: usize, cap_bucket_max: Option<(usize, usize)>)
7309 -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
7310 let mut tok_out = e.stream().alloc_zeros::<u32>(1)?;
7311 self.gemma4_decode_step_dc_into(e, token_d, pos_d, embd_gpu, embd_qt, embd_rb, cache,
7312 n_vocab, cap_bucket_max, &mut tok_out)?;
7313 Ok(tok_out)
7314 }
7315
7316 #[allow(clippy::too_many_arguments)]
7319 pub fn gemma4_decode_step_dc_into(&self, e: &Engine, token_d: &CudaSlice<u32>,
7320 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
7321 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
7322 n_vocab: usize, cap_bucket_max: Option<(usize, usize)>,
7323 tok_out: &mut CudaSlice<u32>)
7324 -> Result<(), Box<dyn std::error::Error>> {
7325 let n_embd = self.cfg.n_embd as usize;
7326 let eps = self.cfg.rms_eps;
7327 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
7328 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
7329 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
7330 let n_layers = self.layers.len();
7331 for (il, layer) in self.layers.iter().enumerate() {
7332 let (hq, hdq) = match h_carry.take() {
7333 Some(p) => p,
7334 None => e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, 1, eps)?,
7335 };
7336 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
7337 let o = self.gemma4_decode_attn_dc(e, fa, il, &hq, &hdq, pos_d, cache, cap_bucket_max)?;
7338 let mut cur = e.uninit(n_embd)?;
7339 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, 1, eps)?;
7340 let next_norm = if il + 1 < n_layers {
7341 Some(self.layers[il + 1].attn_norm.float_data())
7342 } else { None };
7343 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, 1, next_norm)?;
7344 x = xn;
7345 h_carry = hn;
7346 }
7347 let mut hn = e.uninit(n_embd)?;
7348 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
7349 let mut logits = e.matmul(&self.output, &hn, 1)?;
7350 self.gemma4_suppress(e, &mut logits, 1)?; e.argmax_token_device_into(&logits, tok_out, n_vocab)?;
7352 e.inc_seqlen(pos_d)?;
7353 if cap_bucket_max.is_none() { cache.pos += 1; }
7354 Ok(())
7355 }
7356
7357 pub fn g4_dc_slots(&self, e: &Engine) -> Result<G4DcSlots, Box<dyn std::error::Error>> {
7364 let n_embd = self.cfg.n_embd as usize;
7365 let n_vocab = self.output.out_features();
7366 let n_layers = self.layers.len();
7367 let (mut qmax, mut kvmax, mut ffmax) = (0usize, 0usize, 0usize);
7368 for il in 0..n_layers {
7369 let (hd, nkv, nh, _b, _s, _w) = self.gemma4_geom(il);
7370 qmax = qmax.max(nh * hd);
7371 kvmax = kvmax.max(nkv * hd);
7372 if let crate::hybrid::Ffn::Dense { ffn_gate, .. } = &self.layers[il].ffn {
7373 ffmax = ffmax.max(ffn_gate.out_features());
7374 }
7375 }
7376 Ok(G4DcSlots {
7377 x: e.uninit(n_embd)?, xn: e.uninit(n_embd)?, cur: e.uninit(n_embd)?,
7378 hq: e.alloc_i8_uninit(n_embd)?, hd_: e.uninit(n_embd / 32)?,
7379 q0: e.uninit(qmax)?, k0: e.uninit(kvmax)?, v0: e.uninit(kvmax)?,
7380 q: e.uninit(qmax)?, k: e.uninit(kvmax)?, v: e.uninit(kvmax)?,
7381 attn: e.uninit(qmax)?, o: e.uninit(n_embd)?,
7382 attn_out: e.uninit(n_embd)?, zsh: e.uninit(n_embd)?,
7383 zq: e.alloc_i8_uninit(n_embd.max(qmax))?, zd: e.uninit(n_embd.max(qmax) / 32)?,
7386 gate: e.uninit(ffmax)?, up: e.uninit(ffmax)?,
7387 act: e.uninit(ffmax)?, actq: e.alloc_i8_uninit(ffmax)?, actd: e.uninit(ffmax / 32)?,
7388 f0: e.uninit(n_embd)?, sn: e.uninit(n_embd)?,
7389 hn: e.uninit(n_embd)?, logits: e.uninit(n_vocab)?,
7390 })
7391 }
7392
7393 fn g4_matvec_m1_into(&self, e: &Engine, w: &crate::model::GpuTensor,
7396 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, y: &mut CudaSlice<f32>)
7397 -> Result<(), Box<dyn std::error::Error>> {
7398 use crate::model::GpuTensor;
7399 let (bytes, qtype, row_bytes, scale, rp) = match w {
7400 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } =>
7401 (bytes, *qtype, *row_bytes, *scale, *rp),
7402 _ => return Err("g4_matvec_m1_into: non-quant tensor".into()),
7403 };
7404 let (mbytes, mrp) = match w {
7405 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
7406 _ => (bytes, rp),
7407 };
7408 e.qmatvec_mmvq_into(mbytes, aq, ad, 1, w.in_features(), w.out_features(),
7409 qtype, row_bytes, scale, mrp, y)
7410 }
7411
7412 #[allow(clippy::too_many_arguments)]
7416 pub fn gemma4_decode_step_dc_slotted(&self, e: &Engine, token_d: &CudaSlice<u32>,
7417 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
7418 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
7419 n_vocab: usize, cap_bucket_max: Option<(usize, usize)>,
7420 sl: &mut G4DcSlots, tok_out: &mut CudaSlice<u32>,
7421 ring: Option<(&mut CudaSlice<u32>, usize)>)
7422 -> Result<(), Box<dyn std::error::Error>> {
7423 let n_embd = self.cfg.n_embd as usize;
7424 let eps = self.cfg.rms_eps;
7425 e.embed_gather_device_into(embd_gpu, token_d, &mut sl.x, n_embd, embd_qt, embd_rb)?;
7426 e.scale_inplace(&mut sl.x, (n_embd as f32).sqrt(), n_embd)?;
7427 let n_layers = self.layers.len();
7428 let mut has_carry = false;
7429 for il in 0..n_layers {
7430 if !has_carry {
7431 e.rms_norm_q8_1_into(&sl.x, self.layers[il].attn_norm.float_data(), n_embd, 1,
7432 eps, &mut sl.hq, &mut sl.hd_)?;
7433 }
7434 has_carry = true;
7435 let layer = &self.layers[il];
7436 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
7437 self.gemma4_decode_attn_dc_slotted(e, fa, il, pos_d, cache, cap_bucket_max, sl)?;
7438 e.rms_norm(&sl.o, layer.post_attn_norm.float_data(), &mut sl.cur, n_embd, 1, eps)?;
7439 let next_norm = if il + 1 < n_layers {
7440 Some(self.layers[il + 1].attn_norm.float_data())
7441 } else { None };
7442 self.gemma4_layer_tail_slotted(e, layer, next_norm, sl)?;
7443 std::mem::swap(&mut sl.x, &mut sl.xn);
7444 }
7445 e.rms_norm(&sl.x, self.output_norm.float_data(), &mut sl.hn, n_embd, 1, eps)?;
7446 e.quantize_q8_1_into(&sl.hn, 1, n_embd, &mut sl.zq, &mut sl.zd)?;
7447 {
7449 let (zq, zd) = (&sl.zq, &sl.zd);
7450 let zq = unsafe { &*(zq as *const CudaSlice<i8>) };
7451 let zd = unsafe { &*(zd as *const CudaSlice<f32>) };
7452 self.g4_matvec_m1_into(e, &self.output, zq, zd, &mut sl.logits)?;
7453 }
7454 self.gemma4_suppress(e, &mut sl.logits, 1)?;
7455 e.argmax_token_device_into(&sl.logits, tok_out, n_vocab)?;
7456 if let Some((ring, base)) = ring {
7457 e.plain_tok_ring(tok_out, pos_d, base, ring)?;
7461 }
7462 e.inc_seqlen(pos_d)?;
7463 if cap_bucket_max.is_none() { cache.pos += 1; }
7464 Ok(())
7465 }
7466
7467 #[allow(clippy::too_many_arguments)]
7469 fn gemma4_decode_attn_dc_slotted(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer,
7470 il: usize, pos_d: &CudaSlice<i32>, cache: &mut Cache,
7471 cap_bucket_max: Option<(usize, usize)>, sl: &mut G4DcSlots)
7472 -> Result<(), Box<dyn std::error::Error>> {
7473 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
7474 let eps = self.cfg.rms_eps;
7475 let aux = self.gemma4_aux.as_ref().unwrap();
7476 {
7477 let hq = unsafe { &*(&sl.hq as *const CudaSlice<i8>) };
7478 let hdq = unsafe { &*(&sl.hd_ as *const CudaSlice<f32>) };
7479 if swa {
7480 if !e.matmul_q4_fused3_into(&fa.wq, &fa.wk, &fa.wv, hq, hdq,
7481 &mut sl.q0, &mut sl.k0, &mut sl.v0)? {
7482 return Err("slotted step: fused3 unavailable (non-uniform trunk)".into());
7483 }
7484 } else {
7485 if !e.matmul_q4_fused2_into(&fa.wq, &fa.wk, hq, hdq, &mut sl.q0, &mut sl.k0)? {
7486 return Err("slotted step: fused2 unavailable".into());
7487 }
7488 let k0r = unsafe { &*(&sl.k0 as *const CudaSlice<f32>) };
7489 e.copy_into(&mut sl.v0, 0, k0r, nkv * hd)?;
7490 }
7491 }
7492 let ff = if swa { None } else {
7495 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
7496 };
7497 let kvl = cache.kv[il].as_mut().unwrap();
7498 let kv_fp8 = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
7499 if crate::Engine::qkv_append_on() {
7500 e.rms_norm_qkv_rope_append_dc(&sl.q0, &sl.k0, &sl.v0, fa.q_norm.float_data(),
7502 fa.k_norm.float_data(), &aux.ones, &mut sl.q, &mut sl.k, &mut sl.v, hd, nh, nkv,
7503 pos_d, nh, nkv, base, 1.0, ff, eps,
7504 &mut kvl.k, &mut kvl.v, &kvl.len_d, kvl.k_tok_bytes, kvl.v_tok_bytes, kv_fp8)?;
7505 } else {
7506 e.rms_norm_qkv_rope(&sl.q0, &sl.k0, &sl.v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
7507 &aux.ones, &mut sl.q, &mut sl.k, &mut sl.v, hd, nh, nkv,
7508 pos_d, nh, nkv, base, 1.0, ff, eps)?;
7509 e.append_kv_quantized_dc(&sl.k, &sl.v, &mut kvl.k, &mut kvl.v, &kvl.len_d,
7510 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes,
7511 kv_fp8)?;
7512 }
7513 e.inc_seqlen(&mut kvl.len_d)?;
7514 let (b_swa, b_glob) = cap_bucket_max.expect("slotted step is capture-only");
7515 let k_view = e.view_u8(&kvl.k, kvl.k.len());
7516 let v_view = e.view_u8(&kvl.v, kvl.v.len());
7517 let rows_on = std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0");
7518 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
7519 let mut fa_q8 = false;
7523 if !swa && hd == 512 && b_glob >= crate::fa512_min_tkv() && rows_on {
7524 e.fa_decode_rows(&sl.q, &k_view, &v_view, &mut sl.attn, hd, nh, nkv, b_glob - 1,
7525 1, scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
7526 Some((&kvl.len_d, -1)), false, false,
7527 Some((&mut sl.zq, &mut sl.zd)))?;
7528 fa_q8 = true;
7529 } else if swa && b_swa > win && hd == 256 && rows_on {
7530 e.fa_decode_rows_w(&sl.q, &k_view, &v_view, &mut sl.attn, hd, nh, nkv,
7531 &kvl.len_d, -1, 1, scale, win,
7532 kvl.k_tok_bytes, kvl.v_tok_bytes,
7533 Some((&mut sl.zq, &mut sl.zd)))?;
7534 fa_q8 = true;
7535 } else {
7536 let b = if swa { b_swa } else { b_glob };
7537 e.fa_decode_dc(&sl.q, &k_view, &v_view, &mut sl.attn, hd, nh, nkv, &kvl.len_d, b,
7538 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
7539 swa && crate::Engine::wkv_on())?;
7540 }
7541 if !fa_q8 {
7542 let aq = unsafe { &*(&sl.attn as *const CudaSlice<f32>) };
7543 e.quantize_q8_1_into(aq, 1, nh * hd, &mut sl.zq, &mut sl.zd)?;
7544 }
7545 {
7546 let zq = unsafe { &*(&sl.zq as *const CudaSlice<i8>) };
7547 let zd = unsafe { &*(&sl.zd as *const CudaSlice<f32>) };
7548 self.g4_matvec_m1_into(e, &fa.wo, zq, zd, &mut sl.o)?;
7549 }
7550 Ok(())
7551 }
7552
7553 fn gemma4_layer_tail_slotted(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
7556 next_norm: Option<&CudaSlice<f32>>, sl: &mut G4DcSlots)
7557 -> Result<(), Box<dyn std::error::Error>> {
7558 let n_embd = self.cfg.n_embd as usize;
7559 let eps = self.cfg.rms_eps;
7560 let bits = layer.gemma4.as_ref().unwrap();
7561 let crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } = &layer.ffn
7562 else { return Err("slotted tail: dense ffn only".into()) };
7563 e.add_rms_norm(&sl.cur, &sl.x, bits.ffn_norm.float_data(), &mut sl.attn_out,
7564 &mut sl.zsh, n_embd, 1, eps)?;
7565 let n_ff = ffn_gate.out_features();
7566 {
7567 let zshr = unsafe { &*(&sl.zsh as *const CudaSlice<f32>) };
7568 e.quantize_q8_1_into(zshr, 1, n_embd, &mut sl.zq, &mut sl.zd)?;
7569 }
7570 {
7571 let zq = unsafe { &*(&sl.zq as *const CudaSlice<i8>) };
7572 let zd = unsafe { &*(&sl.zd as *const CudaSlice<f32>) };
7573 if !e.matmul_q4_fused2_into(ffn_gate, ffn_up, zq, zd, &mut sl.gate, &mut sl.up)? {
7574 return Err("slotted tail: ffn fused2 unavailable".into());
7575 }
7576 }
7577 debug_assert!(e.uses_q8_1_fast(ffn_down));
7578 {
7579 let upr = unsafe { &*(&sl.up as *const CudaSlice<f32>) };
7580 let upv = e.view(upr, n_ff);
7581 let up_all = upv.slice(0..n_ff);
7582 let gr = unsafe { &*(&sl.gate as *const CudaSlice<f32>) };
7583 e.gelu_tanh_mul_q8_1_into(gr, &up_all, &mut sl.act, n_ff, 1,
7584 &mut sl.actq, &mut sl.actd)?;
7585 }
7586 {
7587 let aq = unsafe { &*(&sl.actq as *const CudaSlice<i8>) };
7588 let ad = unsafe { &*(&sl.actd as *const CudaSlice<f32>) };
7589 self.g4_matvec_m1_into(e, ffn_down, aq, ad, &mut sl.f0)?;
7590 }
7591 e.rms_norm(&sl.f0, bits.post_ffw_norm.float_data(), &mut sl.sn, n_embd, 1, eps)?;
7592 match next_norm {
7593 Some(w) => {
7594 e.add_scale_rms_norm_q8_1_into(&sl.sn, &sl.attn_out, bits.layer_scale, w,
7595 &mut sl.xn, n_embd, 1, eps,
7596 &mut sl.hq, &mut sl.hd_)?;
7597 }
7598 None => {
7599 e.add_scale(&sl.sn, &sl.attn_out, bits.layer_scale, &mut sl.xn, n_embd)?;
7600 }
7601 }
7602 Ok(())
7603 }
7604
7605 #[allow(clippy::too_many_arguments)]
7607 fn gemma4_decode_attn_dc(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
7608 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
7609 pos_d: &CudaSlice<i32>, cache: &mut Cache,
7610 cap_bucket_max: Option<(usize, usize)>)
7611 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7612 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
7613 let eps = self.cfg.rms_eps;
7614 let aux = self.gemma4_aux.as_ref().unwrap();
7615 let (q0, k0, v0) = if swa {
7616 match e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)? {
7617 Some(t3) => t3,
7618 None => {
7619 let h0 = e.zeros(0)?;
7620 (e.matmul_pre(&fa.wq, hq, hdq, &h0, 1)?,
7621 e.matmul_pre(&fa.wk, hq, hdq, &h0, 1)?,
7622 e.matmul_pre(&fa.wv, hq, hdq, &h0, 1)?)
7623 }
7624 }
7625 } else {
7626 let (q0, k0) = match e.matmul_q4_fused2(&fa.wq, &fa.wk, hq, hdq)? {
7627 Some(p) => p,
7628 None => {
7629 let h0 = e.zeros(0)?;
7630 (e.matmul_pre(&fa.wq, hq, hdq, &h0, 1)?,
7631 e.matmul_pre(&fa.wk, hq, hdq, &h0, 1)?)
7632 }
7633 };
7634 let v0 = e.clone_dtod(&k0)?;
7635 (q0, k0, v0)
7636 };
7637 let mut q = e.uninit(nh * hd)?;
7638 let mut k = e.uninit(nkv * hd)?;
7639 let mut v = e.uninit(nkv * hd)?;
7640 let ff = if swa { None } else {
7642 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
7643 };
7644 let kvl = cache.kv[il].as_mut().unwrap();
7645 let kv_fp8 = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
7646 if crate::Engine::qkv_append_on() {
7647 e.rms_norm_qkv_rope_append_dc(&q0, &k0, &v0, fa.q_norm.float_data(),
7649 fa.k_norm.float_data(), &aux.ones, &mut q, &mut k, &mut v, hd, nh, nkv,
7650 pos_d, nh, nkv, base, 1.0, ff, eps,
7651 &mut kvl.k, &mut kvl.v, &kvl.len_d, kvl.k_tok_bytes, kvl.v_tok_bytes, kv_fp8)?;
7652 } else {
7653 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
7654 &aux.ones, &mut q, &mut k, &mut v, hd, nh, nkv,
7655 pos_d, nh, nkv, base, 1.0, ff, eps)?;
7656 e.append_kv_quantized_dc(&k, &v, &mut kvl.k, &mut kvl.v, &kvl.len_d,
7657 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes, kv_fp8)?;
7658 }
7659 e.inc_seqlen(&mut kvl.len_d)?;
7660 let mut attn = e.uninit(nh * hd)?;
7661 let mut fa_q8: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
7664 match cap_bucket_max {
7669 None => {
7670 kvl.len += 1;
7674 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
7675 if !swa && hd == 512 && kvl.len >= crate::fa512_min_tkv()
7676 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
7677 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
7680 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
7681 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
7682 e.fa_decode_rows(&q, &kp, &vp, &mut attn, hd, nh, nkv, kvl.len - 1, 1,
7683 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
7684 Some((&kvl.len_d, -1)), false, false,
7685 Some((&mut aq8, &mut ad8)))?;
7686 fa_q8 = Some((aq8, ad8));
7687 } else if swa && kvl.len > win && hd == 256
7688 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
7689 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
7691 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
7692 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
7693 e.fa_decode_rows_w(&q, &kp, &vp, &mut attn, hd, nh, nkv, &kvl.len_d, -1,
7694 1, scale, win, kvl.k_tok_bytes, kvl.v_tok_bytes,
7695 Some((&mut aq8, &mut ad8)))?;
7696 fa_q8 = Some((aq8, ad8));
7697 } else {
7698 let (off_tok, t_kv) = if swa && kvl.len > win { (kvl.len - win, win) }
7699 else { (0, kvl.len) };
7700 let k_view = e.view_u8_range(&kvl.k, off_tok * kvl.k_tok_bytes,
7701 (off_tok + t_kv) * kvl.k_tok_bytes);
7702 let v_view = e.view_u8_range(&kvl.v, off_tok * kvl.v_tok_bytes,
7703 (off_tok + t_kv) * kvl.v_tok_bytes);
7704 e.fa_decode_kvmod(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, t_kv, scale,
7705 kvl.k_tok_bytes, kvl.v_tok_bytes, swa && crate::Engine::wkv_on())?;
7706 }
7707 }
7708 Some((b_swa, b_glob)) => {
7709 let k_view = e.view_u8(&kvl.k, kvl.k.len());
7715 let v_view = e.view_u8(&kvl.v, kvl.v.len());
7716 let rows_on = std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0");
7717 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
7718 if !swa && hd == 512 && b_glob >= crate::fa512_min_tkv() && rows_on {
7719 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
7720 e.fa_decode_rows(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, b_glob - 1,
7721 1, scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
7722 Some((&kvl.len_d, -1)), false, false,
7723 Some((&mut aq8, &mut ad8)))?;
7724 fa_q8 = Some((aq8, ad8));
7725 } else if swa && b_swa > win && hd == 256 && rows_on {
7726 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
7727 e.fa_decode_rows_w(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
7728 &kvl.len_d, -1, 1, scale, win,
7729 kvl.k_tok_bytes, kvl.v_tok_bytes,
7730 Some((&mut aq8, &mut ad8)))?;
7731 fa_q8 = Some((aq8, ad8));
7732 } else {
7733 let b = if swa { b_swa } else { b_glob };
7734 e.fa_decode_dc(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, &kvl.len_d, b,
7735 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
7736 swa && crate::Engine::wkv_on())?;
7737 }
7738 }
7739 }
7740 if let Some((aq8, ad8)) = fa_q8 {
7743 let mut y = e.uninit(fa.wo.out_features())?;
7744 self.g4_matvec_m1_into(e, &fa.wo, &aq8, &ad8, &mut y)?;
7745 return Ok(y);
7746 }
7747 Ok(e.matmul(&fa.wo, &attn, 1)?)
7748 }
7749
7750 pub fn gemma4_generate_graph(&self, e: &Engine, prompt_pos: usize, first_token: u32,
7755 cache: &mut Cache, max_new: usize, eos: &[u32],
7756 mut on_token: impl FnMut(u32) -> bool)
7757 -> Result<(Vec<u32>, crate::decode::StopReason), Box<dyn std::error::Error>> {
7758 if self.is_gemma4_e4b() {
7759 return Err("E4B graph serving is unwired (HANDOVER-E4B.md) — dc-eager is the serving arm".into());
7760 }
7761 use crate::decode::StopReason;
7762 let n_vocab = self.output.out_features();
7763 let n_embd = self.cfg.n_embd as usize;
7764 let embd_gpu = self.embd_gpu.get_or_init(|| {
7765 e.upload_u8(&self.embd.raw).expect("embed table upload")
7766 });
7767 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
7768 for kvl in cache.kv.iter_mut().flatten() {
7769 e.set_i32_one(&mut kvl.len_d, kvl.len as i32)?;
7770 }
7771 let mut token_d = e.stream().clone_htod(&[first_token])?;
7772 let mut pos_d = e.htod_i32(&[prompt_pos as i32])?;
7773 let g4 = self.cfg.gemma4.as_ref().unwrap();
7774 let (hd_s, hd_g) = (g4.key_length_swa as usize, g4.key_length_global as usize);
7775 let nkv_s = g4.head_count_kv.iter().zip(g4.swa_pattern.iter())
7777 .find(|p| *p.1).map(|p| *p.0 as usize).unwrap_or(8);
7778 let nkv_g = g4.head_count_kv.iter().zip(g4.swa_pattern.iter())
7779 .find(|p| !*p.1).map(|p| *p.0 as usize).unwrap_or(2);
7780 let mut graphs: std::collections::HashMap<((bool, usize), (bool, usize), bool, bool),
7781 (cudarc::driver::CudaGraph,
7782 Vec<Box<dyn std::any::Any + Send>>)> = Default::default();
7783 let mut slots = self.g4_dc_slots(e)?;
7786 const RING: usize = 64;
7789 const DRAIN: usize = 1;
7795 let mut ring = e.stream().alloc_zeros::<u32>(RING)?;
7796 let ring_base = prompt_pos;
7797 let mut out = Vec::with_capacity(max_new);
7798 let mut reason = StopReason::MaxNew;
7799 let mut next = first_token;
7800 let mut captures = 0usize;
7801 for _ in 0..max_new {
7802 out.push(next);
7803 if eos.contains(&next) { reason = StopReason::Eos; break; }
7804 if !on_token(next) { reason = StopReason::Callback; break; }
7805 let t_kv = cache.pos + 1;
7806 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
7814 let f512 = crate::fa512_min_tkv();
7815 let key_s = if t_kv > win { (true, usize::MAX) }
7816 else { e.fa_bucket_key(t_kv, hd_s, nkv_s, crate::Engine::wkv_on()) };
7817 let (key_g, rung_end) = if t_kv >= f512 {
7818 let end = (t_kv + 1).next_power_of_two().max(f512 * 2);
7821 ((true, end), end)
7822 } else { (e.fa_bucket_key(t_kv, hd_g, nkv_g, false), t_kv) };
7823 let key = (key_s, key_g, t_kv >= f512, t_kv > win);
7824 if !graphs.contains_key(&key) {
7825 let bucket_max = (t_kv, rung_end);
7826 let snap = cache.snapshot(e)?;
7828 let pos_save = e.dtoh_i32_one(&pos_d)?;
7829 let len_save: Vec<Option<i32>> = cache.kv.iter()
7830 .map(|k| k.as_ref().map(|kvl| e.dtoh_i32_one(&kvl.len_d).unwrap())).collect();
7831 let tok_save = e.dtoh_u32_one(&token_d)?;
7832 let graph = {
7837 let tok_ref = &mut token_d;
7838 let pos_ref = &mut pos_d;
7839 let cache_ref = &mut *cache;
7840 let slots_ref = &mut slots;
7841 let ring_ref = &mut ring;
7842 e.capture_graph_retained_flags(
7843 cudarc::driver::sys::CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_USE_NODE_PRIORITY,
7844 |e| {
7845 let tok_in = unsafe { &*(tok_ref as *const CudaSlice<u32>) };
7847 let sl = unsafe { &mut *(slots_ref as *mut G4DcSlots) };
7848 let rg = unsafe { &mut *(ring_ref as *mut CudaSlice<u32>) };
7849 self.gemma4_decode_step_dc_slotted(e, tok_in, pos_ref, embd_gpu, qt, rb,
7850 cache_ref, n_vocab, Some(bucket_max),
7851 sl, tok_ref, Some((rg, ring_base)))
7852 })?
7853 };
7854 cache.rollback(e, &snap, 0)?;
7855 e.set_i32_one(&mut pos_d, pos_save)?;
7856 for (il, ls) in len_save.iter().enumerate() {
7857 if let (Some(kvl), Some(v)) = (cache.kv[il].as_mut(), ls) {
7858 e.set_i32_one(&mut kvl.len_d, *v)?;
7859 }
7860 }
7861 e.set_u32_one(&mut token_d, tok_save)?;
7862 if std::env::var("MEMRA_GRAPH_CENSUS").as_deref() == Ok("1") {
7863 if let Ok(c) = crate::graph_update::node_census(&graph.0) {
7864 eprintln!("[graph-census] {c:?}");
7865 }
7866 }
7867 graphs.insert(key, graph);
7868 captures += 1;
7869 }
7870 let mut chunk = 1usize;
7875 let drain_cap: usize = std::env::var("MEMRA_GRAPH_DRAIN").ok()
7876 .and_then(|v| v.parse().ok()).unwrap_or(DRAIN);
7877 while chunk < drain_cap && out.len() + chunk < max_new {
7878 let t_next = cache.pos + 1 + chunk;
7879 let key_s2 = if t_next > win { (true, usize::MAX) }
7880 else { e.fa_bucket_key(t_next, hd_s, nkv_s, crate::Engine::wkv_on()) };
7881 let key_g2 = if t_next >= f512 {
7882 (true, (t_next + 1).next_power_of_two().max(f512 * 2))
7883 } else { e.fa_bucket_key(t_next, hd_g, nkv_g, false) };
7884 if (key_s2, key_g2, t_next >= f512, t_next > win) != key { break; }
7885 chunk += 1;
7886 }
7887 let g = &graphs.get(&key).unwrap().0;
7888 for _ in 0..chunk { g.launch()?; }
7889 e.stream().synchronize()?;
7890 let ringh = e.dtoh_u32(&ring)?;
7891 for j in 0..chunk {
7892 let pos_j = cache.pos + j;
7893 let tok_j = ringh[(pos_j - ring_base) % RING];
7894 cache.pos += 0; if j + 1 == chunk { next = tok_j; }
7896 else {
7897 out.push(tok_j);
7898 if eos.contains(&tok_j) || !on_token(tok_j) {
7899 reason = if eos.contains(&tok_j) { StopReason::Eos }
7900 else { StopReason::Callback };
7901 let keep = cache.pos + j + 1;
7903 e.set_i32_one(&mut pos_d, keep as i32)?;
7904 for kvl in cache.kv.iter_mut().filter_map(|k| k.as_mut()) {
7905 e.set_i32_one(&mut kvl.len_d, keep as i32)?;
7906 kvl.len = keep;
7907 }
7908 cache.pos = keep;
7909 if std::env::var("MEMRA_GRAPH_STATS").is_ok() {
7910 eprintln!("[gemma-graph] captures={captures} buckets={}", graphs.len());
7911 }
7912 return Ok((out, reason));
7913 }
7914 }
7915 }
7916 cache.pos += chunk;
7917 for kvl in cache.kv.iter_mut().filter_map(|k| k.as_mut()) { kvl.len += chunk; }
7918 }
7919 if std::env::var("MEMRA_GRAPH_STATS").is_ok() {
7920 eprintln!("[gemma-graph] captures={captures} buckets={}", graphs.len());
7921 }
7922 Ok((out, reason))
7923 }
7924
7925 pub(crate) fn gemma4_decode_step_t(&self, e: &Engine, tokens: &[u32], pos0: usize,
7931 cache: &mut Cache)
7932 -> Result<Vec<f32>, Box<dyn std::error::Error>> {
7933 Ok(self.gemma4_decode_step_t_h(e, tokens, pos0, cache)?.0)
7934 }
7935
7936 pub(crate) fn gemma4_decode_step_t_am(&self, e: &Engine, tokens: &[u32], pos0: usize,
7940 cache: &mut Cache)
7941 -> Result<(Vec<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7942 let (ld, hn) = self.gemma4_verify_trunk(e, tokens, pos0, cache, None)?;
7943 let t = tokens.len();
7944 let n_vocab = self.output.out_features();
7945 let mut toks = e.stream().alloc_zeros::<u32>(t)?;
7946 for i in 0..t {
7947 e.argmax_token_device_col(&ld, i, n_vocab, &mut toks, i)?;
7948 }
7949 Ok((e.dtoh_u32(&toks)?, hn))
7950 }
7951
7952 pub(crate) fn gemma4_decode_step_t_am_dev(&self, e: &Engine, tok_d: &CudaSlice<u32>, t: usize,
7955 pos0: usize, cache: &mut Cache)
7956 -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7957 let (ld, hn) = self.gemma4_verify_trunk(e, &vec![0u32; t], pos0, cache, Some(tok_d))?;
7958 let n_vocab = self.output.out_features();
7959 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
7960 for i in 0..t {
7961 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
7962 }
7963 Ok((vam, hn))
7964 }
7965
7966 pub(crate) fn gemma4_decode_step_t_h(&self, e: &Engine, tokens: &[u32], pos0: usize,
7969 cache: &mut Cache)
7970 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7971 let (mut ld, hn) = self.gemma4_verify_trunk(e, tokens, pos0, cache, None)?;
7972 let t = tokens.len();
7973 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
7974 e.softcap(&mut ld, cap, t * self.output.out_features())?;
7975 Ok((e.dtoh(&ld)?, hn))
7976 }
7977
7978 pub(crate) fn verify_stream_scratch(&self, e: &Engine, cap: usize)
7981 -> Result<VerifyStreamScratch, Box<dyn std::error::Error>> {
7982 Ok(VerifyStreamScratch {
7983 pos_d: e.htod_i32(&vec![0i32; cap])?,
7984 row_ctrs: (0..cap).map(|_| e.htod_i32(&[0])).collect::<Result<_, _>>()?,
7985 })
7986 }
7987
7988 pub(crate) fn gemma4_verify_t_am_stream(&self, e: &Engine, tok_d: &CudaSlice<u32>, t: usize,
7996 ctr: &CudaSlice<i32>, hint: usize,
7997 cache: &mut Cache,
7998 scr: &mut VerifyStreamScratch)
7999 -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
8000 let n_embd = self.cfg.n_embd as usize;
8001 let eps = self.cfg.rms_eps;
8002 assert!(t <= scr.row_ctrs.len() && t <= 64);
8003 e.i32_iota_from(ctr, &mut scr.pos_d, t)?;
8004 for i in 0..t {
8005 e.i32_copy_add(ctr, &mut scr.row_ctrs[i], (i + 1) as i32)?;
8006 }
8007 let (pos_d, row_ctrs) = (&scr.pos_d, &scr.row_ctrs);
8008 let embd_gpu = self.embd_gpu.get_or_init(|| {
8009 e.upload_u8(&self.embd.raw).expect("embed table upload")
8010 });
8011 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
8012 let mut x = e.embed_gather_device_td(embd_gpu, tok_d, t, n_embd, qt, rb)?;
8013 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
8014 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
8015 let n_layers = self.layers.len();
8016 for (il, layer) in self.layers.iter().enumerate() {
8017 let (hq, hdq) = match h_carry.take() {
8018 Some(p) => p,
8019 None => e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, t, eps)?,
8020 };
8021 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
8022 let o = self.gemma4_verify_attn_stream(e, fa, il, &hq, &hdq, pos_d, t, cache,
8023 hint, row_ctrs)?;
8024 let mut cur = e.uninit(t * n_embd)?;
8025 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, t, eps)?;
8026 let next_norm = if il + 1 < n_layers {
8027 Some(self.layers[il + 1].attn_norm.float_data())
8028 } else { None };
8029 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, t, next_norm)?;
8030 x = xn;
8031 h_carry = hn;
8032 self.dflash_tap(e, cache, il, &x, t)?;
8033 }
8034 let mut hn = e.uninit(t * n_embd)?;
8035 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
8036 let ld = e.matmul(&self.output, &hn, t)?;
8037 let n_vocab = self.output.out_features();
8038 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
8039 for i in 0..t {
8040 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
8041 }
8042 Ok((vam, hn))
8043 }
8044
8045 fn dflash_tap(&self, e: &Engine, cache: &mut Cache, il: usize, x: &CudaSlice<f32>, t: usize)
8052 -> Result<(), Box<dyn std::error::Error>> {
8053 let Some(taps) = cache.dflash_taps.as_mut() else { return Ok(()) };
8054 let Some(slot) = taps.layer_ids.iter().position(|&l| l == il) else { return Ok(()) };
8055 let h = taps.hidden;
8056 let n_taps = taps.layer_ids.len();
8057 debug_assert_eq!(taps.t, t);
8058 let xv = e.view(x, t * h);
8059 for r in 0..t {
8060 let row = xv.slice(r * h..(r + 1) * h);
8061 e.copy_view_into(&mut taps.buf, r * n_taps * h + slot * h, &row, h)?;
8062 }
8063 Ok(())
8064 }
8065
8066 fn gemma4_verify_trunk(&self, e: &Engine, tokens: &[u32], pos0: usize, cache: &mut Cache,
8067 tok_dev: Option<&CudaSlice<u32>>)
8068 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
8069 let n_embd = self.cfg.n_embd as usize;
8070 let eps = self.cfg.rms_eps;
8071 let t = tokens.len();
8072 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
8073 let pos_d = e.htod_i32(&pos)?;
8074 let mut x = match tok_dev {
8075 Some(td) => {
8076 let embd_gpu = self.embd_gpu.get_or_init(|| {
8077 e.upload_u8(&self.embd.raw).expect("embed table upload")
8078 });
8079 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
8080 e.embed_gather_device_td(embd_gpu, td, t, n_embd, qt, rb)?
8081 }
8082 None => e.htod(&self.embd.gather(n_embd, tokens))?,
8083 };
8084 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
8085 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
8086 let n_layers = self.layers.len();
8087 for (il, layer) in self.layers.iter().enumerate() {
8088 let (hq, hdq) = match h_carry.take() {
8089 Some(p) => p,
8090 None => e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, t, eps)?,
8091 };
8092 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
8093 let o = self.gemma4_verify_attn(e, fa, il, &hq, &hdq, &pos_d, t, cache)?;
8094 let mut cur = e.uninit(t * n_embd)?;
8095 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, t, eps)?;
8096 let next_norm = if il + 1 < n_layers {
8097 Some(self.layers[il + 1].attn_norm.float_data())
8098 } else { None };
8099 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, t, next_norm)?;
8100 x = xn;
8101 h_carry = hn;
8102 self.dflash_tap(e, cache, il, &x, t)?;
8103 }
8104 let mut hn = e.uninit(t * n_embd)?;
8105 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
8106 let mut ld = e.matmul(&self.output, &hn, t)?;
8107 self.gemma4_suppress(e, &mut ld, t)?; cache.pos += t;
8109 Ok((ld, hn))
8110 }
8111
8112 #[allow(clippy::too_many_arguments)]
8120 fn gemma4_verify_attn_stream(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
8121 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
8122 pos_d: &CudaSlice<i32>, t: usize,
8123 cache: &mut Cache, hint: usize,
8124 row_ctrs: &[CudaSlice<i32>])
8125 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8126 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
8127 let eps = self.cfg.rms_eps;
8128 let aux = self.gemma4_aux.as_ref().unwrap();
8129 let h0 = e.zeros(0)?;
8130 let h = &h0;
8131 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8134 let f2b = *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
8135 let fused_qkv = if f2b {
8136 if swa {
8137 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
8138 .map(|(a, b, c)| (a, b, Some(c)))
8139 } else {
8140 e.matmul_q4_fused2_batched(&fa.wq, &fa.wk, hq, hdq, t)?
8141 .map(|(a, b)| (a, b, None))
8142 }
8143 } else { None };
8144 let (q0, k0, v0) = match fused_qkv {
8145 Some((a, b, cv)) => {
8146 let v = match cv { Some(c) => c, None => e.clone_dtod(&b)? };
8147 (a, b, v)
8148 }
8149 None => {
8150 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
8151 let k0 = e.matmul_pre(&fa.wk, hq, hdq, h, t)?;
8152 let v0 = if swa { e.matmul_pre(&fa.wv, hq, hdq, h, t)? }
8153 else { e.clone_dtod(&k0)? };
8154 (q0, k0, v0)
8155 }
8156 };
8157 let mut q = e.uninit(t * nh * hd)?;
8158 let mut k = e.uninit(t * nkv * hd)?;
8159 let mut v = e.uninit(t * nkv * hd)?;
8160 let ff = if swa { None } else {
8163 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
8164 };
8165 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
8166 &aux.ones, &mut q, &mut k, &mut v, hd, nh * t, nkv * t,
8167 pos_d, nh, nkv, base, 1.0, ff, eps)?;
8168 let kvl = cache.kv[il].as_mut().unwrap();
8169 e.append_kv_quantized_rows_dc(&k, &v, &mut kvl.k, &mut kvl.v, &kvl.len_d, t,
8171 kvl.kv_dim_k, kvl.kv_dim_v,
8172 kvl.k_tok_bytes, kvl.v_tok_bytes,
8173 (!swa && crate::Engine::gkv_on())
8174 || (swa && crate::Engine::wkv_on()))?;
8175 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
8178 let mut attn = e.uninit(t * nh * hd)?;
8179 let k_view = e.view_u8(&kvl.k, kvl.k.len());
8180 let v_view = e.view_u8(&kvl.v, kvl.v.len());
8181 if swa && hint + 1 >= win {
8184 e.fa_decode_rows_w(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
8187 &kvl.len_d, 0, t, scale, win,
8188 kvl.k_tok_bytes, kvl.v_tok_bytes, None)?;
8189 } else if hd == 512 && hint + t < crate::fa512_min_tkv() {
8190 let bucket = (hint + t + 2).next_power_of_two()
8203 .min(crate::fa512_min_tkv().saturating_sub(1));
8204 let qv = e.view(&q, t * nh * hd);
8205 for i in 0..t {
8206 let q_row = qv.slice(i * nh * hd..(i + 1) * nh * hd);
8207 let mut q_one = e.uninit(nh * hd)?;
8208 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
8209 let mut a_one = e.uninit(nh * hd)?;
8210 e.fa_decode_dc(&q_one, &k_view, &v_view, &mut a_one, hd, nh, nkv,
8211 &row_ctrs[i], bucket, scale,
8212 kvl.k_tok_bytes, kvl.v_tok_bytes, false)?;
8213 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
8214 }
8215 } else if hd == 512 {
8216 e.fa_decode_rows(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, hint, t, scale,
8219 kvl.k_tok_bytes, kvl.v_tok_bytes,
8220 Some((&kvl.len_d, 0)), false, false, None)?;
8221 } else {
8222 e.fa_decode_rows_dc(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
8224 &kvl.len_d, hint + t, t, scale,
8225 kvl.k_tok_bytes, kvl.v_tok_bytes, 0,
8226 swa && crate::Engine::wkv_on())?;
8227 }
8228 Ok(e.matmul(&fa.wo, &attn, t)?)
8229 }
8230
8231 fn gemma4_verify_attn(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
8232 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
8233 pos_d: &CudaSlice<i32>, t: usize,
8234 cache: &mut Cache)
8235 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8236 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
8237 let eps = self.cfg.rms_eps;
8238 let aux = self.gemma4_aux.as_ref().unwrap();
8239 let n_embd = self.cfg.n_embd as usize;
8240 let _ = n_embd;
8241
8242 let h0 = e.zeros(0)?;
8243 let h = &h0;
8244 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8247 let f2b = *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
8248 let fused_qkv = if f2b {
8249 if swa {
8250 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
8251 .map(|(a, b, c)| (a, b, Some(c)))
8252 } else {
8253 e.matmul_q4_fused2_batched(&fa.wq, &fa.wk, hq, hdq, t)?
8254 .map(|(a, b)| (a, b, None))
8255 }
8256 } else { None };
8257 let (q0, k0, v0) = match fused_qkv {
8258 Some((a, b, cv)) => {
8259 let v = match cv { Some(c) => c, None => e.clone_dtod(&b)? };
8260 (a, b, v)
8261 }
8262 None => {
8263 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
8264 let k0 = e.matmul_pre(&fa.wk, hq, hdq, h, t)?;
8265 let v0 = if swa { e.matmul_pre(&fa.wv, hq, hdq, h, t)? }
8266 else { e.clone_dtod(&k0)? };
8267 (q0, k0, v0)
8268 }
8269 };
8270 let mut q = e.uninit(t * nh * hd)?;
8271 let mut k = e.uninit(t * nkv * hd)?;
8272 let mut v = e.uninit(t * nkv * hd)?;
8273 let ff = if swa { None } else {
8276 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
8277 };
8278 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
8279 &aux.ones, &mut q, &mut k, &mut v, hd, nh * t, nkv * t,
8280 pos_d, nh, nkv, base, 1.0, ff, eps)?;
8281 let kvl = cache.kv[il].as_mut().unwrap();
8282 let base_len = kvl.len;
8283 e.append_kv_quantized_rows(&k, &v, &mut kvl.k, &mut kvl.v, base_len, t,
8284 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes, (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on()))?;
8285 kvl.len += t;
8286 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
8287 let mut attn = e.uninit(t * nh * hd)?;
8288 let rows_ok = (hd == 256 && base_len + 1 >= crate::fa_vec_min_tkv())
8291 || (hd == 512 && !swa && base_len + 1 >= crate::fa512_min_tkv());
8294 if rows_ok && (!swa || base_len + t <= win) {
8295 let k_view = e.view_u8(&kvl.k, (base_len + t) * kvl.k_tok_bytes);
8296 let v_view = e.view_u8(&kvl.v, (base_len + t) * kvl.v_tok_bytes);
8297 if hd == 512 {
8298 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
8300 e.fa_decode_rows(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, base_len, t,
8301 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
8302 Some((&kvl.len_d, 0)), false,
8303 swa && crate::Engine::wkv_on(), None)?;
8304 } else {
8305 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
8309 e.fa_decode_rows_dc(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
8310 &kvl.len_d, base_len + t, t, scale,
8311 kvl.k_tok_bytes, kvl.v_tok_bytes, 0,
8312 swa && crate::Engine::wkv_on())?;
8313 }
8314 return Ok(e.matmul(&fa.wo, &attn, t)?);
8315 }
8316 if hd == 256 && swa && base_len + 1 >= win
8324 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
8325 let k_view = e.view_u8(&kvl.k, (base_len + t) * kvl.k_tok_bytes);
8326 let v_view = e.view_u8(&kvl.v, (base_len + t) * kvl.v_tok_bytes);
8327 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
8328 e.fa_decode_rows_w(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, &kvl.len_d, 0,
8329 t, scale, win, kvl.k_tok_bytes, kvl.v_tok_bytes, None)?;
8330 return Ok(e.matmul(&fa.wo, &attn, t)?);
8331 }
8332 for i in 0..t {
8333 let avail = base_len + i + 1;
8334 let (off_tok, t_kv) = if swa && avail > win { (avail - win, win) } else { (0, avail) };
8335 let k_view = e.view_u8_range(&kvl.k, off_tok * kvl.k_tok_bytes,
8336 (off_tok + t_kv) * kvl.k_tok_bytes);
8337 let v_view = e.view_u8_range(&kvl.v, off_tok * kvl.v_tok_bytes,
8338 (off_tok + t_kv) * kvl.v_tok_bytes);
8339 let qi = e.view(&q, t * nh * hd);
8340 let q_row = qi.slice(i * nh * hd..(i + 1) * nh * hd);
8341 let mut q_one = e.uninit(nh * hd)?;
8342 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
8343 let mut a_one = e.uninit(nh * hd)?;
8344 if swa && avail > win && hd == 256
8348 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
8349 let kp = e.view_u8(&kvl.k, avail * kvl.k_tok_bytes);
8350 let vp = e.view_u8(&kvl.v, avail * kvl.v_tok_bytes);
8351 e.i32_set_k(&mut kvl.len_d, (avail - 1) as i32)?;
8352 e.fa_decode_rows_w(&q_one, &kp, &vp, &mut a_one, hd, nh, nkv, &kvl.len_d, 0,
8353 1, scale, win, kvl.k_tok_bytes, kvl.v_tok_bytes, None)?;
8354 } else if !swa && hd == 512 && avail >= crate::fa512_min_tkv()
8355 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
8356 let kp = e.view_u8(&kvl.k, avail * kvl.k_tok_bytes);
8357 let vp = e.view_u8(&kvl.v, avail * kvl.v_tok_bytes);
8358 e.i32_set_k(&mut kvl.len_d, (avail - 1) as i32)?;
8359 e.fa_decode_rows(&q_one, &kp, &vp, &mut a_one, hd, nh, nkv, avail - 1, 1,
8360 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
8361 Some((&kvl.len_d, 0)), false, false, None)?;
8362 } else {
8363 e.fa_decode_kvmod(&q_one, &k_view, &v_view, &mut a_one, hd, nh, nkv, t_kv, scale,
8364 kvl.k_tok_bytes, kvl.v_tok_bytes, swa && crate::Engine::wkv_on())?;
8365 }
8366 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
8367 }
8368 Ok(e.matmul(&fa.wo, &attn, t)?)
8369 }
8370
8371 pub(crate) fn gemma4_decode_step_h(&self, e: &Engine, token: u32, cache: &mut Cache)
8374 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
8375 if let Some(split) = crate::pp::pp2_split(self.layers.len()) {
8380 return self.gemma4_decode_step_h_pp2(e, token, cache, split);
8381 }
8382 if crate::pp::pp_cuts(self.layers.len()).is_some() {
8383 crate::pp::warn_unwired_once("gemma4 eager decode (N>2)");
8384 }
8385 let n_embd = self.cfg.n_embd as usize;
8386 let eps = self.cfg.rms_eps;
8387 let pos_d = e.htod_i32(&[cache.pos as i32])?;
8388 let mut x = e.htod(&self.embd.gather(n_embd, &[token]))?;
8389 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
8390 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
8393 let n_layers = self.layers.len();
8394 for (il, layer) in self.layers.iter().enumerate() {
8395 let (hq, hdq) = match h_carry.take() {
8396 Some(p) => p,
8397 None => e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, 1, eps)?,
8398 };
8399 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
8400 let o = self.gemma4_decode_attn(e, fa, il, &hq, &hdq, &pos_d, cache)?;
8401 let mut cur = e.uninit(n_embd)?;
8402 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, 1, eps)?;
8403 let next_norm = if il + 1 < n_layers {
8404 Some(self.layers[il + 1].attn_norm.float_data())
8405 } else { None };
8406 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, 1, next_norm)?;
8407 x = xn;
8408 h_carry = hn;
8409 }
8410 let mut hn = e.uninit(n_embd)?;
8411 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
8412 let h_seed = e.clone_dtod(&x)?;
8413 let mut ld = e.matmul(&self.output, &hn, 1)?;
8414 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
8415 e.softcap(&mut ld, cap, self.output.out_features())?; self.gemma4_suppress(e, &mut ld, 1)?;
8417 let logits = e.dtoh(&ld)?;
8418 cache.pos += 1;
8419 Ok((logits, h_seed))
8420 }
8421
8422 fn gemma4_decode_layers(&self, e: &Engine, mut x: CudaSlice<f32>, lo: usize, hi: usize,
8430 pos_d: &CudaSlice<i32>, cache: &mut Cache)
8431 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8432 let n_embd = self.cfg.n_embd as usize;
8433 let eps = self.cfg.rms_eps;
8434 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
8435 for il in lo..hi {
8436 let layer = &self.layers[il];
8437 let (hq, hdq) = match h_carry.take() {
8438 Some(p) => p,
8439 None => e.rms_norm_q8_1(&x, self.layers[il].attn_norm.float_data(), n_embd, 1, eps)?,
8441 };
8442 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
8443 let o = self.gemma4_decode_attn(e, fa, il, &hq, &hdq, pos_d, cache)?;
8444 let mut cur = e.uninit(n_embd)?;
8445 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, 1, eps)?;
8446 let next_norm = if il + 1 < hi {
8447 Some(self.layers[il + 1].attn_norm.float_data())
8448 } else { None };
8449 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, 1, next_norm)?;
8450 x = xn;
8451 h_carry = hn;
8452 }
8453 Ok(x)
8454 }
8455
8456 fn gemma4_decode_step_h_pp2(&self, e: &Engine, token: u32, cache: &mut Cache, split: usize)
8463 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
8464 if crate::pp::pp2_streams_off() {
8465 return self.gemma4_decode_step_h_pp2_samestream(e, token, cache, split);
8466 }
8467 let rt = crate::pp::Pp2Rt::get(e)?;
8468 let e0 = rt.engine(0, e);
8469 let e1 = rt.engine(1, e);
8470 let n_embd = self.cfg.n_embd as usize;
8471 let eps = self.cfg.rms_eps;
8472
8473 let (pos_d, slot) = {
8475 let _st0 = rt.enter(0);
8476 let pos_d = e0.htod_i32(&[cache.pos as i32])?;
8477 let mut x = e0.htod(&self.embd.gather(n_embd, &[token]))?;
8478 e0.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
8479 let x = self.gemma4_decode_layers(e0, x, 0, split, &pos_d, cache)?;
8480 let slot = rt.tx(0, &x, n_embd)?;
8481 (pos_d, slot)
8482 };
8483
8484 let _st1 = rt.enter(1);
8486 let x = rt.rx(0, slot, n_embd)?;
8487 let x = self.gemma4_decode_layers(e1, x, split, self.layers.len(), &pos_d, cache)?;
8488
8489 let mut hn = e1.uninit(n_embd)?;
8490 e1.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
8491 let h_seed = e1.clone_dtod(&x)?;
8492 let mut ld = e1.matmul(&self.output, &hn, 1)?;
8493 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
8494 e1.softcap(&mut ld, cap, self.output.out_features())?;
8495 self.gemma4_suppress(e1, &mut ld, 1)?;
8496 let logits = e1.dtoh(&ld)?;
8497 cache.pos += 1;
8498 Ok((logits, h_seed))
8499 }
8500
8501 fn gemma4_decode_step_h_pp2_samestream(&self, e: &Engine, token: u32, cache: &mut Cache,
8504 split: usize)
8505 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
8506 let n_embd = self.cfg.n_embd as usize;
8507 let eps = self.cfg.rms_eps;
8508 let pos_d = e.htod_i32(&[cache.pos as i32])?;
8509
8510 let mut x = e.htod(&self.embd.gather(n_embd, &[token]))?;
8512 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
8513 let x = self.gemma4_decode_layers(e, x, 0, split, &pos_d, cache)?;
8514
8515 let boundary_tx = e.clone_dtod(&x)?;
8517 let boundary_rx = e.clone_dtod(&boundary_tx)?;
8518
8519 let x = self.gemma4_decode_layers(e, boundary_rx, split, self.layers.len(), &pos_d, cache)?;
8521
8522 let mut hn = e.uninit(n_embd)?;
8523 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
8524 let h_seed = e.clone_dtod(&x)?;
8525 let mut ld = e.matmul(&self.output, &hn, 1)?;
8526 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
8527 e.softcap(&mut ld, cap, self.output.out_features())?;
8528 self.gemma4_suppress(e, &mut ld, 1)?;
8529 let logits = e.dtoh(&ld)?;
8530 cache.pos += 1;
8531 Ok((logits, h_seed))
8532 }
8533}
8534
8535impl HybridModel {
8554 pub(crate) fn step35_geom(&self, il: usize) -> memra_gguf::config::LayerGeometry {
8557 let geometry = self.cfg.layer_geometry(il as u32)
8558 .unwrap_or_else(|| panic!("step35 layer {il} has no geometry-table row"));
8559 debug_assert_eq!(
8560 geometry.attention_gate,
8561 memra_gguf::config::AttentionGateKind::SeparateHead
8562 );
8563 geometry
8564 }
8565
8566 #[allow(clippy::too_many_arguments)]
8626 fn step35_attn_pre_wo(&self, e: &Engine, fa: &FullAttnLayer, mut g3: Vec<CudaSlice<f32>>,
8627 hg: Option<&CudaSlice<f32>>, gt_pre: Option<&CudaSlice<f32>>,
8628 pos_d: &CudaSlice<i32>, t: usize,
8629 cache: Option<&mut Cache>, il: usize, seq_end: usize)
8630 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8631 let geometry = self.step35_geom(il);
8632 let hd = geometry.head_dim_k as usize;
8633 let nkv = geometry.n_head_kv as usize;
8634 let nh = geometry.n_head as usize;
8635 let rbase = geometry.rope_base;
8636 let scale = geometry.attention_scale();
8637 let swa = geometry.window.is_some();
8638 let eps = self.cfg.rms_eps;
8639 let win = geometry.window.unwrap_or(0) as usize;
8640 let n_rot = geometry.n_rot as usize;
8641
8642 let v = g3.pop().unwrap();
8643 let k0 = g3.pop().unwrap();
8644 let q0 = g3.pop().unwrap();
8645
8646 let mut q = e.uninit(t * nh * hd)?;
8650 e.rms_norm(&q0, fa.q_norm.float_data(), &mut q, hd, nh * t, eps)?;
8651 let mut k = e.uninit(t * nkv * hd)?;
8652 e.rms_norm(&k0, fa.k_norm.float_data(), &mut k, hd, nkv * t, eps)?;
8653 let ff = if geometry.rope_factors {
8654 self.step35_aux.as_ref().and_then(|a| a.rope_freqs.as_ref())
8655 } else {
8656 None
8657 };
8658 e.rope_neox2(&mut q, &mut k, pos_d, hd, n_rot, nh, nkv, t, rbase, 1.0, ff)?;
8659
8660 let mut attn = e.uninit(t * nh * hd)?;
8661 match cache {
8662 Some(cache) => {
8663 let base_len = cache.kv[il].as_ref().unwrap().len;
8664 let legacy_tkv = std::env::var("MEMRA_STEP35_SWA_TKV").as_deref() == Ok("1");
8666 let legacy_calllocal =
8667 std::env::var("MEMRA_PRIME_CALLLOCAL").as_deref() == Ok("1");
8668 let off = if swa {
8669 let raw = base_len.saturating_sub(win - 1);
8670 if legacy_tkv || legacy_calllocal { raw } else { raw & !31usize }
8671 } else {
8672 0
8673 };
8674 {
8675 let kvl = cache.kv[il].as_mut().unwrap();
8676 assert!(kvl.len + t <= cache.max_ctx, "step35 prime: KV overflow");
8677 let write_row = e.prepare_kv_append(kvl, off, t)?;
8678 e.append_kv_quantized_rows(&k, &v, &mut kvl.k, &mut kvl.v, write_row, t,
8679 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes,
8680 kvl.v_tok_bytes, crate::Engine::kv_fp8_on())?;
8681 kvl.len += t;
8682 let new_len = kvl.len as i32;
8683 e.set_i32_one(&mut kvl.len_d, new_len)?;
8684 }
8685 let kvl = cache.kv[il].as_ref().unwrap();
8686 let t_kv = base_len + t - off;
8709 let physical = kvl.physical_rows(off, off + t_kv)?;
8710 let k_view = e.view_u8_range(&kvl.k, physical.start * kvl.k_tok_bytes,
8711 physical.end * kvl.k_tok_bytes);
8712 let v_view = e.view_u8_range(&kvl.v, physical.start * kvl.v_tok_bytes,
8713 physical.end * kvl.v_tok_bytes);
8714 let swa_naive = if legacy_tkv { t_kv > win } else { seq_end > win };
8726 if swa && swa_naive {
8727 if std::env::var("MEMRA_STEP35_SWA_FA").as_deref() == Ok("0") {
8740 e.sdpa_naive_w_quantized_view(&q, &k_view, &v_view, &mut attn, hd, nh,
8741 nkv, t, t_kv, scale, true, win,
8742 kvl.k_tok_bytes, kvl.v_tok_bytes)?;
8743 } else {
8744 e.fa_prefill_view_ws_w_hd128(&q, &k_view, &v_view, &mut attn, hd, nh,
8745 nkv, t, t_kv, scale, true, win,
8746 kvl.k_tok_bytes, kvl.v_tok_bytes)?;
8747 }
8748 } else if std::env::var("MEMRA_NOFA").is_ok() {
8749 e.sdpa_naive_quantized_view(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
8750 t, t_kv, scale, true,
8751 kvl.k_tok_bytes, kvl.v_tok_bytes)?;
8752 } else {
8753 e.fa_prefill_view_ws(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
8758 t, t_kv, scale, true,
8759 kvl.k_tok_bytes, kvl.v_tok_bytes,
8760 crate::Engine::kv_fp8_on())?;
8761 }
8762 }
8763 None => {
8764 debug_assert_eq!(seq_end, t, "step35 cacheless prefill is monolithic (seq_end == t)");
8769 if swa && seq_end > win {
8770 e.sdpa_naive_w(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
8771 } else if std::env::var("MEMRA_NOFA").is_ok() {
8772 e.sdpa_naive(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
8773 } else {
8774 e.fa_prefill(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
8775 }
8776 }
8777 }
8778
8779 let gw = fa.attn_gate.as_ref()
8782 .ok_or("step35 layer is missing attn_gate.weight (head-wise attention gate)")?;
8783 let gt_owned = if gt_pre.is_none() {
8784 Some(e.matmul(
8785 gw,
8786 hg.ok_or("step35 attention needs hg when gt_pre is absent")?,
8787 t,
8788 )?)
8789 } else {
8790 None
8791 };
8792 let gt = gt_pre.or(gt_owned.as_ref()).unwrap();
8793 let mut ag = e.uninit(t * nh * hd)?;
8794 e.attn_head_gate(&attn, gt, &mut ag, None, hd, nh, t)?;
8795 Ok(ag)
8796 }
8797
8798 pub(crate) fn step35_attn(&self, e: &Engine, fa: &FullAttnLayer, h: &CudaSlice<f32>,
8801 pos_d: &CudaSlice<i32>, t: usize, il: usize)
8802 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8803 let g3 = e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?;
8804 let ag = self.step35_attn_pre_wo(e, fa, g3, Some(h), None, pos_d, t, None, il, t)?;
8806 Ok(e.matmul(&fa.wo, &ag, t)?)
8807 }
8808
8809 #[allow(clippy::too_many_arguments)]
8816 pub(crate) fn step35_attn_prime(&self, e: &Engine, fa: &FullAttnLayer, h: &CudaSlice<f32>,
8817 hx: Option<&CudaSlice<u8>>, pos_d: &CudaSlice<i32>, t: usize,
8818 cache: &mut Cache, il: usize, seq_end: usize)
8819 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8820 let g3 = match hx {
8821 Some(xh) => e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], h, xh, t)?,
8822 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
8823 };
8824 let ag = self.step35_attn_pre_wo(
8825 e,
8826 fa,
8827 g3,
8828 Some(h),
8829 None,
8830 pos_d,
8831 t,
8832 Some(cache),
8833 il,
8834 seq_end,
8835 )?;
8836 Ok(e.matmul(&fa.wo, &ag, t)?)
8837 }
8838
8839 #[allow(clippy::too_many_arguments)]
8849 pub(crate) fn step35_decode_attn(&self, e: &Engine, fa: &FullAttnLayer, il: usize,
8850 h: &CudaSlice<f32>,
8851 pre_q: Option<(&CudaSlice<i8>, &CudaSlice<f32>)>,
8852 pos_d: &CudaSlice<i32>, cache: &mut Cache)
8853 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8854 let geometry = self.step35_geom(il);
8855 let hd = geometry.head_dim_k as usize;
8856 let nkv = geometry.n_head_kv as usize;
8857 let nh = geometry.n_head as usize;
8858 let rbase = geometry.rope_base;
8859 let scale = geometry.attention_scale();
8860 let swa = geometry.window.is_some();
8861 let eps = self.cfg.rms_eps;
8862 let win = geometry.window.unwrap_or(0) as usize;
8863 let n_rot = geometry.n_rot as usize;
8864 let n_embd = self.cfg.n_embd as usize;
8865 let gw = fa.attn_gate.as_ref()
8866 .ok_or("step35 layer is missing attn_gate.weight (head-wise attention gate)")?;
8867
8868 let (q0, k0, v0, gt) = match pre_q {
8869 Some((hq, hdq)) => {
8870 debug_assert!(e.uses_q8_1_fast(gw),
8871 "step35 pre-quantized decode requires attn_gate on the q8_1 fast path \
8872 (h is a zero-length placeholder here) — see mixer_in_q8_1_fast");
8873 let (a, b, c) = match e.matmul_q8_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)? {
8874 Some(t3) => t3,
8875 None => (e.matmul_pre(&fa.wq, hq, hdq, h, 1)?,
8876 e.matmul_pre(&fa.wk, hq, hdq, h, 1)?,
8877 e.matmul_pre(&fa.wv, hq, hdq, h, 1)?),
8878 };
8879 let gt = e.matmul_pre(gw, hq, hdq, h, 1)?;
8880 (a, b, c, gt)
8881 }
8882 None => {
8883 if e.uses_q8_1_fast(&fa.wq) && e.uses_q8_1_fast(&fa.wk)
8884 && e.uses_q8_1_fast(&fa.wv) && e.uses_q8_1_fast(gw) {
8885 let (hq, hdq) = e.quantize_q8_1(h, 1, n_embd)?;
8886 let (a, b, c) = match e.matmul_q8_fused3(&fa.wq, &fa.wk, &fa.wv, &hq, &hdq)? {
8887 Some(t3) => t3,
8888 None => (e.matmul_pre(&fa.wq, &hq, &hdq, h, 1)?,
8889 e.matmul_pre(&fa.wk, &hq, &hdq, h, 1)?,
8890 e.matmul_pre(&fa.wv, &hq, &hdq, h, 1)?),
8891 };
8892 let gt = e.matmul_pre(gw, &hq, &hdq, h, 1)?;
8893 (a, b, c, gt)
8894 } else {
8895 (e.matmul(&fa.wq, h, 1)?, e.matmul(&fa.wk, h, 1)?,
8896 e.matmul(&fa.wv, h, 1)?, e.matmul(gw, h, 1)?)
8897 }
8898 }
8899 };
8900
8901 let mut q = e.uninit(nh * hd)?;
8902 e.rms_norm(&q0, fa.q_norm.float_data(), &mut q, hd, nh, eps)?;
8903 let mut k = e.uninit(nkv * hd)?;
8904 e.rms_norm(&k0, fa.k_norm.float_data(), &mut k, hd, nkv, eps)?;
8905 let ff = if swa { None } else {
8906 self.step35_aux.as_ref().and_then(|a| a.rope_freqs.as_ref())
8907 };
8908 e.rope_neox2(&mut q, &mut k, pos_d, hd, n_rot, nh, nkv, 1, rbase, 1.0, ff)?;
8909
8910 if std::env::var("MEMRA_NOFA").is_ok() {
8911 return Err("MEMRA_NOFA (naive f32 SDPA) is incompatible with the quantized KV \
8912 cache; unset MEMRA_NOFA to use fa_decode".into());
8913 }
8914 let kvl = cache.kv[il].as_mut().unwrap();
8915 let next_len = kvl.len + 1;
8916 let (off, t_kv) = if swa && next_len > win { (next_len - win, win) } else { (0, next_len) };
8917 let write_row = e.prepare_kv_append(kvl, off & !31usize, 1)?;
8918 e.append_kv_quantized(&k, &v0, &mut kvl.k, &mut kvl.v, write_row,
8919 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes,
8920 crate::Engine::kv_fp8_on())?;
8921 kvl.len = next_len;
8922 let physical = kvl.physical_rows(off, off + t_kv)?;
8923 let k_view = e.view_u8_range(&kvl.k, physical.start * kvl.k_tok_bytes,
8924 physical.end * kvl.k_tok_bytes);
8925 let v_view = e.view_u8_range(&kvl.v, physical.start * kvl.v_tok_bytes,
8926 physical.end * kvl.v_tok_bytes);
8927 let mut attn = e.uninit(nh * hd)?;
8928 e.fa_decode_kvmod(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, t_kv, scale,
8929 kvl.k_tok_bytes, kvl.v_tok_bytes, crate::Engine::kv_fp8_on())?;
8930
8931 let mut ag = e.uninit(nh * hd)?;
8932 e.attn_head_gate(&attn, >, &mut ag, None, hd, nh, 1)?;
8933 Ok(e.matmul(&fa.wo, &ag, 1)?)
8934 }
8935}
8936
8937impl HybridModel {
8946 pub fn is_gemma4_e4b(&self) -> bool {
8947 self.gemma4_aux.as_ref().is_some_and(|a| a.e4b.is_some())
8948 }
8949
8950 fn gemma4_e4b_geom(&self, il: usize) -> (usize, usize, usize, f32, f32, bool) {
8954 let g = self.cfg.gemma4.as_ref().unwrap();
8955 let swa = g.swa_pattern[il];
8956 let hd = if swa { g.key_length_swa } else { g.key_length_global } as usize;
8957 let Mixer::Full(fa) = &self.layers[il].mixer else { panic!("e4b layer {il} not full-attn") };
8958 let nh = fa.wq.out_features() / hd;
8959 let nkv = fa.wk.out_features() / hd;
8960 (hd, nkv, nh, if swa { g.rope_base_swa } else { g.rope_base_global }, 1.0, swa)
8961 }
8962
8963 fn gemma4_e4b_kv_target(&self, il: usize) -> Option<usize> {
8965 self.layers[il].gemma4.as_ref()
8966 .and_then(|b| b.e4b.as_ref())
8967 .and_then(|e4| e4.kv_share.map(|t| t as usize))
8968 }
8969
8970 fn gemma4_e4b_inp_pl(&self, e: &Engine, tokens: &[u32], x_scaled: &CudaSlice<f32>, t: usize)
8975 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8976 let tok_d = e.stream().clone_htod(&tokens.to_vec())?;
8977 self.gemma4_e4b_inp_pl_dev(e, &tok_d, x_scaled, t)
8978 }
8979
8980 fn gemma4_e4b_inp_pl_dev(&self, e: &Engine, tok_d: &CudaSlice<u32>,
8982 x_scaled: &CudaSlice<f32>, t: usize)
8983 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8984 let aux = self.gemma4_aux.as_ref().unwrap();
8985 let m = aux.e4b.as_ref().unwrap();
8986 let n_embd = self.cfg.n_embd as usize;
8987 let n_layer = self.layers.len();
8988 let width = m.n_epl * n_layer;
8989 let tbl = m.tok_tbl_gpu.get_or_init(|| {
8990 e.upload_u8(&m.tok_embd_bytes).expect("e4b per-layer token table upload")
8991 });
8992 let mut a = e.embed_gather_device_td(tbl, tok_d, t, width, m.tok_embd_qt,
8993 m.tok_embd_row_bytes)?;
8994 e.scale_inplace(&mut a, (m.n_epl as f32).sqrt(), t * width)?;
8995 let mut p = e.matmul(&m.model_proj, x_scaled, t)?;
8996 e.scale_inplace(&mut p, 1.0 / (n_embd as f32).sqrt(), t * width)?;
8997 let mut pn = e.uninit(t * width)?;
8998 e.rms_norm(&p, m.proj_norm.float_data(), &mut pn, m.n_epl, t * n_layer,
8999 self.cfg.rms_eps)?;
9000 let mut out = e.uninit(t * width)?;
9001 e.add_scale(&a, &pn, 1.0 / 2f32.sqrt(), &mut out, t * width)?;
9002 Ok(out)
9003 }
9004
9005 #[allow(clippy::too_many_arguments)]
9010 fn gemma4_e4b_attn(&self, e: &Engine, il: usize,
9011 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
9012 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache,
9013 dc_bucket: Option<usize>)
9014 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9015 let (hd, nkv, nh, base, scale, swa) = self.gemma4_e4b_geom(il);
9016 let eps = self.cfg.rms_eps;
9017 let aux = self.gemma4_aux.as_ref().unwrap();
9018 let Mixer::Full(fa) = &self.layers[il].mixer else { unreachable!() };
9019 let h0 = e.zeros(0)?;
9023 let h = &h0;
9024
9025 let ff = if swa { None } else {
9026 Some(aux.rope_freqs.as_ref().expect("e4b global rope needs rope_freqs.weight"))
9027 };
9028 let share = self.gemma4_e4b_kv_target(il);
9029 let mut kv_f32: Option<(CudaSlice<f32>, CudaSlice<f32>)> = None;
9031 let mut q;
9032 if let Some(_tgt) = share {
9033 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
9034 q = e.uninit(t * nh * hd)?;
9035 let mut kdummy = e.uninit(1)?;
9038 let mut vdummy = e.uninit(1)?;
9039 e.rms_norm_qkv_rope(&q0, &q0, &q0, fa.q_norm.float_data(),
9040 fa.q_norm.float_data(), &aux.ones,
9041 &mut q, &mut kdummy, &mut vdummy, hd, nh * t, 0,
9042 pos_d, nh, 1, base, 1.0, ff, eps)?;
9043 } else {
9044 let e4bits = self.layers[il].gemma4.as_ref().and_then(|g| g.e4b.as_ref());
9048 let cat = e4bits.and_then(|e4| e4.qkv_cat.as_ref());
9049 q = e.uninit(t * nh * hd)?;
9050 let mut k = e.uninit(t * nkv * hd)?;
9051 let mut v = e.uninit(t * nkv * hd)?;
9052 if t == 1 && cat.is_some() {
9053 let qkv0 = e.matmul_pre(cat.unwrap(), hq, hdq, h, 1)?;
9054 e.rms_norm_qkv_rope_cat(&qkv0, fa.q_norm.float_data(), fa.k_norm.float_data(),
9055 &aux.ones, &mut q, &mut k, &mut v, hd, nh, nkv,
9056 pos_d, nh, nkv, base, 1.0, ff, eps)?;
9057 } else {
9058 let (q0, k0, v0) = match if t == 1 {
9059 e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)?
9060 } else {
9061 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9064 if *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0")) {
9065 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
9066 } else { None }
9067 } {
9068 Some(triple) => triple,
9069 None => (e.matmul_pre(&fa.wq, hq, hdq, h, t)?,
9070 e.matmul_pre(&fa.wk, hq, hdq, h, t)?,
9071 e.matmul_pre(&fa.wv, hq, hdq, h, t)?), };
9073 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(),
9076 fa.k_norm.float_data(), &aux.ones, &mut q, &mut k, &mut v,
9077 hd, nh * t, nkv * t, pos_d, nh, nkv, base, 1.0, ff, eps)?;
9078 }
9079 let kvl = cache.kv[il].as_mut().unwrap();
9080 let cls = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
9084 if dc_bucket.is_some() {
9085 debug_assert!(t == 1);
9090 e.append_kv_quantized_row_dc_inc(&k, &v, &mut kvl.k, &mut kvl.v,
9092 &mut kvl.len_d, kvl.kv_dim_k, kvl.kv_dim_v,
9093 kvl.k_tok_bytes, kvl.v_tok_bytes, cls)?;
9094 } else {
9095 e.append_kv_quantized_rows(&k, &v, &mut kvl.k, &mut kvl.v, kvl.len, t,
9096 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes,
9097 kvl.v_tok_bytes, cls)?;
9098 kvl.len += t;
9099 }
9100 kv_f32 = Some((k, v));
9101 }
9102 let kvl_idx = share.unwrap_or(il);
9105 let kvl = cache.kv[kvl_idx].as_ref().unwrap();
9106 let base_len = kvl.len - t; let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
9108 let mut attn = e.uninit(t * nh * hd)?;
9109 if t > 1 && base_len == 0 && std::env::var("MEMRA_NOFA").is_err() {
9121 if let Some((kf, vf)) = &kv_f32 {
9122 if hd == 256 && t <= win {
9123 e.fa_prefill(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale, true)?;
9124 return Ok(e.matmul(&fa.wo, &attn, t)?);
9125 }
9126 if hd == 256 && swa && t > win {
9127 e.fa_prefill_w(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale, true,
9128 win)?;
9129 return Ok(e.matmul(&fa.wo, &attn, t)?);
9130 }
9131 if hd == 512 && !swa {
9132 e.fa_prefill_hd512(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale,
9133 true)?;
9134 return Ok(e.matmul(&fa.wo, &attn, t)?);
9135 }
9136 } else if share.is_some() {
9137 let g = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
9138 let k_view = e.view_u8(&kvl.k, kvl.k.len());
9139 let v_view = e.view_u8(&kvl.v, kvl.v.len());
9140 if hd == 256 && (!swa || t <= win) {
9141 e.fa_prefill_view(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, t, t,
9143 scale, true, kvl.k_tok_bytes, kvl.v_tok_bytes, g)?;
9144 return Ok(e.matmul(&fa.wo, &attn, t)?);
9145 }
9146 let kv_dim = nkv * hd;
9149 let mut kf = e.uninit(t * kv_dim)?;
9150 let mut vf = e.uninit(t * kv_dim)?;
9151 e.fa_dequant_kv_view_f32(&k_view, &v_view, &mut kf, &mut vf, kv_dim, kv_dim,
9152 t, kvl.k_tok_bytes, kvl.v_tok_bytes, g)?;
9153 if hd == 512 {
9154 e.fa_prefill_hd512(&q, &kf, &vf, &mut attn, hd, nh, nkv, t, t, scale,
9155 true)?;
9156 } else {
9157 e.fa_prefill_w(&q, &kf, &vf, &mut attn, hd, nh, nkv, t, t, scale, true,
9158 win)?;
9159 }
9160 return Ok(e.matmul(&fa.wo, &attn, t)?);
9161 }
9162 }
9163 if let Some(bucket) = dc_bucket {
9164 assert!(t == 1);
9169 let bucket = if hd == 512 && win <= crate::fa512_min_tkv() {
9175 bucket.min(crate::fa512_min_tkv().saturating_sub(1))
9176 } else { bucket };
9177 let k_view = e.view_u8(&kvl.k, kvl.k.len());
9178 let v_view = e.view_u8(&kvl.v, kvl.v.len());
9179 let g = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
9180 if crate::Engine::wpf_level() >= 1 {
9188 e.prefetch_weight_l2(&fa.wo)?;
9189 }
9190 if e.uses_q8_1_fast(&fa.wo) {
9193 let mut oq = e.alloc_i8_uninit(nh * hd)?;
9194 let mut od = e.zeros(nh * hd / 32)?;
9195 e.fa_decode_dc_q8(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
9196 &kvl.len_d, bucket, scale,
9197 kvl.k_tok_bytes, kvl.v_tok_bytes, g,
9198 Some((&mut oq, &mut od)))?;
9199 return Ok(e.matmul_pre(&fa.wo, &oq, &od, &attn, t)?);
9200 }
9201 e.fa_decode_dc(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
9202 &kvl.len_d, bucket, scale,
9203 kvl.k_tok_bytes, kvl.v_tok_bytes, g)?;
9204 return Ok(e.matmul(&fa.wo, &attn, t)?);
9205 }
9206 for i in 0..t {
9207 let avail = base_len + i + 1;
9208 let (off_tok, t_kv) = if swa && avail > win { (avail - win, win) } else { (0, avail) };
9209 let k_view = e.view_u8_range(&kvl.k, off_tok * kvl.k_tok_bytes,
9210 (off_tok + t_kv) * kvl.k_tok_bytes);
9211 let v_view = e.view_u8_range(&kvl.v, off_tok * kvl.v_tok_bytes,
9212 (off_tok + t_kv) * kvl.v_tok_bytes);
9213 let qv = e.view(&q, t * nh * hd);
9214 let q_row = qv.slice(i * nh * hd..(i + 1) * nh * hd);
9215 let mut q_one = e.uninit(nh * hd)?;
9216 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
9217 let mut a_one = e.uninit(nh * hd)?;
9218 e.fa_decode_kvmod(&q_one, &k_view, &v_view, &mut a_one, hd, nh, nkv, t_kv, scale,
9222 kvl.k_tok_bytes, kvl.v_tok_bytes,
9223 (!swa && crate::Engine::gkv_on())
9224 || (swa && crate::Engine::wkv_on()))?;
9225 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
9226 }
9227 Ok(e.matmul(&fa.wo, &attn, t)?)
9228 }
9229
9230 fn gemma4_e4b_trunk(&self, e: &Engine, tokens: &[u32], pos0: usize, cache: &mut Cache,
9235 head_last: bool)
9236 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
9237 let n_embd = self.cfg.n_embd as usize;
9238 let t = tokens.len();
9239 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
9240 let pos_d = e.htod_i32(&pos)?;
9241 let mut x = e.htod(&self.embd.gather(n_embd, tokens))?;
9242 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
9243 let inp_pl = self.gemma4_e4b_inp_pl(e, tokens, &x, t)?;
9244 self.gemma4_e4b_trunk_core(e, x, inp_pl, &pos_d, t, cache, None, true, head_last)
9245 }
9246
9247 fn gemma4_e4b_trunk_core(&self, e: &Engine, x_in: CudaSlice<f32>, inp_pl: CudaSlice<f32>,
9251 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache,
9252 dc_bucket: Option<usize>, cap_logits: bool, head_last: bool)
9253 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
9254 let n_embd = self.cfg.n_embd as usize;
9255 let eps = self.cfg.rms_eps;
9256 let n_layer = self.layers.len();
9257 let mut x = x_in;
9258 let aux_e4b = self.gemma4_aux.as_ref().unwrap().e4b.as_ref().unwrap();
9259 let n_epl = aux_e4b.n_epl;
9260
9261 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
9267 for il in 0..n_layer {
9268 let layer = &self.layers[il];
9269 let (hq, hdq) = match h_carry.take() {
9270 Some(p) => p,
9271 None => e.rms_norm_q8_1(&x, layer.attn_norm.float_data(), n_embd, t, eps)?,
9272 };
9273 let o = self.gemma4_e4b_attn(e, il, &hq, &hdq, pos_d, t, cache, dc_bucket)?;
9274 let bits = layer.gemma4.as_ref().unwrap();
9277 let e4b = bits.e4b.as_ref().expect("e4b layer bits");
9278 let fuse_exit = e.uses_q8_1_fast(&e4b.inp_gate);
9289 let (sn, attn_out) = self.gemma4_layer_tail_core_pn(
9290 e, layer, &o, &x, t, Some(layer.post_attn_norm.float_data()), fuse_exit)?;
9291 let mut resid = e.uninit(t * n_embd)?;
9292 let g = if fuse_exit {
9298 let (rq, rd) = e.rms_pre_add_q8_1(&sn, bits.post_ffw_norm.float_data(),
9300 &attn_out, &mut resid, n_embd, t,
9301 self.cfg.rms_eps)?;
9302 e.matmul_pre(&e4b.inp_gate, &rq, &rd, &resid, t)?
9303 } else {
9304 e.add(&sn, &attn_out, &mut resid, t * n_embd)?;
9305 e.matmul(&e4b.inp_gate, &resid, t)?
9306 };
9307 let mut act = e.uninit(t * n_epl)?;
9308 let y = if t == 1 && e.uses_q8_1_fast(&e4b.proj) {
9309 let ipv = e.view(&inp_pl, n_epl * n_layer);
9310 let row = ipv.slice(il * n_epl..(il + 1) * n_epl);
9311 let (aq, ad) = e.gelu_tanh_mul_q8_1(&g, &row, &mut act, n_epl, 1)?;
9312 e.matmul_pre(&e4b.proj, &aq, &ad, &act, t)?
9313 } else {
9314 let mut inp_this = e.uninit(t * n_epl)?;
9315 e.copy_rows_strided(&inp_pl, &mut inp_this, n_epl, t, n_epl * n_layer,
9316 il * n_epl)?;
9317 e.gelu_tanh_mul(&g, &inp_this, &mut act, t * n_epl)?;
9318 e.matmul(&e4b.proj, &act, t)?
9319 };
9320 let next_norm = if il + 1 < n_layer {
9323 self.layers[il + 1].attn_norm.float_data()
9324 } else {
9325 self.output_norm.float_data()
9326 };
9327 let mut xn = e.uninit(t * n_embd)?;
9328 let pair = e.rms_pre_add_scale_rms_norm_q8_1(&y, e4b.post_norm.float_data(),
9329 &resid, bits.layer_scale, next_norm,
9330 &mut xn, n_embd, t, eps)?;
9331 h_carry = Some(pair);
9332 x = xn;
9333 }
9334 let (oq, odq) = h_carry.take().unwrap();
9338 let h0 = e.zeros(0)?;
9339 let hm = if head_last { 1 } else { t };
9340 let (hq, hd) = if head_last && t > 1 {
9341 let mut q1 = e.uninit_i8(n_embd)?;
9342 e.dtod_copy_view_i8(&oq.slice((t - 1) * n_embd..t * n_embd), &mut q1)?;
9343 let nb = n_embd / 32;
9344 let mut d1 = e.uninit(nb)?;
9345 e.dtod_copy_view(&odq.slice((t - 1) * nb..t * nb), &mut d1)?;
9346 (q1, d1)
9347 } else {
9348 (oq, odq)
9349 };
9350 let mut ld = e.matmul_pre(&self.output, &hq, &hd, &h0, hm)?;
9351 if cap_logits {
9355 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
9356 e.softcap(&mut ld, cap, hm * self.output.out_features())?;
9357 }
9358 self.gemma4_suppress(e, &mut ld, hm)?; Ok((ld, x))
9360 }
9361
9362 pub fn gemma4_e4b_decode_step_t_am_dev(&self, e: &Engine, tok_d: &CudaSlice<u32>,
9369 t: usize, pos0: usize, cache: &mut Cache)
9370 -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
9371 let n_embd = self.cfg.n_embd as usize;
9372 let eps = self.cfg.rms_eps;
9373 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
9374 let pos_d = e.htod_i32(&pos)?;
9375 let embd_gpu = self.embd_gpu.get_or_init(|| {
9376 e.upload_u8(&self.embd.raw).expect("embed table upload")
9377 });
9378 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
9379 let mut x = e.embed_gather_device_td(embd_gpu, tok_d, t, n_embd, qt, rb)?;
9380 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
9381 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, tok_d, &x, t)?;
9382 let (ld, xp) = self.gemma4_e4b_trunk_core(e, x, inp_pl, &pos_d, t, cache, None, true,
9383 false)?;
9384 let n_vocab = self.output.out_features();
9387 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
9388 for i in 0..t {
9389 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
9390 }
9391 let mut hn = e.uninit(t * n_embd)?;
9392 e.rms_norm(&xp, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
9393 cache.pos += t;
9394 Ok((vam, hn))
9395 }
9396
9397 pub(crate) fn gemma4_e4b_decode_step_t_h(&self, e: &Engine, tokens: &[u32], pos0: usize,
9400 cache: &mut Cache)
9401 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
9402 let n_embd = self.cfg.n_embd as usize;
9403 let eps = self.cfg.rms_eps;
9404 let t = tokens.len();
9405 let (ld, xp) = self.gemma4_e4b_trunk(e, tokens, pos0, cache, false)?;
9406 let mut hn = e.uninit(t * n_embd)?;
9407 e.rms_norm(&xp, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
9408 cache.pos += t;
9409 Ok((e.dtoh(&ld)?, hn))
9410 }
9411
9412 pub fn gemma4_e4b_decode_step_dcg(&self, e: &Engine, token_d: &mut CudaSlice<u32>,
9418 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
9419 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
9420 n_vocab: usize, bucket: usize)
9421 -> Result<(), Box<dyn std::error::Error>> {
9422 let n_embd = self.cfg.n_embd as usize;
9423 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
9424 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
9425 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, token_d, &x, 1)?;
9426 let (ld, _x) = self.gemma4_e4b_trunk_core(e, x, inp_pl, pos_d, 1, cache, Some(bucket),
9427 false, false)?;
9428 e.argmax_token_device_into(&ld, token_d, n_vocab)?;
9429 e.inc_seqlen(pos_d)?;
9430 Ok(())
9431 }
9432
9433 #[allow(clippy::too_many_arguments)]
9441 pub fn gemma4_e4b_decode_step_dc(&self, e: &Engine, token_d: &CudaSlice<u32>,
9442 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
9443 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
9444 n_vocab: usize)
9445 -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
9446 let n_embd = self.cfg.n_embd as usize;
9447 let eps = self.cfg.rms_eps;
9448 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
9449 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
9450 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, token_d, &x, 1)?;
9451 let (ld, _x) = self.gemma4_e4b_trunk_core(e, x, inp_pl, pos_d, 1, cache, None, false,
9452 false)?;
9453 let mut tok_out = e.stream().alloc_zeros::<u32>(1)?;
9454 e.argmax_token_device_into(&ld, &mut tok_out, n_vocab)?;
9455 e.inc_seqlen(pos_d)?;
9456 cache.pos += 1;
9457 let _ = eps;
9458 Ok(tok_out)
9459 }
9460
9461 pub(crate) fn gemma4_e4b_decode_step_h(&self, e: &Engine, token: u32, cache: &mut Cache)
9464 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
9465 let (ld, x) = self.gemma4_e4b_trunk(e, &[token], cache.pos, cache, false)?;
9466 let logits = e.dtoh(&ld)?;
9467 cache.pos += 1;
9468 Ok((logits, x))
9469 }
9470
9471 pub(crate) fn gemma4_e4b_prime(&self, e: &Engine, tokens: &[u32], cache: &mut Cache)
9475 -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
9476 if cache.pos != 0 {
9479 return Err("e4b prime is fresh-prompt only (v0) — prime the full prompt in one \
9480 call or decode tokenwise".into());
9481 }
9482 let n_embd = self.cfg.n_embd as usize;
9483 let t = tokens.len();
9484 let (ld, x) = self.gemma4_e4b_trunk(e, tokens, 0, cache, true)?;
9485 cache.pos += t;
9486 let last = e.dtoh(&ld)?; let xv = e.view(&x, t * n_embd);
9488 let row = xv.slice((t - 1) * n_embd..t * n_embd);
9489 let mut h_seed = e.uninit(n_embd)?;
9490 e.copy_view_into(&mut h_seed, 0, &row, n_embd)?;
9491 Ok((last, h_seed, x))
9492 }
9493
9494 pub(crate) fn gemma4_e4b_forward(&self, e: &Engine, tokens: &[u32], last_only: bool)
9496 -> Result<Vec<f32>, Box<dyn std::error::Error>> {
9497 let mut cache = Cache::new(e, &self.cfg, tokens.len() + 8)?;
9498 let (ld, _x) = self.gemma4_e4b_trunk(e, tokens, 0, &mut cache, last_only)?;
9499 Ok(e.dtoh(&ld)?) }
9501}
9502
9503#[cfg(test)]
9504mod prime_chunk_schedule_tests {
9505 use super::{
9506 dynamic_prime_chunk_ranges, fixed_prime_chunk_ranges, fixed_prime_chunk_ranges_for_ring,
9507 PRIME_MIN_T,
9508 PRIME_PIPE_MIN_CHUNK,
9509 };
9510
9511 fn sizes(ranges: &[(usize, usize)]) -> Vec<usize> {
9512 ranges.iter().map(|(start, end)| end - start).collect()
9513 }
9514
9515 fn auto_chunk(t: usize) -> usize {
9516 t.div_ceil(8).max(PRIME_PIPE_MIN_CHUNK).min(4096)
9517 }
9518
9519 #[test]
9520 fn fixed_schedule_retains_measured_geometry() {
9521 assert_eq!(
9522 sizes(&fixed_prime_chunk_ranges(461, 128)),
9523 vec![128, 128, 128, 77]
9524 );
9525 assert_eq!(
9526 sizes(&fixed_prime_chunk_ranges(1833, 230)),
9527 vec![230, 230, 230, 230, 230, 230, 230, 223]
9528 );
9529 assert_eq!(
9530 sizes(&fixed_prime_chunk_ranges(4096, 512)),
9531 vec![512; 8]
9532 );
9533 let capped = sizes(&fixed_prime_chunk_ranges_for_ring(8200, 4096, true));
9534 assert_eq!(capped, vec![4096, 4088, 16]);
9535 assert!(capped.iter().all(|&rows| rows <= 4096));
9536 assert_eq!(
9537 sizes(&fixed_prime_chunk_ranges_for_ring(4100, 4096, false)),
9538 vec![4100],
9539 "flag-off schedule remains byte-for-byte the legacy monolithic tail",
9540 );
9541 }
9542
9543 #[test]
9544 fn dynamic_schedule_matches_registered_shapes() {
9545 let cases = [
9546 (461, vec![64, 141, 132, 124]),
9547 (1833, vec![115, 269, 260, 252, 244, 237, 231, 225]),
9548 (4096, vec![256, 602, 580, 563, 545, 531, 516, 503]),
9549 ];
9550 for (t, expected) in cases {
9551 let chunk = auto_chunk(t);
9552 let fixed = fixed_prime_chunk_ranges(t, chunk);
9553 assert_eq!(
9554 sizes(&dynamic_prime_chunk_ranges(t, chunk, &fixed)),
9555 expected
9556 );
9557 }
9558 }
9559
9560 #[test]
9561 fn dynamic_schedule_covers_exactly_and_shrinks_after_fill() {
9562 for t in 256..=8192 {
9563 let chunk = auto_chunk(t);
9564 let fixed = fixed_prime_chunk_ranges(t, chunk);
9565 let dynamic = dynamic_prime_chunk_ranges(t, chunk, &fixed);
9566 assert_eq!(dynamic.len(), fixed.len(), "T={t}");
9567 assert_eq!(dynamic.first().unwrap().0, 0, "T={t}");
9568 assert_eq!(dynamic.last().unwrap().1, t, "T={t}");
9569 for pair in dynamic.windows(2) {
9570 assert_eq!(pair[0].1, pair[1].0, "T={t}");
9571 }
9572 assert!(
9573 dynamic
9574 .iter()
9575 .all(|(start, end)| end - start >= PRIME_MIN_T),
9576 "T={t} sizes={:?}",
9577 sizes(&dynamic)
9578 );
9579 if dynamic.len() >= 3 {
9580 let chunk_sizes = sizes(&dynamic);
9581 assert!(
9582 chunk_sizes[0] < chunk_sizes[1],
9583 "T={t} sizes={chunk_sizes:?}"
9584 );
9585 assert!(
9586 chunk_sizes[1..].windows(2).all(|pair| pair[0] >= pair[1]),
9587 "T={t} sizes={chunk_sizes:?}"
9588 );
9589 }
9590 }
9591 }
9592}
9593
9594#[cfg(test)]
9595mod page_prefetch_tests {
9596 use super::{
9597 grouped_worker_prefetch_position, page_prefetch_positions,
9598 page_prefetch_window_from_values, worker_prefetch_positions,
9599 };
9600
9601 #[test]
9602 fn page_prefetch_window_keeps_existing_opt_in_default() {
9603 assert_eq!(page_prefetch_window_from_values(false, None), 0);
9604 assert_eq!(page_prefetch_window_from_values(false, Some("8")), 0);
9605 assert_eq!(page_prefetch_window_from_values(true, None), 1);
9606 assert_eq!(page_prefetch_window_from_values(true, Some("bad")), 1);
9607 assert_eq!(page_prefetch_window_from_values(true, Some("0")), 0);
9608 assert_eq!(page_prefetch_window_from_values(true, Some("8")), 8);
9609 }
9610
9611 #[test]
9612 fn rolling_page_prefetch_advises_each_future_expert_once() {
9613 let advised: Vec<_> = (0..7)
9614 .flat_map(|position| page_prefetch_positions(position, 7, 3))
9615 .collect();
9616 assert_eq!(advised, vec![1, 2, 3, 4, 5, 6]);
9617
9618 let one_ahead: Vec<_> = (0..4)
9619 .flat_map(|position| page_prefetch_positions(position, 4, 1))
9620 .collect();
9621 assert_eq!(one_ahead, vec![1, 2, 3]);
9622 assert!(page_prefetch_positions(0, 4, 0).is_empty());
9623 }
9624
9625 #[test]
9626 fn grouped_worker_prefetch_primes_first_then_each_known_next_once() {
9627 assert_eq!(grouped_worker_prefetch_position(0, None), None);
9628 let positions: Vec<_> = std::iter::once(grouped_worker_prefetch_position(4, None).unwrap())
9629 .chain((0..4).filter_map(|position| {
9630 grouped_worker_prefetch_position(4, Some(position))
9631 }))
9632 .collect();
9633 assert_eq!(positions, vec![0, 1, 2, 3]);
9634 assert_eq!(grouped_worker_prefetch_position(1, Some(0)), None);
9635 }
9636
9637 #[test]
9638 fn rolling_worker_prefetch_primes_current_and_each_future_expert_once() {
9639 let queued: Vec<_> = (0..8)
9640 .flat_map(|position| worker_prefetch_positions(position, 8, 5))
9641 .collect();
9642 assert_eq!(queued, (0..8).collect::<Vec<_>>());
9643
9644 let one_at_a_time: Vec<_> = (0..4)
9645 .flat_map(|position| worker_prefetch_positions(position, 4, 1))
9646 .collect();
9647 assert_eq!(one_at_a_time, vec![0, 1, 2, 3]);
9648 assert!(worker_prefetch_positions(0, 4, 0).is_empty());
9649 }
9650}
9651
9652pub struct G4DcSlots {
9653 x: CudaSlice<f32>, xn: CudaSlice<f32>, cur: CudaSlice<f32>,
9654 hq: CudaSlice<i8>, hd_: CudaSlice<f32>,
9655 q0: CudaSlice<f32>, k0: CudaSlice<f32>, v0: CudaSlice<f32>,
9656 q: CudaSlice<f32>, k: CudaSlice<f32>, v: CudaSlice<f32>,
9657 attn: CudaSlice<f32>, o: CudaSlice<f32>,
9658 attn_out: CudaSlice<f32>, zsh: CudaSlice<f32>,
9659 zq: CudaSlice<i8>, zd: CudaSlice<f32>,
9660 gate: CudaSlice<f32>, up: CudaSlice<f32>,
9661 act: CudaSlice<f32>, actq: CudaSlice<i8>, actd: CudaSlice<f32>,
9662 f0: CudaSlice<f32>, sn: CudaSlice<f32>,
9663 hn: CudaSlice<f32>, logits: CudaSlice<f32>,
9664}