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
40
41pub(crate) struct AttnPre {
43 pub q: cudarc::driver::CudaSlice<f32>,
44 pub k: cudarc::driver::CudaSlice<f32>,
45 pub v: cudarc::driver::CudaSlice<f32>,
46 pub gate: Option<cudarc::driver::CudaSlice<f32>>,
47}
48
49pub(crate) struct GdnPrep {
51 pub hk: usize,
52 pub q_l2: cudarc::driver::CudaSlice<f32>,
53 pub k_l2: cudarc::driver::CudaSlice<f32>,
54 pub v_g: cudarc::driver::CudaSlice<f32>,
55 pub beta: cudarc::driver::CudaSlice<f32>,
56 pub g_log: cudarc::driver::CudaSlice<f32>,
57 pub kb16: Option<cudarc::driver::CudaSlice<u8>>,
58 pub qb16: Option<cudarc::driver::CudaSlice<u8>>,
59}
60
61pub(crate) struct VerifyStreamScratch {
63 pub pos_d: CudaSlice<i32>,
64 pub row_ctrs: Vec<CudaSlice<i32>>,
65}
66use crate::hybrid::{HybridModel, Mixer, FullAttnLayer, LinearAttnLayer, MoeWeights};
67
68struct MoeInputTraceWriter {
69 dir: std::path::PathBuf,
70 index: std::fs::File,
71 payloads: std::collections::HashMap<u16, (std::fs::File, u64)>,
72}
73
74static MOE_INPUT_TRACE_WRITER: std::sync::OnceLock<
75 std::sync::Mutex<Option<MoeInputTraceWriter>>,
76> = std::sync::OnceLock::new();
77
78fn gdec_enabled() -> bool {
81 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
82 *E.get_or_init(|| std::env::var("MEMRA_MOE_GDEC").map(|v| v != "0").unwrap_or(true))
83}
84
85fn moe_prefetch_enabled() -> bool {
88 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
89 *E.get_or_init(|| std::env::var("MEMRA_MOE_PREFETCH").as_deref() == Ok("1")
90 || crate::spill_pread::worker_enabled())
91}
92
93fn moe_page_prefetch_window() -> usize {
98 static W: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
99 *W.get_or_init(|| page_prefetch_window_from_values(
100 std::env::var("MEMRA_MOE_PAGE_PREFETCH").as_deref() == Ok("1"),
101 std::env::var("MEMRA_MOE_PAGE_PREFETCH_WINDOW").ok().as_deref(),
102 ))
103}
104
105fn page_prefetch_window_from_values(enabled: bool, raw_window: Option<&str>) -> usize {
106 if !enabled {
107 return 0;
108 }
109 raw_window
110 .and_then(|value| value.parse().ok())
111 .unwrap_or(1)
112}
113
114fn page_prefetch_positions(
118 position: usize,
119 len: usize,
120 window: usize,
121) -> std::ops::Range<usize> {
122 if window == 0 || position >= len {
123 return len..len;
124 }
125 let (start, count) = if position == 0 {
126 (1, window)
127 } else {
128 (position.saturating_add(window), 1)
129 };
130 let start = start.min(len);
131 start..start.saturating_add(count).min(len)
132}
133
134fn grouped_worker_prefetch_position(order_len: usize, current: Option<usize>) -> Option<usize> {
137 let position = current.map_or(0, |position| position.saturating_add(1));
138 (position < order_len).then_some(position)
139}
140
141fn worker_prefetch_window() -> usize {
146 static WINDOW: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
147 *WINDOW.get_or_init(|| {
148 let automatic = crate::spill_pread::configured_depth().saturating_sub(1) / 3;
149 std::env::var("MEMRA_SPILL_WORKER_EXPERT_WINDOW")
150 .ok()
151 .and_then(|value| value.parse::<usize>().ok())
152 .unwrap_or(automatic.max(1))
153 })
154}
155
156fn worker_prefetch_positions(position: usize, len: usize, window: usize) -> std::ops::Range<usize> {
160 if window == 0 || position >= len {
161 return len..len;
162 }
163 let (start, count) = if position == 0 {
164 (0, window)
165 } else {
166 (position.saturating_add(window).saturating_sub(1), 1)
167 };
168 let start = start.min(len);
169 start..start.saturating_add(count).min(len)
170}
171
172fn moe_dev_enabled() -> bool {
177 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
178 *E.get_or_init(|| std::env::var("MEMRA_MOE_DEV").map(|v| v != "0").unwrap_or(true)
179 && !matches!(std::env::var("MEMRA_FUSED_ROUTER").as_deref(), Ok("0")))
180}
181
182fn moe_q8_enabled() -> bool {
187 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
188 *E.get_or_init(|| std::env::var("MEMRA_MOE_Q8").map(|v| v != "0").unwrap_or(true))
189}
190
191fn expert_dp4a_supported(qt: i32) -> bool {
194 qt == crate::QT_Q4_0 || qt == crate::QT_IQ3_S || qt == crate::QT_IQ4_XS
195 || qt == crate::QT_Q3_K || qt == crate::QT_Q4_K || qt == crate::QT_Q6_K
196}
197
198fn q8_expert_supported(qt: i32) -> bool {
199 static KQ: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
205 let kq = *KQ.get_or_init(|| {
206 std::env::var("MEMRA_MOE_Q8_KQ").map(|v| v != "0").unwrap_or(true)
207 });
208 let nvfp4_q8 = std::env::var("MEMRA_MOE_Q8_NVFP4").map(|v| v != "0").unwrap_or(true);
215 qt == crate::QT_IQ3_S || qt == crate::QT_IQ4_XS || (nvfp4_q8 && qt == crate::QT_NVFP4)
216 || (kq && (qt == crate::QT_Q3_K || qt == crate::QT_Q4_K || qt == crate::QT_Q6_K))
217}
218
219fn q8_expert_dec_supported(qt: i32) -> bool {
222 qt == crate::QT_IQ3_S || qt == crate::QT_IQ4_XS || qt == crate::QT_Q4_0
223}
224
225fn f16g_proj_ok(qt: i32, in_f: usize) -> bool {
231 match qt {
232 crate::QT_Q4_0 => in_f % 32 == 0,
233 crate::QT_IQ4_XS | crate::QT_IQ3_S | crate::QT_Q3_K | crate::QT_Q4_K
234 | crate::QT_Q6_K => in_f % 256 == 0,
235 _ => false,
236 }
237}
238
239fn moe_prewarm_enabled() -> bool {
242 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
243 *E.get_or_init(|| std::env::var("MEMRA_MOE_PREWARM").map(|v| v != "0").unwrap_or(true))
244}
245
246fn cpu_expert_profile_admit_enabled() -> bool {
250 static E: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
251 *E.get_or_init(|| std::env::var("MEMRA_CPU_EXPERT_FREEZE_PROFILE_ADMIT").as_deref() == Ok("1"))
252}
253
254pub const PRIME_MIN_T: usize = 16;
258
259impl HybridModel {
260 fn prime_trace_path() -> Option<&'static str> {
265 static P: std::sync::OnceLock<Option<String>> = std::sync::OnceLock::new();
266 P.get_or_init(|| std::env::var("MEMRA_PRIME_TRACE").ok())
267 .as_deref()
268 }
269
270 pub fn forward(&self, e: &Engine, tokens: &[u32]) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
272 if self.is_gemma4_e4b() { return self.gemma4_e4b_forward(e, tokens, false); }
273 if self.cfg.gemma4.is_some() { return self.gemma4_forward(e, tokens, false); }
274 let cfg = &self.cfg;
275 let n_embd = cfg.n_embd as usize;
276 let t = tokens.len();
277 let eps = cfg.rms_eps;
278 let pos: Vec<i32> = (0..t as i32).collect();
279 let pos_d = e.htod_i32(&pos)?;
280
281 let mut x = self.embed(e, tokens)?; for (il, layer) in self.layers.iter().enumerate() {
284 let mut h = e.uninit(t * n_embd)?;
286 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
287
288 let mixed = match &layer.mixer {
289 Mixer::Full(fa) => self.full_attn(e, fa, &h, &pos_d, t)?,
290 Mixer::Linear(la) => self.linear_attn(e, la, &h, t)?,
291 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
292 };
293
294 let mut x1 = e.uninit(t * n_embd)?;
296 e.add(&x, &mixed, &mut x1, t * n_embd)?;
297
298 let mut z = e.uninit(t * n_embd)?;
300 e.rms_norm(&x1, layer.post_attn_norm.float_data(), &mut z, n_embd, t, eps)?;
301 let ffn_out = match &layer.ffn {
302 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
303 let n_ff = ffn_gate.out_features();
304 let mut g2 = e.matmul_group(&[ffn_gate, ffn_up], &z, t)?;
305 let up = g2.pop().unwrap();
306 let gate = g2.pop().unwrap();
307 let mut act = e.uninit(t * n_ff)?;
308 Self::ffn_act(e, &self.cfg, &gate, &up, &mut act, t * n_ff)?;
309 e.matmul(ffn_down, &act, t)?
310 }
311 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il(e, m, &z, t, il as u16)?,
312 };
313 let mut x2 = e.uninit(t * n_embd)?;
314 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
315 x = x2;
316 }
317
318 let mut hn = e.uninit(t * n_embd)?;
319 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
320 let logits = e.matmul(&self.output, &hn, t)?;
321 Ok(e.dtoh(&logits)?)
322 }
323
324 pub fn forward_last(&self, e: &Engine, tokens: &[u32]) -> Result<Vec<f32>, Box<dyn std::error::Error>> {
330 if self.cfg.gemma4.is_some() { return self.gemma4_forward(e, tokens, true); }
331 let cfg = &self.cfg;
332 let n_embd = cfg.n_embd as usize;
333 let t = tokens.len();
334 let eps = cfg.rms_eps;
335 let pos: Vec<i32> = (0..t as i32).collect();
336 let pos_d = e.htod_i32(&pos)?;
337
338 let mut x = self.embed(e, tokens)?; let probe = std::env::var("MEMRA_LAYER_PROBE").is_ok();
342 for (il, layer) in self.layers.iter().enumerate() {
343 let mut h = e.uninit(t * n_embd)?;
344 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
345 if probe { e.stream().synchronize()?; eprintln!("[probe] L{il} norm ok"); }
346 let mixed = match &layer.mixer {
347 Mixer::Full(fa) => self.full_attn(e, fa, &h, &pos_d, t)?,
348 Mixer::Linear(la) => self.linear_attn(e, la, &h, t)?,
349 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
350 };
351 if probe { e.stream().synchronize()?; eprintln!("[probe] L{il} mixer ok"); }
352 let mut x1 = e.uninit(t * n_embd)?;
353 e.add(&x, &mixed, &mut x1, t * n_embd)?;
354 let mut z = e.uninit(t * n_embd)?;
355 e.rms_norm(&x1, layer.post_attn_norm.float_data(), &mut z, n_embd, t, eps)?;
356 let ffn_out = match &layer.ffn {
357 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
358 let n_ff = ffn_gate.out_features();
359 let mut g2 = e.matmul_group(&[ffn_gate, ffn_up], &z, t)?;
360 let up = g2.pop().unwrap();
361 let gate = g2.pop().unwrap();
362 let mut act = e.uninit(t * n_ff)?;
363 Self::ffn_act(e, &self.cfg, &gate, &up, &mut act, t * n_ff)?;
364 e.matmul(ffn_down, &act, t)?
365 }
366 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il(e, m, &z, t, il as u16)?,
367 };
368 if probe { e.stream().synchronize()?; eprintln!("[probe] L{il} ffn ok"); }
369 let mut x2 = e.uninit(t * n_embd)?;
370 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
371 x = x2;
372 }
373 let mut hn = e.uninit(t * n_embd)?;
375 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
376 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)?;
379 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
380 let logits = e.matmul(&self.output, &hlast, 1)?; Ok(e.dtoh(&logits)?)
382 }
383
384 pub fn prime_cache(&self, e: &Engine, tokens: &[u32], cache: &mut Cache)
401 -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
402 let n_embd = self.cfg.n_embd as usize;
403 let t = tokens.len();
404 assert!(t >= PRIME_MIN_T, "prime_cache needs T >= {PRIME_MIN_T} (caller gates)");
408 assert!(cache.pos + t <= cache.max_ctx, "prime_cache: prompt exceeds cache max_ctx");
409
410 if self.is_gemma4_e4b() {
422 return self.gemma4_e4b_prime(e, tokens, cache);
423 }
424 if self.cfg.gemma4.is_some() {
425 return self.gemma4_prime(e, tokens, cache);
427 }
428 let chunk: usize = std::env::var("MEMRA_PRIME_CHUNK").ok()
429 .and_then(|v| v.parse().ok()).unwrap_or(4096);
430 if chunk == 0 || t <= chunk {
455 return self.prime_chunk(e, tokens, cache);
456 }
457 let mut hiddens = e.uninit(t * n_embd)?;
458 let mut last: Option<(Vec<f32>, CudaSlice<f32>)> = None;
459 let mut start = 0usize;
460 while start < t {
461 let mut end = (start + chunk).min(t);
463 if t - end > 0 && t - end < PRIME_MIN_T { end = t; }
464 let (l, hs, x) = self.prime_chunk(e, &tokens[start..end], cache)?;
465 e.copy_into(&mut hiddens, start * n_embd, &x, (end - start) * n_embd)?;
466 last = Some((l, hs));
467 start = end;
468 }
469 let (logits, h_seed) = last.unwrap();
470 Ok((logits, h_seed, hiddens))
471 }
472
473 fn gdn_hk(e: &Engine, t: usize, num_v: usize, num_k: usize) -> usize {
480 if Engine::gdn_db_on()
481 && Engine::gdn_chunked_enabled() && t >= 16
482 && e.gdn_mma_enabled(Engine::gdn_chunk_size())
483 && num_k * 2 == num_v
484 {
485 num_k
486 } else {
487 num_v
488 }
489 }
490
491 fn f16out_on(e: &Engine, t: usize) -> bool {
496 crate::f16_ffi::pp_f16_enabled() && t >= 16 && !e.verify_exact_on()
497 && std::env::var("MEMRA_F16OUT").as_deref() != Ok("0")
498 }
499
500 pub fn prime_slabs_get(&self, e: &Engine, t: usize, n_embd: usize, n_ff_max: usize)
503 -> Result<std::sync::MutexGuard<'_, Option<PrimeSlabs>>, Box<dyn std::error::Error>> {
504 let mut g = self.prime_slabs.lock().unwrap();
505 let need_new = match g.as_ref() { None => true, Some(sl) => sl.t_cap < t };
506 if need_new {
507 *g = Some(PrimeSlabs {
508 t_cap: t,
509 h: e.uninit(t * n_embd)?,
510 x1: e.uninit(t * n_embd)?,
511 z: e.uninit(t * n_embd)?,
512 act: e.uninit(t * n_ff_max)?,
513 xa: e.uninit(t * n_embd)?,
514 xb: e.uninit(t * n_embd)?,
515 h16: e.alloc_u8_uninit(t * n_embd * 2)?,
516 z16: e.alloc_u8_uninit(t * n_embd * 2)?,
517 gate: e.uninit(t * n_ff_max)?,
518 up: e.uninit(t * n_ff_max)?,
519 ffn_out: e.uninit(t * n_embd)?,
520 seg_glue: Vec::new(),
521 mixed: e.uninit(t * n_embd)?,
522 seg_mid: Vec::new(),
523 seg_t: 0,
524 });
525 }
526 Ok(g)
527 }
528
529 fn prime_chunk(&self, e: &Engine, tokens: &[u32], cache: &mut Cache)
530 -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
531 let cfg = &self.cfg;
532 let n_embd = cfg.n_embd as usize;
533 let t = tokens.len();
534 let eps = cfg.rms_eps;
535 let base = cache.pos;
536 let pos: Vec<i32> = (base as i32..(base + t) as i32).collect();
537 let pos_d = e.htod_i32(&pos)?;
538
539 let x_embed = self.embed(e, tokens)?; let f16fuse = crate::f16_ffi::pp_f16_enabled() && t >= 16;
544 let n_ff_max = self.layers.iter().map(|l| match &l.ffn {
550 crate::hybrid::Ffn::Dense { ffn_gate, .. } => ffn_gate.out_features(),
551 _ => n_embd,
552 }).max().unwrap_or(n_embd).max(n_embd);
553 let use_slabs = std::env::var("MEMRA_PRIME_SLABS").as_deref() != Ok("0");
554 let mut slab_guard = if use_slabs {
555 Some(self.prime_slabs_get(e, t, n_embd, n_ff_max)?)
556 } else {
557 None
558 };
559 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>);
561 let (mut x_cur, mut x_nxt, sl): (&mut CudaSlice<f32>, &mut CudaSlice<f32>, Option<SlabRefs>);
562 let mut seg: Option<(&mut Vec<Option<cudarc::driver::CudaGraph>>, &mut Vec<Option<cudarc::driver::CudaGraph>>, &mut CudaSlice<f32>, &mut usize)> = None;
563 let mut x_own2;
564 match slab_guard.as_mut() {
565 Some(g) => {
566 let slabs = g.as_mut().unwrap();
567 e.copy_into(&mut slabs.xa, 0, &x_embed, t * n_embd)?;
568 let PrimeSlabs { xa, xb, h, x1, z, act, h16, z16, gate, up, ffn_out, seg_glue, mixed, seg_mid, seg_t, .. } = slabs;
569 x_cur = xa;
570 x_nxt = xb;
571 seg = Some((seg_glue, seg_mid, mixed, seg_t));
572 sl = Some((h, x1, z, act, h16, z16, gate, up, ffn_out));
573 }
574 None => {
575 x_own = x_embed;
576 x_own2 = e.uninit(t * n_embd)?;
577 x_cur = &mut x_own;
578 x_nxt = &mut x_own2;
579 sl = None;
580 }
581 }
582 let mut alloc_h; let mut alloc_x1; let mut alloc_z; let mut alloc_act;
583 let mut alloc_h16; let mut alloc_z16;
584 let mut alloc_gate; let mut alloc_up; let mut alloc_fo;
585 let (h, x1, z, act): (&mut CudaSlice<f32>, &mut CudaSlice<f32>, &mut CudaSlice<f32>, &mut CudaSlice<f32>);
586 let (h16, z16): (&mut CudaSlice<u8>, &mut CudaSlice<u8>);
587 let (sl_gate, sl_up, sl_fo): (&mut CudaSlice<f32>, &mut CudaSlice<f32>, &mut CudaSlice<f32>);
588 match sl {
589 Some((a, b, c, d, e16, f16b, g, u, fo)) => {
590 h = a; x1 = b; z = c; act = d; h16 = e16; z16 = f16b;
591 sl_gate = g; sl_up = u; sl_fo = fo;
592 }
593 None => {
594 alloc_h = e.uninit(t * n_embd)?;
595 alloc_x1 = e.uninit(t * n_embd)?;
596 alloc_z = e.uninit(t * n_embd)?;
597 alloc_act = e.uninit(t * n_ff_max)?;
598 alloc_h16 = e.alloc_u8_uninit(t * n_embd * 2)?;
599 alloc_z16 = e.alloc_u8_uninit(t * n_embd * 2)?;
600 alloc_gate = e.uninit(t * n_ff_max)?;
601 alloc_up = e.uninit(t * n_ff_max)?;
602 alloc_fo = e.uninit(t * n_embd)?;
603 h = &mut alloc_h; x1 = &mut alloc_x1; z = &mut alloc_z; act = &mut alloc_act;
604 h16 = &mut alloc_h16; z16 = &mut alloc_z16;
605 sl_gate = &mut alloc_gate; sl_up = &mut alloc_up; sl_fo = &mut alloc_fo;
606 }
607 }
608 let n_layers = self.layers.len();
613 let use_seg = f16fuse && seg.is_some()
620 && std::env::var("MEMRA_PRIME_SEG").as_deref() == Ok("1");
621 if let Some((sg, sm, _, st)) = seg.as_mut() {
622 if **st != t {
623 sg.clear();
624 sg.extend((0..n_layers).map(|_| None));
625 sm.clear();
626 sm.extend((0..n_layers).map(|_| None));
627 **st = t;
628 }
629 }
630 {
631 let layer0 = &self.layers[0];
632 if f16fuse {
633 e.rms_norm_f16out(x_cur, layer0.attn_norm.float_data(), h, h16, n_embd, t, eps)?;
634 } else {
635 e.rms_norm(x_cur, layer0.attn_norm.float_data(), h, n_embd, t, eps)?;
636 }
637 }
638 for (il, layer) in self.layers.iter().enumerate() {
639 let hx16 = if f16fuse { Some(&*h16) } else { None };
640 if use_seg {
641 let (pre, pre16, w_out) = match &layer.mixer {
644 Mixer::Full(fa) => {
645 let g3 = match hx16 {
646 Some(xh) => e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], h, xh, t)?,
647 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
648 };
649 let (pre, pre16) = self.full_attn_prime_core_inner(e, fa, g3, &pos_d, t, cache, il)?;
650 (pre, pre16, &fa.wo)
651 }
652 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
653 Mixer::Linear(la) => {
654 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
655 let g4 = match hx16 {
656 Some(xh) => e.matmul_group_xh(&ws, h, xh, t)?,
657 None => e.matmul_group(&ws, h, t)?,
658 };
659 let (pre, pre16) = self.linear_attn_prime_core_pad_inner(e, la, g4, t, cache, il, None)?;
660 (pre, pre16, &la.ssm_out)
661 }
662 };
663 {
664 let (_, sm, mslab, _) = seg.as_mut().unwrap();
665 let pre_n = pre.len() / t;
666 let xh_pre = match pre16 {
667 Some(x) => x,
668 None => e.f16_act(&pre, t * pre_n, pre_n)?,
669 };
670 if !e.try_f16_gemm_pre_into(w_out, &xh_pre, t, mslab)? {
671 let y = e.matmul(w_out, &pre, t)?;
672 e.copy_into(mslab, 0, &y, t * n_embd)?;
673 }
674 if sm[il].is_none() {
675 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
676 let w_post = layer.post_attn_norm.float_data();
677 e.stream().synchronize()?;
678 e.stream().begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
679 let r = (|| -> Result<(), Box<dyn std::error::Error>> {
680 e.add(x_cur, mslab, x1, t * n_embd)?;
681 e.rms_norm_f16out(x1, w_post, z, z16, n_embd, t, eps)?;
682 Ok(())
683 })();
684 let g = e.stream().end_capture(
685 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH);
686 r?;
687 sm[il] = Some(g?.ok_or("S-mid capture produced no graph")?);
688 }
689 sm[il].as_ref().unwrap().launch()?;
690 }
691 } else {
692 let mixed = match &layer.mixer {
693 Mixer::Full(fa) => self.full_attn_prime(e, fa, h, hx16, &pos_d, t, cache, il)?,
694 Mixer::Linear(la) => self.linear_attn_prime(e, la, h, hx16, t, cache, il)?,
695 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
696 };
697 if f16fuse {
698 e.add_rms_norm_f16out(x_cur, &mixed, layer.post_attn_norm.float_data(),
701 x1, z, z16, n_embd, t, eps)?;
702 } else {
703 e.add(x_cur, &mixed, x1, t * n_embd)?;
704 e.rms_norm(x1, layer.post_attn_norm.float_data(), z, n_embd, t, eps)?;
705 }
706 }
707 let zx16 = if f16fuse { Some(&*z16) } else { None };
708 match &layer.ffn {
709 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
710 let n_ff = ffn_gate.out_features();
711 let mut into_ok = false;
714 if let Some(xh) = zx16 {
715 into_ok = e.try_f16_gemm_pre_into(ffn_gate, xh, t, sl_gate)?
716 && e.try_f16_gemm_pre_into(ffn_up, xh, t, sl_up)?;
717 }
718 if !into_ok {
719 let mut g2 = match zx16 {
720 Some(xh) => e.matmul_group_xh(&[ffn_gate, ffn_up], z, xh, t)?,
721 None => e.matmul_group(&[ffn_gate, ffn_up], z, t)?,
722 };
723 let up_y = g2.pop().unwrap();
724 let gate_y = g2.pop().unwrap();
725 e.copy_into(sl_gate, 0, &gate_y, t * n_ff)?;
726 e.copy_into(sl_up, 0, &up_y, t * n_ff)?;
727 }
728 let act16 = if Self::f16out_on(e, t) && self.cfg.m3.is_none() {
731 let mut a16 = e.alloc_u8_uninit(t * n_ff * 2)?;
732 e.silu_mul_f16out(sl_gate, sl_up, act, &mut a16, t * n_ff)?;
733 Some(a16)
734 } else {
735 Self::ffn_act(e, &self.cfg, sl_gate, sl_up, act, t * n_ff)?;
736 None
737 };
738 let xh_act = match act16 {
740 Some(x) => x,
741 None => e.f16_act(act, t * n_ff, n_ff)?,
742 };
743 if !e.try_f16_gemm_pre_into(ffn_down, &xh_act, t, sl_fo)? {
744 let y = e.matmul(ffn_down, &*act, t)?;
745 e.copy_into(sl_fo, 0, &y, t * n_embd)?;
746 }
747 }
748 crate::hybrid::Ffn::Moe(m) => {
749 let y = self.moe_ffn_il(e, m, z, t, il as u16)?;
750 e.copy_into(sl_fo, 0, &y, t * n_embd)?;
751 }
752 }
753 if use_seg && il + 1 < n_layers {
754 let w_next = self.layers[il + 1].attn_norm.float_data();
756 let (sg, _, _, _) = seg.as_mut().unwrap();
757 if sg[il].is_none() {
758 use cudarc::driver::sys::{CUgraphInstantiate_flags, CUstreamCaptureMode};
759 e.stream().synchronize()?;
760 e.stream().begin_capture(CUstreamCaptureMode::CU_STREAM_CAPTURE_MODE_RELAXED)?;
761 let r = (|| -> Result<(), Box<dyn std::error::Error>> {
762 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
763 e.rms_norm_f16out(x_nxt, w_next, h, h16, n_embd, t, eps)?;
764 Ok(())
765 })();
766 let g = e.stream().end_capture(
767 CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_AUTO_FREE_ON_LAUNCH);
768 r?;
769 sg[il] = Some(g?.ok_or("S-glue capture produced no graph")?);
770 }
771 sg[il].as_ref().unwrap().launch()?;
772 } else {
773 if il + 1 < n_layers {
774 let w_next = self.layers[il + 1].attn_norm.float_data();
775 if f16fuse {
776 e.add_rms_norm_f16out(x1, sl_fo, w_next, x_nxt, h, h16, n_embd, t, eps)?;
777 } else {
778 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
779 e.rms_norm(x_nxt, w_next, h, n_embd, t, eps)?;
780 }
781 } else {
782 e.add(x1, sl_fo, x_nxt, t * n_embd)?;
783 }
784 }
785 if let Some(path) = Self::prime_trace_path() {
791 let row = (base + t - 1) as usize;
792 let host = e.dtoh(x_nxt)?;
793 let last = &host[(t - 1) * n_embd..t * n_embd];
794 use std::io::Write as _;
795 let mut f = std::fs::OpenOptions::new().create(true).append(true).open(path)?;
796 let mut h64: u64 = 0xcbf29ce484222325;
797 for v in last {
798 h64 ^= v.to_bits() as u64;
799 h64 = h64.wrapping_mul(0x100000001b3);
800 }
801 writeln!(f, "{{\"pos\":{row},\"layer\":{il},\"t\":{t},\"base\":{base},\
802 \"hash\":\"{h64:016x}\",\"v0\":{:.9e},\"v1\":{:.9e},\"v2\":{:.9e}}}",
803 last[0], last[1], last[2])?;
804 }
805 std::mem::swap(&mut x_cur, &mut x_nxt);
806 }
807 let mut x = e.uninit(t * n_embd)?;
809 e.copy_into(&mut x, 0, x_cur, t * n_embd)?;
810 drop(slab_guard);
811
812 let mut h_seed = e.uninit(n_embd)?;
816 if !crate::spec::spec_hpost() {
817 e.copy_view_into(&mut h_seed, 0, &x.slice((t - 1) * n_embd..t * n_embd), n_embd)?;
818 }
819 let mut hn = e.uninit(t * n_embd)?;
821 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
822 if crate::spec::spec_hpost() {
823 e.copy_view_into(&mut h_seed, 0, &hn.slice((t - 1) * n_embd..t * n_embd), n_embd)?;
824 }
825 let last = e.view(&hn, t * n_embd);
826 let last_row = last.slice((t - 1) * n_embd..t * n_embd);
827 let mut hlast = e.uninit(n_embd)?;
828 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
829 let logits = e.matmul(&self.output, &hlast, 1)?;
830 cache.pos += t;
831 Ok((e.dtoh(&logits)?, h_seed, if crate::spec::spec_hpost() { hn } else { x }))
834 }
835
836 pub fn prime_chunk_captured(&self, e: &Engine, x_in: &CudaSlice<f32>, pos_d: &CudaSlice<i32>,
852 t: usize, cache: &mut Cache,
853 len_d: &CudaSlice<i32>,
854 logits_out: &mut CudaSlice<f32>, h_seed_out: &mut CudaSlice<f32>)
855 -> Result<(), Box<dyn std::error::Error>> {
856 let cfg = &self.cfg;
857 let n_embd = cfg.n_embd as usize;
858 let eps = cfg.rms_eps;
859 let f16fuse = crate::f16_ffi::pp_f16_enabled() && t >= 16;
860 let mut x = e.uninit(t * n_embd)?;
861 e.copy_into(&mut x, 0, x_in, t * n_embd)?;
862 for (il, layer) in self.layers.iter().enumerate() {
863 let mut h = e.uninit(t * n_embd)?;
864 let mut hx16: Option<CudaSlice<u8>> = None;
865 if f16fuse {
866 let mut b16 = e.alloc_u8_uninit(t * n_embd * 2)?;
867 e.rms_norm_f16out(&x, layer.attn_norm.float_data(), &mut h, &mut b16, n_embd, t, eps)?;
868 hx16 = Some(b16);
869 } else {
870 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
871 }
872 let mixed = match &layer.mixer {
873 Mixer::Full(fa) => self.full_attn_prime(e, fa, &h, hx16.as_ref(), pos_d, t, cache, il)?,
874 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
875 Mixer::Linear(la) => {
876 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
877 let g4 = match hx16.as_ref() {
878 Some(xh) => e.matmul_group_xh(&ws, &h, xh, t)?,
879 None => e.matmul_group(&ws, &h, t)?,
880 };
881 self.linear_attn_prime_core_pad(e, la, g4, t, cache, il, Some(len_d))?
882 }
883 };
884 let mut x1 = e.uninit(t * n_embd)?;
885 e.add(&x, &mixed, &mut x1, t * n_embd)?;
886 let mut z = e.uninit(t * n_embd)?;
887 let mut zx16: Option<CudaSlice<u8>> = None;
888 if f16fuse {
889 let mut b16 = e.alloc_u8_uninit(t * n_embd * 2)?;
890 e.rms_norm_f16out(&x1, layer.post_attn_norm.float_data(), &mut z, &mut b16, n_embd, t, eps)?;
891 zx16 = Some(b16);
892 } else {
893 e.rms_norm(&x1, layer.post_attn_norm.float_data(), &mut z, n_embd, t, eps)?;
894 }
895 let ffn_out = match &layer.ffn {
896 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
897 let n_ff = ffn_gate.out_features();
898 let mut g2 = match &zx16 {
899 Some(xh) => e.matmul_group_xh(&[ffn_gate, ffn_up], &z, xh, t)?,
900 None => e.matmul_group(&[ffn_gate, ffn_up], &z, t)?,
901 };
902 let up = g2.pop().unwrap();
903 let gate = g2.pop().unwrap();
904 let mut act = e.uninit(t * n_ff)?;
905 Self::ffn_act(e, &self.cfg, &gate, &up, &mut act, t * n_ff)?;
906 e.matmul(ffn_down, &act, t)?
907 }
908 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il(e, m, &z, t, il as u16)?,
909 };
910 let mut x2 = e.uninit(t * n_embd)?;
911 e.add(&x1, &ffn_out, &mut x2, t * n_embd)?;
912 x = x2;
913 }
914 if !crate::spec::spec_hpost() {
916 e.row_gather_dev(&x, h_seed_out, len_d, n_embd)?;
917 }
918 let mut hn = e.uninit(t * n_embd)?;
919 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
920 if crate::spec::spec_hpost() {
921 e.row_gather_dev(&hn, h_seed_out, len_d, n_embd)?;
922 }
923 let mut hlast = e.uninit(n_embd)?;
924 e.row_gather_dev(&hn, &mut hlast, len_d, n_embd)?;
925 let logits = e.matmul(&self.output, &hlast, 1)?;
926 let nv = logits.len();
927 e.copy_into(logits_out, 0, &logits, nv)?;
928 Ok(())
929 }
930
931 pub fn prime_cache_batch(&self, e: &Engine, prompts: &[&[u32]], caches: &mut [&mut Cache])
948 -> Result<Vec<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>)>, Box<dyn std::error::Error>> {
949 let cfg = &self.cfg;
950 let n_embd = cfg.n_embd as usize;
951 let eps = cfg.rms_eps;
952 let b = prompts.len();
953 assert!(b >= 1 && b == caches.len());
954 let pos0s: Vec<usize> = caches.iter().map(|c| c.pos).collect();
955 let carried = pos0s.iter().any(|&p| p > 0);
956 if carried && cfg.gemma4.is_some() {
957 return Err("prime_cache_batch: gemma4 has no continuation prime (v0 fresh-only)".into());
958 }
959 let ts: Vec<usize> = prompts.iter().map(|p| p.len()).collect();
960 for &t in &ts { assert!(t >= PRIME_MIN_T, "prime_cache_batch needs T >= {PRIME_MIN_T}"); }
961 for (s, c) in caches.iter().enumerate() {
962 assert!(c.pos + ts[s] <= c.max_ctx, "prime_cache_batch: prompt exceeds cache max_ctx");
963 }
964 let total: usize = ts.iter().sum();
965 let offs: Vec<usize> = ts.iter().scan(0usize, |a, &t| { let o = *a; *a += t; Some(o) }).collect();
966 let pos_ds: Vec<CudaSlice<i32>> = ts.iter().zip(&pos0s)
968 .map(|(&t, &p0)| e.htod_i32(&(p0 as i32..(p0 + t) as i32).collect::<Vec<_>>()))
969 .collect::<Result<_, _>>()?;
970 let split = |e: &Engine, y: &CudaSlice<f32>, dim: usize|
972 -> Result<Vec<CudaSlice<f32>>, Box<dyn std::error::Error>> {
973 let mut out = Vec::with_capacity(b);
974 for s in 0..b {
975 let mut ys = e.uninit(ts[s] * dim)?;
976 e.copy_view_into(&mut ys, 0, &y.slice(offs[s] * dim..(offs[s] + ts[s]) * dim), ts[s] * dim)?;
977 out.push(ys);
978 }
979 Ok(out)
980 };
981
982 let cat_tokens: Vec<u32> = prompts.iter().flat_map(|p| p.iter().copied()).collect();
983 let mut x = self.embed(e, &cat_tokens)?; for (il, layer) in self.layers.iter().enumerate() {
985 let mut h = e.uninit(total * n_embd)?;
986 let mut hx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
987 e.rms_norm_f16out(&x, layer.attn_norm.float_data(), &mut h, &mut hx16, n_embd, total, eps)?;
988 let mut mixed = e.uninit(total * n_embd)?;
990 match &layer.mixer {
991 Mixer::Full(fa) => {
992 let g3 = e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], &h, &hx16, total)?;
993 let (n_head, n_head_kv, head_dim) =
999 (self.cfg.n_head as usize, self.cfg.n_head_kv as usize, self.cfg.head_dim_k as usize);
1000 let fa_scale = 1.0 / (head_dim as f32).sqrt();
1001 let use_favl = !carried
1002 && (2..=8).contains(&b)
1003 && (head_dim == 256 || head_dim == 128)
1004 && self.cfg.attn_out_gate()
1005 && std::env::var("MEMRA_NOFA").is_err()
1006 && std::env::var("MEMRA_FA_FLOOR").is_err()
1007 && std::env::var("MEMRA_FA_PP_W2").as_deref() != Ok("1")
1008 && std::env::var("MEMRA_FA_BF16KV").as_deref() != Ok("0")
1009 && std::env::var("MEMRA_FA_VL").as_deref() != Ok("0");
1010 if use_favl {
1011 let (qf_w, kf_w, vf_w) =
1012 (fa.wq.out_features(), fa.wk.out_features(), fa.wv.out_features());
1013 struct APre {
1014 q: CudaSlice<f32>, gate: Option<CudaSlice<f32>>,
1015 qn: CudaSlice<f32>, kn: CudaSlice<f32>,
1016 }
1017 let mut aps = Vec::with_capacity(b);
1018 for &t in ts.iter().take(b) {
1019 aps.push(APre {
1020 q: e.uninit(t * n_head * head_dim)?,
1021 gate: Some(e.uninit(t * n_head * head_dim)?),
1022 qn: e.uninit(t * n_head * head_dim)?,
1023 kn: e.uninit(t * n_head_kv * head_dim)?,
1024 });
1025 }
1026 let (kv_dim_k, kv_dim_v, ktb, vtb) = {
1027 let kvl = caches[0].kv[il].as_ref().unwrap();
1028 (kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes)
1029 };
1030 let pargs: Vec<crate::AttnPreVl> = (0..b).map(|s| {
1031 let (o, t) = (offs[s], ts[s]);
1032 let kvl = caches[s].kv[il].as_ref().unwrap();
1033 assert!(kvl.len == 0 && kvl.len + t <= caches[s].max_ctx,
1034 "prime_cache_batch attn vl: fresh + capacity");
1035 crate::AttnPreVl {
1036 qf: e.addr_f32v(&g3[0].slice(o * qf_w..(o + t) * qf_w)),
1037 kf: e.addr_f32v(&g3[1].slice(o * kf_w..(o + t) * kf_w)),
1038 vf: e.addr_f32v(&g3[2].slice(o * vf_w..(o + t) * vf_w)),
1039 q: e.addr_f32(&aps[s].q),
1040 gate: e.addr_f32(aps[s].gate.as_ref().unwrap()),
1041 qn: e.addr_f32(&aps[s].qn), kn: e.addr_f32(&aps[s].kn),
1042 kc: e.addr_u8(&kvl.k), vc: e.addr_u8(&kvl.v),
1043 t: t as i32, pad: 0,
1044 }
1045 }).collect();
1046 e.attn_pre_vl8(&pargs, fa.q_norm.float_data(), fa.k_norm.float_data(),
1047 head_dim, self.cfg.rope_dim_count as usize, n_head, n_head_kv,
1048 self.cfg.rms_eps, self.cfg.rope_freq_base, 1.0,
1049 kv_dim_k, kv_dim_v, ktb, vtb)?;
1050 for s in 0..b {
1051 let kvl = caches[s].kv[il].as_mut().unwrap();
1052 kvl.len += ts[s];
1053 let new_len = kvl.len as i32;
1054 e.set_i32_one(&mut kvl.len_d, new_len)?;
1055 }
1056 let mut attns = Vec::with_capacity(b);
1057 let mut mirrors = Vec::with_capacity(b);
1058 for &t in ts.iter().take(b) {
1059 attns.push(e.uninit(t * n_head * head_dim)?);
1060 let n = t * n_head_kv * head_dim;
1061 mirrors.push((e.alloc_u8_uninit(n * 2)?, e.alloc_u8_uninit(n * 2)?));
1062 }
1063 let fa3_on = match std::env::var("MEMRA_FA3").as_deref() {
1066 Ok("0") => false,
1067 Ok("1") => true,
1068 _ => cfg!(memra_hopper_mma),
1069 };
1070 if fa3_on {
1071 let mut q16s = Vec::with_capacity(b);
1072 let mut v16s = Vec::with_capacity(b);
1073 for s in 0..b {
1074 let t = ts[s];
1075 let mut q16 = e.alloc_u8_uninit(t * n_head * head_dim * 2)?;
1076 e.f32_to_bf16_into(&aps[s].qn, &mut q16, t * n_head * head_dim)?;
1077 let mut k16 = e.alloc_u8_uninit(t * n_head_kv * head_dim * 2)?;
1078 e.f32_to_bf16_into(&aps[s].kn, &mut k16, t * n_head_kv * head_dim)?;
1079 let mut v16 = e.alloc_u8_uninit(t * n_head_kv * head_dim * 2)?;
1080 e.f32_to_bf16_v(&g3[2].slice(offs[s] * vf_w..(offs[s] + t) * vf_w),
1081 &mut v16, t * n_head_kv * head_dim)?;
1082 q16s.push(q16);
1083 v16s.push((k16, v16));
1084 }
1085 let mut qp = [core::ptr::null::<core::ffi::c_void>(); 8];
1086 let mut kp = qp;
1087 let mut vp = qp;
1088 let mut op = [core::ptr::null_mut::<f32>(); 8];
1089 let mut tsv = [0i32; 8];
1090 for s in 0..b {
1091 qp[s] = e.addr_u8(&q16s[s]) as *const core::ffi::c_void;
1092 kp[s] = e.addr_u8(&v16s[s].0) as *const core::ffi::c_void;
1093 vp[s] = e.addr_u8(&v16s[s].1) as *const core::ffi::c_void;
1094 op[s] = e.addr_f32(&attns[s]) as *mut f32;
1095 tsv[s] = ts[s] as i32;
1096 }
1097 let rc = unsafe {
1098 crate::fa3_vl_raw(qp.as_ptr(), kp.as_ptr(), vp.as_ptr(), op.as_ptr(),
1099 tsv.as_ptr(), b as i32, n_head as i32,
1100 n_head_kv as i32, head_dim as i32, fa_scale,
1101 e.stream().cu_stream() as *mut core::ffi::c_void)
1102 };
1103 if rc != 0 {
1104 return Err(format!("memra_fa3_vl rc={rc}").into());
1105 }
1106 } else {
1107 let fargs: Vec<crate::FaSeqVl> = (0..b).map(|s| crate::FaSeqVl {
1108 q: e.addr_f32(&aps[s].qn), k16: e.addr_u8(&mirrors[s].0),
1109 v16: e.addr_u8(&mirrors[s].1), o: e.addr_f32(&attns[s]),
1110 kf: e.addr_f32(&aps[s].kn),
1111 vf: e.addr_f32v(&g3[2].slice(offs[s] * vf_w..(offs[s] + ts[s]) * vf_w)),
1112 t: ts[s] as i32, pad: 0,
1113 }).collect();
1114 e.fa_prefill_vl8(&fargs, head_dim, n_head, n_head_kv, fa_scale)?;
1115 }
1116 for (s, attn) in attns.into_iter().enumerate() {
1117 let (attn_g, ag16) = self.full_attn_prime_post_fa(
1118 e, attn, &aps[s].gate, ts[s], n_head, head_dim)?;
1119 let mut done = false;
1120 if let Some(xh) = &ag16 {
1121 done = e.try_f16_gemm_pre_into_off(&fa.wo, xh, ts[s], &mut mixed, offs[s] * n_embd)?;
1122 }
1123 if !done {
1124 let m = e.matmul(&fa.wo, &attn_g, ts[s])?;
1125 e.copy_into(&mut mixed, offs[s] * n_embd, &m, ts[s] * n_embd)?;
1126 }
1127 }
1128 } else {
1129 let mut parts: Vec<Vec<CudaSlice<f32>>> = (0..b).map(|_| Vec::new()).collect();
1130 for (w, y) in [&fa.wq, &fa.wk, &fa.wv].iter().zip(g3) {
1131 for (s, ys) in split(e, &y, w.out_features())?.into_iter().enumerate() {
1132 parts[s].push(ys);
1133 }
1134 }
1135 for (s, g3s) in parts.into_iter().enumerate() {
1136 let (attn_g, ag16) = self.full_attn_prime_core_inner(
1138 e, fa, g3s, &pos_ds[s], ts[s], caches[s], il)?;
1139 let mut done = false;
1140 if let Some(xh) = &ag16 {
1141 done = e.try_f16_gemm_pre_into_off(&fa.wo, xh, ts[s], &mut mixed, offs[s] * n_embd)?;
1142 }
1143 if !done {
1144 let m = e.matmul(&fa.wo, &attn_g, ts[s])?;
1145 e.copy_into(&mut mixed, offs[s] * n_embd, &m, ts[s] * n_embd)?;
1146 }
1147 }
1148 }
1149 }
1150 Mixer::Mla(_) => crate::hybrid::mla_forward_unimplemented(),
1151 Mixer::Linear(la) => {
1152 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
1157 let g4 = e.matmul_group_xh(&ws, &h, &hx16, total)?;
1158 let outs = self.linear_attn_prime_core_batch(e, la, &g4, &offs, &ts, caches, il)?;
1159 for (s, (gn, gn16)) in outs.into_iter().enumerate() {
1160 let (o, t) = (offs[s], ts[s]);
1161 let mut done = false;
1162 if let Some(xh) = &gn16 {
1163 done = e.try_f16_gemm_pre_into_off(&la.ssm_out, xh, t, &mut mixed, o * n_embd)?;
1164 }
1165 if !done {
1166 let m = e.matmul(&la.ssm_out, &gn, t)?;
1167 e.copy_into(&mut mixed, o * n_embd, &m, t * n_embd)?;
1168 }
1169 }
1170 }
1171 }
1172 let mut x1 = e.uninit(total * n_embd)?;
1173 let mut z = e.uninit(total * n_embd)?;
1174 let mut zx16 = e.alloc_u8_uninit(total * n_embd * 2)?;
1175 e.add_rms_norm_f16out(&x, &mixed, layer.post_attn_norm.float_data(),
1176 &mut x1, &mut z, &mut zx16, n_embd, total, eps)?;
1177 let ffn_out = match &layer.ffn {
1178 crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } => {
1179 let n_ff = ffn_gate.out_features();
1180 let mut g2 = e.matmul_group_xh(&[ffn_gate, ffn_up], &z, &zx16, total)?;
1181 let up = g2.pop().unwrap();
1182 let gate = g2.pop().unwrap();
1183 let mut act = e.uninit(total * n_ff)?;
1184 if Self::f16out_on(e, total) && self.cfg.m3.is_none() {
1187 let mut a16 = e.alloc_u8_uninit(total * n_ff * 2)?;
1188 e.silu_mul_f16out(&gate, &up, &mut act, &mut a16, total * n_ff)?;
1189 match e.try_f16_gemm_pre(ffn_down, &a16, total)? {
1190 Some(y) => y,
1191 None => e.matmul(ffn_down, &act, total)?,
1192 }
1193 } else {
1194 Self::ffn_act(e, &self.cfg, &gate, &up, &mut act, total * n_ff)?;
1195 e.matmul(ffn_down, &act, total)?
1196 }
1197 }
1198 crate::hybrid::Ffn::Moe(m) => self.moe_ffn_il(e, m, &z, total, il as u16)?,
1199 };
1200 let mut x2 = e.uninit(total * n_embd)?;
1201 e.add(&x1, &ffn_out, &mut x2, total * n_embd)?;
1202 x = x2;
1203 }
1204 let mut hn = e.uninit(total * n_embd)?;
1206 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, total, eps)?;
1207 let mut hcat = e.uninit(b * n_embd)?;
1213 for s in 0..b {
1214 let last0 = (offs[s] + ts[s] - 1) * n_embd;
1215 e.copy_view_into(&mut hcat, s * n_embd, &hn.slice(last0..last0 + n_embd), n_embd)?;
1216 }
1217 let logits_cat = if b >= 2 { e.try_f16_gemm(&self.output, &hcat, b)? } else { None };
1218 let logits_host: Option<Vec<f32>> = match &logits_cat {
1219 Some(lc) => Some(e.dtoh(lc)?),
1220 None => None,
1221 };
1222 let n_vocab = self.output.out_features();
1223 let mut hidden_all = if crate::spec::spec_hpost() {
1224 split(e, &hn, n_embd)?
1225 } else {
1226 split(e, &x, n_embd)?
1227 };
1228 let mut out = Vec::with_capacity(b);
1229 for s in 0..b {
1230 let last0 = (offs[s] + ts[s] - 1) * n_embd;
1231 let mut h_seed = e.uninit(n_embd)?;
1232 if !crate::spec::spec_hpost() {
1233 e.copy_view_into(&mut h_seed, 0, &x.slice(last0..last0 + n_embd), n_embd)?;
1234 } else {
1235 e.copy_view_into(&mut h_seed, 0, &hn.slice(last0..last0 + n_embd), n_embd)?;
1236 }
1237 let logits = match &logits_host {
1238 Some(lh) => lh[s * n_vocab..(s + 1) * n_vocab].to_vec(),
1239 None => {
1240 let mut hlast = e.uninit(n_embd)?;
1241 e.copy_view_into(&mut hlast, 0, &hn.slice(last0..last0 + n_embd), n_embd)?;
1242 e.dtoh(&e.matmul(&self.output, &hlast, 1)?)?
1243 }
1244 };
1245 caches[s].pos += ts[s];
1246 out.push((logits, h_seed, hidden_all.remove(0)));
1247 }
1248 Ok(out)
1249 }
1250
1251 fn full_attn_prime(&self, e: &Engine, fa: &FullAttnLayer, h: &CudaSlice<f32>,
1257 hx: Option<&CudaSlice<u8>>,
1258 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache, il: usize)
1259 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1260 let g3 = match hx {
1265 Some(xh) => e.matmul_group_xh(&[&fa.wq, &fa.wk, &fa.wv], h, xh, t)?,
1266 None => e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?,
1267 };
1268 self.full_attn_prime_core(e, fa, g3, pos_d, t, cache, il)
1269 }
1270
1271 fn full_attn_prime_core(&self, e: &Engine, fa: &FullAttnLayer, g3: Vec<CudaSlice<f32>>,
1275 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache, il: usize)
1276 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1277 let (attn_g, ag16) = self.full_attn_prime_core_inner(e, fa, g3, pos_d, t, cache, il)?;
1278 if let Some(xh) = &ag16 {
1279 if let Some(y) = e.try_f16_gemm_pre(&fa.wo, xh, t)? {
1280 return Ok(y);
1281 }
1282 }
1283 Ok(e.matmul(&fa.wo, &attn_g, t)?)
1284 }
1285
1286 fn full_attn_prime_core_inner(&self, e: &Engine, fa: &FullAttnLayer, g3: Vec<CudaSlice<f32>>,
1287 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache, il: usize)
1288 -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
1289 let cfg = &self.cfg;
1290 let n_head = cfg.n_head as usize;
1291 let n_head_kv = cfg.n_head_kv as usize;
1292 let head_dim = cfg.head_dim_k as usize;
1293 let scale = 1.0 / (head_dim as f32).sqrt();
1294 let (pre, base_len) = self.full_attn_prime_pre_fa(e, fa, g3, pos_d, t, cache, il)?;
1295 let AttnPre { q, k, v, gate } = pre;
1296 let mut attn = e.uninit(t * n_head * head_dim)?;
1297 self.full_attn_prime_fa_dispatch(e, &q, &k, &v, &mut attn, base_len, t, cache, il,
1298 head_dim, n_head, n_head_kv, scale)?;
1299 self.full_attn_prime_post_fa(e, attn, &gate, t, n_head, head_dim)
1300 }
1301
1302 #[allow(clippy::type_complexity)]
1306 fn full_attn_prime_pre_fa(&self, e: &Engine, fa: &FullAttnLayer, mut g3: Vec<CudaSlice<f32>>,
1307 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache, il: usize)
1308 -> Result<(AttnPre, usize), Box<dyn std::error::Error>> {
1309 let cfg = &self.cfg;
1310 let n_head = cfg.n_head as usize;
1311 let n_head_kv = cfg.n_head_kv as usize;
1312 let head_dim = cfg.head_dim_k as usize;
1313 let eps = cfg.rms_eps;
1314
1315 let gated = cfg.attn_out_gate();
1319 let v = g3.pop().unwrap();
1320 let mut k = g3.pop().unwrap();
1321 let qf = g3.pop().unwrap();
1322 let (mut q, gate) = if gated {
1323 let mut q = e.uninit(t * n_head * head_dim)?;
1324 let mut gate = e.uninit(t * n_head * head_dim)?;
1325 e.q_gate_split(&qf, &mut q, &mut gate, head_dim, n_head, t)?;
1326 (q, Some(gate))
1327 } else {
1328 (qf, None)
1329 };
1330
1331 let mut qn = e.uninit(t * n_head * head_dim)?;
1332 e.rms_norm(&q, fa.q_norm.float_data(), &mut qn, head_dim, n_head * t, eps)?;
1333 q = qn;
1334 let mut kn = e.uninit(t * n_head_kv * head_dim)?;
1335 e.rms_norm(&k, fa.k_norm.float_data(), &mut kn, head_dim, n_head_kv * t, eps)?;
1336 k = kn;
1337 let rope_dims = cfg.rope_dim_count as usize;
1338 e.rope_neox(&mut q, pos_d, head_dim, rope_dims, n_head, t, cfg.rope_freq_base, 1.0)?;
1339 e.rope_neox(&mut k, pos_d, head_dim, rope_dims, n_head_kv, t, cfg.rope_freq_base, 1.0)?;
1340
1341 {
1344 let kvl = cache.kv[il].as_mut().unwrap();
1345 assert!(kvl.len + t <= cache.max_ctx, "prime_cache: KV overflow");
1346 e.append_kv_quantized_rows(&k, &v, &mut kvl.k, &mut kvl.v, kvl.len, t,
1347 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes,
1348 crate::Engine::kv_fp8_on())?;
1349 kvl.len += t;
1350 let new_len = kvl.len as i32;
1351 e.set_i32_one(&mut kvl.len_d, new_len)?;
1352 }
1353
1354 let base_len = {
1355 let kvl = cache.kv[il].as_ref().unwrap();
1356 kvl.len - t };
1358 Ok((AttnPre { q, k, v, gate }, base_len))
1359 }
1360
1361 #[allow(clippy::too_many_arguments)]
1368 fn full_attn_prime_fa_dispatch(&self, e: &Engine, q: &CudaSlice<f32>, k: &CudaSlice<f32>,
1369 v: &CudaSlice<f32>, attn: &mut CudaSlice<f32>, base_len: usize,
1370 t: usize, cache: &mut Cache, il: usize,
1371 head_dim: usize, n_head: usize, n_head_kv: usize, scale: f32)
1372 -> Result<(), Box<dyn std::error::Error>> {
1373 if base_len == 0 && std::env::var("MEMRA_PRIME_F32CHUNK0").as_deref() == Ok("1") {
1386 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
1387 e.sdpa_naive(q, k, v, attn, head_dim, n_head, n_head_kv, t, t, scale, true)?;
1388 } else {
1389 e.fa_prefill(q, k, v, attn, head_dim, n_head, n_head_kv, t, t, scale, true)?;
1390 }
1391 return Ok(());
1392 }
1393 let kvl = cache.kv[il].as_ref().unwrap();
1394 let t_kv = base_len + t;
1395 let k_view = e.view_u8(&kvl.k, t_kv * kvl.k_tok_bytes);
1396 let v_view = e.view_u8(&kvl.v, t_kv * kvl.v_tok_bytes);
1397 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
1401 e.sdpa_naive_quantized_view(q, &k_view, &v_view, attn, head_dim, n_head,
1402 n_head_kv, t, t_kv, scale, true,
1403 kvl.k_tok_bytes, kvl.v_tok_bytes)?;
1404 return Ok(());
1405 }
1406 let deqw = std::env::var("MEMRA_PRIME_DEQW").map(|v| v != "0").unwrap_or(true);
1414 if deqw {
1415 e.fa_prefill_view_ws(q, &k_view, &v_view, attn, head_dim, n_head, n_head_kv,
1416 t, t_kv, scale, true, kvl.k_tok_bytes, kvl.v_tok_bytes,
1417 crate::Engine::kv_fp8_on())?;
1418 } else {
1419 e.fa_prefill_view(q, &k_view, &v_view, attn, head_dim, n_head, n_head_kv,
1420 t, t_kv, scale, true, kvl.k_tok_bytes, kvl.v_tok_bytes,
1421 crate::Engine::kv_fp8_on())?;
1422 }
1423 Ok(())
1424 }
1425
1426 fn full_attn_prime_post_fa(&self, e: &Engine, attn: CudaSlice<f32>,
1429 gate: &Option<CudaSlice<f32>>, t: usize,
1430 n_head: usize, head_dim: usize)
1431 -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
1432 let (attn_g, ag16) = match gate {
1433 Some(gate) => {
1434 let n = t * n_head * head_dim;
1435 let mut ag = e.uninit(n)?;
1436 if Self::f16out_on(e, t) {
1437 let mut a16 = e.alloc_u8_uninit(n * 2)?;
1438 e.sig_mul_f16out(&attn, gate, &mut ag, &mut a16, n)?;
1439 (ag, Some(a16))
1440 } else {
1441 let mut gsig = e.uninit(n)?;
1442 e.sigmoid(gate, &mut gsig, n)?;
1443 e.mul(&attn, &gsig, &mut ag, n)?;
1444 (ag, None)
1445 }
1446 }
1447 None => (attn, None),
1448 };
1449 Ok((attn_g, ag16))
1450 }
1451
1452 fn linear_attn_prime(&self, e: &Engine, la: &LinearAttnLayer, h: &CudaSlice<f32>,
1459 hx: Option<&CudaSlice<u8>>, t: usize,
1460 cache: &mut Cache, il: usize)
1461 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1462 let ws = [&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha];
1464 let g4 = match hx {
1465 Some(xh) => e.matmul_group_xh(&ws, h, xh, t)?,
1466 None => e.matmul_group(&ws, h, t)?,
1467 };
1468 self.linear_attn_prime_core(e, la, g4, t, cache, il)
1469 }
1470
1471 fn linear_attn_prime_core(&self, e: &Engine, la: &LinearAttnLayer, mut g4: Vec<CudaSlice<f32>>,
1473 t: usize, cache: &mut Cache, il: usize)
1474 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1475 self.linear_attn_prime_core_pad(e, la, g4.drain(..).collect(), t, cache, il, None)
1476 }
1477
1478 #[allow(clippy::too_many_arguments)]
1482 fn linear_attn_prime_core_pad_inner(&self, e: &Engine, la: &LinearAttnLayer, mut g4: Vec<CudaSlice<f32>>,
1483 t: usize, cache: &mut Cache, il: usize,
1484 pad_len: Option<&CudaSlice<i32>>)
1485 -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
1486 let ssm = self.cfg.ssm.as_ref().unwrap();
1488 let d_state = ssm.state_size as usize;
1489 let num_k = ssm.group_count as usize;
1490 let num_v = ssm.time_step_rank as usize;
1491 let key_dim = d_state * num_k;
1492 let value_dim = d_state * num_v;
1493 let conv_dim = key_dim * 2 + value_dim;
1494 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(
1499 e, la,
1500 &qkv_mixed.slice(0..t * conv_dim), &z.slice(0..t * value_dim),
1501 &beta_raw.slice(0..t * num_v), &alpha.slice(0..t * num_v),
1502 t, cache, il, pad_len)
1503 }
1504
1505 #[allow(clippy::too_many_arguments)]
1508 fn linear_attn_gdn_prep(&self, e: &Engine, la: &LinearAttnLayer,
1509 qkv_mixed: &cudarc::driver::CudaView<f32>,
1510 beta_raw: &cudarc::driver::CudaView<f32>,
1511 alpha: &cudarc::driver::CudaView<f32>,
1512 t: usize, cache: &mut Cache, il: usize,
1513 pad_len: Option<&CudaSlice<i32>>)
1514 -> Result<GdnPrep, Box<dyn std::error::Error>> {
1515 let cfg = &self.cfg;
1516 let ssm = cfg.ssm.as_ref().unwrap();
1517 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;
1525 debug_assert!(t >= d_conv - 1, "stateful conv needs T >= pad (PRIME_MIN_T gates)");
1526
1527 let rl = cache.recur[il].as_mut().unwrap();
1532 let hk = Self::gdn_hk(e, t, num_v, num_k);
1533 let conv_fuse = std::env::var("MEMRA_CONV_FUSE").as_deref() != Ok("0");
1534 let hk = if conv_fuse { hk } else { num_v }; let mut q_g = e.uninit(d_state * hk * t)?;
1536 let mut k_g = e.uninit(d_state * hk * t)?;
1537 let mut v_g = e.uninit(d_state * num_v * t)?;
1538 if conv_fuse {
1539 e.ssm_conv1d_gdn_state_pad(qkv_mixed, &mut rl.conv_state, la.ssm_conv1d.float_data(),
1540 &mut q_g, &mut k_g, &mut v_g,
1541 conv_dim, t, d_conv, d_state, num_v, num_k, key_dim, hk, pad_len)?;
1542 } else {
1543 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(),
1545 &mut conv_out, conv_dim, t, d_conv, pad_len)?;
1546 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)?;
1547 }
1548 let mut q_l2 = e.uninit(d_state * hk * t)?;
1549 let qb16 = if Engine::l2_v2_on(d_state) && e.gdn_wgmma_on(32) {
1553 let mut qb = e.alloc_u8_uninit(d_state * hk * t * 2)?;
1554 e.l2_norm_pp(&q_g, &mut q_l2, Some(&mut qb), d_state, hk * t, eps)?;
1555 Some(qb)
1556 } else {
1557 e.l2_norm_pp(&q_g, &mut q_l2, None, d_state, hk * t, eps)?;
1558 None
1559 };
1560 let mut k_l2 = e.uninit(d_state * hk * t)?;
1561 let kb16 = if Engine::l2_v2_on(d_state) {
1563 let mut kb = e.alloc_u8_uninit(d_state * hk * t * 2)?;
1564 e.l2_norm_pp(&k_g, &mut k_l2, Some(&mut kb), d_state, hk * t, eps)?;
1565 Some(kb)
1566 } else {
1567 e.l2_norm_pp(&k_g, &mut k_l2, None, d_state, hk * t, eps)?;
1568 None
1569 };
1570 let mut beta = e.uninit(t * num_v)?;
1571 e.sigmoid_v(beta_raw, &mut beta, t * num_v)?;
1572 let mut g_log = e.uninit(t * num_v)?;
1573 e.gdn_glog_v(alpha, la.ssm_dt.float_data(), la.ssm_a.float_data(), &mut g_log, num_v, t)?;
1574 if let Some(len_d) = pad_len {
1575 e.gdn_pad_mask(&mut beta, &mut g_log, len_d, num_v, t)?;
1576 }
1577 Ok(GdnPrep { hk, q_l2, k_l2, v_g, beta, g_log, kb16, qb16 })
1578 }
1579
1580 #[allow(clippy::too_many_arguments)]
1585 fn linear_attn_prime_core_batch(&self, e: &Engine, la: &LinearAttnLayer,
1586 g4: &[CudaSlice<f32>], offs: &[usize], ts: &[usize],
1587 caches: &mut [&mut Cache], il: usize)
1588 -> Result<Vec<(CudaSlice<f32>, Option<CudaSlice<u8>>)>, Box<dyn std::error::Error>> {
1589 let ssm = self.cfg.ssm.as_ref().unwrap();
1590 let d_state = ssm.state_size as usize;
1591 let num_k = ssm.group_count as usize;
1592 let num_v = ssm.time_step_rank as usize;
1593 let key_dim = d_state * num_k;
1594 let value_dim = d_state * num_v;
1595 let conv_dim = key_dim * 2 + value_dim;
1596 let eps = self.cfg.rms_eps;
1597 let scale = 1.0 / (d_state as f32).sqrt();
1598 let b = ts.len();
1599 let c = Engine::gdn_chunk_size();
1600 let carried = caches.iter().any(|c| c.pos > 0);
1603 let use_vl = !carried
1604 && (2..=8).contains(&b)
1605 && Engine::gdn_chunked_enabled() && ts.iter().all(|&t| t >= 16)
1606 && e.gdn_mma_enabled(c)
1607 && std::env::var("MEMRA_GDN_VL").as_deref() != Ok("0");
1608 if !use_vl {
1609 return (0..b).map(|s| {
1610 let (o, t) = (offs[s], ts[s]);
1611 self.linear_attn_prime_core_pad_view(
1612 e, la,
1613 &g4[0].slice(o * conv_dim..(o + t) * conv_dim),
1614 &g4[1].slice(o * value_dim..(o + t) * value_dim),
1615 &g4[2].slice(o * num_v..(o + t) * num_v),
1616 &g4[3].slice(o * num_v..(o + t) * num_v),
1617 t, caches[s], il, None)
1618 }).collect();
1619 }
1620 struct SeqBufs {
1624 conv_out: CudaSlice<f32>, q_g: CudaSlice<f32>, k_g: CudaSlice<f32>, v_g: CudaSlice<f32>,
1625 q_l2: CudaSlice<f32>, k_l2: CudaSlice<f32>, beta: CudaSlice<f32>, g_log: CudaSlice<f32>,
1626 gn: CudaSlice<f32>, gn16: CudaSlice<u8>,
1627 }
1628 let d_conv = ssm.conv_kernel as usize;
1629 let f16o = Self::f16out_on(e, 16);
1630 let hk = Self::gdn_hk(e, 16, num_v, num_k); let mut sb = Vec::with_capacity(b);
1632 let mut pres = Vec::with_capacity(b);
1633 for &t in ts.iter().take(b) {
1634 sb.push(SeqBufs {
1635 conv_out: e.uninit(conv_dim * t)?,
1636 q_g: e.uninit(d_state * hk * t)?,
1637 k_g: e.uninit(d_state * hk * t)?,
1638 v_g: e.uninit(d_state * num_v * t)?,
1639 q_l2: e.uninit(d_state * hk * t)?,
1640 k_l2: e.uninit(d_state * hk * t)?,
1641 beta: e.uninit(t * num_v)?,
1642 g_log: e.uninit(t * num_v)?,
1643 gn: e.uninit(d_state * num_v * t)?,
1644 gn16: e.alloc_u8_uninit(d_state * num_v * t * 2)?,
1645 });
1646 pres.push(e.gdn_chunk_alloc(num_v, t, c, hk)?);
1647 }
1648 let prep_args: Vec<crate::GdnPrepVl> = (0..b).map(|s| {
1649 let (o, t) = (offs[s], ts[s]);
1650 let rl = caches[s].recur[il].as_ref().unwrap();
1651 crate::GdnPrepVl {
1652 qkv: e.addr_f32v(&g4[0].slice(o * conv_dim..(o + t) * conv_dim)),
1653 conv_state: e.addr_f32(&rl.conv_state),
1654 conv_out: e.addr_f32(&sb[s].conv_out),
1655 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),
1656 q_l2: e.addr_f32(&sb[s].q_l2), k_l2: e.addr_f32(&sb[s].k_l2),
1657 beta_raw: e.addr_f32v(&g4[2].slice(o * num_v..(o + t) * num_v)),
1658 alpha: e.addr_f32v(&g4[3].slice(o * num_v..(o + t) * num_v)),
1659 beta: e.addr_f32(&sb[s].beta), g_log: e.addr_f32(&sb[s].g_log),
1660 o: e.addr_f32(&pres[s].o),
1661 z: e.addr_f32v(&g4[1].slice(o * value_dim..(o + t) * value_dim)),
1662 gn: e.addr_f32(&sb[s].gn), gn16: e.addr_u8(&sb[s].gn16),
1663 kb16: if Engine::l2_v2_on(d_state) { e.addr_u8(&pres[s].kb16) } else { 0 },
1664 qb16: if Engine::l2_v2_on(d_state) && e.gdn_wgmma_on(c) { e.addr_u8(&pres[s].qb16) } else { 0 },
1665 t: t as i32, pad: 0,
1666 }
1667 }).collect();
1668 let args: Vec<crate::GdnSeqVl> = (0..b).map(|s| {
1669 let rl = caches[s].recur[il].as_ref().unwrap();
1670 crate::GdnSeqVl {
1671 kb16: e.addr_u8(&pres[s].kb16), gcum: e.addr_f32(&pres[s].gcum),
1672 beta: e.addr_f32(&sb[s].beta), u: e.addr_f32(&pres[s].u),
1673 wb16: e.addr_u8(&pres[s].wb16), y: e.addr_u8(&pres[s].y16),
1674 ssnap: e.addr_u8(&pres[s].ssnap16),
1675 state_in: e.addr_f32(&rl.ssm_state), state_out: e.addr_f32(&rl.ssm_state_alt),
1676 q: e.addr_f32(&sb[s].q_l2), p: e.addr_f32(&pres[s].p),
1677 o: e.addr_f32(&pres[s].o),
1678 k: e.addr_f32(&sb[s].k_l2), v: e.addr_f32(&sb[s].v_g),
1679 g: e.addr_f32(&sb[s].g_log), a: e.addr_f32(&pres[s].a),
1680 w: e.addr_f32(&pres[s].w),
1681 t: ts[s] as i32, nc: pres[s].nc as i32,
1682 }
1683 }).collect();
1684 e.gdn_prep_vl8(&prep_args, la.ssm_conv1d.float_data(), la.ssm_dt.float_data(),
1685 la.ssm_a.float_data(), conv_dim, d_conv, d_state, num_v, num_k, key_dim, hk, eps)?;
1686 if !Engine::l2_v2_on(d_state) {
1689 e.gdn_mirror_vl8(&args, num_v, 0, hk)?;
1690 }
1691 let wq8: Option<crate::GdnWVl8> = if e.gdn_wgmma_on(c) {
1693 if !Engine::l2_v2_on(d_state) {
1695 for s in 0..b {
1696 e.f32_to_bf16_into(&sb[s].q_l2, &mut pres[s].qb16, d_state * hk * ts[s])?;
1697 }
1698 }
1699 let mut wa = [crate::GdnWVl::default(); 8];
1700 for s in 0..b {
1701 wa[s] = crate::GdnWVl { qb16: e.addr_u8(&pres[s].qb16), pb16: e.addr_u8(&pres[s].pb16) };
1702 }
1703 Some(crate::GdnWVl8(wa))
1704 } else { None };
1705 e.gdn_chunk_k123_vl8(&args, num_v, hk, wq8.as_ref())?;
1706 e.gdn_chunk_vl8(&args, num_v, scale, hk, wq8.as_ref())?;
1707 if f16o {
1708 e.gdn_tail_vl8(&prep_args, la.ssm_norm.float_data(), d_state, num_v, eps)?;
1709 }
1710 let mut out = Vec::with_capacity(b);
1712 for (s, bufs) in sb.into_iter().enumerate() {
1713 let rl = caches[s].recur[il].as_mut().unwrap();
1714 std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
1715 let (o, t) = (offs[s], ts[s]);
1716 let SeqBufs { mut gn, gn16, .. } = bufs;
1717 if f16o {
1718 out.push((gn, Some(gn16)));
1719 } else {
1720 let z_v = g4[1].slice(o * value_dim..(o + t) * value_dim);
1721 e.gated_rmsnorm_zv(&pres[s].o, la.ssm_norm.float_data(), &z_v, &mut gn,
1722 d_state, num_v * t, eps)?;
1723 out.push((gn, None));
1724 }
1725 }
1726 Ok(out)
1727 }
1728
1729 #[allow(clippy::too_many_arguments)]
1733 fn linear_attn_prime_core_pad_view(&self, e: &Engine, la: &LinearAttnLayer,
1734 qkv_mixed: &cudarc::driver::CudaView<f32>,
1735 z: &cudarc::driver::CudaView<f32>,
1736 beta_raw: &cudarc::driver::CudaView<f32>,
1737 alpha: &cudarc::driver::CudaView<f32>,
1738 t: usize, cache: &mut Cache, il: usize,
1739 pad_len: Option<&CudaSlice<i32>>)
1740 -> Result<(CudaSlice<f32>, Option<CudaSlice<u8>>), Box<dyn std::error::Error>> {
1741 let cfg = &self.cfg;
1742 let ssm = cfg.ssm.as_ref().unwrap();
1743 let d_state = ssm.state_size as usize; let num_v = ssm.time_step_rank as usize; let eps = cfg.rms_eps;
1746 let scale = 1.0 / (d_state as f32).sqrt();
1747
1748 let prep = self.linear_attn_gdn_prep(e, la, qkv_mixed, beta_raw, alpha, t, cache, il, pad_len)?;
1749
1750 let mut o = e.uninit(d_state * num_v * t)?;
1756 let rl = cache.recur[il].as_mut().unwrap();
1757 {
1758 let crate::cache::RecurLayer { ssm_state, ssm_state_alt, .. } = rl;
1759 e.gdn_scan_prefill(&prep.q_l2, &prep.k_l2, &prep.v_g, &prep.g_log, &prep.beta,
1760 prep.kb16.as_ref(), prep.qb16.as_ref(), ssm_state, ssm_state_alt, &mut o, num_v, t, scale,
1761 prep.hk)?;
1762 }
1763 std::mem::swap(&mut rl.ssm_state, &mut rl.ssm_state_alt);
1764
1765 let mut gn = e.uninit(d_state * num_v * t)?;
1768 let gn16 = if Self::f16out_on(e, t) {
1769 let mut g16 = e.alloc_u8_uninit(d_state * num_v * t * 2)?;
1770 e.gated_rmsnorm_f16out_zv(&o, la.ssm_norm.float_data(), z, &mut gn, &mut g16,
1771 d_state, num_v * t, eps)?;
1772 Some(g16)
1773 } else {
1774 e.gated_rmsnorm_zv(&o, la.ssm_norm.float_data(), z, &mut gn, d_state, num_v * t, eps)?;
1775 None
1776 };
1777 Ok((gn, gn16))
1778 }
1779
1780 #[allow(clippy::too_many_arguments)]
1782 fn linear_attn_prime_core_pad(&self, e: &Engine, la: &LinearAttnLayer, g4: Vec<CudaSlice<f32>>,
1783 t: usize, cache: &mut Cache, il: usize,
1784 pad_len: Option<&CudaSlice<i32>>)
1785 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1786 let (gn, gn16) = self.linear_attn_prime_core_pad_inner(e, la, g4, t, cache, il, pad_len)?;
1787 if let Some(xh) = &gn16 {
1788 if let Some(y) = e.try_f16_gemm_pre(&la.ssm_out, xh, t)? {
1789 return Ok(y);
1790 }
1791 }
1792 Ok(e.matmul(&la.ssm_out, &gn, t)?)
1793 }
1794
1795 pub fn full_attn(&self, e: &Engine, fa: &FullAttnLayer, h: &CudaSlice<f32>, pos_d: &CudaSlice<i32>, t: usize)
1797 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1798 let cfg = &self.cfg;
1799 let _n_embd = cfg.n_embd as usize;
1800 let n_head = cfg.n_head as usize;
1801 let n_head_kv = cfg.n_head_kv as usize;
1802 let head_dim = cfg.head_dim_k as usize;
1803 let eps = cfg.rms_eps;
1804 let scale = 1.0 / (head_dim as f32).sqrt();
1805
1806 let gated = cfg.attn_out_gate();
1809 let mut g3 = e.matmul_group(&[&fa.wq, &fa.wk, &fa.wv], h, t)?;
1811 let v = g3.pop().unwrap();
1812 let mut k = g3.pop().unwrap();
1813 let qf = g3.pop().unwrap();
1814 let (mut q, gate) = if gated {
1815 let mut q = e.uninit(t * n_head * head_dim)?;
1816 let mut gate = e.uninit(t * n_head * head_dim)?;
1817 e.q_gate_split(&qf, &mut q, &mut gate, head_dim, n_head, t)?;
1818 (q, Some(gate))
1819 } else {
1820 (qf, None)
1821 };
1822
1823 let mut qn = e.uninit(t * n_head * head_dim)?;
1825 e.rms_norm(&q, fa.q_norm.float_data(), &mut qn, head_dim, n_head * t, eps)?;
1826 q = qn;
1827 let mut kn = e.uninit(t * n_head_kv * head_dim)?;
1828 e.rms_norm(&k, fa.k_norm.float_data(), &mut kn, head_dim, n_head_kv * t, eps)?;
1829 k = kn;
1830 let rope_dims = cfg.rope_dim_count as usize;
1831 e.rope_neox(&mut q, pos_d, head_dim, rope_dims, n_head, t, cfg.rope_freq_base, 1.0)?;
1832 e.rope_neox(&mut k, pos_d, head_dim, rope_dims, n_head_kv, t, cfg.rope_freq_base, 1.0)?;
1833
1834 let mut attn = e.uninit(t * n_head * head_dim)?;
1836 if std::env::var("MEMRA_NOFA").is_ok() || !(head_dim == 256 || head_dim == 128) {
1839 e.sdpa_naive(&q, &k, &v, &mut attn, head_dim, n_head, n_head_kv, t, t, scale, true)?;
1841 } else {
1842 e.fa_prefill(&q, &k, &v, &mut attn, head_dim, n_head, n_head_kv, t, t, scale, true)?;
1843 }
1844
1845 let attn_g = match &gate {
1847 Some(gate) => {
1848 let mut gsig = e.uninit(t * n_head * head_dim)?;
1849 e.sigmoid(gate, &mut gsig, t * n_head * head_dim)?;
1850 let mut ag = e.uninit(t * n_head * head_dim)?;
1851 e.mul(&attn, &gsig, &mut ag, t * n_head * head_dim)?;
1852 ag
1853 }
1854 None => attn,
1855 };
1856
1857 let o = e.matmul(&fa.wo, &attn_g, t)?;
1859 Ok(o)
1860 }
1861
1862 pub fn linear_attn(&self, e: &Engine, la: &LinearAttnLayer, h: &CudaSlice<f32>, t: usize)
1864 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1865 let cfg = &self.cfg;
1866 let _n_embd = cfg.n_embd as usize;
1867 let ssm = cfg.ssm.as_ref().unwrap();
1868 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;
1873 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;
1877 let scale = 1.0 / (d_state as f32).sqrt();
1878
1879 let mut g4 = e.matmul_group(&[&la.wqkv, &la.wqkv_gate, &la.ssm_beta, &la.ssm_alpha], h, t)?;
1882 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);
1894 let mut q_g = e.uninit(d_state * num_v * t)?;
1895 let mut k_g = e.uninit(d_state * num_v * t)?;
1896 let mut v_g = e.uninit(d_state * num_v * t)?;
1897 e.ssm_conv1d_gdn(&qkv_mixed, la.ssm_conv1d.float_data(), &mut q_g, &mut k_g, &mut v_g,
1898 conv_dim, t, d_conv, d_state, num_v, num_k, key_dim)?;
1899 let mut q_l2 = e.uninit(d_state * num_v * t)?;
1901 e.l2_norm(&q_g, &mut q_l2, d_state, num_v * t, eps)?;
1902 let mut k_l2 = e.uninit(d_state * num_v * t)?;
1903 e.l2_norm(&k_g, &mut k_l2, d_state, num_v * t, eps)?;
1904 let v_gd = v_g;
1905
1906 let mut beta = e.uninit(t * num_v)?;
1909 e.sigmoid(&beta_raw, &mut beta, t * num_v)?;
1910 let mut g_log = e.uninit(t * num_v)?;
1912 e.gdn_glog(&alpha, la.ssm_dt.float_data(), la.ssm_a.float_data(), &mut g_log, num_v, t)?;
1913
1914 let state_in = e.zeros(d_state * d_state * num_v)?; let mut state_out = e.zeros(d_state * d_state * num_v)?;
1917 let mut o = e.uninit(d_state * num_v * t)?;
1918 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)?;
1919
1920 let mut gn = e.uninit(d_state * num_v * t)?;
1925 e.gated_rmsnorm(&o, la.ssm_norm.float_data(), &z, &mut gn, d_state, num_v * t, eps)?;
1926
1927 let out = e.matmul(&la.ssm_out, &gn, t)?;
1931 Ok(out)
1932 }
1933}
1934
1935impl HybridModel {
1936 pub fn moe_ffn_il(&self, e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, t: usize, il: u16)
1947 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1948 Self::moe_ffn(e, m, z, t, &self.cfg, il, self.max_moe_block())
1949 }
1950
1951 pub fn moe_ffn_il_zq8(&self, e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>,
1955 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>, t: usize, il: u16)
1956 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1957 Self::moe_ffn_inner(e, m, z, zq8, t, &self.cfg, il, self.max_moe_block())
1958 }
1959
1960 pub(crate) fn moe_ffn(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, t: usize,
1968 cfg: &ModelConfig, il: u16, max_block: usize)
1969 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1970 Self::moe_ffn_inner(e, m, z, None, t, cfg, il, max_block)
1971 }
1972
1973 #[allow(clippy::too_many_arguments)]
1974 pub(crate) fn moe_ffn_inner(
1975 e: &Engine,
1976 m: &MoeWeights,
1977 z: &CudaSlice<f32>,
1978 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
1979 t: usize,
1980 cfg: &ModelConfig,
1981 il: u16,
1982 max_block: usize,
1983 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
1984 let worker_io = crate::spill_pread::worker_enabled();
1985 let epoch_lfu = std::env::var_os("MEMRA_MOE_LFU_DECAY").is_some();
1986 if Engine::moe_cache_enabled() && (worker_io || epoch_lfu) {
1987 e.with_moe_cache(max_block, |cache, _| {
1988 cache.begin_forward_epoch(il, t);
1989 if worker_io {
1990 cache.begin_worker_scope();
1991 }
1992 Ok(())
1993 })?;
1994 }
1995 if t > 1 && std::env::var("MEMRA_MOE_GROUPED").is_ok() {
1997 let grouped_out = Self::moe_ffn_grouped(e, m, z, t, cfg, il, max_block)?;
1998 if std::env::var("MEMRA_MOE_GATE").is_ok() {
2005 let seq_out = Self::moe_ffn_sequential(e, m, z, t, cfg, il, max_block)?;
2006 let g_host = e.dtoh(&grouped_out)?;
2007 let s_host = e.dtoh(&seq_out)?;
2008 let g_bytes: &[u8] = unsafe { std::slice::from_raw_parts(g_host.as_ptr() as *const u8, g_host.len() * 4) };
2009 let s_bytes: &[u8] = unsafe { std::slice::from_raw_parts(s_host.as_ptr() as *const u8, s_host.len() * 4) };
2010 if g_bytes == s_bytes {
2011 if il == 0 { println!("moe-gate il={il} t={t} BYTE-IDENTICAL (first layer only printed)"); }
2012 } else {
2013 let diffs = g_host.iter().zip(s_host.iter()).enumerate()
2014 .filter(|(_, (a, b))| a != b).count();
2015 let maxdiff = g_host.iter().zip(s_host.iter())
2016 .map(|(a, b)| (a - b).abs()).fold(0.0f32, f32::max);
2017 panic!("moe-gate il={il} t={t} MISMATCH: {diffs}/{} elems differ, maxdiff={maxdiff:.6e}", g_host.len());
2018 }
2019 }
2020 return Ok(grouped_out);
2021 }
2022 Self::moe_ffn_sequential_zq8(e, m, z, zq8, t, cfg, il, max_block)
2023 }
2024
2025 pub(crate) fn moe_ffn_sequential(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, t: usize,
2027 cfg: &ModelConfig, il: u16, max_block: usize)
2028 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2029 Self::moe_ffn_sequential_zq8(e, m, z, None, t, cfg, il, max_block)
2030 }
2031
2032 fn trace_moe_routes(il: u16, t: usize, sel_all: &[u32], weights: &[f32])
2036 -> Result<(), Box<dyn std::error::Error>> {
2037 use std::io::Write as _;
2038 if let Ok(path) = std::env::var("MEMRA_MOE_TRACE") {
2039 let mut f = std::fs::OpenOptions::new().create(true).append(true).open(path)?;
2040 let ids: Vec<String> = sel_all.iter().map(|s| s.to_string()).collect();
2041 writeln!(f, "{} {} {}", il, t, ids.join(","))?;
2042 }
2043 if let Ok(path) = std::env::var("MEMRA_MOE_WEIGHT_TRACE") {
2044 let mut f = std::fs::OpenOptions::new().create(true).append(true).open(path)?;
2045 let pairs: Vec<String> = sel_all.iter().zip(weights)
2046 .map(|(expert, weight)| format!("{expert}:{weight:.9}"))
2047 .collect();
2048 writeln!(f, "{} {} {}", il, t, pairs.join(","))?;
2049 }
2050 Ok(())
2051 }
2052
2053 fn trace_moe_input(e: &Engine, il: u16, t: usize, n_embd: usize, z: &CudaSlice<f32>)
2058 -> Result<(), Box<dyn std::error::Error>> {
2059 use std::io::Write as _;
2060 let Ok(dir) = std::env::var("MEMRA_MOE_INPUT_TRACE_DIR") else { return Ok(()) };
2061 let host = e.dtoh(z)?;
2062 if host.len() != t * n_embd {
2063 return Err(format!(
2064 "MoE input trace shape mismatch at layer {il}: got {} values, expected {}x{}",
2065 host.len(), t, n_embd
2066 ).into());
2067 }
2068 let bytes = unsafe {
2069 std::slice::from_raw_parts(
2070 host.as_ptr().cast::<u8>(), host.len() * std::mem::size_of::<f32>()
2071 )
2072 };
2073 let state = MOE_INPUT_TRACE_WRITER.get_or_init(|| std::sync::Mutex::new(None));
2074 let mut state = state.lock().map_err(|_| "MoE input trace writer lock is poisoned")?;
2075 if state.is_none() {
2076 let dir = std::path::PathBuf::from(&dir);
2077 std::fs::create_dir_all(&dir)?;
2078 let index = std::fs::OpenOptions::new().create(true).append(true)
2079 .open(dir.join("index.jsonl"))?;
2080 *state = Some(MoeInputTraceWriter {
2081 dir,
2082 index,
2083 payloads: std::collections::HashMap::new(),
2084 });
2085 }
2086 let writer = state.as_mut().unwrap();
2087 if writer.dir != std::path::Path::new(&dir) {
2088 return Err("MEMRA_MOE_INPUT_TRACE_DIR changed after capture started".into());
2089 }
2090 let file_name = format!("layer-{il:03}.f32");
2091 if !writer.payloads.contains_key(&il) {
2092 let payload = std::fs::OpenOptions::new().create(true).append(true)
2093 .open(writer.dir.join(&file_name))?;
2094 let offset = payload.metadata()?.len();
2095 writer.payloads.insert(il, (payload, offset));
2096 }
2097 let (payload, offset) = writer.payloads.get_mut(&il).unwrap();
2098 let row_offset = *offset;
2099 payload.write_all(bytes)?;
2100 *offset += bytes.len() as u64;
2101 writeln!(
2102 writer.index,
2103 "{{\"format\":\"memra-moe-input-trace-v1\",\"layer\":{il},\"tokens\":{t},\
2104 \"hidden_size\":{n_embd},\"file\":\"{file_name}\",\"offset\":{row_offset},\
2105 \"payload_bytes\":{}}}",
2106 bytes.len()
2107 )?;
2108 Ok(())
2109 }
2110
2111 #[allow(clippy::too_many_arguments)]
2112 pub(crate) fn moe_ffn_sequential_zq8(
2113 e: &Engine,
2114 m: &MoeWeights,
2115 z: &CudaSlice<f32>,
2116 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>,
2117 t: usize,
2118 cfg: &ModelConfig,
2119 il: u16,
2120 max_block: usize,
2121 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
2122 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
2123 let moe = cfg.moe.as_ref().unwrap();
2124 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);
2131 debug_assert_eq!(m.gate_exps.out_f, n_ff_exp);
2132 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);
2135
2136 let use_cache = Engine::moe_cache_enabled();
2137 let uniform_experts = m.has_uniform_expert_layout();
2138 let moe_q8 = uniform_experts && moe_q8_enabled()
2139 && q8_expert_supported(m.gate_exps.qtype) && q8_expert_supported(m.up_exps.qtype)
2140 && q8_expert_supported(m.down_exps.qtype);
2141 let cpu_expert_requested = crate::cpu_experts::configured();
2148 if cpu_expert_requested && (cfg.hy3.is_none() || cfg.m3.is_some()) {
2149 return Err(std::io::Error::other(
2150 "MEMRA_CPU_EXPERT_LIB is experimental and currently gated to Hy3",
2151 )
2152 .into());
2153 }
2154 let cpu_hybrid = cpu_expert_requested && t < PRIME_MIN_T && m.dev_exps.is_none();
2155 let freeze_cpu_residency = cpu_expert_requested
2161 && std::env::var("MEMRA_CPU_EXPERT_FREEZE_CACHE").as_deref() == Ok("1");
2162 let caller_warms_before_freeze = std::env::var("MEMRA_CPU_EXPERT_FREEZE_WARMUP_TOKENS")
2163 .ok()
2164 .and_then(|value| value.parse::<usize>().ok())
2165 .is_some_and(|tokens| tokens > 0);
2166 if cpu_hybrid && freeze_cpu_residency && !caller_warms_before_freeze {
2167 e.freeze_moe_cache();
2168 }
2169 let cache_frozen = use_cache && e.moe_cache_frozen();
2170 let cache_dispatch = use_cache && (!cache_frozen || cpu_hybrid);
2171
2172 let logits = if t < PRIME_MIN_T {
2179 if crate::router_kernel_on() {
2183 e.router_gemv(m.gate_inp.float_data(), z, cfg.n_embd as usize,
2186 m.gate_exps.n_expert, t)?
2187 } else {
2188 e.matmul_decode_exact(&m.gate_inp, z, t)?
2189 }
2190 } else if crate::router_prefill_exact_on() && crate::router_kernel_on() {
2191 e.router_gemv(m.gate_inp.float_data(), z, cfg.n_embd as usize,
2206 m.gate_exps.n_expert, t)?
2207 } else {
2208 e.matmul(&m.gate_inp, z, t)?
2209 };
2210
2211 let no_exp_macros = m.gate_exps.macros.is_none() && m.up_exps.macros.is_none()
2249 && m.down_exps.macros.is_none();
2250 if cfg.sigmoid_router().is_none() && cfg.m3.is_none() && cfg.hy3.is_none()
2251 && no_exp_macros
2252 && t >= PRIME_MIN_T && m.dev_exps.is_some() && moe_q8_enabled()
2253 && q8_expert_supported(m.gate_exps.qtype) && q8_expert_supported(m.up_exps.qtype)
2254 && q8_expert_supported(m.down_exps.qtype)
2255 && std::env::var("MEMRA_MOE_PAIRS").map(|v| v != "0").unwrap_or(true)
2256 && std::env::var("MEMRA_MOE_STATS").is_err() {
2257 return Self::moe_ffn_pairs(e, m, z, &logits, t, cfg);
2258 }
2259
2260 let dev_ok = uniform_experts && cfg.m3.is_none() && cfg.hy3.is_none();
2270 let observe_routes = std::env::var("MEMRA_MOE_STATS").is_ok()
2274 || std::env::var("MEMRA_MOE_TRACE").is_ok()
2275 || std::env::var("MEMRA_MOE_WEIGHT_TRACE").is_ok()
2276 || std::env::var("MEMRA_MOE_INPUT_TRACE_DIR").is_ok();
2277 if dev_ok && t < PRIME_MIN_T && m.dev_exps.is_some() && n_used <= 8 && moe_dev_enabled()
2278 && !observe_routes {
2279 return Self::moe_ffn_dev(e, m, z, zq8, &logits, t, cfg, il, max_block);
2280 }
2281 if dev_ok && use_cache && n_used <= 8 && moe_dev_enabled()
2282 && !observe_routes {
2283 let row_ok = e.with_moe_cache(max_block, |c, eng| {
2284 if moe_prewarm_enabled() { c.prewarm_layer(il, m, eng)?; }
2285 Ok(c.layer_dev_row(il, n_expert, eng)?.is_some())
2286 })?;
2287 if row_ok {
2288 return Self::moe_ffn_dev(e, m, z, zq8, &logits, t, cfg, il, max_block);
2289 }
2290 }
2291
2292 let (sel_all, w_all, routed_cpu_input) = if let Some(sig) = cfg.sigmoid_router() {
2294 if cpu_hybrid {
2295 let (sel, w, input) = Self::moe_route_sigmoid_with_input(
2296 e,
2297 &logits,
2298 z,
2299 t,
2300 n_expert,
2301 n_used,
2302 m.exp_probs_b.as_deref(),
2303 sig,
2304 m.active_experts.as_deref(),
2305 )?;
2306 (sel, w, Some(input))
2307 } else {
2308 let (sel, w) = Self::moe_route_cfg(
2309 e,
2310 &logits,
2311 t,
2312 n_expert,
2313 n_used,
2314 m.exp_probs_b.as_deref(),
2315 Some(sig),
2316 m.active_experts.as_deref(),
2317 )?;
2318 (sel, w, None)
2319 }
2320 } else {
2321 let (sel, w) = Self::moe_route_cfg(
2322 e,
2323 &logits,
2324 t,
2325 n_expert,
2326 n_used,
2327 None,
2328 None,
2329 m.active_experts.as_deref(),
2330 )?;
2331 (sel, w, None)
2332 };
2333
2334 Self::trace_moe_routes(il, t, &sel_all, &w_all)?;
2338 Self::trace_moe_input(e, il, t, n_embd, z)?;
2339
2340 let worker_disk_prefetch =
2352 cache_dispatch && crate::spill_pread::worker_enabled() && !cpu_hybrid;
2353 let promote_worker_h2d =
2354 t == 1 && worker_disk_prefetch && crate::spill_pread::copy_h2d_enabled();
2355 if promote_worker_h2d {
2356 let mut selected_blocks = Vec::with_capacity(n_used * 3);
2357 for &ex in sel_all.iter().take(n_used) {
2358 let ex = ex as u16;
2359 selected_blocks.extend([
2360 BlockId::new(il, PROJ_GATE, ex),
2361 BlockId::new(il, PROJ_UP, ex),
2362 BlockId::new(il, PROJ_DOWN, ex),
2363 ]);
2364 }
2365 for &ex in sel_all.iter().take(n_used) {
2366 Self::moe_prefetch_disk_expert(e, il, ex as usize, m, max_block, &selected_blocks)?;
2367 }
2368 e.with_moe_cache(max_block, |cache, eng| {
2369 cache.promote_worker_reads_at_safe_boundary(
2370 &selected_blocks,
2371 &selected_blocks,
2372 eng,
2373 )?;
2374 Ok(())
2375 })?;
2376 }
2377
2378 if t > 1 && std::env::var("MEMRA_MOE_STATS").is_ok() {
2381 let mut cnt = vec![0u32; n_expert];
2382 for &s in sel_all.iter() { cnt[s as usize] += 1; }
2383 let total = sel_all.len() as f64;
2384 let mut h = 0.0f64;
2385 let mut active = 0usize;
2386 for &c in &cnt { if c > 0 { active += 1; let p = c as f64 / total; h -= p * p.log2(); } }
2387 let maxc = cnt.iter().copied().max().unwrap_or(0);
2388 println!("moe-stats il={} t={} assignments={} active={}/{} entropy={:.3}b (max {:.3}b) mean_tok_per_active={:.2} max_tok_per_expert={}",
2389 il, t, sel_all.len(), active, n_expert, h, (n_expert as f64).log2(), total / active.max(1) as f64, maxc);
2390 }
2391
2392 let gdec_may_fire = uniform_experts && use_cache && n_used <= 8 && gdec_enabled();
2401 let mut moe_out = if gdec_may_fire {
2402 e.uninit(t * n_embd)?
2403 } else {
2404 e.zeros(t * n_embd)?
2405 };
2406 let cpu_input = if cpu_hybrid {
2409 Some(routed_cpu_input.ok_or("CPU expert routing did not return the MoE input")?)
2410 } else {
2411 None
2412 };
2413
2414 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;
2422 let mut scratch_u: Option<CudaSlice<u8>> = None;
2423 let mut scratch_d: Option<CudaSlice<u8>> = None;
2424 let page_window = moe_page_prefetch_window();
2432
2433 for tok in 0..t {
2436 let sel = &sel_all[tok * n_used..(tok + 1) * n_used];
2437 let w = &w_all[tok * n_used..(tok + 1) * n_used];
2438 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd); let mut tok_q8: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
2440
2441 let no_macros = m.gate_exps.macros.is_none() && m.up_exps.macros.is_none()
2455 && m.down_exps.macros.is_none();
2456 if gdec_may_fire && moe_q8 && cfg.m3.is_none() && no_macros {
2457 if tok_q8.is_none() {
2458 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
2459 }
2460 let (zq, zd) = tok_q8.as_ref().unwrap();
2461 if Self::moe_gdec_token_q8(e, m, il, max_block, zq, zd, sel, w,
2462 &mut moe_out, tok, n_embd, n_ff_exp, n_used)? {
2463 continue;
2464 }
2465 } else if gdec_may_fire && cfg.m3.is_none() && no_macros
2466 && Self::moe_gdec_token(e, m, il, max_block, &zt, sel, w,
2467 &mut moe_out, tok, n_embd, n_ff_exp, n_used)? {
2468 continue;
2469 }
2470
2471 if gdec_may_fire {
2475 let mut row = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
2476 e.memset_zeros_view(&mut row)?;
2477 }
2478
2479 let mut cpu_mask = vec![false; sel.len()];
2485 let cpu_worker = if let Some(host_input) = cpu_input.as_ref() {
2486 let gpu_resident = if use_cache {
2487 e.with_moe_cache(max_block, |cache, _| {
2488 Ok(sel
2489 .iter()
2490 .map(|&expert| {
2491 let expert = expert as u16;
2492 [PROJ_GATE, PROJ_UP, PROJ_DOWN]
2493 .into_iter()
2494 .filter(|&projection| {
2495 cache
2496 .resident(BlockId::new(il, projection, expert))
2497 .is_some()
2498 })
2499 .count()
2500 })
2501 .collect::<Vec<_>>())
2502 })?
2503 } else {
2504 vec![0; sel.len()]
2505 };
2506 let mut cpu_selected = Vec::new();
2507 for (index, (&expert, &route_weight)) in sel.iter().zip(w).enumerate() {
2508 if gpu_resident[index] != 3 {
2509 cpu_mask[index] = true;
2510 crate::cpu_experts::record_incomplete_gpu_residency(gpu_resident[index]);
2511 let expert = expert as usize;
2512 cpu_selected.push((expert, route_weight));
2513 }
2514 }
2515 if crate::cpu_experts::predictor_enabled() {
2516 let row = &host_input[tok * n_embd..(tok + 1) * n_embd];
2520 crate::cpu_experts::predictor_submit(il, row);
2521 }
2522 if cpu_selected.is_empty() {
2523 None
2524 } else {
2525 let row = &host_input[tok * n_embd..(tok + 1) * n_embd];
2526 let job = crate::cpu_experts::prepare_job(m, il, &cpu_selected, row)
2527 .map_err(std::io::Error::other)?;
2528 Some(crate::cpu_experts::submit(job).map_err(std::io::Error::other)?)
2529 }
2530 } else {
2531 None
2532 };
2533
2534 let worker_window = worker_disk_prefetch
2535 .then(worker_prefetch_window)
2536 .unwrap_or(0);
2537 for (j, &ex) in sel.iter().enumerate() {
2538 if cpu_mask[j] {
2539 continue;
2540 }
2541 let ex = ex as usize;
2542 for next in page_prefetch_positions(j, sel.len(), page_window) {
2543 Self::moe_prefetch_host_expert(sel[next] as usize, m);
2544 }
2545 let keep = [
2546 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_GATE, ex as u16),
2547 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_UP, ex as u16),
2548 crate::moe_cache::BlockId::new(il, crate::moe_cache::PROJ_DOWN, ex as u16),
2549 ];
2550 if worker_disk_prefetch && worker_window > 0 {
2551 for next in worker_prefetch_positions(j, sel.len(), worker_window) {
2552 Self::moe_prefetch_disk_expert(
2553 e,
2554 il,
2555 sel[next] as usize,
2556 m,
2557 max_block,
2558 &keep,
2559 )?;
2560 }
2561 } else if cache_dispatch
2562 && !cpu_hybrid
2563 && moe_prefetch_enabled()
2564 && j + 1 < sel.len()
2565 {
2566 let next = sel[j + 1] as usize;
2567 Self::moe_prefetch_expert(e, il, next, m, max_block, &keep)?;
2568 }
2569 let [gate_q8, up_q8, down_q8] = [moe_q8; 3];
2570 if cache_dispatch && (gate_q8 || up_q8 || down_q8) {
2571 if (gate_q8 || up_q8) && tok_q8.is_none() {
2574 tok_q8 = Some(e.quantize_q8_1_view(&zt, 1, n_embd)?);
2575 }
2576 let gate = if gate_q8 {
2577 let (zq, zd) = tok_q8.as_ref().unwrap();
2578 Self::moe_cached_gemm_q8(e, il, PROJ_GATE, ex, m, max_block, zq, zd)?
2579 } else {
2580 Self::moe_cached_gemm(e, il, PROJ_GATE, ex, m, max_block, &zt)?
2581 };
2582 let up = if up_q8 {
2583 let (zq, zd) = tok_q8.as_ref().unwrap();
2584 Self::moe_cached_gemm_q8(e, il, PROJ_UP, ex, m, max_block, zq, zd)?
2585 } else {
2586 Self::moe_cached_gemm(e, il, PROJ_UP, ex, m, max_block, &zt)?
2587 };
2588 let mut act = e.uninit(n_ff_exp)?;
2589 Self::ffn_act_scaled(
2590 e,
2591 cfg,
2592 &gate,
2593 &up,
2594 m.gate_exps.macro_scale(ex),
2595 m.up_exps.macro_scale(ex),
2596 &mut act,
2597 n_ff_exp,
2598 )?;
2599 let y = if down_q8 {
2600 let (aq2, ad2) = e.quantize_q8_1(&act, 1, n_ff_exp)?;
2601 Self::moe_cached_gemm_q8(e, il, PROJ_DOWN, ex, m, max_block, &aq2, &ad2)?
2602 } else {
2603 let actv = act.slice(0..n_ff_exp);
2604 Self::moe_cached_gemm(e, il, PROJ_DOWN, ex, m, max_block, &actv)?
2605 };
2606 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
2607 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
2609 } else if cache_dispatch {
2610 let gate = Self::moe_cached_gemm(e, il, PROJ_GATE, ex, m, max_block, &zt)?;
2615 let up = Self::moe_cached_gemm(e, il, PROJ_UP, ex, m, max_block, &zt)?;
2616 let mut act = e.uninit(n_ff_exp)?; Self::ffn_act_scaled(e, cfg, &gate, &up,
2618 m.gate_exps.macro_scale(ex), m.up_exps.macro_scale(ex), &mut act, n_ff_exp)?;
2619 let actv = act.slice(0..n_ff_exp);
2620 let y = Self::moe_cached_gemm(e, il, PROJ_DOWN, ex, m, max_block, &actv)?;
2621 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
2622 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
2624 } else if cache_frozen {
2625 let gate = Self::moe_frozen_gemm(
2630 e,
2631 il,
2632 PROJ_GATE,
2633 ex,
2634 m,
2635 max_block,
2636 &zt,
2637 &mut scratch_g,
2638 g_len,
2639 )?;
2640 let up = Self::moe_frozen_gemm(
2641 e,
2642 il,
2643 PROJ_UP,
2644 ex,
2645 m,
2646 max_block,
2647 &zt,
2648 &mut scratch_u,
2649 u_len,
2650 )?;
2651 let mut act = e.uninit(n_ff_exp)?;
2652 Self::ffn_act_scaled(
2653 e,
2654 cfg,
2655 &gate,
2656 &up,
2657 m.gate_exps.macro_scale(ex),
2658 m.up_exps.macro_scale(ex),
2659 &mut act,
2660 n_ff_exp,
2661 )?;
2662 let actv = act.slice(0..n_ff_exp);
2663 let y = Self::moe_frozen_gemm(
2664 e,
2665 il,
2666 PROJ_DOWN,
2667 ex,
2668 m,
2669 max_block,
2670 &actv,
2671 &mut scratch_d,
2672 d_len,
2673 )?;
2674 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
2675 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
2676 } else {
2677 if scratch_g.is_none() {
2681 scratch_g = Some(e.alloc_u8_uninit(g_len)?);
2682 scratch_u = Some(e.alloc_u8_uninit(u_len)?);
2683 scratch_d = Some(e.alloc_u8_uninit(d_len)?);
2684 }
2685 let (sg, su, sd) = (scratch_g.as_mut().unwrap(), scratch_u.as_mut().unwrap(),
2686 scratch_d.as_mut().unwrap());
2687 let gl = m.gate_exps.expert_layout(ex);
2688 let ul = m.up_exps.expert_layout(ex);
2689 let dl = m.down_exps.expert_layout(ex);
2690 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
2691 let gate = e.qmatvec_view(sg, 0..gl.len, &zt, 1,
2692 m.gate_exps.in_f, m.gate_exps.out_f, gl.qtype, gl.row_bytes)?;
2693
2694 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
2695 let up = e.qmatvec_view(su, 0..ul.len, &zt, 1,
2696 m.up_exps.in_f, m.up_exps.out_f, ul.qtype, ul.row_bytes)?;
2697
2698 let mut act = e.uninit(n_ff_exp)?; Self::ffn_act_scaled(e, cfg, &gate, &up,
2700 m.gate_exps.macro_scale(ex), m.up_exps.macro_scale(ex), &mut act, n_ff_exp)?;
2701
2702 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
2703 let actv = act.slice(0..n_ff_exp);
2704 let y = e.qmatvec_view(sd, 0..dl.len, &actv, 1,
2705 m.down_exps.in_f, m.down_exps.out_f, dl.qtype, dl.row_bytes)?;
2706
2707 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
2708 e.axpy_into(&y, w[j] * m.down_exps.macro_scale(ex), &mut dst, n_embd)?;
2709 }
2710 }
2711 if let Some(worker) = cpu_worker {
2712 let cpu_output = worker.wait().map_err(std::io::Error::other)?;
2713 let cpu_output = e.htod(&cpu_output)?;
2714 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
2715 e.axpy_into(&cpu_output, 1.0, &mut dst, n_embd)?;
2716 }
2717 if cpu_hybrid && !cache_frozen && cpu_expert_profile_admit_enabled() {
2718 for (j, &ex) in sel.iter().enumerate() {
2719 if cpu_mask[j] {
2720 Self::moe_profile_admit_expert(e, il, ex as usize, m, max_block)?;
2721 }
2722 }
2723 }
2724 }
2725
2726 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
2731 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
2732 {
2733 let n_ff_sh = gate_shexp.out_features(); let verify_t = t > 1 && t < PRIME_MIN_T;
2742 let (sg_gate, sg_up) = if t == 1 {
2743 match e.matmul_q8_fused2_x(gate_shexp, up_shexp, z)? {
2744 Some(pair) => pair,
2745 None => (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?),
2746 }
2747 } else if verify_t {
2748 (e.matmul_decode_exact(gate_shexp, z, t)?, e.matmul_decode_exact(up_shexp, z, t)?)
2749 } else {
2750 (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?) };
2752 let mut sa = e.uninit(t * n_ff_sh)?; Self::ffn_act(e, cfg, &sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
2754 let sh = if verify_t { e.matmul_decode_exact(down_shexp, &sa, t)? }
2755 else { e.matmul(down_shexp, &sa, t)? }; let g = match &m.gate_inp_shexp {
2769 Some(gate_inp_shexp) => {
2770 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
2771 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
2772 } else {
2773 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
2774 let mut g = e.uninit(t)?; e.sigmoid(&gs, &mut g, t)?;
2776 g
2777 }
2778 }
2779 None => e.htod(&vec![1.0f32; t])?,
2780 };
2781 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
2783 }
2784
2785 Ok(moe_out)
2786 }
2787
2788 pub fn stage1_h2d_per_token(&self) -> u64 {
2791 use crate::hybrid::Ffn;
2792 let n_used = self.cfg.moe.as_ref().map(|m| m.expert_used_count as u64).unwrap_or(0);
2793 let mut bytes = 0u64;
2794 for l in self.layers.iter() {
2795 if let Ffn::Moe(m) = &l.ffn {
2796 bytes += n_used * (m.gate_exps.max_expert_bytes() + m.up_exps.max_expert_bytes()
2797 + m.down_exps.max_expert_bytes()) as u64;
2798 }
2799 }
2800 bytes
2801 }
2802
2803 pub(crate) fn max_moe_block(&self) -> usize {
2807 use crate::hybrid::Ffn;
2808 let mut mx = 0usize;
2809 let mut scan = |ffn: &Ffn| {
2810 if let Ffn::Moe(m) = ffn {
2811 mx = mx.max(m.gate_exps.max_expert_bytes())
2812 .max(m.up_exps.max_expert_bytes())
2813 .max(m.down_exps.max_expert_bytes());
2814 }
2815 };
2816 for l in self.layers.iter() { scan(&l.ffn); }
2817 if let Some(mtp) = self.mtp.as_ref() { scan(&mtp.ffn); }
2818 mx
2819 }
2820
2821 pub(crate) fn moe_cache_block_sizes(&self) -> Vec<usize> {
2824 use crate::hybrid::Ffn;
2825 let mut sizes = Vec::new();
2826 let mut scan = |ffn: &Ffn| {
2827 let Ffn::Moe(m) = ffn else { return };
2828 for ex in 0..m.gate_exps.n_expert {
2829 if m.active_experts.as_ref().is_some_and(|active| !active[ex]) {
2830 continue;
2831 }
2832 for exps in [&m.gate_exps, &m.up_exps, &m.down_exps] {
2833 let len = exps.expert_layout(ex).len;
2834 if len > 0 {
2835 sizes.push(len);
2836 }
2837 }
2838 }
2839 };
2840 for layer in &self.layers {
2841 scan(&layer.ffn);
2842 }
2843 if let Some(mtp) = &self.mtp {
2844 scan(&mtp.ffn);
2845 }
2846 sizes
2847 }
2848
2849 pub fn save_cpu_expert_residency_profile(
2855 &self,
2856 e: &Engine,
2857 path: &std::path::Path,
2858 ) -> Result<(), Box<dyn std::error::Error>> {
2859 let Some(ids) = e.export_moe_residency() else {
2860 return Err("no MoE residency cache to persist".into());
2861 };
2862 let mut body = format!(
2863 "memra-freeze-profile v1 max_block={} blocks={}\n",
2864 self.max_moe_block(),
2865 ids.len()
2866 );
2867 for (layer, proj, ex) in &ids {
2868 body.push_str(&format!("{layer} {proj} {ex}\n"));
2869 }
2870 let tmp = path.with_extension("tmp");
2871 std::fs::write(&tmp, body)?;
2872 std::fs::rename(&tmp, path)?;
2873 println!(
2874 "[moe-cache] freeze profile saved: {} blocks -> {}",
2875 ids.len(),
2876 path.display()
2877 );
2878 Ok(())
2879 }
2880
2881 pub fn restore_cpu_expert_residency_profile(
2885 &self,
2886 e: &Engine,
2887 path: &std::path::Path,
2888 ) -> Result<bool, Box<dyn std::error::Error>> {
2889 use crate::hybrid::Ffn;
2890 use crate::moe_cache::BlockId;
2891 let Ok(content) = std::fs::read_to_string(path) else {
2892 return Ok(false);
2893 };
2894 let mut lines = content.lines();
2895 let Some(header) = lines.next() else { return Ok(false) };
2896 let expected = format!("memra-freeze-profile v1 max_block={}", self.max_moe_block());
2897 if !header.starts_with(&expected) {
2898 println!(
2899 "[moe-cache] freeze profile ignored (geometry mismatch): {}",
2900 path.display()
2901 );
2902 return Ok(false);
2903 }
2904 let mut by_layer: std::collections::HashMap<u16, Vec<BlockId>> =
2905 std::collections::HashMap::new();
2906 for line in lines {
2907 let mut fields = line.split_whitespace();
2908 let (Some(layer), Some(proj), Some(ex)) =
2909 (fields.next(), fields.next(), fields.next())
2910 else {
2911 continue;
2912 };
2913 let (Ok(layer), Ok(proj), Ok(ex)) =
2914 (layer.parse::<u16>(), proj.parse::<u8>(), ex.parse::<u16>())
2915 else {
2916 continue;
2917 };
2918 by_layer
2919 .entry(layer)
2920 .or_default()
2921 .push(BlockId::new(layer, proj, ex));
2922 }
2923 let requested: usize = by_layer.values().map(Vec::len).sum();
2924 if requested == 0 {
2925 return Ok(false);
2926 }
2927 let max_block = self.max_moe_block();
2928 let mut restaged = 0usize;
2929 let mut stage_layer = |layer_index: u16,
2930 ffn: &Ffn|
2931 -> Result<(), Box<dyn std::error::Error>> {
2932 let Ffn::Moe(m) = ffn else { return Ok(()) };
2933 let Some(ids) = by_layer.get(&layer_index) else {
2934 return Ok(());
2935 };
2936 e.with_moe_cache(max_block, |cache, eng| {
2937 for id in ids {
2938 if cache.restage_block(*id, m, eng)? {
2939 restaged += 1;
2940 }
2941 }
2942 Ok(())
2943 })
2944 };
2945 for (index, layer) in self.layers.iter().enumerate() {
2946 stage_layer(index as u16, &layer.ffn)?;
2947 }
2948 if let Some(mtp) = self.mtp.as_ref() {
2949 stage_layer(u16::MAX, &mtp.ffn)?;
2950 }
2951 e.freeze_moe_cache();
2952 println!(
2953 "[moe-cache] freeze profile restored: {restaged}/{requested} blocks restaged from {}",
2954 path.display()
2955 );
2956 Ok(true)
2957 }
2958
2959 pub fn freeze_cpu_expert_residency(
2961 &self,
2962 e: &Engine,
2963 ) -> Result<(), Box<dyn std::error::Error>> {
2964 e.freeze_moe_cache();
2965 Ok(())
2966 }
2967
2968 pub fn ffn_act(e: &Engine, cfg: &ModelConfig, gate: &CudaSlice<f32>, up: &CudaSlice<f32>,
2972 act: &mut CudaSlice<f32>, n: usize) -> Result<(), Box<dyn std::error::Error>> {
2973 Self::ffn_act_scaled(e, cfg, gate, up, 1.0, 1.0, act, n)
2974 }
2975
2976 #[allow(clippy::too_many_arguments)]
2980 pub(crate) fn ffn_act_scaled(e: &Engine, cfg: &ModelConfig, gate: &CudaSlice<f32>, up: &CudaSlice<f32>,
2981 gs: f32, us: f32, act: &mut CudaSlice<f32>, n: usize)
2982 -> Result<(), Box<dyn std::error::Error>> {
2983 if let Some(m3) = cfg.m3.as_ref() {
2984 return e.swigluoai_mul_scaled(gate, up, gs, us, m3.swiglu_alpha, m3.swiglu_limit, act, n);
2985 }
2986 if gs == 1.0 && us == 1.0 { return e.silu_mul(gate, up, act, n); }
2987 e.silu_mul_scaled(gate, up, gs, us, act, n)
2988 }
2989
2990 fn moe_route(e: &Engine, logits: &CudaSlice<f32>, t: usize, n_expert: usize, n_used: usize)
2996 -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
2997 Self::moe_route_cfg(e, logits, t, n_expert, n_used, None, None, None)
2998 }
2999
3000 fn moe_route_cfg(e: &Engine, logits: &CudaSlice<f32>, t: usize, n_expert: usize, n_used: usize,
3008 bias: Option<&[f32]>, sig: Option<(f32, bool)>, active: Option<&[bool]>)
3009 -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
3010 if let Some((sf, route_norm)) = sig {
3011 let lg = e.dtoh(logits)?;
3013 return Self::moe_route_sigmoid_host(
3014 &lg, t, n_expert, n_used, bias, sf, route_norm, active,
3015 );
3016 }
3017 if active.is_none() && !matches!(std::env::var("MEMRA_FUSED_ROUTER").as_deref(), Ok("0")) {
3021 return e.moe_router_topk_host(logits, t, n_expert, n_used);
3022 }
3023 let lg = e.dtoh(logits)?; let mut sel = vec![0u32; t * n_used];
3026 let mut w_out = vec![0f32; t * n_used];
3027 for tok in 0..t {
3028 let row = &lg[tok * n_expert..(tok + 1) * n_expert];
3029 let maxl = row.iter().enumerate()
3031 .filter(|(i, _)| active.is_none_or(|mask| mask[*i]))
3032 .map(|(_, &x)| x).fold(f32::NEG_INFINITY, f32::max);
3033 let mut probs = vec![0f32; n_expert];
3034 let mut den = 0f32;
3035 for i in 0..n_expert {
3036 if active.is_some_and(|mask| !mask[i]) { continue; }
3037 let x = (row[i] - maxl).exp(); probs[i] = x; den += x;
3038 }
3039 for p in probs.iter_mut() { *p /= den; }
3040 let mut idx: Vec<usize> = (0..n_expert)
3042 .filter(|&i| active.is_none_or(|mask| mask[i])).collect();
3043 idx.sort_by(|&a, &b| probs[b].total_cmp(&probs[a]).then(a.cmp(&b)));
3044 let sl = &idx[..n_used];
3045 let mut wv: Vec<f32> = sl.iter().map(|&i| probs[i]).collect();
3046 let mut ws: f32 = wv.iter().sum();
3047 ws = ws.max(6.103515625e-5_f32); for x in wv.iter_mut() { *x /= ws; }
3049 for j in 0..n_used {
3050 sel[tok * n_used + j] = sl[j] as u32;
3051 w_out[tok * n_used + j] = wv[j];
3052 }
3053 }
3054 Ok((sel, w_out))
3055 }
3056
3057 #[allow(clippy::too_many_arguments)]
3058 fn moe_route_sigmoid_with_input(
3059 e: &Engine,
3060 logits: &CudaSlice<f32>,
3061 input: &CudaSlice<f32>,
3062 t: usize,
3063 n_expert: usize,
3064 n_used: usize,
3065 bias: Option<&[f32]>,
3066 (sf, route_norm): (f32, bool),
3067 active: Option<&[bool]>,
3068 ) -> Result<(Vec<u32>, Vec<f32>, Vec<f32>), Box<dyn std::error::Error>> {
3069 let (lg, input) = e.dtoh_pair(logits, input)?;
3070 let (sel, w) =
3071 Self::moe_route_sigmoid_host(&lg, t, n_expert, n_used, bias, sf, route_norm, active)?;
3072 Ok((sel, w, input))
3073 }
3074
3075 pub fn start_moe_prefetch_predictor(
3080 &self,
3081 e: &Engine,
3082 cfg: &ModelConfig,
3083 ) -> Result<(), Box<dyn std::error::Error>> {
3084 use crate::hybrid::Ffn;
3085 let Some(sig) = cfg.sigmoid_router() else {
3086 return Err("prefetch predictor requires a sigmoid-router arch".into());
3087 };
3088 let resident: std::collections::HashSet<(u16, u8, u16)> = e
3089 .export_moe_residency()
3090 .ok_or("prefetch predictor needs the frozen MoE residency cache")?
3091 .into_iter()
3092 .collect();
3093 let mut layers = Vec::new();
3094 for (index, layer) in self.layers.iter().enumerate() {
3095 let Ffn::Moe(m) = &layer.ffn else { continue };
3096 let crate::model::GpuTensor::Float { data, .. } = &m.gate_inp else { continue };
3097 let router = e.dtoh(data)?;
3098 let n_expert = m.gate_exps.n_expert;
3099 let n_embd = m.gate_exps.in_f;
3100 if router.len() != n_embd * n_expert {
3101 continue;
3102 }
3103 let build = |exps: &crate::model::HostExps| {
3104 (0..n_expert)
3105 .map(|expert| crate::cpu_experts::predictor_projection(exps, expert))
3106 .collect::<Vec<_>>()
3107 };
3108 layers.push((index as u16, crate::cpu_experts::PredictLayerInit {
3109 router,
3110 bias: m.exp_probs_b.clone(),
3111 active: m.active_experts.clone(),
3112 n_embd,
3113 n_used: cfg
3114 .moe
3115 .as_ref()
3116 .map(|moe| moe.expert_used_count as usize)
3117 .ok_or("prefetch predictor requires MoE config")?,
3118 sig,
3119 weights_n_expert: n_expert,
3120 gate: build(&m.gate_exps),
3121 up: build(&m.up_exps),
3122 down: build(&m.down_exps),
3123 }));
3124 }
3125 crate::cpu_experts::start_prefetch_predictor(layers, resident)
3126 .map_err(|error| error.into())
3127 }
3128
3129 #[allow(clippy::too_many_arguments)]
3132 pub(crate) fn moe_route_sigmoid_host_public(
3133 logits: &[f32],
3134 t: usize,
3135 n_expert: usize,
3136 n_used: usize,
3137 bias: Option<&[f32]>,
3138 sf: f32,
3139 route_norm: bool,
3140 active: Option<&[bool]>,
3141 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
3142 Self::moe_route_sigmoid_host(logits, t, n_expert, n_used, bias, sf, route_norm, active)
3143 }
3144
3145 #[allow(clippy::too_many_arguments)]
3146 fn moe_route_sigmoid_host(
3147 lg: &[f32],
3148 t: usize,
3149 n_expert: usize,
3150 n_used: usize,
3151 bias: Option<&[f32]>,
3152 sf: f32,
3153 route_norm: bool,
3154 active: Option<&[bool]>,
3155 ) -> Result<(Vec<u32>, Vec<f32>), Box<dyn std::error::Error>> {
3156 if lg.len() != t * n_expert {
3157 return Err(format!(
3158 "sigmoid router logits length mismatch: got {}, expected {}",
3159 lg.len(),
3160 t * n_expert,
3161 )
3162 .into());
3163 }
3164 let mut sel = vec![0u32; t * n_used];
3165 let mut w_out = vec![0f32; t * n_used];
3166 for tok in 0..t {
3167 let row = &lg[tok * n_expert..(tok + 1) * n_expert];
3168 let scores: Vec<f32> = row.iter().map(|&x| 1.0 / (1.0 + (-x).exp())).collect();
3169 let selsc: Vec<f32> = match bias {
3171 Some(b) => scores.iter().zip(b).map(|(s, bb)| s + bb).collect(),
3172 None => scores.clone(),
3173 };
3174 let mut idx: Vec<usize> = (0..n_expert)
3175 .filter(|&i| active.is_none_or(|mask| mask[i]))
3176 .collect();
3177 idx.sort_by(|&a, &b| selsc[b].total_cmp(&selsc[a]).then(a.cmp(&b)));
3178 let sl = &idx[..n_used];
3179 let mut wv: Vec<f32> = sl.iter().map(|&i| scores[i]).collect();
3180 if route_norm {
3181 let ws: f32 = wv.iter().sum::<f32>().max(1e-20);
3182 for x in wv.iter_mut() {
3183 *x = *x / ws * sf;
3184 }
3185 } else {
3186 for x in wv.iter_mut() {
3187 *x *= sf;
3188 }
3189 }
3190 for j in 0..n_used {
3191 sel[tok * n_used + j] = sl[j] as u32;
3192 w_out[tok * n_used + j] = wv[j];
3193 }
3194 }
3195 Ok((sel, w_out))
3196 }
3197
3198 fn moe_ffn_pairs(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, logits: &CudaSlice<f32>,
3207 t: usize, cfg: &ModelConfig)
3208 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3209 let moe = cfg.moe.as_ref().unwrap();
3210 let n_embd = cfg.n_embd as usize;
3211 let n_expert = moe.expert_count as usize;
3212 let n_used = moe.expert_used_count as usize;
3213 let n_ff_exp = moe.expert_ff_length as usize;
3214 let dev = m.dev_exps.as_ref().unwrap();
3215 let (rbg_d, rbu_d) = if dev.gu_il {
3217 let sxx = m.gate_exps.row_bytes + m.up_exps.row_bytes; (sxx, sxx)
3218 } else { (m.gate_exps.row_bytes, m.up_exps.row_bytes) };
3219
3220 let (sel_all, w_all) = Self::moe_route(e, logits, t, n_expert, n_used)?;
3221 let n_pairs = t * n_used;
3222 let pair_tok: Vec<i32> = (0..n_pairs).map(|p| (p / n_used) as i32).collect();
3225 let pair_ex: Vec<i32> = sel_all.iter().map(|&x| x as i32).collect();
3226 let pair_w: Vec<f32> = w_all.clone();
3227 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
3228 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
3229 let pt = e.htod_i32(&pair_tok)?;
3230 let px = e.htod_i32(&pair_ex)?;
3231 let pw = e.htod(&pair_w)?;
3232 let toff = e.htod_i32(&tok_off)?;
3233 let tids = e.htod_i32(&tok_ids)?;
3234
3235 let mut by_ex: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
3239 for p in 0..n_pairs { by_ex[pair_ex[p] as usize].push(p as i32); }
3240 let mut ex_ids: Vec<i32> = Vec::new();
3241 let mut ex_off: Vec<i32> = vec![0];
3242 let mut ex_pairs: Vec<i32> = Vec::with_capacity(n_pairs);
3243 for (ex, list) in by_ex.iter().enumerate() {
3244 if list.is_empty() { continue; }
3245 ex_ids.push(ex as i32);
3246 ex_pairs.extend_from_slice(list);
3247 ex_off.push(ex_pairs.len() as i32);
3248 }
3249 let n_active = ex_ids.len();
3250 let exi = e.htod_i32(&ex_ids)?;
3251 let exo = e.htod_i32(&ex_off)?;
3252 let exp_d = e.htod_i32(&ex_pairs)?;
3253 let _ = &px; static MMA_T: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
3274 let mma_t = *MMA_T.get_or_init(|| {
3275 std::env::var("MEMRA_MOE_MMA_T").ok().and_then(|v| v.parse().ok()).unwrap_or(16)
3276 });
3277 let use_mma = std::env::var("MEMRA_MOE_MMA").map(|v| v != "0").unwrap_or(true)
3278 && t >= mma_t
3279 && q8_expert_dec_supported(m.gate_exps.qtype) && q8_expert_dec_supported(m.up_exps.qtype)
3280 && q8_expert_dec_supported(m.down_exps.qtype)
3281 && n_embd % 256 == 0 && n_ff_exp % 256 == 0;
3282 let mma_capable = q8_expert_dec_supported(m.gate_exps.qtype)
3298 && q8_expert_dec_supported(m.up_exps.qtype)
3299 && q8_expert_dec_supported(m.down_exps.qtype)
3300 && n_embd % 256 == 0 && n_ff_exp % 256 == 0;
3301 let f16g_mode = crate::moe_f16g_mode();
3302 let f16g = f16g_mode != 0 && t >= mma_t
3303 && (f16g_mode != 3 || !mma_capable)
3304 && f16g_proj_ok(m.gate_exps.qtype, n_embd)
3305 && f16g_proj_ok(m.up_exps.qtype, n_embd)
3306 && f16g_proj_ok(m.down_exps.qtype, n_ff_exp);
3307 if use_mma || f16g {
3308 let y_down = if f16g {
3316 let csr_tok: Vec<i32> = ex_pairs.iter().map(|&p| p / n_used as i32).collect();
3320 let csr_tok_d = e.htod_i32(&csr_tok)?;
3321 let (z_f16, z_s) = e.moe_f16g_act(z, Some(&csr_tok_d), n_embd, n_pairs)?;
3322 let g_csr = e.moe_f16_grouped(&dev.ptr_row, 0, n_expert, &exi, &ex_off, &exo,
3323 &z_f16, &z_s, n_embd, n_ff_exp, n_active, n_pairs,
3324 m.gate_exps.qtype, rbg_d)?;
3325 let u_csr = e.moe_f16_grouped(&dev.ptr_row, 1, n_expert, &exi, &ex_off, &exo,
3326 &z_f16, &z_s, n_embd, n_ff_exp, n_active, n_pairs,
3327 m.up_exps.qtype, rbu_d)?;
3328 let act_csr = e.moe_pairs_silu_mul(&g_csr, &u_csr, n_pairs * n_ff_exp)?;
3329 let (a_f16, a_s) = e.moe_f16g_act(&act_csr, None, n_ff_exp, n_pairs)?;
3330 let d_csr = e.moe_f16_grouped(&dev.ptr_row, 2, n_expert, &exi, &ex_off, &exo,
3331 &a_f16, &a_s, n_ff_exp, n_embd, n_active, n_pairs,
3332 m.down_exps.qtype, m.down_exps.row_bytes)?;
3333 e.rows_permute(&d_csr, &exp_d, n_pairs, n_embd)?
3334 } else {
3335 let z_scr = e.mmq_iq_quantize_act(z, n_embd, t)?;
3337 let gate = e.mmq_iq_experts(&dev.ptr_row, 0, n_expert, &exi, &exo, &exp_d, &pt, &z_scr,
3338 n_embd, n_ff_exp, n_active, n_pairs, t,
3339 m.gate_exps.qtype, rbg_d)?;
3340 let up = e.mmq_iq_experts(&dev.ptr_row, 1, n_expert, &exi, &exo, &exp_d, &pt, &z_scr,
3341 n_embd, n_ff_exp, n_active, n_pairs, t,
3342 m.up_exps.qtype, rbu_d)?;
3343 let a_scr = if crate::moe_fuse_actq_on() {
3349 e.mmq_iq_fused_act_quant(&gate, &up, n_ff_exp, n_pairs, 0)?
3350 } else {
3351 let act = e.moe_pairs_silu_mul(&gate, &up, n_pairs * n_ff_exp)?;
3352 e.mmq_iq_quantize_act(&act, n_ff_exp, n_pairs)?
3353 };
3354 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
3355 let pself = e.htod_i32(&pair_self)?;
3356 e.mmq_iq_experts(&dev.ptr_row, 2, n_expert, &exi, &exo, &exp_d, &pself, &a_scr,
3357 n_ff_exp, n_embd, n_active, n_pairs, n_pairs,
3358 m.down_exps.qtype, m.down_exps.row_bytes)?
3359 };
3360 let mut moe_out = e.uninit(t * n_embd)?;
3361 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
3362 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
3363 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
3364 {
3365 let n_ff_sh = gate_shexp.out_features();
3366 let sg_gate = e.matmul(gate_shexp, z, t)?;
3367 let sg_up = e.matmul(up_shexp, z, t)?;
3368 let mut sa = e.uninit(t * n_ff_sh)?;
3369 Self::ffn_act(e, cfg, &sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
3370 let sh = e.matmul(down_shexp, &sa, t)?;
3371 let g = match &m.gate_inp_shexp {
3377 Some(gate_inp_shexp) if crate::router_prefill_exact_on() => {
3378 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
3379 }
3380 Some(gate_inp_shexp) => {
3381 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
3382 let mut g = e.uninit(t)?;
3383 e.sigmoid(&gs, &mut g, t)?;
3384 g
3385 }
3386 None => e.htod(&vec![1.0f32; t])?,
3387 };
3388 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
3389 }
3390 return Ok(moe_out);
3391 }
3392
3393 let dec = std::env::var("MEMRA_MOE_DEC").map(|v| v != "0").unwrap_or(true);
3396 let matvec = |proj, exi: &_, exo: &_, exp_d: &_, pt: &_, aq: &_, ad: &_,
3397 inf, outf, qtype, rb| -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3398 let dec = dec && q8_expert_dec_supported(qtype);
3400 if dec { e.moe_pairs_matvec_q8_dec(&dev.ptr_row, proj, exi, exo, exp_d, pt, aq, ad,
3401 inf, outf, n_expert, n_active, n_pairs, qtype, rb) }
3402 else { e.moe_pairs_matvec_q8_em (&dev.ptr_row, proj, exi, exo, exp_d, pt, aq, ad,
3403 inf, outf, n_expert, n_active, n_pairs, qtype, rb) }
3404 };
3405 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
3406 let gate = matvec(0, &exi, &exo, &exp_d, &pt, &zq, &zd,
3407 n_embd, n_ff_exp, m.gate_exps.qtype, rbg_d)?;
3408 let up = matvec(1, &exi, &exo, &exp_d, &pt, &zq, &zd,
3409 n_embd, n_ff_exp, m.up_exps.qtype, rbu_d)?;
3410 let act = e.moe_pairs_silu_mul(&gate, &up, n_pairs * n_ff_exp)?;
3411 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
3412 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
3414 let pself = e.htod_i32(&pair_self)?;
3415 let y_down = matvec(2, &exi, &exo, &exp_d, &pself, &aq2, &ad2,
3416 n_ff_exp, n_embd, m.down_exps.qtype, m.down_exps.row_bytes)?;
3417 let mut moe_out = e.uninit(t * n_embd)?; e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
3419
3420 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
3424 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
3425 {
3426 let n_ff_sh = gate_shexp.out_features();
3427 let sg_gate = e.matmul(gate_shexp, z, t)?;
3428 let sg_up = e.matmul(up_shexp, z, t)?;
3429 let mut sa = e.uninit(t * n_ff_sh)?;
3430 e.silu_mul(&sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
3431 let sh = e.matmul(down_shexp, &sa, t)?;
3432 let g = match &m.gate_inp_shexp {
3437 Some(gate_inp_shexp) if crate::router_prefill_exact_on() => {
3438 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
3439 }
3440 Some(gate_inp_shexp) => {
3441 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
3442 let mut g = e.uninit(t)?;
3443 e.sigmoid(&gs, &mut g, t)?;
3444 g
3445 }
3446 None => e.htod(&vec![1.0f32; t])?,
3447 };
3448 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
3449 }
3450 Ok(moe_out)
3451 }
3452
3453 #[allow(clippy::too_many_arguments)]
3455 #[allow(clippy::too_many_arguments)]
3456 fn moe_ffn_dev(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>,
3457 zq8: Option<&(CudaSlice<i8>, CudaSlice<f32>)>, logits: &CudaSlice<f32>,
3458 t: usize, cfg: &ModelConfig, il: u16, max_block: usize)
3459 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3460 let moe = cfg.moe.as_ref().unwrap();
3461 let n_embd = cfg.n_embd as usize;
3462 let n_expert = moe.expert_count as usize;
3463 let n_used = moe.expert_used_count as usize;
3464 let n_ff_exp = moe.expert_ff_length as usize;
3465
3466 let (sel_d, mut w_d) = e.moe_router_topk(logits, t, n_expert, n_used)?;
3468 if m.has_macros {
3471 e.moe_w_scale_by_expert(&mut w_d, &sel_d, &m.dev_macros, n_expert, t * n_used)?;
3472 }
3473
3474 let mut moe_out = e.uninit(t * n_embd)?;
3476
3477 if let Some(dev) = m.dev_exps.as_ref() {
3480 let (rbg_d, rbu_d) = if dev.gu_il {
3483 let sxx = m.gate_exps.row_bytes + m.up_exps.row_bytes; (sxx, sxx)
3484 } else { (m.gate_exps.row_bytes, m.up_exps.row_bytes) };
3485 let q8 = moe_q8_enabled()
3486 && q8_expert_supported(m.gate_exps.qtype) && q8_expert_supported(m.up_exps.qtype)
3487 && q8_expert_supported(m.down_exps.qtype);
3488 let rows_arm = q8 && t > 1 && crate::spec::spec_m2()
3497 && n_ff_exp == 512 && n_used <= 8
3498 && std::env::var("MEMRA_MOE_DEVQ8_GU").map(|v| v.is_empty() || v == "v").unwrap_or(true)
3499 && std::env::var("MEMRA_MOE_DEVQ8_DOWN").map(|v| v.is_empty() || v == "w8h2v").unwrap_or(true);
3500 let csr_mode = std::env::var("MEMRA_MOE_CSR").ok()
3509 .and_then(|v| v.parse::<i32>().ok()).unwrap_or(1);
3510 let csr_qt = |qt: i32| qt == crate::QT_IQ4_XS || qt == crate::QT_IQ3_S;
3511 let csr_arm = rows_arm && csr_mode > 0 && t <= 10
3512 && csr_qt(m.gate_exps.qtype) && csr_qt(m.up_exps.qtype)
3513 && csr_qt(m.down_exps.qtype);
3514 if csr_arm {
3515 if csr_mode == 2 {
3516 static ENGAGED: std::sync::Once = std::sync::Once::new();
3517 ENGAGED.call_once(|| eprintln!("[memra] moe CSR byte-compare mode ON (t={t})"));
3518 }
3519 let n_pairs = t * n_used;
3520 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
3521 let act = e.moe_gate_up_silu8_dev_q8_csr(&dev.ptr_row, &sel_d, &zq, &zd, n_pairs,
3522 n_embd, n_ff_exp, n_used, n_expert,
3523 m.gate_exps.qtype, m.up_exps.qtype,
3524 rbg_d, rbu_d)?;
3525 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
3526 e.moe_down8_fma_dev_q8_rows(&dev.ptr_row, &sel_d, &w_d, &aq2, &ad2, &mut moe_out,
3530 t, n_ff_exp, n_embd, n_used, n_expert,
3531 m.down_exps.qtype, m.down_exps.row_bytes)?;
3532 if csr_mode == 2 {
3533 let act_r = e.moe_gate_up_silu8_dev_q8_rows(&dev.ptr_row, &sel_d, &zq, &zd, t,
3535 n_embd, n_ff_exp, n_used, n_expert,
3536 m.gate_exps.qtype, m.up_exps.qtype,
3537 rbg_d, rbu_d, &m.dev_macros)?;
3538 let mut out_r = e.uninit(t * n_embd)?;
3539 let (aq2r, ad2r) = e.quantize_q8_1(&act_r, n_pairs, n_ff_exp)?;
3540 e.moe_down8_fma_dev_q8_rows(&dev.ptr_row, &sel_d, &w_d, &aq2r, &ad2r, &mut out_r,
3541 t, n_ff_exp, n_embd, n_used, n_expert,
3542 m.down_exps.qtype, m.down_exps.row_bytes)?;
3543 let (a1, a2) = (e.dtoh(&act)?, e.dtoh(&act_r)?);
3544 let (o1, o2) = (e.dtoh(&moe_out)?, e.dtoh(&out_r)?);
3545 let ba = a1.iter().zip(&a2).filter(|(x, y)| x.to_bits() != y.to_bits()).count();
3546 let bo = o1.iter().zip(&o2).filter(|(x, y)| x.to_bits() != y.to_bits()).count();
3547 if ba + bo > 0 {
3548 eprintln!("[csr-check] il={il} t={t} ACT diffs={ba}/{} OUT diffs={bo}/{}",
3549 a1.len(), o1.len());
3550 let sel_h = e.dtoh_i32(&sel_d)?;
3552 let mut shown = 0;
3553 for (i, (x, y)) in a1.iter().zip(&a2).enumerate() {
3554 if x.to_bits() != y.to_bits() && shown < 4 {
3555 let (p, o) = (i / n_ff_exp, i % n_ff_exp);
3556 let ex = sel_h[p];
3557 let npx = sel_h.iter().filter(|&&v| v == ex).count();
3558 eprintln!(" ACT p={p} ex={ex} np={npx} o={o} csr={x:e} rows={y:e}");
3559 shown += 1;
3560 }
3561 }
3562 std::process::exit(3);
3563 }
3564 }
3565 } else if rows_arm {
3566 if std::env::var("MEMRA_MOE_OVERLAP").as_deref() == Ok("1") {
3569 use std::sync::atomic::{AtomicU64, Ordering};
3570 static PAIRS: AtomicU64 = AtomicU64::new(0);
3571 static UNIQ: AtomicU64 = AtomicU64::new(0);
3572 static CALLS: AtomicU64 = AtomicU64::new(0);
3573 let sel_h = e.dtoh_i32(&sel_d)?;
3574 let mut u: Vec<i32> = sel_h.clone(); u.sort_unstable(); u.dedup();
3575 PAIRS.fetch_add(sel_h.len() as u64, Ordering::Relaxed);
3576 UNIQ.fetch_add(u.len() as u64, Ordering::Relaxed);
3577 let c = CALLS.fetch_add(1, Ordering::Relaxed) + 1;
3578 if c % 480 == 0 {
3579 let p = PAIRS.load(Ordering::Relaxed); let q = UNIQ.load(Ordering::Relaxed);
3580 eprintln!("[overlap] calls={c} pairs={p} unique={q} ratio={:.3} (t={t})",
3581 q as f64 / p as f64);
3582 }
3583 }
3584 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
3585 let act = e.moe_gate_up_silu8_dev_q8_rows(&dev.ptr_row, &sel_d, &zq, &zd, t,
3586 n_embd, n_ff_exp, n_used, n_expert,
3587 m.gate_exps.qtype, m.up_exps.qtype,
3588 rbg_d, rbu_d, &m.dev_macros)?;
3589 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
3590 e.moe_down8_fma_dev_q8_rows(&dev.ptr_row, &sel_d, &w_d, &aq2, &ad2, &mut moe_out,
3591 t, n_ff_exp, n_embd, n_used, n_expert,
3592 m.down_exps.qtype, m.down_exps.row_bytes)?;
3593 } else {
3594 for tok in 0..t {
3595 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
3596 let selt = sel_d.slice(tok * n_used..(tok + 1) * n_used);
3597 let wt = w_d.slice(tok * n_used..(tok + 1) * n_used);
3598 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
3599 if q8 {
3600 let (zq, zd) = match (t, zq8) {
3601 (1, Some((q, d))) => (q.clone(), d.clone()),
3602 _ => e.quantize_q8_1_view(&zt, 1, n_embd)?,
3603 };
3604 let act = e.moe_gate_up_silu8_dev_q8(&dev.ptr_row, &selt, &zq, &zd,
3605 n_embd, n_ff_exp, n_used, n_expert,
3606 m.gate_exps.qtype, m.up_exps.qtype,
3607 rbg_d, rbu_d, &m.dev_macros)?;
3608 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
3609 e.moe_down8_fma_dev_q8(&dev.ptr_row, &selt, &wt, &aq2, &ad2, &mut dst,
3610 n_ff_exp, n_embd, n_used, n_expert,
3611 m.down_exps.qtype, m.down_exps.row_bytes)?;
3612 } else {
3613 let act = e.moe_gate_up_silu8_dev(&dev.ptr_row, &selt, &zt, n_embd, n_ff_exp,
3614 n_used, n_expert,
3615 m.gate_exps.qtype, m.up_exps.qtype,
3616 rbg_d, rbu_d, &m.dev_macros)?;
3617 e.moe_down8_fma_dev(&dev.ptr_row, &selt, &wt, &act, &mut dst,
3618 n_ff_exp, n_embd, n_used, n_expert,
3619 m.down_exps.qtype, m.down_exps.row_bytes)?;
3620 }
3621 }
3622 }
3623 } else {
3624 let q8 = moe_q8_enabled()
3631 && q8_expert_supported(m.gate_exps.qtype) && q8_expert_supported(m.up_exps.qtype)
3632 && q8_expert_supported(m.down_exps.qtype);
3633 e.with_moe_cache(max_block, |c, eng| {
3634 let row = c.layer_dev_row(il, n_expert, eng)?
3635 .ok_or("moe_ffn_dev: layer row vanished under the lock")?;
3636 for tok in 0..t {
3637 let zt = z.slice(tok * n_embd..(tok + 1) * n_embd);
3638 let selt = sel_d.slice(tok * n_used..(tok + 1) * n_used);
3639 let wt = w_d.slice(tok * n_used..(tok + 1) * n_used);
3640 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
3641 if q8 {
3642 let (zq, zd) = match (t, zq8) {
3643 (1, Some((q, d))) => (q.clone(), d.clone()),
3644 _ => eng.quantize_q8_1_view(&zt, 1, n_embd)?,
3645 };
3646 let act = eng.moe_gate_up_silu8_dev_q8(row, &selt, &zq, &zd,
3647 n_embd, n_ff_exp, n_used, n_expert,
3648 m.gate_exps.qtype, m.up_exps.qtype,
3649 m.gate_exps.row_bytes, m.up_exps.row_bytes,
3650 &m.dev_macros)?;
3651 let (aq2, ad2) = eng.quantize_q8_1(&act, n_used, n_ff_exp)?;
3652 eng.moe_down8_fma_dev_q8(row, &selt, &wt, &aq2, &ad2, &mut dst,
3653 n_ff_exp, n_embd, n_used, n_expert,
3654 m.down_exps.qtype, m.down_exps.row_bytes)?;
3655 } else {
3656 let act = eng.moe_gate_up_silu8_dev(row, &selt, &zt, n_embd, n_ff_exp,
3657 n_used, n_expert,
3658 m.gate_exps.qtype, m.up_exps.qtype,
3659 m.gate_exps.row_bytes, m.up_exps.row_bytes,
3660 &m.dev_macros)?;
3661 eng.moe_down8_fma_dev(row, &selt, &wt, &act, &mut dst,
3662 n_ff_exp, n_embd, n_used, n_expert,
3663 m.down_exps.qtype, m.down_exps.row_bytes)?;
3664 }
3665 }
3666 c.hits += (t * 3 * n_used) as u64;
3668 Ok(())
3669 })?;
3670 }
3671
3672 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
3677 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
3678 {
3679 let n_ff_sh = gate_shexp.out_features();
3680 let verify_t = t > 1 && t < PRIME_MIN_T;
3683 let (sg_gate, sg_up) = if t == 1 {
3684 match e.matmul_q8_fused2_x(gate_shexp, up_shexp, z)? {
3685 Some(pair) => pair,
3686 None => (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?),
3687 }
3688 } else if verify_t {
3689 let mut fused = None;
3693 if crate::spec::spec_fused_t() && (2..=4).contains(&t)
3694 && e.uses_q8_1_fast(gate_shexp) && e.uses_q8_1_fast(up_shexp) {
3695 let (zq, zd) = e.quantize_q8_1(z, t, n_embd)?;
3696 fused = e.matmul_q8_fused2_t(gate_shexp, up_shexp, &zq, &zd, t)?;
3697 }
3698 match fused {
3699 Some(pair) => pair,
3700 None => (e.matmul_decode_exact(gate_shexp, z, t)?,
3701 e.matmul_decode_exact(up_shexp, z, t)?),
3702 }
3703 } else {
3704 (e.matmul(gate_shexp, z, t)?, e.matmul(up_shexp, z, t)?)
3705 };
3706 let mut sa = e.uninit(t * n_ff_sh)?; e.silu_mul(&sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
3708 let sh = if verify_t { e.matmul_decode_exact(down_shexp, &sa, t)? }
3709 else { e.matmul(down_shexp, &sa, t)? };
3710 let g = match &m.gate_inp_shexp {
3714 Some(gate_inp_shexp) => {
3715 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
3718 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
3719 } else {
3720 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
3721 let mut g = e.uninit(t)?;
3722 e.sigmoid(&gs, &mut g, t)?;
3723 g
3724 }
3725 }
3726 None => e.htod(&vec![1.0f32; t])?,
3727 };
3728 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
3729 }
3730
3731 Ok(moe_out)
3732 }
3733
3734 #[allow(clippy::too_many_arguments)]
3744 #[allow(clippy::too_many_arguments)]
3747 fn moe_gdec_token_q8(e: &Engine, m: &MoeWeights, il: u16, max_block: usize,
3748 zq: &CudaSlice<i8>, zd: &CudaSlice<f32>, sel: &[u32], w: &[f32],
3749 moe_out: &mut CudaSlice<f32>, tok: usize,
3750 n_embd: usize, n_ff_exp: usize, n_used: usize)
3751 -> Result<bool, Box<dyn std::error::Error>> {
3752 use crate::moe_cache::{BlockId, PROJ_GATE, PROJ_UP, PROJ_DOWN};
3753 use cudarc::driver::DevicePtr;
3754 let ptrs = e.with_moe_cache(max_block, |c, eng| {
3755 let mut g = [0u64; 8];
3756 let mut u = [0u64; 8];
3757 let mut d = [0u64; 8];
3758 for (j, &ex) in sel.iter().enumerate() {
3759 let ex = ex as u16;
3760 let (Some(sg), Some(su), Some(sd)) = (c.resident(BlockId::new(il, PROJ_GATE, ex)),
3761 c.resident(BlockId::new(il, PROJ_UP, ex)),
3762 c.resident(BlockId::new(il, PROJ_DOWN, ex)))
3763 else { return Ok(None); };
3764 let __s = eng.stream();
3765 let (pg, _e0) = c.slot(sg).device_ptr(&__s);
3766 let (pu, _e1) = c.slot(su).device_ptr(&__s);
3767 let (pd, _e2) = c.slot(sd).device_ptr(&__s);
3768 g[j] = pg as u64; u[j] = pu as u64; d[j] = pd as u64;
3769 }
3770 if cpu_expert_profile_admit_enabled() && !c.is_frozen() {
3771 for &ex in sel {
3772 let ex = ex as u16;
3773 for proj in [PROJ_GATE, PROJ_UP, PROJ_DOWN] {
3774 c.note_profile_hit(BlockId::new(il, proj, ex));
3775 }
3776 }
3777 }
3778 c.hits += (3 * n_used) as u64;
3779 Ok(Some((g, u, d)))
3780 })?;
3781 let Some((g, u, d)) = ptrs else { return Ok(false) };
3782 let mut wv = [0f32; 8];
3783 wv[..n_used].copy_from_slice(w);
3784 let act = e.moe_gate_up_silu8_q8(crate::WPtr8(g), crate::WPtr8(u), zq, zd,
3785 n_embd, n_ff_exp, n_used,
3786 m.gate_exps.qtype, m.up_exps.qtype,
3787 m.gate_exps.row_bytes, m.up_exps.row_bytes)?;
3788 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
3790 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
3791 e.moe_down8_fma_q8(crate::WPtr8(d), crate::F32x8(wv), &aq2, &ad2, &mut dst,
3792 n_ff_exp, n_embd, n_used,
3793 m.down_exps.qtype, m.down_exps.row_bytes)?;
3794 Ok(true)
3795 }
3796
3797 fn moe_gdec_token(e: &Engine, m: &MoeWeights, il: u16, max_block: usize,
3798 zt: &cudarc::driver::CudaView<f32>, sel: &[u32], w: &[f32],
3799 moe_out: &mut CudaSlice<f32>, tok: usize,
3800 n_embd: usize, n_ff_exp: usize, n_used: usize)
3801 -> Result<bool, Box<dyn std::error::Error>> {
3802 use crate::moe_cache::{BlockId, PROJ_GATE, PROJ_UP, PROJ_DOWN};
3803 use cudarc::driver::DevicePtr;
3804 let ptrs = e.with_moe_cache(max_block, |c, eng| {
3806 let mut g = [0u64; 8];
3807 let mut u = [0u64; 8];
3808 let mut d = [0u64; 8];
3809 for (j, &ex) in sel.iter().enumerate() {
3810 let ex = ex as u16;
3811 let (Some(sg), Some(su), Some(sd)) = (c.resident(BlockId::new(il, PROJ_GATE, ex)),
3812 c.resident(BlockId::new(il, PROJ_UP, ex)),
3813 c.resident(BlockId::new(il, PROJ_DOWN, ex)))
3814 else { return Ok(None); };
3815 let __s = eng.stream();
3816 let (pg, _e0) = c.slot(sg).device_ptr(&__s);
3817 let (pu, _e1) = c.slot(su).device_ptr(&__s);
3818 let (pd, _e2) = c.slot(sd).device_ptr(&__s);
3819 g[j] = pg as u64; u[j] = pu as u64; d[j] = pd as u64;
3820 }
3821 if cpu_expert_profile_admit_enabled() && !c.is_frozen() {
3822 for &ex in sel {
3823 let ex = ex as u16;
3824 for proj in [PROJ_GATE, PROJ_UP, PROJ_DOWN] {
3825 c.note_profile_hit(BlockId::new(il, proj, ex));
3826 }
3827 }
3828 }
3829 c.hits += (3 * n_used) as u64; Ok(Some((g, u, d)))
3831 })?;
3832 let Some((g, u, d)) = ptrs else { return Ok(false) };
3833 let mut wv = [0f32; 8];
3834 wv[..n_used].copy_from_slice(w);
3835 let act = e.moe_gate_up_silu8(crate::WPtr8(g), crate::WPtr8(u), zt,
3837 n_embd, n_ff_exp, n_used,
3838 m.gate_exps.qtype, m.up_exps.qtype,
3839 m.gate_exps.row_bytes, m.up_exps.row_bytes)?;
3840 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
3841 e.moe_down8_fma_into(crate::WPtr8(d), crate::F32x8(wv), &act, &mut dst,
3842 n_ff_exp, n_embd, n_used,
3843 m.down_exps.qtype, m.down_exps.row_bytes)?;
3844 Ok(true)
3845 }
3846
3847 fn moe_cached_gemm_q8(e: &Engine, il: u16, proj: u8, ex: usize, m: &MoeWeights,
3852 max_block: usize, aq: &CudaSlice<i8>, ad: &CudaSlice<f32>)
3853 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3854 use crate::moe_cache::{BlockId, DispatchSlot, PROJ_GATE, PROJ_UP};
3855 let exps = match proj { PROJ_GATE => &m.gate_exps, PROJ_UP => &m.up_exps, _ => &m.down_exps };
3856 let layout = exps.expert_layout(ex);
3857 let id = BlockId::new(il, proj, ex as u16);
3858 let source = exps.expert_source(ex);
3859 e.with_moe_cache(max_block, |c, eng| {
3860 let slot = c.dispatch_source(id, source, eng)?;
3861 let DispatchSlot::Resident(sl) = slot;
3862 let buf = c.slot(sl);
3863 eng.qmatvec_expert_q8(buf, 0..layout.len, aq, ad, 1, exps.in_f, exps.out_f,
3864 layout.qtype, layout.row_bytes)
3865 })
3866 }
3867
3868 fn moe_cached_gemm(e: &Engine, il: u16, proj: u8, ex: usize, m: &MoeWeights,
3869 max_block: usize, x: &cudarc::driver::CudaView<f32>)
3870 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3871 use crate::moe_cache::{BlockId, DispatchSlot, PROJ_GATE, PROJ_UP};
3872 let exps = match proj { PROJ_GATE => &m.gate_exps, PROJ_UP => &m.up_exps, _ => &m.down_exps };
3873 let layout = exps.expert_layout(ex);
3874 let id = BlockId::new(il, proj, ex as u16);
3875 let source = exps.expert_source(ex);
3876 e.with_moe_cache(max_block, |c, eng| {
3878 let slot = c.dispatch_source(id, source, eng)?;
3879 let DispatchSlot::Resident(sl) = slot;
3882 let buf = c.slot(sl);
3883 eng.qmatvec_view(buf, 0..layout.len, x, 1, exps.in_f, exps.out_f,
3884 layout.qtype, layout.row_bytes)
3885 })
3886 }
3887
3888 fn moe_profile_admit_expert(
3892 e: &Engine,
3893 il: u16,
3894 ex: usize,
3895 m: &MoeWeights,
3896 max_block: usize,
3897 ) -> Result<(), Box<dyn std::error::Error>> {
3898 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
3899 e.with_moe_cache(max_block, |cache, eng| {
3900 for (proj, exps) in [
3901 (PROJ_GATE, &m.gate_exps),
3902 (PROJ_UP, &m.up_exps),
3903 (PROJ_DOWN, &m.down_exps),
3904 ] {
3905 let id = BlockId::new(il, proj, ex as u16);
3906 let _ = cache.dispatch_source(id, exps.expert_source(ex), eng)?;
3907 }
3908 Ok(())
3909 })
3910 }
3911
3912 #[allow(clippy::too_many_arguments)]
3915 fn moe_frozen_gemm(
3916 e: &Engine,
3917 il: u16,
3918 proj: u8,
3919 ex: usize,
3920 m: &MoeWeights,
3921 max_block: usize,
3922 x: &cudarc::driver::CudaView<f32>,
3923 scratch: &mut Option<CudaSlice<u8>>,
3924 scratch_len: usize,
3925 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
3926 use crate::moe_cache::{BlockId, PROJ_GATE, PROJ_UP};
3927 let exps = match proj {
3928 PROJ_GATE => &m.gate_exps,
3929 PROJ_UP => &m.up_exps,
3930 _ => &m.down_exps,
3931 };
3932 let layout = exps.expert_layout(ex);
3933 let id = BlockId::new(il, proj, ex as u16);
3934 if let Some(output) = e.with_moe_cache(max_block, |cache, eng| {
3935 let Some(slot) = cache.resident(id) else {
3936 return Ok(None);
3937 };
3938 let buf = cache.slot(slot);
3939 Ok(Some(eng.qmatvec_view(
3940 buf,
3941 0..layout.len,
3942 x,
3943 1,
3944 exps.in_f,
3945 exps.out_f,
3946 layout.qtype,
3947 layout.row_bytes,
3948 )?))
3949 })? {
3950 return Ok(output);
3951 }
3952 if scratch.is_none() {
3953 *scratch = Some(e.alloc_u8_uninit(scratch_len)?);
3954 }
3955 let scratch = scratch.as_mut().unwrap();
3956 e.stage_expert(exps.expert_bytes(ex), scratch, 0)?;
3957 e.qmatvec_view(
3958 scratch,
3959 0..layout.len,
3960 x,
3961 1,
3962 exps.in_f,
3963 exps.out_f,
3964 layout.qtype,
3965 layout.row_bytes,
3966 )
3967 }
3968
3969 fn moe_prefetch_expert(
3970 e: &Engine,
3971 il: u16,
3972 ex: usize,
3973 m: &MoeWeights,
3974 max_block: usize,
3975 keep: &[crate::moe_cache::BlockId],
3976 ) -> Result<(), Box<dyn std::error::Error>> {
3977 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
3978 e.with_moe_cache(max_block, |c, eng| {
3979 for (proj, exps) in [(PROJ_GATE, &m.gate_exps), (PROJ_UP, &m.up_exps),
3980 (PROJ_DOWN, &m.down_exps)] {
3981 let id = BlockId::new(il, proj, ex as u16);
3982 let _ = c.prefetch_source(id, exps.expert_source(ex), keep, eng)?;
3983 }
3984 Ok(())
3985 })
3986 }
3987
3988 fn moe_prefetch_disk_expert(e: &Engine, il: u16, ex: usize, m: &MoeWeights,
3991 max_block: usize, keep: &[crate::moe_cache::BlockId])
3992 -> Result<(), Box<dyn std::error::Error>> {
3993 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
3994 e.with_moe_cache(max_block, |c, eng| {
3995 for (proj, exps) in [(PROJ_GATE, &m.gate_exps), (PROJ_UP, &m.up_exps),
3996 (PROJ_DOWN, &m.down_exps)] {
3997 let source = exps.expert_source(ex);
3998 if let crate::model::ExpertSource::Disk { .. } = &source {
3999 let id = BlockId::new(il, proj, ex as u16);
4000 let _ = c.prefetch_source(id, source, keep, eng)?;
4001 }
4002 }
4003 Ok(())
4004 })
4005 }
4006
4007 #[inline]
4008 fn moe_prefetch_host_expert(ex: usize, m: &MoeWeights) {
4009 let _ = m.gate_exps.prefetch_expert_pages(ex);
4010 let _ = m.up_exps.prefetch_expert_pages(ex);
4011 let _ = m.down_exps.prefetch_expert_pages(ex);
4012 }
4013}
4014
4015impl HybridModel {
4032 pub(crate) fn moe_ffn_grouped(e: &Engine, m: &MoeWeights, z: &CudaSlice<f32>, t: usize,
4035 cfg: &ModelConfig, il: u16, _max_block: usize)
4036 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4037 let moe = cfg.moe.as_ref().unwrap();
4038 let n_embd = cfg.n_embd as usize;
4039 let n_expert = moe.expert_count as usize;
4040 let n_used = moe.expert_used_count as usize;
4041 let n_ff_exp = moe.expert_ff_length as usize;
4042
4043 let logits = e.matmul(&m.gate_inp, z, t)?;
4045 let (sel_all, w_all) = if let Some(sig) = cfg.sigmoid_router() {
4046 Self::moe_route_cfg(e, &logits, t, n_expert, n_used,
4047 m.exp_probs_b.as_deref(), Some(sig), m.active_experts.as_deref())?
4048 } else {
4049 Self::moe_route_cfg(e, &logits, t, n_expert, n_used,
4050 None, None, m.active_experts.as_deref())?
4051 };
4052 Self::trace_moe_routes(il, t, &sel_all, &w_all)?;
4053
4054 struct ExpertGroup {
4058 tok_indices: Vec<i32>, slot_indices: Vec<i32>, weights: Vec<f32>, }
4062 let mut groups: Vec<ExpertGroup> = (0..n_expert).map(|_| ExpertGroup {
4063 tok_indices: Vec::new(), slot_indices: Vec::new(), weights: Vec::new(),
4064 }).collect();
4065
4066 for tok in 0..t {
4067 for j in 0..n_used {
4068 let ex = sel_all[tok * n_used + j] as usize;
4069 let w = w_all[tok * n_used + j];
4070 groups[ex].tok_indices.push(tok as i32);
4071 groups[ex].slot_indices.push(j as i32);
4072 groups[ex].weights.push(w);
4073 }
4074 }
4075
4076 let mut slot_buf = e.zeros(t * n_used * n_embd)?;
4079 let mut wbuf = e.zeros(t * n_used)?; let g_len = m.gate_exps.max_expert_bytes();
4083 let u_len = m.up_exps.max_expert_bytes();
4084 let d_len = m.down_exps.max_expert_bytes();
4085 let use_cache = Engine::moe_cache_enabled();
4086 let max_block = _max_block;
4087
4088 let (mut scratch_g, mut scratch_u, mut scratch_d) = if !use_cache {
4090 (Some(e.alloc_u8(g_len)?), Some(e.alloc_u8(u_len)?), Some(e.alloc_u8(d_len)?))
4091 } else {
4092 (None, None, None)
4093 };
4094
4095 let mut order: Vec<usize> =
4106 (0..n_expert).filter(|&ex| !groups[ex].tok_indices.is_empty()).collect();
4107 order.sort_by(|&a, &b| groups[b].tok_indices.len()
4108 .cmp(&groups[a].tok_indices.len()).then(a.cmp(&b)));
4109 let mut m_dist: Vec<usize> = Vec::new(); let page_window = moe_page_prefetch_window();
4111 let worker_disk_prefetch = use_cache && crate::spill_pread::worker_enabled();
4112 if worker_disk_prefetch {
4113 if let Some(first) = grouped_worker_prefetch_position(order.len(), None) {
4114 Self::moe_prefetch_disk_expert(e, il, order[first], m, max_block, &[])?;
4115 }
4116 }
4117 for (order_pos, &ex) in order.iter().enumerate() {
4118 for next in page_prefetch_positions(order_pos, order.len(), page_window) {
4119 Self::moe_prefetch_host_expert(order[next], m);
4120 }
4121 if worker_disk_prefetch {
4122 if let Some(next) = grouped_worker_prefetch_position(order.len(), Some(order_pos)) {
4123 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
4124 let keep = [
4125 BlockId::new(il, PROJ_GATE, ex as u16),
4126 BlockId::new(il, PROJ_UP, ex as u16),
4127 BlockId::new(il, PROJ_DOWN, ex as u16),
4128 ];
4129 Self::moe_prefetch_disk_expert(e, il, order[next], m, max_block, &keep)?;
4130 }
4131 }
4132 let grp = &groups[ex];
4133 let m_e = grp.tok_indices.len();
4134 m_dist.push(m_e);
4135 let gl = m.gate_exps.expert_layout(ex);
4136 let ul = m.up_exps.expert_layout(ex);
4137 let dl = m.down_exps.expert_layout(ex);
4138
4139 let tok_idx_d = e.htod_i32(&grp.tok_indices)?;
4143 let slot_idx_d = e.htod_i32(&grp.slot_indices)?;
4144 let dmac = m.down_exps.macro_scale(ex);
4145 let weight_d = if dmac == 1.0 { e.htod(&grp.weights)? } else {
4146 let scaled: Vec<f32> = grp.weights.iter().map(|&w| w * dmac).collect();
4147 e.htod(&scaled)?
4148 };
4149
4150 let mut gathered = e.zeros(m_e * n_embd)?;
4152 e.gather_rows(z, &tok_idx_d, &mut gathered, n_embd, m_e)?;
4153 let gv = gathered.slice(0..m_e * n_embd);
4154
4155 let y = if use_cache {
4157 use crate::moe_cache::{BlockId, PROJ_GATE, PROJ_UP, PROJ_DOWN};
4158 let gate = e.with_moe_cache(max_block, |c, eng| {
4160 let id = BlockId::new(il, PROJ_GATE, ex as u16);
4161 let slot = c.dispatch_source(id, m.gate_exps.expert_source(ex), eng)?;
4162 let buf = c.buf(slot);
4163 eng.qmatvec_view(buf, 0..gl.len, &gv, m_e,
4164 m.gate_exps.in_f, m.gate_exps.out_f, gl.qtype, gl.row_bytes)
4165 })?;
4166 let up = e.with_moe_cache(max_block, |c, eng| {
4167 let id = BlockId::new(il, PROJ_UP, ex as u16);
4168 let slot = c.dispatch_source(id, m.up_exps.expert_source(ex), eng)?;
4169 let buf = c.buf(slot);
4170 eng.qmatvec_view(buf, 0..ul.len, &gv, m_e,
4171 m.up_exps.in_f, m.up_exps.out_f, ul.qtype, ul.row_bytes)
4172 })?;
4173 let mut act = e.zeros(m_e * n_ff_exp)?;
4175 Self::ffn_act_scaled(e, cfg, &gate, &up,
4176 m.gate_exps.macro_scale(ex), m.up_exps.macro_scale(ex), &mut act, m_e * n_ff_exp)?;
4177 let actv = act.slice(0..m_e * n_ff_exp);
4178 e.with_moe_cache(max_block, |c, eng| {
4179 let id = BlockId::new(il, PROJ_DOWN, ex as u16);
4180 let slot = c.dispatch_source(id, m.down_exps.expert_source(ex), eng)?;
4181 let buf = c.buf(slot);
4182 eng.qmatvec_view(buf, 0..dl.len, &actv, m_e,
4183 m.down_exps.in_f, m.down_exps.out_f, dl.qtype, dl.row_bytes)
4184 })?
4185 } else {
4186 let sg = scratch_g.as_mut().unwrap();
4188 let su = scratch_u.as_mut().unwrap();
4189 let sd = scratch_d.as_mut().unwrap();
4190 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
4191 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
4192 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
4193 let gate = e.qmatvec_view(sg, 0..gl.len, &gv, m_e,
4194 m.gate_exps.in_f, m.gate_exps.out_f, gl.qtype, gl.row_bytes)?;
4195 let up = e.qmatvec_view(su, 0..ul.len, &gv, m_e,
4196 m.up_exps.in_f, m.up_exps.out_f, ul.qtype, ul.row_bytes)?;
4197 let mut act = e.zeros(m_e * n_ff_exp)?;
4199 Self::ffn_act_scaled(e, cfg, &gate, &up,
4200 m.gate_exps.macro_scale(ex), m.up_exps.macro_scale(ex), &mut act, m_e * n_ff_exp)?;
4201 let actv = act.slice(0..m_e * n_ff_exp);
4202 e.qmatvec_view(sd, 0..dl.len, &actv, m_e,
4203 m.down_exps.in_f, m.down_exps.out_f, dl.qtype, dl.row_bytes)?
4204 };
4205
4206 e.scatter_slot(&y, &tok_idx_d, &slot_idx_d, &weight_d,
4208 &mut slot_buf, &mut wbuf, n_embd, n_used, m_e)?;
4209 }
4210
4211 let mut moe_out = e.zeros(t * n_embd)?;
4213 e.reduce_slots(&slot_buf, &wbuf, &mut moe_out, n_embd, n_used, t)?;
4214
4215 if std::env::var("MEMRA_MOE_STATS").is_ok() && !m_dist.is_empty() {
4217 m_dist.sort_unstable();
4218 let active = m_dist.len();
4219 let mean = m_dist.iter().sum::<usize>() as f64 / active as f64;
4220 let median = m_dist[active / 2];
4221 let max_m = *m_dist.last().unwrap();
4222 let min_m = m_dist[0];
4223 let above16 = m_dist.iter().filter(|&&x| x >= 16).count();
4224 println!("moe-grouped il={il} t={t} active={active}/{n_expert} \
4225 m_e: min={min_m} median={median} mean={mean:.1} max={max_m} \
4226 above_gemm_threshold(>=16)={above16}/{active}");
4227 }
4228
4229 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
4233 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
4234 {
4235 let n_ff_sh = gate_shexp.out_features();
4236 let sg_gate = e.matmul(gate_shexp, z, t)?;
4237 let sg_up = e.matmul(up_shexp, z, t)?;
4238 let mut sa = e.zeros(t * n_ff_sh)?;
4239 Self::ffn_act(e, cfg, &sg_gate, &sg_up, &mut sa, t * n_ff_sh)?;
4240 let sh = e.matmul(down_shexp, &sa, t)?;
4241 let g = match &m.gate_inp_shexp {
4245 Some(gate_inp_shexp) => {
4246 if t < PRIME_MIN_T || crate::router_prefill_exact_on() {
4249 e.sigmoid_dot_rows(z, gate_inp_shexp.float_data(), n_embd, t)?
4250 } else {
4251 let gs = e.linear(z, gate_inp_shexp.float_data(), t, n_embd, 1)?;
4252 let mut g = e.uninit(t)?;
4253 e.sigmoid(&gs, &mut g, t)?;
4254 g
4255 }
4256 }
4257 None => e.htod(&vec![1.0f32; t])?,
4258 };
4259 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, t)?;
4260 }
4261
4262 Ok(moe_out)
4263 }
4264
4265 pub(crate) fn moe_ffn_lockstep(
4272 &self,
4273 e: &Engine,
4274 m: &MoeWeights,
4275 zbatch: &CudaSlice<f32>,
4276 mrows: usize,
4277 il: u16,
4278 max_block: usize,
4279 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4280 use crate::moe_cache::{BlockId, PROJ_DOWN, PROJ_GATE, PROJ_UP};
4281 let cfg = &self.cfg;
4282 let moe = cfg.moe.as_ref().unwrap();
4283 let n_embd = cfg.n_embd as usize;
4284 let n_expert = moe.expert_count as usize;
4285 let n_used = moe.expert_used_count as usize;
4286 let n_ff_exp = moe.expert_ff_length as usize;
4287
4288 let logits = e.matmul(&m.gate_inp, zbatch, mrows)?;
4289 let (sel_all, w_all) = if let Some(sig) = cfg.sigmoid_router() {
4290 Self::moe_route_cfg(e, &logits, mrows, n_expert, n_used,
4291 m.exp_probs_b.as_deref(), Some(sig), m.active_experts.as_deref())?
4292 } else {
4293 Self::moe_route_cfg(e, &logits, mrows, n_expert, n_used,
4294 None, None, m.active_experts.as_deref())?
4295 };
4296 Self::trace_moe_routes(il, mrows, &sel_all, &w_all)?;
4297
4298 let resident_expert: Vec<bool> = e.with_moe_cache(max_block, |c, _| {
4300 Ok((0..n_expert)
4301 .map(|ex| {
4302 [PROJ_GATE, PROJ_UP, PROJ_DOWN].into_iter().all(|p| {
4303 c.resident(BlockId::new(il, p, ex as u16)).is_some()
4304 })
4305 })
4306 .collect())
4307 })?;
4308
4309 struct Group {
4310 rows: Vec<i32>,
4311 slots: Vec<i32>,
4312 weights: Vec<f32>,
4313 }
4314 let mut groups: std::collections::HashMap<usize, Group> = Default::default();
4315 let mut cpu_rows: Vec<Vec<(usize, f32)>> = vec![Vec::new(); mrows];
4316 let mut cpu_by_expert: std::collections::HashMap<usize, Vec<(usize, f32)>> =
4317 Default::default();
4318 for row in 0..mrows {
4319 for j in 0..n_used {
4320 let ex = sel_all[row * n_used + j] as usize;
4321 let w = w_all[row * n_used + j];
4322 if resident_expert[ex] {
4323 let group = groups.entry(ex).or_insert_with(|| Group {
4324 rows: Vec::new(),
4325 slots: Vec::new(),
4326 weights: Vec::new(),
4327 });
4328 group.rows.push(row as i32);
4329 group.slots.push(j as i32);
4330 group.weights.push(w);
4331 } else {
4332 crate::cpu_experts::record_incomplete_gpu_residency(0);
4333 cpu_rows[row].push((ex, w));
4334 cpu_by_expert.entry(ex).or_default().push((row, w));
4335 }
4336 }
4337 }
4338
4339 let host_rows = e.dtoh(zbatch)?;
4345 let rows_ok = crate::cpu_experts::rows_supported();
4346 enum CpuPart {
4347 Single { row: usize },
4348 Rows { rows: Vec<usize> },
4349 }
4350 let mut tickets: Vec<(CpuPart, crate::cpu_experts::CpuExpertTicket)> = Vec::new();
4351 let mut rows_served: std::collections::HashSet<(usize, usize)> = Default::default();
4352 if rows_ok {
4353 let mut shared: Vec<(usize, Vec<(usize, f32)>)> = cpu_by_expert
4354 .into_iter()
4355 .filter(|(_, rows)| rows.len() >= 2)
4356 .collect();
4357 shared.sort_by_key(|(ex, _)| *ex);
4358 for (ex, mut row_weights) in shared {
4359 row_weights.sort_by_key(|(row, _)| *row);
4360 let inputs: Vec<(&[f32], f32)> = row_weights
4361 .iter()
4362 .map(|&(row, w)| (&host_rows[row * n_embd..(row + 1) * n_embd], w))
4363 .collect();
4364 let job = crate::cpu_experts::prepare_rows_job(m, ex, &inputs)
4365 .map_err(std::io::Error::other)?;
4366 for &(row, _) in &row_weights {
4367 rows_served.insert((row, ex));
4368 }
4369 tickets.push((
4370 CpuPart::Rows {
4371 rows: row_weights.iter().map(|&(row, _)| row).collect(),
4372 },
4373 crate::cpu_experts::submit_rows(job).map_err(std::io::Error::other)?,
4374 ));
4375 }
4376 }
4377 for (row, selected) in cpu_rows.iter().enumerate() {
4378 let leftover: Vec<(usize, f32)> = selected
4379 .iter()
4380 .copied()
4381 .filter(|&(ex, _)| !rows_served.contains(&(row, ex)))
4382 .collect();
4383 if leftover.is_empty() {
4384 continue;
4385 }
4386 let host_row = &host_rows[row * n_embd..(row + 1) * n_embd];
4387 let job = crate::cpu_experts::prepare_job(m, il, &leftover, host_row)
4388 .map_err(std::io::Error::other)?;
4389 tickets.push((
4390 CpuPart::Single { row },
4391 crate::cpu_experts::submit(job).map_err(std::io::Error::other)?,
4392 ));
4393 }
4394
4395 let mut slot_buf = e.zeros(mrows * n_used * n_embd)?;
4396 let mut wbuf = e.zeros(mrows * n_used)?;
4397 let mut order: Vec<usize> = groups.keys().copied().collect();
4398 order.sort_by(|&a, &b| {
4399 groups[&b].rows.len().cmp(&groups[&a].rows.len()).then(a.cmp(&b))
4400 });
4401 for &ex in &order {
4402 let group = &groups[&ex];
4403 let m_e = group.rows.len();
4404 let gl = m.gate_exps.expert_layout(ex);
4405 let ul = m.up_exps.expert_layout(ex);
4406 let dl = m.down_exps.expert_layout(ex);
4407 let row_idx_d = e.htod_i32(&group.rows)?;
4408 let slot_idx_d = e.htod_i32(&group.slots)?;
4409 let dmac = m.down_exps.macro_scale(ex);
4410 let weight_d = if dmac == 1.0 {
4411 e.htod(&group.weights)?
4412 } else {
4413 let scaled: Vec<f32> = group.weights.iter().map(|&w| w * dmac).collect();
4414 e.htod(&scaled)?
4415 };
4416 let mut gathered = e.zeros(m_e * n_embd)?;
4417 e.gather_rows(zbatch, &row_idx_d, &mut gathered, n_embd, m_e)?;
4418 let gv = gathered.slice(0..m_e * n_embd);
4419 let gate = e.with_moe_cache(max_block, |c, eng| {
4420 let slot = c
4421 .resident(BlockId::new(il, PROJ_GATE, ex as u16))
4422 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
4423 eng.qmatvec_view(c.buf(crate::moe_cache::DispatchSlot::Resident(slot)), 0..gl.len, &gv, m_e,
4424 m.gate_exps.in_f, m.gate_exps.out_f, gl.qtype, gl.row_bytes)
4425 })?;
4426 let up = e.with_moe_cache(max_block, |c, eng| {
4427 let slot = c
4428 .resident(BlockId::new(il, PROJ_UP, ex as u16))
4429 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
4430 eng.qmatvec_view(c.buf(crate::moe_cache::DispatchSlot::Resident(slot)), 0..ul.len, &gv, m_e,
4431 m.up_exps.in_f, m.up_exps.out_f, ul.qtype, ul.row_bytes)
4432 })?;
4433 let mut act = e.zeros(m_e * n_ff_exp)?;
4434 Self::ffn_act_scaled(e, cfg, &gate, &up,
4435 m.gate_exps.macro_scale(ex), m.up_exps.macro_scale(ex), &mut act, m_e * n_ff_exp)?;
4436 let actv = act.slice(0..m_e * n_ff_exp);
4437 let y = e.with_moe_cache(max_block, |c, eng| {
4438 let slot = c
4439 .resident(BlockId::new(il, PROJ_DOWN, ex as u16))
4440 .ok_or("lockstep resident expert vanished (cache not frozen?)")?;
4441 eng.qmatvec_view(c.buf(crate::moe_cache::DispatchSlot::Resident(slot)), 0..dl.len, &actv, m_e,
4442 m.down_exps.in_f, m.down_exps.out_f, dl.qtype, dl.row_bytes)
4443 })?;
4444 e.scatter_slot(&y, &row_idx_d, &slot_idx_d, &weight_d,
4445 &mut slot_buf, &mut wbuf, n_embd, n_used, m_e)?;
4446 }
4447 let mut moe_out = e.zeros(mrows * n_embd)?;
4448 e.reduce_slots(&slot_buf, &wbuf, &mut moe_out, n_embd, n_used, mrows)?;
4449
4450 let mut row_sums: Vec<Option<Vec<f32>>> = vec![None; mrows];
4452 for (part, ticket) in tickets {
4453 let cpu_output = ticket.wait().map_err(std::io::Error::other)?;
4454 let mut add_row = |row: usize, chunk: &[f32]| {
4455 let sum = row_sums[row].get_or_insert_with(|| vec![0.0f32; n_embd]);
4456 for (accumulator, value) in sum.iter_mut().zip(chunk) {
4457 *accumulator += value;
4458 }
4459 };
4460 match part {
4461 CpuPart::Single { row } => add_row(row, &cpu_output),
4462 CpuPart::Rows { rows } => {
4463 for (slot, row) in rows.into_iter().enumerate() {
4464 add_row(row, &cpu_output[slot * n_embd..(slot + 1) * n_embd]);
4465 }
4466 }
4467 }
4468 }
4469 for (row, sum) in row_sums.into_iter().enumerate() {
4470 let Some(sum) = sum else { continue };
4471 let cpu_output = e.htod(&sum)?;
4472 let mut dst = moe_out.slice_mut(row * n_embd..(row + 1) * n_embd);
4473 e.axpy_into(&cpu_output, 1.0, &mut dst, n_embd)?;
4474 }
4475
4476 if let (Some(gate_shexp), Some(up_shexp), Some(down_shexp)) =
4477 (&m.gate_shexp, &m.up_shexp, &m.down_shexp)
4478 {
4479 let n_ff_sh = gate_shexp.out_features();
4480 let sg_gate = e.matmul(gate_shexp, zbatch, mrows)?;
4481 let sg_up = e.matmul(up_shexp, zbatch, mrows)?;
4482 let mut sa = e.zeros(mrows * n_ff_sh)?;
4483 Self::ffn_act(e, cfg, &sg_gate, &sg_up, &mut sa, mrows * n_ff_sh)?;
4484 let sh = e.matmul(down_shexp, &sa, mrows)?;
4485 let g = match &m.gate_inp_shexp {
4488 Some(gate_inp_shexp) => {
4489 e.sigmoid_dot_rows(zbatch, gate_inp_shexp.float_data(), n_embd, mrows)?
4490 }
4491 None => e.htod(&vec![1.0f32; mrows])?,
4492 };
4493 e.add_scaled_rows(&sh, &g, &mut moe_out, n_embd, mrows)?;
4494 }
4495
4496 Ok(moe_out)
4497 }
4498}
4499
4500impl HybridModel {
4506 pub(crate) fn gemma4_geom(&self, il: usize) -> (usize, usize, usize, f32, f32, bool) {
4508 let g = self.cfg.gemma4.as_ref().unwrap();
4509 let swa = g.swa_pattern[il];
4510 let hd = if swa { g.key_length_swa } else { g.key_length_global } as usize;
4511 (hd, g.head_count_kv[il] as usize, self.cfg.n_head as usize,
4515 if swa { g.rope_base_swa } else { g.rope_base_global },
4516 1.0, swa)
4517 }
4518
4519 fn gemma4_suppress(&self, e: &Engine, ld: &mut CudaSlice<f32>, t: usize)
4523 -> Result<(), Box<dyn std::error::Error>> {
4524 if let Some((ids, n)) = self.gemma4_aux.as_ref().and_then(|a| a.suppress_d.as_ref()) {
4525 e.mask_ids_rows(ld, ids, *n, self.output.out_features(), t)?;
4526 }
4527 Ok(())
4528 }
4529
4530 fn gemma4_attn_prime(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
4535 h: &CudaSlice<f32>, pos_d: &CudaSlice<i32>, t: usize,
4536 cache: Option<&mut Cache>)
4537 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4538 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
4539 let eps = self.cfg.rms_eps;
4540 let aux = self.gemma4_aux.as_ref().unwrap();
4541
4542 e.mmq_act_begin();
4545 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)? };
4550
4551 let mut q = e.uninit(t * nh * hd)?;
4552 let mut k = e.uninit(t * nkv * hd)?;
4553 let mut v = e.uninit(t * nkv * hd)?;
4555 static EMIT: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
4559 let emit = t >= 16 && crate::Engine::qkvnorm_w_on_prefill(nh * t + 2 * nkv * t, hd)
4560 && *EMIT.get_or_init(|| std::env::var("MEMRA_FA_EMIT").map(|s| s != "0").unwrap_or(true));
4561 let mut qb = e.alloc_uninit::<u8>(if emit { t * nh * hd * 2 } else { 1 })?;
4562 let mut kb = e.alloc_uninit::<u8>(if emit { t * nkv * hd * 2 } else { 1 })?;
4563 let mut vb = e.alloc_uninit::<u8>(if emit { t * nkv * hd * 2 } else { 1 })?;
4564 let v_f16 = emit && crate::fa_f16pv_on() && match hd {
4567 512 => true,
4568 256 => swa && crate::faw_hp_on() && nh % 2 == 0 && (nh / nkv) % 2 == 0,
4569 _ => false,
4570 };
4571 if emit {
4572 e.rms_norm_qkv_w4b(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
4573 &aux.ones, &mut q, &mut k, &mut v, &mut vb,
4574 hd, nh * t, nkv * t, eps, v_f16)?;
4575 } else {
4576 e.rms_norm_qkv(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
4577 &aux.ones, &mut q, &mut k, &mut v, hd, nh * t, nkv * t, eps)?;
4578 }
4579
4580 let ff = if swa { None } else {
4581 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
4582 };
4583 if emit {
4584 e.rope_neox2_bf16e(&mut q, &mut k, &mut qb, &mut kb, pos_d, hd, hd, nh, nkv, t,
4585 base, 1.0, ff)?;
4586 } else {
4587 e.rope_neox2(&mut q, &mut k, pos_d, hd, hd, nh, nkv, t, base, 1.0, ff)?;
4588 }
4589
4590 if let Some(cache) = cache {
4591 let kvl = cache.kv[il].as_mut().unwrap();
4592 assert_eq!(kvl.len, 0, "gemma4 prime is fresh-prompt only (v0)");
4593 e.append_kv_quantized_rows(&k, &v, &mut kvl.k, &mut kvl.v, kvl.len, t,
4594 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()))?;
4595 kvl.len += t;
4596 }
4597 let mut attn = e.zeros(t * nh * hd)?;
4598 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
4602 if swa && t > win {
4603 if hd == 256 && std::env::var("MEMRA_NOFA").is_err() {
4604 if emit { e.fa_prefill_w_pre(&qb, &kb, &vb, &mut attn, hd, nh, nkv, t, t,
4605 scale, true, win, v_f16)?; }
4606 else { e.fa_prefill_w(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true,
4607 win)?; }
4608 } else {
4609 e.sdpa_naive_w(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true, win)?;
4610 }
4611 } else if hd == 256 && std::env::var("MEMRA_NOFA").is_err() {
4612 e.fa_prefill(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
4613 } else if hd == 512 && std::env::var("MEMRA_NOFA").is_err() {
4614 if emit { e.fa_prefill_hd512_pre(&qb, &kb, &vb, &mut attn, hd, nh, nkv, t, t,
4615 scale, true, v_f16)?; }
4616 else { e.fa_prefill_hd512(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?; }
4617 } else {
4618 e.sdpa_naive(&q, &k, &v, &mut attn, hd, nh, nkv, t, t, scale, true)?;
4619 }
4620 Ok(e.matmul(&fa.wo, &attn, t)?)
4621 }
4622
4623 fn gemma4_attn(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
4625 h: &CudaSlice<f32>, pos_d: &CudaSlice<i32>, t: usize)
4626 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4627 self.gemma4_attn_prime(e, fa, il, h, pos_d, t, None)
4628 }
4629
4630 fn gemma4_moe_q8(&self, e: &Engine, m: &crate::hybrid::MoeWeights,
4635 bits: &crate::hybrid::Gemma4MoeBits,
4636 mq: &(CudaSlice<i8>, CudaSlice<f32>),
4637 router_in: &CudaSlice<f32>, t: usize)
4638 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4639 let cfg = &self.cfg;
4640 let moe = cfg.moe.as_ref().unwrap();
4641 let n_embd = cfg.n_embd as usize;
4642 let n_expert = moe.expert_count as usize;
4643 let n_used = moe.expert_used_count as usize;
4644 let n_ff_exp = moe.expert_ff_length as usize;
4645 let logits = if crate::router_kernel_on() {
4649 e.router_gemv(m.gate_inp.float_data(), router_in, n_embd, n_expert, t)?
4650 } else {
4651 e.matmul(&m.gate_inp, router_in, t)?
4652 };
4653 let dev = m.dev_exps.as_ref().unwrap();
4654 let (sel_d, w_d) = e.moe_router_topk_scaled(&logits, t, n_expert, n_used,
4655 &bits.per_expert_scale_d)?;
4656 let (zq, zd) = mq;
4657 if t == 1 {
4658 let selv = sel_d.slice(0..n_used);
4659 let wv = w_d.slice(0..n_used);
4660 let act = e.moe_gate_up_gelu8_dev_q8(&dev.ptr_row, &selv, zq, zd,
4661 n_embd, n_ff_exp, n_used, n_expert,
4662 m.gate_exps.qtype, m.up_exps.qtype,
4663 m.gate_exps.row_bytes, m.up_exps.row_bytes)?;
4664 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
4665 let mut moe_out = e.uninit(n_embd)?;
4666 e.moe_down8_fma_dev_q8(&dev.ptr_row, &selv, &wv, &aq2, &ad2,
4667 &mut moe_out.slice_mut(0..n_embd), n_ff_exp, n_embd,
4668 n_used, n_expert, m.down_exps.qtype, m.down_exps.row_bytes)?;
4669 return Ok(moe_out);
4670 }
4671 let csr = t <= 10 && std::env::var("MEMRA_GEMMA_CSR").as_deref() != Ok("0");
4672 let act = if csr {
4673 e.moe_gate_up_gelu8_dev_q8_csr(&dev.ptr_row, &sel_d, zq, zd, t * n_used,
4674 n_embd, n_ff_exp, n_used, n_expert,
4675 m.gate_exps.qtype, m.up_exps.qtype,
4676 m.gate_exps.row_bytes, m.up_exps.row_bytes)?
4677 } else {
4678 e.moe_gate_up_gelu8_dev_q8_rows(&dev.ptr_row, &sel_d, zq, zd, t,
4679 n_embd, n_ff_exp, n_used, n_expert,
4680 m.gate_exps.qtype, m.up_exps.qtype,
4681 m.gate_exps.row_bytes, m.up_exps.row_bytes)?
4682 };
4683 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
4684 let mut moe_out = e.uninit(t * n_embd)?;
4685 e.moe_down8_fma_dev_q8_rows_g(&dev.ptr_row, &sel_d, &w_d, &aq2, &ad2, &mut moe_out, t,
4688 n_ff_exp, n_embd, n_used, n_expert,
4689 m.down_exps.qtype, m.down_exps.row_bytes)?;
4690 Ok(moe_out)
4691 }
4692
4693 fn gemma4_moe(&self, e: &Engine, m: &crate::hybrid::MoeWeights,
4697 bits: &crate::hybrid::Gemma4MoeBits, moe_in: &CudaSlice<f32>,
4698 router_in: &CudaSlice<f32>, t: usize)
4699 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4700 let cfg = &self.cfg;
4701 let moe = cfg.moe.as_ref().unwrap();
4702 let n_embd = cfg.n_embd as usize;
4703 let n_expert = moe.expert_count as usize;
4704 let n_used = moe.expert_used_count as usize;
4705 let n_ff_exp = moe.expert_ff_length as usize;
4706
4707 let logits = if t < PRIME_MIN_T && crate::router_kernel_on() {
4711 e.router_gemv(m.gate_inp.float_data(), router_in, n_embd, n_expert, t)?
4712 } else {
4713 e.matmul(&m.gate_inp, router_in, t)?
4714 };
4715
4716 if t < PRIME_MIN_T && m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
4721 && expert_dp4a_supported(m.gate_exps.qtype) && expert_dp4a_supported(m.up_exps.qtype)
4722 && expert_dp4a_supported(m.down_exps.qtype)
4723 && std::env::var("MEMRA_GEMMA_MOE_FAST").as_deref() != Ok("0") {
4724 let dev = m.dev_exps.as_ref().unwrap();
4725 let (sel_d, w_d) = e.moe_router_topk_scaled(&logits, t, n_expert, n_used,
4726 &bits.per_expert_scale_d)?;
4727 if t == 1 {
4728 let (zq, zd) = e.quantize_q8_1(moe_in, 1, n_embd)?;
4729 let selv = sel_d.slice(0..n_used);
4730 let wv = w_d.slice(0..n_used);
4731 let act = e.moe_gate_up_gelu8_dev_q8(&dev.ptr_row, &selv, &zq, &zd,
4732 n_embd, n_ff_exp, n_used, n_expert,
4733 m.gate_exps.qtype, m.up_exps.qtype,
4734 m.gate_exps.row_bytes, m.up_exps.row_bytes)?;
4735 let (aq2, ad2) = e.quantize_q8_1(&act, n_used, n_ff_exp)?;
4736 let mut moe_out = e.uninit(n_embd)?;
4737 e.moe_down8_fma_dev_q8(&dev.ptr_row, &selv, &wv, &aq2, &ad2,
4738 &mut moe_out.slice_mut(0..n_embd), n_ff_exp, n_embd,
4739 n_used, n_expert, m.down_exps.qtype, m.down_exps.row_bytes)?;
4740 return Ok(moe_out);
4741 }
4742 let (zq, zd) = e.quantize_q8_1(moe_in, t, n_embd)?;
4747 let csr = t <= 10 && std::env::var("MEMRA_GEMMA_CSR").as_deref() != Ok("0");
4748 let act = if csr {
4749 e.moe_gate_up_gelu8_dev_q8_csr(&dev.ptr_row, &sel_d, &zq, &zd, t * n_used,
4750 n_embd, n_ff_exp, n_used, n_expert,
4751 m.gate_exps.qtype, m.up_exps.qtype,
4752 m.gate_exps.row_bytes, m.up_exps.row_bytes)?
4753 } else {
4754 e.moe_gate_up_gelu8_dev_q8_rows(&dev.ptr_row, &sel_d, &zq, &zd, t,
4755 n_embd, n_ff_exp, n_used, n_expert,
4756 m.gate_exps.qtype, m.up_exps.qtype,
4757 m.gate_exps.row_bytes, m.up_exps.row_bytes)?
4758 };
4759 let (aq2, ad2) = e.quantize_q8_1(&act, t * n_used, n_ff_exp)?;
4760 let mut moe_out = e.uninit(t * n_embd)?;
4761 e.moe_down8_fma_dev_q8_rows_g(&dev.ptr_row, &sel_d, &w_d, &aq2, &ad2, &mut moe_out, t,
4762 n_ff_exp, n_embd, n_used, n_expert,
4763 m.down_exps.qtype, m.down_exps.row_bytes)?;
4764 return Ok(moe_out);
4765 }
4766
4767 let (sel_all, mut w_all) = Self::moe_route(e, &logits, t, n_expert, n_used)?;
4768 for (i, &sx) in sel_all.iter().enumerate() {
4769 w_all[i] *= bits.per_expert_scale[sx as usize];
4770 }
4771
4772 if t >= PRIME_MIN_T && m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
4776 && expert_dp4a_supported(m.gate_exps.qtype) && expert_dp4a_supported(m.up_exps.qtype)
4777 && expert_dp4a_supported(m.down_exps.qtype)
4778 && std::env::var("MEMRA_GEMMA_MOE_PAIRS").as_deref() != Ok("0") {
4779 let dev = m.dev_exps.as_ref().unwrap();
4780 let n_pairs = t * n_used;
4781 let pair_ex: Vec<i32> = sel_all.iter().map(|&x| x as i32).collect();
4782 let pair_tok: Vec<i32> = (0..n_pairs).map(|p| (p / n_used) as i32).collect();
4783 let tok_off: Vec<i32> = (0..=t).map(|tok| (tok * n_used) as i32).collect();
4784 let tok_ids: Vec<i32> = (0..n_pairs as i32).collect();
4785 let pt = e.htod_i32(&pair_tok)?;
4786 let pw = e.htod(&w_all)?;
4787 let toff = e.htod_i32(&tok_off)?;
4788 let tids = e.htod_i32(&tok_ids)?;
4789 let mut by_ex: Vec<Vec<i32>> = vec![Vec::new(); n_expert];
4790 for p in 0..n_pairs { by_ex[pair_ex[p] as usize].push(p as i32); }
4791 let mut ex_ids: Vec<i32> = Vec::new();
4792 let mut ex_off: Vec<i32> = vec![0];
4793 let mut ex_pairs: Vec<i32> = Vec::with_capacity(n_pairs);
4794 for (ex, list) in by_ex.iter().enumerate() {
4795 if list.is_empty() { continue; }
4796 ex_ids.push(ex as i32);
4797 ex_pairs.extend_from_slice(list);
4798 ex_off.push(ex_pairs.len() as i32);
4799 }
4800 let n_active = ex_ids.len();
4801 let exi = e.htod_i32(&ex_ids)?;
4802 let exo = e.htod_i32(&ex_off)?;
4803 let exp_d = e.htod_i32(&ex_pairs)?;
4804 if crate::moe_f16g_gemma_on()
4812 && f16g_proj_ok(m.gate_exps.qtype, n_embd)
4813 && f16g_proj_ok(m.up_exps.qtype, n_embd)
4814 && f16g_proj_ok(m.down_exps.qtype, n_ff_exp) {
4815 let csr_tok: Vec<i32> = ex_pairs.iter().map(|&p| p / n_used as i32).collect();
4816 let csr_tok_d = e.htod_i32(&csr_tok)?;
4817 let (z_f16, z_s) = e.moe_f16g_act(moe_in, Some(&csr_tok_d), n_embd, n_pairs)?;
4818 let g_csr = e.moe_f16_grouped(&dev.ptr_row, 0, n_expert, &exi, &ex_off, &exo,
4819 &z_f16, &z_s, n_embd, n_ff_exp, n_active, n_pairs,
4820 m.gate_exps.qtype, m.gate_exps.row_bytes)?;
4821 let u_csr = e.moe_f16_grouped(&dev.ptr_row, 1, n_expert, &exi, &ex_off, &exo,
4822 &z_f16, &z_s, n_embd, n_ff_exp, n_active, n_pairs,
4823 m.up_exps.qtype, m.up_exps.row_bytes)?;
4824 let act_csr = e.moe_pairs_gelu_mul(&g_csr, &u_csr, n_pairs * n_ff_exp)?;
4825 let (a_f16, a_s) = e.moe_f16g_act(&act_csr, None, n_ff_exp, n_pairs)?;
4826 let d_csr = e.moe_f16_grouped(&dev.ptr_row, 2, n_expert, &exi, &ex_off, &exo,
4827 &a_f16, &a_s, n_ff_exp, n_embd, n_active, n_pairs,
4828 m.down_exps.qtype, m.down_exps.row_bytes)?;
4829 let y_down = e.rows_permute(&d_csr, &exp_d, n_pairs, n_embd)?;
4830 let mut moe_out = e.uninit(t * n_embd)?;
4831 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
4832 if std::env::var("MEMRA_F16G_DEBUG").is_ok() {
4833 let scan = |v: &[f32]| v.iter().filter(|x| !x.is_finite()).count();
4834 let (yd, mo) = (e.dtoh(&y_down)?, e.dtoh(&moe_out)?);
4835 eprintln!("[f16g-debug] post-permute bad={} post-scatter bad={}",
4836 scan(&yd), scan(&mo));
4837 }
4838 return Ok(moe_out);
4839 }
4840 let mma = n_embd % 256 == 0
4843 && std::env::var("MEMRA_GEMMA_MOE_MMA").as_deref() != Ok("0");
4844 let (gate, up) = if mma {
4845 let z_scr = e.mmq_iq_quantize_act(moe_in, n_embd, t)?;
4846 (e.mmq_iq_experts(&dev.ptr_row, 0, n_expert, &exi, &exo, &exp_d, &pt, &z_scr,
4847 n_embd, n_ff_exp, n_active, n_pairs, t,
4848 m.gate_exps.qtype, m.gate_exps.row_bytes)?,
4849 e.mmq_iq_experts(&dev.ptr_row, 1, n_expert, &exi, &exo, &exp_d, &pt, &z_scr,
4850 n_embd, n_ff_exp, n_active, n_pairs, t,
4851 m.up_exps.qtype, m.up_exps.row_bytes)?)
4852 } else {
4853 let (zq, zd) = e.quantize_q8_1(moe_in, t, n_embd)?;
4854 (e.moe_pairs_matvec_q8_dec(&dev.ptr_row, 0, &exi, &exo, &exp_d, &pt, &zq, &zd,
4855 n_embd, n_ff_exp, n_expert, n_active, n_pairs,
4856 m.gate_exps.qtype, m.gate_exps.row_bytes)?,
4857 e.moe_pairs_matvec_q8_dec(&dev.ptr_row, 1, &exi, &exo, &exp_d, &pt, &zq, &zd,
4858 n_embd, n_ff_exp, n_expert, n_active, n_pairs,
4859 m.up_exps.qtype, m.up_exps.row_bytes)?)
4860 };
4861 let pair_self: Vec<i32> = (0..n_pairs as i32).collect();
4862 let pself = e.htod_i32(&pair_self)?;
4863 let y_down = if mma {
4875 let in_pad = n_ff_exp.div_ceil(256) * 256;
4876 let a_scr = if crate::moe_fuse_actq_on() {
4877 e.mmq_iq_fused_act_quant(&gate, &up, n_ff_exp, n_pairs, 1)?
4878 } else {
4879 let act = e.moe_pairs_gelu_mul(&gate, &up, n_pairs * n_ff_exp)?;
4880 e.mmq_iq_quantize_act(&act, n_ff_exp, n_pairs)?
4881 };
4882 e.mmq_iq_experts(&dev.ptr_row, 2, n_expert, &exi, &exo, &exp_d, &pself, &a_scr,
4883 in_pad, n_embd, n_active, n_pairs, n_pairs,
4884 m.down_exps.qtype, m.down_exps.row_bytes)?
4885 } else {
4886 let act = e.moe_pairs_gelu_mul(&gate, &up, n_pairs * n_ff_exp)?;
4887 let (aq2, ad2) = e.quantize_q8_1(&act, n_pairs, n_ff_exp)?;
4888 e.moe_pairs_matvec_q8_dec(&dev.ptr_row, 2, &exi, &exo, &exp_d, &pself, &aq2, &ad2,
4889 n_ff_exp, n_embd, n_expert, n_active, n_pairs,
4890 m.down_exps.qtype, m.down_exps.row_bytes)?
4891 };
4892 let mut moe_out = e.uninit(t * n_embd)?;
4893 e.moe_pairs_scatter(&y_down, &pw, &toff, &tids, &mut moe_out, t, n_embd)?;
4894 return Ok(moe_out);
4895 }
4896
4897 let g_len = m.gate_exps.expert_stride;
4898 let u_len = m.up_exps.expert_stride;
4899 let d_len = m.down_exps.expert_stride;
4900 let dev = m.dev_exps.as_ref().filter(|d| !d.gu_il);
4904 let (mut sg, mut su, mut sd) = if dev.is_some() { (None, None, None) } else {
4905 (Some(e.alloc_u8_uninit(g_len)?), Some(e.alloc_u8_uninit(u_len)?), Some(e.alloc_u8_uninit(d_len)?))
4906 };
4907 let mut moe_out = e.zeros(t * n_embd)?;
4908 for tok in 0..t {
4909 let sel = &sel_all[tok * n_used..(tok + 1) * n_used];
4910 let w = &w_all[tok * n_used..(tok + 1) * n_used];
4911 let zt = moe_in.slice(tok * n_embd..(tok + 1) * n_embd);
4912 for (j, &ex) in sel.iter().enumerate() {
4913 let ex = ex as usize;
4914 let gate = match dev {
4915 Some(d) => e.qmatvec_view(&d.gate, ex * g_len..(ex + 1) * g_len, &zt, 1,
4916 m.gate_exps.in_f, m.gate_exps.out_f, m.gate_exps.qtype, m.gate_exps.row_bytes)?,
4917 None => {
4918 let sg = sg.as_mut().unwrap();
4919 e.stage_expert(m.gate_exps.expert_bytes(ex), sg, 0)?;
4920 e.qmatvec_view(sg, 0..g_len, &zt, 1,
4921 m.gate_exps.in_f, m.gate_exps.out_f, m.gate_exps.qtype, m.gate_exps.row_bytes)?
4922 }
4923 };
4924 let up = match dev {
4925 Some(d) => e.qmatvec_view(&d.up, ex * u_len..(ex + 1) * u_len, &zt, 1,
4926 m.up_exps.in_f, m.up_exps.out_f, m.up_exps.qtype, m.up_exps.row_bytes)?,
4927 None => {
4928 let su = su.as_mut().unwrap();
4929 e.stage_expert(m.up_exps.expert_bytes(ex), su, 0)?;
4930 e.qmatvec_view(su, 0..u_len, &zt, 1,
4931 m.up_exps.in_f, m.up_exps.out_f, m.up_exps.qtype, m.up_exps.row_bytes)?
4932 }
4933 };
4934 let mut act = e.uninit(n_ff_exp)?;
4935 e.gelu_tanh_mul(&gate, &up, &mut act, n_ff_exp)?;
4936 let actv = act.slice(0..n_ff_exp);
4937 let y = match dev {
4938 Some(d) => e.qmatvec_view(&d.down, ex * d_len..(ex + 1) * d_len, &actv, 1,
4939 m.down_exps.in_f, m.down_exps.out_f, m.down_exps.qtype, m.down_exps.row_bytes)?,
4940 None => {
4941 let sd = sd.as_mut().unwrap();
4942 e.stage_expert(m.down_exps.expert_bytes(ex), sd, 0)?;
4943 e.qmatvec_view(sd, 0..d_len, &actv, 1,
4944 m.down_exps.in_f, m.down_exps.out_f, m.down_exps.qtype, m.down_exps.row_bytes)?
4945 }
4946 };
4947 let mut dst = moe_out.slice_mut(tok * n_embd..(tok + 1) * n_embd);
4948 e.axpy_into(&y, w[j], &mut dst, n_embd)?;
4949 }
4950 }
4951 Ok(moe_out)
4952 }
4953
4954 fn gemma4_layer(&self, e: &Engine, il: usize, layer: &crate::hybrid::HybridLayer,
4956 x: &CudaSlice<f32>, pos_d: &CudaSlice<i32>, t: usize)
4957 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4958 let n_embd = self.cfg.n_embd as usize;
4959 let eps = self.cfg.rms_eps;
4960
4961 let mut h = e.zeros(t * n_embd)?;
4962 e.rms_norm(x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
4963 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
4964 let o = self.gemma4_attn(e, fa, il, &h, pos_d, t)?;
4965 let mut cur = e.zeros(t * n_embd)?;
4967 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, t, eps)?;
4968 self.gemma4_layer_tail_add(e, layer, &cur, x, t)
4969 }
4970
4971 fn gemma4_layer_tail_add(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
4975 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize)
4976 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
4977 Ok(self.gemma4_layer_tail_add_n(e, layer, cur, x, t, None)?.0)
4978 }
4979
4980 fn gemma4_layer_tail_add_n(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
4983 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize,
4984 next_norm: Option<&CudaSlice<f32>>)
4985 -> Result<(CudaSlice<f32>, Option<CudaSlice<f32>>), Box<dyn std::error::Error>> {
4986 let n_embd = self.cfg.n_embd as usize;
4987 let bits = layer.gemma4.as_ref().unwrap();
4988 let (sn, attn_out) = self.gemma4_layer_tail_core(e, layer, cur, x, t)?;
4989 let mut xn = e.uninit(t * n_embd)?;
4990 match next_norm {
4991 Some(w) => {
4992 let mut hn = e.uninit(t * n_embd)?;
4993 e.add_scale_rms_norm(&sn, &attn_out, bits.layer_scale, w, &mut xn, &mut hn,
4994 n_embd, t, self.cfg.rms_eps)?;
4995 Ok((xn, Some(hn)))
4996 }
4997 None => {
4998 e.add_scale(&sn, &attn_out, bits.layer_scale, &mut xn, t * n_embd)?;
4999 Ok((xn, None))
5000 }
5001 }
5002 }
5003
5004 fn gemma4_layer_tail_core(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
5007 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize)
5008 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5009 self.gemma4_layer_tail_core_pn(e, layer, cur, x, t, None, false)
5010 }
5011
5012 fn gemma4_layer_tail_core_pn(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
5019 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize,
5020 pre_norm: Option<&CudaSlice<f32>>, defer_post_norm: bool)
5021 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5022 let n_embd = self.cfg.n_embd as usize;
5023 let eps = self.cfg.rms_eps;
5024 let bits = layer.gemma4.as_ref().unwrap();
5025
5026 let Some(mbits) = bits.moe_bits.as_ref() else {
5029 let crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } = &layer.ffn
5030 else { panic!("gemma4 dense layer without Dense ffn") };
5031 let mut attn_out = e.uninit(t * n_embd)?;
5032 let mut zsh = e.uninit(t * n_embd)?;
5033 let mut zpair: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
5036 match pre_norm {
5037 Some(wa) if t == 1 => {
5038 zpair = Some(e.rms_pre_add_rms_norm_q8z(cur, wa, x,
5039 bits.ffn_norm.float_data(),
5040 &mut attn_out, &mut zsh,
5041 n_embd, t, eps)?);
5042 }
5043 Some(wa) => e.rms_pre_add_rms_norm(cur, wa, x, bits.ffn_norm.float_data(),
5044 &mut attn_out, &mut zsh, n_embd, t, eps)?,
5045 None => e.add_rms_norm(cur, x, bits.ffn_norm.float_data(), &mut attn_out,
5046 &mut zsh, n_embd, t, eps)?,
5047 }
5048 let n_ff = ffn_gate.out_features();
5049 let (gate, up) = if t == 1 {
5055 let (zq, zd) = match zpair {
5056 Some(p) => p,
5057 None => e.quantize_q8_1(&zsh, 1, n_embd)?,
5058 };
5059 match e.matmul_q4_fused2(ffn_gate, ffn_up, &zq, &zd)? {
5060 Some(p) => p,
5061 None => (e.matmul_pre(ffn_gate, &zq, &zd, &zsh, 1)?,
5062 e.matmul_pre(ffn_up, &zq, &zd, &zsh, 1)?),
5063 }
5064 } else {
5065 static F2B: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
5070 let f2b = *F2B.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
5071 let fused = if f2b {
5072 let (zq, zd) = e.quantize_q8_1(&zsh, t, n_embd)?;
5073 e.matmul_q4_fused2_batched(ffn_gate, ffn_up, &zq, &zd, t)?
5074 } else { None };
5075 match fused {
5076 Some(p) => p,
5077 None => {
5078 e.mmq_act_begin();
5080 (e.matmul(ffn_gate, &zsh, t)?, e.matmul(ffn_up, &zsh, t)?)
5081 }
5082 }
5083 };
5084 let mut act = e.uninit(t * n_ff)?;
5085 let f0 = if e.uses_q8_1_fast(ffn_down) {
5088 let upv = e.view(&up, t * n_ff);
5089 let up_all = upv.slice(0..t * n_ff);
5090 let (aq, ad) = e.gelu_tanh_mul_q8_1(&gate, &up_all, &mut act, n_ff, t)?;
5091 e.matmul_pre(ffn_down, &aq, &ad, &act, t)?
5092 } else {
5093 e.gelu_tanh_mul(&gate, &up, &mut act, t * n_ff)?;
5094 e.matmul(ffn_down, &act, t)?
5095 };
5096 if defer_post_norm { return Ok((f0, attn_out)); }
5097 let mut sn = e.uninit(t * n_embd)?;
5098 e.rms_norm(&f0, bits.post_ffw_norm.float_data(), &mut sn, n_embd, t, eps)?;
5099 return Ok((sn, attn_out));
5100 };
5101
5102 assert!(pre_norm.is_none(), "pre-norm fold is dense-entry only");
5103 let mut attn_out = e.uninit(t * n_embd)?;
5108 let mut router_in = e.uninit(t * n_embd)?;
5109 let fast_moe = match &layer.ffn {
5110 crate::hybrid::Ffn::Moe(m) => m.dev_exps.as_ref().is_some_and(|d| !d.gu_il)
5111 && expert_dp4a_supported(m.gate_exps.qtype)
5112 && expert_dp4a_supported(m.up_exps.qtype)
5113 && expert_dp4a_supported(m.down_exps.qtype)
5114 && std::env::var("MEMRA_GEMMA_MOE_FAST").as_deref() != Ok("0"),
5115 _ => false,
5116 };
5117 let q8z = t < PRIME_MIN_T && fast_moe;
5118 let (zsh_f32, zsh_q8, moe_q8) = if q8z {
5119 let (z0, m2) = e.add_rms_norm3_q8z(cur, x, bits.ffn_norm.float_data(),
5120 &mbits.router_scale_pre,
5121 mbits.pre_ffw_norm_2.float_data(),
5122 &mut attn_out, &mut router_in, n_embd, t, eps)?;
5123 (None, Some(z0), Some(m2))
5124 } else {
5125 let mut zsh = e.uninit(t * n_embd)?;
5126 let mut moe_in = e.uninit(t * n_embd)?;
5127 e.add_rms_norm3(cur, x, bits.ffn_norm.float_data(), &mbits.router_scale_pre,
5128 mbits.pre_ffw_norm_2.float_data(), &mut attn_out, &mut zsh,
5129 &mut router_in, &mut moe_in, n_embd, t, eps)?;
5130 (Some((zsh, moe_in)), None, None)
5131 };
5132 let attn_out2 = attn_out;
5133 #[allow(unused_variables)]
5134 let attn_out = &attn_out2;
5135 let n_ff = mbits.shared_gate.out_features();
5136 let (gate, up) = if let Some((zq, zd)) = zsh_q8.as_ref() {
5137 if t == 1 {
5138 match e.matmul_q4_fused2(&mbits.shared_gate, &mbits.shared_up, zq, zd)? {
5139 Some(p) => p,
5140 None => {
5141 let h0 = e.zeros(0)?;
5142 (e.matmul_pre(&mbits.shared_gate, zq, zd, &h0, 1)?,
5143 e.matmul_pre(&mbits.shared_up, zq, zd, &h0, 1)?)
5144 }
5145 }
5146 } else {
5147 let h0 = e.zeros(0)?;
5149 (e.matmul_pre(&mbits.shared_gate, zq, zd, &h0, t)?,
5150 e.matmul_pre(&mbits.shared_up, zq, zd, &h0, t)?)
5151 }
5152 } else {
5153 let (zsh, _) = zsh_f32.as_ref().unwrap();
5154 (e.matmul(&mbits.shared_gate, zsh, t)?, e.matmul(&mbits.shared_up, zsh, t)?)
5155 };
5156 let mut act = e.uninit(t * n_ff)?;
5157 e.gelu_tanh_mul(&gate, &up, &mut act, t * n_ff)?;
5158 let mlp0 = e.matmul(&mbits.shared_down, &act, t)?;
5159 let crate::hybrid::Ffn::Moe(m) = &layer.ffn else { panic!("gemma4 layer not MoE") };
5160 let moe0 = match (&moe_q8, &zsh_f32) {
5161 (Some(mq), _) => self.gemma4_moe_q8(e, m, mbits, mq, &router_in, t)?,
5162 (None, Some((_, moe_in))) => self.gemma4_moe(e, m, mbits, moe_in, &router_in, t)?,
5163 _ => unreachable!(),
5164 };
5165 let mut mlp = e.uninit(t * n_embd)?;
5167 let mut moe = e.uninit(t * n_embd)?;
5168 e.rms_norm2x(&mlp0, &moe0, mbits.post_ffw_norm_1.float_data(),
5169 mbits.post_ffw_norm_2.float_data(), &mut mlp, &mut moe, n_embd, t, eps)?;
5170
5171 let mut sum = e.uninit(t * n_embd)?;
5174 let mut sn = e.uninit(t * n_embd)?;
5175 e.add_rms_norm(&mlp, &moe, bits.post_ffw_norm.float_data(), &mut sum, &mut sn,
5176 n_embd, t, eps)?;
5177 Ok((sn, attn_out2))
5178 }
5179
5180 fn gemma4_layer_tail_add_nq(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
5182 cur: &CudaSlice<f32>, x: &CudaSlice<f32>, t: usize,
5183 next_norm: Option<&CudaSlice<f32>>)
5184 -> Result<(CudaSlice<f32>, Option<(CudaSlice<i8>, CudaSlice<f32>)>), Box<dyn std::error::Error>> {
5185 let n_embd = self.cfg.n_embd as usize;
5186 let bits = layer.gemma4.as_ref().unwrap();
5187 let (sn, attn_out) = self.gemma4_layer_tail_core(e, layer, cur, x, t)?;
5188 let mut xn = e.uninit(t * n_embd)?;
5189 match next_norm {
5190 Some(w) => {
5191 let pair = e.add_scale_rms_norm_q8_1(&sn, &attn_out, bits.layer_scale, w, &mut xn,
5192 n_embd, t, self.cfg.rms_eps)?;
5193 Ok((xn, Some(pair)))
5194 }
5195 None => {
5196 e.add_scale(&sn, &attn_out, bits.layer_scale, &mut xn, t * n_embd)?;
5197 Ok((xn, None))
5198 }
5199 }
5200 }
5201
5202 fn gemma4_forward(&self, e: &Engine, tokens: &[u32], last_only: bool)
5205 -> Result<Vec<f32>, Box<dyn std::error::Error>> {
5206 if self.is_gemma4_e4b() { return self.gemma4_e4b_forward(e, tokens, last_only); }
5209 let n_embd = self.cfg.n_embd as usize;
5210 let t = tokens.len();
5211 let pos: Vec<i32> = (0..t as i32).collect();
5212 let pos_d = e.htod_i32(&pos)?;
5213
5214 let mut x = self.embed(e, tokens)?;
5215 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
5216 let probe = std::env::var("MEMRA_GEMMA_PROBE").is_ok();
5219 let stat = |e: &Engine, x: &CudaSlice<f32>, tag: &str| -> Result<(), Box<dyn std::error::Error>> {
5220 let h = e.dtoh(x)?;
5221 let bad = h.iter().filter(|v| !v.is_finite()).count();
5222 let mx = h.iter().filter(|v| v.is_finite()).fold(0.0f32, |m, v| m.max(v.abs()));
5223 eprintln!("[gemma-probe] {tag}: tok0_first3={:?} bad={bad} max={mx:.3e}", &h[..3]);
5224 Ok(())
5225 };
5226 if probe { stat(e, &x, "embed")?; }
5227 for (il, layer) in self.layers.iter().enumerate() {
5228 x = self.gemma4_layer(e, il, layer, &x, &pos_d, t)?;
5229 if probe { stat(e, &x, &format!("L{il}"))?; }
5230 }
5231 let mut hn = e.zeros(t * n_embd)?;
5232 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, self.cfg.rms_eps)?;
5233 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
5234 let n_vocab = self.output.out_features();
5235 let logits = if last_only {
5236 let hv = e.view(&hn, t * n_embd);
5237 let last_row = hv.slice((t - 1) * n_embd..t * n_embd);
5238 let mut hlast = e.zeros(n_embd)?;
5239 e.copy_view_into(&mut hlast, 0, &last_row, n_embd)?;
5240 let mut ld = e.matmul(&self.output, &hlast, 1)?;
5241 e.softcap(&mut ld, cap, n_vocab)?;
5242 self.gemma4_suppress(e, &mut ld, 1)?;
5243 e.dtoh(&ld)?
5244 } else {
5245 let mut ld = e.matmul(&self.output, &hn, t)?;
5246 e.softcap(&mut ld, cap, t * n_vocab)?;
5247 self.gemma4_suppress(e, &mut ld, t)?;
5248 e.dtoh(&ld)?
5249 };
5250 Ok(logits)
5251 }
5252
5253 pub(crate) fn gemma4_prime(&self, e: &Engine, tokens: &[u32], cache: &mut Cache)
5258 -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
5259 assert_eq!(cache.pos, 0, "gemma4 prime v0 is fresh-prompt only");
5260 let n_embd = self.cfg.n_embd as usize;
5261 let eps = self.cfg.rms_eps;
5262 let t = tokens.len();
5263 let pos: Vec<i32> = (0..t as i32).collect();
5264 let pos_d = e.htod_i32(&pos)?;
5265 let mut x = self.embed(e, tokens)?;
5266 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
5267 for (il, layer) in self.layers.iter().enumerate() {
5268 let mut h = e.zeros(t * n_embd)?;
5269 e.rms_norm(&x, layer.attn_norm.float_data(), &mut h, n_embd, t, eps)?;
5270 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer not full-attn") };
5271 let o = self.gemma4_attn_prime(e, fa, il, &h, &pos_d, t, Some(cache))?;
5272 let mut cur = e.zeros(t * n_embd)?;
5273 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, t, eps)?;
5274 x = self.gemma4_layer_tail_add(e, layer, &cur, &x, t)?;
5275 self.dflash_tap(e, cache, il, &x, t)?;
5276 }
5277 cache.pos += t;
5278 let hiddens = e.clone_dtod(&x)?;
5279 let xv = e.view(&x, t * n_embd);
5280 let last_row = xv.slice((t - 1) * n_embd..t * n_embd);
5281 let mut h_seed = e.zeros(n_embd)?;
5282 e.copy_view_into(&mut h_seed, 0, &last_row, n_embd)?;
5283 let mut hn = e.uninit(n_embd)?;
5284 e.rms_norm(&h_seed, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
5285 let mut ld = e.matmul(&self.output, &hn, 1)?;
5286 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
5287 e.softcap(&mut ld, cap, self.output.out_features())?;
5288 self.gemma4_suppress(e, &mut ld, 1)?;
5289 let logits = e.dtoh(&ld)?;
5290 Ok((logits, h_seed, hiddens))
5291 }
5292
5293 fn gemma4_decode_attn(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
5298 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
5299 pos_d: &CudaSlice<i32>, cache: &mut Cache)
5300 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5301 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
5302 let eps = self.cfg.rms_eps;
5303 let aux = self.gemma4_aux.as_ref().unwrap();
5304 let (hq, hdq) = (hq, hdq);
5305 let h0 = e.zeros(0)?;
5306 let h = &h0;
5307 let (q0, k0, v0) = if swa {
5308 match e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, &hq, &hdq)? {
5309 Some(t3) => t3,
5310 None => (e.matmul_pre(&fa.wq, &hq, &hdq, h, 1)?,
5311 e.matmul_pre(&fa.wk, &hq, &hdq, h, 1)?,
5312 e.matmul_pre(&fa.wv, &hq, &hdq, h, 1)?),
5313 }
5314 } else {
5315 let (q0, k0) = match e.matmul_q4_fused2(&fa.wq, &fa.wk, &hq, &hdq)? {
5316 Some(p) => p,
5317 None => (e.matmul_pre(&fa.wq, &hq, &hdq, h, 1)?,
5318 e.matmul_pre(&fa.wk, &hq, &hdq, h, 1)?),
5319 };
5320 let v0 = e.clone_dtod(&k0)?;
5321 (q0, k0, v0)
5322 };
5323 let mut q = e.uninit(nh * hd)?;
5324 let mut k = e.uninit(nkv * hd)?;
5325 let mut v = e.uninit(nkv * hd)?;
5326 let ff = if swa { None } else {
5329 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
5330 };
5331 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
5332 &aux.ones, &mut q, &mut k, &mut v, hd, nh, nkv,
5333 pos_d, nh, nkv, base, 1.0, ff, eps)?;
5334 let kvl = cache.kv[il].as_mut().unwrap();
5335 e.append_kv_quantized(&k, &v, &mut kvl.k, &mut kvl.v, kvl.len,
5336 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()))?;
5337 kvl.len += 1;
5338 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
5342 let mut attn = e.uninit(nh * hd)?;
5343 if !swa && hd == 512 && kvl.len >= crate::fa512_min_tkv()
5345 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
5346 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
5347 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
5348 let base = kvl.len as i32;
5350 e.i32_set_k(&mut kvl.len_d, base)?;
5351 e.fa_decode_rows(&q, &kp, &vp, &mut attn, hd, nh, nkv, kvl.len - 1, 1, scale,
5352 kvl.k_tok_bytes, kvl.v_tok_bytes, Some((&kvl.len_d, -1)), false,
5353 false, None)?;
5354 return Ok(e.matmul(&fa.wo, &attn, 1)?);
5355 }
5356 if swa && kvl.len > win && hd == 256
5358 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
5359 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
5360 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
5361 let base = kvl.len as i32;
5362 e.i32_set_k(&mut kvl.len_d, base)?;
5363 e.fa_decode_rows_w(&q, &kp, &vp, &mut attn, hd, nh, nkv, &kvl.len_d, -1, 1, scale,
5364 win, kvl.k_tok_bytes, kvl.v_tok_bytes, None)?;
5365 return Ok(e.matmul(&fa.wo, &attn, 1)?);
5366 }
5367 let (off_tok, t_kv) = if swa && kvl.len > win { (kvl.len - win, win) } else { (0, kvl.len) };
5368 let k_view = e.view_u8_range(&kvl.k, off_tok * kvl.k_tok_bytes,
5369 (off_tok + t_kv) * kvl.k_tok_bytes);
5370 let v_view = e.view_u8_range(&kvl.v, off_tok * kvl.v_tok_bytes,
5371 (off_tok + t_kv) * kvl.v_tok_bytes);
5372 e.fa_decode_kvmod(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, t_kv, scale,
5373 kvl.k_tok_bytes, kvl.v_tok_bytes, swa && crate::Engine::wkv_on())?;
5374 Ok(e.matmul(&fa.wo, &attn, 1)?)
5375 }
5376
5377 #[allow(clippy::too_many_arguments)]
5384 pub fn gemma4_decode_step_dc(&self, e: &Engine, token_d: &CudaSlice<u32>,
5385 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
5386 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
5387 n_vocab: usize, cap_bucket_max: Option<(usize, usize)>)
5388 -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
5389 let mut tok_out = e.stream().alloc_zeros::<u32>(1)?;
5390 self.gemma4_decode_step_dc_into(e, token_d, pos_d, embd_gpu, embd_qt, embd_rb, cache,
5391 n_vocab, cap_bucket_max, &mut tok_out)?;
5392 Ok(tok_out)
5393 }
5394
5395 #[allow(clippy::too_many_arguments)]
5398 pub fn gemma4_decode_step_dc_into(&self, e: &Engine, token_d: &CudaSlice<u32>,
5399 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
5400 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
5401 n_vocab: usize, cap_bucket_max: Option<(usize, usize)>,
5402 tok_out: &mut CudaSlice<u32>)
5403 -> Result<(), Box<dyn std::error::Error>> {
5404 let n_embd = self.cfg.n_embd as usize;
5405 let eps = self.cfg.rms_eps;
5406 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
5407 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
5408 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
5409 let n_layers = self.layers.len();
5410 for (il, layer) in self.layers.iter().enumerate() {
5411 let (hq, hdq) = match h_carry.take() {
5412 Some(p) => p,
5413 None => e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, 1, eps)?,
5414 };
5415 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
5416 let o = self.gemma4_decode_attn_dc(e, fa, il, &hq, &hdq, pos_d, cache, cap_bucket_max)?;
5417 let mut cur = e.uninit(n_embd)?;
5418 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, 1, eps)?;
5419 let next_norm = if il + 1 < n_layers {
5420 Some(self.layers[il + 1].attn_norm.float_data())
5421 } else { None };
5422 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, 1, next_norm)?;
5423 x = xn;
5424 h_carry = hn;
5425 }
5426 let mut hn = e.uninit(n_embd)?;
5427 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
5428 let mut logits = e.matmul(&self.output, &hn, 1)?;
5429 self.gemma4_suppress(e, &mut logits, 1)?; e.argmax_token_device_into(&logits, tok_out, n_vocab)?;
5431 e.inc_seqlen(pos_d)?;
5432 if cap_bucket_max.is_none() { cache.pos += 1; }
5433 Ok(())
5434 }
5435
5436 pub fn g4_dc_slots(&self, e: &Engine) -> Result<G4DcSlots, Box<dyn std::error::Error>> {
5443 let n_embd = self.cfg.n_embd as usize;
5444 let n_vocab = self.output.out_features();
5445 let n_layers = self.layers.len();
5446 let (mut qmax, mut kvmax, mut ffmax) = (0usize, 0usize, 0usize);
5447 for il in 0..n_layers {
5448 let (hd, nkv, nh, _b, _s, _w) = self.gemma4_geom(il);
5449 qmax = qmax.max(nh * hd);
5450 kvmax = kvmax.max(nkv * hd);
5451 if let crate::hybrid::Ffn::Dense { ffn_gate, .. } = &self.layers[il].ffn {
5452 ffmax = ffmax.max(ffn_gate.out_features());
5453 }
5454 }
5455 Ok(G4DcSlots {
5456 x: e.uninit(n_embd)?, xn: e.uninit(n_embd)?, cur: e.uninit(n_embd)?,
5457 hq: e.alloc_i8_uninit(n_embd)?, hd_: e.uninit(n_embd / 32)?,
5458 q0: e.uninit(qmax)?, k0: e.uninit(kvmax)?, v0: e.uninit(kvmax)?,
5459 q: e.uninit(qmax)?, k: e.uninit(kvmax)?, v: e.uninit(kvmax)?,
5460 attn: e.uninit(qmax)?, o: e.uninit(n_embd)?,
5461 attn_out: e.uninit(n_embd)?, zsh: e.uninit(n_embd)?,
5462 zq: e.alloc_i8_uninit(n_embd.max(qmax))?, zd: e.uninit(n_embd.max(qmax) / 32)?,
5465 gate: e.uninit(ffmax)?, up: e.uninit(ffmax)?,
5466 act: e.uninit(ffmax)?, actq: e.alloc_i8_uninit(ffmax)?, actd: e.uninit(ffmax / 32)?,
5467 f0: e.uninit(n_embd)?, sn: e.uninit(n_embd)?,
5468 hn: e.uninit(n_embd)?, logits: e.uninit(n_vocab)?,
5469 })
5470 }
5471
5472 fn g4_matvec_m1_into(&self, e: &Engine, w: &crate::model::GpuTensor,
5475 aq: &CudaSlice<i8>, ad: &CudaSlice<f32>, y: &mut CudaSlice<f32>)
5476 -> Result<(), Box<dyn std::error::Error>> {
5477 use crate::model::GpuTensor;
5478 let (bytes, qtype, row_bytes, scale, rp) = match w {
5479 GpuTensor::Quant { bytes, qtype, row_bytes, scale, rp, .. } =>
5480 (bytes, *qtype, *row_bytes, *scale, *rp),
5481 _ => return Err("g4_matvec_m1_into: non-quant tensor".into()),
5482 };
5483 let (mbytes, mrp) = match w {
5484 GpuTensor::Quant { rp4: Some(m4), .. } => (m4, true),
5485 _ => (bytes, rp),
5486 };
5487 e.qmatvec_mmvq_into(mbytes, aq, ad, 1, w.in_features(), w.out_features(),
5488 qtype, row_bytes, scale, mrp, y)
5489 }
5490
5491 #[allow(clippy::too_many_arguments)]
5495 pub fn gemma4_decode_step_dc_slotted(&self, e: &Engine, token_d: &CudaSlice<u32>,
5496 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
5497 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
5498 n_vocab: usize, cap_bucket_max: Option<(usize, usize)>,
5499 sl: &mut G4DcSlots, tok_out: &mut CudaSlice<u32>,
5500 ring: Option<(&mut CudaSlice<u32>, usize)>)
5501 -> Result<(), Box<dyn std::error::Error>> {
5502 let n_embd = self.cfg.n_embd as usize;
5503 let eps = self.cfg.rms_eps;
5504 e.embed_gather_device_into(embd_gpu, token_d, &mut sl.x, n_embd, embd_qt, embd_rb)?;
5505 e.scale_inplace(&mut sl.x, (n_embd as f32).sqrt(), n_embd)?;
5506 let n_layers = self.layers.len();
5507 let mut has_carry = false;
5508 for il in 0..n_layers {
5509 if !has_carry {
5510 e.rms_norm_q8_1_into(&sl.x, self.layers[il].attn_norm.float_data(), n_embd, 1,
5511 eps, &mut sl.hq, &mut sl.hd_)?;
5512 }
5513 has_carry = true;
5514 let layer = &self.layers[il];
5515 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
5516 self.gemma4_decode_attn_dc_slotted(e, fa, il, pos_d, cache, cap_bucket_max, sl)?;
5517 e.rms_norm(&sl.o, layer.post_attn_norm.float_data(), &mut sl.cur, n_embd, 1, eps)?;
5518 let next_norm = if il + 1 < n_layers {
5519 Some(self.layers[il + 1].attn_norm.float_data())
5520 } else { None };
5521 self.gemma4_layer_tail_slotted(e, layer, next_norm, sl)?;
5522 std::mem::swap(&mut sl.x, &mut sl.xn);
5523 }
5524 e.rms_norm(&sl.x, self.output_norm.float_data(), &mut sl.hn, n_embd, 1, eps)?;
5525 e.quantize_q8_1_into(&sl.hn, 1, n_embd, &mut sl.zq, &mut sl.zd)?;
5526 {
5528 let (zq, zd) = (&sl.zq, &sl.zd);
5529 let zq = unsafe { &*(zq as *const CudaSlice<i8>) };
5530 let zd = unsafe { &*(zd as *const CudaSlice<f32>) };
5531 self.g4_matvec_m1_into(e, &self.output, zq, zd, &mut sl.logits)?;
5532 }
5533 self.gemma4_suppress(e, &mut sl.logits, 1)?;
5534 e.argmax_token_device_into(&sl.logits, tok_out, n_vocab)?;
5535 if let Some((ring, base)) = ring {
5536 e.plain_tok_ring(tok_out, pos_d, base, ring)?;
5540 }
5541 e.inc_seqlen(pos_d)?;
5542 if cap_bucket_max.is_none() { cache.pos += 1; }
5543 Ok(())
5544 }
5545
5546 #[allow(clippy::too_many_arguments)]
5548 fn gemma4_decode_attn_dc_slotted(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer,
5549 il: usize, pos_d: &CudaSlice<i32>, cache: &mut Cache,
5550 cap_bucket_max: Option<(usize, usize)>, sl: &mut G4DcSlots)
5551 -> Result<(), Box<dyn std::error::Error>> {
5552 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
5553 let eps = self.cfg.rms_eps;
5554 let aux = self.gemma4_aux.as_ref().unwrap();
5555 {
5556 let hq = unsafe { &*(&sl.hq as *const CudaSlice<i8>) };
5557 let hdq = unsafe { &*(&sl.hd_ as *const CudaSlice<f32>) };
5558 if swa {
5559 if !e.matmul_q4_fused3_into(&fa.wq, &fa.wk, &fa.wv, hq, hdq,
5560 &mut sl.q0, &mut sl.k0, &mut sl.v0)? {
5561 return Err("slotted step: fused3 unavailable (non-uniform trunk)".into());
5562 }
5563 } else {
5564 if !e.matmul_q4_fused2_into(&fa.wq, &fa.wk, hq, hdq, &mut sl.q0, &mut sl.k0)? {
5565 return Err("slotted step: fused2 unavailable".into());
5566 }
5567 let k0r = unsafe { &*(&sl.k0 as *const CudaSlice<f32>) };
5568 e.copy_into(&mut sl.v0, 0, k0r, nkv * hd)?;
5569 }
5570 }
5571 let ff = if swa { None } else {
5574 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
5575 };
5576 let kvl = cache.kv[il].as_mut().unwrap();
5577 let kv_fp8 = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
5578 if crate::Engine::qkv_append_on() {
5579 e.rms_norm_qkv_rope_append_dc(&sl.q0, &sl.k0, &sl.v0, fa.q_norm.float_data(),
5581 fa.k_norm.float_data(), &aux.ones, &mut sl.q, &mut sl.k, &mut sl.v, hd, nh, nkv,
5582 pos_d, nh, nkv, base, 1.0, ff, eps,
5583 &mut kvl.k, &mut kvl.v, &kvl.len_d, kvl.k_tok_bytes, kvl.v_tok_bytes, kv_fp8)?;
5584 } else {
5585 e.rms_norm_qkv_rope(&sl.q0, &sl.k0, &sl.v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
5586 &aux.ones, &mut sl.q, &mut sl.k, &mut sl.v, hd, nh, nkv,
5587 pos_d, nh, nkv, base, 1.0, ff, eps)?;
5588 e.append_kv_quantized_dc(&sl.k, &sl.v, &mut kvl.k, &mut kvl.v, &kvl.len_d,
5589 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes,
5590 kv_fp8)?;
5591 }
5592 e.inc_seqlen(&mut kvl.len_d)?;
5593 let (b_swa, b_glob) = cap_bucket_max.expect("slotted step is capture-only");
5594 let k_view = e.view_u8(&kvl.k, kvl.k.len());
5595 let v_view = e.view_u8(&kvl.v, kvl.v.len());
5596 let rows_on = std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0");
5597 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
5598 let mut fa_q8 = false;
5602 if !swa && hd == 512 && b_glob >= crate::fa512_min_tkv() && rows_on {
5603 e.fa_decode_rows(&sl.q, &k_view, &v_view, &mut sl.attn, hd, nh, nkv, b_glob - 1,
5604 1, scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
5605 Some((&kvl.len_d, -1)), false, false,
5606 Some((&mut sl.zq, &mut sl.zd)))?;
5607 fa_q8 = true;
5608 } else if swa && b_swa > win && hd == 256 && rows_on {
5609 e.fa_decode_rows_w(&sl.q, &k_view, &v_view, &mut sl.attn, hd, nh, nkv,
5610 &kvl.len_d, -1, 1, scale, win,
5611 kvl.k_tok_bytes, kvl.v_tok_bytes,
5612 Some((&mut sl.zq, &mut sl.zd)))?;
5613 fa_q8 = true;
5614 } else {
5615 let b = if swa { b_swa } else { b_glob };
5616 e.fa_decode_dc(&sl.q, &k_view, &v_view, &mut sl.attn, hd, nh, nkv, &kvl.len_d, b,
5617 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
5618 swa && crate::Engine::wkv_on())?;
5619 }
5620 if !fa_q8 {
5621 let aq = unsafe { &*(&sl.attn as *const CudaSlice<f32>) };
5622 e.quantize_q8_1_into(aq, 1, nh * hd, &mut sl.zq, &mut sl.zd)?;
5623 }
5624 {
5625 let zq = unsafe { &*(&sl.zq as *const CudaSlice<i8>) };
5626 let zd = unsafe { &*(&sl.zd as *const CudaSlice<f32>) };
5627 self.g4_matvec_m1_into(e, &fa.wo, zq, zd, &mut sl.o)?;
5628 }
5629 Ok(())
5630 }
5631
5632 fn gemma4_layer_tail_slotted(&self, e: &Engine, layer: &crate::hybrid::HybridLayer,
5635 next_norm: Option<&CudaSlice<f32>>, sl: &mut G4DcSlots)
5636 -> Result<(), Box<dyn std::error::Error>> {
5637 let n_embd = self.cfg.n_embd as usize;
5638 let eps = self.cfg.rms_eps;
5639 let bits = layer.gemma4.as_ref().unwrap();
5640 let crate::hybrid::Ffn::Dense { ffn_gate, ffn_up, ffn_down } = &layer.ffn
5641 else { return Err("slotted tail: dense ffn only".into()) };
5642 e.add_rms_norm(&sl.cur, &sl.x, bits.ffn_norm.float_data(), &mut sl.attn_out,
5643 &mut sl.zsh, n_embd, 1, eps)?;
5644 let n_ff = ffn_gate.out_features();
5645 {
5646 let zshr = unsafe { &*(&sl.zsh as *const CudaSlice<f32>) };
5647 e.quantize_q8_1_into(zshr, 1, n_embd, &mut sl.zq, &mut sl.zd)?;
5648 }
5649 {
5650 let zq = unsafe { &*(&sl.zq as *const CudaSlice<i8>) };
5651 let zd = unsafe { &*(&sl.zd as *const CudaSlice<f32>) };
5652 if !e.matmul_q4_fused2_into(ffn_gate, ffn_up, zq, zd, &mut sl.gate, &mut sl.up)? {
5653 return Err("slotted tail: ffn fused2 unavailable".into());
5654 }
5655 }
5656 debug_assert!(e.uses_q8_1_fast(ffn_down));
5657 {
5658 let upr = unsafe { &*(&sl.up as *const CudaSlice<f32>) };
5659 let upv = e.view(upr, n_ff);
5660 let up_all = upv.slice(0..n_ff);
5661 let gr = unsafe { &*(&sl.gate as *const CudaSlice<f32>) };
5662 e.gelu_tanh_mul_q8_1_into(gr, &up_all, &mut sl.act, n_ff, 1,
5663 &mut sl.actq, &mut sl.actd)?;
5664 }
5665 {
5666 let aq = unsafe { &*(&sl.actq as *const CudaSlice<i8>) };
5667 let ad = unsafe { &*(&sl.actd as *const CudaSlice<f32>) };
5668 self.g4_matvec_m1_into(e, ffn_down, aq, ad, &mut sl.f0)?;
5669 }
5670 e.rms_norm(&sl.f0, bits.post_ffw_norm.float_data(), &mut sl.sn, n_embd, 1, eps)?;
5671 match next_norm {
5672 Some(w) => {
5673 e.add_scale_rms_norm_q8_1_into(&sl.sn, &sl.attn_out, bits.layer_scale, w,
5674 &mut sl.xn, n_embd, 1, eps,
5675 &mut sl.hq, &mut sl.hd_)?;
5676 }
5677 None => {
5678 e.add_scale(&sl.sn, &sl.attn_out, bits.layer_scale, &mut sl.xn, n_embd)?;
5679 }
5680 }
5681 Ok(())
5682 }
5683
5684 #[allow(clippy::too_many_arguments)]
5686 fn gemma4_decode_attn_dc(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
5687 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
5688 pos_d: &CudaSlice<i32>, cache: &mut Cache,
5689 cap_bucket_max: Option<(usize, usize)>)
5690 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
5691 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
5692 let eps = self.cfg.rms_eps;
5693 let aux = self.gemma4_aux.as_ref().unwrap();
5694 let (q0, k0, v0) = if swa {
5695 match e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)? {
5696 Some(t3) => t3,
5697 None => {
5698 let h0 = e.zeros(0)?;
5699 (e.matmul_pre(&fa.wq, hq, hdq, &h0, 1)?,
5700 e.matmul_pre(&fa.wk, hq, hdq, &h0, 1)?,
5701 e.matmul_pre(&fa.wv, hq, hdq, &h0, 1)?)
5702 }
5703 }
5704 } else {
5705 let (q0, k0) = match e.matmul_q4_fused2(&fa.wq, &fa.wk, hq, hdq)? {
5706 Some(p) => p,
5707 None => {
5708 let h0 = e.zeros(0)?;
5709 (e.matmul_pre(&fa.wq, hq, hdq, &h0, 1)?,
5710 e.matmul_pre(&fa.wk, hq, hdq, &h0, 1)?)
5711 }
5712 };
5713 let v0 = e.clone_dtod(&k0)?;
5714 (q0, k0, v0)
5715 };
5716 let mut q = e.uninit(nh * hd)?;
5717 let mut k = e.uninit(nkv * hd)?;
5718 let mut v = e.uninit(nkv * hd)?;
5719 let ff = if swa { None } else {
5721 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
5722 };
5723 let kvl = cache.kv[il].as_mut().unwrap();
5724 let kv_fp8 = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
5725 if crate::Engine::qkv_append_on() {
5726 e.rms_norm_qkv_rope_append_dc(&q0, &k0, &v0, fa.q_norm.float_data(),
5728 fa.k_norm.float_data(), &aux.ones, &mut q, &mut k, &mut v, hd, nh, nkv,
5729 pos_d, nh, nkv, base, 1.0, ff, eps,
5730 &mut kvl.k, &mut kvl.v, &kvl.len_d, kvl.k_tok_bytes, kvl.v_tok_bytes, kv_fp8)?;
5731 } else {
5732 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
5733 &aux.ones, &mut q, &mut k, &mut v, hd, nh, nkv,
5734 pos_d, nh, nkv, base, 1.0, ff, eps)?;
5735 e.append_kv_quantized_dc(&k, &v, &mut kvl.k, &mut kvl.v, &kvl.len_d,
5736 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes, kvl.v_tok_bytes, kv_fp8)?;
5737 }
5738 e.inc_seqlen(&mut kvl.len_d)?;
5739 let mut attn = e.uninit(nh * hd)?;
5740 let mut fa_q8: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
5743 match cap_bucket_max {
5748 None => {
5749 kvl.len += 1;
5753 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
5754 if !swa && hd == 512 && kvl.len >= crate::fa512_min_tkv()
5755 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
5756 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
5759 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
5760 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
5761 e.fa_decode_rows(&q, &kp, &vp, &mut attn, hd, nh, nkv, kvl.len - 1, 1,
5762 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
5763 Some((&kvl.len_d, -1)), false, false,
5764 Some((&mut aq8, &mut ad8)))?;
5765 fa_q8 = Some((aq8, ad8));
5766 } else if swa && kvl.len > win && hd == 256
5767 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
5768 let kp = e.view_u8(&kvl.k, kvl.len * kvl.k_tok_bytes);
5770 let vp = e.view_u8(&kvl.v, kvl.len * kvl.v_tok_bytes);
5771 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
5772 e.fa_decode_rows_w(&q, &kp, &vp, &mut attn, hd, nh, nkv, &kvl.len_d, -1,
5773 1, scale, win, kvl.k_tok_bytes, kvl.v_tok_bytes,
5774 Some((&mut aq8, &mut ad8)))?;
5775 fa_q8 = Some((aq8, ad8));
5776 } else {
5777 let (off_tok, t_kv) = if swa && kvl.len > win { (kvl.len - win, win) }
5778 else { (0, kvl.len) };
5779 let k_view = e.view_u8_range(&kvl.k, off_tok * kvl.k_tok_bytes,
5780 (off_tok + t_kv) * kvl.k_tok_bytes);
5781 let v_view = e.view_u8_range(&kvl.v, off_tok * kvl.v_tok_bytes,
5782 (off_tok + t_kv) * kvl.v_tok_bytes);
5783 e.fa_decode_kvmod(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, t_kv, scale,
5784 kvl.k_tok_bytes, kvl.v_tok_bytes, swa && crate::Engine::wkv_on())?;
5785 }
5786 }
5787 Some((b_swa, b_glob)) => {
5788 let k_view = e.view_u8(&kvl.k, kvl.k.len());
5794 let v_view = e.view_u8(&kvl.v, kvl.v.len());
5795 let rows_on = std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0");
5796 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
5797 if !swa && hd == 512 && b_glob >= crate::fa512_min_tkv() && rows_on {
5798 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
5799 e.fa_decode_rows(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, b_glob - 1,
5800 1, scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
5801 Some((&kvl.len_d, -1)), false, false,
5802 Some((&mut aq8, &mut ad8)))?;
5803 fa_q8 = Some((aq8, ad8));
5804 } else if swa && b_swa > win && hd == 256 && rows_on {
5805 let (mut aq8, mut ad8) = e.uninit_q8_pair(nh * hd)?;
5806 e.fa_decode_rows_w(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
5807 &kvl.len_d, -1, 1, scale, win,
5808 kvl.k_tok_bytes, kvl.v_tok_bytes,
5809 Some((&mut aq8, &mut ad8)))?;
5810 fa_q8 = Some((aq8, ad8));
5811 } else {
5812 let b = if swa { b_swa } else { b_glob };
5813 e.fa_decode_dc(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, &kvl.len_d, b,
5814 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
5815 swa && crate::Engine::wkv_on())?;
5816 }
5817 }
5818 }
5819 if let Some((aq8, ad8)) = fa_q8 {
5822 let mut y = e.uninit(fa.wo.out_features())?;
5823 self.g4_matvec_m1_into(e, &fa.wo, &aq8, &ad8, &mut y)?;
5824 return Ok(y);
5825 }
5826 Ok(e.matmul(&fa.wo, &attn, 1)?)
5827 }
5828
5829 pub fn gemma4_generate_graph(&self, e: &Engine, prompt_pos: usize, first_token: u32,
5834 cache: &mut Cache, max_new: usize, eos: &[u32],
5835 mut on_token: impl FnMut(u32) -> bool)
5836 -> Result<(Vec<u32>, crate::decode::StopReason), Box<dyn std::error::Error>> {
5837 if self.is_gemma4_e4b() {
5838 return Err("E4B graph serving is unwired (HANDOVER-E4B.md) — dc-eager is the serving arm".into());
5839 }
5840 use crate::decode::StopReason;
5841 let n_vocab = self.output.out_features();
5842 let n_embd = self.cfg.n_embd as usize;
5843 let embd_gpu = self.embd_gpu.get_or_init(|| {
5844 e.upload_u8(&self.embd.raw).expect("embed table upload")
5845 });
5846 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
5847 for kvl in cache.kv.iter_mut().flatten() {
5848 e.set_i32_one(&mut kvl.len_d, kvl.len as i32)?;
5849 }
5850 let mut token_d = e.stream().clone_htod(&[first_token])?;
5851 let mut pos_d = e.htod_i32(&[prompt_pos as i32])?;
5852 let g4 = self.cfg.gemma4.as_ref().unwrap();
5853 let (hd_s, hd_g) = (g4.key_length_swa as usize, g4.key_length_global as usize);
5854 let nkv_s = g4.head_count_kv.iter().zip(g4.swa_pattern.iter())
5856 .find(|p| *p.1).map(|p| *p.0 as usize).unwrap_or(8);
5857 let nkv_g = g4.head_count_kv.iter().zip(g4.swa_pattern.iter())
5858 .find(|p| !*p.1).map(|p| *p.0 as usize).unwrap_or(2);
5859 let mut graphs: std::collections::HashMap<((bool, usize), (bool, usize), bool, bool),
5860 (cudarc::driver::CudaGraph,
5861 Vec<Box<dyn std::any::Any + Send>>)> = Default::default();
5862 let mut slots = self.g4_dc_slots(e)?;
5865 const RING: usize = 64;
5868 const DRAIN: usize = 1;
5874 let mut ring = e.stream().alloc_zeros::<u32>(RING)?;
5875 let ring_base = prompt_pos;
5876 let mut out = Vec::with_capacity(max_new);
5877 let mut reason = StopReason::MaxNew;
5878 let mut next = first_token;
5879 let mut captures = 0usize;
5880 for _ in 0..max_new {
5881 out.push(next);
5882 if eos.contains(&next) { reason = StopReason::Eos; break; }
5883 if !on_token(next) { reason = StopReason::Callback; break; }
5884 let t_kv = cache.pos + 1;
5885 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
5893 let f512 = crate::fa512_min_tkv();
5894 let key_s = if t_kv > win { (true, usize::MAX) }
5895 else { e.fa_bucket_key(t_kv, hd_s, nkv_s, crate::Engine::wkv_on()) };
5896 let (key_g, rung_end) = if t_kv >= f512 {
5897 let end = (t_kv + 1).next_power_of_two().max(f512 * 2);
5900 ((true, end), end)
5901 } else { (e.fa_bucket_key(t_kv, hd_g, nkv_g, false), t_kv) };
5902 let key = (key_s, key_g, t_kv >= f512, t_kv > win);
5903 if !graphs.contains_key(&key) {
5904 let bucket_max = (t_kv, rung_end);
5905 let snap = cache.snapshot(e)?;
5907 let pos_save = e.dtoh_i32_one(&pos_d)?;
5908 let len_save: Vec<Option<i32>> = cache.kv.iter()
5909 .map(|k| k.as_ref().map(|kvl| e.dtoh_i32_one(&kvl.len_d).unwrap())).collect();
5910 let tok_save = e.dtoh_u32_one(&token_d)?;
5911 let graph = {
5916 let tok_ref = &mut token_d;
5917 let pos_ref = &mut pos_d;
5918 let cache_ref = &mut *cache;
5919 let slots_ref = &mut slots;
5920 let ring_ref = &mut ring;
5921 e.capture_graph_retained_flags(
5922 cudarc::driver::sys::CUgraphInstantiate_flags::CUDA_GRAPH_INSTANTIATE_FLAG_USE_NODE_PRIORITY,
5923 |e| {
5924 let tok_in = unsafe { &*(tok_ref as *const CudaSlice<u32>) };
5926 let sl = unsafe { &mut *(slots_ref as *mut G4DcSlots) };
5927 let rg = unsafe { &mut *(ring_ref as *mut CudaSlice<u32>) };
5928 self.gemma4_decode_step_dc_slotted(e, tok_in, pos_ref, embd_gpu, qt, rb,
5929 cache_ref, n_vocab, Some(bucket_max),
5930 sl, tok_ref, Some((rg, ring_base)))
5931 })?
5932 };
5933 cache.rollback(e, &snap, 0)?;
5934 e.set_i32_one(&mut pos_d, pos_save)?;
5935 for (il, ls) in len_save.iter().enumerate() {
5936 if let (Some(kvl), Some(v)) = (cache.kv[il].as_mut(), ls) {
5937 e.set_i32_one(&mut kvl.len_d, *v)?;
5938 }
5939 }
5940 e.set_u32_one(&mut token_d, tok_save)?;
5941 if std::env::var("MEMRA_GRAPH_CENSUS").as_deref() == Ok("1") {
5942 if let Ok(c) = crate::graph_update::node_census(&graph.0) {
5943 eprintln!("[graph-census] {c:?}");
5944 }
5945 }
5946 graphs.insert(key, graph);
5947 captures += 1;
5948 }
5949 let mut chunk = 1usize;
5954 let drain_cap: usize = std::env::var("MEMRA_GRAPH_DRAIN").ok()
5955 .and_then(|v| v.parse().ok()).unwrap_or(DRAIN);
5956 while chunk < drain_cap && out.len() + chunk < max_new {
5957 let t_next = cache.pos + 1 + chunk;
5958 let key_s2 = if t_next > win { (true, usize::MAX) }
5959 else { e.fa_bucket_key(t_next, hd_s, nkv_s, crate::Engine::wkv_on()) };
5960 let key_g2 = if t_next >= f512 {
5961 (true, (t_next + 1).next_power_of_two().max(f512 * 2))
5962 } else { e.fa_bucket_key(t_next, hd_g, nkv_g, false) };
5963 if (key_s2, key_g2, t_next >= f512, t_next > win) != key { break; }
5964 chunk += 1;
5965 }
5966 let g = &graphs.get(&key).unwrap().0;
5967 for _ in 0..chunk { g.launch()?; }
5968 e.stream().synchronize()?;
5969 let ringh = e.dtoh_u32(&ring)?;
5970 for j in 0..chunk {
5971 let pos_j = cache.pos + j;
5972 let tok_j = ringh[(pos_j - ring_base) % RING];
5973 cache.pos += 0; if j + 1 == chunk { next = tok_j; }
5975 else {
5976 out.push(tok_j);
5977 if eos.contains(&tok_j) || !on_token(tok_j) {
5978 reason = if eos.contains(&tok_j) { StopReason::Eos }
5979 else { StopReason::Callback };
5980 let keep = cache.pos + j + 1;
5982 e.set_i32_one(&mut pos_d, keep as i32)?;
5983 for kvl in cache.kv.iter_mut().filter_map(|k| k.as_mut()) {
5984 e.set_i32_one(&mut kvl.len_d, keep as i32)?;
5985 kvl.len = keep;
5986 }
5987 cache.pos = keep;
5988 if std::env::var("MEMRA_GRAPH_STATS").is_ok() {
5989 eprintln!("[gemma-graph] captures={captures} buckets={}", graphs.len());
5990 }
5991 return Ok((out, reason));
5992 }
5993 }
5994 }
5995 cache.pos += chunk;
5996 for kvl in cache.kv.iter_mut().filter_map(|k| k.as_mut()) { kvl.len += chunk; }
5997 }
5998 if std::env::var("MEMRA_GRAPH_STATS").is_ok() {
5999 eprintln!("[gemma-graph] captures={captures} buckets={}", graphs.len());
6000 }
6001 Ok((out, reason))
6002 }
6003
6004 pub(crate) fn gemma4_decode_step_t(&self, e: &Engine, tokens: &[u32], pos0: usize,
6010 cache: &mut Cache)
6011 -> Result<Vec<f32>, Box<dyn std::error::Error>> {
6012 Ok(self.gemma4_decode_step_t_h(e, tokens, pos0, cache)?.0)
6013 }
6014
6015 pub(crate) fn gemma4_decode_step_t_am(&self, e: &Engine, tokens: &[u32], pos0: usize,
6019 cache: &mut Cache)
6020 -> Result<(Vec<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6021 let (ld, hn) = self.gemma4_verify_trunk(e, tokens, pos0, cache, None)?;
6022 let t = tokens.len();
6023 let n_vocab = self.output.out_features();
6024 let mut toks = e.stream().alloc_zeros::<u32>(t)?;
6025 for i in 0..t {
6026 e.argmax_token_device_col(&ld, i, n_vocab, &mut toks, i)?;
6027 }
6028 Ok((e.dtoh_u32(&toks)?, hn))
6029 }
6030
6031 pub(crate) fn gemma4_decode_step_t_am_dev(&self, e: &Engine, tok_d: &CudaSlice<u32>, t: usize,
6034 pos0: usize, cache: &mut Cache)
6035 -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6036 let (ld, hn) = self.gemma4_verify_trunk(e, &vec![0u32; t], pos0, cache, Some(tok_d))?;
6037 let n_vocab = self.output.out_features();
6038 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
6039 for i in 0..t {
6040 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
6041 }
6042 Ok((vam, hn))
6043 }
6044
6045 pub(crate) fn gemma4_decode_step_t_h(&self, e: &Engine, tokens: &[u32], pos0: usize,
6048 cache: &mut Cache)
6049 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6050 let (mut ld, hn) = self.gemma4_verify_trunk(e, tokens, pos0, cache, None)?;
6051 let t = tokens.len();
6052 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
6053 e.softcap(&mut ld, cap, t * self.output.out_features())?;
6054 Ok((e.dtoh(&ld)?, hn))
6055 }
6056
6057 pub(crate) fn verify_stream_scratch(&self, e: &Engine, cap: usize)
6060 -> Result<VerifyStreamScratch, Box<dyn std::error::Error>> {
6061 Ok(VerifyStreamScratch {
6062 pos_d: e.htod_i32(&vec![0i32; cap])?,
6063 row_ctrs: (0..cap).map(|_| e.htod_i32(&[0])).collect::<Result<_, _>>()?,
6064 })
6065 }
6066
6067 pub(crate) fn gemma4_verify_t_am_stream(&self, e: &Engine, tok_d: &CudaSlice<u32>, t: usize,
6075 ctr: &CudaSlice<i32>, hint: usize,
6076 cache: &mut Cache,
6077 scr: &mut VerifyStreamScratch)
6078 -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6079 let n_embd = self.cfg.n_embd as usize;
6080 let eps = self.cfg.rms_eps;
6081 assert!(t <= scr.row_ctrs.len() && t <= 64);
6082 e.i32_iota_from(ctr, &mut scr.pos_d, t)?;
6083 for i in 0..t {
6084 e.i32_copy_add(ctr, &mut scr.row_ctrs[i], (i + 1) as i32)?;
6085 }
6086 let (pos_d, row_ctrs) = (&scr.pos_d, &scr.row_ctrs);
6087 let embd_gpu = self.embd_gpu.get_or_init(|| {
6088 e.upload_u8(&self.embd.raw).expect("embed table upload")
6089 });
6090 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
6091 let mut x = e.embed_gather_device_td(embd_gpu, tok_d, t, n_embd, qt, rb)?;
6092 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
6093 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
6094 let n_layers = self.layers.len();
6095 for (il, layer) in self.layers.iter().enumerate() {
6096 let (hq, hdq) = match h_carry.take() {
6097 Some(p) => p,
6098 None => e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, t, eps)?,
6099 };
6100 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
6101 let o = self.gemma4_verify_attn_stream(e, fa, il, &hq, &hdq, pos_d, t, cache,
6102 hint, row_ctrs)?;
6103 let mut cur = e.uninit(t * n_embd)?;
6104 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, t, eps)?;
6105 let next_norm = if il + 1 < n_layers {
6106 Some(self.layers[il + 1].attn_norm.float_data())
6107 } else { None };
6108 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, t, next_norm)?;
6109 x = xn;
6110 h_carry = hn;
6111 self.dflash_tap(e, cache, il, &x, t)?;
6112 }
6113 let mut hn = e.uninit(t * n_embd)?;
6114 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
6115 let ld = e.matmul(&self.output, &hn, t)?;
6116 let n_vocab = self.output.out_features();
6117 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
6118 for i in 0..t {
6119 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
6120 }
6121 Ok((vam, hn))
6122 }
6123
6124 fn dflash_tap(&self, e: &Engine, cache: &mut Cache, il: usize, x: &CudaSlice<f32>, t: usize)
6131 -> Result<(), Box<dyn std::error::Error>> {
6132 let Some(taps) = cache.dflash_taps.as_mut() else { return Ok(()) };
6133 let Some(slot) = taps.layer_ids.iter().position(|&l| l == il) else { return Ok(()) };
6134 let h = taps.hidden;
6135 let n_taps = taps.layer_ids.len();
6136 debug_assert_eq!(taps.t, t);
6137 let xv = e.view(x, t * h);
6138 for r in 0..t {
6139 let row = xv.slice(r * h..(r + 1) * h);
6140 e.copy_view_into(&mut taps.buf, r * n_taps * h + slot * h, &row, h)?;
6141 }
6142 Ok(())
6143 }
6144
6145 fn gemma4_verify_trunk(&self, e: &Engine, tokens: &[u32], pos0: usize, cache: &mut Cache,
6146 tok_dev: Option<&CudaSlice<u32>>)
6147 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6148 let n_embd = self.cfg.n_embd as usize;
6149 let eps = self.cfg.rms_eps;
6150 let t = tokens.len();
6151 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
6152 let pos_d = e.htod_i32(&pos)?;
6153 let mut x = match tok_dev {
6154 Some(td) => {
6155 let embd_gpu = self.embd_gpu.get_or_init(|| {
6156 e.upload_u8(&self.embd.raw).expect("embed table upload")
6157 });
6158 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
6159 e.embed_gather_device_td(embd_gpu, td, t, n_embd, qt, rb)?
6160 }
6161 None => e.htod(&self.embd.gather(n_embd, tokens))?,
6162 };
6163 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
6164 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
6165 let n_layers = self.layers.len();
6166 for (il, layer) in self.layers.iter().enumerate() {
6167 let (hq, hdq) = match h_carry.take() {
6168 Some(p) => p,
6169 None => e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, t, eps)?,
6170 };
6171 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
6172 let o = self.gemma4_verify_attn(e, fa, il, &hq, &hdq, &pos_d, t, cache)?;
6173 let mut cur = e.uninit(t * n_embd)?;
6174 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, t, eps)?;
6175 let next_norm = if il + 1 < n_layers {
6176 Some(self.layers[il + 1].attn_norm.float_data())
6177 } else { None };
6178 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, t, next_norm)?;
6179 x = xn;
6180 h_carry = hn;
6181 self.dflash_tap(e, cache, il, &x, t)?;
6182 }
6183 let mut hn = e.uninit(t * n_embd)?;
6184 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
6185 let mut ld = e.matmul(&self.output, &hn, t)?;
6186 self.gemma4_suppress(e, &mut ld, t)?; cache.pos += t;
6188 Ok((ld, hn))
6189 }
6190
6191 #[allow(clippy::too_many_arguments)]
6199 fn gemma4_verify_attn_stream(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
6200 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
6201 pos_d: &CudaSlice<i32>, t: usize,
6202 cache: &mut Cache, hint: usize,
6203 row_ctrs: &[CudaSlice<i32>])
6204 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6205 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
6206 let eps = self.cfg.rms_eps;
6207 let aux = self.gemma4_aux.as_ref().unwrap();
6208 let h0 = e.zeros(0)?;
6209 let h = &h0;
6210 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6213 let f2b = *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
6214 let fused_qkv = if f2b {
6215 if swa {
6216 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
6217 .map(|(a, b, c)| (a, b, Some(c)))
6218 } else {
6219 e.matmul_q4_fused2_batched(&fa.wq, &fa.wk, hq, hdq, t)?
6220 .map(|(a, b)| (a, b, None))
6221 }
6222 } else { None };
6223 let (q0, k0, v0) = match fused_qkv {
6224 Some((a, b, cv)) => {
6225 let v = match cv { Some(c) => c, None => e.clone_dtod(&b)? };
6226 (a, b, v)
6227 }
6228 None => {
6229 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
6230 let k0 = e.matmul_pre(&fa.wk, hq, hdq, h, t)?;
6231 let v0 = if swa { e.matmul_pre(&fa.wv, hq, hdq, h, t)? }
6232 else { e.clone_dtod(&k0)? };
6233 (q0, k0, v0)
6234 }
6235 };
6236 let mut q = e.uninit(t * nh * hd)?;
6237 let mut k = e.uninit(t * nkv * hd)?;
6238 let mut v = e.uninit(t * nkv * hd)?;
6239 let ff = if swa { None } else {
6242 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
6243 };
6244 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
6245 &aux.ones, &mut q, &mut k, &mut v, hd, nh * t, nkv * t,
6246 pos_d, nh, nkv, base, 1.0, ff, eps)?;
6247 let kvl = cache.kv[il].as_mut().unwrap();
6248 e.append_kv_quantized_rows_dc(&k, &v, &mut kvl.k, &mut kvl.v, &kvl.len_d, t,
6250 kvl.kv_dim_k, kvl.kv_dim_v,
6251 kvl.k_tok_bytes, kvl.v_tok_bytes,
6252 (!swa && crate::Engine::gkv_on())
6253 || (swa && crate::Engine::wkv_on()))?;
6254 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
6257 let mut attn = e.uninit(t * nh * hd)?;
6258 let k_view = e.view_u8(&kvl.k, kvl.k.len());
6259 let v_view = e.view_u8(&kvl.v, kvl.v.len());
6260 if swa && hint + 1 >= win {
6263 e.fa_decode_rows_w(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
6266 &kvl.len_d, 0, t, scale, win,
6267 kvl.k_tok_bytes, kvl.v_tok_bytes, None)?;
6268 } else if hd == 512 && hint + t < crate::fa512_min_tkv() {
6269 let bucket = (hint + t + 2).next_power_of_two()
6282 .min(crate::fa512_min_tkv().saturating_sub(1));
6283 let qv = e.view(&q, t * nh * hd);
6284 for i in 0..t {
6285 let q_row = qv.slice(i * nh * hd..(i + 1) * nh * hd);
6286 let mut q_one = e.uninit(nh * hd)?;
6287 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
6288 let mut a_one = e.uninit(nh * hd)?;
6289 e.fa_decode_dc(&q_one, &k_view, &v_view, &mut a_one, hd, nh, nkv,
6290 &row_ctrs[i], bucket, scale,
6291 kvl.k_tok_bytes, kvl.v_tok_bytes, false)?;
6292 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
6293 }
6294 } else if hd == 512 {
6295 e.fa_decode_rows(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, hint, t, scale,
6298 kvl.k_tok_bytes, kvl.v_tok_bytes,
6299 Some((&kvl.len_d, 0)), false, false, None)?;
6300 } else {
6301 e.fa_decode_rows_dc(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
6303 &kvl.len_d, hint + t, t, scale,
6304 kvl.k_tok_bytes, kvl.v_tok_bytes, 0,
6305 swa && crate::Engine::wkv_on())?;
6306 }
6307 Ok(e.matmul(&fa.wo, &attn, t)?)
6308 }
6309
6310 fn gemma4_verify_attn(&self, e: &Engine, fa: &crate::hybrid::FullAttnLayer, il: usize,
6311 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
6312 pos_d: &CudaSlice<i32>, t: usize,
6313 cache: &mut Cache)
6314 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6315 let (hd, nkv, nh, base, scale, swa) = self.gemma4_geom(il);
6316 let eps = self.cfg.rms_eps;
6317 let aux = self.gemma4_aux.as_ref().unwrap();
6318 let n_embd = self.cfg.n_embd as usize;
6319 let _ = n_embd;
6320
6321 let h0 = e.zeros(0)?;
6322 let h = &h0;
6323 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6326 let f2b = *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0"));
6327 let fused_qkv = if f2b {
6328 if swa {
6329 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
6330 .map(|(a, b, c)| (a, b, Some(c)))
6331 } else {
6332 e.matmul_q4_fused2_batched(&fa.wq, &fa.wk, hq, hdq, t)?
6333 .map(|(a, b)| (a, b, None))
6334 }
6335 } else { None };
6336 let (q0, k0, v0) = match fused_qkv {
6337 Some((a, b, cv)) => {
6338 let v = match cv { Some(c) => c, None => e.clone_dtod(&b)? };
6339 (a, b, v)
6340 }
6341 None => {
6342 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
6343 let k0 = e.matmul_pre(&fa.wk, hq, hdq, h, t)?;
6344 let v0 = if swa { e.matmul_pre(&fa.wv, hq, hdq, h, t)? }
6345 else { e.clone_dtod(&k0)? };
6346 (q0, k0, v0)
6347 }
6348 };
6349 let mut q = e.uninit(t * nh * hd)?;
6350 let mut k = e.uninit(t * nkv * hd)?;
6351 let mut v = e.uninit(t * nkv * hd)?;
6352 let ff = if swa { None } else {
6355 Some(aux.rope_freqs.as_ref().expect("gemma4 global rope needs rope_freqs.weight"))
6356 };
6357 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(), fa.k_norm.float_data(),
6358 &aux.ones, &mut q, &mut k, &mut v, hd, nh * t, nkv * t,
6359 pos_d, nh, nkv, base, 1.0, ff, eps)?;
6360 let kvl = cache.kv[il].as_mut().unwrap();
6361 let base_len = kvl.len;
6362 e.append_kv_quantized_rows(&k, &v, &mut kvl.k, &mut kvl.v, base_len, t,
6363 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()))?;
6364 kvl.len += t;
6365 let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
6366 let mut attn = e.uninit(t * nh * hd)?;
6367 let rows_ok = (hd == 256 && base_len + 1 >= crate::fa_vec_min_tkv())
6370 || (hd == 512 && !swa && base_len + 1 >= crate::fa512_min_tkv());
6373 if rows_ok && (!swa || base_len + t <= win) {
6374 let k_view = e.view_u8(&kvl.k, (base_len + t) * kvl.k_tok_bytes);
6375 let v_view = e.view_u8(&kvl.v, (base_len + t) * kvl.v_tok_bytes);
6376 if hd == 512 {
6377 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
6379 e.fa_decode_rows(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, base_len, t,
6380 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
6381 Some((&kvl.len_d, 0)), false,
6382 swa && crate::Engine::wkv_on(), None)?;
6383 } else {
6384 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
6388 e.fa_decode_rows_dc(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
6389 &kvl.len_d, base_len + t, t, scale,
6390 kvl.k_tok_bytes, kvl.v_tok_bytes, 0,
6391 swa && crate::Engine::wkv_on())?;
6392 }
6393 return Ok(e.matmul(&fa.wo, &attn, t)?);
6394 }
6395 if hd == 256 && swa && base_len + 1 >= win
6403 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
6404 let k_view = e.view_u8(&kvl.k, (base_len + t) * kvl.k_tok_bytes);
6405 let v_view = e.view_u8(&kvl.v, (base_len + t) * kvl.v_tok_bytes);
6406 e.i32_set_k(&mut kvl.len_d, base_len as i32)?;
6407 e.fa_decode_rows_w(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, &kvl.len_d, 0,
6408 t, scale, win, kvl.k_tok_bytes, kvl.v_tok_bytes, None)?;
6409 return Ok(e.matmul(&fa.wo, &attn, t)?);
6410 }
6411 for i in 0..t {
6412 let avail = base_len + i + 1;
6413 let (off_tok, t_kv) = if swa && avail > win { (avail - win, win) } else { (0, avail) };
6414 let k_view = e.view_u8_range(&kvl.k, off_tok * kvl.k_tok_bytes,
6415 (off_tok + t_kv) * kvl.k_tok_bytes);
6416 let v_view = e.view_u8_range(&kvl.v, off_tok * kvl.v_tok_bytes,
6417 (off_tok + t_kv) * kvl.v_tok_bytes);
6418 let qi = e.view(&q, t * nh * hd);
6419 let q_row = qi.slice(i * nh * hd..(i + 1) * nh * hd);
6420 let mut q_one = e.uninit(nh * hd)?;
6421 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
6422 let mut a_one = e.uninit(nh * hd)?;
6423 if swa && avail > win && hd == 256
6427 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
6428 let kp = e.view_u8(&kvl.k, avail * kvl.k_tok_bytes);
6429 let vp = e.view_u8(&kvl.v, avail * kvl.v_tok_bytes);
6430 e.i32_set_k(&mut kvl.len_d, (avail - 1) as i32)?;
6431 e.fa_decode_rows_w(&q_one, &kp, &vp, &mut a_one, hd, nh, nkv, &kvl.len_d, 0,
6432 1, scale, win, kvl.k_tok_bytes, kvl.v_tok_bytes, None)?;
6433 } else if !swa && hd == 512 && avail >= crate::fa512_min_tkv()
6434 && std::env::var("MEMRA_GEMMA_ROWS_W").as_deref() != Ok("0") {
6435 let kp = e.view_u8(&kvl.k, avail * kvl.k_tok_bytes);
6436 let vp = e.view_u8(&kvl.v, avail * kvl.v_tok_bytes);
6437 e.i32_set_k(&mut kvl.len_d, (avail - 1) as i32)?;
6438 e.fa_decode_rows(&q_one, &kp, &vp, &mut a_one, hd, nh, nkv, avail - 1, 1,
6439 scale, kvl.k_tok_bytes, kvl.v_tok_bytes,
6440 Some((&kvl.len_d, 0)), false, false, None)?;
6441 } else {
6442 e.fa_decode_kvmod(&q_one, &k_view, &v_view, &mut a_one, hd, nh, nkv, t_kv, scale,
6443 kvl.k_tok_bytes, kvl.v_tok_bytes, swa && crate::Engine::wkv_on())?;
6444 }
6445 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
6446 }
6447 Ok(e.matmul(&fa.wo, &attn, t)?)
6448 }
6449
6450 pub(crate) fn gemma4_decode_step_h(&self, e: &Engine, token: u32, cache: &mut Cache)
6453 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6454 if let Some(split) = crate::pp::pp2_split(self.layers.len()) {
6459 return self.gemma4_decode_step_h_pp2(e, token, cache, split);
6460 }
6461 if crate::pp::pp_cuts(self.layers.len()).is_some() {
6462 crate::pp::warn_unwired_once("gemma4 eager decode (N>2)");
6463 }
6464 let n_embd = self.cfg.n_embd as usize;
6465 let eps = self.cfg.rms_eps;
6466 let pos_d = e.htod_i32(&[cache.pos as i32])?;
6467 let mut x = e.htod(&self.embd.gather(n_embd, &[token]))?;
6468 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
6469 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
6472 let n_layers = self.layers.len();
6473 for (il, layer) in self.layers.iter().enumerate() {
6474 let (hq, hdq) = match h_carry.take() {
6475 Some(p) => p,
6476 None => e.rms_norm_q8_1(&x, self.layers[0].attn_norm.float_data(), n_embd, 1, eps)?,
6477 };
6478 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
6479 let o = self.gemma4_decode_attn(e, fa, il, &hq, &hdq, &pos_d, cache)?;
6480 let mut cur = e.uninit(n_embd)?;
6481 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, 1, eps)?;
6482 let next_norm = if il + 1 < n_layers {
6483 Some(self.layers[il + 1].attn_norm.float_data())
6484 } else { None };
6485 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, 1, next_norm)?;
6486 x = xn;
6487 h_carry = hn;
6488 }
6489 let mut hn = e.uninit(n_embd)?;
6490 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
6491 let h_seed = e.clone_dtod(&x)?;
6492 let mut ld = e.matmul(&self.output, &hn, 1)?;
6493 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
6494 e.softcap(&mut ld, cap, self.output.out_features())?; self.gemma4_suppress(e, &mut ld, 1)?;
6496 let logits = e.dtoh(&ld)?;
6497 cache.pos += 1;
6498 Ok((logits, h_seed))
6499 }
6500
6501 fn gemma4_decode_layers(&self, e: &Engine, mut x: CudaSlice<f32>, lo: usize, hi: usize,
6509 pos_d: &CudaSlice<i32>, cache: &mut Cache)
6510 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6511 let n_embd = self.cfg.n_embd as usize;
6512 let eps = self.cfg.rms_eps;
6513 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
6514 for il in lo..hi {
6515 let layer = &self.layers[il];
6516 let (hq, hdq) = match h_carry.take() {
6517 Some(p) => p,
6518 None => e.rms_norm_q8_1(&x, self.layers[il].attn_norm.float_data(), n_embd, 1, eps)?,
6520 };
6521 let Mixer::Full(fa) = &layer.mixer else { panic!("gemma4 layer {il} not full-attn") };
6522 let o = self.gemma4_decode_attn(e, fa, il, &hq, &hdq, pos_d, cache)?;
6523 let mut cur = e.uninit(n_embd)?;
6524 e.rms_norm(&o, layer.post_attn_norm.float_data(), &mut cur, n_embd, 1, eps)?;
6525 let next_norm = if il + 1 < hi {
6526 Some(self.layers[il + 1].attn_norm.float_data())
6527 } else { None };
6528 let (xn, hn) = self.gemma4_layer_tail_add_nq(e, layer, &cur, &x, 1, next_norm)?;
6529 x = xn;
6530 h_carry = hn;
6531 }
6532 Ok(x)
6533 }
6534
6535 fn gemma4_decode_step_h_pp2(&self, e: &Engine, token: u32, cache: &mut Cache, split: usize)
6542 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6543 if crate::pp::pp2_streams_off() {
6544 return self.gemma4_decode_step_h_pp2_samestream(e, token, cache, split);
6545 }
6546 let rt = crate::pp::Pp2Rt::get(e)?;
6547 let e0 = rt.engine(0, e);
6548 let e1 = rt.engine(1, e);
6549 let n_embd = self.cfg.n_embd as usize;
6550 let eps = self.cfg.rms_eps;
6551
6552 let (pos_d, slot) = {
6554 let _st0 = rt.enter(0);
6555 let pos_d = e0.htod_i32(&[cache.pos as i32])?;
6556 let mut x = e0.htod(&self.embd.gather(n_embd, &[token]))?;
6557 e0.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
6558 let x = self.gemma4_decode_layers(e0, x, 0, split, &pos_d, cache)?;
6559 let slot = rt.tx(0, &x, n_embd)?;
6560 (pos_d, slot)
6561 };
6562
6563 let _st1 = rt.enter(1);
6565 let x = rt.rx(0, slot, n_embd)?;
6566 let x = self.gemma4_decode_layers(e1, x, split, self.layers.len(), &pos_d, cache)?;
6567
6568 let mut hn = e1.uninit(n_embd)?;
6569 e1.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
6570 let h_seed = e1.clone_dtod(&x)?;
6571 let mut ld = e1.matmul(&self.output, &hn, 1)?;
6572 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
6573 e1.softcap(&mut ld, cap, self.output.out_features())?;
6574 self.gemma4_suppress(e1, &mut ld, 1)?;
6575 let logits = e1.dtoh(&ld)?;
6576 cache.pos += 1;
6577 Ok((logits, h_seed))
6578 }
6579
6580 fn gemma4_decode_step_h_pp2_samestream(&self, e: &Engine, token: u32, cache: &mut Cache,
6583 split: usize)
6584 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6585 let n_embd = self.cfg.n_embd as usize;
6586 let eps = self.cfg.rms_eps;
6587 let pos_d = e.htod_i32(&[cache.pos as i32])?;
6588
6589 let mut x = e.htod(&self.embd.gather(n_embd, &[token]))?;
6591 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
6592 let x = self.gemma4_decode_layers(e, x, 0, split, &pos_d, cache)?;
6593
6594 let boundary_tx = e.clone_dtod(&x)?;
6596 let boundary_rx = e.clone_dtod(&boundary_tx)?;
6597
6598 let x = self.gemma4_decode_layers(e, boundary_rx, split, self.layers.len(), &pos_d, cache)?;
6600
6601 let mut hn = e.uninit(n_embd)?;
6602 e.rms_norm(&x, self.output_norm.float_data(), &mut hn, n_embd, 1, eps)?;
6603 let h_seed = e.clone_dtod(&x)?;
6604 let mut ld = e.matmul(&self.output, &hn, 1)?;
6605 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
6606 e.softcap(&mut ld, cap, self.output.out_features())?;
6607 self.gemma4_suppress(e, &mut ld, 1)?;
6608 let logits = e.dtoh(&ld)?;
6609 cache.pos += 1;
6610 Ok((logits, h_seed))
6611 }
6612}
6613
6614impl HybridModel {
6623 pub fn is_gemma4_e4b(&self) -> bool {
6624 self.gemma4_aux.as_ref().is_some_and(|a| a.e4b.is_some())
6625 }
6626
6627 fn gemma4_e4b_geom(&self, il: usize) -> (usize, usize, usize, f32, f32, bool) {
6631 let g = self.cfg.gemma4.as_ref().unwrap();
6632 let swa = g.swa_pattern[il];
6633 let hd = if swa { g.key_length_swa } else { g.key_length_global } as usize;
6634 let Mixer::Full(fa) = &self.layers[il].mixer else { panic!("e4b layer {il} not full-attn") };
6635 let nh = fa.wq.out_features() / hd;
6636 let nkv = fa.wk.out_features() / hd;
6637 (hd, nkv, nh, if swa { g.rope_base_swa } else { g.rope_base_global }, 1.0, swa)
6638 }
6639
6640 fn gemma4_e4b_kv_target(&self, il: usize) -> Option<usize> {
6642 self.layers[il].gemma4.as_ref()
6643 .and_then(|b| b.e4b.as_ref())
6644 .and_then(|e4| e4.kv_share.map(|t| t as usize))
6645 }
6646
6647 fn gemma4_e4b_inp_pl(&self, e: &Engine, tokens: &[u32], x_scaled: &CudaSlice<f32>, t: usize)
6652 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6653 let tok_d = e.stream().clone_htod(&tokens.to_vec())?;
6654 self.gemma4_e4b_inp_pl_dev(e, &tok_d, x_scaled, t)
6655 }
6656
6657 fn gemma4_e4b_inp_pl_dev(&self, e: &Engine, tok_d: &CudaSlice<u32>,
6659 x_scaled: &CudaSlice<f32>, t: usize)
6660 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6661 let aux = self.gemma4_aux.as_ref().unwrap();
6662 let m = aux.e4b.as_ref().unwrap();
6663 let n_embd = self.cfg.n_embd as usize;
6664 let n_layer = self.layers.len();
6665 let width = m.n_epl * n_layer;
6666 let tbl = m.tok_tbl_gpu.get_or_init(|| {
6667 e.upload_u8(&m.tok_embd_bytes).expect("e4b per-layer token table upload")
6668 });
6669 let mut a = e.embed_gather_device_td(tbl, tok_d, t, width, m.tok_embd_qt,
6670 m.tok_embd_row_bytes)?;
6671 e.scale_inplace(&mut a, (m.n_epl as f32).sqrt(), t * width)?;
6672 let mut p = e.matmul(&m.model_proj, x_scaled, t)?;
6673 e.scale_inplace(&mut p, 1.0 / (n_embd as f32).sqrt(), t * width)?;
6674 let mut pn = e.uninit(t * width)?;
6675 e.rms_norm(&p, m.proj_norm.float_data(), &mut pn, m.n_epl, t * n_layer,
6676 self.cfg.rms_eps)?;
6677 let mut out = e.uninit(t * width)?;
6678 e.add_scale(&a, &pn, 1.0 / 2f32.sqrt(), &mut out, t * width)?;
6679 Ok(out)
6680 }
6681
6682 #[allow(clippy::too_many_arguments)]
6687 fn gemma4_e4b_attn(&self, e: &Engine, il: usize,
6688 hq: &CudaSlice<i8>, hdq: &CudaSlice<f32>,
6689 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache,
6690 dc_bucket: Option<usize>)
6691 -> Result<CudaSlice<f32>, Box<dyn std::error::Error>> {
6692 let (hd, nkv, nh, base, scale, swa) = self.gemma4_e4b_geom(il);
6693 let eps = self.cfg.rms_eps;
6694 let aux = self.gemma4_aux.as_ref().unwrap();
6695 let Mixer::Full(fa) = &self.layers[il].mixer else { unreachable!() };
6696 let h0 = e.zeros(0)?;
6700 let h = &h0;
6701
6702 let ff = if swa { None } else {
6703 Some(aux.rope_freqs.as_ref().expect("e4b global rope needs rope_freqs.weight"))
6704 };
6705 let share = self.gemma4_e4b_kv_target(il);
6706 let mut kv_f32: Option<(CudaSlice<f32>, CudaSlice<f32>)> = None;
6708 let mut q;
6709 if let Some(_tgt) = share {
6710 let q0 = e.matmul_pre(&fa.wq, hq, hdq, h, t)?;
6711 q = e.uninit(t * nh * hd)?;
6712 let mut kdummy = e.uninit(1)?;
6715 let mut vdummy = e.uninit(1)?;
6716 e.rms_norm_qkv_rope(&q0, &q0, &q0, fa.q_norm.float_data(),
6717 fa.q_norm.float_data(), &aux.ones,
6718 &mut q, &mut kdummy, &mut vdummy, hd, nh * t, 0,
6719 pos_d, nh, 1, base, 1.0, ff, eps)?;
6720 } else {
6721 let e4bits = self.layers[il].gemma4.as_ref().and_then(|g| g.e4b.as_ref());
6725 let cat = e4bits.and_then(|e4| e4.qkv_cat.as_ref());
6726 q = e.uninit(t * nh * hd)?;
6727 let mut k = e.uninit(t * nkv * hd)?;
6728 let mut v = e.uninit(t * nkv * hd)?;
6729 if t == 1 && cat.is_some() {
6730 let qkv0 = e.matmul_pre(cat.unwrap(), hq, hdq, h, 1)?;
6731 e.rms_norm_qkv_rope_cat(&qkv0, fa.q_norm.float_data(), fa.k_norm.float_data(),
6732 &aux.ones, &mut q, &mut k, &mut v, hd, nh, nkv,
6733 pos_d, nh, nkv, base, 1.0, ff, eps)?;
6734 } else {
6735 let (q0, k0, v0) = match if t == 1 {
6736 e.matmul_q4_fused3(&fa.wq, &fa.wk, &fa.wv, hq, hdq)?
6737 } else {
6738 static F2B_QKV: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6741 if *F2B_QKV.get_or_init(|| std::env::var("MEMRA_F2B").as_deref() != Ok("0")) {
6742 e.matmul_q4_fused3_batched(&fa.wq, &fa.wk, &fa.wv, hq, hdq, t)?
6743 } else { None }
6744 } {
6745 Some(triple) => triple,
6746 None => (e.matmul_pre(&fa.wq, hq, hdq, h, t)?,
6747 e.matmul_pre(&fa.wk, hq, hdq, h, t)?,
6748 e.matmul_pre(&fa.wv, hq, hdq, h, t)?), };
6750 e.rms_norm_qkv_rope(&q0, &k0, &v0, fa.q_norm.float_data(),
6753 fa.k_norm.float_data(), &aux.ones, &mut q, &mut k, &mut v,
6754 hd, nh * t, nkv * t, pos_d, nh, nkv, base, 1.0, ff, eps)?;
6755 }
6756 let kvl = cache.kv[il].as_mut().unwrap();
6757 let cls = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
6761 if dc_bucket.is_some() {
6762 debug_assert!(t == 1);
6767 e.append_kv_quantized_row_dc_inc(&k, &v, &mut kvl.k, &mut kvl.v,
6769 &mut kvl.len_d, kvl.kv_dim_k, kvl.kv_dim_v,
6770 kvl.k_tok_bytes, kvl.v_tok_bytes, cls)?;
6771 } else {
6772 e.append_kv_quantized_rows(&k, &v, &mut kvl.k, &mut kvl.v, kvl.len, t,
6773 kvl.kv_dim_k, kvl.kv_dim_v, kvl.k_tok_bytes,
6774 kvl.v_tok_bytes, cls)?;
6775 kvl.len += t;
6776 }
6777 kv_f32 = Some((k, v));
6778 }
6779 let kvl_idx = share.unwrap_or(il);
6782 let kvl = cache.kv[kvl_idx].as_ref().unwrap();
6783 let base_len = kvl.len - t; let win = self.cfg.gemma4.as_ref().unwrap().sliding_window as usize;
6785 let mut attn = e.uninit(t * nh * hd)?;
6786 if t > 1 && base_len == 0 && std::env::var("MEMRA_NOFA").is_err() {
6798 if let Some((kf, vf)) = &kv_f32 {
6799 if hd == 256 && t <= win {
6800 e.fa_prefill(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale, true)?;
6801 return Ok(e.matmul(&fa.wo, &attn, t)?);
6802 }
6803 if hd == 256 && swa && t > win {
6804 e.fa_prefill_w(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale, true,
6805 win)?;
6806 return Ok(e.matmul(&fa.wo, &attn, t)?);
6807 }
6808 if hd == 512 && !swa {
6809 e.fa_prefill_hd512(&q, kf, vf, &mut attn, hd, nh, nkv, t, t, scale,
6810 true)?;
6811 return Ok(e.matmul(&fa.wo, &attn, t)?);
6812 }
6813 } else if share.is_some() {
6814 let g = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
6815 let k_view = e.view_u8(&kvl.k, kvl.k.len());
6816 let v_view = e.view_u8(&kvl.v, kvl.v.len());
6817 if hd == 256 && (!swa || t <= win) {
6818 e.fa_prefill_view(&q, &k_view, &v_view, &mut attn, hd, nh, nkv, t, t,
6820 scale, true, kvl.k_tok_bytes, kvl.v_tok_bytes, g)?;
6821 return Ok(e.matmul(&fa.wo, &attn, t)?);
6822 }
6823 let kv_dim = nkv * hd;
6826 let mut kf = e.uninit(t * kv_dim)?;
6827 let mut vf = e.uninit(t * kv_dim)?;
6828 e.fa_dequant_kv_view_f32(&k_view, &v_view, &mut kf, &mut vf, kv_dim, kv_dim,
6829 t, kvl.k_tok_bytes, kvl.v_tok_bytes, g)?;
6830 if hd == 512 {
6831 e.fa_prefill_hd512(&q, &kf, &vf, &mut attn, hd, nh, nkv, t, t, scale,
6832 true)?;
6833 } else {
6834 e.fa_prefill_w(&q, &kf, &vf, &mut attn, hd, nh, nkv, t, t, scale, true,
6835 win)?;
6836 }
6837 return Ok(e.matmul(&fa.wo, &attn, t)?);
6838 }
6839 }
6840 if let Some(bucket) = dc_bucket {
6841 assert!(t == 1);
6846 let bucket = if hd == 512 && win <= crate::fa512_min_tkv() {
6852 bucket.min(crate::fa512_min_tkv().saturating_sub(1))
6853 } else { bucket };
6854 let k_view = e.view_u8(&kvl.k, kvl.k.len());
6855 let v_view = e.view_u8(&kvl.v, kvl.v.len());
6856 let g = (!swa && crate::Engine::gkv_on()) || (swa && crate::Engine::wkv_on());
6857 if crate::Engine::wpf_level() >= 1 {
6865 e.prefetch_weight_l2(&fa.wo)?;
6866 }
6867 if e.uses_q8_1_fast(&fa.wo) {
6870 let mut oq = e.alloc_i8_uninit(nh * hd)?;
6871 let mut od = e.zeros(nh * hd / 32)?;
6872 e.fa_decode_dc_q8(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
6873 &kvl.len_d, bucket, scale,
6874 kvl.k_tok_bytes, kvl.v_tok_bytes, g,
6875 Some((&mut oq, &mut od)))?;
6876 return Ok(e.matmul_pre(&fa.wo, &oq, &od, &attn, t)?);
6877 }
6878 e.fa_decode_dc(&q, &k_view, &v_view, &mut attn, hd, nh, nkv,
6879 &kvl.len_d, bucket, scale,
6880 kvl.k_tok_bytes, kvl.v_tok_bytes, g)?;
6881 return Ok(e.matmul(&fa.wo, &attn, t)?);
6882 }
6883 for i in 0..t {
6884 let avail = base_len + i + 1;
6885 let (off_tok, t_kv) = if swa && avail > win { (avail - win, win) } else { (0, avail) };
6886 let k_view = e.view_u8_range(&kvl.k, off_tok * kvl.k_tok_bytes,
6887 (off_tok + t_kv) * kvl.k_tok_bytes);
6888 let v_view = e.view_u8_range(&kvl.v, off_tok * kvl.v_tok_bytes,
6889 (off_tok + t_kv) * kvl.v_tok_bytes);
6890 let qv = e.view(&q, t * nh * hd);
6891 let q_row = qv.slice(i * nh * hd..(i + 1) * nh * hd);
6892 let mut q_one = e.uninit(nh * hd)?;
6893 e.copy_view_into(&mut q_one, 0, &q_row, nh * hd)?;
6894 let mut a_one = e.uninit(nh * hd)?;
6895 e.fa_decode_kvmod(&q_one, &k_view, &v_view, &mut a_one, hd, nh, nkv, t_kv, scale,
6899 kvl.k_tok_bytes, kvl.v_tok_bytes,
6900 (!swa && crate::Engine::gkv_on())
6901 || (swa && crate::Engine::wkv_on()))?;
6902 e.copy_into(&mut attn, i * nh * hd, &a_one, nh * hd)?;
6903 }
6904 Ok(e.matmul(&fa.wo, &attn, t)?)
6905 }
6906
6907 fn gemma4_e4b_trunk(&self, e: &Engine, tokens: &[u32], pos0: usize, cache: &mut Cache,
6912 head_last: bool)
6913 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6914 let n_embd = self.cfg.n_embd as usize;
6915 let t = tokens.len();
6916 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
6917 let pos_d = e.htod_i32(&pos)?;
6918 let mut x = e.htod(&self.embd.gather(n_embd, tokens))?;
6919 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
6920 let inp_pl = self.gemma4_e4b_inp_pl(e, tokens, &x, t)?;
6921 self.gemma4_e4b_trunk_core(e, x, inp_pl, &pos_d, t, cache, None, true, head_last)
6922 }
6923
6924 fn gemma4_e4b_trunk_core(&self, e: &Engine, x_in: CudaSlice<f32>, inp_pl: CudaSlice<f32>,
6928 pos_d: &CudaSlice<i32>, t: usize, cache: &mut Cache,
6929 dc_bucket: Option<usize>, cap_logits: bool, head_last: bool)
6930 -> Result<(CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
6931 let n_embd = self.cfg.n_embd as usize;
6932 let eps = self.cfg.rms_eps;
6933 let n_layer = self.layers.len();
6934 let mut x = x_in;
6935 let aux_e4b = self.gemma4_aux.as_ref().unwrap().e4b.as_ref().unwrap();
6936 let n_epl = aux_e4b.n_epl;
6937
6938 let mut h_carry: Option<(CudaSlice<i8>, CudaSlice<f32>)> = None;
6944 for il in 0..n_layer {
6945 let layer = &self.layers[il];
6946 let (hq, hdq) = match h_carry.take() {
6947 Some(p) => p,
6948 None => e.rms_norm_q8_1(&x, layer.attn_norm.float_data(), n_embd, t, eps)?,
6949 };
6950 let o = self.gemma4_e4b_attn(e, il, &hq, &hdq, pos_d, t, cache, dc_bucket)?;
6951 let bits = layer.gemma4.as_ref().unwrap();
6954 let e4b = bits.e4b.as_ref().expect("e4b layer bits");
6955 let fuse_exit = e.uses_q8_1_fast(&e4b.inp_gate);
6966 let (sn, attn_out) = self.gemma4_layer_tail_core_pn(
6967 e, layer, &o, &x, t, Some(layer.post_attn_norm.float_data()), fuse_exit)?;
6968 let mut resid = e.uninit(t * n_embd)?;
6969 let g = if fuse_exit {
6975 let (rq, rd) = e.rms_pre_add_q8_1(&sn, bits.post_ffw_norm.float_data(),
6977 &attn_out, &mut resid, n_embd, t,
6978 self.cfg.rms_eps)?;
6979 e.matmul_pre(&e4b.inp_gate, &rq, &rd, &resid, t)?
6980 } else {
6981 e.add(&sn, &attn_out, &mut resid, t * n_embd)?;
6982 e.matmul(&e4b.inp_gate, &resid, t)?
6983 };
6984 let mut act = e.uninit(t * n_epl)?;
6985 let y = if t == 1 && e.uses_q8_1_fast(&e4b.proj) {
6986 let ipv = e.view(&inp_pl, n_epl * n_layer);
6987 let row = ipv.slice(il * n_epl..(il + 1) * n_epl);
6988 let (aq, ad) = e.gelu_tanh_mul_q8_1(&g, &row, &mut act, n_epl, 1)?;
6989 e.matmul_pre(&e4b.proj, &aq, &ad, &act, t)?
6990 } else {
6991 let mut inp_this = e.uninit(t * n_epl)?;
6992 e.copy_rows_strided(&inp_pl, &mut inp_this, n_epl, t, n_epl * n_layer,
6993 il * n_epl)?;
6994 e.gelu_tanh_mul(&g, &inp_this, &mut act, t * n_epl)?;
6995 e.matmul(&e4b.proj, &act, t)?
6996 };
6997 let next_norm = if il + 1 < n_layer {
7000 self.layers[il + 1].attn_norm.float_data()
7001 } else {
7002 self.output_norm.float_data()
7003 };
7004 let mut xn = e.uninit(t * n_embd)?;
7005 let pair = e.rms_pre_add_scale_rms_norm_q8_1(&y, e4b.post_norm.float_data(),
7006 &resid, bits.layer_scale, next_norm,
7007 &mut xn, n_embd, t, eps)?;
7008 h_carry = Some(pair);
7009 x = xn;
7010 }
7011 let (oq, odq) = h_carry.take().unwrap();
7015 let h0 = e.zeros(0)?;
7016 let hm = if head_last { 1 } else { t };
7017 let (hq, hd) = if head_last && t > 1 {
7018 let mut q1 = e.uninit_i8(n_embd)?;
7019 e.dtod_copy_view_i8(&oq.slice((t - 1) * n_embd..t * n_embd), &mut q1)?;
7020 let nb = n_embd / 32;
7021 let mut d1 = e.uninit(nb)?;
7022 e.dtod_copy_view(&odq.slice((t - 1) * nb..t * nb), &mut d1)?;
7023 (q1, d1)
7024 } else {
7025 (oq, odq)
7026 };
7027 let mut ld = e.matmul_pre(&self.output, &hq, &hd, &h0, hm)?;
7028 if cap_logits {
7032 let cap = self.cfg.gemma4.as_ref().unwrap().final_logit_softcapping;
7033 e.softcap(&mut ld, cap, hm * self.output.out_features())?;
7034 }
7035 self.gemma4_suppress(e, &mut ld, hm)?; Ok((ld, x))
7037 }
7038
7039 pub fn gemma4_e4b_decode_step_t_am_dev(&self, e: &Engine, tok_d: &CudaSlice<u32>,
7046 t: usize, pos0: usize, cache: &mut Cache)
7047 -> Result<(CudaSlice<u32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7048 let n_embd = self.cfg.n_embd as usize;
7049 let eps = self.cfg.rms_eps;
7050 let pos: Vec<i32> = (0..t).map(|i| (pos0 + i) as i32).collect();
7051 let pos_d = e.htod_i32(&pos)?;
7052 let embd_gpu = self.embd_gpu.get_or_init(|| {
7053 e.upload_u8(&self.embd.raw).expect("embed table upload")
7054 });
7055 let (qt, rb) = self.embd.qt_and_row_bytes(n_embd);
7056 let mut x = e.embed_gather_device_td(embd_gpu, tok_d, t, n_embd, qt, rb)?;
7057 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), t * n_embd)?;
7058 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, tok_d, &x, t)?;
7059 let (ld, xp) = self.gemma4_e4b_trunk_core(e, x, inp_pl, &pos_d, t, cache, None, true,
7060 false)?;
7061 let n_vocab = self.output.out_features();
7064 let mut vam = e.stream().alloc_zeros::<u32>(t)?;
7065 for i in 0..t {
7066 e.argmax_token_device_col(&ld, i, n_vocab, &mut vam, i)?;
7067 }
7068 let mut hn = e.uninit(t * n_embd)?;
7069 e.rms_norm(&xp, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
7070 cache.pos += t;
7071 Ok((vam, hn))
7072 }
7073
7074 pub(crate) fn gemma4_e4b_decode_step_t_h(&self, e: &Engine, tokens: &[u32], pos0: usize,
7077 cache: &mut Cache)
7078 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7079 let n_embd = self.cfg.n_embd as usize;
7080 let eps = self.cfg.rms_eps;
7081 let t = tokens.len();
7082 let (ld, xp) = self.gemma4_e4b_trunk(e, tokens, pos0, cache, false)?;
7083 let mut hn = e.uninit(t * n_embd)?;
7084 e.rms_norm(&xp, self.output_norm.float_data(), &mut hn, n_embd, t, eps)?;
7085 cache.pos += t;
7086 Ok((e.dtoh(&ld)?, hn))
7087 }
7088
7089 pub fn gemma4_e4b_decode_step_dcg(&self, e: &Engine, token_d: &mut CudaSlice<u32>,
7095 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
7096 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
7097 n_vocab: usize, bucket: usize)
7098 -> Result<(), Box<dyn std::error::Error>> {
7099 let n_embd = self.cfg.n_embd as usize;
7100 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
7101 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
7102 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, token_d, &x, 1)?;
7103 let (ld, _x) = self.gemma4_e4b_trunk_core(e, x, inp_pl, pos_d, 1, cache, Some(bucket),
7104 false, false)?;
7105 e.argmax_token_device_into(&ld, token_d, n_vocab)?;
7106 e.inc_seqlen(pos_d)?;
7107 Ok(())
7108 }
7109
7110 #[allow(clippy::too_many_arguments)]
7118 pub fn gemma4_e4b_decode_step_dc(&self, e: &Engine, token_d: &CudaSlice<u32>,
7119 pos_d: &mut CudaSlice<i32>, embd_gpu: &CudaSlice<u8>,
7120 embd_qt: i32, embd_rb: usize, cache: &mut Cache,
7121 n_vocab: usize)
7122 -> Result<CudaSlice<u32>, Box<dyn std::error::Error>> {
7123 let n_embd = self.cfg.n_embd as usize;
7124 let eps = self.cfg.rms_eps;
7125 let mut x = e.embed_gather_device(embd_gpu, token_d, n_embd, embd_qt, embd_rb)?;
7126 e.scale_inplace(&mut x, (n_embd as f32).sqrt(), n_embd)?;
7127 let inp_pl = self.gemma4_e4b_inp_pl_dev(e, token_d, &x, 1)?;
7128 let (ld, _x) = self.gemma4_e4b_trunk_core(e, x, inp_pl, pos_d, 1, cache, None, false,
7129 false)?;
7130 let mut tok_out = e.stream().alloc_zeros::<u32>(1)?;
7131 e.argmax_token_device_into(&ld, &mut tok_out, n_vocab)?;
7132 e.inc_seqlen(pos_d)?;
7133 cache.pos += 1;
7134 let _ = eps;
7135 Ok(tok_out)
7136 }
7137
7138 pub(crate) fn gemma4_e4b_decode_step_h(&self, e: &Engine, token: u32, cache: &mut Cache)
7141 -> Result<(Vec<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7142 let (ld, x) = self.gemma4_e4b_trunk(e, &[token], cache.pos, cache, false)?;
7143 let logits = e.dtoh(&ld)?;
7144 cache.pos += 1;
7145 Ok((logits, x))
7146 }
7147
7148 pub(crate) fn gemma4_e4b_prime(&self, e: &Engine, tokens: &[u32], cache: &mut Cache)
7152 -> Result<(Vec<f32>, CudaSlice<f32>, CudaSlice<f32>), Box<dyn std::error::Error>> {
7153 assert_eq!(cache.pos, 0, "e4b prime is fresh-prompt only (v0)");
7154 let n_embd = self.cfg.n_embd as usize;
7155 let t = tokens.len();
7156 let (ld, x) = self.gemma4_e4b_trunk(e, tokens, 0, cache, true)?;
7157 cache.pos += t;
7158 let last = e.dtoh(&ld)?; let xv = e.view(&x, t * n_embd);
7160 let row = xv.slice((t - 1) * n_embd..t * n_embd);
7161 let mut h_seed = e.uninit(n_embd)?;
7162 e.copy_view_into(&mut h_seed, 0, &row, n_embd)?;
7163 Ok((last, h_seed, x))
7164 }
7165
7166 pub(crate) fn gemma4_e4b_forward(&self, e: &Engine, tokens: &[u32], last_only: bool)
7168 -> Result<Vec<f32>, Box<dyn std::error::Error>> {
7169 let mut cache = Cache::new(e, &self.cfg, tokens.len() + 8)?;
7170 let (ld, _x) = self.gemma4_e4b_trunk(e, tokens, 0, &mut cache, last_only)?;
7171 Ok(e.dtoh(&ld)?) }
7173}
7174
7175#[cfg(test)]
7176mod page_prefetch_tests {
7177 use super::{
7178 grouped_worker_prefetch_position, page_prefetch_positions,
7179 page_prefetch_window_from_values, worker_prefetch_positions,
7180 };
7181
7182 #[test]
7183 fn page_prefetch_window_keeps_existing_opt_in_default() {
7184 assert_eq!(page_prefetch_window_from_values(false, None), 0);
7185 assert_eq!(page_prefetch_window_from_values(false, Some("8")), 0);
7186 assert_eq!(page_prefetch_window_from_values(true, None), 1);
7187 assert_eq!(page_prefetch_window_from_values(true, Some("bad")), 1);
7188 assert_eq!(page_prefetch_window_from_values(true, Some("0")), 0);
7189 assert_eq!(page_prefetch_window_from_values(true, Some("8")), 8);
7190 }
7191
7192 #[test]
7193 fn rolling_page_prefetch_advises_each_future_expert_once() {
7194 let advised: Vec<_> = (0..7)
7195 .flat_map(|position| page_prefetch_positions(position, 7, 3))
7196 .collect();
7197 assert_eq!(advised, vec![1, 2, 3, 4, 5, 6]);
7198
7199 let one_ahead: Vec<_> = (0..4)
7200 .flat_map(|position| page_prefetch_positions(position, 4, 1))
7201 .collect();
7202 assert_eq!(one_ahead, vec![1, 2, 3]);
7203 assert!(page_prefetch_positions(0, 4, 0).is_empty());
7204 }
7205
7206 #[test]
7207 fn grouped_worker_prefetch_primes_first_then_each_known_next_once() {
7208 assert_eq!(grouped_worker_prefetch_position(0, None), None);
7209 let positions: Vec<_> = std::iter::once(grouped_worker_prefetch_position(4, None).unwrap())
7210 .chain((0..4).filter_map(|position| {
7211 grouped_worker_prefetch_position(4, Some(position))
7212 }))
7213 .collect();
7214 assert_eq!(positions, vec![0, 1, 2, 3]);
7215 assert_eq!(grouped_worker_prefetch_position(1, Some(0)), None);
7216 }
7217
7218 #[test]
7219 fn rolling_worker_prefetch_primes_current_and_each_future_expert_once() {
7220 let queued: Vec<_> = (0..8)
7221 .flat_map(|position| worker_prefetch_positions(position, 8, 5))
7222 .collect();
7223 assert_eq!(queued, (0..8).collect::<Vec<_>>());
7224
7225 let one_at_a_time: Vec<_> = (0..4)
7226 .flat_map(|position| worker_prefetch_positions(position, 4, 1))
7227 .collect();
7228 assert_eq!(one_at_a_time, vec![0, 1, 2, 3]);
7229 assert!(worker_prefetch_positions(0, 4, 0).is_empty());
7230 }
7231}
7232
7233pub struct G4DcSlots {
7234 x: CudaSlice<f32>, xn: CudaSlice<f32>, cur: CudaSlice<f32>,
7235 hq: CudaSlice<i8>, hd_: CudaSlice<f32>,
7236 q0: CudaSlice<f32>, k0: CudaSlice<f32>, v0: CudaSlice<f32>,
7237 q: CudaSlice<f32>, k: CudaSlice<f32>, v: CudaSlice<f32>,
7238 attn: CudaSlice<f32>, o: CudaSlice<f32>,
7239 attn_out: CudaSlice<f32>, zsh: CudaSlice<f32>,
7240 zq: CudaSlice<i8>, zd: CudaSlice<f32>,
7241 gate: CudaSlice<f32>, up: CudaSlice<f32>,
7242 act: CudaSlice<f32>, actq: CudaSlice<i8>, actd: CudaSlice<f32>,
7243 f0: CudaSlice<f32>, sn: CudaSlice<f32>,
7244 hn: CudaSlice<f32>, logits: CudaSlice<f32>,
7245}
7246