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 crate::pp::pp_host_bounce_active()
998 && (self.cfg.gemma4.is_some() || !crate::pp::prime_pp_on())
999 {
1000 return Err(
1001 "prime_chunk: refused with MEMRA_PP_HOST_BOUNCE=1 because this configuration \
1002 has no active prime stage split and would peer-read remote weights; keep \
1003 MEMRA_PRIME_PP enabled and use a PP-prime-supported model"
1004 .into(),
1005 );
1006 }
1007 if self.cfg.gemma4.is_none()
1016 && !crate::pp::pp2_streams_off()
1017 && crate::pp::prime_pp_on()
1018 {
1019 if let Some(fence) = crate::pp::pp_cuts(self.layers.len()) {
1020 return self.prime_chunk_ppn(e, tokens, cache, seq_end, &fence);
1021 }
1022 }
1023 if crate::pp::pp_host_bounce_active() {
1024 return Err(
1025 "prime_chunk: MEMRA_PP_HOST_BOUNCE=1 found no valid prime stage split; \
1026 refusing an unsplit remote-weight walk"
1027 .into(),
1028 );
1029 }
1030 let t = tokens.len();
1031 let base = cache.pos;
1032 debug_assert!(seq_end >= base + t, "prime_chunk: seq_end must cover this chunk");
1033 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
1034 let pos_d = e.htod_i32(&pos)?;
1035
1036 let x_embed = self.embed(e, tokens)?; let x = self.prime_layers(
1038 e, x_embed, 0, self.layers.len(), &pos_d, t, base, cache, seq_end,
1039 )?;
1040 self.prime_chunk_epilogue(e, x, t, cache)
1041 }
1042
1043 #[allow(clippy::too_many_arguments)]
1059 fn prime_layers(&self, e: &Engine, x_in: CudaSlice<f32>, lo: usize, hi: usize,
1060 pos_d: &CudaSlice<i32>, t: usize, base: usize, cache: &mut Cache,
1061 seq_end: usize)
1062 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1063 let cfg = &self.cfg;
1064 let n_embd = cfg.n_embd as usize;
1065 let eps = cfg.rms_eps;
1066 let f16fuse = crate::f16_ffi::pp_f16_enabled() && t >= 16;
1070 let n_ff_max = self.layers.iter().map(|l| match &l.ffn {
1076 crate::hybrid::Ffn::Dense { ffn_gate, .. } => ffn_gate.out_features(),
1077 _ => n_embd,
1078 }).max().unwrap_or(n_embd).max(n_embd);
1079 let use_slabs = std::env::var("MEMRA_PRIME_SLABS").as_deref() != Ok("0");
1080 let slab = if use_slabs {
1081 Some(self.prime_slabs_get(e, t, n_embd, n_ff_max)?)
1082 } else {
1083 None
1084 };
1085 let mut slab_guard = slab.as_ref().map(|sl| sl.lock().unwrap());
1086 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>);
1088 let (mut x_cur, mut x_nxt, sl): (&mut CudaSlice<f32>, &mut CudaSlice<f32>, Option<SlabRefs>);
1089 let mut seg: Option<(&mut Vec<Option<cudarc::driver::CudaGraph>>, &mut Vec<Option<cudarc::driver::CudaGraph>>, &mut CudaSlice<f32>, &mut usize)> = None;
1090 let mut x_own2;
1091 match slab_guard.as_mut() {
1092 Some(g) => {
1093 let slabs = &mut **g;
1094 e.copy_into(&mut slabs.xa, 0, &x_in, t * n_embd)?;
1095 let PrimeSlabs { xa, xb, h, x1, z, act, h16, z16, gate, up, ffn_out, seg_glue, mixed, seg_mid, seg_t, .. } = slabs;
1096 x_cur = xa;
1097 x_nxt = xb;
1098 seg = Some((seg_glue, seg_mid, mixed, seg_t));
1099 sl = Some((h, x1, z, act, h16, z16, gate, up, ffn_out));
1100 }
1101 None => {
1102 x_own = x_in;
1103 x_own2 = e.uninit(t * n_embd)?;
1104 x_cur = &mut x_own;
1105 x_nxt = &mut x_own2;
1106 sl = None;
1107 }
1108 }
1109 let mut alloc_h; let mut alloc_x1; let mut alloc_z; let mut alloc_act;
1110 let mut alloc_h16; let mut alloc_z16;
1111 let mut alloc_gate; let mut alloc_up; let mut alloc_fo;
1112 let (h, x1, z, act): (&mut CudaSlice<f32>, &mut CudaSlice<f32>, &mut CudaSlice<f32>, &mut CudaSlice<f32>);
1113 let (h16, z16): (&mut CudaSlice<u8>, &mut CudaSlice<u8>);
1114 let (sl_gate, sl_up, sl_fo): (&mut CudaSlice<f32>, &mut CudaSlice<f32>, &mut CudaSlice<f32>);
1115 match sl {
1116 Some((a, b, c, d, e16, f16b, g, u, fo)) => {
1117 h = a; x1 = b; z = c; act = d; h16 = e16; z16 = f16b;
1118 sl_gate = g; sl_up = u; sl_fo = fo;
1119 }
1120 None => {
1121 alloc_h = e.uninit(t * n_embd)?;
1122 alloc_x1 = e.uninit(t * n_embd)?;
1123 alloc_z = e.uninit(t * n_embd)?;
1124 alloc_act = e.uninit(t * n_ff_max)?;
1125 alloc_h16 = e.alloc_u8_uninit(t * n_embd * 2)?;
1126 alloc_z16 = e.alloc_u8_uninit(t * n_embd * 2)?;
1127 alloc_gate = e.uninit(t * n_ff_max)?;
1128 alloc_up = e.uninit(t * n_ff_max)?;
1129 alloc_fo = e.uninit(t * n_embd)?;
1130 h = &mut alloc_h; x1 = &mut alloc_x1; z = &mut alloc_z; act = &mut alloc_act;
1131 h16 = &mut alloc_h16; z16 = &mut alloc_z16;
1132 sl_gate = &mut alloc_gate; sl_up = &mut alloc_up; sl_fo = &mut alloc_fo;
1133 }
1134 }
1135 let n_layers = self.layers.len();
1140 let use_seg = f16fuse && seg.is_some() && self.cfg.step35.is_none()
1150 && lo == 0 && hi == n_layers
1151 && std::env::var("MEMRA_PRIME_SEG").as_deref() == Ok("1");
1152 if let Some((sg, sm, _, st)) = seg.as_mut() {
1153 if **st != t {
1154 sg.clear();
1155 sg.extend((0..n_layers).map(|_| None));
1156 sm.clear();
1157 sm.extend((0..n_layers).map(|_| None));
1158 **st = t;
1159 }
1160 }
1161 {
1162 let layer_lo = &self.layers[lo];
1163 if f16fuse {
1164 e.rms_norm_f16out(x_cur, layer_lo.attn_norm.float_data(), h, h16, n_embd, t, eps)?;
1165 } else {
1166 e.rms_norm(x_cur, layer_lo.attn_norm.float_data(), h, n_embd, t, eps)?;
1167 }
1168 }
1169 for il in lo..hi {
1170 let layer = &self.layers[il];
1171 let hx16 = if f16fuse { Some(&*h16) } else { None };
1172 if use_seg {
1173 let (pre, pre16, w_out) = match &layer.mixer {
1176 Mixer::Full(fa) => {
1177 let g3 = match hx16 {
1178 Some(xh) => e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], h, xh, t)?,
1179 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
1180 };
1181 let (pre, pre16) = self.full_attn_prime_core_inner(e, fa, g3, &pos_d, t, cache, il)?;
1182 (pre, pre16, &fa.wo)
1183 }
1184 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
1185 Mixer::Linear(la) => {
1186 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
1187 let g4 = match hx16 {
1188 Some(xh) => e.matmul_group_xh(&ws, h, xh, t)?,
1189 None => e.matmul_group(&ws, h, t)?,
1190 };
1191 let (pre, pre16) = self.linear_attn_prime_core_pad_inner(e, la, g4, t, cache, il, None)?;
1192 (pre, pre16, &la.ssm_out)
1193 }
1194 };
1195 {
1196 let (_, sm, mslab, _) = seg.as_mut().unwrap();
1197 let pre_n = pre.len() / t;
1198 let xh_pre = match pre16 {
1199 Some(x) => x,
1200 None => e.f16_act(&pre, t * pre_n, pre_n)?,
1201 };
1202 if !e.try_f16_gemm_pre_into(w_out, &xh_pre, t, mslab)? {
1203 let y = e.matmul(w_out, &pre, t)?;
1204 e.copy_into(mslab, 0, &y, t * n_embd)?;
1205 }
1206 if sm[il].is_none() {
1207 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
1208 let w_post = layer.post_attn_norm.float_data();
1209 e.stream().synchronize()?;
1210 e.stream().begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
1211 let r = (|| -> Result<(), Box<dyn std::error::Error>> {
1212 e.add(x_cur, mslab, x1, t * n_embd)?;
1213 e.rms_norm_f16out(x1, w_post, z, z16, n_embd, t, eps)?;
1214 Ok(())
1215 })();
1216 let g = e.stream().end_capture(
1217 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH);
1218 r?;
1219 sm[il] = Some(g?.ok_or("S-mid capture produced no graph")?);
1220 }
1221 sm[il].as_ref().unwrap().launch()?;
1222 }
1223 } else {
1224 let mixed = match &layer.mixer {
1225 Mixer::Full(fa) => self.full_attn_prime(e, fa, h, hx16, &pos_d, t, cache, il,
1226 seq_end)?,
1227 Mixer::Linear(la) => self.linear_attn_prime(e, la, h, hx16, t, cache, il)?,
1228 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
1229 };
1230 if f16fuse {
1231 e.add_rms_norm_f16out(x_cur, &mixed, layer.post_attn_norm.float_data(),
1234 x1, z, z16, n_embd, t, eps)?;
1235 } else {
1236 e.add(x_cur, &mixed, x1, t * n_embd)?;
1237 e.rms_norm(x1, layer.post_attn_norm.float_data(), z, n_embd, t, eps)?;
1238 }
1239 }
1240 let zx16 = if f16fuse { Some(&*z16) } else { None };
1241 match &layer.ffn {
1242 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
1243 let n_ff = ffn_gate.out_features();
1244 let mut into_ok = false;
1247 if let Some(xh) = zx16 {
1248 into_ok = e.try_f16_gemm_pre_into(ffn_gate, xh, t, sl_gate)?
1249 && e.try_f16_gemm_pre_into(ffn_up, xh, t, sl_up)?;
1250 }
1251 if !into_ok {
1252 let mut g2 = match zx16 {
1253 Some(xh) => e.matmul_group_xh(&[ffn_gate, ffn_up], z, xh, t)?,
1254 None => e.matmul_group(&[ffn_gate, ffn_up], z, t)?,
1255 };
1256 let up_y = g2.pop().unwrap();
1257 let gate_y = g2.pop().unwrap();
1258 e.copy_into(sl_gate, 0, &gate_y, t * n_ff)?;
1259 e.copy_into(sl_up, 0, &up_y, t * n_ff)?;
1260 }
1261 let d_lim = self.cfg.clamp_shexp_at(il as u32);
1266 let act16 = if Self::f16out_on(e, t) && self.cfg.m3.is_none()
1267 && d_lim.is_none() {
1268 let mut a16 = e.alloc_u8_uninit(t * n_ff * 2)?;
1269 e.silu_mul_f16out(sl_gate, sl_up, act, &mut a16, t * n_ff)?;
1270 Some(a16)
1271 } else {
1272 Self::ffn_act_lim(e, &self.cfg, sl_gate, sl_up, 1.0, 1.0, d_lim,
1273 act, t * n_ff)?;
1274 None
1275 };
1276 let xh_act = match act16 {
1278 Some(x) => x,
1279 None => e.f16_act(act, t * n_ff, n_ff)?,
1280 };
1281 if !e.try_f16_gemm_pre_into(ffn_down, &xh_act, t, sl_fo)? {
1282 let y = e.matmul(ffn_down, &*act, t)?;
1283 e.copy_into(sl_fo, 0, &y, t * n_embd)?;
1284 }
1285 }
1286 crate::hybrid::Ffn::Moe(m) => {
1287 let y = self.moe_ffn_il_prefill(e, m, z, t, il as u16)?;
1288 e.copy_into(sl_fo, 0, &y, t * n_embd)?;
1289 }
1290 }
1291 if use_seg && il + 1 < hi {
1292 let w_next = self.layers[il + 1].attn_norm.float_data();
1294 let (sg, _, _, _) = seg.as_mut().unwrap();
1295 if sg[il].is_none() {
1296 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
1297 e.stream().synchronize()?;
1298 e.stream().begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
1299 let r = (|| -> Result<(), Box<dyn std::error::Error>> {
1300 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
1301 e.rms_norm_f16out(x_nxt, w_next, h, h16, n_embd, t, eps)?;
1302 Ok(())
1303 })();
1304 let g = e.stream().end_capture(
1305 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH);
1306 r?;
1307 sg[il] = Some(g?.ok_or("S-glue capture produced no graph")?);
1308 }
1309 sg[il].as_ref().unwrap().launch()?;
1310 } else {
1311 if il + 1 < hi {
1312 let w_next = self.layers[il + 1].attn_norm.float_data();
1313 if f16fuse {
1314 e.add_rms_norm_f16out(x1, sl_fo, w_next, x_nxt, h, h16, n_embd, t, eps)?;
1315 } else {
1316 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
1317 e.rms_norm(x_nxt, w_next, h, n_embd, t, eps)?;
1318 }
1319 } else {
1320 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
1321 }
1322 }
1323 if let Some(path) = Self::prime_trace_path() {
1329 let row = (base + t - 1) as usize;
1330 let host = e.dtoh(x_nxt)?;
1331 let last = &host[(t - 1) * n_embd..t * n_embd];
1332 use std::io::Write as _;
1333 let mut f = std::fs::OpenOptions::new().create(true).append(true).open(path)?;
1334 let mut h64: u64 = 0xcbf29ce484222325;
1335 for v in last {
1336 h64 ^= v.to_bits() as u64;
1337 h64 = h64.wrapping_mul(0x100000001b3);
1338 }
1339 writeln!(f, "{{\"pos\":{row},\"layer\":{il},\"t\":{t},\"base\":{base},\
1340 \"hash\":\"{h64:016x}\",\"v0\":{:.9e},\"v1\":{:.9e},\"v2\":{:.9e}}}",
1341 last[0], last[1], last[2])?;
1342 }
1343 std::mem::swap(&mut x_cur, &mut x_nxt);
1344 }
1345 let mut x = e.uninit(t * n_embd)?;
1347 e.copy_into(&mut x, 0, x_cur, t * n_embd)?;
1348 drop(slab_guard);
1349 Ok(x)
1350 }
1351
1352 fn prime_chunk_epilogue(&self, e: &Engine, x: CudaSlice<f32>, t: usize, cache: &mut Cache)
1357 -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
1358 let n_embd = self.cfg.n_embd as usize;
1359 let eps = self.cfg.rms_eps;
1360 let mut h_seed = e.uninit(n_embd)?;
1364 if !crate::spec::spec_hpost() {
1365 e.copy_view_into(&mut h_seed, 0, &x.slice((t - 1) * n_embd..t * n_embd), n_embd)?;
1366 }
1367 let mut hn = e.uninit(t * n_embd)?;
1369 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
1370 if crate::spec::spec_hpost() {
1371 e.copy_view_into(&mut h_seed, 0, &hn.slice((t - 1) * n_embd..t * n_embd), n_embd)?;
1372 }
1373 let last = e.view(&hn, t * n_embd);
1374 let last_row = last.slice((t - 1) * n_embd..t * n_embd);
1375 let mut hlast = e.uninit(n_embd)?;
1376 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
1377 let logits = e.matmul(&self.output, &hlast, 1)?;
1378 cache.pos += t;
1379 Ok((e.dtoh(&logits)?, h_seed, if crate::spec::spec_hpost() { hn } else { x }))
1382 }
1383
1384 fn prime_chunk_ppn(&self, e: &Engine, tokens: &[u32], cache: &mut Cache, seq_end: usize,
1408 fence: &[usize])
1409 -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
1410 let rt = crate::pp::PpNRt::get(e)?;
1411 let n_st = fence.len() - 1;
1412 assert_eq!(
1413 rt.n_stages(), n_st,
1414 "PpNRt stage count {} != fence stages {n_st}", rt.n_stages()
1415 );
1416 let n_embd = self.cfg.n_embd as usize;
1417 let t = tokens.len();
1418 let base = cache.pos;
1419 debug_assert!(seq_end >= base + t, "prime_chunk_ppn: seq_end must cover this chunk");
1420 let payload = t * n_embd;
1421 let caller_stream = e.stream();
1425 rt.fence_stages_behind(&caller_stream)?;
1426
1427 if n_st == 2 {
1428 let slot = self.prime_pp2_stage0_enqueue(
1429 e, rt, tokens, cache, seq_end, fence, base, false,
1430 )?;
1431 let x = self.prime_pp2_stage1_enqueue(
1432 e, rt, slot, t, cache, seq_end, fence, base, false,
1433 )?;
1434 let out = {
1435 rt.bind_stage(1)?;
1436 let _st1 = rt.enter(1);
1437 let e1 = rt.engine(1, e);
1438 self.prime_chunk_epilogue(e1, x, t, cache)?
1439 };
1440 rt.publish_to(1, &caller_stream)?;
1441 crate::pp::PRIME_SPLIT_CHUNKS
1442 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
1443 return Ok(out);
1444 }
1445
1446 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
1447
1448 let mut slot = {
1450 let _st0 = rt.enter(0);
1451 let e0 = rt.engine(0, e);
1452 let pos_d = e0.htod_i32(&pos)?;
1453 let x = self.embed(e0, tokens)?;
1454 let x = self.prime_layers(
1455 e0, x, fence[0], fence[1], &pos_d, t, base, cache, seq_end,
1456 )?;
1457 rt.tx(0, &x, payload)?
1458 };
1460
1461 for s in 1..n_st - 1 {
1463 let _st = rt.enter(s);
1464 let es = rt.engine(s, e);
1465 let pos_d = es.htod_i32(&pos)?;
1466 let x = rt.rx(s - 1, slot, payload)?;
1467 let x = self.prime_layers(
1468 es, x, fence[s], fence[s + 1], &pos_d, t, base, cache, seq_end,
1469 )?;
1470 slot = rt.tx(s, &x, payload)?;
1471 }
1472
1473 let _stl = rt.enter(n_st - 1);
1475 let el = rt.engine(n_st - 1, e);
1476 let pos_d = el.htod_i32(&pos)?;
1477 let x = rt.rx(n_st - 2, slot, payload)?;
1478 let x = self.prime_layers(
1479 el, x, fence[n_st - 1], fence[n_st], &pos_d, t, base, cache, seq_end,
1480 )?;
1481 let out = self.prime_chunk_epilogue(el, x, t, cache)?;
1482 rt.publish_to(n_st - 1, &caller_stream)?;
1488 crate::pp::PRIME_SPLIT_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
1489 Ok(out)
1490 }
1491
1492 fn prime_pp2_stage0_enqueue(
1493 &self,
1494 e: &Engine,
1495 rt: &crate::pp::PpNRt,
1496 tokens: &[u32],
1497 cache: &mut Cache,
1498 seq_end: usize,
1499 fence: &[usize],
1500 base: usize,
1501 pipelined: bool,
1502 ) -> Result<usize, Box<dyn std::error::Error>> {
1503 let t = tokens.len();
1504 let n_embd = self.cfg.n_embd as usize;
1505 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
1506 rt.bind_stage(0)?;
1507 let _st0 = rt.enter(0);
1508 let e0 = rt.engine(0, e);
1509 let pos_d = e0.htod_i32(&pos)?;
1510 let x = self.embed(e0, tokens)?;
1511 let _overlap = pipelined.then(crate::pp::enter_prime_pipe_stage);
1512 let x = self.prime_layers(
1513 e0, x, fence[0], fence[1], &pos_d, t, base, cache, seq_end,
1514 )?;
1515 if pipelined {
1516 rt.tx_pipelined(0, &x, t * n_embd)
1517 } else {
1518 rt.tx(0, &x, t * n_embd)
1519 }
1520 }
1521
1522 fn prime_pp2_stage1_enqueue(
1523 &self,
1524 e: &Engine,
1525 rt: &crate::pp::PpNRt,
1526 slot: usize,
1527 t: usize,
1528 cache: &mut Cache,
1529 seq_end: usize,
1530 fence: &[usize],
1531 base: usize,
1532 pipelined: bool,
1533 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1534 let n_embd = self.cfg.n_embd as usize;
1535 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
1536 rt.bind_stage(1)?;
1537 let _st1 = rt.enter(1);
1538 let e1 = rt.engine(1, e);
1539 let pos_d = e1.htod_i32(&pos)?;
1540 let x = rt.rx(0, slot, t * n_embd)?;
1541 let _overlap = pipelined.then(crate::pp::enter_prime_pipe_stage);
1542 self.prime_layers(
1543 e1, x, fence[1], fence[2], &pos_d, t, base, cache, seq_end,
1544 )
1545 }
1546
1547 pub fn prime_chunk_captured(&self, e: &Engine, x_in: &CudaSlice<f32>, pos_d: &CudaSlice<i32>,
1563 t: usize, cache: &mut Cache,
1564 len_d: &CudaSlice<i32>,
1565 logits_out: &mut CudaSlice<f32>, h_seed_out: &mut CudaSlice<f32>)
1566 -> Result<(), Box<dyn std::error::Error>> {
1567 let cfg = &self.cfg;
1568 let n_embd = cfg.n_embd as usize;
1569 let eps = cfg.rms_eps;
1570 let f16fuse = crate::f16_ffi::pp_f16_enabled() && t >= 16;
1571 let mut x = e.uninit(t * n_embd)?;
1572 e.copy_into(&mut x, 0, x_in, t * n_embd)?;
1573 for (il, layer) in self.layers.iter().enumerate() {
1574 let mut h = e.uninit(t * n_embd)?;
1575 let mut hx16: Option<CudaSlice<u8>> = None;
1576 if f16fuse {
1577 let mut b16 = e.alloc_u8_uninit(t * n_embd * 2)?;
1578 e.rms_norm_f16out(&x, layer.attn_norm.float_data(), &mut h, &mut b16, n_embd, t, eps)?;
1579 hx16 = Some(b16);
1580 } else {
1581 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
1582 }
1583 let mixed = match &layer.mixer {
1584 Mixer::Full(fa) => self.full_attn_prime(e, fa, &h, hx16.as_ref(), pos_d, t, cache,
1588 il, t)?,
1589 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
1590 Mixer::Linear(la) => {
1591 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
1592 let g4 = match hx16.as_ref() {
1593 Some(xh) => e.matmul_group_xh(&ws, &h, xh, t)?,
1594 None => e.matmul_group(&ws, &h, t)?,
1595 };
1596 self.linear_attn_prime_core_pad(e, la, g4, t, cache, il, Some(len_d))?
1597 }
1598 };
1599 let mut x1 = e.uninit(t * n_embd)?;
1600 e.add(&x, &mixed, &mut x1, t * n_embd)?;
1601 let mut z = e.uninit(t * n_embd)?;
1602 let mut zx16: Option<CudaSlice<u8>> = None;
1603 if f16fuse {
1604 let mut b16 = e.alloc_u8_uninit(t * n_embd * 2)?;
1605 e.rms_norm_f16out(&x1, layer.post_attn_norm.float_data(), &mut z, &mut b16, n_embd, t, eps)?;
1606 zx16 = Some(b16);
1607 } else {
1608 e.rms_norm(&x1, layer.post_attn_norm.float_data(), &mut z, n_embd, t, eps)?;
1609 }
1610 let ffn_out = match &layer.ffn {
1611 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
1612 let n_ff = ffn_gate.out_features();
1613 let mut g2 = match &zx16 {
1614 Some(xh) => e.matmul_group_xh(&[ffn_gate, ffn_up], &z, xh, t)?,
1615 None => e.matmul_group(&[ffn_gate, ffn_up], &z, t)?,
1616 };
1617 let up = g2.pop().unwrap();
1618 let gate = g2.pop().unwrap();
1619 let mut act = e.uninit(t * n_ff)?;
1620 Self::ffn_act_lim(e, &self.cfg, &gate, &up, 1.0, 1.0,
1622 self.cfg.clamp_shexp_at(il as u32), &mut act, t * n_ff)?;
1623 e.matmul(ffn_down, &act, t)?
1624 }
1625 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il_prefill(e, m, &z, t, il as u16)?,
1626 };
1627 let mut x2 = e.uninit(t * n_embd)?;
1628 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
1629 x = x2;
1630 }
1631 if !crate::spec::spec_hpost() {
1633 e.row_gather_dev(&x, h_seed_out, len_d, n_embd)?;
1634 }
1635 let mut hn = e.uninit(t * n_embd)?;
1636 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
1637 if crate::spec::spec_hpost() {
1638 e.row_gather_dev(&hn, h_seed_out, len_d, n_embd)?;
1639 }
1640 let mut hlast = e.uninit(n_embd)?;
1641 e.row_gather_dev(&hn, &mut hlast, len_d, n_embd)?;
1642 let logits = e.matmul(&self.output, &hlast, 1)?;
1643 let nv = logits.len();
1644 e.copy_into(logits_out, 0, &logits, nv)?;
1645 Ok(())
1646 }
1647
1648 fn step35_prime_batch_on() -> bool {
1649 std::env::var("MEMRA_STEP35_PRIME_BATCH").as_deref() != Ok("0")
1650 }
1651
1652 #[allow(clippy::too_many_arguments)]
1655 fn step35_prime_batch_layers(
1656 &self,
1657 e: &Engine,
1658 mut x: CudaSlice<f32>,
1659 lo: usize,
1660 hi: usize,
1661 ts: &[usize],
1662 offs: &[usize],
1663 pos_ds: &[CudaSlice<i32>],
1664 caches: &mut [&mut Cache],
1665 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1666 let cfg = &self.cfg;
1667 let n_embd = cfg.n_embd as usize;
1668 let eps = cfg.rms_eps;
1669 let b = ts.len();
1670 let total: usize = ts.iter().sum();
1671 let f16fuse = crate::f16_ffi::pp_f16_enabled() && total >= 16;
1672
1673 let split = |e: &Engine, y: &CudaSlice<f32>, dim: usize|
1674 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
1675 let mut out = Vec::with_capacity(b);
1676 for s in 0..b {
1677 let mut ys = e.uninit(ts[s] * dim)?;
1678 e.copy_view_into(
1679 &mut ys,
1680 0,
1681 &y.slice(offs[s] * dim..(offs[s] + ts[s]) * dim),
1682 ts[s] * dim,
1683 )?;
1684 out.push(ys);
1685 }
1686 Ok(out)
1687 };
1688
1689 for il in lo..hi {
1690 let layer = &self.layers[il];
1691 let Mixer::Full(fa) = &layer.mixer else {
1692 return Err(format!("step35 layer {il} is not full-attn — corrupt config").into());
1693 };
1694
1695 let mut h = e.uninit(total * n_embd)?;
1696 let mut hx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
1697 if f16fuse {
1698 e.rms_norm_f16out(
1699 &x,
1700 layer.attn_norm.float_data(),
1701 &mut h,
1702 &mut hx16,
1703 n_embd,
1704 total,
1705 eps,
1706 )?;
1707 } else {
1708 e.rms_norm(
1709 &x,
1710 layer.attn_norm.float_data(),
1711 &mut h,
1712 n_embd,
1713 total,
1714 eps,
1715 )?;
1716 }
1717
1718 let gate_w = fa
1722 .attn_gate
1723 .as_ref()
1724 .ok_or("step35 layer is missing attn_gate.weight (head-wise attention gate)")?;
1725 let mut g4 = if f16fuse {
1726 e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv, gate_w], &h, &hx16, total)?
1727 } else {
1728 e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv, gate_w], &h, total)?
1729 };
1730 let gate = g4.pop().unwrap();
1731 let mut parts: Vec<Vec<CudaSlice<f32>>> =
1732 (0..b).map(|_| Vec::with_capacity(3)).collect();
1733 for (w, y) in [&fa.wq, &fa.wk, &fa.wv].iter().zip(g4) {
1734 for (s, ys) in split(e, &y, w.out_features())?.into_iter().enumerate() {
1735 parts[s].push(ys);
1736 }
1737 }
1738 let gates = split(e, &gate, gate_w.out_features())?;
1739 let geometry = self.step35_geom(il);
1740 let hd = geometry.head_dim_k as usize;
1741 let nh = geometry.n_head as usize;
1742 let mut ag_cat = e.uninit(total * nh * hd)?;
1743 for (s, (g3s, gate)) in parts.into_iter().zip(gates).enumerate() {
1744 let ag = self.step35_attn_pre_wo(
1745 e,
1746 fa,
1747 g3s,
1748 None,
1749 Some(&gate),
1750 &pos_ds[s],
1751 ts[s],
1752 Some(&mut *caches[s]),
1753 il,
1754 ts[s],
1755 )?;
1756 e.copy_into(
1757 &mut ag_cat,
1758 offs[s] * nh * hd,
1759 &ag,
1760 ts[s] * nh * hd,
1761 )?;
1762 }
1763 let mixed = e.matmul(&fa.wo, &ag_cat, total)?;
1764
1765 let mut x1 = e.uninit(total * n_embd)?;
1766 let mut z = e.uninit(total * n_embd)?;
1767 let mut zx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
1768 if f16fuse {
1769 e.add_rms_norm_f16out(
1770 &x,
1771 &mixed,
1772 layer.post_attn_norm.float_data(),
1773 &mut x1,
1774 &mut z,
1775 &mut zx16,
1776 n_embd,
1777 total,
1778 eps,
1779 )?;
1780 } else {
1781 e.add(&x, &mixed, &mut x1, total * n_embd)?;
1782 e.rms_norm(
1783 &x1,
1784 layer.post_attn_norm.float_data(),
1785 &mut z,
1786 n_embd,
1787 total,
1788 eps,
1789 )?;
1790 }
1791
1792 let ffn_out = match &layer.ffn {
1793 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
1794 let n_ff = ffn_gate.out_features();
1795 let mut g2 = if f16fuse {
1796 e.matmul_group_xh(&[ffn_gate, ffn_up], &z, &zx16, total)?
1797 } else {
1798 e.matmul_group(&[ffn_gate, ffn_up], &z, total)?
1799 };
1800 let up = g2.pop().unwrap();
1801 let gate = g2.pop().unwrap();
1802 let mut act = e.uninit(total * n_ff)?;
1803 let d_lim = cfg.clamp_shexp_at(il as u32);
1804 if Self::f16out_on(e, total) && cfg.m3.is_none() && d_lim.is_none() {
1805 let mut a16 = e.alloc_u8_uninit(total * n_ff * 2)?;
1806 e.silu_mul_f16out(&gate, &up, &mut act, &mut a16, total * n_ff)?;
1807 match e.try_f16_gemm_pre(ffn_down, &a16, total)? {
1808 Some(y) => y,
1809 None => e.matmul(ffn_down, &act, total)?,
1810 }
1811 } else {
1812 Self::ffn_act_lim(
1813 e,
1814 cfg,
1815 &gate,
1816 &up,
1817 1.0,
1818 1.0,
1819 d_lim,
1820 &mut act,
1821 total * n_ff,
1822 )?;
1823 e.matmul(ffn_down, &act, total)?
1824 }
1825 }
1826 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il(e, m, &z, total, il as u16)?,
1827 };
1828 let mut x2 = e.uninit(total * n_embd)?;
1829 e.add(&x1, &ffn_out, &mut x2, total * n_embd)?;
1830 x = x2;
1831 }
1832 Ok(x)
1833 }
1834
1835 fn step35_prime_batch_epilogue(
1836 &self,
1837 e: &Engine,
1838 x: CudaSlice<f32>,
1839 ts: &[usize],
1840 offs: &[usize],
1841 caches: &mut [&mut Cache],
1842 ) -> Result<Vec<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
1843 let n_embd = self.cfg.n_embd as usize;
1844 let total: usize = ts.iter().sum();
1845 let mut hn = e.uninit(total * n_embd)?;
1846 e.rms_norm(
1847 &x,
1848 self.output_norm.float_data(),
1849 &mut hn,
1850 n_embd,
1851 total,
1852 self.cfg.rms_eps,
1853 )?;
1854
1855 let hidden_src = if crate::spec::spec_hpost() { &hn } else { &x };
1856 let mut out = Vec::with_capacity(ts.len());
1857 for s in 0..ts.len() {
1858 let mut hidden = e.uninit(ts[s] * n_embd)?;
1859 e.copy_view_into(
1860 &mut hidden,
1861 0,
1862 &hidden_src.slice(offs[s] * n_embd..(offs[s] + ts[s]) * n_embd),
1863 ts[s] * n_embd,
1864 )?;
1865 let last0 = (offs[s] + ts[s] - 1) * n_embd;
1866 let mut h_seed = e.uninit(n_embd)?;
1867 e.copy_view_into(
1868 &mut h_seed,
1869 0,
1870 &hidden_src.slice(last0..last0 + n_embd),
1871 n_embd,
1872 )?;
1873 let mut hlast = e.uninit(n_embd)?;
1875 e.copy_view_into(
1876 &mut hlast,
1877 0,
1878 &hn.slice(last0..last0 + n_embd),
1879 n_embd,
1880 )?;
1881 let logits = e.dtoh(&e.matmul(&self.output, &hlast, 1)?)?;
1882 caches[s].pos += ts[s];
1883 out.push((logits, h_seed, hidden));
1884 }
1885 Ok(out)
1886 }
1887
1888 fn step35_prime_cache_batch(
1889 &self,
1890 e: &Engine,
1891 prompts: &[&[u32]],
1892 caches: &mut [&mut Cache],
1893 ) -> Result<Vec<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
1894 if crate::pp::pp_host_bounce_active()
1895 && (!crate::pp::prime_pp_on()
1896 || crate::pp::pp_cuts(self.layers.len()).is_none())
1897 {
1898 return Err(
1899 "step35_prime_cache_batch: MEMRA_PP_HOST_BOUNCE=1 requires a valid prime \
1900 stage split; refusing an unsplit remote-weight walk"
1901 .into(),
1902 );
1903 }
1904 if !Self::step35_prime_batch_on() {
1905 return Err("step35 batched prime is disabled (MEMRA_STEP35_PRIME_BATCH=0)".into());
1906 }
1907 if caches.iter().any(|c| c.pos != 0) {
1908 return Err(
1909 "step35 batched prime currently supports complete fresh prompts only; \
1910 continuation/tick chunks require per-request queued_after"
1911 .into(),
1912 );
1913 }
1914
1915 let ts: Vec<usize> = prompts.iter().map(|p| p.len()).collect();
1916 for &t in &ts {
1917 assert!(t >= PRIME_MIN_T, "step35 batched prime needs T >= {PRIME_MIN_T}");
1918 }
1919 for (s, c) in caches.iter().enumerate() {
1920 assert!(ts[s] <= c.max_ctx, "step35 batched prime exceeds cache max_ctx");
1921 }
1922 let offs: Vec<usize> = ts
1923 .iter()
1924 .scan(0usize, |a, &t| {
1925 let o = *a;
1926 *a += t;
1927 Some(o)
1928 })
1929 .collect();
1930 let total: usize = ts.iter().sum();
1931 let payload = total * self.cfg.n_embd as usize;
1932 let cat_tokens: Vec<u32> = prompts.iter().flat_map(|p| p.iter().copied()).collect();
1933 let positions: Vec<Vec<i32>> = ts
1934 .iter()
1935 .map(|&t| (0..t as i32).collect())
1936 .collect();
1937 let upload_positions = |e: &Engine|
1938 -> Result<Vec<CudaSlice<i32>>, Box<dyn std::error::Error>> {
1939 positions
1940 .iter()
1941 .map(|p| e.htod_i32(p))
1942 .collect::<Result<_, _>>()
1943 };
1944
1945 static ONCE: std::sync::Once = std::sync::Once::new();
1946 ONCE.call_once(|| {
1947 eprintln!(
1948 "[step35-prime-batch] first concat prime: B={} tokens={total}",
1949 prompts.len()
1950 );
1951 });
1952
1953 let out = if !crate::pp::pp2_streams_off() && crate::pp::prime_pp_on() {
1954 if let Some(fence) = crate::pp::pp_cuts(self.layers.len()) {
1955 let rt = crate::pp::PpNRt::get(e)?;
1956 let n_st = fence.len() - 1;
1957 assert_eq!(rt.n_stages(), n_st, "step35 prime batch stage count mismatch");
1958 let caller_stream = e.stream();
1959 rt.fence_stages_behind(&caller_stream)?;
1960
1961 let mut slot = {
1962 let _st0 = rt.enter(0);
1963 let e0 = rt.engine(0, e);
1964 let pos_ds = upload_positions(e0)?;
1965 let x = self.embed(e0, &cat_tokens)?;
1966 let x = self.step35_prime_batch_layers(
1967 e0,
1968 x,
1969 fence[0],
1970 fence[1],
1971 &ts,
1972 &offs,
1973 &pos_ds,
1974 caches,
1975 )?;
1976 rt.tx(0, &x, payload)?
1977 };
1978 for s in 1..n_st - 1 {
1979 let _st = rt.enter(s);
1980 let es = rt.engine(s, e);
1981 let pos_ds = upload_positions(es)?;
1982 let x = rt.rx(s - 1, slot, payload)?;
1983 let x = self.step35_prime_batch_layers(
1984 es,
1985 x,
1986 fence[s],
1987 fence[s + 1],
1988 &ts,
1989 &offs,
1990 &pos_ds,
1991 caches,
1992 )?;
1993 slot = rt.tx(s, &x, payload)?;
1994 }
1995
1996 let _stl = rt.enter(n_st - 1);
1997 let el = rt.engine(n_st - 1, e);
1998 let pos_ds = upload_positions(el)?;
1999 let x = rt.rx(n_st - 2, slot, payload)?;
2000 let x = self.step35_prime_batch_layers(
2001 el,
2002 x,
2003 fence[n_st - 1],
2004 fence[n_st],
2005 &ts,
2006 &offs,
2007 &pos_ds,
2008 caches,
2009 )?;
2010 let out = self.step35_prime_batch_epilogue(el, x, &ts, &offs, caches)?;
2011 rt.publish_to(n_st - 1, &caller_stream)?;
2012 crate::pp::STEP35_PRIME_BATCH_SPLITS.fetch_add(
2013 1,
2014 std::sync::atomic::Ordering::Relaxed,
2015 );
2016 out
2017 } else {
2018 let pos_ds = upload_positions(e)?;
2019 let x = self.embed(e, &cat_tokens)?;
2020 let x = self.step35_prime_batch_layers(
2021 e,
2022 x,
2023 0,
2024 self.layers.len(),
2025 &ts,
2026 &offs,
2027 &pos_ds,
2028 caches,
2029 )?;
2030 self.step35_prime_batch_epilogue(e, x, &ts, &offs, caches)?
2031 }
2032 } else {
2033 let pos_ds = upload_positions(e)?;
2034 let x = self.embed(e, &cat_tokens)?;
2035 let x = self.step35_prime_batch_layers(
2036 e,
2037 x,
2038 0,
2039 self.layers.len(),
2040 &ts,
2041 &offs,
2042 &pos_ds,
2043 caches,
2044 )?;
2045 self.step35_prime_batch_epilogue(e, x, &ts, &offs, caches)?
2046 };
2047 crate::pp::STEP35_PRIME_BATCHES.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2048 Ok(out)
2049 }
2050
2051 pub fn prime_cache_batch(&self, e: &Engine, prompts: &[&[u32]], caches: &mut [&mut Cache])
2068 -> Result<Vec<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
2069 let cfg = &self.cfg;
2070 let n_embd = cfg.n_embd as usize;
2071 let eps = cfg.rms_eps;
2072 let b = prompts.len();
2073 assert!(b >= 1 && b == caches.len());
2074 let pos0s: Vec<usize> = caches.iter().map(|c| c.pos).collect();
2075 let carried = pos0s.iter().any(|&p| p > 0);
2076 if cfg.gemma4.is_some() {
2082 return Err("prime_cache_batch: gemma4 has no batched prime core (per-layer \
2083 swa/global geometry, softcapped head) — use gemma4_prime per sequence".into());
2084 }
2085 if cfg.step35.is_some() {
2088 return self.step35_prime_cache_batch(e, prompts, caches);
2089 }
2090 let ts: Vec<usize> = prompts.iter().map(|p| p.len()).collect();
2091 for &t in &ts { assert!(t >= PRIME_MIN_T, "prime_cache_batch needs T >= {PRIME_MIN_T}"); }
2092 for (s, c) in caches.iter().enumerate() {
2093 assert!(c.pos + ts[s] <= c.max_ctx, "prime_cache_batch: prompt exceeds cache max_ctx");
2094 }
2095 let total: usize = ts.iter().sum();
2096 let offs: Vec<usize> = ts.iter().scan(0usize, |a, &t| { let o = *a; *a += t; Some(o) }).collect();
2097 let pos_ds: Vec<CudaSlice<i32>> = ts.iter().zip(&pos0s)
2099 .map(|(&t, &p0)| e.htod_i32(&(p0 as i32..(p0 + t) as i32).collect::<Vec<_>>()))
2100 .collect::<Result<_, _>>()?;
2101 let split = |e: &Engine, y: &CudaSlice<f32>, dim: usize|
2103 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
2104 let mut out = Vec::with_capacity(b);
2105 for s in 0..b {
2106 let mut ys = e.uninit(ts[s] * dim)?;
2107 e.copy_view_into(&mut ys, 0, &y.slice(offs[s] * dim..(offs[s] + ts[s]) * dim), ts[s] * dim)?;
2108 out.push(ys);
2109 }
2110 Ok(out)
2111 };
2112
2113 let cat_tokens: Vec<u32> = prompts.iter().flat_map(|p| p.iter().copied()).collect();
2114 let mut x = self.embed(e, &cat_tokens)?; for (il, layer) in self.layers.iter().enumerate() {
2116 let mut h = e.uninit(total * n_embd)?;
2117 let mut hx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
2118 e.rms_norm_f16out(&x, layer.attn_norm.float_data(), &mut h, &mut hx16, n_embd, total, eps)?;
2119 let mut mixed = e.uninit(total * n_embd)?;
2121 match &layer.mixer {
2122 Mixer::Full(fa) => {
2123 let g3 = e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], &h, &hx16, total)?;
2124 let geometry = self.cfg.full_attention_geometry_at(il as u32);
2130 let (n_head, n_head_kv, head_dim) = (
2131 geometry.n_head as usize,
2132 geometry.n_head_kv as usize,
2133 geometry.head_dim_k as usize,
2134 );
2135 let fa_scale = geometry.attention_scale();
2136 let use_favl = !carried
2137 && (2..=8).contains(&b)
2138 && (head_dim == 256 || head_dim == 128)
2139 && geometry.attention_gate
2140 == memra_gguf::config::AttentionGateKind::FusedQ
2141 && std::env::var("MEMRA_NOFA").is_err()
2142 && std::env::var("MEMRA_FA_FLOOR").is_err()
2143 && std::env::var("MEMRA_FA_PP_W2").as_deref() != Ok("1")
2144 && std::env::var("MEMRA_FA_BF16KV").as_deref() != Ok("0")
2145 && std::env::var("MEMRA_FA_VL").as_deref() != Ok("0");
2146 if use_favl {
2147 let (qf_w, kf_w, vf_w) =
2148 (fa.wq.out_features(), fa.wk.out_features(), fa.wv.out_features());
2149 struct APre {
2150 q: CudaSlice<f32>, gate: Option<CudaSlice<f32>>,
2151 qn: CudaSlice<f32>, kn: CudaSlice<f32>,
2152 }
2153 let mut aps = Vec::with_capacity(b);
2154 for &t in ts.iter().take(b) {
2155 aps.push(APre {
2156 q: e.uninit(t * n_head * head_dim)?,
2157 gate: Some(e.uninit(t * n_head * head_dim)?),
2158 qn: e.uninit(t * n_head * head_dim)?,
2159 kn: e.uninit(t * n_head_kv * head_dim)?,
2160 });
2161 }
2162 let (kv_dim_k, kv_dim_v, ktb, vtb) = {
2163 let kvl = caches[0].kv[il].as_ref().unwrap();
2164 (kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes)
2165 };
2166 let pargs: Vec<crate::AttnPreVl> = (0..b).map(|s| {
2167 let (o, t) = (offs[s], ts[s]);
2168 let kvl = caches[s].kv[il].as_ref().unwrap();
2169 assert!(kvl.len == 0 && kvl.len + t <= caches[s].max_ctx,
2170 "prime_cache_batch attn vl: fresh + capacity");
2171 crate::AttnPreVl {
2172 qf: e.addr_f32v(&g3[0].slice(o * qf_w..(o + t) * qf_w)),
2173 kf: e.addr_f32v(&g3[1].slice(o * kf_w..(o + t) * kf_w)),
2174 vf: e.addr_f32v(&g3[2].slice(o * vf_w..(o + t) * vf_w)),
2175 q: e.addr_f32(&aps[s].q),
2176 gate: e.addr_f32(aps[s].gate.as_ref().unwrap()),
2177 qn: e.addr_f32(&aps[s].qn), kn: e.addr_f32(&aps[s].kn),
2178 kc: e.addr_u8(&kvl.k), vc: e.addr_u8(&kvl.v),
2179 t: t as i32, pad: 0,
2180 }
2181 }).collect();
2182 e.attn_pre_vl8(&pargs, fa.q_norm.float_data(), fa.k_norm.float_data(),
2183 head_dim, geometry.n_rot as usize, n_head, n_head_kv,
2184 self.cfg.rms_eps, geometry.rope_base, 1.0,
2185 kv_dim_k, kv_dim_v, ktb, vtb)?;
2186 for s in 0..b {
2187 let kvl = caches[s].kv[il].as_mut().unwrap();
2188 kvl.len += ts[s];
2189 let new_len = kvl.len as i32;
2190 e.set_i32_one(&mut kvl.len_d, new_len)?;
2191 }
2192 let mut attns = Vec::with_capacity(b);
2193 let mut mirrors = Vec::with_capacity(b);
2194 for &t in ts.iter().take(b) {
2195 attns.push(e.uninit(t * n_head * head_dim)?);
2196 let n = t * n_head_kv * head_dim;
2197 mirrors.push((e.alloc_u8_uninit(n * 2)?, e.alloc_u8_uninit(n * 2)?));
2198 }
2199 let fa3_on = match std::env::var("MEMRA_FA3").as_deref() {
2202 Ok("0") => false,
2203 Ok("1") => true,
2204 _ => cfg!(memra_hopper_mma),
2205 };
2206 if fa3_on {
2207 let mut q16s = Vec::with_capacity(b);
2208 let mut v16s = Vec::with_capacity(b);
2209 for s in 0..b {
2210 let t = ts[s];
2211 let mut q16 = e.alloc_u8_uninit(t * n_head * head_dim * 2)?;
2212 e.f32_to_bf16_into(&aps[s].qn, &mut q16, t * n_head * head_dim)?;
2213 let mut k16 = e.alloc_u8_uninit(t * n_head_kv * head_dim * 2)?;
2214 e.f32_to_bf16_into(&aps[s].kn, &mut k16, t * n_head_kv * head_dim)?;
2215 let mut v16 = e.alloc_u8_uninit(t * n_head_kv * head_dim * 2)?;
2216 e.f32_to_bf16_v(&g3[2].slice(offs[s] * vf_w..(offs[s] + t) * vf_w),
2217 &mut v16, t * n_head_kv * head_dim)?;
2218 q16s.push(q16);
2219 v16s.push((k16, v16));
2220 }
2221 let mut qp = [core::ptr::null::<core::ffi::c_void>(); 8];
2222 let mut kp = qp;
2223 let mut vp = qp;
2224 let mut op = [core::ptr::null_mut::<f32>(); 8];
2225 let mut tsv = [0i32; 8];
2226 for s in 0..b {
2227 qp[s] = e.addr_u8(&q16s[s]) as *const core::ffi::c_void;
2228 kp[s] = e.addr_u8(&v16s[s].0) as *const core::ffi::c_void;
2229 vp[s] = e.addr_u8(&v16s[s].1) as *const core::ffi::c_void;
2230 op[s] = e.addr_f32(&attns[s]) as *mut f32;
2231 tsv[s] = ts[s] as i32;
2232 }
2233 let rc = unsafe {
2234 crate::fa3_vl_raw(qp.as_ptr(), kp.as_ptr(), vp.as_ptr(), op.as_ptr(),
2235 tsv.as_ptr(), b as i32, n_head as i32,
2236 n_head_kv as i32, head_dim as i32, fa_scale,
2237 e.stream().cu_stream() as *mut core::ffi::c_void)
2238 };
2239 if rc != 0 {
2240 return Err(format!("memra_fa3_vl rc={rc}").into());
2241 }
2242 } else {
2243 let fargs: Vec<crate::FaSeqVl> = (0..b).map(|s| crate::FaSeqVl {
2244 q: e.addr_f32(&aps[s].qn), k16: e.addr_u8(&mirrors[s].0),
2245 v16: e.addr_u8(&mirrors[s].1), o: e.addr_f32(&attns[s]),
2246 kf: e.addr_f32(&aps[s].kn),
2247 vf: e.addr_f32v(&g3[2].slice(offs[s] * vf_w..(offs[s] + ts[s]) * vf_w)),
2248 t: ts[s] as i32, pad: 0,
2249 }).collect();
2250 e.fa_prefill_vl8(&fargs, head_dim, n_head, n_head_kv, fa_scale)?;
2251 }
2252 for (s, attn) in attns.into_iter().enumerate() {
2253 let (attn_g, ag16) = self.full_attn_prime_post_fa(
2254 e, attn, &aps[s].gate, ts[s], n_head, head_dim)?;
2255 let mut done = false;
2256 if let Some(xh) = &ag16 {
2257 done = e.try_f16_gemm_pre_into_off(&fa.wo, xh, ts[s], &mut mixed, offs[s] * n_embd)?;
2258 }
2259 if !done {
2260 let m = e.matmul(&fa.wo, &attn_g, ts[s])?;
2261 e.copy_into(&mut mixed, offs[s] * n_embd, &m, ts[s] * n_embd)?;
2262 }
2263 }
2264 } else {
2265 let mut parts: Vec<Vec<CudaSlice<f32>>> = (0..b).map(|_| Vec::new()).collect();
2266 for (w, y) in [&fa.wq, &fa.wk, &fa.wv].iter().zip(g3) {
2267 for (s, ys) in split(e, &y, w.out_features())?.into_iter().enumerate() {
2268 parts[s].push(ys);
2269 }
2270 }
2271 for (s, g3s) in parts.into_iter().enumerate() {
2272 let (attn_g, ag16) = self.full_attn_prime_core_inner(
2274 e, fa, g3s, &pos_ds[s], ts[s], caches[s], il)?;
2275 let mut done = false;
2276 if let Some(xh) = &ag16 {
2277 done = e.try_f16_gemm_pre_into_off(&fa.wo, xh, ts[s], &mut mixed, offs[s] * n_embd)?;
2278 }
2279 if !done {
2280 let m = e.matmul(&fa.wo, &attn_g, ts[s])?;
2281 e.copy_into(&mut mixed, offs[s] * n_embd, &m, ts[s] * n_embd)?;
2282 }
2283 }
2284 }
2285 }
2286 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
2287 Mixer::Linear(la) => {
2288 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
2293 let g4 = e.matmul_group_xh(&ws, &h, &hx16, total)?;
2294 let outs = self.linear_attn_prime_core_batch(e, la, &g4, &offs, &ts, caches, il)?;
2295 for (s, (gn, gn16)) in outs.into_iter().enumerate() {
2296 let (o, t) = (offs[s], ts[s]);
2297 let mut done = false;
2298 if let Some(xh) = &gn16 {
2299 done = e.try_f16_gemm_pre_into_off(&la.ssm_out, xh, t, &mut mixed, o * n_embd)?;
2300 }
2301 if !done {
2302 let m = e.matmul(&la.ssm_out, &gn, t)?;
2303 e.copy_into(&mut mixed, o * n_embd, &m, t * n_embd)?;
2304 }
2305 }
2306 }
2307 }
2308 let mut x1 = e.uninit(total * n_embd)?;
2309 let mut z = e.uninit(total * n_embd)?;
2310 let mut zx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
2311 e.add_rms_norm_f16out(&x, &mixed, layer.post_attn_norm.float_data(),
2312 &mut x1, &mut z, &mut zx16, n_embd, total, eps)?;
2313 let ffn_out = match &layer.ffn {
2314 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
2315 let n_ff = ffn_gate.out_features();
2316 let mut g2 = e.matmul_group_xh(&[ffn_gate, ffn_up], &z, &zx16, total)?;
2317 let up = g2.pop().unwrap();
2318 let gate = g2.pop().unwrap();
2319 let mut act = e.uninit(total * n_ff)?;
2320 let d_lim = self.cfg.clamp_shexp_at(il as u32);
2324 if Self::f16out_on(e, total) && self.cfg.m3.is_none() && d_lim.is_none() {
2325 let mut a16 = e.alloc_u8_uninit(total * n_ff * 2)?;
2326 e.silu_mul_f16out(&gate, &up, &mut act, &mut a16, total * n_ff)?;
2327 match e.try_f16_gemm_pre(ffn_down, &a16, total)? {
2328 Some(y) => y,
2329 None => e.matmul(ffn_down, &act, total)?,
2330 }
2331 } else {
2332 Self::ffn_act_lim(e, &self.cfg, &gate, &up, 1.0, 1.0, d_lim,
2333 &mut act, total * n_ff)?;
2334 e.matmul(ffn_down, &act, total)?
2335 }
2336 }
2337 crate::hybrid::Ffn::Moe(m) => {
2338 self.moe_ffn_il_prefill(e, m, &z, total, il as u16)?
2339 }
2340 };
2341 let mut x2 = e.uninit(total * n_embd)?;
2342 e.add(&x1, &ffn_out, &mut x2, total * n_embd)?;
2343 x = x2;
2344 }
2345 let mut hn = e.uninit(total * n_embd)?;
2347 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, total, eps)?;
2348 let mut hcat = e.uninit(b * n_embd)?;
2354 for s in 0..b {
2355 let last0 = (offs[s] + ts[s] - 1) * n_embd;
2356 e.copy_view_into(&mut hcat, s * n_embd, &hn.slice(last0..last0 + n_embd), n_embd)?;
2357 }
2358 let logits_cat = if b >= 2 { e.try_f16_gemm(&self.output, &hcat, b)? } else { None };
2359 let logits_host: Option<Vec<f32>> = match &logits_cat {
2360 Some(lc) => Some(e.dtoh(lc)?),
2361 None => None,
2362 };
2363 let n_vocab = self.output.out_features();
2364 let mut hidden_all = if crate::spec::spec_hpost() {
2365 split(e, &hn, n_embd)?
2366 } else {
2367 split(e, &x, n_embd)?
2368 };
2369 let mut out = Vec::with_capacity(b);
2370 for s in 0..b {
2371 let last0 = (offs[s] + ts[s] - 1) * n_embd;
2372 let mut h_seed = e.uninit(n_embd)?;
2373 if !crate::spec::spec_hpost() {
2374 e.copy_view_into(&mut h_seed, 0, &x.slice(last0..last0 + n_embd), n_embd)?;
2375 } else {
2376 e.copy_view_into(&mut h_seed, 0, &hn.slice(last0..last0 + n_embd), n_embd)?;
2377 }
2378 let logits = match &logits_host {
2379 Some(lh) => lh[s * n_vocab..(s + 1) * n_vocab].to_vec(),
2380 None => {
2381 let mut hlast = e.uninit(n_embd)?;
2382 e.copy_view_into(&mut hlast, 0, &hn.slice(last0..last0 + n_embd), n_embd)?;
2383 e.dtoh(&e.matmul(&self.output, &hlast, 1)?)?
2384 }
2385 };
2386 caches[s].pos += ts[s];
2387 out.push((logits, h_seed, hidden_all.remove(0)));
2388 }
2389 Ok(out)
2390 }
2391
2392 #[allow(clippy::too_many_arguments)]
2403 fn full_attn_prime(&self, e: &Engine, fa: &FullAttnLayer, h: &CudaSlice<f32>,
2404 hx: Option<&CudaSlice<u8>>,
2405 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache, il: usize,
2406 seq_end: usize)
2407 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2408 if self.cfg.step35.is_some() {
2409 return self.step35_attn_prime(e, fa, h, hx, pos_d, t, cache, il, seq_end);
2410 }
2411 let g3 = match hx {
2416 Some(xh) => e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], h, xh, t)?,
2417 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
2418 };
2419 self.full_attn_prime_core(e, fa, g3, pos_d, t, cache, il)
2420 }
2421
2422 fn full_attn_prime_core(&self, e: &Engine, fa: &FullAttnLayer, g3: Vec<CudaSlice<f32>>,
2426 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache, il: usize)
2427 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2428 let (attn_g, ag16) = self.full_attn_prime_core_inner(e, fa, g3, pos_d, t, cache, il)?;
2429 if let Some(xh) = &ag16 {
2430 if let Some(y) = e.try_f16_gemm_pre(&fa.wo, xh, t)? {
2431 return Ok(y);
2432 }
2433 }
2434 Ok(e.matmul(&fa.wo, &attn_g, t)?)
2435 }
2436
2437 fn full_attn_prime_core_inner(&self, e: &Engine, fa: &FullAttnLayer, g3: Vec<CudaSlice<f32>>,
2438 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache, il: usize)
2439 -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
2440 let cfg = &self.cfg;
2441 let geometry = cfg.full_attention_geometry_at(il as u32);
2442 let n_head = geometry.n_head as usize;
2443 let n_head_kv = geometry.n_head_kv as usize;
2444 let head_dim = geometry.head_dim_k as usize;
2445 let scale = geometry.attention_scale();
2446 let (pre, base_len) = self.full_attn_prime_pre_fa(e, fa, g3, pos_d, t, cache, il)?;
2447 let AttnPre { q, k, v, gate } = pre;
2448 let mut attn = e.uninit(t * n_head * head_dim)?;
2449 self.full_attn_prime_fa_dispatch(e, &q, &k, &v, &mut attn, base_len, t, cache, il,
2450 head_dim, n_head, n_head_kv, scale)?;
2451 self.full_attn_prime_post_fa(e, attn, &gate, t, n_head, head_dim)
2452 }
2453
2454 #[allow(clippy::type_complexity)]
2458 fn full_attn_prime_pre_fa(&self, e: &Engine, fa: &FullAttnLayer, mut g3: Vec<CudaSlice<f32>>,
2459 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache, il: usize)
2460 -> Result<(AttnPre, usize), Box<dyn std::error::Error>> {
2461 let cfg = &self.cfg;
2462 let geometry = cfg.full_attention_geometry_at(il as u32);
2463 let n_head = geometry.n_head as usize;
2464 let n_head_kv = geometry.n_head_kv as usize;
2465 let head_dim = geometry.head_dim_k as usize;
2466 let eps = cfg.rms_eps;
2467
2468 let gated = geometry.attention_gate
2472 == memra_gguf::config::AttentionGateKind::FusedQ;
2473 let v = g3.pop().unwrap();
2474 let mut k = g3.pop().unwrap();
2475 let qf = g3.pop().unwrap();
2476 let (mut q, gate) = if gated {
2477 let mut q = e.uninit(t * n_head * head_dim)?;
2478 let mut gate = e.uninit(t * n_head * head_dim)?;
2479 e.q_gate_split(&qf, &mut q, &mut gate, head_dim, n_head, t)?;
2480 (q, Some(gate))
2481 } else {
2482 (qf, None)
2483 };
2484
2485 let mut qn = e.uninit(t * n_head * head_dim)?;
2486 e.rms_norm(&q, fa.q_norm.float_data(), &mut qn, head_dim, n_head * t, eps)?;
2487 q = qn;
2488 let mut kn = e.uninit(t * n_head_kv * head_dim)?;
2489 e.rms_norm(&k, fa.k_norm.float_data(), &mut kn, head_dim, n_head_kv * t, eps)?;
2490 k = kn;
2491 let rope_dims = geometry.n_rot as usize;
2492 e.rope_neox(&mut q, pos_d, head_dim, rope_dims, n_head, t, geometry.rope_base, 1.0)?;
2493 e.rope_neox(&mut k, pos_d, head_dim, rope_dims, n_head_kv, t, geometry.rope_base, 1.0)?;
2494
2495 {
2498 let kvl = cache.kv[il].as_mut().unwrap();
2499 assert!(kvl.len + t <= cache.max_ctx, "prime_cache: KV overflow");
2500 e.append_kv_quantized_rows(&k, &v, &mut kvl.k, &mut kvl.v, kvl.len, t,
2501 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes,
2502 crate::Engine::kv_fp8_on())?;
2503 kvl.len += t;
2504 let new_len = kvl.len as i32;
2505 e.set_i32_one(&mut kvl.len_d, new_len)?;
2506 }
2507
2508 let base_len = {
2509 let kvl = cache.kv[il].as_ref().unwrap();
2510 kvl.len - t };
2512 Ok((AttnPre { q, k, v, gate }, base_len))
2513 }
2514
2515 #[allow(clippy::too_many_arguments)]
2522 fn full_attn_prime_fa_dispatch(&self, e: &Engine, q: &CudaSlice<f32>, k: &CudaSlice<f32>,
2523 v: &CudaSlice<f32>, attn: &mut CudaSlice<f32>, base_len: usize,
2524 t: usize, cache: &mut Cache, il: usize,
2525 head_dim: usize, n_head: usize, n_head_kv: usize, scale: f32)
2526 -> Result<(), Box<dyn std::error::Error>> {
2527 if base_len == 0 && std::env::var("MEMRA_PRIME_F32CHUNK0").as_deref() == Ok("1") {
2540 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
2541 e.sdpa_naive(q, k, v, attn, head_dim, n_head, n_head_kv, t, t, scale, true)?;
2542 } else {
2543 e.fa_prefill(q, k, v, attn, head_dim, n_head, n_head_kv, t, t, scale, true)?;
2544 }
2545 return Ok(());
2546 }
2547 let kvl = cache.kv[il].as_ref().unwrap();
2548 let t_kv = base_len + t;
2549 let k_view = e.view_u8(&kvl.k, t_kv * kvl.k_tok_bytes);
2550 let v_view = e.view_u8(&kvl.v, t_kv * kvl.v_tok_bytes);
2551 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
2555 e.sdpa_naive_quantized_view(q, &k_view, &v_view, attn, head_dim, n_head,
2556 n_head_kv, t, t_kv, scale, true,
2557 kvl.k_tok_bytes, kvl.v_tok_bytes)?;
2558 return Ok(());
2559 }
2560 let deqw = std::env::var("MEMRA_PRIME_DEQW").map(|v| v != "0").unwrap_or(true);
2568 if deqw {
2569 e.fa_prefill_view_ws(q, &k_view, &v_view, attn, head_dim, n_head, n_head_kv,
2570 t, t_kv, scale, true, kvl.k_tok_bytes, kvl.v_tok_bytes,
2571 crate::Engine::kv_fp8_on())?;
2572 } else {
2573 e.fa_prefill_view(q, &k_view, &v_view, attn, head_dim, n_head, n_head_kv,
2574 t, t_kv, scale, true, kvl.k_tok_bytes, kvl.v_tok_bytes,
2575 crate::Engine::kv_fp8_on())?;
2576 }
2577 Ok(())
2578 }
2579
2580 fn full_attn_prime_post_fa(&self, e: &Engine, attn: CudaSlice<f32>,
2583 gate: &Option<CudaSlice<f32>>, t: usize,
2584 n_head: usize, head_dim: usize)
2585 -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
2586 let (attn_g, ag16) = match gate {
2587 Some(gate) => {
2588 let n = t * n_head * head_dim;
2589 let mut ag = e.uninit(n)?;
2590 if Self::f16out_on(e, t) {
2591 let mut a16 = e.alloc_u8_uninit(n * 2)?;
2592 e.sig_mul_f16out(&attn, gate, &mut ag, &mut a16, n)?;
2593 (ag, Some(a16))
2594 } else {
2595 let mut gsig = e.uninit(n)?;
2596 e.sigmoid(gate, &mut gsig, n)?;
2597 e.mul(&attn, &gsig, &mut ag, n)?;
2598 (ag, None)
2599 }
2600 }
2601 None => (attn, None),
2602 };
2603 Ok((attn_g, ag16))
2604 }
2605
2606 fn linear_attn_prime(&self, e: &Engine, la: &LinearAttnLayer, h: &CudaSlice<f32>,
2613 hx: Option<&CudaSlice<u8>>, t: usize,
2614 cache: &mut Cache, il: usize)
2615 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2616 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
2618 let g4 = match hx {
2619 Some(xh) => e.matmul_group_xh(&ws, h, xh, t)?,
2620 None => e.matmul_group(&ws, h, t)?,
2621 };
2622 self.linear_attn_prime_core(e, la, g4, t, cache, il)
2623 }
2624
2625 fn linear_attn_prime_core(&self, e: &Engine, la: &LinearAttnLayer, mut g4: Vec<CudaSlice<f32>>,
2627 t: usize, cache: &mut Cache, il: usize)
2628 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2629 self.linear_attn_prime_core_pad(e, la, g4.drain(..).collect(), t, cache, il, None)
2630 }
2631
2632 #[allow(clippy::too_many_arguments)]
2636 fn linear_attn_prime_core_pad_inner(&self, e: &Engine, la: &LinearAttnLayer, mut g4: Vec<CudaSlice<f32>>,
2637 t: usize, cache: &mut Cache, il: usize,
2638 pad_len: Option<&CudaSlice<i32>>)
2639 -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
2640 let ssm = self.cfg.ssm.as_ref().unwrap();
2642 let d_state = ssm.state_size as usize;
2643 let num_k = ssm.group_count as usize;
2644 let num_v = ssm.time_step_rank as usize;
2645 let key_dim = d_state * num_k;
2646 let value_dim = d_state * num_v;
2647 let conv_dim = key_dim * 2 + value_dim;
2648 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(
2653 e, la,
2654 &qkv_mixed.slice(0..t * conv_dim), &z.slice(0..t * value_dim),
2655 &beta_raw.slice(0..t * num_v), &alpha.slice(0..t * num_v),
2656 t, cache, il, pad_len)
2657 }
2658
2659 #[allow(clippy::too_many_arguments)]
2662 fn linear_attn_gdn_prep(&self, e: &Engine, la: &LinearAttnLayer,
2663 qkv_mixed: &cudarc::driver::CudaView<f32>,
2664 beta_raw: &cudarc::driver::CudaView<f32>,
2665 alpha: &cudarc::driver::CudaView<f32>,
2666 t: usize, cache: &mut Cache, il: usize,
2667 pad_len: Option<&CudaSlice<i32>>)
2668 -> Result<GdnPrep, Box<dyn std::error::Error>> {
2669 let cfg = &self.cfg;
2670 let ssm = cfg.ssm.as_ref().unwrap();
2671 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;
2679 debug_assert!(t >= d_conv - 1, "stateful conv needs T >= pad (PRIME_MIN_T gates)");
2680
2681 let rl = cache.recur[il].as_mut().unwrap();
2686 let hk = Self::gdn_hk(e, t, num_v, num_k);
2687 let conv_fuse = std::env::var("MEMRA_CONV_FUSE").as_deref() != Ok("0");
2688 let hk = if conv_fuse { hk } else { num_v }; let mut q_g = e.uninit(d_state * hk * t)?;
2690 let mut k_g = e.uninit(d_state * hk * t)?;
2691 let mut v_g = e.uninit(d_state * num_v * t)?;
2692 if conv_fuse {
2693 e.ssm_conv1d_gdn_state_pad(qkv_mixed, &mut rl.conv_state, la.ssm_conv1d.float_data(),
2694 &mut q_g, &mut k_g, &mut v_g,
2695 conv_dim, t, d_conv, d_state, num_v, num_k, key_dim, hk, pad_len)?;
2696 } else {
2697 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(),
2699 &mut conv_out, conv_dim, t, d_conv, pad_len)?;
2700 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)?;
2701 }
2702 let mut q_l2 = e.uninit(d_state * hk * t)?;
2703 let qb16 = if Engine::l2_v2_on(d_state) && e.gdn_wgmma_on(32) {
2707 let mut qb = e.alloc_u8_uninit(d_state * hk * t * 2)?;
2708 e.l2_norm_pp(&q_g, &mut q_l2, Some(&mut qb), d_state, hk * t, eps)?;
2709 Some(qb)
2710 } else {
2711 e.l2_norm_pp(&q_g, &mut q_l2, None, d_state, hk * t, eps)?;
2712 None
2713 };
2714 let mut k_l2 = e.uninit(d_state * hk * t)?;
2715 let kb16 = if Engine::l2_v2_on(d_state) {
2717 let mut kb = e.alloc_u8_uninit(d_state * hk * t * 2)?;
2718 e.l2_norm_pp(&k_g, &mut k_l2, Some(&mut kb), d_state, hk * t, eps)?;
2719 Some(kb)
2720 } else {
2721 e.l2_norm_pp(&k_g, &mut k_l2, None, d_state, hk * t, eps)?;
2722 None
2723 };
2724 let mut beta = e.uninit(t * num_v)?;
2725 e.sigmoid_v(beta_raw, &mut beta, t * num_v)?;
2726 let mut g_log = e.uninit(t * num_v)?;
2727 e.gdn_glog_v(alpha, la.ssm_dt.float_data(), la.ssm_a.float_data(), &mut g_log, num_v, t)?;
2728 if let Some(len_d) = pad_len {
2729 e.gdn_pad_mask(&mut beta, &mut g_log, len_d, num_v, t)?;
2730 }
2731 Ok(GdnPrep { hk, q_l2, k_l2, v_g, beta, g_log, kb16, qb16 })
2732 }
2733
2734 #[allow(clippy::too_many_arguments)]
2739 fn linear_attn_prime_core_batch(&self, e: &Engine, la: &LinearAttnLayer,
2740 g4: &[CudaSlice<f32>], offs: &[usize], ts: &[usize],
2741 caches: &mut [&mut Cache], il: usize)
2742 -> Result<Vec<(CudaSlice<f32>, Option<CudaSlice<u8>>)>, Box<dyn std::error::Error>> {
2743 let ssm = self.cfg.ssm.as_ref().unwrap();
2744 let d_state = ssm.state_size as usize;
2745 let num_k = ssm.group_count as usize;
2746 let num_v = ssm.time_step_rank as usize;
2747 let key_dim = d_state * num_k;
2748 let value_dim = d_state * num_v;
2749 let conv_dim = key_dim * 2 + value_dim;
2750 let eps = self.cfg.rms_eps;
2751 let scale = 1.0 / (d_state as f32).sqrt();
2752 let b = ts.len();
2753 let c = Engine::gdn_chunk_size();
2754 let carried = caches.iter().any(|c| c.pos > 0);
2757 let use_vl = !carried
2758 && (2..=8).contains(&b)
2759 && Engine::gdn_chunked_enabled() && ts.iter().all(|&t| t >= 16)
2760 && e.gdn_mma_enabled(c)
2761 && std::env::var("MEMRA_GDN_VL").as_deref() != Ok("0");
2762 if !use_vl {
2763 return (0..b).map(|s| {
2764 let (o, t) = (offs[s], ts[s]);
2765 self.linear_attn_prime_core_pad_view(
2766 e, la,
2767 &g4[0].slice(o * conv_dim..(o + t) * conv_dim),
2768 &g4[1].slice(o * value_dim..(o + t) * value_dim),
2769 &g4[2].slice(o * num_v..(o + t) * num_v),
2770 &g4[3].slice(o * num_v..(o + t) * num_v),
2771 t, caches[s], il, None)
2772 }).collect();
2773 }
2774 struct SeqBufs {
2778 conv_out: CudaSlice<f32>, q_g: CudaSlice<f32>, k_g: CudaSlice<f32>, v_g: CudaSlice<f32>,
2779 q_l2: CudaSlice<f32>, k_l2: CudaSlice<f32>, beta: CudaSlice<f32>, g_log: CudaSlice<f32>,
2780 gn: CudaSlice<f32>, gn16: CudaSlice<u8>,
2781 }
2782 let d_conv = ssm.conv_kernel as usize;
2783 let f16o = Self::f16out_on(e, 16);
2784 let hk = Self::gdn_hk(e, 16, num_v, num_k); let mut sb = Vec::with_capacity(b);
2786 let mut pres = Vec::with_capacity(b);
2787 for &t in ts.iter().take(b) {
2788 sb.push(SeqBufs {
2789 conv_out: e.uninit(conv_dim * t)?,
2790 q_g: e.uninit(d_state * hk * t)?,
2791 k_g: e.uninit(d_state * hk * t)?,
2792 v_g: e.uninit(d_state * num_v * t)?,
2793 q_l2: e.uninit(d_state * hk * t)?,
2794 k_l2: e.uninit(d_state * hk * t)?,
2795 beta: e.uninit(t * num_v)?,
2796 g_log: e.uninit(t * num_v)?,
2797 gn: e.uninit(d_state * num_v * t)?,
2798 gn16: e.alloc_u8_uninit(d_state * num_v * t * 2)?,
2799 });
2800 pres.push(e.gdn_chunk_alloc(num_v, t, c, hk)?);
2801 }
2802 let prep_args: Vec<crate::GdnPrepVl> = (0..b).map(|s| {
2803 let (o, t) = (offs[s], ts[s]);
2804 let rl = caches[s].recur[il].as_ref().unwrap();
2805 crate::GdnPrepVl {
2806 qkv: e.addr_f32v(&g4[0].slice(o * conv_dim..(o + t) * conv_dim)),
2807 conv_state: e.addr_f32(&rl.conv_state),
2808 conv_out: e.addr_f32(&sb[s].conv_out),
2809 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),
2810 q_l2: e.addr_f32(&sb[s].q_l2), k_l2: e.addr_f32(&sb[s].k_l2),
2811 beta_raw: e.addr_f32v(&g4[2].slice(o * num_v..(o + t) * num_v)),
2812 alpha: e.addr_f32v(&g4[3].slice(o * num_v..(o + t) * num_v)),
2813 beta: e.addr_f32(&sb[s].beta), g_log: e.addr_f32(&sb[s].g_log),
2814 o: e.addr_f32(&pres[s].o),
2815 z: e.addr_f32v(&g4[1].slice(o * value_dim..(o + t) * value_dim)),
2816 gn: e.addr_f32(&sb[s].gn), gn16: e.addr_u8(&sb[s].gn16),
2817 kb16: if Engine::l2_v2_on(d_state) { e.addr_u8(&pres[s].kb16) } else { 0 },
2818 qb16: if Engine::l2_v2_on(d_state) && e.gdn_wgmma_on(c) { e.addr_u8(&pres[s].qb16) } else { 0 },
2819 t: t as i32, pad: 0,
2820 }
2821 }).collect();
2822 let args: Vec<crate::GdnSeqVl> = (0..b).map(|s| {
2823 let rl = caches[s].recur[il].as_ref().unwrap();
2824 crate::GdnSeqVl {
2825 kb16: e.addr_u8(&pres[s].kb16), gcum: e.addr_f32(&pres[s].gcum),
2826 beta: e.addr_f32(&sb[s].beta), u: e.addr_f32(&pres[s].u),
2827 wb16: e.addr_u8(&pres[s].wb16), y: e.addr_u8(&pres[s].y16),
2828 ssnap: e.addr_u8(&pres[s].ssnap16),
2829 state_in: e.addr_f32(&rl.ssm_state), state_out: e.addr_f32(&rl.ssm_state_alt),
2830 q: e.addr_f32(&sb[s].q_l2), p: e.addr_f32(&pres[s].p),
2831 o: e.addr_f32(&pres[s].o),
2832 k: e.addr_f32(&sb[s].k_l2), v: e.addr_f32(&sb[s].v_g),
2833 g: e.addr_f32(&sb[s].g_log), a: e.addr_f32(&pres[s].a),
2834 w: e.addr_f32(&pres[s].w),
2835 t: ts[s] as i32, nc: pres[s].nc as i32,
2836 }
2837 }).collect();
2838 e.gdn_prep_vl8(&prep_args, la.ssm_conv1d.float_data(), la.ssm_dt.float_data(),
2839 la.ssm_a.float_data(), conv_dim, d_conv, d_state, num_v, num_k, key_dim, hk, eps)?;
2840 if !Engine::l2_v2_on(d_state) {
2843 e.gdn_mirror_vl8(&args, num_v, 0, hk)?;
2844 }
2845 let wq8: Option<crate::GdnWVl8> = if e.gdn_wgmma_on(c) {
2847 if !Engine::l2_v2_on(d_state) {
2849 for s in 0..b {
2850 e.f32_to_bf16_into(&sb[s].q_l2, &mut pres[s].qb16, d_state * hk * ts[s])?;
2851 }
2852 }
2853 let mut wa = [crate::GdnWVl::default(); 8];
2854 for s in 0..b {
2855 wa[s] = crate::GdnWVl { qb16: e.addr_u8(&pres[s].qb16), pb16: e.addr_u8(&pres[s].pb16) };
2856 }
2857 Some(crate::GdnWVl8(wa))
2858 } else { None };
2859 e.gdn_chunk_k123_vl8(&args, num_v, hk, wq8.as_ref())?;
2860 e.gdn_chunk_vl8(&args, num_v, scale, hk, wq8.as_ref())?;
2861 if f16o {
2862 e.gdn_tail_vl8(&prep_args, la.ssm_norm.float_data(), d_state, num_v, eps)?;
2863 }
2864 let mut out = Vec::with_capacity(b);
2866 for (s, bufs) in sb.into_iter().enumerate() {
2867 let rl = caches[s].recur[il].as_mut().unwrap();
2868 std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
2869 let (o, t) = (offs[s], ts[s]);
2870 let SeqBufs { mut gn, gn16, .. } = bufs;
2871 if f16o {
2872 out.push((gn, Some(gn16)));
2873 } else {
2874 let z_v = g4[1].slice(o * value_dim..(o + t) * value_dim);
2875 e.gated_rmsnorm_zv(&pres[s].o, la.ssm_norm.float_data(), &z_v, &mut gn,
2876 d_state, num_v * t, eps)?;
2877 out.push((gn, None));
2878 }
2879 }
2880 Ok(out)
2881 }
2882
2883 #[allow(clippy::too_many_arguments)]
2887 fn linear_attn_prime_core_pad_view(&self, e: &Engine, la: &LinearAttnLayer,
2888 qkv_mixed: &cudarc::driver::CudaView<f32>,
2889 z: &cudarc::driver::CudaView<f32>,
2890 beta_raw: &cudarc::driver::CudaView<f32>,
2891 alpha: &cudarc::driver::CudaView<f32>,
2892 t: usize, cache: &mut Cache, il: usize,
2893 pad_len: Option<&CudaSlice<i32>>)
2894 -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
2895 let cfg = &self.cfg;
2896 let ssm = cfg.ssm.as_ref().unwrap();
2897 let d_state = ssm.state_size as usize; let num_v = ssm.time_step_rank as usize; let eps = cfg.rms_eps;
2900 let scale = 1.0 / (d_state as f32).sqrt();
2901
2902 let prep = self.linear_attn_gdn_prep(e, la, qkv_mixed, beta_raw, alpha, t, cache, il, pad_len)?;
2903
2904 let mut o = e.uninit(d_state * num_v * t)?;
2910 let rl = cache.recur[il].as_mut().unwrap();
2911 {
2912 let crate::cache::RecurLayer { ssm_state, ssm_state_alt, .. } = rl;
2913 e.gdn_scan_prefill(&prep.q_l2, &prep.k_l2, &prep.v_g, &prep.g_log, &prep.beta,
2914 prep.kb16.as_ref(), prep.qb16.as_ref(), ssm_state, ssm_state_alt, &mut o, num_v, t, scale,
2915 prep.hk)?;
2916 }
2917 std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
2918
2919 let mut gn = e.uninit(d_state * num_v * t)?;
2922 let gn16 = if Self::f16out_on(e, t) {
2923 let mut g16 = e.alloc_u8_uninit(d_state * num_v * t * 2)?;
2924 e.gated_rmsnorm_f16out_zv(&o, la.ssm_norm.float_data(), z, &mut gn, &mut g16,
2925 d_state, num_v * t, eps)?;
2926 Some(g16)
2927 } else {
2928 e.gated_rmsnorm_zv(&o, la.ssm_norm.float_data(), z, &mut gn, d_state, num_v * t, eps)?;
2929 None
2930 };
2931 Ok((gn, gn16))
2932 }
2933
2934 #[allow(clippy::too_many_arguments)]
2936 fn linear_attn_prime_core_pad(&self, e: &Engine, la: &LinearAttnLayer, g4: Vec<CudaSlice<f32>>,
2937 t: usize, cache: &mut Cache, il: usize,
2938 pad_len: Option<&CudaSlice<i32>>)
2939 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2940 let (gn, gn16) = self.linear_attn_prime_core_pad_inner(e, la, g4, t, cache, il, pad_len)?;
2941 if let Some(xh) = &gn16 {
2942 if let Some(y) = e.try_f16_gemm_pre(&la.ssm_out, xh, t)? {
2943 return Ok(y);
2944 }
2945 }
2946 Ok(e.matmul(&la.ssm_out, &gn, t)?)
2947 }
2948
2949 pub fn full_attn(&self, e: &Engine, fa: &FullAttnLayer, h: &CudaSlice<f32>, pos_d: &CudaSlice<i32>, t: usize, il: usize)
2954 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2955 if self.cfg.step35.is_some() {
2956 return self.step35_attn(e, fa, h, pos_d, t, il);
2957 }
2958 let cfg = &self.cfg;
2959 let _n_embd = cfg.n_embd as usize;
2960 let geometry = cfg.full_attention_geometry_at(il as u32);
2961 let n_head = geometry.n_head as usize;
2962 let n_head_kv = geometry.n_head_kv as usize;
2963 let head_dim = geometry.head_dim_k as usize;
2964 let eps = cfg.rms_eps;
2965 let scale = geometry.attention_scale();
2966
2967 let gated = geometry.attention_gate
2970 == memra_gguf::config::AttentionGateKind::FusedQ;
2971 let mut g3 = e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?;
2973 let v = g3.pop().unwrap();
2974 let mut k = g3.pop().unwrap();
2975 let qf = g3.pop().unwrap();
2976 let (mut q, gate) = if gated {
2977 let mut q = e.uninit(t * n_head * head_dim)?;
2978 let mut gate = e.uninit(t * n_head * head_dim)?;
2979 e.q_gate_split(&qf, &mut q, &mut gate, head_dim, n_head, t)?;
2980 (q, Some(gate))
2981 } else {
2982 (qf, None)
2983 };
2984
2985 let mut qn = e.uninit(t * n_head * head_dim)?;
2987 e.rms_norm(&q, fa.q_norm.float_data(), &mut qn, head_dim, n_head * t, eps)?;
2988 q = qn;
2989 let mut kn = e.uninit(t * n_head_kv * head_dim)?;
2990 e.rms_norm(&k, fa.k_norm.float_data(), &mut kn, head_dim, n_head_kv * t, eps)?;
2991 k = kn;
2992 let rope_dims = geometry.n_rot as usize;
2993 e.rope_neox(&mut q, pos_d, head_dim, rope_dims, n_head, t, geometry.rope_base, 1.0)?;
2994 e.rope_neox(&mut k, pos_d, head_dim, rope_dims, n_head_kv, t, geometry.rope_base, 1.0)?;
2995
2996 let mut attn = e.uninit(t * n_head * head_dim)?;
2998 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
3001 e.sdpa_naive(&q, &k, &v, &mut attn, head_dim, n_head, n_head_kv, t, t, scale, true)?;
3003 } else {
3004 e.fa_prefill(&q, &k, &v, &mut attn, head_dim, n_head, n_head_kv, t, t, scale, true)?;
3005 }
3006
3007 let attn_g = match &gate {
3009 Some(gate) => {
3010 let mut gsig = e.uninit(t * n_head * head_dim)?;
3011 e.sigmoid(gate, &mut gsig, t * n_head * head_dim)?;
3012 let mut ag = e.uninit(t * n_head * head_dim)?;
3013 e.mul(&attn, &gsig, &mut ag, t * n_head * head_dim)?;
3014 ag
3015 }
3016 None => attn,
3017 };
3018
3019 let o = e.matmul(&fa.wo, &attn_g, t)?;
3021 Ok(o)
3022 }
3023
3024 pub fn linear_attn(&self, e: &Engine, la: &LinearAttnLayer, h: &CudaSlice<f32>, t: usize)
3026 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3027 let cfg = &self.cfg;
3028 let _n_embd = cfg.n_embd as usize;
3029 let ssm = cfg.ssm.as_ref().unwrap();
3030 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;
3035 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;
3039 let scale = 1.0 / (d_state as f32).sqrt();
3040
3041 let mut g4 = e.matmul_group(&[&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha], h, t)?;
3044 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);
3056 let mut q_g = e.uninit(d_state * num_v * t)?;
3057 let mut k_g = e.uninit(d_state * num_v * t)?;
3058 let mut v_g = e.uninit(d_state * num_v * t)?;
3059 e.ssm_conv1d_gdn(&qkv_mixed, la.ssm_conv1d.float_data(), &mut q_g, &mut k_g, &mut v_g,
3060 conv_dim, t, d_conv, d_state, num_v, num_k, key_dim)?;
3061 let mut q_l2 = e.uninit(d_state * num_v * t)?;
3063 e.l2_norm(&q_g, &mut q_l2, d_state, num_v * t, eps)?;
3064 let mut k_l2 = e.uninit(d_state * num_v * t)?;
3065 e.l2_norm(&k_g, &mut k_l2, d_state, num_v * t, eps)?;
3066 let v_gd = v_g;
3067
3068 let mut beta = e.uninit(t * num_v)?;
3071 e.sigmoid(&beta_raw, &mut beta, t * num_v)?;
3072 let mut g_log = e.uninit(t * num_v)?;
3074 e.gdn_glog(&alpha, la.ssm_dt.float_data(), la.ssm_a.float_data(), &mut g_log, num_v, t)?;
3075
3076 let state_in = e.zeros(d_state * d_state * num_v)?; let mut state_out = e.zeros(d_state * d_state * num_v)?;
3079 let mut o = e.uninit(d_state * num_v * t)?;
3080 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)?;
3081
3082 let mut gn = e.uninit(d_state * num_v * t)?;
3087 e.gated_rmsnorm(&o, la.ssm_norm.float_data(), &z, &mut gn, d_state, num_v * t, eps)?;
3088
3089 let out = e.matmul(&la.ssm_out, &gn, t)?;
3093 Ok(out)
3094 }
3095}
3096
3097impl HybridModel {
3098 pub fn moe_ffn_il(&self, e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, t: usize, il: u16)
3109 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3110 Self::moe_ffn_inner(e, m, z, None, t, &self.cfg, il, self.max_moe_block(), false)
3111 }
3112
3113 pub fn moe_ffn_il_prefill(
3116 &self,
3117 e: &Engine,
3118 m: &MoeWeights,
3119 z: &CudaSlice<f32>,
3120 t: usize,
3121 il: u16,
3122 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3123 Self::moe_ffn_inner(e, m, z, None, t, &self.cfg, il, self.max_moe_block(), true)
3124 }
3125
3126 pub fn moe_ffn_il_zq8(&self, e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>,
3130 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>, t: usize, il: u16)
3131 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3132 Self::moe_ffn_inner(
3133 e, m, z, zq8, t, &self.cfg, il, self.max_moe_block(), false,
3134 )
3135 }
3136
3137 pub(crate) fn moe_ffn(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, t: usize,
3145 cfg: &ModelConfig, il: u16, max_block: usize)
3146 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3147 Self::moe_ffn_inner(e, m, z, None, t, cfg, il, max_block, false)
3148 }
3149
3150 #[allow(clippy::too_many_arguments)]
3151 pub(crate) fn moe_ffn_inner(
3152 e: &Engine,
3153 m: &MoeWeights,
3154 z: &CudaSlice<f32>,
3155 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
3156 t: usize,
3157 cfg: &ModelConfig,
3158 il: u16,
3159 max_block: usize,
3160 prefill: bool,
3161 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3162 let worker_io = crate::spill_pread::worker_enabled();
3163 let epoch_lfu = std::env::var_os("MEMRA_MOE_LFU_DECAY").is_some();
3164 if Engine::moe_cache_enabled() && (worker_io || epoch_lfu) {
3165 e.with_moe_cache(max_block, |cache, _| {
3166 cache.begin_forward_epoch(il, t);
3167 if worker_io {
3168 cache.begin_worker_scope();
3169 }
3170 Ok(())
3171 })?;
3172 }
3173 if t > 1 && moe_grouped_enabled(cfg, prefill) {
3176 let grouped_out = Self::moe_ffn_grouped(e, m, z, t, cfg, il, max_block)?;
3177 if std::env::var("MEMRA_MOE_GATE").is_ok() {
3182 let seq_out = Self::moe_ffn_sequential(e, m, z, t, cfg, il, max_block)?;
3183 let g_host = e.dtoh(&grouped_out)?;
3184 let s_host = e.dtoh(&seq_out)?;
3185 let g_bytes: &[u8] = unsafe { std::slice::from_raw_parts(g_host.as_ptr() as *const u8, g_host.len() * 4) };
3186 let s_bytes: &[u8] = unsafe { std::slice::from_raw_parts(s_host.as_ptr() as *const u8, s_host.len() * 4) };
3187 if g_bytes == s_bytes {
3188 println!("moe-gate il={il} t={t} BYTE-IDENTICAL");
3189 } else {
3190 let diffs = g_host.iter().zip(s_host.iter()).enumerate()
3191 .filter(|(_, (a, b))| a != b).count();
3192 let maxdiff = g_host.iter().zip(s_host.iter())
3193 .map(|(a, b)| (a - b).abs()).fold(0.0f32, f32::max);
3194 panic!("moe-gate il={il} t={t} MISMATCH: {diffs}/{} elems differ, maxdiff={maxdiff:.6e}", g_host.len());
3195 }
3196 }
3197 return Ok(grouped_out);
3198 }
3199 Self::moe_ffn_sequential_zq8(e, m, z, zq8, t, cfg, il, max_block)
3200 }
3201
3202 pub(crate) fn moe_ffn_sequential(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, t: usize,
3204 cfg: &ModelConfig, il: u16, max_block: usize)
3205 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3206 Self::moe_ffn_sequential_zq8(e, m, z, None, t, cfg, il, max_block)
3207 }
3208
3209 fn moe_router_logits(
3213 e: &Engine,
3214 m: &MoeWeights,
3215 z: &CudaSlice<f32>,
3216 t: usize,
3217 cfg: &ModelConfig,
3218 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3219 if t < PRIME_MIN_T {
3220 if crate::router_kernel_on() {
3222 e.router_gemv(
3223 m.gate_inp.float_data(),
3224 z,
3225 cfg.n_embd as usize,
3226 m.gate_exps.n_expert,
3227 t,
3228 )
3229 } else {
3230 e.matmul_decode_exact(&m.gate_inp, z, t)
3231 }
3232 } else if crate::router_prefill_exact_on() && crate::router_kernel_on() {
3233 e.router_gemv(
3234 m.gate_inp.float_data(),
3235 z,
3236 cfg.n_embd as usize,
3237 m.gate_exps.n_expert,
3238 t,
3239 )
3240 } else {
3241 e.matmul(&m.gate_inp, z, t)
3242 }
3243 }
3244
3245 fn trace_moe_routes(il: u16, t: usize, sel_all: &[u32], weights: &[f32])
3249 -> Result<(), Box<dyn std::error::Error>> {
3250 use std::io::Write as _;
3251 if let Ok(path) = std::env::var("MEMRA_MOE_TRACE") {
3252 let mut f = std::fs::OpenOptions::new().create(true).append(true).open(path)?;
3253 let ids: Vec<String> = sel_all.iter().map(|s| s.to_string()).collect();
3254 writeln!(f, "{} {} {}", il, t, ids.join(","))?;
3255 }
3256 if let Ok(path) = std::env::var("MEMRA_MOE_WEIGHT_TRACE") {
3257 let mut f = std::fs::OpenOptions::new().create(true).append(true).open(path)?;
3258 let pairs: Vec<String> = sel_all.iter().zip(weights)
3259 .map(|(expert, weight)| format!("{expert}:{weight:.9}"))
3260 .collect();
3261 writeln!(f, "{} {} {}", il, t, pairs.join(","))?;
3262 }
3263 Ok(())
3264 }
3265
3266 fn trace_moe_input(e: &Engine, il: u16, t: usize, n_embd: usize, z: &CudaSlice<f32>)
3271 -> Result<(), Box<dyn std::error::Error>> {
3272 use std::io::Write as _;
3273 let Ok(dir) = std::env::var("MEMRA_MOE_INPUT_TRACE_DIR") else { return Ok(()) };
3274 let host = e.dtoh(z)?;
3275 if host.len() != t * n_embd {
3276 return Err(format!(
3277 "MoE input trace shape mismatch at layer {il}: got {} values, expected {}x{}",
3278 host.len(), t, n_embd
3279 ).into());
3280 }
3281 let bytes = unsafe {
3282 std::slice::from_raw_parts(
3283 host.as_ptr().cast::<u8>(), host.len() * std::mem::size_of::<f32>()
3284 )
3285 };
3286 let state = MOE_INPUT_TRACE_WRITER.get_or_init(|| std::sync::Mutex::new(None));
3287 let mut state = state.lock().map_err(|_| "MoE input trace writer lock is poisoned")?;
3288 if state.is_none() {
3289 let dir = std::path::PathBuf::from(&dir);
3290 std::fs::create_dir_all(&dir)?;
3291 let index = std::fs::OpenOptions::new().create(true).append(true)
3292 .open(dir.join("index.jsonl"))?;
3293 *state = Some(MoeInputTraceWriter {
3294 dir,
3295 index,
3296 payloads: std::collections::HashMap::new(),
3297 });
3298 }
3299 let writer = state.as_mut().unwrap();
3300 if writer.dir != std::path::Path::new(&dir) {
3301 return Err("MEMRA_MOE_INPUT_TRACE_DIR changed after capture started".into());
3302 }
3303 let file_name = format!("layer-{il:03}.f32");
3304 if !writer.payloads.contains_key(&il) {
3305 let payload = std::fs::OpenOptions::new().create(true).append(true)
3306 .open(writer.dir.join(&file_name))?;
3307 let offset = payload.metadata()?.len();
3308 writer.payloads.insert(il, (payload, offset));
3309 }
3310 let (payload, offset) = writer.payloads.get_mut(&il).unwrap();
3311 let row_offset = *offset;
3312 payload.write_all(bytes)?;
3313 *offset += bytes.len() as u64;
3314 writeln!(
3315 writer.index,
3316 "{{\"format\":\"memra-moe-input-trace-v1\",\"layer\":{il},\"tokens\":{t},\
3317 \"hidden_size\":{n_embd},\"file\":\"{file_name}\",\"offset\":{row_offset},\
3318 \"payload_bytes\":{}}}",
3319 bytes.len()
3320 )?;
3321 Ok(())
3322 }
3323
3324 #[allow(clippy::too_many_arguments)]
3325 pub(crate) fn moe_ffn_sequential_zq8(
3326 e: &Engine,
3327 m: &MoeWeights,
3328 z: &CudaSlice<f32>,
3329 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
3330 t: usize,
3331 cfg: &ModelConfig,
3332 il: u16,
3333 max_block: usize,
3334 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3335 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
3336 let moe = cfg.moe.as_ref().unwrap();
3337 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);
3344 debug_assert_eq!(m.gate_exps.out_f, n_ff_exp);
3345 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);
3348
3349 let lim_exp = cfg.clamp_exp_at(il as u32);
3352 let lim_shexp = cfg.clamp_shexp_at(il as u32);
3353 let use_cache = Engine::moe_cache_enabled();
3354 let uniform_experts = m.has_uniform_expert_layout();
3355 let moe_q8 = uniform_experts && moe_q8_enabled()
3356 && q8_expert_supported(m.gate_exps.qtype) && q8_expert_supported(m.up_exps.qtype)
3357 && q8_expert_supported(m.down_exps.qtype);
3358 let cpu_expert_requested = crate::cpu_experts::configured();
3365 if cpu_expert_requested && (cfg.hy3.is_none() || cfg.m3.is_some()) {
3366 return Err(std::io::Error::other(
3367 "MEMRA_CPU_EXPERT_LIB is experimental and currently gated to Hy3",
3368 )
3369 .into());
3370 }
3371 let cpu_hybrid = cpu_expert_requested && t < PRIME_MIN_T && m.dev_exps.is_none();
3372 let freeze_cpu_residency = cpu_expert_requested
3378 && std::env::var("MEMRA_CPU_EXPERT_FREEZE_CACHE").as_deref() == Ok("1");
3379 let caller_warms_before_freeze = std::env::var("MEMRA_CPU_EXPERT_FREEZE_WARMUP_TOKENS")
3380 .ok()
3381 .and_then(|value| value.parse::<usize>().ok())
3382 .is_some_and(|tokens| tokens > 0);
3383 if cpu_hybrid && freeze_cpu_residency && !caller_warms_before_freeze {
3384 e.freeze_moe_cache();
3385 }
3386 let cache_frozen = use_cache && e.moe_cache_frozen();
3387 let cache_dispatch = use_cache && (!cache_frozen || cpu_hybrid);
3388
3389 let logits = Self::moe_router_logits(e, m, z, t, cfg)?;
3392
3393 let no_exp_macros = m.gate_exps.macros.is_none() && m.up_exps.macros.is_none()
3431 && m.down_exps.macros.is_none();
3432 if cfg.sigmoid_router().is_none() && cfg.m3.is_none() && cfg.hy3.is_none()
3436 && !cfg.swiglu_clamped_at(il as u32)
3437 && no_exp_macros
3438 && t >= PRIME_MIN_T && m.dev_exps.is_some() && moe_q8_enabled()
3439 && q8_expert_supported(m.gate_exps.qtype) && q8_expert_supported(m.up_exps.qtype)
3440 && q8_expert_supported(m.down_exps.qtype)
3441 && std::env::var("MEMRA_MOE_PAIRS").map(|v| v != "0").unwrap_or(true)
3442 && std::env::var("MEMRA_MOE_STATS").is_err() {
3443 return Self::moe_ffn_pairs(e, m, z, &logits, t, cfg);
3444 }
3445
3446 let dev_ok = uniform_experts && cfg.sigmoid_router().is_none()
3463 && cfg.m3.is_none() && cfg.hy3.is_none()
3464 && !cfg.swiglu_clamped_at(il as u32);
3465 let observe_routes = std::env::var("MEMRA_MOE_STATS").is_ok()
3469 || std::env::var("MEMRA_MOE_TRACE").is_ok()
3470 || std::env::var("MEMRA_MOE_WEIGHT_TRACE").is_ok()
3471 || std::env::var("MEMRA_MOE_INPUT_TRACE_DIR").is_ok();
3472 if dev_ok && t < PRIME_MIN_T && m.dev_exps.is_some() && n_used <= 8 && moe_dev_enabled()
3473 && !observe_routes {
3474 return Self::moe_ffn_dev(e, m, z, zq8, &logits, t, cfg, il, max_block);
3475 }
3476 if dev_ok && use_cache && n_used <= 8 && moe_dev_enabled()
3477 && !observe_routes {
3478 let row_ok = e.with_moe_cache(max_block, |c, eng| {
3479 if moe_prewarm_enabled() { c.prewarm_layer(il, m, eng)?; }
3480 Ok(c.layer_dev_row(il, n_expert, eng)?.is_some())
3481 })?;
3482 if row_ok {
3483 return Self::moe_ffn_dev(e, m, z, zq8, &logits, t, cfg, il, max_block);
3484 }
3485 }
3486
3487 let (sel_all, w_all, routed_cpu_input) = if let Some(sig) = cfg.sigmoid_router() {
3489 if cpu_hybrid {
3490 let (sel, w, input) = Self::moe_route_sigmoid_with_input(
3491 e,
3492 &logits,
3493 z,
3494 t,
3495 n_expert,
3496 n_used,
3497 m.exp_probs_b.as_deref(),
3498 sig,
3499 m.active_experts.as_deref(),
3500 )?;
3501 (sel, w, Some(input))
3502 } else {
3503 let (sel, w) = Self::moe_route_cfg(
3504 e,
3505 &logits,
3506 t,
3507 n_expert,
3508 n_used,
3509 m.exp_probs_b.as_deref(),
3510 Some(sig),
3511 m.active_experts.as_deref(),
3512 )?;
3513 (sel, w, None)
3514 }
3515 } else {
3516 let (sel, w) = Self::moe_route_cfg(
3517 e,
3518 &logits,
3519 t,
3520 n_expert,
3521 n_used,
3522 None,
3523 None,
3524 m.active_experts.as_deref(),
3525 )?;
3526 (sel, w, None)
3527 };
3528
3529 Self::trace_moe_routes(il, t, &sel_all, &w_all)?;
3533 Self::trace_moe_input(e, il, t, n_embd, z)?;
3534
3535 let worker_disk_prefetch =
3547 cache_dispatch && crate::spill_pread::worker_enabled() && !cpu_hybrid;
3548 let promote_worker_h2d =
3549 t == 1 && worker_disk_prefetch && crate::spill_pread::copy_h2d_enabled();
3550 if promote_worker_h2d {
3551 let mut selected_blocks = Vec::with_capacity(n_used * 3);
3552 for &ex in sel_all.iter().take(n_used) {
3553 let ex = ex as u16;
3554 selected_blocks.extend([
3555 BlockId::new(il, PROJ_GATE, ex),
3556 BlockId::new(il, PROJ_UP, ex),
3557 BlockId::new(il, PROJ_DOWN, ex),
3558 ]);
3559 }
3560 for &ex in sel_all.iter().take(n_used) {
3561 Self::moe_prefetch_disk_expert(e, il, ex as usize, m, max_block, &selected_blocks)?;
3562 }
3563 e.with_moe_cache(max_block, |cache, eng| {
3564 cache.promote_worker_reads_at_safe_boundary(
3565 &selected_blocks,
3566 &selected_blocks,
3567 eng,
3568 )?;
3569 Ok(())
3570 })?;
3571 }
3572
3573 if t > 1 && std::env::var("MEMRA_MOE_STATS").is_ok() {
3576 let mut cnt = vec![0u32; n_expert];
3577 for &s in sel_all.iter() { cnt[s as usize] += 1; }
3578 let total = sel_all.len() as f64;
3579 let mut h = 0.0f64;
3580 let mut active = 0usize;
3581 for &c in &cnt { if c > 0 { active += 1; let p = c as f64 / total; h -= p * p.log2(); } }
3582 let maxc = cnt.iter().copied().max().unwrap_or(0);
3583 println!("moe-stats il={} t={} assignments={} active={}/{} entropy={:.3}b (max {:.3}b) mean_tok_per_active={:.2} max_tok_per_expert={}",
3584 il, t, sel_all.len(), active, n_expert, h, (n_expert as f64).log2(), total / active.max(1) as f64, maxc);
3585 }
3586
3587 let gdec_may_fire = uniform_experts && use_cache && n_used <= 8 && gdec_enabled()
3600 && !cfg.swiglu_clamped_at(il as u32);
3601 let slab_local = m.dev_exps.as_ref()
3617 .filter(|d| !d.gu_il && moe_slab_enabled() && d.dev == e.ctx().ordinal());
3618 let slab_bases = slab_local.map(|d| {
3619 use cudarc::driver::DevicePtr;
3620 let s = e.stream();
3621 let (pg, _g0) = d.gate.device_ptr(&s);
3622 let (pu, _g1) = d.up.device_ptr(&s);
3623 let (pd, _g2) = d.down.device_ptr(&s);
3624 (pg as u64, pu as u64, pd as u64)
3625 });
3626 let slab_fused_may_fire = slab_bases.is_some() && n_used <= 8 && gdec_enabled()
3636 && !cfg.swiglu_clamped_at(il as u32) && cfg.m3.is_none()
3637 && no_exp_macros && moe_q8;
3638 let mut moe_out = if gdec_may_fire || slab_fused_may_fire {
3641 e.uninit(t * n_embd)?
3642 } else {
3643 e.zeros(t * n_embd)?
3644 };
3645 let cpu_input = if cpu_hybrid {
3648 Some(routed_cpu_input.ok_or("CPU expert routing did not return the MoE input")?)
3649 } else {
3650 None
3651 };
3652
3653 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;
3661 let mut scratch_u: Option<CudaSlice<u8>> = None;
3662 let mut scratch_d: Option<CudaSlice<u8>> = None;
3663 let page_window = moe_page_prefetch_window();
3671
3672 for tok in 0..t {
3675 let sel = &sel_all[tok * n_used..(tok + 1) * n_used];
3676 let w = &w_all[tok * n_used..(tok + 1) * n_used];
3677 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd); let mut tok_q8: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
3679
3680 let no_macros = m.gate_exps.macros.is_none() && m.up_exps.macros.is_none()
3694 && m.down_exps.macros.is_none();
3695 if slab_fused_may_fire {
3705 let (pg, pu, pd) = slab_bases.unwrap();
3706 let mut gp = [0u64; 8];
3707 let mut up = [0u64; 8];
3708 let mut dp = [0u64; 8];
3709 for (j, &ex) in sel.iter().enumerate() {
3710 let ex = ex as usize;
3711 gp[j] = pg + (ex * m.gate_exps.expert_stride) as u64;
3712 up[j] = pu + (ex * m.up_exps.expert_stride) as u64;
3713 dp[j] = pd + (ex * m.down_exps.expert_stride) as u64;
3714 }
3715 let mut wv = [0f32; 8];
3716 wv[..n_used].copy_from_slice(w);
3717 if tok_q8.is_none() {
3718 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
3719 }
3720 let (zq, zd) = tok_q8.as_ref().unwrap();
3721 let act = e.moe_gate_up_silu8_q8(crate::WPtr8(gp), crate::WPtr8(up), zq, zd,
3722 n_embd, n_ff_exp, n_used,
3723 m.gate_exps.qtype, m.up_exps.qtype,
3724 m.gate_exps.row_bytes, m.up_exps.row_bytes)?;
3725 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
3726 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
3727 e.moe_down8_fma_q8(crate::WPtr8(dp), crate::F32x8(wv), &aq2, &ad2, &mut dst,
3728 n_ff_exp, n_embd, n_used,
3729 m.down_exps.qtype, m.down_exps.row_bytes)?;
3730 continue;
3731 }
3732 if gdec_may_fire && moe_q8 && cfg.m3.is_none() && no_macros {
3733 if tok_q8.is_none() {
3734 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
3735 }
3736 let (zq, zd) = tok_q8.as_ref().unwrap();
3737 if Self::moe_gdec_token_q8(e, m, il, max_block, zq, zd, sel, w,
3738 &mut moe_out, tok, n_embd, n_ff_exp, n_used)? {
3739 continue;
3740 }
3741 } else if gdec_may_fire && cfg.m3.is_none() && no_macros
3742 && Self::moe_gdec_token(e, m, il, max_block, &zt, sel, w,
3743 &mut moe_out, tok, n_embd, n_ff_exp, n_used)? {
3744 continue;
3745 }
3746
3747 if gdec_may_fire || slab_fused_may_fire {
3753 let mut row = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
3754 e.memset_zeros_view(&mut row)?;
3755 }
3756
3757 let mut cpu_mask = vec![false; sel.len()];
3763 let cpu_worker = if let Some(host_input) = cpu_input.as_ref() {
3764 let gpu_resident = if use_cache {
3765 e.with_moe_cache(max_block, |cache, _| {
3766 Ok(sel
3767 .iter()
3768 .map(|&expert| {
3769 let expert = expert as u16;
3770 [PROJ_GATE, PROJ_UP, PROJ_DOWN]
3771 .into_iter()
3772 .filter(|&projection| {
3773 cache
3774 .resident(BlockId::new(il, projection, expert))
3775 .is_some()
3776 })
3777 .count()
3778 })
3779 .collect::<Vec<_>>())
3780 })?
3781 } else {
3782 vec![0; sel.len()]
3783 };
3784 let mut cpu_selected = Vec::new();
3785 for (index, (&expert, &route_weight)) in sel.iter().zip(w).enumerate() {
3786 if gpu_resident[index] != 3 {
3787 cpu_mask[index] = true;
3788 crate::cpu_experts::record_incomplete_gpu_residency(gpu_resident[index]);
3789 let expert = expert as usize;
3790 cpu_selected.push((expert, route_weight));
3791 }
3792 }
3793 if crate::cpu_experts::predictor_enabled() {
3794 let row = &host_input[tok * n_embd..(tok + 1) * n_embd];
3798 crate::cpu_experts::predictor_submit(il, row);
3799 }
3800 if cpu_selected.is_empty() {
3801 None
3802 } else {
3803 let row = &host_input[tok * n_embd..(tok + 1) * n_embd];
3804 let job = crate::cpu_experts::prepare_job(m, il, &cpu_selected, row)
3805 .map_err(std::io::Error::other)?;
3806 Some(crate::cpu_experts::submit(job).map_err(std::io::Error::other)?)
3807 }
3808 } else {
3809 None
3810 };
3811
3812 let worker_window = worker_disk_prefetch
3813 .then(worker_prefetch_window)
3814 .unwrap_or(0);
3815 for (j, &ex) in sel.iter().enumerate() {
3816 if cpu_mask[j] {
3817 continue;
3818 }
3819 let ex = ex as usize;
3820 if let Some(d) = slab_local {
3827 let gl = m.gate_exps.expert_layout(ex);
3828 let ul = m.up_exps.expert_layout(ex);
3829 let dl = m.down_exps.expert_layout(ex);
3830 let (g0, u0, d0) = (ex * m.gate_exps.expert_stride,
3831 ex * m.up_exps.expert_stride,
3832 ex * m.down_exps.expert_stride);
3833 let (gate, up) = if moe_q8 {
3834 if tok_q8.is_none() {
3835 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
3836 }
3837 let (zq, zd) = tok_q8.as_ref().unwrap();
3838 (e.qmatvec_expert_q8(&d.gate, g0..g0 + gl.len, zq, zd, 1,
3839 m.gate_exps.in_f, m.gate_exps.out_f,
3840 gl.qtype, gl.row_bytes)?,
3841 e.qmatvec_expert_q8(&d.up, u0..u0 + ul.len, zq, zd, 1,
3842 m.up_exps.in_f, m.up_exps.out_f,
3843 ul.qtype, ul.row_bytes)?)
3844 } else {
3845 (e.qmatvec_view(&d.gate, g0..g0 + gl.len, &zt, 1,
3846 m.gate_exps.in_f, m.gate_exps.out_f,
3847 gl.qtype, gl.row_bytes)?,
3848 e.qmatvec_view(&d.up, u0..u0 + ul.len, &zt, 1,
3849 m.up_exps.in_f, m.up_exps.out_f,
3850 ul.qtype, ul.row_bytes)?)
3851 };
3852 let mut act = e.uninit(n_ff_exp)?;
3853 Self::ffn_act_lim(e, cfg, &gate, &up, m.gate_exps.macro_scale(ex),
3854 m.up_exps.macro_scale(ex), lim_exp, &mut act, n_ff_exp)?;
3855 let y = if moe_q8 {
3856 let (aq2, ad2) = e.quantize_q8_1(&act, 1, n_ff_exp)?;
3857 e.qmatvec_expert_q8(&d.down, d0..d0 + dl.len, &aq2, &ad2, 1,
3858 m.down_exps.in_f, m.down_exps.out_f,
3859 dl.qtype, dl.row_bytes)?
3860 } else {
3861 let actv = act.slice(0..n_ff_exp);
3862 e.qmatvec_view(&d.down, d0..d0 + dl.len, &actv, 1,
3863 m.down_exps.in_f, m.down_exps.out_f,
3864 dl.qtype, dl.row_bytes)?
3865 };
3866 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
3867 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
3868 continue;
3869 }
3870 for next in page_prefetch_positions(j, sel.len(), page_window) {
3871 Self::moe_prefetch_host_expert(sel[next] as usize, m);
3872 }
3873 let keep = [
3874 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_GATE, ex as u16),
3875 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_UP, ex as u16),
3876 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_DOWN, ex as u16),
3877 ];
3878 if worker_disk_prefetch && worker_window > 0 {
3879 for next in worker_prefetch_positions(j, sel.len(), worker_window) {
3880 Self::moe_prefetch_disk_expert(
3881 e,
3882 il,
3883 sel[next] as usize,
3884 m,
3885 max_block,
3886 &keep,
3887 )?;
3888 }
3889 } else if cache_dispatch
3890 && !cpu_hybrid
3891 && moe_prefetch_enabled()
3892 && j + 1 < sel.len()
3893 {
3894 let next = sel[j + 1] as usize;
3895 Self::moe_prefetch_expert(e, il, next, m, max_block, &keep)?;
3896 }
3897 let [gate_q8, up_q8, down_q8] = [moe_q8; 3];
3898 if cache_dispatch && (gate_q8 || up_q8 || down_q8) {
3899 if (gate_q8 || up_q8) && tok_q8.is_none() {
3902 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
3903 }
3904 let gate = if gate_q8 {
3905 let (zq, zd) = tok_q8.as_ref().unwrap();
3906 Self::moe_cached_gemm_q8(e, il, PROJ_GATE, ex, m, max_block, zq, zd)?
3907 } else {
3908 Self::moe_cached_gemm(e, il, PROJ_GATE, ex, m, max_block, &zt)?
3909 };
3910 let up = if up_q8 {
3911 let (zq, zd) = tok_q8.as_ref().unwrap();
3912 Self::moe_cached_gemm_q8(e, il, PROJ_UP, ex, m, max_block, zq, zd)?
3913 } else {
3914 Self::moe_cached_gemm(e, il, PROJ_UP, ex, m, max_block, &zt)?
3915 };
3916 let mut act = e.uninit(n_ff_exp)?;
3917 Self::ffn_act_lim(
3918 e,
3919 cfg,
3920 &gate,
3921 &up,
3922 m.gate_exps.macro_scale(ex),
3923 m.up_exps.macro_scale(ex),
3924 lim_exp,
3925 &mut act,
3926 n_ff_exp,
3927 )?;
3928 let y = if down_q8 {
3929 let (aq2, ad2) = e.quantize_q8_1(&act, 1, n_ff_exp)?;
3930 Self::moe_cached_gemm_q8(e, il, PROJ_DOWN, ex, m, max_block, &aq2, &ad2)?
3931 } else {
3932 let actv = act.slice(0..n_ff_exp);
3933 Self::moe_cached_gemm(e, il, PROJ_DOWN, ex, m, max_block, &actv)?
3934 };
3935 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
3936 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
3938 } else if cache_dispatch {
3939 let gate = Self::moe_cached_gemm(e, il, PROJ_GATE, ex, m, max_block, &zt)?;
3944 let up = Self::moe_cached_gemm(e, il, PROJ_UP, ex, m, max_block, &zt)?;
3945 let mut act = e.uninit(n_ff_exp)?; Self::ffn_act_lim(e, cfg, &gate, &up, m.gate_exps.macro_scale(ex),
3947 m.up_exps.macro_scale(ex), lim_exp, &mut act, n_ff_exp)?;
3948 let actv = act.slice(0..n_ff_exp);
3949 let y = Self::moe_cached_gemm(e, il, PROJ_DOWN, ex, m, max_block, &actv)?;
3950 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
3951 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
3953 } else if cache_frozen {
3954 let gate = Self::moe_frozen_gemm(
3959 e,
3960 il,
3961 PROJ_GATE,
3962 ex,
3963 m,
3964 max_block,
3965 &zt,
3966 &mut scratch_g,
3967 g_len,
3968 )?;
3969 let up = Self::moe_frozen_gemm(
3970 e,
3971 il,
3972 PROJ_UP,
3973 ex,
3974 m,
3975 max_block,
3976 &zt,
3977 &mut scratch_u,
3978 u_len,
3979 )?;
3980 let mut act = e.uninit(n_ff_exp)?;
3981 Self::ffn_act_lim(
3982 e,
3983 cfg,
3984 &gate,
3985 &up,
3986 m.gate_exps.macro_scale(ex),
3987 m.up_exps.macro_scale(ex),
3988 lim_exp,
3989 &mut act,
3990 n_ff_exp,
3991 )?;
3992 let actv = act.slice(0..n_ff_exp);
3993 let y = Self::moe_frozen_gemm(
3994 e,
3995 il,
3996 PROJ_DOWN,
3997 ex,
3998 m,
3999 max_block,
4000 &actv,
4001 &mut scratch_d,
4002 d_len,
4003 )?;
4004 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
4005 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
4006 } else {
4007 if scratch_g.is_none() {
4011 scratch_g = Some(e.alloc_u8_uninit(g_len)?);
4012 scratch_u = Some(e.alloc_u8_uninit(u_len)?);
4013 scratch_d = Some(e.alloc_u8_uninit(d_len)?);
4014 }
4015 let (sg, su, sd) = (scratch_g.as_mut().unwrap(), scratch_u.as_mut().unwrap(),
4016 scratch_d.as_mut().unwrap());
4017 let gl = m.gate_exps.expert_layout(ex);
4018 let ul = m.up_exps.expert_layout(ex);
4019 let dl = m.down_exps.expert_layout(ex);
4020 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
4021 let gate = e.qmatvec_view(sg, 0..gl.len, &zt, 1,
4022 m.gate_exps.in_f, m.gate_exps.out_f, gl.qtype, gl.row_bytes)?;
4023
4024 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
4025 let up = e.qmatvec_view(su, 0..ul.len, &zt, 1,
4026 m.up_exps.in_f, m.up_exps.out_f, ul.qtype, ul.row_bytes)?;
4027
4028 let mut act = e.uninit(n_ff_exp)?; Self::ffn_act_lim(e, cfg, &gate, &up, m.gate_exps.macro_scale(ex),
4030 m.up_exps.macro_scale(ex), lim_exp, &mut act, n_ff_exp)?;
4031
4032 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
4033 let actv = act.slice(0..n_ff_exp);
4034 let y = e.qmatvec_view(sd, 0..dl.len, &actv, 1,
4035 m.down_exps.in_f, m.down_exps.out_f, dl.qtype, dl.row_bytes)?;
4036
4037 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
4038 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
4039 }
4040 }
4041 if let Some(worker) = cpu_worker {
4042 let cpu_output = worker.wait().map_err(std::io::Error::other)?;
4043 let cpu_output = e.htod(&cpu_output)?;
4044 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
4045 e.axpy_into(&cpu_output, 1.0, &mut dst, n_embd)?;
4046 }
4047 if cpu_hybrid && !cache_frozen && cpu_expert_profile_admit_enabled() {
4048 for (j, &ex) in sel.iter().enumerate() {
4049 if cpu_mask[j] {
4050 Self::moe_profile_admit_expert(e, il, ex as usize, m, max_block)?;
4051 }
4052 }
4053 }
4054 }
4055
4056 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
4061 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
4062 {
4063 let n_ff_sh = gate_shexp.out_features(); let verify_t = t > 1 && t < PRIME_MIN_T;
4072 let (sg_gate, sg_up) = if t == 1 {
4073 match e.matmul_q8_fused2_x(gate_shexp, up_shexp, z)? {
4074 Some(pair) => pair,
4075 None => (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?),
4076 }
4077 } else if verify_t {
4078 (e.matmul_decode_exact(gate_shexp, z, t)?, e.matmul_decode_exact(up_shexp, z, t)?)
4079 } else {
4080 (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?) };
4082 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)?;
4084 let sh = if verify_t { e.matmul_decode_exact(down_shexp, &sa, t)? }
4085 else { e.matmul(down_shexp, &sa, t)? }; let g = match &m.gate_inp_shexp {
4099 Some(gate_inp_shexp) => {
4100 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
4101 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
4102 } else {
4103 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
4104 let mut g = e.uninit(t)?; e.sigmoid(&gs, &mut g, t)?;
4106 g
4107 }
4108 }
4109 None => e.htod(&vec![1.0f32; t])?,
4110 };
4111 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
4113 }
4114
4115 Ok(moe_out)
4116 }
4117
4118 pub fn stage1_h2d_per_token(&self) -> u64 {
4121 use crate::hybrid::Ffn;
4122 let n_used = self.cfg.moe.as_ref().map(|m| m.expert_used_count as u64).unwrap_or(0);
4123 let mut bytes = 0u64;
4124 for l in self.layers.iter() {
4125 if let Ffn::Moe(m) = &l.ffn {
4126 bytes += n_used * (m.gate_exps.max_expert_bytes() + m.up_exps.max_expert_bytes()
4127 + m.down_exps.max_expert_bytes()) as u64;
4128 }
4129 }
4130 bytes
4131 }
4132
4133 pub(crate) fn max_moe_block(&self) -> usize {
4137 use crate::hybrid::Ffn;
4138 let mut mx = 0usize;
4139 let mut scan = |ffn: &Ffn| {
4140 if let Ffn::Moe(m) = ffn {
4141 mx = mx.max(m.gate_exps.max_expert_bytes())
4142 .max(m.up_exps.max_expert_bytes())
4143 .max(m.down_exps.max_expert_bytes());
4144 }
4145 };
4146 for l in self.layers.iter() { scan(&l.ffn); }
4147 if let Some(mtp) = self.mtp.as_ref() { scan(&mtp.ffn); }
4148 mx
4149 }
4150
4151 pub(crate) fn moe_cache_block_sizes(&self) -> Vec<usize> {
4154 use crate::hybrid::Ffn;
4155 let mut sizes = Vec::new();
4156 let mut scan = |ffn: &Ffn| {
4157 let Ffn::Moe(m) = ffn else { return };
4158 for ex in 0..m.gate_exps.n_expert {
4159 if m.active_experts.as_ref().is_some_and(|active| !active[ex]) {
4160 continue;
4161 }
4162 for exps in [&m.gate_exps, &m.up_exps, &m.down_exps] {
4163 let len = exps.expert_layout(ex).len;
4164 if len > 0 {
4165 sizes.push(len);
4166 }
4167 }
4168 }
4169 };
4170 for layer in &self.layers {
4171 scan(&layer.ffn);
4172 }
4173 if let Some(mtp) = &self.mtp {
4174 scan(&mtp.ffn);
4175 }
4176 sizes
4177 }
4178
4179 pub fn save_cpu_expert_residency_profile(
4185 &self,
4186 e: &Engine,
4187 path: &std::path::Path,
4188 ) -> Result<(), Box<dyn std::error::Error>> {
4189 let Some(ids) = e.export_moe_residency() else {
4190 return Err("no MoE residency cache to persist".into());
4191 };
4192 let mut body = format!(
4193 "memra-freeze-profile v1 max_block={} blocks={}\n",
4194 self.max_moe_block(),
4195 ids.len()
4196 );
4197 for (layer, proj, ex) in &ids {
4198 body.push_str(&format!("{layer} {proj} {ex}\n"));
4199 }
4200 let tmp = path.with_extension("tmp");
4201 std::fs::write(&tmp, body)?;
4202 std::fs::rename(&tmp, path)?;
4203 println!(
4204 "[moe-cache] freeze profile saved: {} blocks -> {}",
4205 ids.len(),
4206 path.display()
4207 );
4208 Ok(())
4209 }
4210
4211 pub fn restore_cpu_expert_residency_profile(
4215 &self,
4216 e: &Engine,
4217 path: &std::path::Path,
4218 ) -> Result<bool, Box<dyn std::error::Error>> {
4219 use crate::hybrid::Ffn;
4220 use crate::moe_cache::BlockId;
4221 let Ok(content) = std::fs::read_to_string(path) else {
4222 return Ok(false);
4223 };
4224 let mut lines = content.lines();
4225 let Some(header) = lines.next() else { return Ok(false) };
4226 let expected = format!("memra-freeze-profile v1 max_block={}", self.max_moe_block());
4227 if !header.starts_with(&expected) {
4228 println!(
4229 "[moe-cache] freeze profile ignored (geometry mismatch): {}",
4230 path.display()
4231 );
4232 return Ok(false);
4233 }
4234 let mut by_layer: std::collections::HashMap<u16, Vec<BlockId>> =
4235 std::collections::HashMap::new();
4236 for line in lines {
4237 let mut fields = line.split_whitespace();
4238 let (Some(layer), Some(proj), Some(ex)) =
4239 (fields.next(), fields.next(), fields.next())
4240 else {
4241 continue;
4242 };
4243 let (Ok(layer), Ok(proj), Ok(ex)) =
4244 (layer.parse::<u16>(), proj.parse::<u8>(), ex.parse::<u16>())
4245 else {
4246 continue;
4247 };
4248 by_layer
4249 .entry(layer)
4250 .or_default()
4251 .push(BlockId::new(layer, proj, ex));
4252 }
4253 let requested: usize = by_layer.values().map(Vec::len).sum();
4254 if requested == 0 {
4255 return Ok(false);
4256 }
4257 let max_block = self.max_moe_block();
4258 let mut restaged = 0usize;
4259 let mut stage_layer = |layer_index: u16,
4260 ffn: &Ffn|
4261 -> Result<(), Box<dyn std::error::Error>> {
4262 let Ffn::Moe(m) = ffn else { return Ok(()) };
4263 let Some(ids) = by_layer.get(&layer_index) else {
4264 return Ok(());
4265 };
4266 e.with_moe_cache(max_block, |cache, eng| {
4267 for id in ids {
4268 if cache.restage_block(*id, m, eng)? {
4269 restaged += 1;
4270 }
4271 }
4272 Ok(())
4273 })
4274 };
4275 for (index, layer) in self.layers.iter().enumerate() {
4276 stage_layer(index as u16, &layer.ffn)?;
4277 }
4278 if let Some(mtp) = self.mtp.as_ref() {
4279 stage_layer(u16::MAX, &mtp.ffn)?;
4280 }
4281 e.freeze_moe_cache();
4282 println!(
4283 "[moe-cache] freeze profile restored: {restaged}/{requested} blocks restaged from {}",
4284 path.display()
4285 );
4286 Ok(true)
4287 }
4288
4289 pub fn freeze_cpu_expert_residency(
4291 &self,
4292 e: &Engine,
4293 ) -> Result<(), Box<dyn std::error::Error>> {
4294 e.freeze_moe_cache();
4295 Ok(())
4296 }
4297
4298 pub fn ffn_act(e: &Engine, cfg: &ModelConfig, gate: &CudaSlice<f32>, up: &CudaSlice<f32>,
4306 act: &mut CudaSlice<f32>, n: usize) -> Result<(), Box<dyn std::error::Error>> {
4307 Self::ffn_act_scaled(e, cfg, gate, up, 1.0, 1.0, act, n)
4308 }
4309
4310 #[allow(clippy::too_many_arguments)]
4314 pub(crate) fn ffn_act_scaled(e: &Engine, cfg: &ModelConfig, gate: &CudaSlice<f32>, up: &CudaSlice<f32>,
4315 gs: f32, us: f32, act: &mut CudaSlice<f32>, n: usize)
4316 -> Result<(), Box<dyn std::error::Error>> {
4317 Self::ffn_act_lim(e, cfg, gate, up, gs, us, None, act, n)
4318 }
4319
4320 #[allow(clippy::too_many_arguments)]
4329 pub(crate) fn ffn_act_lim(e: &Engine, cfg: &ModelConfig, gate: &CudaSlice<f32>, up: &CudaSlice<f32>,
4330 gs: f32, us: f32, limit: Option<f32>, act: &mut CudaSlice<f32>, n: usize)
4331 -> Result<(), Box<dyn std::error::Error>> {
4332 if let Some(m3) = cfg.m3.as_ref() {
4333 debug_assert!(limit.is_none(), "m3 swigluoai and step35 clamp are different archs");
4334 return e.swigluoai_mul_scaled(gate, up, gs, us, m3.swiglu_alpha, m3.swiglu_limit, act, n);
4335 }
4336 if let Some(l) = limit {
4337 return e.swiglu_clamped_mul_scaled(gate, up, gs, us, l, act, n);
4338 }
4339 if gs == 1.0 && us == 1.0 { return e.silu_mul(gate, up, act, n); }
4340 e.silu_mul_scaled(gate, up, gs, us, act, n)
4341 }
4342
4343 fn moe_route(e: &Engine, logits: &CudaSlice<f32>, t: usize, n_expert: usize, n_used: usize)
4349 -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
4350 Self::moe_route_cfg(e, logits, t, n_expert, n_used, None, None, None)
4351 }
4352
4353 fn moe_route_cfg(e: &Engine, logits: &CudaSlice<f32>, t: usize, n_expert: usize, n_used: usize,
4361 bias: Option<&[f32]>, sig: Option<(f32, bool)>, active: Option<&[bool]>)
4362 -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
4363 if let Some((sf, route_norm)) = sig {
4364 let lg = e.dtoh(logits)?;
4366 return Self::moe_route_sigmoid_host(
4367 &lg, t, n_expert, n_used, bias, sf, route_norm, active,
4368 );
4369 }
4370 if active.is_none() && !matches!(std::env::var("MEMRA_FUSED_ROUTER").as_deref(), Ok("0")) {
4374 return e.moe_router_topk_host(logits, t, n_expert, n_used);
4375 }
4376 let lg = e.dtoh(logits)?; let mut sel = vec![0u32; t * n_used];
4379 let mut w_out = vec![0f32; t * n_used];
4380 for tok in 0..t {
4381 let row = &lg[tok * n_expert..(tok + 1) * n_expert];
4382 let maxl = row.iter().enumerate()
4384 .filter(|(i, _)| active.is_none_or(|mask| mask[*i]))
4385 .map(|(_, &x)| x).fold(f32::NEG_INFINITY, f32::max);
4386 let mut probs = vec![0f32; n_expert];
4387 let mut den = 0f32;
4388 for i in 0..n_expert {
4389 if active.is_some_and(|mask| !mask[i]) { continue; }
4390 let x = (row[i] - maxl).exp(); probs[i] = x; den += x;
4391 }
4392 for p in probs.iter_mut() { *p /= den; }
4393 let mut idx: Vec<usize> = (0..n_expert)
4395 .filter(|&i| active.is_none_or(|mask| mask[i])).collect();
4396 idx.sort_by(|&a, &b| probs[b].total_cmp(&probs[a]).then(a.cmp(&b)));
4397 let sl = &idx[..n_used];
4398 let mut wv: Vec<f32> = sl.iter().map(|&i| probs[i]).collect();
4399 let mut ws: f32 = wv.iter().sum();
4400 ws = ws.max(6.103515625e-5_f32); for x in wv.iter_mut() { *x /= ws; }
4402 for j in 0..n_used {
4403 sel[tok * n_used + j] = sl[j] as u32;
4404 w_out[tok * n_used + j] = wv[j];
4405 }
4406 }
4407 Ok((sel, w_out))
4408 }
4409
4410 #[allow(clippy::too_many_arguments)]
4411 fn moe_route_sigmoid_with_input(
4412 e: &Engine,
4413 logits: &CudaSlice<f32>,
4414 input: &CudaSlice<f32>,
4415 t: usize,
4416 n_expert: usize,
4417 n_used: usize,
4418 bias: Option<&[f32]>,
4419 (sf, route_norm): (f32, bool),
4420 active: Option<&[bool]>,
4421 ) -> Result<(Vec<u32>, Vec<f32>, Vec<f32>), Box<dyn std::error::Error>> {
4422 let (lg, input) = e.dtoh_pair(logits, input)?;
4423 let (sel, w) =
4424 Self::moe_route_sigmoid_host(&lg, t, n_expert, n_used, bias, sf, route_norm, active)?;
4425 Ok((sel, w, input))
4426 }
4427
4428 pub fn start_moe_prefetch_predictor(
4433 &self,
4434 e: &Engine,
4435 cfg: &ModelConfig,
4436 ) -> Result<(), Box<dyn std::error::Error>> {
4437 use crate::hybrid::Ffn;
4438 let Some(sig) = cfg.sigmoid_router() else {
4439 return Err("prefetch predictor requires a sigmoid-router arch".into());
4440 };
4441 let resident: std::collections::HashSet<(u16, u8, u16)> = e
4442 .export_moe_residency()
4443 .ok_or("prefetch predictor needs the frozen MoE residency cache")?
4444 .into_iter()
4445 .collect();
4446 let mut layers = Vec::new();
4447 for (index, layer) in self.layers.iter().enumerate() {
4448 let Ffn::Moe(m) = &layer.ffn else { continue };
4449 let crate::model::GpuTensor::Float { data, .. } = &m.gate_inp else { continue };
4450 let router = e.dtoh(data)?;
4451 let n_expert = m.gate_exps.n_expert;
4452 let n_embd = m.gate_exps.in_f;
4453 if router.len() != n_embd * n_expert {
4454 continue;
4455 }
4456 let build = |exps: &crate::model::HostExps| {
4457 (0..n_expert)
4458 .map(|expert| crate::cpu_experts::predictor_projection(exps, expert))
4459 .collect::<Vec<_>>()
4460 };
4461 layers.push((index as u16, crate::cpu_experts::PredictLayerInit {
4462 router,
4463 bias: m.exp_probs_b.clone(),
4464 active: m.active_experts.clone(),
4465 n_embd,
4466 n_used: cfg
4467 .moe
4468 .as_ref()
4469 .map(|moe| moe.expert_used_count as usize)
4470 .ok_or("prefetch predictor requires MoE config")?,
4471 sig,
4472 weights_n_expert: n_expert,
4473 gate: build(&m.gate_exps),
4474 up: build(&m.up_exps),
4475 down: build(&m.down_exps),
4476 }));
4477 }
4478 crate::cpu_experts::start_prefetch_predictor(layers, resident)
4479 .map_err(|error| error.into())
4480 }
4481
4482 #[allow(clippy::too_many_arguments)]
4485 pub(crate) fn moe_route_sigmoid_host_public(
4486 logits: &[f32],
4487 t: usize,
4488 n_expert: usize,
4489 n_used: usize,
4490 bias: Option<&[f32]>,
4491 sf: f32,
4492 route_norm: bool,
4493 active: Option<&[bool]>,
4494 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
4495 Self::moe_route_sigmoid_host(logits, t, n_expert, n_used, bias, sf, route_norm, active)
4496 }
4497
4498 #[allow(clippy::too_many_arguments)]
4499 fn moe_route_sigmoid_host(
4500 lg: &[f32],
4501 t: usize,
4502 n_expert: usize,
4503 n_used: usize,
4504 bias: Option<&[f32]>,
4505 sf: f32,
4506 route_norm: bool,
4507 active: Option<&[bool]>,
4508 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
4509 if lg.len() != t * n_expert {
4510 return Err(format!(
4511 "sigmoid router logits length mismatch: got {}, expected {}",
4512 lg.len(),
4513 t * n_expert,
4514 )
4515 .into());
4516 }
4517 let mut sel = vec![0u32; t * n_used];
4518 let mut w_out = vec![0f32; t * n_used];
4519 for tok in 0..t {
4520 let row = &lg[tok * n_expert..(tok + 1) * n_expert];
4521 let scores: Vec<f32> = row.iter().map(|&x| 1.0 / (1.0 + (-x).exp())).collect();
4522 let selsc: Vec<f32> = match bias {
4524 Some(b) => scores.iter().zip(b).map(|(s, bb)| s + bb).collect(),
4525 None => scores.clone(),
4526 };
4527 let mut idx: Vec<usize> = (0..n_expert)
4528 .filter(|&i| active.is_none_or(|mask| mask[i]))
4529 .collect();
4530 idx.sort_by(|&a, &b| selsc[b].total_cmp(&selsc[a]).then(a.cmp(&b)));
4531 let sl = &idx[..n_used];
4532 let mut wv: Vec<f32> = sl.iter().map(|&i| scores[i]).collect();
4533 if route_norm {
4534 let ws: f32 = wv.iter().sum::<f32>().max(1e-20);
4535 for x in wv.iter_mut() {
4536 *x = *x / ws * sf;
4537 }
4538 } else {
4539 for x in wv.iter_mut() {
4540 *x *= sf;
4541 }
4542 }
4543 for j in 0..n_used {
4544 sel[tok * n_used + j] = sl[j] as u32;
4545 w_out[tok * n_used + j] = wv[j];
4546 }
4547 }
4548 Ok((sel, w_out))
4549 }
4550
4551 fn moe_ffn_pairs(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, logits: &CudaSlice<f32>,
4560 t: usize, cfg: &ModelConfig)
4561 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4562 let moe = cfg.moe.as_ref().unwrap();
4563 let n_embd = cfg.n_embd as usize;
4564 let n_expert = moe.expert_count as usize;
4565 let n_used = moe.expert_used_count as usize;
4566 let n_ff_exp = moe.expert_ff_length as usize;
4567 debug_assert!(!cfg.swiglu_clamped_anywhere(),
4572 "moe_ffn_pairs has no per-layer clamp: fused epilogues are plain SiLU");
4573 let dev = m.dev_exps.as_ref().unwrap();
4574 let (rbg_d, rbu_d) = if dev.gu_il {
4576 let sxx = m.gate_exps.row_bytes + m.up_exps.row_bytes; (sxx, sxx)
4577 } else { (m.gate_exps.row_bytes, m.up_exps.row_bytes) };
4578
4579 let (sel_all, w_all) = Self::moe_route(e, logits, t, n_expert, n_used)?;
4580 let n_pairs = t * n_used;
4581 let pair_tok: Vec<i32> = (0..n_pairs).map(|p| (p / n_used) as i32).collect();
4584 let pair_ex: Vec<i32> = sel_all.iter().map(|&x| x as i32).collect();
4585 let pair_w: Vec<f32> = w_all.clone();
4586 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
4587 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
4588 let pt = e.htod_i32(&pair_tok)?;
4589 let px = e.htod_i32(&pair_ex)?;
4590 let pw = e.htod(&pair_w)?;
4591 let toff = e.htod_i32(&tok_off)?;
4592 let tids = e.htod_i32(&tok_ids)?;
4593
4594 let mut by_ex: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
4598 for p in 0..n_pairs { by_ex[pair_ex[p] as usize].push(p as i32); }
4599 let mut ex_ids: Vec<i32> = Vec::new();
4600 let mut ex_off: Vec<i32> = vec![0];
4601 let mut ex_pairs: Vec<i32> = Vec::with_capacity(n_pairs);
4602 for (ex, list) in by_ex.iter().enumerate() {
4603 if list.is_empty() { continue; }
4604 ex_ids.push(ex as i32);
4605 ex_pairs.extend_from_slice(list);
4606 ex_off.push(ex_pairs.len() as i32);
4607 }
4608 let n_active = ex_ids.len();
4609 let exi = e.htod_i32(&ex_ids)?;
4610 let exo = e.htod_i32(&ex_off)?;
4611 let exp_d = e.htod_i32(&ex_pairs)?;
4612 let _ = &px; static MMA_T: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
4633 let mma_t = *MMA_T.get_or_init(|| {
4634 std::env::var("MEMRA_MOE_MMA_T").ok().and_then(|v| v.parse().ok()).unwrap_or(16)
4635 });
4636 let use_mma = std::env::var("MEMRA_MOE_MMA").map(|v| v != "0").unwrap_or(true)
4637 && t >= mma_t
4638 && q8_expert_dec_supported(m.gate_exps.qtype) && q8_expert_dec_supported(m.up_exps.qtype)
4639 && q8_expert_dec_supported(m.down_exps.qtype)
4640 && n_embd % 256 == 0 && n_ff_exp % 256 == 0;
4641 let mma_capable = q8_expert_dec_supported(m.gate_exps.qtype)
4657 && q8_expert_dec_supported(m.up_exps.qtype)
4658 && q8_expert_dec_supported(m.down_exps.qtype)
4659 && n_embd % 256 == 0 && n_ff_exp % 256 == 0;
4660 let f16g_mode = crate::moe_f16g_mode();
4661 let f16g = f16g_mode != 0 && t >= mma_t
4662 && (f16g_mode != 3 || !mma_capable)
4663 && f16g_proj_ok(m.gate_exps.qtype, n_embd)
4664 && f16g_proj_ok(m.up_exps.qtype, n_embd)
4665 && f16g_proj_ok(m.down_exps.qtype, n_ff_exp);
4666 if use_mma || f16g {
4667 let y_down = if f16g {
4675 let csr_tok: Vec<i32> = ex_pairs.iter().map(|&p| p / n_used as i32).collect();
4679 let csr_tok_d = e.htod_i32(&csr_tok)?;
4680 let (z_f16, z_s) = e.moe_f16g_act(z, Some(&csr_tok_d), n_embd, n_pairs)?;
4681 let g_csr = e.moe_f16_grouped(&dev.ptr_row, 0, n_expert, &exi, &ex_off, &exo,
4682 &z_f16, &z_s, n_embd, n_ff_exp, n_active, n_pairs,
4683 m.gate_exps.qtype, rbg_d)?;
4684 let u_csr = e.moe_f16_grouped(&dev.ptr_row, 1, n_expert, &exi, &ex_off, &exo,
4685 &z_f16, &z_s, n_embd, n_ff_exp, n_active, n_pairs,
4686 m.up_exps.qtype, rbu_d)?;
4687 let act_csr = e.moe_pairs_silu_mul(&g_csr, &u_csr, n_pairs * n_ff_exp)?;
4688 let (a_f16, a_s) = e.moe_f16g_act(&act_csr, None, n_ff_exp, n_pairs)?;
4689 let d_csr = e.moe_f16_grouped(&dev.ptr_row, 2, n_expert, &exi, &ex_off, &exo,
4690 &a_f16, &a_s, n_ff_exp, n_embd, n_active, n_pairs,
4691 m.down_exps.qtype, m.down_exps.row_bytes)?;
4692 e.rows_permute(&d_csr, &exp_d, n_pairs, n_embd)?
4693 } else {
4694 let z_scr = e.mmq_iq_quantize_act(z, n_embd, t)?;
4696 let gate = e.mmq_iq_experts(&dev.ptr_row, 0, n_expert, &exi, &exo, &exp_d, &pt, &z_scr,
4697 n_embd, n_ff_exp, n_active, n_pairs, t,
4698 m.gate_exps.qtype, rbg_d)?;
4699 let up = e.mmq_iq_experts(&dev.ptr_row, 1, n_expert, &exi, &exo, &exp_d, &pt, &z_scr,
4700 n_embd, n_ff_exp, n_active, n_pairs, t,
4701 m.up_exps.qtype, rbu_d)?;
4702 let a_scr = if crate::moe_fuse_actq_on() {
4708 e.mmq_iq_fused_act_quant(&gate, &up, n_ff_exp, n_pairs, 0)?
4709 } else {
4710 let act = e.moe_pairs_silu_mul(&gate, &up, n_pairs * n_ff_exp)?;
4711 e.mmq_iq_quantize_act(&act, n_ff_exp, n_pairs)?
4712 };
4713 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
4714 let pself = e.htod_i32(&pair_self)?;
4715 e.mmq_iq_experts(&dev.ptr_row, 2, n_expert, &exi, &exo, &exp_d, &pself, &a_scr,
4716 n_ff_exp, n_embd, n_active, n_pairs, n_pairs,
4717 m.down_exps.qtype, m.down_exps.row_bytes)?
4718 };
4719 let mut moe_out = e.uninit(t * n_embd)?;
4720 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
4721 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
4722 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
4723 {
4724 let n_ff_sh = gate_shexp.out_features();
4725 let sg_gate = e.matmul(gate_shexp, z, t)?;
4726 let sg_up = e.matmul(up_shexp, z, t)?;
4727 let mut sa = e.uninit(t * n_ff_sh)?;
4728 Self::ffn_act(e, cfg, &sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
4729 let sh = e.matmul(down_shexp, &sa, t)?;
4730 let g = match &m.gate_inp_shexp {
4736 Some(gate_inp_shexp) if crate::router_prefill_exact_on() => {
4737 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
4738 }
4739 Some(gate_inp_shexp) => {
4740 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
4741 let mut g = e.uninit(t)?;
4742 e.sigmoid(&gs, &mut g, t)?;
4743 g
4744 }
4745 None => e.htod(&vec![1.0f32; t])?,
4746 };
4747 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
4748 }
4749 return Ok(moe_out);
4750 }
4751
4752 let dec = std::env::var("MEMRA_MOE_DEC").map(|v| v != "0").unwrap_or(true);
4755 let matvec = |proj, exi: &_, exo: &_, exp_d: &_, pt: &_, aq: &_, ad: &_,
4756 inf, outf, qtype, rb| -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4757 let dec = dec && q8_expert_dec_supported(qtype);
4759 if dec { e.moe_pairs_matvec_q8_dec(&dev.ptr_row, proj, exi, exo, exp_d, pt, aq, ad,
4760 inf, outf, n_expert, n_active, n_pairs, qtype, rb) }
4761 else { e.moe_pairs_matvec_q8_em (&dev.ptr_row, proj, exi, exo, exp_d, pt, aq, ad,
4762 inf, outf, n_expert, n_active, n_pairs, qtype, rb) }
4763 };
4764 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
4765 let gate = matvec(0, &exi, &exo, &exp_d, &pt, &zq, &zd,
4766 n_embd, n_ff_exp, m.gate_exps.qtype, rbg_d)?;
4767 let up = matvec(1, &exi, &exo, &exp_d, &pt, &zq, &zd,
4768 n_embd, n_ff_exp, m.up_exps.qtype, rbu_d)?;
4769 let act = e.moe_pairs_silu_mul(&gate, &up, n_pairs * n_ff_exp)?;
4770 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
4771 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
4773 let pself = e.htod_i32(&pair_self)?;
4774 let y_down = matvec(2, &exi, &exo, &exp_d, &pself, &aq2, &ad2,
4775 n_ff_exp, n_embd, m.down_exps.qtype, m.down_exps.row_bytes)?;
4776 let mut moe_out = e.uninit(t * n_embd)?; e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
4778
4779 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
4783 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
4784 {
4785 let n_ff_sh = gate_shexp.out_features();
4786 let sg_gate = e.matmul(gate_shexp, z, t)?;
4787 let sg_up = e.matmul(up_shexp, z, t)?;
4788 let mut sa = e.uninit(t * n_ff_sh)?;
4789 e.silu_mul(&sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
4790 let sh = e.matmul(down_shexp, &sa, t)?;
4791 let g = match &m.gate_inp_shexp {
4796 Some(gate_inp_shexp) if crate::router_prefill_exact_on() => {
4797 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
4798 }
4799 Some(gate_inp_shexp) => {
4800 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
4801 let mut g = e.uninit(t)?;
4802 e.sigmoid(&gs, &mut g, t)?;
4803 g
4804 }
4805 None => e.htod(&vec![1.0f32; t])?,
4806 };
4807 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
4808 }
4809 Ok(moe_out)
4810 }
4811
4812 #[allow(clippy::too_many_arguments)]
4814 #[allow(clippy::too_many_arguments)]
4815 fn moe_ffn_dev(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>,
4816 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>, logits: &CudaSlice<f32>,
4817 t: usize, cfg: &ModelConfig, il: u16, max_block: usize)
4818 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4819 let moe = cfg.moe.as_ref().unwrap();
4820 let n_embd = cfg.n_embd as usize;
4821 let n_expert = moe.expert_count as usize;
4822 let n_used = moe.expert_used_count as usize;
4823 let n_ff_exp = moe.expert_ff_length as usize;
4824 debug_assert!(cfg.sigmoid_router().is_none(),
4828 "moe_ffn_dev routes SOFTMAX: a sigmoid-router arch would pick wrong experts");
4829 debug_assert!(!cfg.swiglu_clamped_at(il as u32),
4830 "moe_ffn_dev's fused epilogue is plain SiLU: no clamped form");
4831
4832 let (sel_d, mut w_d) = e.moe_router_topk(logits, t, n_expert, n_used)?;
4834 if m.has_macros {
4837 e.moe_w_scale_by_expert(&mut w_d, &sel_d, &m.dev_macros, n_expert, t * n_used)?;
4838 }
4839
4840 let mut moe_out = e.uninit(t * n_embd)?;
4842
4843 if let Some(dev) = m.dev_exps.as_ref() {
4846 let (rbg_d, rbu_d) = if dev.gu_il {
4849 let sxx = m.gate_exps.row_bytes + m.up_exps.row_bytes; (sxx, sxx)
4850 } else { (m.gate_exps.row_bytes, m.up_exps.row_bytes) };
4851 let q8 = moe_q8_enabled()
4852 && q8_expert_supported(m.gate_exps.qtype) && q8_expert_supported(m.up_exps.qtype)
4853 && q8_expert_supported(m.down_exps.qtype);
4854 let rows_arm = q8 && t > 1 && crate::spec::spec_m2()
4863 && n_ff_exp == 512 && n_used <= 8
4864 && std::env::var("MEMRA_MOE_DEVQ8_GU").map(|v| v.is_empty() || v == "v").unwrap_or(true)
4865 && std::env::var("MEMRA_MOE_DEVQ8_DOWN").map(|v| v.is_empty() || v == "w8h2v").unwrap_or(true);
4866 let csr_mode = std::env::var("MEMRA_MOE_CSR").ok()
4875 .and_then(|v| v.parse::<i32>().ok()).unwrap_or(1);
4876 let csr_qt = |qt: i32| qt == crate::QT_IQ4_XS || qt == crate::QT_IQ3_S;
4877 let csr_arm = rows_arm && csr_mode > 0 && t <= 10
4878 && csr_qt(m.gate_exps.qtype) && csr_qt(m.up_exps.qtype)
4879 && csr_qt(m.down_exps.qtype);
4880 if csr_arm {
4881 if csr_mode == 2 {
4882 static ENGAGED: std::sync::Once = std::sync::Once::new();
4883 ENGAGED.call_once(|| eprintln!("[memra] moe CSR byte-compare mode ON (t={t})"));
4884 }
4885 let n_pairs = t * n_used;
4886 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
4887 let act = e.moe_gate_up_silu8_dev_q8_csr(&dev.ptr_row, &sel_d, &zq, &zd, n_pairs,
4888 n_embd, n_ff_exp, n_used, n_expert,
4889 m.gate_exps.qtype, m.up_exps.qtype,
4890 rbg_d, rbu_d)?;
4891 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
4892 e.moe_down8_fma_dev_q8_rows(&dev.ptr_row, &sel_d, &w_d, &aq2, &ad2, &mut moe_out,
4896 t, n_ff_exp, n_embd, n_used, n_expert,
4897 m.down_exps.qtype, m.down_exps.row_bytes)?;
4898 if csr_mode == 2 {
4899 let act_r = e.moe_gate_up_silu8_dev_q8_rows(&dev.ptr_row, &sel_d, &zq, &zd, t,
4901 n_embd, n_ff_exp, n_used, n_expert,
4902 m.gate_exps.qtype, m.up_exps.qtype,
4903 rbg_d, rbu_d, &m.dev_macros)?;
4904 let mut out_r = e.uninit(t * n_embd)?;
4905 let (aq2r, ad2r) = e.quantize_q8_1(&act_r, n_pairs, n_ff_exp)?;
4906 e.moe_down8_fma_dev_q8_rows(&dev.ptr_row, &sel_d, &w_d, &aq2r, &ad2r, &mut out_r,
4907 t, n_ff_exp, n_embd, n_used, n_expert,
4908 m.down_exps.qtype, m.down_exps.row_bytes)?;
4909 let (a1, a2) = (e.dtoh(&act)?, e.dtoh(&act_r)?);
4910 let (o1, o2) = (e.dtoh(&moe_out)?, e.dtoh(&out_r)?);
4911 let ba = a1.iter().zip(&a2).filter(|(x, y)| x.to_bits() != y.to_bits()).count();
4912 let bo = o1.iter().zip(&o2).filter(|(x, y)| x.to_bits() != y.to_bits()).count();
4913 if ba + bo > 0 {
4914 eprintln!("[csr-check] il={il} t={t} ACT diffs={ba}/{} OUT diffs={bo}/{}",
4915 a1.len(), o1.len());
4916 let sel_h = e.dtoh_i32(&sel_d)?;
4918 let mut shown = 0;
4919 for (i, (x, y)) in a1.iter().zip(&a2).enumerate() {
4920 if x.to_bits() != y.to_bits() && shown < 4 {
4921 let (p, o) = (i / n_ff_exp, i % n_ff_exp);
4922 let ex = sel_h[p];
4923 let npx = sel_h.iter().filter(|&&v| v == ex).count();
4924 eprintln!(" ACT p={p} ex={ex} np={npx} o={o} csr={x:e} rows={y:e}");
4925 shown += 1;
4926 }
4927 }
4928 std::process::exit(3);
4929 }
4930 }
4931 } else if rows_arm {
4932 if std::env::var("MEMRA_MOE_OVERLAP").as_deref() == Ok("1") {
4935 use std::sync::atomic::{AtomicU64, Ordering};
4936 static PAIRS: AtomicU64 = AtomicU64::new(0);
4937 static UNIQ: AtomicU64 = AtomicU64::new(0);
4938 static CALLS: AtomicU64 = AtomicU64::new(0);
4939 let sel_h = e.dtoh_i32(&sel_d)?;
4940 let mut u: Vec<i32> = sel_h.clone(); u.sort_unstable(); u.dedup();
4941 PAIRS.fetch_add(sel_h.len() as u64, Ordering::Relaxed);
4942 UNIQ.fetch_add(u.len() as u64, Ordering::Relaxed);
4943 let c = CALLS.fetch_add(1, Ordering::Relaxed) + 1;
4944 if c % 480 == 0 {
4945 let p = PAIRS.load(Ordering::Relaxed); let q = UNIQ.load(Ordering::Relaxed);
4946 eprintln!("[overlap] calls={c} pairs={p} unique={q} ratio={:.3} (t={t})",
4947 q as f64 / p as f64);
4948 }
4949 }
4950 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
4951 let act = e.moe_gate_up_silu8_dev_q8_rows(&dev.ptr_row, &sel_d, &zq, &zd, t,
4952 n_embd, n_ff_exp, n_used, n_expert,
4953 m.gate_exps.qtype, m.up_exps.qtype,
4954 rbg_d, rbu_d, &m.dev_macros)?;
4955 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
4956 e.moe_down8_fma_dev_q8_rows(&dev.ptr_row, &sel_d, &w_d, &aq2, &ad2, &mut moe_out,
4957 t, n_ff_exp, n_embd, n_used, n_expert,
4958 m.down_exps.qtype, m.down_exps.row_bytes)?;
4959 } else {
4960 for tok in 0..t {
4961 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
4962 let selt = sel_d.slice(tok * n_used..(tok + 1) * n_used);
4963 let wt = w_d.slice(tok * n_used..(tok + 1) * n_used);
4964 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
4965 if q8 {
4966 let (zq, zd) = match (t, zq8) {
4967 (1, Some((q, d))) => (q.clone(), d.clone()),
4968 _ => e.quantize_q8_1_view(&zt, 1, n_embd)?,
4969 };
4970 let act = e.moe_gate_up_silu8_dev_q8(&dev.ptr_row, &selt, &zq, &zd,
4971 n_embd, n_ff_exp, n_used, n_expert,
4972 m.gate_exps.qtype, m.up_exps.qtype,
4973 rbg_d, rbu_d, &m.dev_macros)?;
4974 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
4975 e.moe_down8_fma_dev_q8(&dev.ptr_row, &selt, &wt, &aq2, &ad2, &mut dst,
4976 n_ff_exp, n_embd, n_used, n_expert,
4977 m.down_exps.qtype, m.down_exps.row_bytes)?;
4978 } else {
4979 let act = e.moe_gate_up_silu8_dev(&dev.ptr_row, &selt, &zt, n_embd, n_ff_exp,
4980 n_used, n_expert,
4981 m.gate_exps.qtype, m.up_exps.qtype,
4982 rbg_d, rbu_d, &m.dev_macros)?;
4983 e.moe_down8_fma_dev(&dev.ptr_row, &selt, &wt, &act, &mut dst,
4984 n_ff_exp, n_embd, n_used, n_expert,
4985 m.down_exps.qtype, m.down_exps.row_bytes)?;
4986 }
4987 }
4988 }
4989 } else {
4990 let q8 = moe_q8_enabled()
4997 && q8_expert_supported(m.gate_exps.qtype) && q8_expert_supported(m.up_exps.qtype)
4998 && q8_expert_supported(m.down_exps.qtype);
4999 e.with_moe_cache(max_block, |c, eng| {
5000 let row = c.layer_dev_row(il, n_expert, eng)?
5001 .ok_or("moe_ffn_dev: layer row vanished under the lock")?;
5002 for tok in 0..t {
5003 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
5004 let selt = sel_d.slice(tok * n_used..(tok + 1) * n_used);
5005 let wt = w_d.slice(tok * n_used..(tok + 1) * n_used);
5006 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
5007 if q8 {
5008 let (zq, zd) = match (t, zq8) {
5009 (1, Some((q, d))) => (q.clone(), d.clone()),
5010 _ => eng.quantize_q8_1_view(&zt, 1, n_embd)?,
5011 };
5012 let act = eng.moe_gate_up_silu8_dev_q8(row, &selt, &zq, &zd,
5013 n_embd, n_ff_exp, n_used, n_expert,
5014 m.gate_exps.qtype, m.up_exps.qtype,
5015 m.gate_exps.row_bytes, m.up_exps.row_bytes,
5016 &m.dev_macros)?;
5017 let (aq2, ad2) = eng.quantize_q8_1(&act, n_used, n_ff_exp)?;
5018 eng.moe_down8_fma_dev_q8(row, &selt, &wt, &aq2, &ad2, &mut dst,
5019 n_ff_exp, n_embd, n_used, n_expert,
5020 m.down_exps.qtype, m.down_exps.row_bytes)?;
5021 } else {
5022 let act = eng.moe_gate_up_silu8_dev(row, &selt, &zt, n_embd, n_ff_exp,
5023 n_used, n_expert,
5024 m.gate_exps.qtype, m.up_exps.qtype,
5025 m.gate_exps.row_bytes, m.up_exps.row_bytes,
5026 &m.dev_macros)?;
5027 eng.moe_down8_fma_dev(row, &selt, &wt, &act, &mut dst,
5028 n_ff_exp, n_embd, n_used, n_expert,
5029 m.down_exps.qtype, m.down_exps.row_bytes)?;
5030 }
5031 }
5032 c.hits += (t * 3 * n_used) as u64;
5034 Ok(())
5035 })?;
5036 }
5037
5038 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
5043 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
5044 {
5045 let n_ff_sh = gate_shexp.out_features();
5046 let verify_t = t > 1 && t < PRIME_MIN_T;
5049 let (sg_gate, sg_up) = if t == 1 {
5050 match e.matmul_q8_fused2_x(gate_shexp, up_shexp, z)? {
5051 Some(pair) => pair,
5052 None => (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?),
5053 }
5054 } else if verify_t {
5055 let mut fused = None;
5059 if crate::spec::spec_fused_t() && (2..=4).contains(&t)
5060 && e.uses_q8_1_fast(gate_shexp) && e.uses_q8_1_fast(up_shexp) {
5061 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
5062 fused = e.matmul_q8_fused2_t(gate_shexp, up_shexp, &zq, &zd, t)?;
5063 }
5064 match fused {
5065 Some(pair) => pair,
5066 None => (e.matmul_decode_exact(gate_shexp, z, t)?,
5067 e.matmul_decode_exact(up_shexp, z, t)?),
5068 }
5069 } else {
5070 (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?)
5071 };
5072 let mut sa = e.uninit(t * n_ff_sh)?; e.silu_mul(&sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
5074 let sh = if verify_t { e.matmul_decode_exact(down_shexp, &sa, t)? }
5075 else { e.matmul(down_shexp, &sa, t)? };
5076 let g = match &m.gate_inp_shexp {
5080 Some(gate_inp_shexp) => {
5081 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
5084 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
5085 } else {
5086 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
5087 let mut g = e.uninit(t)?;
5088 e.sigmoid(&gs, &mut g, t)?;
5089 g
5090 }
5091 }
5092 None => e.htod(&vec![1.0f32; t])?,
5093 };
5094 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
5095 }
5096
5097 Ok(moe_out)
5098 }
5099
5100 #[allow(clippy::too_many_arguments)]
5110 #[allow(clippy::too_many_arguments)]
5113 fn moe_gdec_token_q8(e: &Engine, m: &MoeWeights, il: u16, max_block: usize,
5114 zq: &CudaSlice<i8>, zd: &CudaSlice<f32>, sel: &[u32], w: &[f32],
5115 moe_out: &mut CudaSlice<f32>, tok: usize,
5116 n_embd: usize, n_ff_exp: usize, n_used: usize)
5117 -> Result<bool, Box<dyn std::error::Error>> {
5118 use crate::moe_cache::{BlockId, PROJ_GATE, PROJ_UP, PROJ_DOWN};
5119 use cudarc::driver::DevicePtr;
5120 let ptrs = e.with_moe_cache(max_block, |c, eng| {
5121 let mut g = [0u64; 8];
5122 let mut u = [0u64; 8];
5123 let mut d = [0u64; 8];
5124 for (j, &ex) in sel.iter().enumerate() {
5125 let ex = ex as u16;
5126 let (Some(sg), Some(su), Some(sd)) = (c.resident(BlockId::new(il, PROJ_GATE, ex)),
5127 c.resident(BlockId::new(il, PROJ_UP, ex)),
5128 c.resident(BlockId::new(il, PROJ_DOWN, ex)))
5129 else { return Ok(None); };
5130 let __s = eng.stream();
5131 let (pg, _e0) = c.slot(sg).device_ptr(&__s);
5132 let (pu, _e1) = c.slot(su).device_ptr(&__s);
5133 let (pd, _e2) = c.slot(sd).device_ptr(&__s);
5134 g[j] = pg as u64; u[j] = pu as u64; d[j] = pd as u64;
5135 }
5136 if cpu_expert_profile_admit_enabled() && !c.is_frozen() {
5137 for &ex in sel {
5138 let ex = ex as u16;
5139 for proj in [PROJ_GATE, PROJ_UP, PROJ_DOWN] {
5140 c.note_profile_hit(BlockId::new(il, proj, ex));
5141 }
5142 }
5143 }
5144 c.hits += (3 * n_used) as u64;
5145 Ok(Some((g, u, d)))
5146 })?;
5147 let Some((g, u, d)) = ptrs else { return Ok(false) };
5148 let mut wv = [0f32; 8];
5149 wv[..n_used].copy_from_slice(w);
5150 let act = e.moe_gate_up_silu8_q8(crate::WPtr8(g), crate::WPtr8(u), zq, zd,
5151 n_embd, n_ff_exp, n_used,
5152 m.gate_exps.qtype, m.up_exps.qtype,
5153 m.gate_exps.row_bytes, m.up_exps.row_bytes)?;
5154 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
5156 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
5157 e.moe_down8_fma_q8(crate::WPtr8(d), crate::F32x8(wv), &aq2, &ad2, &mut dst,
5158 n_ff_exp, n_embd, n_used,
5159 m.down_exps.qtype, m.down_exps.row_bytes)?;
5160 Ok(true)
5161 }
5162
5163 fn moe_gdec_token(e: &Engine, m: &MoeWeights, il: u16, max_block: usize,
5164 zt: &cudarc::driver::CudaView<f32>, sel: &[u32], w: &[f32],
5165 moe_out: &mut CudaSlice<f32>, tok: usize,
5166 n_embd: usize, n_ff_exp: usize, n_used: usize)
5167 -> Result<bool, Box<dyn std::error::Error>> {
5168 use crate::moe_cache::{BlockId, PROJ_GATE, PROJ_UP, PROJ_DOWN};
5169 use cudarc::driver::DevicePtr;
5170 let ptrs = e.with_moe_cache(max_block, |c, eng| {
5172 let mut g = [0u64; 8];
5173 let mut u = [0u64; 8];
5174 let mut d = [0u64; 8];
5175 for (j, &ex) in sel.iter().enumerate() {
5176 let ex = ex as u16;
5177 let (Some(sg), Some(su), Some(sd)) = (c.resident(BlockId::new(il, PROJ_GATE, ex)),
5178 c.resident(BlockId::new(il, PROJ_UP, ex)),
5179 c.resident(BlockId::new(il, PROJ_DOWN, ex)))
5180 else { return Ok(None); };
5181 let __s = eng.stream();
5182 let (pg, _e0) = c.slot(sg).device_ptr(&__s);
5183 let (pu, _e1) = c.slot(su).device_ptr(&__s);
5184 let (pd, _e2) = c.slot(sd).device_ptr(&__s);
5185 g[j] = pg as u64; u[j] = pu as u64; d[j] = pd as u64;
5186 }
5187 if cpu_expert_profile_admit_enabled() && !c.is_frozen() {
5188 for &ex in sel {
5189 let ex = ex as u16;
5190 for proj in [PROJ_GATE, PROJ_UP, PROJ_DOWN] {
5191 c.note_profile_hit(BlockId::new(il, proj, ex));
5192 }
5193 }
5194 }
5195 c.hits += (3 * n_used) as u64; Ok(Some((g, u, d)))
5197 })?;
5198 let Some((g, u, d)) = ptrs else { return Ok(false) };
5199 let mut wv = [0f32; 8];
5200 wv[..n_used].copy_from_slice(w);
5201 let act = e.moe_gate_up_silu8(crate::WPtr8(g), crate::WPtr8(u), zt,
5203 n_embd, n_ff_exp, n_used,
5204 m.gate_exps.qtype, m.up_exps.qtype,
5205 m.gate_exps.row_bytes, m.up_exps.row_bytes)?;
5206 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
5207 e.moe_down8_fma_into(crate::WPtr8(d), crate::F32x8(wv), &act, &mut dst,
5208 n_ff_exp, n_embd, n_used,
5209 m.down_exps.qtype, m.down_exps.row_bytes)?;
5210 Ok(true)
5211 }
5212
5213 fn moe_cached_gemm_q8(e: &Engine, il: u16, proj: u8, ex: usize, m: &MoeWeights,
5218 max_block: usize, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
5219 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5220 use crate::moe_cache::{BlockId, DispatchSlot, PROJ_GATE, PROJ_UP};
5221 let exps = match proj { PROJ_GATE => &m.gate_exps, PROJ_UP => &m.up_exps, _ => &m.down_exps };
5222 let layout = exps.expert_layout(ex);
5223 let id = BlockId::new(il, proj, ex as u16);
5224 let source = exps.expert_source(ex);
5225 e.with_moe_cache(max_block, |c, eng| {
5226 let slot = c.dispatch_source(id, source, eng)?;
5227 let DispatchSlot::Resident(sl) = slot;
5228 let buf = c.slot(sl);
5229 eng.qmatvec_expert_q8(buf, 0..layout.len, aq, ad, 1, exps.in_f, exps.out_f,
5230 layout.qtype, layout.row_bytes)
5231 })
5232 }
5233
5234 fn moe_cached_gemm(e: &Engine, il: u16, proj: u8, ex: usize, m: &MoeWeights,
5235 max_block: usize, x: &cudarc::driver::CudaView<f32>)
5236 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5237 use crate::moe_cache::{BlockId, DispatchSlot, PROJ_GATE, PROJ_UP};
5238 let exps = match proj { PROJ_GATE => &m.gate_exps, PROJ_UP => &m.up_exps, _ => &m.down_exps };
5239 let layout = exps.expert_layout(ex);
5240 let id = BlockId::new(il, proj, ex as u16);
5241 let source = exps.expert_source(ex);
5242 e.with_moe_cache(max_block, |c, eng| {
5244 let slot = c.dispatch_source(id, source, eng)?;
5245 let DispatchSlot::Resident(sl) = slot;
5248 let buf = c.slot(sl);
5249 eng.qmatvec_view(buf, 0..layout.len, x, 1, exps.in_f, exps.out_f,
5250 layout.qtype, layout.row_bytes)
5251 })
5252 }
5253
5254 fn moe_profile_admit_expert(
5258 e: &Engine,
5259 il: u16,
5260 ex: usize,
5261 m: &MoeWeights,
5262 max_block: usize,
5263 ) -> Result<(), Box<dyn std::error::Error>> {
5264 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
5265 e.with_moe_cache(max_block, |cache, eng| {
5266 for (proj, exps) in [
5267 (PROJ_GATE, &m.gate_exps),
5268 (PROJ_UP, &m.up_exps),
5269 (PROJ_DOWN, &m.down_exps),
5270 ] {
5271 let id = BlockId::new(il, proj, ex as u16);
5272 let _ = cache.dispatch_source(id, exps.expert_source(ex), eng)?;
5273 }
5274 Ok(())
5275 })
5276 }
5277
5278 #[allow(clippy::too_many_arguments)]
5281 fn moe_frozen_gemm(
5282 e: &Engine,
5283 il: u16,
5284 proj: u8,
5285 ex: usize,
5286 m: &MoeWeights,
5287 max_block: usize,
5288 x: &cudarc::driver::CudaView<f32>,
5289 scratch: &mut Option<CudaSlice<u8>>,
5290 scratch_len: usize,
5291 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5292 use crate::moe_cache::{BlockId, PROJ_GATE, PROJ_UP};
5293 let exps = match proj {
5294 PROJ_GATE => &m.gate_exps,
5295 PROJ_UP => &m.up_exps,
5296 _ => &m.down_exps,
5297 };
5298 let layout = exps.expert_layout(ex);
5299 let id = BlockId::new(il, proj, ex as u16);
5300 if let Some(output) = e.with_moe_cache(max_block, |cache, eng| {
5301 let Some(slot) = cache.resident(id) else {
5302 return Ok(None);
5303 };
5304 let buf = cache.slot(slot);
5305 Ok(Some(eng.qmatvec_view(
5306 buf,
5307 0..layout.len,
5308 x,
5309 1,
5310 exps.in_f,
5311 exps.out_f,
5312 layout.qtype,
5313 layout.row_bytes,
5314 )?))
5315 })? {
5316 return Ok(output);
5317 }
5318 if scratch.is_none() {
5319 *scratch = Some(e.alloc_u8_uninit(scratch_len)?);
5320 }
5321 let scratch = scratch.as_mut().unwrap();
5322 e.stage_expert(exps.expert_bytes(ex), scratch, 0)?;
5323 e.qmatvec_view(
5324 scratch,
5325 0..layout.len,
5326 x,
5327 1,
5328 exps.in_f,
5329 exps.out_f,
5330 layout.qtype,
5331 layout.row_bytes,
5332 )
5333 }
5334
5335 fn moe_prefetch_expert(
5336 e: &Engine,
5337 il: u16,
5338 ex: usize,
5339 m: &MoeWeights,
5340 max_block: usize,
5341 keep: &[crate::moe_cache::BlockId],
5342 ) -> Result<(), Box<dyn std::error::Error>> {
5343 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
5344 e.with_moe_cache(max_block, |c, eng| {
5345 for (proj, exps) in [(PROJ_GATE, &m.gate_exps), (PROJ_UP, &m.up_exps),
5346 (PROJ_DOWN, &m.down_exps)] {
5347 let id = BlockId::new(il, proj, ex as u16);
5348 let _ = c.prefetch_source(id, exps.expert_source(ex), keep, eng)?;
5349 }
5350 Ok(())
5351 })
5352 }
5353
5354 fn moe_prefetch_disk_expert(e: &Engine, il: u16, ex: usize, m: &MoeWeights,
5357 max_block: usize, keep: &[crate::moe_cache::BlockId])
5358 -> Result<(), Box<dyn std::error::Error>> {
5359 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
5360 e.with_moe_cache(max_block, |c, eng| {
5361 for (proj, exps) in [(PROJ_GATE, &m.gate_exps), (PROJ_UP, &m.up_exps),
5362 (PROJ_DOWN, &m.down_exps)] {
5363 let source = exps.expert_source(ex);
5364 if let crate::model::ExpertSource::Disk { .. } = &source {
5365 let id = BlockId::new(il, proj, ex as u16);
5366 let _ = c.prefetch_source(id, source, keep, eng)?;
5367 }
5368 }
5369 Ok(())
5370 })
5371 }
5372
5373 #[inline]
5374 fn moe_prefetch_host_expert(ex: usize, m: &MoeWeights) {
5375 let _ = m.gate_exps.prefetch_expert_pages(ex);
5376 let _ = m.up_exps.prefetch_expert_pages(ex);
5377 let _ = m.down_exps.prefetch_expert_pages(ex);
5378 }
5379}
5380
5381impl HybridModel {
5398 #[allow(clippy::too_many_arguments)]
5402 fn moe_ffn_grouped_resident_q8(
5403 e: &Engine,
5404 m: &MoeWeights,
5405 z: &CudaSlice<f32>,
5406 t: usize,
5407 cfg: &ModelConfig,
5408 il: u16,
5409 sel_all: &[u32],
5410 w_all: &[f32],
5411 table: &CudaSlice<u64>,
5412 gu_il: bool,
5413 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5414 let moe = cfg.moe.as_ref().unwrap();
5415 let n_embd = cfg.n_embd as usize;
5416 let n_expert = moe.expert_count as usize;
5417 let n_used = moe.expert_used_count as usize;
5418 let n_ff_exp = moe.expert_ff_length as usize;
5419 let n_pairs = t * n_used;
5420 debug_assert_eq!(sel_all.len(), n_pairs);
5421 debug_assert_eq!(w_all.len(), n_pairs);
5422 debug_assert!(
5423 m.gate_exps.macros.is_none()
5424 && m.up_exps.macros.is_none()
5425 && m.down_exps.macros.is_none(),
5426 "resident grouped q8 does not fold per-expert macro scales",
5427 );
5428
5429 if !cfg.swiglu_clamped_at(il as u32) {
5434 let sel: Vec<i32> = sel_all.iter().map(|&expert| expert as i32).collect();
5435 let sel_d = e.htod_i32(&sel)?;
5436 let w_d = e.htod(w_all)?;
5437 let (gate_row_bytes, up_row_bytes) = if gu_il {
5438 let combined = m.gate_exps.row_bytes + m.up_exps.row_bytes;
5439 (combined, combined)
5440 } else {
5441 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
5442 };
5443 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
5444 let act = e.moe_gate_up_silu8_dev_q8_rows(
5445 table,
5446 &sel_d,
5447 &zq,
5448 &zd,
5449 t,
5450 n_embd,
5451 n_ff_exp,
5452 n_used,
5453 n_expert,
5454 m.gate_exps.qtype,
5455 m.up_exps.qtype,
5456 gate_row_bytes,
5457 up_row_bytes,
5458 &m.dev_macros,
5459 )?;
5460 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
5461 let mut moe_out = e.uninit(t * n_embd)?;
5462 e.moe_down8_fma_dev_q8_rows_g(
5463 table,
5464 &sel_d,
5465 &w_d,
5466 &aq2,
5467 &ad2,
5468 &mut moe_out,
5469 t,
5470 n_ff_exp,
5471 n_embd,
5472 n_used,
5473 n_expert,
5474 m.down_exps.qtype,
5475 m.down_exps.row_bytes,
5476 )?;
5477
5478 if std::env::var("MEMRA_MOE_STATS").is_ok() {
5479 let mut counts = vec![0usize; n_expert];
5480 for &expert in sel_all {
5481 counts[expert as usize] += 1;
5482 }
5483 let mut sizes: Vec<usize> =
5484 counts.into_iter().filter(|&count| count != 0).collect();
5485 sizes.sort_unstable();
5486 let mean = sizes.iter().sum::<usize>() as f64 / sizes.len().max(1) as f64;
5487 println!(
5488 "moe-grouped il={il} t={t} dispatch=resident-q8-rows active={}/{} \
5489 m_e: min={} median={} mean={mean:.1} max={}",
5490 sizes.len(),
5491 n_expert,
5492 sizes.first().copied().unwrap_or(0),
5493 sizes.get(sizes.len() / 2).copied().unwrap_or(0),
5494 sizes.last().copied().unwrap_or(0),
5495 );
5496 }
5497 return Ok(moe_out);
5498 }
5499
5500 let pair_tok: Vec<i32> = (0..n_pairs).map(|pair| (pair / n_used) as i32).collect();
5504 let pair_ex: Vec<i32> = sel_all.iter().map(|&expert| expert as i32).collect();
5505 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
5506 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
5507
5508 let mut by_expert: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
5509 for (pair, &expert) in pair_ex.iter().enumerate() {
5510 by_expert[expert as usize].push(pair as i32);
5511 }
5512
5513 let pair_tok_d = e.htod_i32(&pair_tok)?;
5514 let pair_ex_d = e.htod_i32(&pair_ex)?;
5515 let pair_w_d = e.htod(w_all)?;
5516 let tok_off_d = e.htod_i32(&tok_off)?;
5517 let tok_ids_d = e.htod_i32(&tok_ids)?;
5518
5519 let matvec = |
5520 proj: i32,
5521 pair_rows: &CudaSlice<i32>,
5522 aq: &CudaSlice<i8>,
5523 ad: &CudaSlice<f32>,
5524 in_f: usize,
5525 out_f: usize,
5526 qtype: i32,
5527 row_bytes: usize,
5528 | -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5529 e.moe_pairs_matvec_q8(
5530 table,
5531 proj,
5532 pair_rows,
5533 &pair_ex_d,
5534 aq,
5535 ad,
5536 in_f,
5537 out_f,
5538 n_expert,
5539 n_pairs,
5540 qtype,
5541 row_bytes,
5542 )
5543 };
5544
5545 let (gate_row_bytes, up_row_bytes) = if gu_il {
5546 let combined = m.gate_exps.row_bytes + m.up_exps.row_bytes;
5547 (combined, combined)
5548 } else {
5549 (m.gate_exps.row_bytes, m.up_exps.row_bytes)
5550 };
5551 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
5552 let gate = matvec(
5553 0,
5554 &pair_tok_d,
5555 &zq,
5556 &zd,
5557 n_embd,
5558 n_ff_exp,
5559 m.gate_exps.qtype,
5560 gate_row_bytes,
5561 )?;
5562 let up = matvec(
5563 1,
5564 &pair_tok_d,
5565 &zq,
5566 &zd,
5567 n_embd,
5568 n_ff_exp,
5569 m.up_exps.qtype,
5570 up_row_bytes,
5571 )?;
5572 let mut act = e.uninit(n_pairs * n_ff_exp)?;
5573 Self::ffn_act_lim(
5574 e,
5575 cfg,
5576 &gate,
5577 &up,
5578 1.0,
5579 1.0,
5580 cfg.clamp_exp_at(il as u32),
5581 &mut act,
5582 n_pairs * n_ff_exp,
5583 )?;
5584 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
5585 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
5586 let pair_self_d = e.htod_i32(&pair_self)?;
5587 let down = matvec(
5588 2,
5589 &pair_self_d,
5590 &aq2,
5591 &ad2,
5592 n_ff_exp,
5593 n_embd,
5594 m.down_exps.qtype,
5595 m.down_exps.row_bytes,
5596 )?;
5597 let mut moe_out = e.uninit(t * n_embd)?;
5598 e.moe_pairs_scatter(
5599 &down,
5600 &pair_w_d,
5601 &tok_off_d,
5602 &tok_ids_d,
5603 &mut moe_out,
5604 t,
5605 n_embd,
5606 )?;
5607
5608 if std::env::var("MEMRA_MOE_STATS").is_ok() {
5609 let mut sizes: Vec<usize> = by_expert
5610 .iter()
5611 .filter_map(|pairs| (!pairs.is_empty()).then_some(pairs.len()))
5612 .collect();
5613 sizes.sort_unstable();
5614 let mean = sizes.iter().sum::<usize>() as f64 / sizes.len().max(1) as f64;
5615 println!(
5616 "moe-grouped il={il} t={t} dispatch=resident-q8-clamped-pairs active={}/{} \
5617 m_e: min={} median={} mean={mean:.1} max={}",
5618 sizes.len(),
5619 n_expert,
5620 sizes.first().copied().unwrap_or(0),
5621 sizes.get(sizes.len() / 2).copied().unwrap_or(0),
5622 sizes.last().copied().unwrap_or(0),
5623 );
5624 }
5625 Ok(moe_out)
5626 }
5627
5628 fn moe_ffn_grouped_add_shared(
5629 e: &Engine,
5630 m: &MoeWeights,
5631 z: &CudaSlice<f32>,
5632 t: usize,
5633 cfg: &ModelConfig,
5634 il: u16,
5635 moe_out: &mut CudaSlice<f32>,
5636 ) -> Result<(), Box<dyn std::error::Error>> {
5637 let n_embd = cfg.n_embd as usize;
5638 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
5639 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
5640 {
5641 let n_ff_sh = gate_shexp.out_features();
5642 let sg_gate = e.matmul(gate_shexp, z, t)?;
5643 let sg_up = e.matmul(up_shexp, z, t)?;
5644 let mut sa = e.uninit(t * n_ff_sh)?;
5645 Self::ffn_act_lim(
5646 e,
5647 cfg,
5648 &sg_gate,
5649 &sg_up,
5650 1.0,
5651 1.0,
5652 cfg.clamp_shexp_at(il as u32),
5653 &mut sa,
5654 t * n_ff_sh,
5655 )?;
5656 let sh = e.matmul(down_shexp, &sa, t)?;
5657 let gate = match &m.gate_inp_shexp {
5658 Some(gate_inp_shexp) => {
5659 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
5660 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
5661 } else {
5662 let raw = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
5663 let mut gate = e.uninit(t)?;
5664 e.sigmoid(&raw, &mut gate, t)?;
5665 gate
5666 }
5667 }
5668 None => e.htod(&vec![1.0f32; t])?,
5669 };
5670 e.add_scaled_rows(&sh, &gate, moe_out, n_embd, t)?;
5671 }
5672 Ok(())
5673 }
5674
5675 pub(crate) fn moe_ffn_grouped(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, t: usize,
5678 cfg: &ModelConfig, il: u16, max_block: usize)
5679 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5680 let moe = cfg.moe.as_ref().unwrap();
5681 let n_embd = cfg.n_embd as usize;
5682 let n_expert = moe.expert_count as usize;
5683 let n_used = moe.expert_used_count as usize;
5684 let n_ff_exp = moe.expert_ff_length as usize;
5685 let lim_exp = cfg.clamp_exp_at(il as u32);
5687
5688 let logits = Self::moe_router_logits(e, m, z, t, cfg)?;
5691 let (sel_all, w_all) = if let Some(sig) = cfg.sigmoid_router() {
5692 Self::moe_route_cfg(e, &logits, t, n_expert, n_used,
5693 m.exp_probs_b.as_deref(), Some(sig), m.active_experts.as_deref())?
5694 } else {
5695 Self::moe_route_cfg(e, &logits, t, n_expert, n_used,
5696 None, None, m.active_experts.as_deref())?
5697 };
5698 Self::trace_moe_routes(il, t, &sel_all, &w_all)?;
5699 Self::trace_moe_input(e, il, t, n_embd, z)?;
5700
5701 let no_exp_macros = m.gate_exps.macros.is_none()
5706 && m.up_exps.macros.is_none()
5707 && m.down_exps.macros.is_none();
5708 let resident_q8 = m.dev_exps.as_ref().filter(|dev| {
5709 m.has_uniform_expert_layout()
5710 && no_exp_macros
5711 && moe_q8_enabled()
5712 && q8_expert_supported(m.gate_exps.qtype)
5713 && q8_expert_supported(m.up_exps.qtype)
5714 && q8_expert_supported(m.down_exps.qtype)
5715 && moe_slab_enabled()
5716 && dev.dev == e.ctx().ordinal()
5717 });
5718 if let Some(dev) = resident_q8 {
5719 let mut moe_out = Self::moe_ffn_grouped_resident_q8(
5720 e,
5721 m,
5722 z,
5723 t,
5724 cfg,
5725 il,
5726 &sel_all,
5727 &w_all,
5728 &dev.ptr_row,
5729 dev.gu_il,
5730 )?;
5731 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut moe_out)?;
5732 return Ok(moe_out);
5733 }
5734
5735 struct ExpertGroup {
5739 tok_indices: Vec<i32>, slot_indices: Vec<i32>, weights: Vec<f32>, }
5743 let mut groups: Vec<ExpertGroup> = (0..n_expert).map(|_| ExpertGroup {
5744 tok_indices: Vec::new(), slot_indices: Vec::new(), weights: Vec::new(),
5745 }).collect();
5746
5747 for tok in 0..t {
5748 for j in 0..n_used {
5749 let ex = sel_all[tok * n_used + j] as usize;
5750 let w = w_all[tok * n_used + j];
5751 groups[ex].tok_indices.push(tok as i32);
5752 groups[ex].slot_indices.push(j as i32);
5753 groups[ex].weights.push(w);
5754 }
5755 }
5756
5757 let mut slot_buf = e.zeros(t * n_used * n_embd)?;
5760 let mut wbuf = e.zeros(t * n_used)?; let g_len = m.gate_exps.max_expert_bytes();
5764 let u_len = m.up_exps.max_expert_bytes();
5765 let d_len = m.down_exps.max_expert_bytes();
5766 let moe_q8 = m.has_uniform_expert_layout()
5767 && moe_q8_enabled()
5768 && q8_expert_supported(m.gate_exps.qtype)
5769 && q8_expert_supported(m.up_exps.qtype)
5770 && q8_expert_supported(m.down_exps.qtype);
5771 let slab_local = m.dev_exps.as_ref().filter(|dev| {
5774 !dev.gu_il && moe_slab_enabled() && dev.dev == e.ctx().ordinal()
5775 });
5776 let use_cache =
5777 slab_local.is_none() && Engine::moe_cache_enabled() && !e.moe_cache_frozen();
5778 let grouped_q8 = moe_q8 && (slab_local.is_some() || use_cache);
5781
5782 let (mut scratch_g, mut scratch_u, mut scratch_d) = if slab_local.is_none() && !use_cache {
5784 (Some(e.alloc_u8(g_len)?), Some(e.alloc_u8(u_len)?), Some(e.alloc_u8(d_len)?))
5785 } else {
5786 (None, None, None)
5787 };
5788
5789 let mut order: Vec<usize> =
5800 (0..n_expert).filter(|&ex| !groups[ex].tok_indices.is_empty()).collect();
5801 order.sort_by(|&a, &b| groups[b].tok_indices.len()
5802 .cmp(&groups[a].tok_indices.len()).then(a.cmp(&b)));
5803 let mut m_dist: Vec<usize> = Vec::new(); let page_window = moe_page_prefetch_window();
5805 let worker_disk_prefetch = use_cache && crate::spill_pread::worker_enabled();
5806 if worker_disk_prefetch {
5807 if let Some(first) = grouped_worker_prefetch_position(order.len(), None) {
5808 Self::moe_prefetch_disk_expert(e, il, order[first], m, max_block, &[])?;
5809 }
5810 }
5811 for (order_pos, &ex) in order.iter().enumerate() {
5812 for next in page_prefetch_positions(order_pos, order.len(), page_window) {
5813 Self::moe_prefetch_host_expert(order[next], m);
5814 }
5815 if worker_disk_prefetch {
5816 if let Some(next) = grouped_worker_prefetch_position(order.len(), Some(order_pos)) {
5817 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
5818 let keep = [
5819 BlockId::new(il, PROJ_GATE, ex as u16),
5820 BlockId::new(il, PROJ_UP, ex as u16),
5821 BlockId::new(il, PROJ_DOWN, ex as u16),
5822 ];
5823 Self::moe_prefetch_disk_expert(e, il, order[next], m, max_block, &keep)?;
5824 }
5825 }
5826 let grp = &groups[ex];
5827 let m_e = grp.tok_indices.len();
5828 m_dist.push(m_e);
5829 let gl = m.gate_exps.expert_layout(ex);
5830 let ul = m.up_exps.expert_layout(ex);
5831 let dl = m.down_exps.expert_layout(ex);
5832
5833 let tok_idx_d = e.htod_i32(&grp.tok_indices)?;
5837 let slot_idx_d = e.htod_i32(&grp.slot_indices)?;
5838 let dmac = m.down_exps.macro_scale(ex);
5839 let weight_d = if dmac == 1.0 { e.htod(&grp.weights)? } else {
5840 let scaled: Vec<f32> = grp.weights.iter().map(|&w| w * dmac).collect();
5841 e.htod(&scaled)?
5842 };
5843
5844 let mut gathered = e.zeros(m_e * n_embd)?;
5846 e.gather_rows(z, &tok_idx_d, &mut gathered, n_embd, m_e)?;
5847 let gv = gathered.slice(0..m_e * n_embd);
5848
5849 let y = if let Some(dev) = slab_local {
5852 let gate_start = ex * m.gate_exps.expert_stride;
5853 let up_start = ex * m.up_exps.expert_stride;
5854 let down_start = ex * m.down_exps.expert_stride;
5855 if grouped_q8 {
5856 let (zq, zd) = e.quantize_q8_1(&gathered, m_e, n_embd)?;
5857 let gate = e.qmatvec_expert_q8(
5858 &dev.gate,
5859 gate_start..gate_start + gl.len,
5860 &zq,
5861 &zd,
5862 m_e,
5863 m.gate_exps.in_f,
5864 m.gate_exps.out_f,
5865 gl.qtype,
5866 gl.row_bytes,
5867 )?;
5868 let up = e.qmatvec_expert_q8(
5869 &dev.up,
5870 up_start..up_start + ul.len,
5871 &zq,
5872 &zd,
5873 m_e,
5874 m.up_exps.in_f,
5875 m.up_exps.out_f,
5876 ul.qtype,
5877 ul.row_bytes,
5878 )?;
5879 let mut act = e.uninit(m_e * n_ff_exp)?;
5880 Self::ffn_act_lim(
5881 e,
5882 cfg,
5883 &gate,
5884 &up,
5885 m.gate_exps.macro_scale(ex),
5886 m.up_exps.macro_scale(ex),
5887 lim_exp,
5888 &mut act,
5889 m_e * n_ff_exp,
5890 )?;
5891 let (aq2, ad2) = e.quantize_q8_1(&act, m_e, n_ff_exp)?;
5892 e.qmatvec_expert_q8(
5893 &dev.down,
5894 down_start..down_start + dl.len,
5895 &aq2,
5896 &ad2,
5897 m_e,
5898 m.down_exps.in_f,
5899 m.down_exps.out_f,
5900 dl.qtype,
5901 dl.row_bytes,
5902 )?
5903 } else {
5904 let gate = e.qmatvec_view(
5905 &dev.gate,
5906 gate_start..gate_start + gl.len,
5907 &gv,
5908 m_e,
5909 m.gate_exps.in_f,
5910 m.gate_exps.out_f,
5911 gl.qtype,
5912 gl.row_bytes,
5913 )?;
5914 let up = e.qmatvec_view(
5915 &dev.up,
5916 up_start..up_start + ul.len,
5917 &gv,
5918 m_e,
5919 m.up_exps.in_f,
5920 m.up_exps.out_f,
5921 ul.qtype,
5922 ul.row_bytes,
5923 )?;
5924 let mut act = e.uninit(m_e * n_ff_exp)?;
5925 Self::ffn_act_lim(
5926 e,
5927 cfg,
5928 &gate,
5929 &up,
5930 m.gate_exps.macro_scale(ex),
5931 m.up_exps.macro_scale(ex),
5932 lim_exp,
5933 &mut act,
5934 m_e * n_ff_exp,
5935 )?;
5936 let actv = act.slice(0..m_e * n_ff_exp);
5937 e.qmatvec_view(
5938 &dev.down,
5939 down_start..down_start + dl.len,
5940 &actv,
5941 m_e,
5942 m.down_exps.in_f,
5943 m.down_exps.out_f,
5944 dl.qtype,
5945 dl.row_bytes,
5946 )?
5947 }
5948 } else if use_cache {
5949 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
5950 if grouped_q8 {
5951 let (zq, zd) = e.quantize_q8_1(&gathered, m_e, n_embd)?;
5952 let gate = e.with_moe_cache(max_block, |cache, eng| {
5953 let id = BlockId::new(il, PROJ_GATE, ex as u16);
5954 let slot =
5955 cache.dispatch_source(id, m.gate_exps.expert_source(ex), eng)?;
5956 eng.qmatvec_expert_q8(
5957 cache.buf(slot),
5958 0..gl.len,
5959 &zq,
5960 &zd,
5961 m_e,
5962 m.gate_exps.in_f,
5963 m.gate_exps.out_f,
5964 gl.qtype,
5965 gl.row_bytes,
5966 )
5967 })?;
5968 let up = e.with_moe_cache(max_block, |cache, eng| {
5969 let id = BlockId::new(il, PROJ_UP, ex as u16);
5970 let slot =
5971 cache.dispatch_source(id, m.up_exps.expert_source(ex), eng)?;
5972 eng.qmatvec_expert_q8(
5973 cache.buf(slot),
5974 0..ul.len,
5975 &zq,
5976 &zd,
5977 m_e,
5978 m.up_exps.in_f,
5979 m.up_exps.out_f,
5980 ul.qtype,
5981 ul.row_bytes,
5982 )
5983 })?;
5984 let mut act = e.uninit(m_e * n_ff_exp)?;
5985 Self::ffn_act_lim(
5986 e,
5987 cfg,
5988 &gate,
5989 &up,
5990 m.gate_exps.macro_scale(ex),
5991 m.up_exps.macro_scale(ex),
5992 lim_exp,
5993 &mut act,
5994 m_e * n_ff_exp,
5995 )?;
5996 let (aq2, ad2) = e.quantize_q8_1(&act, m_e, n_ff_exp)?;
5997 e.with_moe_cache(max_block, |cache, eng| {
5998 let id = BlockId::new(il, PROJ_DOWN, ex as u16);
5999 let slot =
6000 cache.dispatch_source(id, m.down_exps.expert_source(ex), eng)?;
6001 eng.qmatvec_expert_q8(
6002 cache.buf(slot),
6003 0..dl.len,
6004 &aq2,
6005 &ad2,
6006 m_e,
6007 m.down_exps.in_f,
6008 m.down_exps.out_f,
6009 dl.qtype,
6010 dl.row_bytes,
6011 )
6012 })?
6013 } else {
6014 let gate = e.with_moe_cache(max_block, |cache, eng| {
6015 let id = BlockId::new(il, PROJ_GATE, ex as u16);
6016 let slot =
6017 cache.dispatch_source(id, m.gate_exps.expert_source(ex), eng)?;
6018 eng.qmatvec_view(
6019 cache.buf(slot),
6020 0..gl.len,
6021 &gv,
6022 m_e,
6023 m.gate_exps.in_f,
6024 m.gate_exps.out_f,
6025 gl.qtype,
6026 gl.row_bytes,
6027 )
6028 })?;
6029 let up = e.with_moe_cache(max_block, |cache, eng| {
6030 let id = BlockId::new(il, PROJ_UP, ex as u16);
6031 let slot =
6032 cache.dispatch_source(id, m.up_exps.expert_source(ex), eng)?;
6033 eng.qmatvec_view(
6034 cache.buf(slot),
6035 0..ul.len,
6036 &gv,
6037 m_e,
6038 m.up_exps.in_f,
6039 m.up_exps.out_f,
6040 ul.qtype,
6041 ul.row_bytes,
6042 )
6043 })?;
6044 let mut act = e.uninit(m_e * n_ff_exp)?;
6045 Self::ffn_act_lim(
6046 e,
6047 cfg,
6048 &gate,
6049 &up,
6050 m.gate_exps.macro_scale(ex),
6051 m.up_exps.macro_scale(ex),
6052 lim_exp,
6053 &mut act,
6054 m_e * n_ff_exp,
6055 )?;
6056 let actv = act.slice(0..m_e * n_ff_exp);
6057 e.with_moe_cache(max_block, |cache, eng| {
6058 let id = BlockId::new(il, PROJ_DOWN, ex as u16);
6059 let slot =
6060 cache.dispatch_source(id, m.down_exps.expert_source(ex), eng)?;
6061 eng.qmatvec_view(
6062 cache.buf(slot),
6063 0..dl.len,
6064 &actv,
6065 m_e,
6066 m.down_exps.in_f,
6067 m.down_exps.out_f,
6068 dl.qtype,
6069 dl.row_bytes,
6070 )
6071 })?
6072 }
6073 } else {
6074 let sg = scratch_g.as_mut().unwrap();
6075 let su = scratch_u.as_mut().unwrap();
6076 let sd = scratch_d.as_mut().unwrap();
6077 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
6078 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
6079 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
6080 if grouped_q8 {
6081 let (zq, zd) = e.quantize_q8_1(&gathered, m_e, n_embd)?;
6082 let gate = e.qmatvec_expert_q8(
6083 sg,
6084 0..gl.len,
6085 &zq,
6086 &zd,
6087 m_e,
6088 m.gate_exps.in_f,
6089 m.gate_exps.out_f,
6090 gl.qtype,
6091 gl.row_bytes,
6092 )?;
6093 let up = e.qmatvec_expert_q8(
6094 su,
6095 0..ul.len,
6096 &zq,
6097 &zd,
6098 m_e,
6099 m.up_exps.in_f,
6100 m.up_exps.out_f,
6101 ul.qtype,
6102 ul.row_bytes,
6103 )?;
6104 let mut act = e.uninit(m_e * n_ff_exp)?;
6105 Self::ffn_act_lim(
6106 e,
6107 cfg,
6108 &gate,
6109 &up,
6110 m.gate_exps.macro_scale(ex),
6111 m.up_exps.macro_scale(ex),
6112 lim_exp,
6113 &mut act,
6114 m_e * n_ff_exp,
6115 )?;
6116 let (aq2, ad2) = e.quantize_q8_1(&act, m_e, n_ff_exp)?;
6117 e.qmatvec_expert_q8(
6118 sd,
6119 0..dl.len,
6120 &aq2,
6121 &ad2,
6122 m_e,
6123 m.down_exps.in_f,
6124 m.down_exps.out_f,
6125 dl.qtype,
6126 dl.row_bytes,
6127 )?
6128 } else {
6129 let gate = e.qmatvec_view(
6130 sg,
6131 0..gl.len,
6132 &gv,
6133 m_e,
6134 m.gate_exps.in_f,
6135 m.gate_exps.out_f,
6136 gl.qtype,
6137 gl.row_bytes,
6138 )?;
6139 let up = e.qmatvec_view(
6140 su,
6141 0..ul.len,
6142 &gv,
6143 m_e,
6144 m.up_exps.in_f,
6145 m.up_exps.out_f,
6146 ul.qtype,
6147 ul.row_bytes,
6148 )?;
6149 let mut act = e.uninit(m_e * n_ff_exp)?;
6150 Self::ffn_act_lim(
6151 e,
6152 cfg,
6153 &gate,
6154 &up,
6155 m.gate_exps.macro_scale(ex),
6156 m.up_exps.macro_scale(ex),
6157 lim_exp,
6158 &mut act,
6159 m_e * n_ff_exp,
6160 )?;
6161 let actv = act.slice(0..m_e * n_ff_exp);
6162 e.qmatvec_view(
6163 sd,
6164 0..dl.len,
6165 &actv,
6166 m_e,
6167 m.down_exps.in_f,
6168 m.down_exps.out_f,
6169 dl.qtype,
6170 dl.row_bytes,
6171 )?
6172 }
6173 };
6174
6175 e.scatter_slot(&y, &tok_idx_d, &slot_idx_d, &weight_d,
6177 &mut slot_buf, &mut wbuf, n_embd, n_used, m_e)?;
6178 }
6179
6180 let mut moe_out = e.zeros(t * n_embd)?;
6182 e.reduce_slots(&slot_buf, &wbuf, &mut moe_out, n_embd, n_used, t)?;
6183
6184 if std::env::var("MEMRA_MOE_STATS").is_ok() && !m_dist.is_empty() {
6186 m_dist.sort_unstable();
6187 let active = m_dist.len();
6188 let mean = m_dist.iter().sum::<usize>() as f64 / active as f64;
6189 let median = m_dist[active / 2];
6190 let max_m = *m_dist.last().unwrap();
6191 let min_m = m_dist[0];
6192 let above16 = m_dist.iter().filter(|&&x| x >= 16).count();
6193 println!("moe-grouped il={il} t={t} active={active}/{n_expert} \
6194 m_e: min={min_m} median={median} mean={mean:.1} max={max_m} \
6195 above_gemm_threshold(>=16)={above16}/{active}");
6196 }
6197
6198 Self::moe_ffn_grouped_add_shared(e, m, z, t, cfg, il, &mut moe_out)?;
6199 Ok(moe_out)
6200 }
6201
6202 pub(crate) fn moe_ffn_lockstep(
6209 &self,
6210 e: &Engine,
6211 m: &MoeWeights,
6212 zbatch: &CudaSlice<f32>,
6213 mrows: usize,
6214 il: u16,
6215 max_block: usize,
6216 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6217 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
6218 let cfg = &self.cfg;
6219 let moe = cfg.moe.as_ref().unwrap();
6220 let n_embd = cfg.n_embd as usize;
6221 let n_expert = moe.expert_count as usize;
6222 let n_used = moe.expert_used_count as usize;
6223 let n_ff_exp = moe.expert_ff_length as usize;
6224 let lim_exp = cfg.clamp_exp_at(il as u32);
6226 let lim_shexp = cfg.clamp_shexp_at(il as u32);
6227
6228 let logits = e.matmul(&m.gate_inp, zbatch, mrows)?;
6229 let (sel_all, w_all) = if let Some(sig) = cfg.sigmoid_router() {
6230 Self::moe_route_cfg(e, &logits, mrows, n_expert, n_used,
6231 m.exp_probs_b.as_deref(), Some(sig), m.active_experts.as_deref())?
6232 } else {
6233 Self::moe_route_cfg(e, &logits, mrows, n_expert, n_used,
6234 None, None, m.active_experts.as_deref())?
6235 };
6236 Self::trace_moe_routes(il, mrows, &sel_all, &w_all)?;
6237
6238 let resident_expert: Vec<bool> = e.with_moe_cache(max_block, |c, _| {
6240 Ok((0..n_expert)
6241 .map(|ex| {
6242 [PROJ_GATE, PROJ_UP, PROJ_DOWN].into_iter().all(|p| {
6243 c.resident(BlockId::new(il, p, ex as u16)).is_some()
6244 })
6245 })
6246 .collect())
6247 })?;
6248
6249 struct Group {
6250 rows: Vec<i32>,
6251 slots: Vec<i32>,
6252 weights: Vec<f32>,
6253 }
6254 let mut groups: std::collections::HashMap<usize, Group> = Default::default();
6255 let mut cpu_rows: Vec<Vec<(usize, f32)>> = vec![Vec::new(); mrows];
6256 let mut cpu_by_expert: std::collections::HashMap<usize, Vec<(usize, f32)>> =
6257 Default::default();
6258 for row in 0..mrows {
6259 for j in 0..n_used {
6260 let ex = sel_all[row * n_used + j] as usize;
6261 let w = w_all[row * n_used + j];
6262 if resident_expert[ex] {
6263 let group = groups.entry(ex).or_insert_with(|| Group {
6264 rows: Vec::new(),
6265 slots: Vec::new(),
6266 weights: Vec::new(),
6267 });
6268 group.rows.push(row as i32);
6269 group.slots.push(j as i32);
6270 group.weights.push(w);
6271 } else {
6272 crate::cpu_experts::record_incomplete_gpu_residency(0);
6273 cpu_rows[row].push((ex, w));
6274 cpu_by_expert.entry(ex).or_default().push((row, w));
6275 }
6276 }
6277 }
6278
6279 let host_rows = e.dtoh(zbatch)?;
6285 let rows_ok = crate::cpu_experts::rows_supported();
6286 enum CpuPart {
6287 Single { row: usize },
6288 Rows { rows: Vec<usize> },
6289 }
6290 let mut tickets: Vec<(CpuPart, crate::cpu_experts::CpuExpertTicket)> = Vec::new();
6291 let mut rows_served: std::collections::HashSet<(usize, usize)> = Default::default();
6292 if rows_ok {
6293 let mut shared: Vec<(usize, Vec<(usize, f32)>)> = cpu_by_expert
6294 .into_iter()
6295 .filter(|(_, rows)| rows.len() >= 2)
6296 .collect();
6297 shared.sort_by_key(|(ex, _)| *ex);
6298 for (ex, mut row_weights) in shared {
6299 row_weights.sort_by_key(|(row, _)| *row);
6300 let inputs: Vec<(&[f32], f32)> = row_weights
6301 .iter()
6302 .map(|&(row, w)| (&host_rows[row * n_embd..(row + 1) * n_embd], w))
6303 .collect();
6304 let job = crate::cpu_experts::prepare_rows_job(m, ex, &inputs)
6305 .map_err(std::io::Error::other)?;
6306 for &(row, _) in &row_weights {
6307 rows_served.insert((row, ex));
6308 }
6309 tickets.push((
6310 CpuPart::Rows {
6311 rows: row_weights.iter().map(|&(row, _)| row).collect(),
6312 },
6313 crate::cpu_experts::submit_rows(job).map_err(std::io::Error::other)?,
6314 ));
6315 }
6316 }
6317 for (row, selected) in cpu_rows.iter().enumerate() {
6318 let leftover: Vec<(usize, f32)> = selected
6319 .iter()
6320 .copied()
6321 .filter(|&(ex, _)| !rows_served.contains(&(row, ex)))
6322 .collect();
6323 if leftover.is_empty() {
6324 continue;
6325 }
6326 let host_row = &host_rows[row * n_embd..(row + 1) * n_embd];
6327 let job = crate::cpu_experts::prepare_job(m, il, &leftover, host_row)
6328 .map_err(std::io::Error::other)?;
6329 tickets.push((
6330 CpuPart::Single { row },
6331 crate::cpu_experts::submit(job).map_err(std::io::Error::other)?,
6332 ));
6333 }
6334
6335 let mut slot_buf = e.zeros(mrows * n_used * n_embd)?;
6336 let mut wbuf = e.zeros(mrows * n_used)?;
6337 let mut order: Vec<usize> = groups.keys().copied().collect();
6338 order.sort_by(|&a, &b| {
6339 groups[&b].rows.len().cmp(&groups[&a].rows.len()).then(a.cmp(&b))
6340 });
6341 for &ex in &order {
6342 let group = &groups[&ex];
6343 let m_e = group.rows.len();
6344 let gl = m.gate_exps.expert_layout(ex);
6345 let ul = m.up_exps.expert_layout(ex);
6346 let dl = m.down_exps.expert_layout(ex);
6347 let row_idx_d = e.htod_i32(&group.rows)?;
6348 let slot_idx_d = e.htod_i32(&group.slots)?;
6349 let dmac = m.down_exps.macro_scale(ex);
6350 let weight_d = if dmac == 1.0 {
6351 e.htod(&group.weights)?
6352 } else {
6353 let scaled: Vec<f32> = group.weights.iter().map(|&w| w * dmac).collect();
6354 e.htod(&scaled)?
6355 };
6356 let mut gathered = e.zeros(m_e * n_embd)?;
6357 e.gather_rows(zbatch, &row_idx_d, &mut gathered, n_embd, m_e)?;
6358 let gv = gathered.slice(0..m_e * n_embd);
6359 let gate = e.with_moe_cache(max_block, |c, eng| {
6360 let slot = c
6361 .resident(BlockId::new(il, PROJ_GATE, ex as u16))
6362 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
6363 eng.qmatvec_view(c.buf(crate::moe_cache::DispatchSlot::Resident(slot)), 0..gl.len, &gv, m_e,
6364 m.gate_exps.in_f, m.gate_exps.out_f, gl.qtype, gl.row_bytes)
6365 })?;
6366 let up = e.with_moe_cache(max_block, |c, eng| {
6367 let slot = c
6368 .resident(BlockId::new(il, PROJ_UP, ex as u16))
6369 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
6370 eng.qmatvec_view(c.buf(crate::moe_cache::DispatchSlot::Resident(slot)), 0..ul.len, &gv, m_e,
6371 m.up_exps.in_f, m.up_exps.out_f, ul.qtype, ul.row_bytes)
6372 })?;
6373 let mut act = e.zeros(m_e * n_ff_exp)?;
6374 Self::ffn_act_lim(e, cfg, &gate, &up, m.gate_exps.macro_scale(ex),
6375 m.up_exps.macro_scale(ex), lim_exp, &mut act, m_e * n_ff_exp)?;
6376 let actv = act.slice(0..m_e * n_ff_exp);
6377 let y = e.with_moe_cache(max_block, |c, eng| {
6378 let slot = c
6379 .resident(BlockId::new(il, PROJ_DOWN, ex as u16))
6380 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
6381 eng.qmatvec_view(c.buf(crate::moe_cache::DispatchSlot::Resident(slot)), 0..dl.len, &actv, m_e,
6382 m.down_exps.in_f, m.down_exps.out_f, dl.qtype, dl.row_bytes)
6383 })?;
6384 e.scatter_slot(&y, &row_idx_d, &slot_idx_d, &weight_d,
6385 &mut slot_buf, &mut wbuf, n_embd, n_used, m_e)?;
6386 }
6387 let mut moe_out = e.zeros(mrows * n_embd)?;
6388 e.reduce_slots(&slot_buf, &wbuf, &mut moe_out, n_embd, n_used, mrows)?;
6389
6390 let mut row_sums: Vec<Option<Vec<f32>>> = vec![None; mrows];
6392 for (part, ticket) in tickets {
6393 let cpu_output = ticket.wait().map_err(std::io::Error::other)?;
6394 let mut add_row = |row: usize, chunk: &[f32]| {
6395 let sum = row_sums[row].get_or_insert_with(|| vec![0.0f32; n_embd]);
6396 for (accumulator, value) in sum.iter_mut().zip(chunk) {
6397 *accumulator += value;
6398 }
6399 };
6400 match part {
6401 CpuPart::Single { row } => add_row(row, &cpu_output),
6402 CpuPart::Rows { rows } => {
6403 for (slot, row) in rows.into_iter().enumerate() {
6404 add_row(row, &cpu_output[slot * n_embd..(slot + 1) * n_embd]);
6405 }
6406 }
6407 }
6408 }
6409 for (row, sum) in row_sums.into_iter().enumerate() {
6410 let Some(sum) = sum else { continue };
6411 let cpu_output = e.htod(&sum)?;
6412 let mut dst = moe_out.slice_mut(row * n_embd..(row + 1) * n_embd);
6413 e.axpy_into(&cpu_output, 1.0, &mut dst, n_embd)?;
6414 }
6415
6416 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
6417 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
6418 {
6419 let n_ff_sh = gate_shexp.out_features();
6420 let sg_gate = e.matmul(gate_shexp, zbatch, mrows)?;
6421 let sg_up = e.matmul(up_shexp, zbatch, mrows)?;
6422 let mut sa = e.zeros(mrows * n_ff_sh)?;
6423 Self::ffn_act_lim(e, cfg, &sg_gate, &sg_up, 1.0, 1.0, lim_shexp,
6424 &mut sa, mrows * n_ff_sh)?;
6425 let sh = e.matmul(down_shexp, &sa, mrows)?;
6426 let g = match &m.gate_inp_shexp {
6429 Some(gate_inp_shexp) => {
6430 e.sigmoid_dot_rows(zbatch, gate_inp_shexp.float_data(), n_embd, mrows)?
6431 }
6432 None => e.htod(&vec![1.0f32; mrows])?,
6433 };
6434 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, mrows)?;
6435 }
6436
6437 Ok(moe_out)
6438 }
6439}
6440
6441impl HybridModel {
6447 pub(crate) fn gemma4_geom(&self, il: usize) -> (usize, usize, usize, f32, f32, bool) {
6449 let g = self.cfg.gemma4.as_ref().unwrap();
6450 let swa = g.swa_pattern[il];
6451 let hd = if swa { g.key_length_swa } else { g.key_length_global } as usize;
6452 (hd, g.head_count_kv[il] as usize, self.cfg.n_head as usize,
6456 if swa { g.rope_base_swa } else { g.rope_base_global },
6457 1.0, swa)
6458 }
6459
6460 fn gemma4_suppress(&self, e: &Engine, ld: &mut CudaSlice<f32>, t: usize)
6464 -> Result<(), Box<dyn std::error::Error>> {
6465 if let Some((ids, n)) = self.gemma4_aux.as_ref().and_then(|a| a.suppress_d.as_ref()) {
6466 e.mask_ids_rows(ld, ids, *n, self.output.out_features(), t)?;
6467 }
6468 Ok(())
6469 }
6470
6471 fn gemma4_attn_prime(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
6476 h: &CudaSlice<f32>, pos_d: &CudaSlice<i32>, t: usize,
6477 cache: Option<&mut Cache>)
6478 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6479 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
6480 let eps = self.cfg.rms_eps;
6481 let aux = self.gemma4_aux.as_ref().unwrap();
6482
6483 e.mmq_act_begin();
6486 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)? };
6491
6492 let mut q = e.uninit(t * nh * hd)?;
6493 let mut k = e.uninit(t * nkv * hd)?;
6494 let mut v = e.uninit(t * nkv * hd)?;
6496 static EMIT: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6500 let emit = t >= 16 && crate::Engine::qkvnorm_w_on_prefill(nh * t + 2 * nkv * t, hd)
6501 && *EMIT.get_or_init(|| std::env::var("MEMRA_FA_EMIT").map(|s| s != "0").unwrap_or(true));
6502 let mut qb = e.alloc_uninit::<u8>(if emit { t * nh * hd * 2 } else { 1 })?;
6503 let mut kb = e.alloc_uninit::<u8>(if emit { t * nkv * hd * 2 } else { 1 })?;
6504 let mut vb = e.alloc_uninit::<u8>(if emit { t * nkv * hd * 2 } else { 1 })?;
6505 let v_f16 = emit && crate::fa_f16pv_on() && match hd {
6508 512 => true,
6509 256 => swa && crate::faw_hp_on() && nh % 2 == 0 && (nh / nkv) % 2 == 0,
6510 _ => false,
6511 };
6512 if emit {
6513 e.rms_norm_qkv_w4b(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
6514 &aux.ones, &mut q, &mut k, &mut v, &mut vb,
6515 hd, nh * t, nkv * t, eps, v_f16)?;
6516 } else {
6517 e.rms_norm_qkv(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
6518 &aux.ones, &mut q, &mut k, &mut v, hd, nh * t, nkv * t, eps)?;
6519 }
6520
6521 let ff = if swa { None } else {
6522 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
6523 };
6524 if emit {
6525 e.rope_neox2_bf16e(&mut q, &mut k, &mut qb, &mut kb, pos_d, hd, hd, nh, nkv, t,
6526 base, 1.0, ff)?;
6527 } else {
6528 e.rope_neox2(&mut q, &mut k, pos_d, hd, hd, nh, nkv, t, base, 1.0, ff)?;
6529 }
6530
6531 if let Some(cache) = cache {
6532 let kvl = cache.kv[il].as_mut().unwrap();
6533 assert_eq!(kvl.len, 0, "gemma4 prime is fresh-prompt only (v0)");
6534 e.append_kv_quantized_rows(&k, &v, &mut kvl.k, &mut kvl.v, kvl.len, t,
6535 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()))?;
6536 kvl.len += t;
6537 }
6538 let mut attn = e.zeros(t * nh * hd)?;
6539 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
6543 if swa && t > win {
6544 if hd == 256 && std::env::var("MEMRA_NOFA").is_err() {
6545 if emit { e.fa_prefill_w_pre(&qb, &kb, &vb, &mut attn, hd, nh, nkv, t, t,
6546 scale, true, win, v_f16)?; }
6547 else { e.fa_prefill_w(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true,
6548 win)?; }
6549 } else {
6550 e.sdpa_naive_w(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
6551 }
6552 } else if hd == 256 && std::env::var("MEMRA_NOFA").is_err() {
6553 e.fa_prefill(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
6554 } else if hd == 512 && std::env::var("MEMRA_NOFA").is_err() {
6555 if emit { e.fa_prefill_hd512_pre(&qb, &kb, &vb, &mut attn, hd, nh, nkv, t, t,
6556 scale, true, v_f16)?; }
6557 else { e.fa_prefill_hd512(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?; }
6558 } else {
6559 e.sdpa_naive(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
6560 }
6561 Ok(e.matmul(&fa.wo, &attn, t)?)
6562 }
6563
6564 fn gemma4_attn(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
6566 h: &CudaSlice<f32>, pos_d: &CudaSlice<i32>, t: usize)
6567 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6568 self.gemma4_attn_prime(e, fa, il, h, pos_d, t, None)
6569 }
6570
6571 fn gemma4_moe_q8(&self, e: &Engine, m: &crate::hybrid::MoeWeights,
6576 bits: &crate::hybrid::Gemma4MoeBits,
6577 mq: &(CudaSlice<i8>, CudaSlice<f32>),
6578 router_in: &CudaSlice<f32>, t: usize)
6579 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6580 let cfg = &self.cfg;
6581 let moe = cfg.moe.as_ref().unwrap();
6582 let n_embd = cfg.n_embd as usize;
6583 let n_expert = moe.expert_count as usize;
6584 let n_used = moe.expert_used_count as usize;
6585 let n_ff_exp = moe.expert_ff_length as usize;
6586 let logits = if crate::router_kernel_on() {
6590 e.router_gemv(m.gate_inp.float_data(), router_in, n_embd, n_expert, t)?
6591 } else {
6592 e.matmul(&m.gate_inp, router_in, t)?
6593 };
6594 let dev = m.dev_exps.as_ref().unwrap();
6595 let (sel_d, w_d) = e.moe_router_topk_scaled(&logits, t, n_expert, n_used,
6596 &bits.per_expert_scale_d)?;
6597 let (zq, zd) = mq;
6598 if t == 1 {
6599 let selv = sel_d.slice(0..n_used);
6600 let wv = w_d.slice(0..n_used);
6601 let act = e.moe_gate_up_gelu8_dev_q8(&dev.ptr_row, &selv, zq, zd,
6602 n_embd, n_ff_exp, n_used, n_expert,
6603 m.gate_exps.qtype, m.up_exps.qtype,
6604 m.gate_exps.row_bytes, m.up_exps.row_bytes)?;
6605 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
6606 let mut moe_out = e.uninit(n_embd)?;
6607 e.moe_down8_fma_dev_q8(&dev.ptr_row, &selv, &wv, &aq2, &ad2,
6608 &mut moe_out.slice_mut(0..n_embd), n_ff_exp, n_embd,
6609 n_used, n_expert, m.down_exps.qtype, m.down_exps.row_bytes)?;
6610 return Ok(moe_out);
6611 }
6612 let csr = t <= 10 && std::env::var("MEMRA_GEMMA_CSR").as_deref() != Ok("0");
6613 let act = if csr {
6614 e.moe_gate_up_gelu8_dev_q8_csr(&dev.ptr_row, &sel_d, zq, zd, t * n_used,
6615 n_embd, n_ff_exp, n_used, n_expert,
6616 m.gate_exps.qtype, m.up_exps.qtype,
6617 m.gate_exps.row_bytes, m.up_exps.row_bytes)?
6618 } else {
6619 e.moe_gate_up_gelu8_dev_q8_rows(&dev.ptr_row, &sel_d, zq, zd, t,
6620 n_embd, n_ff_exp, n_used, n_expert,
6621 m.gate_exps.qtype, m.up_exps.qtype,
6622 m.gate_exps.row_bytes, m.up_exps.row_bytes)?
6623 };
6624 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
6625 let mut moe_out = e.uninit(t * n_embd)?;
6626 e.moe_down8_fma_dev_q8_rows_g(&dev.ptr_row, &sel_d, &w_d, &aq2, &ad2, &mut moe_out, t,
6629 n_ff_exp, n_embd, n_used, n_expert,
6630 m.down_exps.qtype, m.down_exps.row_bytes)?;
6631 Ok(moe_out)
6632 }
6633
6634 fn gemma4_moe(&self, e: &Engine, m: &crate::hybrid::MoeWeights,
6638 bits: &crate::hybrid::Gemma4MoeBits, moe_in: &CudaSlice<f32>,
6639 router_in: &CudaSlice<f32>, t: usize)
6640 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6641 let cfg = &self.cfg;
6642 let moe = cfg.moe.as_ref().unwrap();
6643 let n_embd = cfg.n_embd as usize;
6644 let n_expert = moe.expert_count as usize;
6645 let n_used = moe.expert_used_count as usize;
6646 let n_ff_exp = moe.expert_ff_length as usize;
6647
6648 let logits = if t < PRIME_MIN_T && crate::router_kernel_on() {
6652 e.router_gemv(m.gate_inp.float_data(), router_in, n_embd, n_expert, t)?
6653 } else {
6654 e.matmul(&m.gate_inp, router_in, t)?
6655 };
6656
6657 if t < PRIME_MIN_T && m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
6662 && expert_dp4a_supported(m.gate_exps.qtype) && expert_dp4a_supported(m.up_exps.qtype)
6663 && expert_dp4a_supported(m.down_exps.qtype)
6664 && std::env::var("MEMRA_GEMMA_MOE_FAST").as_deref() != Ok("0") {
6665 let dev = m.dev_exps.as_ref().unwrap();
6666 let (sel_d, w_d) = e.moe_router_topk_scaled(&logits, t, n_expert, n_used,
6667 &bits.per_expert_scale_d)?;
6668 if t == 1 {
6669 let (zq, zd) = e.quantize_q8_1(moe_in, 1, n_embd)?;
6670 let selv = sel_d.slice(0..n_used);
6671 let wv = w_d.slice(0..n_used);
6672 let act = e.moe_gate_up_gelu8_dev_q8(&dev.ptr_row, &selv, &zq, &zd,
6673 n_embd, n_ff_exp, n_used, n_expert,
6674 m.gate_exps.qtype, m.up_exps.qtype,
6675 m.gate_exps.row_bytes, m.up_exps.row_bytes)?;
6676 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
6677 let mut moe_out = e.uninit(n_embd)?;
6678 e.moe_down8_fma_dev_q8(&dev.ptr_row, &selv, &wv, &aq2, &ad2,
6679 &mut moe_out.slice_mut(0..n_embd), n_ff_exp, n_embd,
6680 n_used, n_expert, m.down_exps.qtype, m.down_exps.row_bytes)?;
6681 return Ok(moe_out);
6682 }
6683 let (zq, zd) = e.quantize_q8_1(moe_in, t, n_embd)?;
6688 let csr = t <= 10 && std::env::var("MEMRA_GEMMA_CSR").as_deref() != Ok("0");
6689 let act = if csr {
6690 e.moe_gate_up_gelu8_dev_q8_csr(&dev.ptr_row, &sel_d, &zq, &zd, t * n_used,
6691 n_embd, n_ff_exp, n_used, n_expert,
6692 m.gate_exps.qtype, m.up_exps.qtype,
6693 m.gate_exps.row_bytes, m.up_exps.row_bytes)?
6694 } else {
6695 e.moe_gate_up_gelu8_dev_q8_rows(&dev.ptr_row, &sel_d, &zq, &zd, t,
6696 n_embd, n_ff_exp, n_used, n_expert,
6697 m.gate_exps.qtype, m.up_exps.qtype,
6698 m.gate_exps.row_bytes, m.up_exps.row_bytes)?
6699 };
6700 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
6701 let mut moe_out = e.uninit(t * n_embd)?;
6702 e.moe_down8_fma_dev_q8_rows_g(&dev.ptr_row, &sel_d, &w_d, &aq2, &ad2, &mut moe_out, t,
6703 n_ff_exp, n_embd, n_used, n_expert,
6704 m.down_exps.qtype, m.down_exps.row_bytes)?;
6705 return Ok(moe_out);
6706 }
6707
6708 let (sel_all, mut w_all) = Self::moe_route(e, &logits, t, n_expert, n_used)?;
6709 for (i, &sx) in sel_all.iter().enumerate() {
6710 w_all[i] *= bits.per_expert_scale[sx as usize];
6711 }
6712
6713 if t >= PRIME_MIN_T && m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
6717 && expert_dp4a_supported(m.gate_exps.qtype) && expert_dp4a_supported(m.up_exps.qtype)
6718 && expert_dp4a_supported(m.down_exps.qtype)
6719 && std::env::var("MEMRA_GEMMA_MOE_PAIRS").as_deref() != Ok("0") {
6720 let dev = m.dev_exps.as_ref().unwrap();
6721 let n_pairs = t * n_used;
6722 let pair_ex: Vec<i32> = sel_all.iter().map(|&x| x as i32).collect();
6723 let pair_tok: Vec<i32> = (0..n_pairs).map(|p| (p / n_used) as i32).collect();
6724 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
6725 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
6726 let pt = e.htod_i32(&pair_tok)?;
6727 let pw = e.htod(&w_all)?;
6728 let toff = e.htod_i32(&tok_off)?;
6729 let tids = e.htod_i32(&tok_ids)?;
6730 let mut by_ex: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
6731 for p in 0..n_pairs { by_ex[pair_ex[p] as usize].push(p as i32); }
6732 let mut ex_ids: Vec<i32> = Vec::new();
6733 let mut ex_off: Vec<i32> = vec![0];
6734 let mut ex_pairs: Vec<i32> = Vec::with_capacity(n_pairs);
6735 for (ex, list) in by_ex.iter().enumerate() {
6736 if list.is_empty() { continue; }
6737 ex_ids.push(ex as i32);
6738 ex_pairs.extend_from_slice(list);
6739 ex_off.push(ex_pairs.len() as i32);
6740 }
6741 let n_active = ex_ids.len();
6742 let exi = e.htod_i32(&ex_ids)?;
6743 let exo = e.htod_i32(&ex_off)?;
6744 let exp_d = e.htod_i32(&ex_pairs)?;
6745 if crate::moe_f16g_gemma_on()
6753 && f16g_proj_ok(m.gate_exps.qtype, n_embd)
6754 && f16g_proj_ok(m.up_exps.qtype, n_embd)
6755 && f16g_proj_ok(m.down_exps.qtype, n_ff_exp) {
6756 let csr_tok: Vec<i32> = ex_pairs.iter().map(|&p| p / n_used as i32).collect();
6757 let csr_tok_d = e.htod_i32(&csr_tok)?;
6758 let (z_f16, z_s) = e.moe_f16g_act(moe_in, Some(&csr_tok_d), n_embd, n_pairs)?;
6759 let g_csr = e.moe_f16_grouped(&dev.ptr_row, 0, n_expert, &exi, &ex_off, &exo,
6760 &z_f16, &z_s, n_embd, n_ff_exp, n_active, n_pairs,
6761 m.gate_exps.qtype, m.gate_exps.row_bytes)?;
6762 let u_csr = e.moe_f16_grouped(&dev.ptr_row, 1, n_expert, &exi, &ex_off, &exo,
6763 &z_f16, &z_s, n_embd, n_ff_exp, n_active, n_pairs,
6764 m.up_exps.qtype, m.up_exps.row_bytes)?;
6765 let act_csr = e.moe_pairs_gelu_mul(&g_csr, &u_csr, n_pairs * n_ff_exp)?;
6766 let (a_f16, a_s) = e.moe_f16g_act(&act_csr, None, n_ff_exp, n_pairs)?;
6767 let d_csr = e.moe_f16_grouped(&dev.ptr_row, 2, n_expert, &exi, &ex_off, &exo,
6768 &a_f16, &a_s, n_ff_exp, n_embd, n_active, n_pairs,
6769 m.down_exps.qtype, m.down_exps.row_bytes)?;
6770 let y_down = e.rows_permute(&d_csr, &exp_d, n_pairs, n_embd)?;
6771 let mut moe_out = e.uninit(t * n_embd)?;
6772 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
6773 if std::env::var("MEMRA_F16G_DEBUG").is_ok() {
6774 let scan = |v: &[f32]| v.iter().filter(|x| !x.is_finite()).count();
6775 let (yd, mo) = (e.dtoh(&y_down)?, e.dtoh(&moe_out)?);
6776 eprintln!("[f16g-debug] post-permute bad={} post-scatter bad={}",
6777 scan(&yd), scan(&mo));
6778 }
6779 return Ok(moe_out);
6780 }
6781 let mma = n_embd % 256 == 0
6784 && std::env::var("MEMRA_GEMMA_MOE_MMA").as_deref() != Ok("0");
6785 let (gate, up) = if mma {
6786 let z_scr = e.mmq_iq_quantize_act(moe_in, n_embd, t)?;
6787 (e.mmq_iq_experts(&dev.ptr_row, 0, n_expert, &exi, &exo, &exp_d, &pt, &z_scr,
6788 n_embd, n_ff_exp, n_active, n_pairs, t,
6789 m.gate_exps.qtype, m.gate_exps.row_bytes)?,
6790 e.mmq_iq_experts(&dev.ptr_row, 1, n_expert, &exi, &exo, &exp_d, &pt, &z_scr,
6791 n_embd, n_ff_exp, n_active, n_pairs, t,
6792 m.up_exps.qtype, m.up_exps.row_bytes)?)
6793 } else {
6794 let (zq, zd) = e.quantize_q8_1(moe_in, t, n_embd)?;
6795 (e.moe_pairs_matvec_q8_dec(&dev.ptr_row, 0, &exi, &exo, &exp_d, &pt, &zq, &zd,
6796 n_embd, n_ff_exp, n_expert, n_active, n_pairs,
6797 m.gate_exps.qtype, m.gate_exps.row_bytes)?,
6798 e.moe_pairs_matvec_q8_dec(&dev.ptr_row, 1, &exi, &exo, &exp_d, &pt, &zq, &zd,
6799 n_embd, n_ff_exp, n_expert, n_active, n_pairs,
6800 m.up_exps.qtype, m.up_exps.row_bytes)?)
6801 };
6802 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
6803 let pself = e.htod_i32(&pair_self)?;
6804 let y_down = if mma {
6816 let in_pad = n_ff_exp.div_ceil(256) * 256;
6817 let a_scr = if crate::moe_fuse_actq_on() {
6818 e.mmq_iq_fused_act_quant(&gate, &up, n_ff_exp, n_pairs, 1)?
6819 } else {
6820 let act = e.moe_pairs_gelu_mul(&gate, &up, n_pairs * n_ff_exp)?;
6821 e.mmq_iq_quantize_act(&act, n_ff_exp, n_pairs)?
6822 };
6823 e.mmq_iq_experts(&dev.ptr_row, 2, n_expert, &exi, &exo, &exp_d, &pself, &a_scr,
6824 in_pad, n_embd, n_active, n_pairs, n_pairs,
6825 m.down_exps.qtype, m.down_exps.row_bytes)?
6826 } else {
6827 let act = e.moe_pairs_gelu_mul(&gate, &up, n_pairs * n_ff_exp)?;
6828 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
6829 e.moe_pairs_matvec_q8_dec(&dev.ptr_row, 2, &exi, &exo, &exp_d, &pself, &aq2, &ad2,
6830 n_ff_exp, n_embd, n_expert, n_active, n_pairs,
6831 m.down_exps.qtype, m.down_exps.row_bytes)?
6832 };
6833 let mut moe_out = e.uninit(t * n_embd)?;
6834 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
6835 return Ok(moe_out);
6836 }
6837
6838 let g_len = m.gate_exps.expert_stride;
6839 let u_len = m.up_exps.expert_stride;
6840 let d_len = m.down_exps.expert_stride;
6841 let dev = m.dev_exps.as_ref().filter(|d| !d.gu_il);
6845 let (mut sg, mut su, mut sd) = if dev.is_some() { (None, None, None) } else {
6846 (Some(e.alloc_u8_uninit(g_len)?), Some(e.alloc_u8_uninit(u_len)?), Some(e.alloc_u8_uninit(d_len)?))
6847 };
6848 let mut moe_out = e.zeros(t * n_embd)?;
6849 for tok in 0..t {
6850 let sel = &sel_all[tok * n_used..(tok + 1) * n_used];
6851 let w = &w_all[tok * n_used..(tok + 1) * n_used];
6852 let zt = moe_in.slice(tok * n_embd..(tok + 1) * n_embd);
6853 for (j, &ex) in sel.iter().enumerate() {
6854 let ex = ex as usize;
6855 let gate = match dev {
6856 Some(d) => e.qmatvec_view(&d.gate, ex * g_len..(ex + 1) * g_len, &zt, 1,
6857 m.gate_exps.in_f, m.gate_exps.out_f, m.gate_exps.qtype, m.gate_exps.row_bytes)?,
6858 None => {
6859 let sg = sg.as_mut().unwrap();
6860 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
6861 e.qmatvec_view(sg, 0..g_len, &zt, 1,
6862 m.gate_exps.in_f, m.gate_exps.out_f, m.gate_exps.qtype, m.gate_exps.row_bytes)?
6863 }
6864 };
6865 let up = match dev {
6866 Some(d) => e.qmatvec_view(&d.up, ex * u_len..(ex + 1) * u_len, &zt, 1,
6867 m.up_exps.in_f, m.up_exps.out_f, m.up_exps.qtype, m.up_exps.row_bytes)?,
6868 None => {
6869 let su = su.as_mut().unwrap();
6870 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
6871 e.qmatvec_view(su, 0..u_len, &zt, 1,
6872 m.up_exps.in_f, m.up_exps.out_f, m.up_exps.qtype, m.up_exps.row_bytes)?
6873 }
6874 };
6875 let mut act = e.uninit(n_ff_exp)?;
6876 e.gelu_tanh_mul(&gate, &up, &mut act, n_ff_exp)?;
6877 let actv = act.slice(0..n_ff_exp);
6878 let y = match dev {
6879 Some(d) => e.qmatvec_view(&d.down, ex * d_len..(ex + 1) * d_len, &actv, 1,
6880 m.down_exps.in_f, m.down_exps.out_f, m.down_exps.qtype, m.down_exps.row_bytes)?,
6881 None => {
6882 let sd = sd.as_mut().unwrap();
6883 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
6884 e.qmatvec_view(sd, 0..d_len, &actv, 1,
6885 m.down_exps.in_f, m.down_exps.out_f, m.down_exps.qtype, m.down_exps.row_bytes)?
6886 }
6887 };
6888 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
6889 e.axpy_into(&y, w[j], &mut dst, n_embd)?;
6890 }
6891 }
6892 Ok(moe_out)
6893 }
6894
6895 fn gemma4_layer(&self, e: &Engine, il: usize, layer: &crate::hybrid::HybridLayer,
6897 x: &CudaSlice<f32>, pos_d: &CudaSlice<i32>, t: usize)
6898 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6899 let n_embd = self.cfg.n_embd as usize;
6900 let eps = self.cfg.rms_eps;
6901
6902 let mut h = e.zeros(t * n_embd)?;
6903 e.rms_norm(x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
6904 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
6905 let o = self.gemma4_attn(e, fa, il, &h, pos_d, t)?;
6906 let mut cur = e.zeros(t * n_embd)?;
6908 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, t, eps)?;
6909 self.gemma4_layer_tail_add(e, layer, &cur, x, t)
6910 }
6911
6912 fn gemma4_layer_tail_add(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
6916 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize)
6917 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6918 Ok(self.gemma4_layer_tail_add_n(e, layer, cur, x, t, None)?.0)
6919 }
6920
6921 fn gemma4_layer_tail_add_n(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
6924 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize,
6925 next_norm: Option<&CudaSlice<f32>>)
6926 -> Result<(CudaSlice<f32>, Option<CudaSlice<f32>>), Box<dyn std::error::Error>> {
6927 let n_embd = self.cfg.n_embd as usize;
6928 let bits = layer.gemma4.as_ref().unwrap();
6929 let (sn, attn_out) = self.gemma4_layer_tail_core(e, layer, cur, x, t)?;
6930 let mut xn = e.uninit(t * n_embd)?;
6931 match next_norm {
6932 Some(w) => {
6933 let mut hn = e.uninit(t * n_embd)?;
6934 e.add_scale_rms_norm(&sn, &attn_out, bits.layer_scale, w, &mut xn, &mut hn,
6935 n_embd, t, self.cfg.rms_eps)?;
6936 Ok((xn, Some(hn)))
6937 }
6938 None => {
6939 e.add_scale(&sn, &attn_out, bits.layer_scale, &mut xn, t * n_embd)?;
6940 Ok((xn, None))
6941 }
6942 }
6943 }
6944
6945 fn gemma4_layer_tail_core(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
6948 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize)
6949 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6950 self.gemma4_layer_tail_core_pn(e, layer, cur, x, t, None, false)
6951 }
6952
6953 fn gemma4_layer_tail_core_pn(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
6960 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize,
6961 pre_norm: Option<&CudaSlice<f32>>, defer_post_norm: bool)
6962 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6963 let n_embd = self.cfg.n_embd as usize;
6964 let eps = self.cfg.rms_eps;
6965 let bits = layer.gemma4.as_ref().unwrap();
6966
6967 let Some(mbits) = bits.moe_bits.as_ref() else {
6970 let crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } = &layer.ffn
6971 else { panic!("gemma4 dense layer without Dense ffn") };
6972 let mut attn_out = e.uninit(t * n_embd)?;
6973 let mut zsh = e.uninit(t * n_embd)?;
6974 let mut zpair: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
6977 match pre_norm {
6978 Some(wa) if t == 1 => {
6979 zpair = Some(e.rms_pre_add_rms_norm_q8z(cur, wa, x,
6980 bits.ffn_norm.float_data(),
6981 &mut attn_out, &mut zsh,
6982 n_embd, t, eps)?);
6983 }
6984 Some(wa) => e.rms_pre_add_rms_norm(cur, wa, x, bits.ffn_norm.float_data(),
6985 &mut attn_out, &mut zsh, n_embd, t, eps)?,
6986 None => e.add_rms_norm(cur, x, bits.ffn_norm.float_data(), &mut attn_out,
6987 &mut zsh, n_embd, t, eps)?,
6988 }
6989 let n_ff = ffn_gate.out_features();
6990 let (gate, up) = if t == 1 {
6996 let (zq, zd) = match zpair {
6997 Some(p) => p,
6998 None => e.quantize_q8_1(&zsh, 1, n_embd)?,
6999 };
7000 match e.matmul_q4_fused2(ffn_gate, ffn_up, &zq, &zd)? {
7001 Some(p) => p,
7002 None => (e.matmul_pre(ffn_gate, &zq, &zd, &zsh, 1)?,
7003 e.matmul_pre(ffn_up, &zq, &zd, &zsh, 1)?),
7004 }
7005 } else {
7006 static F2B: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
7011 let f2b = *F2B.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
7012 let fused = if f2b {
7013 let (zq, zd) = e.quantize_q8_1(&zsh, t, n_embd)?;
7014 e.matmul_q4_fused2_batched(ffn_gate, ffn_up, &zq, &zd, t)?
7015 } else { None };
7016 match fused {
7017 Some(p) => p,
7018 None => {
7019 e.mmq_act_begin();
7021 (e.matmul(ffn_gate, &zsh, t)?, e.matmul(ffn_up, &zsh, t)?)
7022 }
7023 }
7024 };
7025 let mut act = e.uninit(t * n_ff)?;
7026 let f0 = if e.uses_q8_1_fast(ffn_down) {
7029 let upv = e.view(&up, t * n_ff);
7030 let up_all = upv.slice(0..t * n_ff);
7031 let (aq, ad) = e.gelu_tanh_mul_q8_1(&gate, &up_all, &mut act, n_ff, t)?;
7032 e.matmul_pre(ffn_down, &aq, &ad, &act, t)?
7033 } else {
7034 e.gelu_tanh_mul(&gate, &up, &mut act, t * n_ff)?;
7035 e.matmul(ffn_down, &act, t)?
7036 };
7037 if defer_post_norm { return Ok((f0, attn_out)); }
7038 let mut sn = e.uninit(t * n_embd)?;
7039 e.rms_norm(&f0, bits.post_ffw_norm.float_data(), &mut sn, n_embd, t, eps)?;
7040 return Ok((sn, attn_out));
7041 };
7042
7043 assert!(pre_norm.is_none(), "pre-norm fold is dense-entry only");
7044 let mut attn_out = e.uninit(t * n_embd)?;
7049 let mut router_in = e.uninit(t * n_embd)?;
7050 let fast_moe = match &layer.ffn {
7051 crate::hybrid::Ffn::Moe(m) => m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
7052 && expert_dp4a_supported(m.gate_exps.qtype)
7053 && expert_dp4a_supported(m.up_exps.qtype)
7054 && expert_dp4a_supported(m.down_exps.qtype)
7055 && std::env::var("MEMRA_GEMMA_MOE_FAST").as_deref() != Ok("0"),
7056 _ => false,
7057 };
7058 let q8z = t < PRIME_MIN_T && fast_moe;
7059 let (zsh_f32, zsh_q8, moe_q8) = if q8z {
7060 let (z0, m2) = e.add_rms_norm3_q8z(cur, x, bits.ffn_norm.float_data(),
7061 &mbits.router_scale_pre,
7062 mbits.pre_ffw_norm_2.float_data(),
7063 &mut attn_out, &mut router_in, n_embd, t, eps)?;
7064 (None, Some(z0), Some(m2))
7065 } else {
7066 let mut zsh = e.uninit(t * n_embd)?;
7067 let mut moe_in = e.uninit(t * n_embd)?;
7068 e.add_rms_norm3(cur, x, bits.ffn_norm.float_data(), &mbits.router_scale_pre,
7069 mbits.pre_ffw_norm_2.float_data(), &mut attn_out, &mut zsh,
7070 &mut router_in, &mut moe_in, n_embd, t, eps)?;
7071 (Some((zsh, moe_in)), None, None)
7072 };
7073 let attn_out2 = attn_out;
7074 #[allow(unused_variables)]
7075 let attn_out = &attn_out2;
7076 let n_ff = mbits.shared_gate.out_features();
7077 let (gate, up) = if let Some((zq, zd)) = zsh_q8.as_ref() {
7078 if t == 1 {
7079 match e.matmul_q4_fused2(&mbits.shared_gate, &mbits.shared_up, zq, zd)? {
7080 Some(p) => p,
7081 None => {
7082 let h0 = e.zeros(0)?;
7083 (e.matmul_pre(&mbits.shared_gate, zq, zd, &h0, 1)?,
7084 e.matmul_pre(&mbits.shared_up, zq, zd, &h0, 1)?)
7085 }
7086 }
7087 } else {
7088 let h0 = e.zeros(0)?;
7090 (e.matmul_pre(&mbits.shared_gate, zq, zd, &h0, t)?,
7091 e.matmul_pre(&mbits.shared_up, zq, zd, &h0, t)?)
7092 }
7093 } else {
7094 let (zsh, _) = zsh_f32.as_ref().unwrap();
7095 (e.matmul(&mbits.shared_gate, zsh, t)?, e.matmul(&mbits.shared_up, zsh, t)?)
7096 };
7097 let mut act = e.uninit(t * n_ff)?;
7098 e.gelu_tanh_mul(&gate, &up, &mut act, t * n_ff)?;
7099 let mlp0 = e.matmul(&mbits.shared_down, &act, t)?;
7100 let crate::hybrid::Ffn::Moe(m) = &layer.ffn else { panic!("gemma4 layer not MoE") };
7101 let moe0 = match (&moe_q8, &zsh_f32) {
7102 (Some(mq), _) => self.gemma4_moe_q8(e, m, mbits, mq, &router_in, t)?,
7103 (None, Some((_, moe_in))) => self.gemma4_moe(e, m, mbits, moe_in, &router_in, t)?,
7104 _ => unreachable!(),
7105 };
7106 let mut mlp = e.uninit(t * n_embd)?;
7108 let mut moe = e.uninit(t * n_embd)?;
7109 e.rms_norm2x(&mlp0, &moe0, mbits.post_ffw_norm_1.float_data(),
7110 mbits.post_ffw_norm_2.float_data(), &mut mlp, &mut moe, n_embd, t, eps)?;
7111
7112 let mut sum = e.uninit(t * n_embd)?;
7115 let mut sn = e.uninit(t * n_embd)?;
7116 e.add_rms_norm(&mlp, &moe, bits.post_ffw_norm.float_data(), &mut sum, &mut sn,
7117 n_embd, t, eps)?;
7118 Ok((sn, attn_out2))
7119 }
7120
7121 fn gemma4_layer_tail_add_nq(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
7123 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize,
7124 next_norm: Option<&CudaSlice<f32>>)
7125 -> Result<(CudaSlice<f32>, Option<(CudaSlice<i8>, CudaSlice<f32>)>), Box<dyn std::error::Error>> {
7126 let n_embd = self.cfg.n_embd as usize;
7127 let bits = layer.gemma4.as_ref().unwrap();
7128 let (sn, attn_out) = self.gemma4_layer_tail_core(e, layer, cur, x, t)?;
7129 let mut xn = e.uninit(t * n_embd)?;
7130 match next_norm {
7131 Some(w) => {
7132 let pair = e.add_scale_rms_norm_q8_1(&sn, &attn_out, bits.layer_scale, w, &mut xn,
7133 n_embd, t, self.cfg.rms_eps)?;
7134 Ok((xn, Some(pair)))
7135 }
7136 None => {
7137 e.add_scale(&sn, &attn_out, bits.layer_scale, &mut xn, t * n_embd)?;
7138 Ok((xn, None))
7139 }
7140 }
7141 }
7142
7143 fn gemma4_forward(&self, e: &Engine, tokens: &[u32], last_only: bool)
7146 -> Result<Vec<f32>, Box<dyn std::error::Error>> {
7147 if self.is_gemma4_e4b() { return self.gemma4_e4b_forward(e, tokens, last_only); }
7150 let n_embd = self.cfg.n_embd as usize;
7151 let t = tokens.len();
7152 let pos: Vec<i32> = (0..t as i32).collect();
7153 let pos_d = e.htod_i32(&pos)?;
7154
7155 let mut x = self.embed(e, tokens)?;
7156 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
7157 let probe = std::env::var("MEMRA_GEMMA_PROBE").is_ok();
7160 let stat = |e: &Engine, x: &CudaSlice<f32>, tag: &str| -> Result<(), Box<dyn std::error::Error>> {
7161 let h = e.dtoh(x)?;
7162 let bad = h.iter().filter(|v| !v.is_finite()).count();
7163 let mx = h.iter().filter(|v| v.is_finite()).fold(0.0f32, |m, v| m.max(v.abs()));
7164 eprintln!("[gemma-probe] {tag}: tok0_first3={:?} bad={bad} max={mx:.3e}", &h[..3]);
7165 Ok(())
7166 };
7167 if probe { stat(e, &x, "embed")?; }
7168 for (il, layer) in self.layers.iter().enumerate() {
7169 x = self.gemma4_layer(e, il, layer, &x, &pos_d, t)?;
7170 if probe { stat(e, &x, &format!("L{il}"))?; }
7171 }
7172 let mut hn = e.zeros(t * n_embd)?;
7173 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, self.cfg.rms_eps)?;
7174 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
7175 let n_vocab = self.output.out_features();
7176 let logits = if last_only {
7177 let hv = e.view(&hn, t * n_embd);
7178 let last_row = hv.slice((t - 1) * n_embd..t * n_embd);
7179 let mut hlast = e.zeros(n_embd)?;
7180 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
7181 let mut ld = e.matmul(&self.output, &hlast, 1)?;
7182 e.softcap(&mut ld, cap, n_vocab)?;
7183 self.gemma4_suppress(e, &mut ld, 1)?;
7184 e.dtoh(&ld)?
7185 } else {
7186 let mut ld = e.matmul(&self.output, &hn, t)?;
7187 e.softcap(&mut ld, cap, t * n_vocab)?;
7188 self.gemma4_suppress(e, &mut ld, t)?;
7189 e.dtoh(&ld)?
7190 };
7191 Ok(logits)
7192 }
7193
7194 pub(crate) fn gemma4_prime(&self, e: &Engine, tokens: &[u32], cache: &mut Cache)
7199 -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7200 if cache.pos != 0 {
7205 return Err("gemma4 prime v0 is fresh-prompt only (no continuation/chunked prime) \
7206 — prime the full prompt in one call or decode tokenwise".into());
7207 }
7208 let n_embd = self.cfg.n_embd as usize;
7209 let eps = self.cfg.rms_eps;
7210 let t = tokens.len();
7211 let pos: Vec<i32> = (0..t as i32).collect();
7212 let pos_d = e.htod_i32(&pos)?;
7213 let mut x = self.embed(e, tokens)?;
7214 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
7215 for (il, layer) in self.layers.iter().enumerate() {
7216 let mut h = e.zeros(t * n_embd)?;
7217 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
7218 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer not full-attn") };
7219 let o = self.gemma4_attn_prime(e, fa, il, &h, &pos_d, t, Some(cache))?;
7220 let mut cur = e.zeros(t * n_embd)?;
7221 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, t, eps)?;
7222 x = self.gemma4_layer_tail_add(e, layer, &cur, &x, t)?;
7223 self.dflash_tap(e, cache, il, &x, t)?;
7224 }
7225 cache.pos += t;
7226 let hiddens = e.clone_dtod(&x)?;
7227 let xv = e.view(&x, t * n_embd);
7228 let last_row = xv.slice((t - 1) * n_embd..t * n_embd);
7229 let mut h_seed = e.zeros(n_embd)?;
7230 e.copy_view_into(&mut h_seed, 0, &last_row, n_embd)?;
7231 let mut hn = e.uninit(n_embd)?;
7232 e.rms_norm(&h_seed, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
7233 let mut ld = e.matmul(&self.output, &hn, 1)?;
7234 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
7235 e.softcap(&mut ld, cap, self.output.out_features())?;
7236 self.gemma4_suppress(e, &mut ld, 1)?;
7237 let logits = e.dtoh(&ld)?;
7238 Ok((logits, h_seed, hiddens))
7239 }
7240
7241 fn gemma4_decode_attn(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
7246 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
7247 pos_d: &CudaSlice<i32>, cache: &mut Cache)
7248 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7249 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
7250 let eps = self.cfg.rms_eps;
7251 let aux = self.gemma4_aux.as_ref().unwrap();
7252 let (hq, hdq) = (hq, hdq);
7253 let h0 = e.zeros(0)?;
7254 let h = &h0;
7255 let (q0, k0, v0) = if swa {
7256 match e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, &hq, &hdq)? {
7257 Some(t3) => t3,
7258 None => (e.matmul_pre(&fa.wq, &hq, &hdq, h, 1)?,
7259 e.matmul_pre(&fa.wk, &hq, &hdq, h, 1)?,
7260 e.matmul_pre(&fa.wv, &hq, &hdq, h, 1)?),
7261 }
7262 } else {
7263 let (q0, k0) = match e.matmul_q4_fused2(&fa.wq, &fa.wk, &hq, &hdq)? {
7264 Some(p) => p,
7265 None => (e.matmul_pre(&fa.wq, &hq, &hdq, h, 1)?,
7266 e.matmul_pre(&fa.wk, &hq, &hdq, h, 1)?),
7267 };
7268 let v0 = e.clone_dtod(&k0)?;
7269 (q0, k0, v0)
7270 };
7271 let mut q = e.uninit(nh * hd)?;
7272 let mut k = e.uninit(nkv * hd)?;
7273 let mut v = e.uninit(nkv * hd)?;
7274 let ff = if swa { None } else {
7277 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
7278 };
7279 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
7280 &aux.ones, &mut q, &mut k, &mut v, hd, nh, nkv,
7281 pos_d, nh, nkv, base, 1.0, ff, eps)?;
7282 let kvl = cache.kv[il].as_mut().unwrap();
7283 e.append_kv_quantized(&k, &v, &mut kvl.k, &mut kvl.v, kvl.len,
7284 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()))?;
7285 kvl.len += 1;
7286 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
7290 let mut attn = e.uninit(nh * hd)?;
7291 if !swa && hd == 512 && kvl.len >= crate::fa512_min_tkv()
7293 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
7294 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
7295 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
7296 let base = kvl.len as i32;
7298 e.i32_set_k(&mut kvl.len_d, base)?;
7299 e.fa_decode_rows(&q, &kp, &vp, &mut attn, hd, nh, nkv, kvl.len - 1, 1, scale,
7300 kvl.k_tok_bytes, kvl.v_tok_bytes, Some((&kvl.len_d, -1)), false,
7301 false, None)?;
7302 return Ok(e.matmul(&fa.wo, &attn, 1)?);
7303 }
7304 if swa && kvl.len > win && hd == 256
7306 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
7307 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
7308 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
7309 let base = kvl.len as i32;
7310 e.i32_set_k(&mut kvl.len_d, base)?;
7311 e.fa_decode_rows_w(&q, &kp, &vp, &mut attn, hd, nh, nkv, &kvl.len_d, -1, 1, scale,
7312 win, kvl.k_tok_bytes, kvl.v_tok_bytes, None)?;
7313 return Ok(e.matmul(&fa.wo, &attn, 1)?);
7314 }
7315 let (off_tok, t_kv) = if swa && kvl.len > win { (kvl.len - win, win) } else { (0, kvl.len) };
7316 let k_view = e.view_u8_range(&kvl.k, off_tok * kvl.k_tok_bytes,
7317 (off_tok + t_kv) * kvl.k_tok_bytes);
7318 let v_view = e.view_u8_range(&kvl.v, off_tok * kvl.v_tok_bytes,
7319 (off_tok + t_kv) * kvl.v_tok_bytes);
7320 e.fa_decode_kvmod(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, t_kv, scale,
7321 kvl.k_tok_bytes, kvl.v_tok_bytes, swa && crate::Engine::wkv_on())?;
7322 Ok(e.matmul(&fa.wo, &attn, 1)?)
7323 }
7324
7325 #[allow(clippy::too_many_arguments)]
7332 pub fn gemma4_decode_step_dc(&self, e: &Engine, token_d: &CudaSlice<u32>,
7333 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
7334 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
7335 n_vocab: usize, cap_bucket_max: Option<(usize, usize)>)
7336 -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
7337 let mut tok_out = e.stream().alloc_zeros::<u32>(1)?;
7338 self.gemma4_decode_step_dc_into(e, token_d, pos_d, embd_gpu, embd_qt, embd_rb, cache,
7339 n_vocab, cap_bucket_max, &mut tok_out)?;
7340 Ok(tok_out)
7341 }
7342
7343 #[allow(clippy::too_many_arguments)]
7346 pub fn gemma4_decode_step_dc_into(&self, e: &Engine, token_d: &CudaSlice<u32>,
7347 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
7348 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
7349 n_vocab: usize, cap_bucket_max: Option<(usize, usize)>,
7350 tok_out: &mut CudaSlice<u32>)
7351 -> Result<(), Box<dyn std::error::Error>> {
7352 let n_embd = self.cfg.n_embd as usize;
7353 let eps = self.cfg.rms_eps;
7354 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
7355 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
7356 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
7357 let n_layers = self.layers.len();
7358 for (il, layer) in self.layers.iter().enumerate() {
7359 let (hq, hdq) = match h_carry.take() {
7360 Some(p) => p,
7361 None => e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, 1, eps)?,
7362 };
7363 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
7364 let o = self.gemma4_decode_attn_dc(e, fa, il, &hq, &hdq, pos_d, cache, cap_bucket_max)?;
7365 let mut cur = e.uninit(n_embd)?;
7366 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, 1, eps)?;
7367 let next_norm = if il + 1 < n_layers {
7368 Some(self.layers[il + 1].attn_norm.float_data())
7369 } else { None };
7370 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, 1, next_norm)?;
7371 x = xn;
7372 h_carry = hn;
7373 }
7374 let mut hn = e.uninit(n_embd)?;
7375 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
7376 let mut logits = e.matmul(&self.output, &hn, 1)?;
7377 self.gemma4_suppress(e, &mut logits, 1)?; e.argmax_token_device_into(&logits, tok_out, n_vocab)?;
7379 e.inc_seqlen(pos_d)?;
7380 if cap_bucket_max.is_none() { cache.pos += 1; }
7381 Ok(())
7382 }
7383
7384 pub fn g4_dc_slots(&self, e: &Engine) -> Result<G4DcSlots, Box<dyn std::error::Error>> {
7391 let n_embd = self.cfg.n_embd as usize;
7392 let n_vocab = self.output.out_features();
7393 let n_layers = self.layers.len();
7394 let (mut qmax, mut kvmax, mut ffmax) = (0usize, 0usize, 0usize);
7395 for il in 0..n_layers {
7396 let (hd, nkv, nh, _b, _s, _w) = self.gemma4_geom(il);
7397 qmax = qmax.max(nh * hd);
7398 kvmax = kvmax.max(nkv * hd);
7399 if let crate::hybrid::Ffn::Dense { ffn_gate, .. } = &self.layers[il].ffn {
7400 ffmax = ffmax.max(ffn_gate.out_features());
7401 }
7402 }
7403 Ok(G4DcSlots {
7404 x: e.uninit(n_embd)?, xn: e.uninit(n_embd)?, cur: e.uninit(n_embd)?,
7405 hq: e.alloc_i8_uninit(n_embd)?, hd_: e.uninit(n_embd / 32)?,
7406 q0: e.uninit(qmax)?, k0: e.uninit(kvmax)?, v0: e.uninit(kvmax)?,
7407 q: e.uninit(qmax)?, k: e.uninit(kvmax)?, v: e.uninit(kvmax)?,
7408 attn: e.uninit(qmax)?, o: e.uninit(n_embd)?,
7409 attn_out: e.uninit(n_embd)?, zsh: e.uninit(n_embd)?,
7410 zq: e.alloc_i8_uninit(n_embd.max(qmax))?, zd: e.uninit(n_embd.max(qmax) / 32)?,
7413 gate: e.uninit(ffmax)?, up: e.uninit(ffmax)?,
7414 act: e.uninit(ffmax)?, actq: e.alloc_i8_uninit(ffmax)?, actd: e.uninit(ffmax / 32)?,
7415 f0: e.uninit(n_embd)?, sn: e.uninit(n_embd)?,
7416 hn: e.uninit(n_embd)?, logits: e.uninit(n_vocab)?,
7417 })
7418 }
7419
7420 fn g4_matvec_m1_into(&self, e: &Engine, w: &crate::model::GpuTensor,
7423 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, y: &mut CudaSlice<f32>)
7424 -> Result<(), Box<dyn std::error::Error>> {
7425 use crate::model::GpuTensor;
7426 let (bytes, qtype, row_bytes, scale, rp) = match w {
7427 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } =>
7428 (bytes, *qtype, *row_bytes, *scale, *rp),
7429 _ => return Err("g4_matvec_m1_into: non-quant tensor".into()),
7430 };
7431 let (mbytes, mrp) = match w {
7432 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
7433 _ => (bytes, rp),
7434 };
7435 e.qmatvec_mmvq_into(mbytes, aq, ad, 1, w.in_features(), w.out_features(),
7436 qtype, row_bytes, scale, mrp, y)
7437 }
7438
7439 #[allow(clippy::too_many_arguments)]
7443 pub fn gemma4_decode_step_dc_slotted(&self, e: &Engine, token_d: &CudaSlice<u32>,
7444 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
7445 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
7446 n_vocab: usize, cap_bucket_max: Option<(usize, usize)>,
7447 sl: &mut G4DcSlots, tok_out: &mut CudaSlice<u32>,
7448 ring: Option<(&mut CudaSlice<u32>, usize)>)
7449 -> Result<(), Box<dyn std::error::Error>> {
7450 let n_embd = self.cfg.n_embd as usize;
7451 let eps = self.cfg.rms_eps;
7452 e.embed_gather_device_into(embd_gpu, token_d, &mut sl.x, n_embd, embd_qt, embd_rb)?;
7453 e.scale_inplace(&mut sl.x, (n_embd as f32).sqrt(), n_embd)?;
7454 let n_layers = self.layers.len();
7455 let mut has_carry = false;
7456 for il in 0..n_layers {
7457 if !has_carry {
7458 e.rms_norm_q8_1_into(&sl.x, self.layers[il].attn_norm.float_data(), n_embd, 1,
7459 eps, &mut sl.hq, &mut sl.hd_)?;
7460 }
7461 has_carry = true;
7462 let layer = &self.layers[il];
7463 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
7464 self.gemma4_decode_attn_dc_slotted(e, fa, il, pos_d, cache, cap_bucket_max, sl)?;
7465 e.rms_norm(&sl.o, layer.post_attn_norm.float_data(), &mut sl.cur, n_embd, 1, eps)?;
7466 let next_norm = if il + 1 < n_layers {
7467 Some(self.layers[il + 1].attn_norm.float_data())
7468 } else { None };
7469 self.gemma4_layer_tail_slotted(e, layer, next_norm, sl)?;
7470 std::mem::swap(&mut sl.x, &mut sl.xn);
7471 }
7472 e.rms_norm(&sl.x, self.output_norm.float_data(), &mut sl.hn, n_embd, 1, eps)?;
7473 e.quantize_q8_1_into(&sl.hn, 1, n_embd, &mut sl.zq, &mut sl.zd)?;
7474 {
7476 let (zq, zd) = (&sl.zq, &sl.zd);
7477 let zq = unsafe { &*(zq as *const CudaSlice<i8>) };
7478 let zd = unsafe { &*(zd as *const CudaSlice<f32>) };
7479 self.g4_matvec_m1_into(e, &self.output, zq, zd, &mut sl.logits)?;
7480 }
7481 self.gemma4_suppress(e, &mut sl.logits, 1)?;
7482 e.argmax_token_device_into(&sl.logits, tok_out, n_vocab)?;
7483 if let Some((ring, base)) = ring {
7484 e.plain_tok_ring(tok_out, pos_d, base, ring)?;
7488 }
7489 e.inc_seqlen(pos_d)?;
7490 if cap_bucket_max.is_none() { cache.pos += 1; }
7491 Ok(())
7492 }
7493
7494 #[allow(clippy::too_many_arguments)]
7496 fn gemma4_decode_attn_dc_slotted(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer,
7497 il: usize, pos_d: &CudaSlice<i32>, cache: &mut Cache,
7498 cap_bucket_max: Option<(usize, usize)>, sl: &mut G4DcSlots)
7499 -> Result<(), Box<dyn std::error::Error>> {
7500 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
7501 let eps = self.cfg.rms_eps;
7502 let aux = self.gemma4_aux.as_ref().unwrap();
7503 {
7504 let hq = unsafe { &*(&sl.hq as *const CudaSlice<i8>) };
7505 let hdq = unsafe { &*(&sl.hd_ as *const CudaSlice<f32>) };
7506 if swa {
7507 if !e.matmul_q4_fused3_into(&fa.wq, &fa.wk, &fa.wv, hq, hdq,
7508 &mut sl.q0, &mut sl.k0, &mut sl.v0)? {
7509 return Err("slotted step: fused3 unavailable (non-uniform trunk)".into());
7510 }
7511 } else {
7512 if !e.matmul_q4_fused2_into(&fa.wq, &fa.wk, hq, hdq, &mut sl.q0, &mut sl.k0)? {
7513 return Err("slotted step: fused2 unavailable".into());
7514 }
7515 let k0r = unsafe { &*(&sl.k0 as *const CudaSlice<f32>) };
7516 e.copy_into(&mut sl.v0, 0, k0r, nkv * hd)?;
7517 }
7518 }
7519 let ff = if swa { None } else {
7522 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
7523 };
7524 let kvl = cache.kv[il].as_mut().unwrap();
7525 let kv_fp8 = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
7526 if crate::Engine::qkv_append_on() {
7527 e.rms_norm_qkv_rope_append_dc(&sl.q0, &sl.k0, &sl.v0, fa.q_norm.float_data(),
7529 fa.k_norm.float_data(), &aux.ones, &mut sl.q, &mut sl.k, &mut sl.v, hd, nh, nkv,
7530 pos_d, nh, nkv, base, 1.0, ff, eps,
7531 &mut kvl.k, &mut kvl.v, &kvl.len_d, kvl.k_tok_bytes, kvl.v_tok_bytes, kv_fp8)?;
7532 } else {
7533 e.rms_norm_qkv_rope(&sl.q0, &sl.k0, &sl.v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
7534 &aux.ones, &mut sl.q, &mut sl.k, &mut sl.v, hd, nh, nkv,
7535 pos_d, nh, nkv, base, 1.0, ff, eps)?;
7536 e.append_kv_quantized_dc(&sl.k, &sl.v, &mut kvl.k, &mut kvl.v, &kvl.len_d,
7537 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes,
7538 kv_fp8)?;
7539 }
7540 e.inc_seqlen(&mut kvl.len_d)?;
7541 let (b_swa, b_glob) = cap_bucket_max.expect("slotted step is capture-only");
7542 let k_view = e.view_u8(&kvl.k, kvl.k.len());
7543 let v_view = e.view_u8(&kvl.v, kvl.v.len());
7544 let rows_on = std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0");
7545 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
7546 let mut fa_q8 = false;
7550 if !swa && hd == 512 && b_glob >= crate::fa512_min_tkv() && rows_on {
7551 e.fa_decode_rows(&sl.q, &k_view, &v_view, &mut sl.attn, hd, nh, nkv, b_glob - 1,
7552 1, scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
7553 Some((&kvl.len_d, -1)), false, false,
7554 Some((&mut sl.zq, &mut sl.zd)))?;
7555 fa_q8 = true;
7556 } else if swa && b_swa > win && hd == 256 && rows_on {
7557 e.fa_decode_rows_w(&sl.q, &k_view, &v_view, &mut sl.attn, hd, nh, nkv,
7558 &kvl.len_d, -1, 1, scale, win,
7559 kvl.k_tok_bytes, kvl.v_tok_bytes,
7560 Some((&mut sl.zq, &mut sl.zd)))?;
7561 fa_q8 = true;
7562 } else {
7563 let b = if swa { b_swa } else { b_glob };
7564 e.fa_decode_dc(&sl.q, &k_view, &v_view, &mut sl.attn, hd, nh, nkv, &kvl.len_d, b,
7565 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
7566 swa && crate::Engine::wkv_on())?;
7567 }
7568 if !fa_q8 {
7569 let aq = unsafe { &*(&sl.attn as *const CudaSlice<f32>) };
7570 e.quantize_q8_1_into(aq, 1, nh * hd, &mut sl.zq, &mut sl.zd)?;
7571 }
7572 {
7573 let zq = unsafe { &*(&sl.zq as *const CudaSlice<i8>) };
7574 let zd = unsafe { &*(&sl.zd as *const CudaSlice<f32>) };
7575 self.g4_matvec_m1_into(e, &fa.wo, zq, zd, &mut sl.o)?;
7576 }
7577 Ok(())
7578 }
7579
7580 fn gemma4_layer_tail_slotted(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
7583 next_norm: Option<&CudaSlice<f32>>, sl: &mut G4DcSlots)
7584 -> Result<(), Box<dyn std::error::Error>> {
7585 let n_embd = self.cfg.n_embd as usize;
7586 let eps = self.cfg.rms_eps;
7587 let bits = layer.gemma4.as_ref().unwrap();
7588 let crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } = &layer.ffn
7589 else { return Err("slotted tail: dense ffn only".into()) };
7590 e.add_rms_norm(&sl.cur, &sl.x, bits.ffn_norm.float_data(), &mut sl.attn_out,
7591 &mut sl.zsh, n_embd, 1, eps)?;
7592 let n_ff = ffn_gate.out_features();
7593 {
7594 let zshr = unsafe { &*(&sl.zsh as *const CudaSlice<f32>) };
7595 e.quantize_q8_1_into(zshr, 1, n_embd, &mut sl.zq, &mut sl.zd)?;
7596 }
7597 {
7598 let zq = unsafe { &*(&sl.zq as *const CudaSlice<i8>) };
7599 let zd = unsafe { &*(&sl.zd as *const CudaSlice<f32>) };
7600 if !e.matmul_q4_fused2_into(ffn_gate, ffn_up, zq, zd, &mut sl.gate, &mut sl.up)? {
7601 return Err("slotted tail: ffn fused2 unavailable".into());
7602 }
7603 }
7604 debug_assert!(e.uses_q8_1_fast(ffn_down));
7605 {
7606 let upr = unsafe { &*(&sl.up as *const CudaSlice<f32>) };
7607 let upv = e.view(upr, n_ff);
7608 let up_all = upv.slice(0..n_ff);
7609 let gr = unsafe { &*(&sl.gate as *const CudaSlice<f32>) };
7610 e.gelu_tanh_mul_q8_1_into(gr, &up_all, &mut sl.act, n_ff, 1,
7611 &mut sl.actq, &mut sl.actd)?;
7612 }
7613 {
7614 let aq = unsafe { &*(&sl.actq as *const CudaSlice<i8>) };
7615 let ad = unsafe { &*(&sl.actd as *const CudaSlice<f32>) };
7616 self.g4_matvec_m1_into(e, ffn_down, aq, ad, &mut sl.f0)?;
7617 }
7618 e.rms_norm(&sl.f0, bits.post_ffw_norm.float_data(), &mut sl.sn, n_embd, 1, eps)?;
7619 match next_norm {
7620 Some(w) => {
7621 e.add_scale_rms_norm_q8_1_into(&sl.sn, &sl.attn_out, bits.layer_scale, w,
7622 &mut sl.xn, n_embd, 1, eps,
7623 &mut sl.hq, &mut sl.hd_)?;
7624 }
7625 None => {
7626 e.add_scale(&sl.sn, &sl.attn_out, bits.layer_scale, &mut sl.xn, n_embd)?;
7627 }
7628 }
7629 Ok(())
7630 }
7631
7632 #[allow(clippy::too_many_arguments)]
7634 fn gemma4_decode_attn_dc(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
7635 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
7636 pos_d: &CudaSlice<i32>, cache: &mut Cache,
7637 cap_bucket_max: Option<(usize, usize)>)
7638 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
7639 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
7640 let eps = self.cfg.rms_eps;
7641 let aux = self.gemma4_aux.as_ref().unwrap();
7642 let (q0, k0, v0) = if swa {
7643 match e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)? {
7644 Some(t3) => t3,
7645 None => {
7646 let h0 = e.zeros(0)?;
7647 (e.matmul_pre(&fa.wq, hq, hdq, &h0, 1)?,
7648 e.matmul_pre(&fa.wk, hq, hdq, &h0, 1)?,
7649 e.matmul_pre(&fa.wv, hq, hdq, &h0, 1)?)
7650 }
7651 }
7652 } else {
7653 let (q0, k0) = match e.matmul_q4_fused2(&fa.wq, &fa.wk, hq, hdq)? {
7654 Some(p) => p,
7655 None => {
7656 let h0 = e.zeros(0)?;
7657 (e.matmul_pre(&fa.wq, hq, hdq, &h0, 1)?,
7658 e.matmul_pre(&fa.wk, hq, hdq, &h0, 1)?)
7659 }
7660 };
7661 let v0 = e.clone_dtod(&k0)?;
7662 (q0, k0, v0)
7663 };
7664 let mut q = e.uninit(nh * hd)?;
7665 let mut k = e.uninit(nkv * hd)?;
7666 let mut v = e.uninit(nkv * hd)?;
7667 let ff = if swa { None } else {
7669 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
7670 };
7671 let kvl = cache.kv[il].as_mut().unwrap();
7672 let kv_fp8 = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
7673 if crate::Engine::qkv_append_on() {
7674 e.rms_norm_qkv_rope_append_dc(&q0, &k0, &v0, fa.q_norm.float_data(),
7676 fa.k_norm.float_data(), &aux.ones, &mut q, &mut k, &mut v, hd, nh, nkv,
7677 pos_d, nh, nkv, base, 1.0, ff, eps,
7678 &mut kvl.k, &mut kvl.v, &kvl.len_d, kvl.k_tok_bytes, kvl.v_tok_bytes, kv_fp8)?;
7679 } else {
7680 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
7681 &aux.ones, &mut q, &mut k, &mut v, hd, nh, nkv,
7682 pos_d, nh, nkv, base, 1.0, ff, eps)?;
7683 e.append_kv_quantized_dc(&k, &v, &mut kvl.k, &mut kvl.v, &kvl.len_d,
7684 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes, kv_fp8)?;
7685 }
7686 e.inc_seqlen(&mut kvl.len_d)?;
7687 let mut attn = e.uninit(nh * hd)?;
7688 let mut fa_q8: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
7691 match cap_bucket_max {
7696 None => {
7697 kvl.len += 1;
7701 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
7702 if !swa && hd == 512 && kvl.len >= crate::fa512_min_tkv()
7703 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
7704 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
7707 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
7708 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
7709 e.fa_decode_rows(&q, &kp, &vp, &mut attn, hd, nh, nkv, kvl.len - 1, 1,
7710 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
7711 Some((&kvl.len_d, -1)), false, false,
7712 Some((&mut aq8, &mut ad8)))?;
7713 fa_q8 = Some((aq8, ad8));
7714 } else if swa && kvl.len > win && hd == 256
7715 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
7716 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
7718 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
7719 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
7720 e.fa_decode_rows_w(&q, &kp, &vp, &mut attn, hd, nh, nkv, &kvl.len_d, -1,
7721 1, scale, win, kvl.k_tok_bytes, kvl.v_tok_bytes,
7722 Some((&mut aq8, &mut ad8)))?;
7723 fa_q8 = Some((aq8, ad8));
7724 } else {
7725 let (off_tok, t_kv) = if swa && kvl.len > win { (kvl.len - win, win) }
7726 else { (0, kvl.len) };
7727 let k_view = e.view_u8_range(&kvl.k, off_tok * kvl.k_tok_bytes,
7728 (off_tok + t_kv) * kvl.k_tok_bytes);
7729 let v_view = e.view_u8_range(&kvl.v, off_tok * kvl.v_tok_bytes,
7730 (off_tok + t_kv) * kvl.v_tok_bytes);
7731 e.fa_decode_kvmod(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, t_kv, scale,
7732 kvl.k_tok_bytes, kvl.v_tok_bytes, swa && crate::Engine::wkv_on())?;
7733 }
7734 }
7735 Some((b_swa, b_glob)) => {
7736 let k_view = e.view_u8(&kvl.k, kvl.k.len());
7742 let v_view = e.view_u8(&kvl.v, kvl.v.len());
7743 let rows_on = std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0");
7744 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
7745 if !swa && hd == 512 && b_glob >= crate::fa512_min_tkv() && rows_on {
7746 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
7747 e.fa_decode_rows(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, b_glob - 1,
7748 1, scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
7749 Some((&kvl.len_d, -1)), false, false,
7750 Some((&mut aq8, &mut ad8)))?;
7751 fa_q8 = Some((aq8, ad8));
7752 } else if swa && b_swa > win && hd == 256 && rows_on {
7753 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
7754 e.fa_decode_rows_w(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
7755 &kvl.len_d, -1, 1, scale, win,
7756 kvl.k_tok_bytes, kvl.v_tok_bytes,
7757 Some((&mut aq8, &mut ad8)))?;
7758 fa_q8 = Some((aq8, ad8));
7759 } else {
7760 let b = if swa { b_swa } else { b_glob };
7761 e.fa_decode_dc(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, &kvl.len_d, b,
7762 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
7763 swa && crate::Engine::wkv_on())?;
7764 }
7765 }
7766 }
7767 if let Some((aq8, ad8)) = fa_q8 {
7770 let mut y = e.uninit(fa.wo.out_features())?;
7771 self.g4_matvec_m1_into(e, &fa.wo, &aq8, &ad8, &mut y)?;
7772 return Ok(y);
7773 }
7774 Ok(e.matmul(&fa.wo, &attn, 1)?)
7775 }
7776
7777 pub fn gemma4_generate_graph(&self, e: &Engine, prompt_pos: usize, first_token: u32,
7782 cache: &mut Cache, max_new: usize, eos: &[u32],
7783 mut on_token: impl FnMut(u32) -> bool)
7784 -> Result<(Vec<u32>, crate::decode::StopReason), Box<dyn std::error::Error>> {
7785 if self.is_gemma4_e4b() {
7786 return Err("E4B graph serving is unwired (HANDOVER-E4B.md) — dc-eager is the serving arm".into());
7787 }
7788 use crate::decode::StopReason;
7789 let n_vocab = self.output.out_features();
7790 let n_embd = self.cfg.n_embd as usize;
7791 let embd_gpu = self.embd_gpu.get_or_init(|| {
7792 e.upload_u8(&self.embd.raw).expect("embed table upload")
7793 });
7794 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
7795 for kvl in cache.kv.iter_mut().flatten() {
7796 e.set_i32_one(&mut kvl.len_d, kvl.len as i32)?;
7797 }
7798 let mut token_d = e.stream().clone_htod(&[first_token])?;
7799 let mut pos_d = e.htod_i32(&[prompt_pos as i32])?;
7800 let g4 = self.cfg.gemma4.as_ref().unwrap();
7801 let (hd_s, hd_g) = (g4.key_length_swa as usize, g4.key_length_global as usize);
7802 let nkv_s = g4.head_count_kv.iter().zip(g4.swa_pattern.iter())
7804 .find(|p| *p.1).map(|p| *p.0 as usize).unwrap_or(8);
7805 let nkv_g = g4.head_count_kv.iter().zip(g4.swa_pattern.iter())
7806 .find(|p| !*p.1).map(|p| *p.0 as usize).unwrap_or(2);
7807 let mut graphs: std::collections::HashMap<((bool, usize), (bool, usize), bool, bool),
7808 (cudarc::driver::CudaGraph,
7809 Vec<Box<dyn std::any::Any + Send>>)> = Default::default();
7810 let mut slots = self.g4_dc_slots(e)?;
7813 const RING: usize = 64;
7816 const DRAIN: usize = 1;
7822 let mut ring = e.stream().alloc_zeros::<u32>(RING)?;
7823 let ring_base = prompt_pos;
7824 let mut out = Vec::with_capacity(max_new);
7825 let mut reason = StopReason::MaxNew;
7826 let mut next = first_token;
7827 let mut captures = 0usize;
7828 for _ in 0..max_new {
7829 out.push(next);
7830 if eos.contains(&next) { reason = StopReason::Eos; break; }
7831 if !on_token(next) { reason = StopReason::Callback; break; }
7832 let t_kv = cache.pos + 1;
7833 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
7841 let f512 = crate::fa512_min_tkv();
7842 let key_s = if t_kv > win { (true, usize::MAX) }
7843 else { e.fa_bucket_key(t_kv, hd_s, nkv_s, crate::Engine::wkv_on()) };
7844 let (key_g, rung_end) = if t_kv >= f512 {
7845 let end = (t_kv + 1).next_power_of_two().max(f512 * 2);
7848 ((true, end), end)
7849 } else { (e.fa_bucket_key(t_kv, hd_g, nkv_g, false), t_kv) };
7850 let key = (key_s, key_g, t_kv >= f512, t_kv > win);
7851 if !graphs.contains_key(&key) {
7852 let bucket_max = (t_kv, rung_end);
7853 let snap = cache.snapshot(e)?;
7855 let pos_save = e.dtoh_i32_one(&pos_d)?;
7856 let len_save: Vec<Option<i32>> = cache.kv.iter()
7857 .map(|k| k.as_ref().map(|kvl| e.dtoh_i32_one(&kvl.len_d).unwrap())).collect();
7858 let tok_save = e.dtoh_u32_one(&token_d)?;
7859 let graph = {
7864 let tok_ref = &mut token_d;
7865 let pos_ref = &mut pos_d;
7866 let cache_ref = &mut *cache;
7867 let slots_ref = &mut slots;
7868 let ring_ref = &mut ring;
7869 e.capture_graph_retained_flags(
7870 cudarc::driver::sys::CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_USE_NODE_PRIORITY,
7871 |e| {
7872 let tok_in = unsafe { &*(tok_ref as *const CudaSlice<u32>) };
7874 let sl = unsafe { &mut *(slots_ref as *mut G4DcSlots) };
7875 let rg = unsafe { &mut *(ring_ref as *mut CudaSlice<u32>) };
7876 self.gemma4_decode_step_dc_slotted(e, tok_in, pos_ref, embd_gpu, qt, rb,
7877 cache_ref, n_vocab, Some(bucket_max),
7878 sl, tok_ref, Some((rg, ring_base)))
7879 })?
7880 };
7881 cache.rollback(e, &snap, 0)?;
7882 e.set_i32_one(&mut pos_d, pos_save)?;
7883 for (il, ls) in len_save.iter().enumerate() {
7884 if let (Some(kvl), Some(v)) = (cache.kv[il].as_mut(), ls) {
7885 e.set_i32_one(&mut kvl.len_d, *v)?;
7886 }
7887 }
7888 e.set_u32_one(&mut token_d, tok_save)?;
7889 if std::env::var("MEMRA_GRAPH_CENSUS").as_deref() == Ok("1") {
7890 if let Ok(c) = crate::graph_update::node_census(&graph.0) {
7891 eprintln!("[graph-census] {c:?}");
7892 }
7893 }
7894 graphs.insert(key, graph);
7895 captures += 1;
7896 }
7897 let mut chunk = 1usize;
7902 let drain_cap: usize = std::env::var("MEMRA_GRAPH_DRAIN").ok()
7903 .and_then(|v| v.parse().ok()).unwrap_or(DRAIN);
7904 while chunk < drain_cap && out.len() + chunk < max_new {
7905 let t_next = cache.pos + 1 + chunk;
7906 let key_s2 = if t_next > win { (true, usize::MAX) }
7907 else { e.fa_bucket_key(t_next, hd_s, nkv_s, crate::Engine::wkv_on()) };
7908 let key_g2 = if t_next >= f512 {
7909 (true, (t_next + 1).next_power_of_two().max(f512 * 2))
7910 } else { e.fa_bucket_key(t_next, hd_g, nkv_g, false) };
7911 if (key_s2, key_g2, t_next >= f512, t_next > win) != key { break; }
7912 chunk += 1;
7913 }
7914 let g = &graphs.get(&key).unwrap().0;
7915 for _ in 0..chunk { g.launch()?; }
7916 e.stream().synchronize()?;
7917 let ringh = e.dtoh_u32(&ring)?;
7918 for j in 0..chunk {
7919 let pos_j = cache.pos + j;
7920 let tok_j = ringh[(pos_j - ring_base) % RING];
7921 cache.pos += 0; if j + 1 == chunk { next = tok_j; }
7923 else {
7924 out.push(tok_j);
7925 if eos.contains(&tok_j) || !on_token(tok_j) {
7926 reason = if eos.contains(&tok_j) { StopReason::Eos }
7927 else { StopReason::Callback };
7928 let keep = cache.pos + j + 1;
7930 e.set_i32_one(&mut pos_d, keep as i32)?;
7931 for kvl in cache.kv.iter_mut().filter_map(|k| k.as_mut()) {
7932 e.set_i32_one(&mut kvl.len_d, keep as i32)?;
7933 kvl.len = keep;
7934 }
7935 cache.pos = keep;
7936 if std::env::var("MEMRA_GRAPH_STATS").is_ok() {
7937 eprintln!("[gemma-graph] captures={captures} buckets={}", graphs.len());
7938 }
7939 return Ok((out, reason));
7940 }
7941 }
7942 }
7943 cache.pos += chunk;
7944 for kvl in cache.kv.iter_mut().filter_map(|k| k.as_mut()) { kvl.len += chunk; }
7945 }
7946 if std::env::var("MEMRA_GRAPH_STATS").is_ok() {
7947 eprintln!("[gemma-graph] captures={captures} buckets={}", graphs.len());
7948 }
7949 Ok((out, reason))
7950 }
7951
7952 pub(crate) fn gemma4_decode_step_t(&self, e: &Engine, tokens: &[u32], pos0: usize,
7958 cache: &mut Cache)
7959 -> Result<Vec<f32>, Box<dyn std::error::Error>> {
7960 Ok(self.gemma4_decode_step_t_h(e, tokens, pos0, cache)?.0)
7961 }
7962
7963 pub(crate) fn gemma4_decode_step_t_am(&self, e: &Engine, tokens: &[u32], pos0: usize,
7967 cache: &mut Cache)
7968 -> Result<(Vec<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7969 let (ld, hn) = self.gemma4_verify_trunk(e, tokens, pos0, cache, None)?;
7970 let t = tokens.len();
7971 let n_vocab = self.output.out_features();
7972 let mut toks = e.stream().alloc_zeros::<u32>(t)?;
7973 for i in 0..t {
7974 e.argmax_token_device_col(&ld, i, n_vocab, &mut toks, i)?;
7975 }
7976 Ok((e.dtoh_u32(&toks)?, hn))
7977 }
7978
7979 pub(crate) fn gemma4_decode_step_t_am_dev(&self, e: &Engine, tok_d: &CudaSlice<u32>, t: usize,
7982 pos0: usize, cache: &mut Cache)
7983 -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7984 let (ld, hn) = self.gemma4_verify_trunk(e, &vec![0u32; t], pos0, cache, Some(tok_d))?;
7985 let n_vocab = self.output.out_features();
7986 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
7987 for i in 0..t {
7988 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
7989 }
7990 Ok((vam, hn))
7991 }
7992
7993 pub(crate) fn gemma4_decode_step_t_h(&self, e: &Engine, tokens: &[u32], pos0: usize,
7996 cache: &mut Cache)
7997 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7998 let (mut ld, hn) = self.gemma4_verify_trunk(e, tokens, pos0, cache, None)?;
7999 let t = tokens.len();
8000 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
8001 e.softcap(&mut ld, cap, t * self.output.out_features())?;
8002 Ok((e.dtoh(&ld)?, hn))
8003 }
8004
8005 pub(crate) fn verify_stream_scratch(&self, e: &Engine, cap: usize)
8008 -> Result<VerifyStreamScratch, Box<dyn std::error::Error>> {
8009 Ok(VerifyStreamScratch {
8010 pos_d: e.htod_i32(&vec![0i32; cap])?,
8011 row_ctrs: (0..cap).map(|_| e.htod_i32(&[0])).collect::<Result<_, _>>()?,
8012 })
8013 }
8014
8015 pub(crate) fn gemma4_verify_t_am_stream(&self, e: &Engine, tok_d: &CudaSlice<u32>, t: usize,
8023 ctr: &CudaSlice<i32>, hint: usize,
8024 cache: &mut Cache,
8025 scr: &mut VerifyStreamScratch)
8026 -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
8027 let n_embd = self.cfg.n_embd as usize;
8028 let eps = self.cfg.rms_eps;
8029 assert!(t <= scr.row_ctrs.len() && t <= 64);
8030 e.i32_iota_from(ctr, &mut scr.pos_d, t)?;
8031 for i in 0..t {
8032 e.i32_copy_add(ctr, &mut scr.row_ctrs[i], (i + 1) as i32)?;
8033 }
8034 let (pos_d, row_ctrs) = (&scr.pos_d, &scr.row_ctrs);
8035 let embd_gpu = self.embd_gpu.get_or_init(|| {
8036 e.upload_u8(&self.embd.raw).expect("embed table upload")
8037 });
8038 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
8039 let mut x = e.embed_gather_device_td(embd_gpu, tok_d, t, n_embd, qt, rb)?;
8040 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
8041 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
8042 let n_layers = self.layers.len();
8043 for (il, layer) in self.layers.iter().enumerate() {
8044 let (hq, hdq) = match h_carry.take() {
8045 Some(p) => p,
8046 None => e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, t, eps)?,
8047 };
8048 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
8049 let o = self.gemma4_verify_attn_stream(e, fa, il, &hq, &hdq, pos_d, t, cache,
8050 hint, row_ctrs)?;
8051 let mut cur = e.uninit(t * n_embd)?;
8052 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, t, eps)?;
8053 let next_norm = if il + 1 < n_layers {
8054 Some(self.layers[il + 1].attn_norm.float_data())
8055 } else { None };
8056 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, t, next_norm)?;
8057 x = xn;
8058 h_carry = hn;
8059 self.dflash_tap(e, cache, il, &x, t)?;
8060 }
8061 let mut hn = e.uninit(t * n_embd)?;
8062 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
8063 let ld = e.matmul(&self.output, &hn, t)?;
8064 let n_vocab = self.output.out_features();
8065 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
8066 for i in 0..t {
8067 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
8068 }
8069 Ok((vam, hn))
8070 }
8071
8072 fn dflash_tap(&self, e: &Engine, cache: &mut Cache, il: usize, x: &CudaSlice<f32>, t: usize)
8079 -> Result<(), Box<dyn std::error::Error>> {
8080 let Some(taps) = cache.dflash_taps.as_mut() else { return Ok(()) };
8081 let Some(slot) = taps.layer_ids.iter().position(|&l| l == il) else { return Ok(()) };
8082 let h = taps.hidden;
8083 let n_taps = taps.layer_ids.len();
8084 debug_assert_eq!(taps.t, t);
8085 let xv = e.view(x, t * h);
8086 for r in 0..t {
8087 let row = xv.slice(r * h..(r + 1) * h);
8088 e.copy_view_into(&mut taps.buf, r * n_taps * h + slot * h, &row, h)?;
8089 }
8090 Ok(())
8091 }
8092
8093 fn gemma4_verify_trunk(&self, e: &Engine, tokens: &[u32], pos0: usize, cache: &mut Cache,
8094 tok_dev: Option<&CudaSlice<u32>>)
8095 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
8096 let n_embd = self.cfg.n_embd as usize;
8097 let eps = self.cfg.rms_eps;
8098 let t = tokens.len();
8099 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
8100 let pos_d = e.htod_i32(&pos)?;
8101 let mut x = match tok_dev {
8102 Some(td) => {
8103 let embd_gpu = self.embd_gpu.get_or_init(|| {
8104 e.upload_u8(&self.embd.raw).expect("embed table upload")
8105 });
8106 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
8107 e.embed_gather_device_td(embd_gpu, td, t, n_embd, qt, rb)?
8108 }
8109 None => e.htod(&self.embd.gather(n_embd, tokens))?,
8110 };
8111 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
8112 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
8113 let n_layers = self.layers.len();
8114 for (il, layer) in self.layers.iter().enumerate() {
8115 let (hq, hdq) = match h_carry.take() {
8116 Some(p) => p,
8117 None => e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, t, eps)?,
8118 };
8119 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
8120 let o = self.gemma4_verify_attn(e, fa, il, &hq, &hdq, &pos_d, t, cache)?;
8121 let mut cur = e.uninit(t * n_embd)?;
8122 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, t, eps)?;
8123 let next_norm = if il + 1 < n_layers {
8124 Some(self.layers[il + 1].attn_norm.float_data())
8125 } else { None };
8126 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, t, next_norm)?;
8127 x = xn;
8128 h_carry = hn;
8129 self.dflash_tap(e, cache, il, &x, t)?;
8130 }
8131 let mut hn = e.uninit(t * n_embd)?;
8132 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
8133 let mut ld = e.matmul(&self.output, &hn, t)?;
8134 self.gemma4_suppress(e, &mut ld, t)?; cache.pos += t;
8136 Ok((ld, hn))
8137 }
8138
8139 #[allow(clippy::too_many_arguments)]
8147 fn gemma4_verify_attn_stream(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
8148 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
8149 pos_d: &CudaSlice<i32>, t: usize,
8150 cache: &mut Cache, hint: usize,
8151 row_ctrs: &[CudaSlice<i32>])
8152 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8153 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
8154 let eps = self.cfg.rms_eps;
8155 let aux = self.gemma4_aux.as_ref().unwrap();
8156 let h0 = e.zeros(0)?;
8157 let h = &h0;
8158 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8161 let f2b = *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
8162 let fused_qkv = if f2b {
8163 if swa {
8164 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
8165 .map(|(a, b, c)| (a, b, Some(c)))
8166 } else {
8167 e.matmul_q4_fused2_batched(&fa.wq, &fa.wk, hq, hdq, t)?
8168 .map(|(a, b)| (a, b, None))
8169 }
8170 } else { None };
8171 let (q0, k0, v0) = match fused_qkv {
8172 Some((a, b, cv)) => {
8173 let v = match cv { Some(c) => c, None => e.clone_dtod(&b)? };
8174 (a, b, v)
8175 }
8176 None => {
8177 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
8178 let k0 = e.matmul_pre(&fa.wk, hq, hdq, h, t)?;
8179 let v0 = if swa { e.matmul_pre(&fa.wv, hq, hdq, h, t)? }
8180 else { e.clone_dtod(&k0)? };
8181 (q0, k0, v0)
8182 }
8183 };
8184 let mut q = e.uninit(t * nh * hd)?;
8185 let mut k = e.uninit(t * nkv * hd)?;
8186 let mut v = e.uninit(t * nkv * hd)?;
8187 let ff = if swa { None } else {
8190 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
8191 };
8192 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
8193 &aux.ones, &mut q, &mut k, &mut v, hd, nh * t, nkv * t,
8194 pos_d, nh, nkv, base, 1.0, ff, eps)?;
8195 let kvl = cache.kv[il].as_mut().unwrap();
8196 e.append_kv_quantized_rows_dc(&k, &v, &mut kvl.k, &mut kvl.v, &kvl.len_d, t,
8198 kvl.kv_dim_k, kvl.kv_dim_v,
8199 kvl.k_tok_bytes, kvl.v_tok_bytes,
8200 (!swa && crate::Engine::gkv_on())
8201 || (swa && crate::Engine::wkv_on()))?;
8202 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
8205 let mut attn = e.uninit(t * nh * hd)?;
8206 let k_view = e.view_u8(&kvl.k, kvl.k.len());
8207 let v_view = e.view_u8(&kvl.v, kvl.v.len());
8208 if swa && hint + 1 >= win {
8211 e.fa_decode_rows_w(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
8214 &kvl.len_d, 0, t, scale, win,
8215 kvl.k_tok_bytes, kvl.v_tok_bytes, None)?;
8216 } else if hd == 512 && hint + t < crate::fa512_min_tkv() {
8217 let bucket = (hint + t + 2).next_power_of_two()
8230 .min(crate::fa512_min_tkv().saturating_sub(1));
8231 let qv = e.view(&q, t * nh * hd);
8232 for i in 0..t {
8233 let q_row = qv.slice(i * nh * hd..(i + 1) * nh * hd);
8234 let mut q_one = e.uninit(nh * hd)?;
8235 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
8236 let mut a_one = e.uninit(nh * hd)?;
8237 e.fa_decode_dc(&q_one, &k_view, &v_view, &mut a_one, hd, nh, nkv,
8238 &row_ctrs[i], bucket, scale,
8239 kvl.k_tok_bytes, kvl.v_tok_bytes, false)?;
8240 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
8241 }
8242 } else if hd == 512 {
8243 e.fa_decode_rows(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, hint, t, scale,
8246 kvl.k_tok_bytes, kvl.v_tok_bytes,
8247 Some((&kvl.len_d, 0)), false, false, None)?;
8248 } else {
8249 e.fa_decode_rows_dc(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
8251 &kvl.len_d, hint + t, t, scale,
8252 kvl.k_tok_bytes, kvl.v_tok_bytes, 0,
8253 swa && crate::Engine::wkv_on())?;
8254 }
8255 Ok(e.matmul(&fa.wo, &attn, t)?)
8256 }
8257
8258 fn gemma4_verify_attn(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
8259 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
8260 pos_d: &CudaSlice<i32>, t: usize,
8261 cache: &mut Cache)
8262 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8263 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
8264 let eps = self.cfg.rms_eps;
8265 let aux = self.gemma4_aux.as_ref().unwrap();
8266 let n_embd = self.cfg.n_embd as usize;
8267 let _ = n_embd;
8268
8269 let h0 = e.zeros(0)?;
8270 let h = &h0;
8271 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
8274 let f2b = *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
8275 let fused_qkv = if f2b {
8276 if swa {
8277 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
8278 .map(|(a, b, c)| (a, b, Some(c)))
8279 } else {
8280 e.matmul_q4_fused2_batched(&fa.wq, &fa.wk, hq, hdq, t)?
8281 .map(|(a, b)| (a, b, None))
8282 }
8283 } else { None };
8284 let (q0, k0, v0) = match fused_qkv {
8285 Some((a, b, cv)) => {
8286 let v = match cv { Some(c) => c, None => e.clone_dtod(&b)? };
8287 (a, b, v)
8288 }
8289 None => {
8290 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
8291 let k0 = e.matmul_pre(&fa.wk, hq, hdq, h, t)?;
8292 let v0 = if swa { e.matmul_pre(&fa.wv, hq, hdq, h, t)? }
8293 else { e.clone_dtod(&k0)? };
8294 (q0, k0, v0)
8295 }
8296 };
8297 let mut q = e.uninit(t * nh * hd)?;
8298 let mut k = e.uninit(t * nkv * hd)?;
8299 let mut v = e.uninit(t * nkv * hd)?;
8300 let ff = if swa { None } else {
8303 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
8304 };
8305 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
8306 &aux.ones, &mut q, &mut k, &mut v, hd, nh * t, nkv * t,
8307 pos_d, nh, nkv, base, 1.0, ff, eps)?;
8308 let kvl = cache.kv[il].as_mut().unwrap();
8309 let base_len = kvl.len;
8310 e.append_kv_quantized_rows(&k, &v, &mut kvl.k, &mut kvl.v, base_len, t,
8311 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()))?;
8312 kvl.len += t;
8313 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
8314 let mut attn = e.uninit(t * nh * hd)?;
8315 let rows_ok = (hd == 256 && base_len + 1 >= crate::fa_vec_min_tkv())
8318 || (hd == 512 && !swa && base_len + 1 >= crate::fa512_min_tkv());
8321 if rows_ok && (!swa || base_len + t <= win) {
8322 let k_view = e.view_u8(&kvl.k, (base_len + t) * kvl.k_tok_bytes);
8323 let v_view = e.view_u8(&kvl.v, (base_len + t) * kvl.v_tok_bytes);
8324 if hd == 512 {
8325 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
8327 e.fa_decode_rows(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, base_len, t,
8328 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
8329 Some((&kvl.len_d, 0)), false,
8330 swa && crate::Engine::wkv_on(), None)?;
8331 } else {
8332 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
8336 e.fa_decode_rows_dc(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
8337 &kvl.len_d, base_len + t, t, scale,
8338 kvl.k_tok_bytes, kvl.v_tok_bytes, 0,
8339 swa && crate::Engine::wkv_on())?;
8340 }
8341 return Ok(e.matmul(&fa.wo, &attn, t)?);
8342 }
8343 if hd == 256 && swa && base_len + 1 >= win
8351 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
8352 let k_view = e.view_u8(&kvl.k, (base_len + t) * kvl.k_tok_bytes);
8353 let v_view = e.view_u8(&kvl.v, (base_len + t) * kvl.v_tok_bytes);
8354 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
8355 e.fa_decode_rows_w(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, &kvl.len_d, 0,
8356 t, scale, win, kvl.k_tok_bytes, kvl.v_tok_bytes, None)?;
8357 return Ok(e.matmul(&fa.wo, &attn, t)?);
8358 }
8359 for i in 0..t {
8360 let avail = base_len + i + 1;
8361 let (off_tok, t_kv) = if swa && avail > win { (avail - win, win) } else { (0, avail) };
8362 let k_view = e.view_u8_range(&kvl.k, off_tok * kvl.k_tok_bytes,
8363 (off_tok + t_kv) * kvl.k_tok_bytes);
8364 let v_view = e.view_u8_range(&kvl.v, off_tok * kvl.v_tok_bytes,
8365 (off_tok + t_kv) * kvl.v_tok_bytes);
8366 let qi = e.view(&q, t * nh * hd);
8367 let q_row = qi.slice(i * nh * hd..(i + 1) * nh * hd);
8368 let mut q_one = e.uninit(nh * hd)?;
8369 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
8370 let mut a_one = e.uninit(nh * hd)?;
8371 if swa && avail > win && hd == 256
8375 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
8376 let kp = e.view_u8(&kvl.k, avail * kvl.k_tok_bytes);
8377 let vp = e.view_u8(&kvl.v, avail * kvl.v_tok_bytes);
8378 e.i32_set_k(&mut kvl.len_d, (avail - 1) as i32)?;
8379 e.fa_decode_rows_w(&q_one, &kp, &vp, &mut a_one, hd, nh, nkv, &kvl.len_d, 0,
8380 1, scale, win, kvl.k_tok_bytes, kvl.v_tok_bytes, None)?;
8381 } else if !swa && hd == 512 && avail >= crate::fa512_min_tkv()
8382 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
8383 let kp = e.view_u8(&kvl.k, avail * kvl.k_tok_bytes);
8384 let vp = e.view_u8(&kvl.v, avail * kvl.v_tok_bytes);
8385 e.i32_set_k(&mut kvl.len_d, (avail - 1) as i32)?;
8386 e.fa_decode_rows(&q_one, &kp, &vp, &mut a_one, hd, nh, nkv, avail - 1, 1,
8387 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
8388 Some((&kvl.len_d, 0)), false, false, None)?;
8389 } else {
8390 e.fa_decode_kvmod(&q_one, &k_view, &v_view, &mut a_one, hd, nh, nkv, t_kv, scale,
8391 kvl.k_tok_bytes, kvl.v_tok_bytes, swa && crate::Engine::wkv_on())?;
8392 }
8393 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
8394 }
8395 Ok(e.matmul(&fa.wo, &attn, t)?)
8396 }
8397
8398 pub(crate) fn gemma4_decode_step_h(&self, e: &Engine, token: u32, cache: &mut Cache)
8401 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
8402 if let Some(split) = crate::pp::pp2_split(self.layers.len()) {
8407 return self.gemma4_decode_step_h_pp2(e, token, cache, split);
8408 }
8409 if crate::pp::pp_cuts(self.layers.len()).is_some() {
8410 crate::pp::warn_unwired_once("gemma4 eager decode (N>2)");
8411 }
8412 let n_embd = self.cfg.n_embd as usize;
8413 let eps = self.cfg.rms_eps;
8414 let pos_d = e.htod_i32(&[cache.pos as i32])?;
8415 let mut x = e.htod(&self.embd.gather(n_embd, &[token]))?;
8416 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
8417 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
8420 let n_layers = self.layers.len();
8421 for (il, layer) in self.layers.iter().enumerate() {
8422 let (hq, hdq) = match h_carry.take() {
8423 Some(p) => p,
8424 None => e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, 1, eps)?,
8425 };
8426 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
8427 let o = self.gemma4_decode_attn(e, fa, il, &hq, &hdq, &pos_d, cache)?;
8428 let mut cur = e.uninit(n_embd)?;
8429 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, 1, eps)?;
8430 let next_norm = if il + 1 < n_layers {
8431 Some(self.layers[il + 1].attn_norm.float_data())
8432 } else { None };
8433 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, 1, next_norm)?;
8434 x = xn;
8435 h_carry = hn;
8436 }
8437 let mut hn = e.uninit(n_embd)?;
8438 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
8439 let h_seed = e.clone_dtod(&x)?;
8440 let mut ld = e.matmul(&self.output, &hn, 1)?;
8441 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
8442 e.softcap(&mut ld, cap, self.output.out_features())?; self.gemma4_suppress(e, &mut ld, 1)?;
8444 let logits = e.dtoh(&ld)?;
8445 cache.pos += 1;
8446 Ok((logits, h_seed))
8447 }
8448
8449 fn gemma4_decode_layers(&self, e: &Engine, mut x: CudaSlice<f32>, lo: usize, hi: usize,
8457 pos_d: &CudaSlice<i32>, cache: &mut Cache)
8458 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8459 let n_embd = self.cfg.n_embd as usize;
8460 let eps = self.cfg.rms_eps;
8461 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
8462 for il in lo..hi {
8463 let layer = &self.layers[il];
8464 let (hq, hdq) = match h_carry.take() {
8465 Some(p) => p,
8466 None => e.rms_norm_q8_1(&x, self.layers[il].attn_norm.float_data(), n_embd, 1, eps)?,
8468 };
8469 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
8470 let o = self.gemma4_decode_attn(e, fa, il, &hq, &hdq, pos_d, cache)?;
8471 let mut cur = e.uninit(n_embd)?;
8472 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, 1, eps)?;
8473 let next_norm = if il + 1 < hi {
8474 Some(self.layers[il + 1].attn_norm.float_data())
8475 } else { None };
8476 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, 1, next_norm)?;
8477 x = xn;
8478 h_carry = hn;
8479 }
8480 Ok(x)
8481 }
8482
8483 fn gemma4_decode_step_h_pp2(&self, e: &Engine, token: u32, cache: &mut Cache, split: usize)
8490 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
8491 if crate::pp::pp_host_bounce_active() {
8492 return Err(
8493 "gemma4_decode_step_h_pp2: refused with MEMRA_PP_HOST_BOUNCE=1 because stage 1 \
8494 still peer-reads stage 0's position buffer; add a stage-local position upload \
8495 before enabling host bounce for gemma4"
8496 .into(),
8497 );
8498 }
8499 if crate::pp::pp2_streams_off() {
8500 return self.gemma4_decode_step_h_pp2_samestream(e, token, cache, split);
8501 }
8502 let rt = crate::pp::Pp2Rt::get(e)?;
8503 let e0 = rt.engine(0, e);
8504 let e1 = rt.engine(1, e);
8505 let n_embd = self.cfg.n_embd as usize;
8506 let eps = self.cfg.rms_eps;
8507
8508 let (pos_d, slot) = {
8510 let _st0 = rt.enter(0);
8511 let pos_d = e0.htod_i32(&[cache.pos as i32])?;
8512 let mut x = e0.htod(&self.embd.gather(n_embd, &[token]))?;
8513 e0.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
8514 let x = self.gemma4_decode_layers(e0, x, 0, split, &pos_d, cache)?;
8515 let slot = rt.tx(0, &x, n_embd)?;
8516 (pos_d, slot)
8517 };
8518
8519 let _st1 = rt.enter(1);
8521 let x = rt.rx(0, slot, n_embd)?;
8522 let x = self.gemma4_decode_layers(e1, x, split, self.layers.len(), &pos_d, cache)?;
8523
8524 let mut hn = e1.uninit(n_embd)?;
8525 e1.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
8526 let h_seed = e1.clone_dtod(&x)?;
8527 let mut ld = e1.matmul(&self.output, &hn, 1)?;
8528 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
8529 e1.softcap(&mut ld, cap, self.output.out_features())?;
8530 self.gemma4_suppress(e1, &mut ld, 1)?;
8531 let logits = e1.dtoh(&ld)?;
8532 cache.pos += 1;
8533 Ok((logits, h_seed))
8534 }
8535
8536 fn gemma4_decode_step_h_pp2_samestream(&self, e: &Engine, token: u32, cache: &mut Cache,
8539 split: usize)
8540 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
8541 let n_embd = self.cfg.n_embd as usize;
8542 let eps = self.cfg.rms_eps;
8543 let pos_d = e.htod_i32(&[cache.pos as i32])?;
8544
8545 let mut x = e.htod(&self.embd.gather(n_embd, &[token]))?;
8547 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
8548 let x = self.gemma4_decode_layers(e, x, 0, split, &pos_d, cache)?;
8549
8550 let boundary_tx = e.clone_dtod(&x)?;
8552 let boundary_rx = e.clone_dtod(&boundary_tx)?;
8553
8554 let x = self.gemma4_decode_layers(e, boundary_rx, split, self.layers.len(), &pos_d, cache)?;
8556
8557 let mut hn = e.uninit(n_embd)?;
8558 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
8559 let h_seed = e.clone_dtod(&x)?;
8560 let mut ld = e.matmul(&self.output, &hn, 1)?;
8561 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
8562 e.softcap(&mut ld, cap, self.output.out_features())?;
8563 self.gemma4_suppress(e, &mut ld, 1)?;
8564 let logits = e.dtoh(&ld)?;
8565 cache.pos += 1;
8566 Ok((logits, h_seed))
8567 }
8568}
8569
8570impl HybridModel {
8589 pub(crate) fn step35_geom(&self, il: usize) -> memra_gguf::config::LayerGeometry {
8592 let geometry = self.cfg.layer_geometry(il as u32)
8593 .unwrap_or_else(|| panic!("step35 layer {il} has no geometry-table row"));
8594 debug_assert_eq!(
8595 geometry.attention_gate,
8596 memra_gguf::config::AttentionGateKind::SeparateHead
8597 );
8598 geometry
8599 }
8600
8601 #[allow(clippy::too_many_arguments)]
8661 fn step35_attn_pre_wo(&self, e: &Engine, fa: &FullAttnLayer, mut g3: Vec<CudaSlice<f32>>,
8662 hg: Option<&CudaSlice<f32>>, gt_pre: Option<&CudaSlice<f32>>,
8663 pos_d: &CudaSlice<i32>, t: usize,
8664 cache: Option<&mut Cache>, il: usize, seq_end: usize)
8665 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8666 let geometry = self.step35_geom(il);
8667 let hd = geometry.head_dim_k as usize;
8668 let nkv = geometry.n_head_kv as usize;
8669 let nh = geometry.n_head as usize;
8670 let rbase = geometry.rope_base;
8671 let scale = geometry.attention_scale();
8672 let swa = geometry.window.is_some();
8673 let eps = self.cfg.rms_eps;
8674 let win = geometry.window.unwrap_or(0) as usize;
8675 let n_rot = geometry.n_rot as usize;
8676
8677 let v = g3.pop().unwrap();
8678 let k0 = g3.pop().unwrap();
8679 let q0 = g3.pop().unwrap();
8680
8681 let mut q = e.uninit(t * nh * hd)?;
8685 e.rms_norm(&q0, fa.q_norm.float_data(), &mut q, hd, nh * t, eps)?;
8686 let mut k = e.uninit(t * nkv * hd)?;
8687 e.rms_norm(&k0, fa.k_norm.float_data(), &mut k, hd, nkv * t, eps)?;
8688 let ff = if geometry.rope_factors {
8689 self.step35_aux.as_ref().and_then(|a| a.rope_freqs(e))
8690 } else {
8691 None
8692 };
8693 e.rope_neox2(&mut q, &mut k, pos_d, hd, n_rot, nh, nkv, t, rbase, 1.0, ff)?;
8694
8695 let mut attn = e.uninit(t * nh * hd)?;
8696 match cache {
8697 Some(cache) => {
8698 let base_len = cache.kv[il].as_ref().unwrap().len;
8699 let legacy_tkv = std::env::var("MEMRA_STEP35_SWA_TKV").as_deref() == Ok("1");
8701 let legacy_calllocal =
8702 std::env::var("MEMRA_PRIME_CALLLOCAL").as_deref() == Ok("1");
8703 let off = if swa {
8704 let raw = base_len.saturating_sub(win - 1);
8705 if legacy_tkv || legacy_calllocal { raw } else { raw & !31usize }
8706 } else {
8707 0
8708 };
8709 {
8710 let kvl = cache.kv[il].as_mut().unwrap();
8711 assert!(kvl.len + t <= cache.max_ctx, "step35 prime: KV overflow");
8712 let write_row = e.prepare_kv_append(kvl, off, t)?;
8713 e.append_kv_quantized_rows(&k, &v, &mut kvl.k, &mut kvl.v, write_row, t,
8714 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes,
8715 kvl.v_tok_bytes, crate::Engine::kv_fp8_on())?;
8716 kvl.len += t;
8717 let new_len = kvl.len as i32;
8718 e.set_i32_one(&mut kvl.len_d, new_len)?;
8719 }
8720 let kvl = cache.kv[il].as_ref().unwrap();
8721 let t_kv = base_len + t - off;
8744 let physical = kvl.physical_rows(off, off + t_kv)?;
8745 let k_view = e.view_u8_range(&kvl.k, physical.start * kvl.k_tok_bytes,
8746 physical.end * kvl.k_tok_bytes);
8747 let v_view = e.view_u8_range(&kvl.v, physical.start * kvl.v_tok_bytes,
8748 physical.end * kvl.v_tok_bytes);
8749 let swa_naive = if legacy_tkv { t_kv > win } else { seq_end > win };
8761 if swa && swa_naive {
8762 if std::env::var("MEMRA_STEP35_SWA_FA").as_deref() == Ok("0") {
8775 e.sdpa_naive_w_quantized_view(&q, &k_view, &v_view, &mut attn, hd, nh,
8776 nkv, t, t_kv, scale, true, win,
8777 kvl.k_tok_bytes, kvl.v_tok_bytes)?;
8778 } else {
8779 e.fa_prefill_view_ws_w_hd128(&q, &k_view, &v_view, &mut attn, hd, nh,
8780 nkv, t, t_kv, scale, true, win,
8781 kvl.k_tok_bytes, kvl.v_tok_bytes)?;
8782 }
8783 } else if std::env::var("MEMRA_NOFA").is_ok() {
8784 e.sdpa_naive_quantized_view(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
8785 t, t_kv, scale, true,
8786 kvl.k_tok_bytes, kvl.v_tok_bytes)?;
8787 } else {
8788 e.fa_prefill_view_ws(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
8793 t, t_kv, scale, true,
8794 kvl.k_tok_bytes, kvl.v_tok_bytes,
8795 crate::Engine::kv_fp8_on())?;
8796 }
8797 }
8798 None => {
8799 debug_assert_eq!(seq_end, t, "step35 cacheless prefill is monolithic (seq_end == t)");
8804 if swa && seq_end > win {
8805 e.sdpa_naive_w(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
8806 } else if std::env::var("MEMRA_NOFA").is_ok() {
8807 e.sdpa_naive(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
8808 } else {
8809 e.fa_prefill(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
8810 }
8811 }
8812 }
8813
8814 let gw = fa.attn_gate.as_ref()
8817 .ok_or("step35 layer is missing attn_gate.weight (head-wise attention gate)")?;
8818 let gt_owned = if gt_pre.is_none() {
8819 Some(e.matmul(
8820 gw,
8821 hg.ok_or("step35 attention needs hg when gt_pre is absent")?,
8822 t,
8823 )?)
8824 } else {
8825 None
8826 };
8827 let gt = gt_pre.or(gt_owned.as_ref()).unwrap();
8828 let mut ag = e.uninit(t * nh * hd)?;
8829 e.attn_head_gate(&attn, gt, &mut ag, None, hd, nh, t)?;
8830 Ok(ag)
8831 }
8832
8833 pub(crate) fn step35_attn(&self, e: &Engine, fa: &FullAttnLayer, h: &CudaSlice<f32>,
8836 pos_d: &CudaSlice<i32>, t: usize, il: usize)
8837 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8838 let g3 = e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?;
8839 let ag = self.step35_attn_pre_wo(e, fa, g3, Some(h), None, pos_d, t, None, il, t)?;
8841 Ok(e.matmul(&fa.wo, &ag, t)?)
8842 }
8843
8844 #[allow(clippy::too_many_arguments)]
8851 pub(crate) fn step35_attn_prime(&self, e: &Engine, fa: &FullAttnLayer, h: &CudaSlice<f32>,
8852 hx: Option<&CudaSlice<u8>>, pos_d: &CudaSlice<i32>, t: usize,
8853 cache: &mut Cache, il: usize, seq_end: usize)
8854 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8855 let g3 = match hx {
8856 Some(xh) => e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], h, xh, t)?,
8857 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
8858 };
8859 let ag = self.step35_attn_pre_wo(
8860 e,
8861 fa,
8862 g3,
8863 Some(h),
8864 None,
8865 pos_d,
8866 t,
8867 Some(cache),
8868 il,
8869 seq_end,
8870 )?;
8871 Ok(e.matmul(&fa.wo, &ag, t)?)
8872 }
8873
8874 #[allow(clippy::too_many_arguments)]
8884 pub(crate) fn step35_decode_attn(&self, e: &Engine, fa: &FullAttnLayer, il: usize,
8885 h: &CudaSlice<f32>,
8886 pre_q: Option<(&CudaSlice<i8>, &CudaSlice<f32>)>,
8887 pos_d: &CudaSlice<i32>, cache: &mut Cache)
8888 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
8889 let geometry = self.step35_geom(il);
8890 let hd = geometry.head_dim_k as usize;
8891 let nkv = geometry.n_head_kv as usize;
8892 let nh = geometry.n_head as usize;
8893 let rbase = geometry.rope_base;
8894 let scale = geometry.attention_scale();
8895 let swa = geometry.window.is_some();
8896 let eps = self.cfg.rms_eps;
8897 let win = geometry.window.unwrap_or(0) as usize;
8898 let n_rot = geometry.n_rot as usize;
8899 let n_embd = self.cfg.n_embd as usize;
8900 let gw = fa.attn_gate.as_ref()
8901 .ok_or("step35 layer is missing attn_gate.weight (head-wise attention gate)")?;
8902
8903 let (q0, k0, v0, gt) = match pre_q {
8904 Some((hq, hdq)) => {
8905 debug_assert!(e.uses_q8_1_fast(gw),
8906 "step35 pre-quantized decode requires attn_gate on the q8_1 fast path \
8907 (h is a zero-length placeholder here) — see mixer_in_q8_1_fast");
8908 let (a, b, c) = match e.matmul_q8_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)? {
8909 Some(t3) => t3,
8910 None => (e.matmul_pre(&fa.wq, hq, hdq, h, 1)?,
8911 e.matmul_pre(&fa.wk, hq, hdq, h, 1)?,
8912 e.matmul_pre(&fa.wv, hq, hdq, h, 1)?),
8913 };
8914 let gt = e.matmul_pre(gw, hq, hdq, h, 1)?;
8915 (a, b, c, gt)
8916 }
8917 None => {
8918 if e.uses_q8_1_fast(&fa.wq) && e.uses_q8_1_fast(&fa.wk)
8919 && e.uses_q8_1_fast(&fa.wv) && e.uses_q8_1_fast(gw) {
8920 let (hq, hdq) = e.quantize_q8_1(h, 1, n_embd)?;
8921 let (a, b, c) = match e.matmul_q8_fused3(&fa.wq, &fa.wk, &fa.wv, &hq, &hdq)? {
8922 Some(t3) => t3,
8923 None => (e.matmul_pre(&fa.wq, &hq, &hdq, h, 1)?,
8924 e.matmul_pre(&fa.wk, &hq, &hdq, h, 1)?,
8925 e.matmul_pre(&fa.wv, &hq, &hdq, h, 1)?),
8926 };
8927 let gt = e.matmul_pre(gw, &hq, &hdq, h, 1)?;
8928 (a, b, c, gt)
8929 } else {
8930 (e.matmul(&fa.wq, h, 1)?, e.matmul(&fa.wk, h, 1)?,
8931 e.matmul(&fa.wv, h, 1)?, e.matmul(gw, h, 1)?)
8932 }
8933 }
8934 };
8935
8936 let mut q = e.uninit(nh * hd)?;
8937 e.rms_norm(&q0, fa.q_norm.float_data(), &mut q, hd, nh, eps)?;
8938 let mut k = e.uninit(nkv * hd)?;
8939 e.rms_norm(&k0, fa.k_norm.float_data(), &mut k, hd, nkv, eps)?;
8940 let ff = if swa { None } else {
8941 self.step35_aux.as_ref().and_then(|a| a.rope_freqs(e))
8942 };
8943 e.rope_neox2(&mut q, &mut k, pos_d, hd, n_rot, nh, nkv, 1, rbase, 1.0, ff)?;
8944
8945 if std::env::var("MEMRA_NOFA").is_ok() {
8946 return Err("MEMRA_NOFA (naive f32 SDPA) is incompatible with the quantized KV \
8947 cache; unset MEMRA_NOFA to use fa_decode".into());
8948 }
8949 let kvl = cache.kv[il].as_mut().unwrap();
8950 let next_len = kvl.len + 1;
8951 let (off, t_kv) = if swa && next_len > win { (next_len - win, win) } else { (0, next_len) };
8952 let write_row = e.prepare_kv_append(kvl, off & !31usize, 1)?;
8953 e.append_kv_quantized(&k, &v0, &mut kvl.k, &mut kvl.v, write_row,
8954 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes,
8955 crate::Engine::kv_fp8_on())?;
8956 kvl.len = next_len;
8957 let physical = kvl.physical_rows(off, off + t_kv)?;
8958 let k_view = e.view_u8_range(&kvl.k, physical.start * kvl.k_tok_bytes,
8959 physical.end * kvl.k_tok_bytes);
8960 let v_view = e.view_u8_range(&kvl.v, physical.start * kvl.v_tok_bytes,
8961 physical.end * kvl.v_tok_bytes);
8962 let mut attn = e.uninit(nh * hd)?;
8963 e.fa_decode_kvmod(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, t_kv, scale,
8964 kvl.k_tok_bytes, kvl.v_tok_bytes, crate::Engine::kv_fp8_on())?;
8965
8966 let mut ag = e.uninit(nh * hd)?;
8967 e.attn_head_gate(&attn, >, &mut ag, None, hd, nh, 1)?;
8968 Ok(e.matmul(&fa.wo, &ag, 1)?)
8969 }
8970}
8971
8972impl HybridModel {
8981 pub fn is_gemma4_e4b(&self) -> bool {
8982 self.gemma4_aux.as_ref().is_some_and(|a| a.e4b.is_some())
8983 }
8984
8985 fn gemma4_e4b_geom(&self, il: usize) -> (usize, usize, usize, f32, f32, bool) {
8989 let g = self.cfg.gemma4.as_ref().unwrap();
8990 let swa = g.swa_pattern[il];
8991 let hd = if swa { g.key_length_swa } else { g.key_length_global } as usize;
8992 let Mixer::Full(fa) = &self.layers[il].mixer else { panic!("e4b layer {il} not full-attn") };
8993 let nh = fa.wq.out_features() / hd;
8994 let nkv = fa.wk.out_features() / hd;
8995 (hd, nkv, nh, if swa { g.rope_base_swa } else { g.rope_base_global }, 1.0, swa)
8996 }
8997
8998 fn gemma4_e4b_kv_target(&self, il: usize) -> Option<usize> {
9000 self.layers[il].gemma4.as_ref()
9001 .and_then(|b| b.e4b.as_ref())
9002 .and_then(|e4| e4.kv_share.map(|t| t as usize))
9003 }
9004
9005 fn gemma4_e4b_inp_pl(&self, e: &Engine, tokens: &[u32], x_scaled: &CudaSlice<f32>, t: usize)
9010 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9011 let tok_d = e.stream().clone_htod(&tokens.to_vec())?;
9012 self.gemma4_e4b_inp_pl_dev(e, &tok_d, x_scaled, t)
9013 }
9014
9015 fn gemma4_e4b_inp_pl_dev(&self, e: &Engine, tok_d: &CudaSlice<u32>,
9017 x_scaled: &CudaSlice<f32>, t: usize)
9018 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9019 let aux = self.gemma4_aux.as_ref().unwrap();
9020 let m = aux.e4b.as_ref().unwrap();
9021 let n_embd = self.cfg.n_embd as usize;
9022 let n_layer = self.layers.len();
9023 let width = m.n_epl * n_layer;
9024 let tbl = m.tok_tbl_gpu.get_or_init(|| {
9025 e.upload_u8(&m.tok_embd_bytes).expect("e4b per-layer token table upload")
9026 });
9027 let mut a = e.embed_gather_device_td(tbl, tok_d, t, width, m.tok_embd_qt,
9028 m.tok_embd_row_bytes)?;
9029 e.scale_inplace(&mut a, (m.n_epl as f32).sqrt(), t * width)?;
9030 let mut p = e.matmul(&m.model_proj, x_scaled, t)?;
9031 e.scale_inplace(&mut p, 1.0 / (n_embd as f32).sqrt(), t * width)?;
9032 let mut pn = e.uninit(t * width)?;
9033 e.rms_norm(&p, m.proj_norm.float_data(), &mut pn, m.n_epl, t * n_layer,
9034 self.cfg.rms_eps)?;
9035 let mut out = e.uninit(t * width)?;
9036 e.add_scale(&a, &pn, 1.0 / 2f32.sqrt(), &mut out, t * width)?;
9037 Ok(out)
9038 }
9039
9040 #[allow(clippy::too_many_arguments)]
9045 fn gemma4_e4b_attn(&self, e: &Engine, il: usize,
9046 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
9047 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache,
9048 dc_bucket: Option<usize>)
9049 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
9050 let (hd, nkv, nh, base, scale, swa) = self.gemma4_e4b_geom(il);
9051 let eps = self.cfg.rms_eps;
9052 let aux = self.gemma4_aux.as_ref().unwrap();
9053 let Mixer::Full(fa) = &self.layers[il].mixer else { unreachable!() };
9054 let h0 = e.zeros(0)?;
9058 let h = &h0;
9059
9060 let ff = if swa { None } else {
9061 Some(aux.rope_freqs.as_ref().expect("e4b global rope needs rope_freqs.weight"))
9062 };
9063 let share = self.gemma4_e4b_kv_target(il);
9064 let mut kv_f32: Option<(CudaSlice<f32>, CudaSlice<f32>)> = None;
9066 let mut q;
9067 if let Some(_tgt) = share {
9068 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
9069 q = e.uninit(t * nh * hd)?;
9070 let mut kdummy = e.uninit(1)?;
9073 let mut vdummy = e.uninit(1)?;
9074 e.rms_norm_qkv_rope(&q0, &q0, &q0, fa.q_norm.float_data(),
9075 fa.q_norm.float_data(), &aux.ones,
9076 &mut q, &mut kdummy, &mut vdummy, hd, nh * t, 0,
9077 pos_d, nh, 1, base, 1.0, ff, eps)?;
9078 } else {
9079 let e4bits = self.layers[il].gemma4.as_ref().and_then(|g| g.e4b.as_ref());
9083 let cat = e4bits.and_then(|e4| e4.qkv_cat.as_ref());
9084 q = e.uninit(t * nh * hd)?;
9085 let mut k = e.uninit(t * nkv * hd)?;
9086 let mut v = e.uninit(t * nkv * hd)?;
9087 if t == 1 && cat.is_some() {
9088 let qkv0 = e.matmul_pre(cat.unwrap(), hq, hdq, h, 1)?;
9089 e.rms_norm_qkv_rope_cat(&qkv0, fa.q_norm.float_data(), fa.k_norm.float_data(),
9090 &aux.ones, &mut q, &mut k, &mut v, hd, nh, nkv,
9091 pos_d, nh, nkv, base, 1.0, ff, eps)?;
9092 } else {
9093 let (q0, k0, v0) = match if t == 1 {
9094 e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)?
9095 } else {
9096 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
9099 if *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0")) {
9100 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
9101 } else { None }
9102 } {
9103 Some(triple) => triple,
9104 None => (e.matmul_pre(&fa.wq, hq, hdq, h, t)?,
9105 e.matmul_pre(&fa.wk, hq, hdq, h, t)?,
9106 e.matmul_pre(&fa.wv, hq, hdq, h, t)?), };
9108 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(),
9111 fa.k_norm.float_data(), &aux.ones, &mut q, &mut k, &mut v,
9112 hd, nh * t, nkv * t, pos_d, nh, nkv, base, 1.0, ff, eps)?;
9113 }
9114 let kvl = cache.kv[il].as_mut().unwrap();
9115 let cls = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
9119 if dc_bucket.is_some() {
9120 debug_assert!(t == 1);
9125 e.append_kv_quantized_row_dc_inc(&k, &v, &mut kvl.k, &mut kvl.v,
9127 &mut kvl.len_d, kvl.kv_dim_k, kvl.kv_dim_v,
9128 kvl.k_tok_bytes, kvl.v_tok_bytes, cls)?;
9129 } else {
9130 e.append_kv_quantized_rows(&k, &v, &mut kvl.k, &mut kvl.v, kvl.len, t,
9131 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes,
9132 kvl.v_tok_bytes, cls)?;
9133 kvl.len += t;
9134 }
9135 kv_f32 = Some((k, v));
9136 }
9137 let kvl_idx = share.unwrap_or(il);
9140 let kvl = cache.kv[kvl_idx].as_ref().unwrap();
9141 let base_len = kvl.len - t; let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
9143 let mut attn = e.uninit(t * nh * hd)?;
9144 if t > 1 && base_len == 0 && std::env::var("MEMRA_NOFA").is_err() {
9156 if let Some((kf, vf)) = &kv_f32 {
9157 if hd == 256 && t <= win {
9158 e.fa_prefill(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale, true)?;
9159 return Ok(e.matmul(&fa.wo, &attn, t)?);
9160 }
9161 if hd == 256 && swa && t > win {
9162 e.fa_prefill_w(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale, true,
9163 win)?;
9164 return Ok(e.matmul(&fa.wo, &attn, t)?);
9165 }
9166 if hd == 512 && !swa {
9167 e.fa_prefill_hd512(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale,
9168 true)?;
9169 return Ok(e.matmul(&fa.wo, &attn, t)?);
9170 }
9171 } else if share.is_some() {
9172 let g = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
9173 let k_view = e.view_u8(&kvl.k, kvl.k.len());
9174 let v_view = e.view_u8(&kvl.v, kvl.v.len());
9175 if hd == 256 && (!swa || t <= win) {
9176 e.fa_prefill_view(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, t, t,
9178 scale, true, kvl.k_tok_bytes, kvl.v_tok_bytes, g)?;
9179 return Ok(e.matmul(&fa.wo, &attn, t)?);
9180 }
9181 let kv_dim = nkv * hd;
9184 let mut kf = e.uninit(t * kv_dim)?;
9185 let mut vf = e.uninit(t * kv_dim)?;
9186 e.fa_dequant_kv_view_f32(&k_view, &v_view, &mut kf, &mut vf, kv_dim, kv_dim,
9187 t, kvl.k_tok_bytes, kvl.v_tok_bytes, g)?;
9188 if hd == 512 {
9189 e.fa_prefill_hd512(&q, &kf, &vf, &mut attn, hd, nh, nkv, t, t, scale,
9190 true)?;
9191 } else {
9192 e.fa_prefill_w(&q, &kf, &vf, &mut attn, hd, nh, nkv, t, t, scale, true,
9193 win)?;
9194 }
9195 return Ok(e.matmul(&fa.wo, &attn, t)?);
9196 }
9197 }
9198 if let Some(bucket) = dc_bucket {
9199 assert!(t == 1);
9204 let bucket = if hd == 512 && win <= crate::fa512_min_tkv() {
9210 bucket.min(crate::fa512_min_tkv().saturating_sub(1))
9211 } else { bucket };
9212 let k_view = e.view_u8(&kvl.k, kvl.k.len());
9213 let v_view = e.view_u8(&kvl.v, kvl.v.len());
9214 let g = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
9215 if crate::Engine::wpf_level() >= 1 {
9223 e.prefetch_weight_l2(&fa.wo)?;
9224 }
9225 if e.uses_q8_1_fast(&fa.wo) {
9228 let mut oq = e.alloc_i8_uninit(nh * hd)?;
9229 let mut od = e.zeros(nh * hd / 32)?;
9230 e.fa_decode_dc_q8(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
9231 &kvl.len_d, bucket, scale,
9232 kvl.k_tok_bytes, kvl.v_tok_bytes, g,
9233 Some((&mut oq, &mut od)))?;
9234 return Ok(e.matmul_pre(&fa.wo, &oq, &od, &attn, t)?);
9235 }
9236 e.fa_decode_dc(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
9237 &kvl.len_d, bucket, scale,
9238 kvl.k_tok_bytes, kvl.v_tok_bytes, g)?;
9239 return Ok(e.matmul(&fa.wo, &attn, t)?);
9240 }
9241 for i in 0..t {
9242 let avail = base_len + i + 1;
9243 let (off_tok, t_kv) = if swa && avail > win { (avail - win, win) } else { (0, avail) };
9244 let k_view = e.view_u8_range(&kvl.k, off_tok * kvl.k_tok_bytes,
9245 (off_tok + t_kv) * kvl.k_tok_bytes);
9246 let v_view = e.view_u8_range(&kvl.v, off_tok * kvl.v_tok_bytes,
9247 (off_tok + t_kv) * kvl.v_tok_bytes);
9248 let qv = e.view(&q, t * nh * hd);
9249 let q_row = qv.slice(i * nh * hd..(i + 1) * nh * hd);
9250 let mut q_one = e.uninit(nh * hd)?;
9251 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
9252 let mut a_one = e.uninit(nh * hd)?;
9253 e.fa_decode_kvmod(&q_one, &k_view, &v_view, &mut a_one, hd, nh, nkv, t_kv, scale,
9257 kvl.k_tok_bytes, kvl.v_tok_bytes,
9258 (!swa && crate::Engine::gkv_on())
9259 || (swa && crate::Engine::wkv_on()))?;
9260 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
9261 }
9262 Ok(e.matmul(&fa.wo, &attn, t)?)
9263 }
9264
9265 fn gemma4_e4b_trunk(&self, e: &Engine, tokens: &[u32], pos0: usize, cache: &mut Cache,
9270 head_last: bool)
9271 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
9272 let n_embd = self.cfg.n_embd as usize;
9273 let t = tokens.len();
9274 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
9275 let pos_d = e.htod_i32(&pos)?;
9276 let mut x = e.htod(&self.embd.gather(n_embd, tokens))?;
9277 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
9278 let inp_pl = self.gemma4_e4b_inp_pl(e, tokens, &x, t)?;
9279 self.gemma4_e4b_trunk_core(e, x, inp_pl, &pos_d, t, cache, None, true, head_last)
9280 }
9281
9282 fn gemma4_e4b_trunk_core(&self, e: &Engine, x_in: CudaSlice<f32>, inp_pl: CudaSlice<f32>,
9286 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache,
9287 dc_bucket: Option<usize>, cap_logits: bool, head_last: bool)
9288 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
9289 let n_embd = self.cfg.n_embd as usize;
9290 let eps = self.cfg.rms_eps;
9291 let n_layer = self.layers.len();
9292 let mut x = x_in;
9293 let aux_e4b = self.gemma4_aux.as_ref().unwrap().e4b.as_ref().unwrap();
9294 let n_epl = aux_e4b.n_epl;
9295
9296 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
9302 for il in 0..n_layer {
9303 let layer = &self.layers[il];
9304 let (hq, hdq) = match h_carry.take() {
9305 Some(p) => p,
9306 None => e.rms_norm_q8_1(&x, layer.attn_norm.float_data(), n_embd, t, eps)?,
9307 };
9308 let o = self.gemma4_e4b_attn(e, il, &hq, &hdq, pos_d, t, cache, dc_bucket)?;
9309 let bits = layer.gemma4.as_ref().unwrap();
9312 let e4b = bits.e4b.as_ref().expect("e4b layer bits");
9313 let fuse_exit = e.uses_q8_1_fast(&e4b.inp_gate);
9324 let (sn, attn_out) = self.gemma4_layer_tail_core_pn(
9325 e, layer, &o, &x, t, Some(layer.post_attn_norm.float_data()), fuse_exit)?;
9326 let mut resid = e.uninit(t * n_embd)?;
9327 let g = if fuse_exit {
9333 let (rq, rd) = e.rms_pre_add_q8_1(&sn, bits.post_ffw_norm.float_data(),
9335 &attn_out, &mut resid, n_embd, t,
9336 self.cfg.rms_eps)?;
9337 e.matmul_pre(&e4b.inp_gate, &rq, &rd, &resid, t)?
9338 } else {
9339 e.add(&sn, &attn_out, &mut resid, t * n_embd)?;
9340 e.matmul(&e4b.inp_gate, &resid, t)?
9341 };
9342 let mut act = e.uninit(t * n_epl)?;
9343 let y = if t == 1 && e.uses_q8_1_fast(&e4b.proj) {
9344 let ipv = e.view(&inp_pl, n_epl * n_layer);
9345 let row = ipv.slice(il * n_epl..(il + 1) * n_epl);
9346 let (aq, ad) = e.gelu_tanh_mul_q8_1(&g, &row, &mut act, n_epl, 1)?;
9347 e.matmul_pre(&e4b.proj, &aq, &ad, &act, t)?
9348 } else {
9349 let mut inp_this = e.uninit(t * n_epl)?;
9350 e.copy_rows_strided(&inp_pl, &mut inp_this, n_epl, t, n_epl * n_layer,
9351 il * n_epl)?;
9352 e.gelu_tanh_mul(&g, &inp_this, &mut act, t * n_epl)?;
9353 e.matmul(&e4b.proj, &act, t)?
9354 };
9355 let next_norm = if il + 1 < n_layer {
9358 self.layers[il + 1].attn_norm.float_data()
9359 } else {
9360 self.output_norm.float_data()
9361 };
9362 let mut xn = e.uninit(t * n_embd)?;
9363 let pair = e.rms_pre_add_scale_rms_norm_q8_1(&y, e4b.post_norm.float_data(),
9364 &resid, bits.layer_scale, next_norm,
9365 &mut xn, n_embd, t, eps)?;
9366 h_carry = Some(pair);
9367 x = xn;
9368 }
9369 let (oq, odq) = h_carry.take().unwrap();
9373 let h0 = e.zeros(0)?;
9374 let hm = if head_last { 1 } else { t };
9375 let (hq, hd) = if head_last && t > 1 {
9376 let mut q1 = e.uninit_i8(n_embd)?;
9377 e.dtod_copy_view_i8(&oq.slice((t - 1) * n_embd..t * n_embd), &mut q1)?;
9378 let nb = n_embd / 32;
9379 let mut d1 = e.uninit(nb)?;
9380 e.dtod_copy_view(&odq.slice((t - 1) * nb..t * nb), &mut d1)?;
9381 (q1, d1)
9382 } else {
9383 (oq, odq)
9384 };
9385 let mut ld = e.matmul_pre(&self.output, &hq, &hd, &h0, hm)?;
9386 if cap_logits {
9390 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
9391 e.softcap(&mut ld, cap, hm * self.output.out_features())?;
9392 }
9393 self.gemma4_suppress(e, &mut ld, hm)?; Ok((ld, x))
9395 }
9396
9397 pub fn gemma4_e4b_decode_step_t_am_dev(&self, e: &Engine, tok_d: &CudaSlice<u32>,
9404 t: usize, pos0: usize, cache: &mut Cache)
9405 -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
9406 let n_embd = self.cfg.n_embd as usize;
9407 let eps = self.cfg.rms_eps;
9408 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
9409 let pos_d = e.htod_i32(&pos)?;
9410 let embd_gpu = self.embd_gpu.get_or_init(|| {
9411 e.upload_u8(&self.embd.raw).expect("embed table upload")
9412 });
9413 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
9414 let mut x = e.embed_gather_device_td(embd_gpu, tok_d, t, n_embd, qt, rb)?;
9415 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
9416 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, tok_d, &x, t)?;
9417 let (ld, xp) = self.gemma4_e4b_trunk_core(e, x, inp_pl, &pos_d, t, cache, None, true,
9418 false)?;
9419 let n_vocab = self.output.out_features();
9422 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
9423 for i in 0..t {
9424 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
9425 }
9426 let mut hn = e.uninit(t * n_embd)?;
9427 e.rms_norm(&xp, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
9428 cache.pos += t;
9429 Ok((vam, hn))
9430 }
9431
9432 pub(crate) fn gemma4_e4b_decode_step_t_h(&self, e: &Engine, tokens: &[u32], pos0: usize,
9435 cache: &mut Cache)
9436 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
9437 let n_embd = self.cfg.n_embd as usize;
9438 let eps = self.cfg.rms_eps;
9439 let t = tokens.len();
9440 let (ld, xp) = self.gemma4_e4b_trunk(e, tokens, pos0, cache, false)?;
9441 let mut hn = e.uninit(t * n_embd)?;
9442 e.rms_norm(&xp, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
9443 cache.pos += t;
9444 Ok((e.dtoh(&ld)?, hn))
9445 }
9446
9447 pub fn gemma4_e4b_decode_step_dcg(&self, e: &Engine, token_d: &mut CudaSlice<u32>,
9453 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
9454 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
9455 n_vocab: usize, bucket: usize)
9456 -> Result<(), Box<dyn std::error::Error>> {
9457 let n_embd = self.cfg.n_embd as usize;
9458 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
9459 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
9460 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, token_d, &x, 1)?;
9461 let (ld, _x) = self.gemma4_e4b_trunk_core(e, x, inp_pl, pos_d, 1, cache, Some(bucket),
9462 false, false)?;
9463 e.argmax_token_device_into(&ld, token_d, n_vocab)?;
9464 e.inc_seqlen(pos_d)?;
9465 Ok(())
9466 }
9467
9468 #[allow(clippy::too_many_arguments)]
9476 pub fn gemma4_e4b_decode_step_dc(&self, e: &Engine, token_d: &CudaSlice<u32>,
9477 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
9478 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
9479 n_vocab: usize)
9480 -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
9481 let n_embd = self.cfg.n_embd as usize;
9482 let eps = self.cfg.rms_eps;
9483 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
9484 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
9485 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, token_d, &x, 1)?;
9486 let (ld, _x) = self.gemma4_e4b_trunk_core(e, x, inp_pl, pos_d, 1, cache, None, false,
9487 false)?;
9488 let mut tok_out = e.stream().alloc_zeros::<u32>(1)?;
9489 e.argmax_token_device_into(&ld, &mut tok_out, n_vocab)?;
9490 e.inc_seqlen(pos_d)?;
9491 cache.pos += 1;
9492 let _ = eps;
9493 Ok(tok_out)
9494 }
9495
9496 pub(crate) fn gemma4_e4b_decode_step_h(&self, e: &Engine, token: u32, cache: &mut Cache)
9499 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
9500 let (ld, x) = self.gemma4_e4b_trunk(e, &[token], cache.pos, cache, false)?;
9501 let logits = e.dtoh(&ld)?;
9502 cache.pos += 1;
9503 Ok((logits, x))
9504 }
9505
9506 pub(crate) fn gemma4_e4b_prime(&self, e: &Engine, tokens: &[u32], cache: &mut Cache)
9510 -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
9511 if cache.pos != 0 {
9514 return Err("e4b prime is fresh-prompt only (v0) — prime the full prompt in one \
9515 call or decode tokenwise".into());
9516 }
9517 let n_embd = self.cfg.n_embd as usize;
9518 let t = tokens.len();
9519 let (ld, x) = self.gemma4_e4b_trunk(e, tokens, 0, cache, true)?;
9520 cache.pos += t;
9521 let last = e.dtoh(&ld)?; let xv = e.view(&x, t * n_embd);
9523 let row = xv.slice((t - 1) * n_embd..t * n_embd);
9524 let mut h_seed = e.uninit(n_embd)?;
9525 e.copy_view_into(&mut h_seed, 0, &row, n_embd)?;
9526 Ok((last, h_seed, x))
9527 }
9528
9529 pub(crate) fn gemma4_e4b_forward(&self, e: &Engine, tokens: &[u32], last_only: bool)
9531 -> Result<Vec<f32>, Box<dyn std::error::Error>> {
9532 let mut cache = Cache::new(e, &self.cfg, tokens.len() + 8)?;
9533 let (ld, _x) = self.gemma4_e4b_trunk(e, tokens, 0, &mut cache, last_only)?;
9534 Ok(e.dtoh(&ld)?) }
9536}
9537
9538#[cfg(test)]
9539mod prime_chunk_schedule_tests {
9540 use super::{
9541 dynamic_prime_chunk_ranges, fixed_prime_chunk_ranges, fixed_prime_chunk_ranges_for_ring,
9542 PRIME_MIN_T,
9543 PRIME_PIPE_MIN_CHUNK,
9544 };
9545
9546 fn sizes(ranges: &[(usize, usize)]) -> Vec<usize> {
9547 ranges.iter().map(|(start, end)| end - start).collect()
9548 }
9549
9550 fn auto_chunk(t: usize) -> usize {
9551 t.div_ceil(8).max(PRIME_PIPE_MIN_CHUNK).min(4096)
9552 }
9553
9554 #[test]
9555 fn fixed_schedule_retains_measured_geometry() {
9556 assert_eq!(
9557 sizes(&fixed_prime_chunk_ranges(461, 128)),
9558 vec![128, 128, 128, 77]
9559 );
9560 assert_eq!(
9561 sizes(&fixed_prime_chunk_ranges(1833, 230)),
9562 vec![230, 230, 230, 230, 230, 230, 230, 223]
9563 );
9564 assert_eq!(
9565 sizes(&fixed_prime_chunk_ranges(4096, 512)),
9566 vec![512; 8]
9567 );
9568 let capped = sizes(&fixed_prime_chunk_ranges_for_ring(8200, 4096, true));
9569 assert_eq!(capped, vec![4096, 4088, 16]);
9570 assert!(capped.iter().all(|&rows| rows <= 4096));
9571 assert_eq!(
9572 sizes(&fixed_prime_chunk_ranges_for_ring(4100, 4096, false)),
9573 vec![4100],
9574 "flag-off schedule remains byte-for-byte the legacy monolithic tail",
9575 );
9576 }
9577
9578 #[test]
9579 fn dynamic_schedule_matches_registered_shapes() {
9580 let cases = [
9581 (461, vec![64, 141, 132, 124]),
9582 (1833, vec![115, 269, 260, 252, 244, 237, 231, 225]),
9583 (4096, vec![256, 602, 580, 563, 545, 531, 516, 503]),
9584 ];
9585 for (t, expected) in cases {
9586 let chunk = auto_chunk(t);
9587 let fixed = fixed_prime_chunk_ranges(t, chunk);
9588 assert_eq!(
9589 sizes(&dynamic_prime_chunk_ranges(t, chunk, &fixed)),
9590 expected
9591 );
9592 }
9593 }
9594
9595 #[test]
9596 fn dynamic_schedule_covers_exactly_and_shrinks_after_fill() {
9597 for t in 256..=8192 {
9598 let chunk = auto_chunk(t);
9599 let fixed = fixed_prime_chunk_ranges(t, chunk);
9600 let dynamic = dynamic_prime_chunk_ranges(t, chunk, &fixed);
9601 assert_eq!(dynamic.len(), fixed.len(), "T={t}");
9602 assert_eq!(dynamic.first().unwrap().0, 0, "T={t}");
9603 assert_eq!(dynamic.last().unwrap().1, t, "T={t}");
9604 for pair in dynamic.windows(2) {
9605 assert_eq!(pair[0].1, pair[1].0, "T={t}");
9606 }
9607 assert!(
9608 dynamic
9609 .iter()
9610 .all(|(start, end)| end - start >= PRIME_MIN_T),
9611 "T={t} sizes={:?}",
9612 sizes(&dynamic)
9613 );
9614 if dynamic.len() >= 3 {
9615 let chunk_sizes = sizes(&dynamic);
9616 assert!(
9617 chunk_sizes[0] < chunk_sizes[1],
9618 "T={t} sizes={chunk_sizes:?}"
9619 );
9620 assert!(
9621 chunk_sizes[1..].windows(2).all(|pair| pair[0] >= pair[1]),
9622 "T={t} sizes={chunk_sizes:?}"
9623 );
9624 }
9625 }
9626 }
9627}
9628
9629#[cfg(test)]
9630mod page_prefetch_tests {
9631 use super::{
9632 grouped_worker_prefetch_position, page_prefetch_positions,
9633 page_prefetch_window_from_values, worker_prefetch_positions,
9634 };
9635
9636 #[test]
9637 fn page_prefetch_window_keeps_existing_opt_in_default() {
9638 assert_eq!(page_prefetch_window_from_values(false, None), 0);
9639 assert_eq!(page_prefetch_window_from_values(false, Some("8")), 0);
9640 assert_eq!(page_prefetch_window_from_values(true, None), 1);
9641 assert_eq!(page_prefetch_window_from_values(true, Some("bad")), 1);
9642 assert_eq!(page_prefetch_window_from_values(true, Some("0")), 0);
9643 assert_eq!(page_prefetch_window_from_values(true, Some("8")), 8);
9644 }
9645
9646 #[test]
9647 fn rolling_page_prefetch_advises_each_future_expert_once() {
9648 let advised: Vec<_> = (0..7)
9649 .flat_map(|position| page_prefetch_positions(position, 7, 3))
9650 .collect();
9651 assert_eq!(advised, vec![1, 2, 3, 4, 5, 6]);
9652
9653 let one_ahead: Vec<_> = (0..4)
9654 .flat_map(|position| page_prefetch_positions(position, 4, 1))
9655 .collect();
9656 assert_eq!(one_ahead, vec![1, 2, 3]);
9657 assert!(page_prefetch_positions(0, 4, 0).is_empty());
9658 }
9659
9660 #[test]
9661 fn grouped_worker_prefetch_primes_first_then_each_known_next_once() {
9662 assert_eq!(grouped_worker_prefetch_position(0, None), None);
9663 let positions: Vec<_> = std::iter::once(grouped_worker_prefetch_position(4, None).unwrap())
9664 .chain((0..4).filter_map(|position| {
9665 grouped_worker_prefetch_position(4, Some(position))
9666 }))
9667 .collect();
9668 assert_eq!(positions, vec![0, 1, 2, 3]);
9669 assert_eq!(grouped_worker_prefetch_position(1, Some(0)), None);
9670 }
9671
9672 #[test]
9673 fn rolling_worker_prefetch_primes_current_and_each_future_expert_once() {
9674 let queued: Vec<_> = (0..8)
9675 .flat_map(|position| worker_prefetch_positions(position, 8, 5))
9676 .collect();
9677 assert_eq!(queued, (0..8).collect::<Vec<_>>());
9678
9679 let one_at_a_time: Vec<_> = (0..4)
9680 .flat_map(|position| worker_prefetch_positions(position, 4, 1))
9681 .collect();
9682 assert_eq!(one_at_a_time, vec![0, 1, 2, 3]);
9683 assert!(worker_prefetch_positions(0, 4, 0).is_empty());
9684 }
9685}
9686
9687pub struct G4DcSlots {
9688 x: CudaSlice<f32>, xn: CudaSlice<f32>, cur: CudaSlice<f32>,
9689 hq: CudaSlice<i8>, hd_: CudaSlice<f32>,
9690 q0: CudaSlice<f32>, k0: CudaSlice<f32>, v0: CudaSlice<f32>,
9691 q: CudaSlice<f32>, k: CudaSlice<f32>, v: CudaSlice<f32>,
9692 attn: CudaSlice<f32>, o: CudaSlice<f32>,
9693 attn_out: CudaSlice<f32>, zsh: CudaSlice<f32>,
9694 zq: CudaSlice<i8>, zd: CudaSlice<f32>,
9695 gate: CudaSlice<f32>, up: CudaSlice<f32>,
9696 act: CudaSlice<f32>, actq: CudaSlice<i8>, actd: CudaSlice<f32>,
9697 f0: CudaSlice<f32>, sn: CudaSlice<f32>,
9698 hn: CudaSlice<f32>, logits: CudaSlice<f32>,
9699}